Merge branch 'main' into perf/async-url-upload-7099409897484162483
This commit is contained in:
@@ -7,6 +7,24 @@ GOTENBERG_URL=http://gotenberg:3000
|
||||
ALLOW_FILE_DELETE=true # Allow deletion of file records
|
||||
COMPLIANCE_ENABLED=true # Enable compliance templates dashboard (GDPR, HIPAA, SOC 2)
|
||||
|
||||
# **Logging**
|
||||
# LOG_LEVEL controls the Python root-logger level.
|
||||
# Accepted values: DEBUG, INFO, WARNING, ERROR, CRITICAL (default: INFO).
|
||||
# When DEBUG=true and LOG_LEVEL is not set, the level is automatically lowered to DEBUG.
|
||||
# LOG_LEVEL=INFO
|
||||
# DEBUG=false
|
||||
|
||||
# Log output format: "text" (human-readable, default) or "json" (structured JSON lines).
|
||||
# Use "json" when shipping logs to Grafana Loki, Splunk, ELK, Datadog, or any SIEM.
|
||||
# LOG_FORMAT=text
|
||||
|
||||
# Forward application logs to a syslog receiver (in addition to stdout).
|
||||
# Useful for traditional (non-container) deployments and centralised SIEM ingestion.
|
||||
# LOG_SYSLOG_ENABLED=false
|
||||
# LOG_SYSLOG_HOST=localhost
|
||||
# LOG_SYSLOG_PORT=514
|
||||
# LOG_SYSLOG_PROTOCOL=udp # udp | tcp
|
||||
|
||||
# **UI / Appearance**
|
||||
# Default colour scheme: system (follow OS), light, or dark
|
||||
# Individual users can always override with the navbar dark-mode toggle.
|
||||
|
||||
+1
-1
@@ -1 +1 @@
|
||||
2026-03-15T21:39:26Z
|
||||
2026-03-16T10:45:13Z
|
||||
|
||||
+223
@@ -10,6 +10,229 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
||||
|
||||
<!-- version list -->
|
||||
|
||||
## Unreleased
|
||||
|
||||
### Code Style
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`46c4031`](https://github.com/christianlouis/DocuElevate/commit/46c403127641c1b5c729ff69ab5505efa5c9b54d))
|
||||
|
||||
### Documentation
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`0f31216`](https://github.com/christianlouis/DocuElevate/commit/0f312160bcb2e46e29c90dd055c4ebc9aaf01ad4))
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`5734df2`](https://github.com/christianlouis/DocuElevate/commit/5734df2d5046158a23b987fdc4235d3f3f6b4042))
|
||||
|
||||
|
||||
## Unreleased
|
||||
|
||||
### Code Style
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`46c4031`](https://github.com/christianlouis/DocuElevate/commit/46c403127641c1b5c729ff69ab5505efa5c9b54d))
|
||||
|
||||
### Documentation
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`5734df2`](https://github.com/christianlouis/DocuElevate/commit/5734df2d5046158a23b987fdc4235d3f3f6b4042))
|
||||
|
||||
|
||||
## Unreleased
|
||||
|
||||
|
||||
## v0.147.1 (2026-03-16)
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- Resolve test failures and mypy errors on main
|
||||
([`c7ff177`](https://github.com/christianlouis/DocuElevate/commit/c7ff177e179b60a4b6adba7f5296522e6da4391f))
|
||||
|
||||
### Documentation
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`dd7c8f0`](https://github.com/christianlouis/DocuElevate/commit/dd7c8f0342f00b80e83394d146e6bc51a3688687))
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`4d4706e`](https://github.com/christianlouis/DocuElevate/commit/4d4706e078d661682326262000e70d7e56fb6d23))
|
||||
|
||||
|
||||
## Unreleased
|
||||
|
||||
### Documentation
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`4d4706e`](https://github.com/christianlouis/DocuElevate/commit/4d4706e078d661682326262000e70d7e56fb6d23))
|
||||
|
||||
|
||||
## Unreleased
|
||||
|
||||
|
||||
## v0.147.0 (2026-03-16)
|
||||
|
||||
### Code Style
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`80a0ddc`](https://github.com/christianlouis/DocuElevate/commit/80a0ddcbfcd89c156d6cce65bb6275bdf35c061e))
|
||||
|
||||
### Documentation
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`5e986cb`](https://github.com/christianlouis/DocuElevate/commit/5e986cb6d66c313753c7ce9cf69137c36b354ae8))
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`8f8e11c`](https://github.com/christianlouis/DocuElevate/commit/8f8e11cffe5a6aaae47c9111bc6a95cb88ef6cd8))
|
||||
|
||||
### Features
|
||||
|
||||
- **config**: Add JSON structured logging and syslog forwarding for application logs
|
||||
([`6da1e6f`](https://github.com/christianlouis/DocuElevate/commit/6da1e6fd816fea1dc0fa37c9ad42ce8cafed8c73))
|
||||
|
||||
- **config**: Add LOG_LEVEL setting and configure root logging at startup
|
||||
([`18c49c6`](https://github.com/christianlouis/DocuElevate/commit/18c49c6b2d214fddefa2f954d57b0358fb4b5f72))
|
||||
|
||||
### Refactoring
|
||||
|
||||
- **main**: Move JSON formatter imports to module level per code review
|
||||
([`df1fa51`](https://github.com/christianlouis/DocuElevate/commit/df1fa51800fb62862271fce4e0fff691f62c3333))
|
||||
|
||||
### Testing
|
||||
|
||||
- Add 500 error test for saved search deletion
|
||||
([`ae9ed6e`](https://github.com/christianlouis/DocuElevate/commit/ae9ed6e9a703f8a80880d0b1cb421de3cf024cd3))
|
||||
|
||||
|
||||
## Unreleased
|
||||
|
||||
### Code Style
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`80a0ddc`](https://github.com/christianlouis/DocuElevate/commit/80a0ddcbfcd89c156d6cce65bb6275bdf35c061e))
|
||||
|
||||
### Documentation
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`8f8e11c`](https://github.com/christianlouis/DocuElevate/commit/8f8e11cffe5a6aaae47c9111bc6a95cb88ef6cd8))
|
||||
|
||||
### Testing
|
||||
|
||||
- Add 500 error test for saved search deletion
|
||||
([`ae9ed6e`](https://github.com/christianlouis/DocuElevate/commit/ae9ed6e9a703f8a80880d0b1cb421de3cf024cd3))
|
||||
|
||||
|
||||
## Unreleased
|
||||
|
||||
### Code Style
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`80a0ddc`](https://github.com/christianlouis/DocuElevate/commit/80a0ddcbfcd89c156d6cce65bb6275bdf35c061e))
|
||||
|
||||
### Testing
|
||||
|
||||
- Add 500 error test for saved search deletion
|
||||
([`ae9ed6e`](https://github.com/christianlouis/DocuElevate/commit/ae9ed6e9a703f8a80880d0b1cb421de3cf024cd3))
|
||||
|
||||
|
||||
## v0.146.0 (2026-03-16)
|
||||
|
||||
### Code Style
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`522cefa`](https://github.com/christianlouis/DocuElevate/commit/522cefad935508fed32e194984726b0a10c99654))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`c76e513`](https://github.com/christianlouis/DocuElevate/commit/c76e51391b14eccbbe74f34ffadc89a477b58a64))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`cc2a07b`](https://github.com/christianlouis/DocuElevate/commit/cc2a07b090883f3cf2affea8a5e254a4163ce966))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`8bb6457`](https://github.com/christianlouis/DocuElevate/commit/8bb6457c65c177a13fd15de7a7b55e24eb39477e))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`275a5ad`](https://github.com/christianlouis/DocuElevate/commit/275a5ad6fa887ce3aa05d50787bd5da435735f26))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`ac35c5e`](https://github.com/christianlouis/DocuElevate/commit/ac35c5e6fa84f2b38cd333de75f202f13bf4c461))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`705c801`](https://github.com/christianlouis/DocuElevate/commit/705c801158394e7d6466f82485f8d872860653bb))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`0425d46`](https://github.com/christianlouis/DocuElevate/commit/0425d46c4440191cd410434ba64c1bc9cdb53cbb))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`b4118f6`](https://github.com/christianlouis/DocuElevate/commit/b4118f61622a35461af2c220d8a948ecb795e65b))
|
||||
|
||||
- Format app/database.py to fix CI failure
|
||||
([`fbd4f83`](https://github.com/christianlouis/DocuElevate/commit/fbd4f837301deb208c7303553bca8facbae85ba6))
|
||||
|
||||
### Features
|
||||
|
||||
- Extract embedded PDF metadata using pypdf in upload_to_email
|
||||
([`726e4df`](https://github.com/christianlouis/DocuElevate/commit/726e4dfdc47507dd08047a27f02c137ed3e7ecaf))
|
||||
|
||||
- Extract embedded PDF metadata using pypdf in upload_to_email
|
||||
([`df64aec`](https://github.com/christianlouis/DocuElevate/commit/df64aece2c19215cf762b2ac62d476888cd21527))
|
||||
|
||||
### Performance Improvements
|
||||
|
||||
- Optimize dropbox token refresh by replacing blocking requests with httpx
|
||||
([`84c6e1c`](https://github.com/christianlouis/DocuElevate/commit/84c6e1c5dd1a427c7f015e7461df59b889e66154))
|
||||
|
||||
- **api**: Optimize reorder_plans to prevent N+1 queries
|
||||
([`d8372c6`](https://github.com/christianlouis/DocuElevate/commit/d8372c6fb83b09ce8d61bc8a870c921372a9d27a))
|
||||
|
||||
- **duplicates**: Fix N+1 query in group listing
|
||||
([`e4e3ac4`](https://github.com/christianlouis/DocuElevate/commit/e4e3ac40771bdd8459cc0876d9d49be441df51a4))
|
||||
|
||||
### Testing
|
||||
|
||||
- Add missing error tests for updating saved searches
|
||||
([`d18c05c`](https://github.com/christianlouis/DocuElevate/commit/d18c05c36dc4324dd24e20e5344bca9c038bbf73))
|
||||
|
||||
- Improve coverage for notify_settings_updated error handling
|
||||
([`d8906ae`](https://github.com/christianlouis/DocuElevate/commit/d8906aece0045aed9dcd0f0cf1891970c7cb00e8))
|
||||
|
||||
|
||||
## v0.145.3 (2026-03-16)
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
- **tests**: Resolve ruff import sorting issue in benchmark test
|
||||
([`fa9b037`](https://github.com/christianlouis/DocuElevate/commit/fa9b037d5a2a67f9115a1bddf0f98ba9020cefd6))
|
||||
|
||||
### Code Style
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`c1657a0`](https://github.com/christianlouis/DocuElevate/commit/c1657a01a77ca6b914bc21a20a42aeb924de8254))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`6f5f4d9`](https://github.com/christianlouis/DocuElevate/commit/6f5f4d9d4948208f65fc5572c86a08642efe54b6))
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`2cfbea2`](https://github.com/christianlouis/DocuElevate/commit/2cfbea29a9211abfc13b74dea0a841fe5a2546b9))
|
||||
|
||||
### Documentation
|
||||
|
||||
- **changelog**: Update changelog [skip ci]
|
||||
([`fa36ec6`](https://github.com/christianlouis/DocuElevate/commit/fa36ec69876b919a33a7471bb837c2a6b69c2a50))
|
||||
|
||||
### Performance Improvements
|
||||
|
||||
- **api**: Fix n+1 query issue in user notification preferences update
|
||||
([`fe20e02`](https://github.com/christianlouis/DocuElevate/commit/fe20e02f78c3c6236b07f50edb66599b8a9c7ed2))
|
||||
|
||||
|
||||
## Unreleased
|
||||
|
||||
### Code Style
|
||||
|
||||
- Apply ruff auto-fix
|
||||
([`2cfbea2`](https://github.com/christianlouis/DocuElevate/commit/2cfbea29a9211abfc13b74dea0a841fe5a2546b9))
|
||||
|
||||
|
||||
## v0.145.2 (2026-03-15)
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
+6
-6
@@ -1,10 +1,10 @@
|
||||
DocuElevate Build Information
|
||||
==============================
|
||||
Version: 0.145.2
|
||||
Build Date: 2026-03-15T21:39:26Z
|
||||
Git Commit: 237af31f5fe598cfc3c08f2bbba79b3d0925787e
|
||||
Git Short SHA: 237af31
|
||||
Version: 0.147.1
|
||||
Build Date: 2026-03-16T10:45:13Z
|
||||
Git Commit: fd15c3666547405bb0a3af37e98be4727ff635bb
|
||||
Git Short SHA: fd15c36
|
||||
Git Branch: main
|
||||
Commit Date: 2026-03-15T22:39:04+01:00
|
||||
Build Timestamp: 2026-03-15T21:39:26Z
|
||||
Commit Date: 2026-03-16T11:44:51+01:00
|
||||
Build Timestamp: 2026-03-16T10:45:13Z
|
||||
==============================
|
||||
|
||||
@@ -20,14 +20,16 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
# Module-level dependency singleton to satisfy Ruff B008 while maintaining default values for manual calls (e.g. in decorators).
|
||||
_db_dep = Depends(get_db)
|
||||
DbSession = Annotated[Session, _db_dep]
|
||||
|
||||
|
||||
@router.get("/audit-logs")
|
||||
@require_login
|
||||
async def list_audit_logs(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
db: DbSession = _db_dep,
|
||||
action: Annotated[str | None, Query(description="Filter by action (exact match)")] = None,
|
||||
user: Annotated[str | None, Query(description="Filter by username")] = None,
|
||||
resource_type: Annotated[str | None, Query(description="Filter by resource type")] = None,
|
||||
@@ -73,7 +75,7 @@ async def list_audit_logs(
|
||||
@require_login
|
||||
async def list_distinct_actions(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
db: DbSession = _db_dep,
|
||||
) -> list[str]:
|
||||
"""Return the distinct action values present in the audit log."""
|
||||
from app.models import AuditLog
|
||||
@@ -86,7 +88,7 @@ async def list_distinct_actions(
|
||||
@require_login
|
||||
async def list_distinct_users(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
db: DbSession = _db_dep,
|
||||
) -> list[str]:
|
||||
"""Return the distinct user values present in the audit log."""
|
||||
from app.models import AuditLog
|
||||
|
||||
+49
-46
@@ -6,7 +6,7 @@ import logging
|
||||
import os
|
||||
from typing import Annotated, Optional
|
||||
|
||||
import requests
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, Form, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -132,57 +132,60 @@ async def test_dropbox_token(request: Request):
|
||||
"message": "Dropbox credentials are not fully configured",
|
||||
}
|
||||
|
||||
# Check token validity by getting current account info
|
||||
headers = {"Authorization": f"Bearer {settings.dropbox_refresh_token}"}
|
||||
response = requests.post(
|
||||
"https://api.dropboxapi.com/2/users/get_current_account",
|
||||
headers=headers,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
# If token is invalid, try refreshing it
|
||||
if response.status_code == 401:
|
||||
logger.info("Dropbox access token invalid or expired, trying to refresh")
|
||||
|
||||
# Get a new access token using the refresh token
|
||||
refresh_url = "https://api.dropbox.com/oauth2/token"
|
||||
refresh_data = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": settings.dropbox_refresh_token,
|
||||
"client_id": settings.dropbox_app_key,
|
||||
"client_secret": settings.dropbox_app_secret,
|
||||
}
|
||||
|
||||
refresh_response = requests.post(refresh_url, data=refresh_data, timeout=settings.http_request_timeout)
|
||||
|
||||
if refresh_response.status_code != 200:
|
||||
logger.error(f"Failed to refresh Dropbox token: {refresh_response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": "Refresh token has expired or is invalid",
|
||||
"needs_reauth": True,
|
||||
}
|
||||
|
||||
token_info = refresh_response.json()
|
||||
access_token = token_info.get("access_token")
|
||||
|
||||
# Try again with the new access token
|
||||
headers = {"Authorization": f"Bearer {access_token}"}
|
||||
response = requests.post(
|
||||
async with httpx.AsyncClient() as client:
|
||||
# Check token validity by getting current account info
|
||||
headers = {"Authorization": f"Bearer {settings.dropbox_refresh_token}"}
|
||||
response = await client.post(
|
||||
"https://api.dropboxapi.com/2/users/get_current_account",
|
||||
headers=headers,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Dropbox token test failed: {response.status_code} {response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": f"Token validation failed with status {response.status_code}: {response.text}",
|
||||
}
|
||||
# If token is invalid, try refreshing it
|
||||
if response.status_code == 401:
|
||||
logger.info("Dropbox access token invalid or expired, trying to refresh")
|
||||
|
||||
# Get account info
|
||||
account_info = response.json()
|
||||
# Get a new access token using the refresh token
|
||||
refresh_url = "https://api.dropbox.com/oauth2/token"
|
||||
refresh_data = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": settings.dropbox_refresh_token,
|
||||
"client_id": settings.dropbox_app_key,
|
||||
"client_secret": settings.dropbox_app_secret,
|
||||
}
|
||||
|
||||
refresh_response = await client.post(
|
||||
refresh_url, data=refresh_data, timeout=settings.http_request_timeout
|
||||
)
|
||||
|
||||
if refresh_response.status_code != 200:
|
||||
logger.error(f"Failed to refresh Dropbox token: {refresh_response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": "Refresh token has expired or is invalid",
|
||||
"needs_reauth": True,
|
||||
}
|
||||
|
||||
token_info = refresh_response.json()
|
||||
access_token = token_info.get("access_token")
|
||||
|
||||
# Try again with the new access token
|
||||
headers = {"Authorization": f"Bearer {access_token}"}
|
||||
response = await client.post(
|
||||
"https://api.dropboxapi.com/2/users/get_current_account",
|
||||
headers=headers,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Dropbox token test failed: {response.status_code} {response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": f"Token validation failed with status {response.status_code}: {response.text}",
|
||||
}
|
||||
|
||||
# Get account info
|
||||
account_info = response.json()
|
||||
account_email = account_info.get("email", "Unknown account")
|
||||
account_name = account_info.get("name", {}).get("display_name", "Unknown user")
|
||||
|
||||
|
||||
+28
-23
@@ -73,33 +73,38 @@ def list_duplicate_groups(
|
||||
groups = []
|
||||
total_duplicate_files = 0
|
||||
|
||||
for filehash in dup_hashes:
|
||||
# Find the original (non-duplicate) record with this hash
|
||||
original = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(False))
|
||||
.order_by(FileRecord.id.asc())
|
||||
.first()
|
||||
if dup_hashes:
|
||||
# Fetch all matching files (both original and duplicates) in a single batch query
|
||||
all_records = (
|
||||
db.query(FileRecord).filter(FileRecord.filehash.in_(dup_hashes)).order_by(FileRecord.id.asc()).all()
|
||||
)
|
||||
|
||||
# Find all duplicate records for this hash
|
||||
duplicates = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(True))
|
||||
.order_by(FileRecord.id.asc())
|
||||
.all()
|
||||
)
|
||||
# Group records by hash
|
||||
originals_by_hash = {}
|
||||
duplicates_by_hash = {h: [] for h in dup_hashes}
|
||||
|
||||
total_duplicate_files += len(duplicates)
|
||||
for record in all_records:
|
||||
h = record.filehash
|
||||
if not record.is_duplicate:
|
||||
# Store only the first original record per hash, matching the old .first() behaviour
|
||||
if h not in originals_by_hash:
|
||||
originals_by_hash[h] = record
|
||||
else:
|
||||
duplicates_by_hash[h].append(record)
|
||||
total_duplicate_files += 1
|
||||
|
||||
groups.append(
|
||||
{
|
||||
"filehash": filehash,
|
||||
"original": _file_record_to_dict(original) if original else None,
|
||||
"duplicates": [_file_record_to_dict(d) for d in duplicates],
|
||||
"duplicate_count": len(duplicates),
|
||||
}
|
||||
)
|
||||
for filehash in dup_hashes:
|
||||
original = originals_by_hash.get(filehash)
|
||||
duplicates = duplicates_by_hash.get(filehash, [])
|
||||
|
||||
groups.append(
|
||||
{
|
||||
"filehash": filehash,
|
||||
"original": _file_record_to_dict(original) if original else None,
|
||||
"duplicates": [_file_record_to_dict(d) for d in duplicates],
|
||||
"duplicate_count": len(duplicates),
|
||||
}
|
||||
)
|
||||
|
||||
total_pages = (total_groups + per_page - 1) // per_page if total_groups > 0 else 1
|
||||
|
||||
|
||||
+4
-3
@@ -11,6 +11,7 @@ import zipfile
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, List, Optional
|
||||
|
||||
import aiofiles
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, UploadFile, status
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import asc, desc
|
||||
@@ -1277,7 +1278,7 @@ async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...
|
||||
# enforcing the size limit during the read so memory usage stays bounded.
|
||||
try:
|
||||
written_size = 0
|
||||
with open(target_path, "wb") as f:
|
||||
async with aiofiles.open(target_path, "wb") as f:
|
||||
chunk_size = 65536 # 64 KB chunks
|
||||
while True:
|
||||
chunk = await file.read(chunk_size)
|
||||
@@ -1286,14 +1287,14 @@ async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...
|
||||
written_size += len(chunk)
|
||||
if written_size > max_size:
|
||||
# Exceeded limit mid-stream; clean up and reject
|
||||
f.close()
|
||||
await f.close()
|
||||
os.remove(target_path)
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
detail=f"File too large: exceeded {max_size} bytes during upload. "
|
||||
f"See SECURITY_AUDIT.md for configuration details.",
|
||||
)
|
||||
f.write(chunk)
|
||||
await f.write(chunk)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
||||
@@ -363,7 +363,7 @@ def format_time_remaining(time_delta):
|
||||
|
||||
@router.post("/google-drive/save-settings")
|
||||
@require_login
|
||||
async def save_dropbox_settings(
|
||||
async def save_google_drive_settings(
|
||||
request: Request,
|
||||
refresh_token: Annotated[str, Form(...)],
|
||||
client_id: Annotated[Optional[str], Form()] = None,
|
||||
|
||||
@@ -452,17 +452,16 @@ async def update_preferences(
|
||||
)
|
||||
|
||||
try:
|
||||
# Pre-fetch existing preferences for this user to avoid N+1 queries
|
||||
existing_prefs = (
|
||||
db.query(UserNotificationPreference).filter(UserNotificationPreference.owner_id == owner_id).all()
|
||||
)
|
||||
|
||||
# Build a fast lookup dictionary keyed by (event_type, channel_type, target_id)
|
||||
prefs_dict = {(pref.event_type, pref.channel_type, pref.target_id): pref for pref in existing_prefs}
|
||||
|
||||
for item in body.preferences:
|
||||
existing = (
|
||||
db.query(UserNotificationPreference)
|
||||
.filter(
|
||||
UserNotificationPreference.owner_id == owner_id,
|
||||
UserNotificationPreference.event_type == item.event_type,
|
||||
UserNotificationPreference.channel_type == item.channel_type,
|
||||
UserNotificationPreference.target_id == item.target_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
existing = prefs_dict.get((item.event_type, item.channel_type, item.target_id))
|
||||
if existing:
|
||||
existing.is_enabled = item.is_enabled
|
||||
else:
|
||||
|
||||
+19
-93
@@ -3,7 +3,6 @@ OneDrive API endpoints
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Annotated, Optional
|
||||
|
||||
@@ -14,6 +13,7 @@ from sqlalchemy.orm import Session
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.utils.env_utils import update_env_file
|
||||
from app.utils.oauth_helper import exchange_oauth_token
|
||||
from app.utils.settings_service import save_setting_to_db
|
||||
from app.utils.settings_sync import notify_settings_updated
|
||||
@@ -115,32 +115,7 @@ async def test_onedrive_token(request: Request):
|
||||
settings.onedrive_refresh_token = new_refresh_token
|
||||
|
||||
# Also try to update .env file if it exists
|
||||
try:
|
||||
env_path = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), ".env")
|
||||
if os.path.exists(env_path):
|
||||
with open(env_path, "r") as f:
|
||||
env_lines = f.readlines()
|
||||
|
||||
updated_lines = []
|
||||
updated = False
|
||||
|
||||
for line in env_lines:
|
||||
if line.startswith("ONEDRIVE_REFRESH_TOKEN="):
|
||||
updated_lines.append(f"ONEDRIVE_REFRESH_TOKEN={new_refresh_token}\n")
|
||||
updated = True
|
||||
else:
|
||||
updated_lines.append(line)
|
||||
|
||||
if not updated:
|
||||
updated_lines.append(f"ONEDRIVE_REFRESH_TOKEN={new_refresh_token}\n")
|
||||
|
||||
with open(env_path, "w") as f:
|
||||
f.writelines(updated_lines)
|
||||
|
||||
logger.info("Updated refresh token in .env file")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to update refresh token in .env file: {e}")
|
||||
update_env_file({"ONEDRIVE_REFRESH_TOKEN": new_refresh_token})
|
||||
|
||||
# Persist the rotated refresh token to the database
|
||||
try:
|
||||
@@ -246,75 +221,26 @@ async def save_onedrive_settings(
|
||||
user.get("preferred_username") or user.get("username") or user.get("email") or user.get("id") or "wizard"
|
||||
)
|
||||
|
||||
# Best-effort .env file write
|
||||
try:
|
||||
env_path = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), ".env")
|
||||
if not os.path.exists(env_path):
|
||||
logger.warning(f".env file not found at {env_path}, skipping file write")
|
||||
else:
|
||||
logger.info(f"Updating OneDrive settings in {env_path}")
|
||||
# Build settings dictionary mapped to database/memory keys
|
||||
onedrive_settings = {
|
||||
"onedrive_refresh_token": refresh_token,
|
||||
"onedrive_client_id": client_id,
|
||||
"onedrive_client_secret": client_secret,
|
||||
"onedrive_tenant_id": tenant_id,
|
||||
"onedrive_folder_path": folder_path,
|
||||
}
|
||||
|
||||
with open(env_path, "r") as f:
|
||||
env_lines = f.readlines()
|
||||
# Filter out None values
|
||||
onedrive_settings = {k: v for k, v in onedrive_settings.items() if v is not None}
|
||||
|
||||
onedrive_settings = {"ONEDRIVE_REFRESH_TOKEN": refresh_token}
|
||||
if client_id:
|
||||
onedrive_settings["ONEDRIVE_CLIENT_ID"] = client_id
|
||||
if client_secret:
|
||||
onedrive_settings["ONEDRIVE_CLIENT_SECRET"] = client_secret
|
||||
if tenant_id:
|
||||
onedrive_settings["ONEDRIVE_TENANT_ID"] = tenant_id
|
||||
if folder_path:
|
||||
onedrive_settings["ONEDRIVE_FOLDER_PATH"] = folder_path
|
||||
# Best-effort .env file write using the new utility
|
||||
env_settings = {k.upper(): v for k, v in onedrive_settings.items()}
|
||||
update_env_file(env_settings)
|
||||
|
||||
updated = set()
|
||||
new_env_lines = []
|
||||
for line in env_lines:
|
||||
stripped_line = line.rstrip()
|
||||
is_updated = False
|
||||
for key, value in onedrive_settings.items():
|
||||
if stripped_line.startswith(f"{key}=") or stripped_line.startswith(f"# {key}="):
|
||||
new_env_lines.append(f"{key}={value}")
|
||||
updated.add(key)
|
||||
is_updated = True
|
||||
break
|
||||
if not is_updated:
|
||||
new_env_lines.append(stripped_line)
|
||||
|
||||
for key, value in onedrive_settings.items():
|
||||
if key not in updated:
|
||||
new_env_lines.append(f"{key}={value}")
|
||||
|
||||
with open(env_path, "w") as f:
|
||||
f.write("\n".join(new_env_lines) + "\n")
|
||||
|
||||
logger.info("Successfully updated OneDrive settings in .env file")
|
||||
except Exception as env_err:
|
||||
logger.warning(f"Failed to write .env file (non-fatal): {env_err}")
|
||||
|
||||
# Update the settings in memory
|
||||
if refresh_token:
|
||||
settings.onedrive_refresh_token = refresh_token
|
||||
if client_id:
|
||||
settings.onedrive_client_id = client_id
|
||||
if client_secret:
|
||||
settings.onedrive_client_secret = client_secret
|
||||
if tenant_id:
|
||||
settings.onedrive_tenant_id = tenant_id
|
||||
if folder_path:
|
||||
settings.onedrive_folder_path = folder_path
|
||||
|
||||
# Persist to database (primary)
|
||||
if refresh_token:
|
||||
save_setting_to_db(db, "onedrive_refresh_token", refresh_token, changed_by=changed_by)
|
||||
if client_id:
|
||||
save_setting_to_db(db, "onedrive_client_id", client_id, changed_by=changed_by)
|
||||
if client_secret:
|
||||
save_setting_to_db(db, "onedrive_client_secret", client_secret, changed_by=changed_by)
|
||||
if tenant_id:
|
||||
save_setting_to_db(db, "onedrive_tenant_id", tenant_id, changed_by=changed_by)
|
||||
if folder_path:
|
||||
save_setting_to_db(db, "onedrive_folder_path", folder_path, changed_by=changed_by)
|
||||
# Update in-memory settings and persist to database dynamically
|
||||
for key, value in onedrive_settings.items():
|
||||
setattr(settings, key, value)
|
||||
save_setting_to_db(db, key, value, changed_by=changed_by)
|
||||
|
||||
notify_settings_updated()
|
||||
|
||||
|
||||
+11
-2
@@ -196,11 +196,20 @@ def seed_plans(db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
def reorder_plans(body: ReorderBody, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Update sort_order for each plan_id in *body.order* (position = index in list)."""
|
||||
updated = 0
|
||||
for sort_order, plan_id in enumerate(body.order):
|
||||
plan = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == plan_id).first()
|
||||
|
||||
# Fetch all requested plans in a single query to avoid N+1
|
||||
plan_ids = body.order
|
||||
plans = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id.in_(plan_ids)).all()
|
||||
|
||||
# Build a map for fast O(1) lookup
|
||||
plan_map = {p.plan_id: p for p in plans}
|
||||
|
||||
for sort_order, plan_id in enumerate(plan_ids):
|
||||
plan = plan_map.get(plan_id)
|
||||
if plan:
|
||||
plan.sort_order = sort_order
|
||||
updated += 1
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
|
||||
@@ -313,16 +313,18 @@ async def list_shared_links(
|
||||
active_only: bool = Query(False, description="When true, only return active (non-revoked) links"),
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List all shared links created by the authenticated user."""
|
||||
q = db.query(SharedLink).filter(SharedLink.owner_id == owner_id)
|
||||
q = (
|
||||
db.query(SharedLink, FileRecord.original_filename)
|
||||
.outerjoin(FileRecord, SharedLink.file_id == FileRecord.id)
|
||||
.filter(SharedLink.owner_id == owner_id)
|
||||
)
|
||||
if active_only:
|
||||
q = q.filter(SharedLink.is_active.is_(True))
|
||||
links = q.order_by(SharedLink.created_at.desc()).all()
|
||||
links_with_filenames = q.order_by(SharedLink.created_at.desc()).all()
|
||||
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
result = []
|
||||
for link in links:
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == link.file_id).first()
|
||||
filename = file_record.original_filename if file_record else None
|
||||
for link, filename in links_with_filenames:
|
||||
result.append(_link_to_dict(link, base_url, filename))
|
||||
return result
|
||||
|
||||
|
||||
+81
-2
@@ -125,8 +125,17 @@ def get_current_user(request: Request):
|
||||
# Check for Bearer token auth first (API tokens)
|
||||
api_user = getattr(request.state, "api_token_user", None)
|
||||
if isinstance(api_user, dict):
|
||||
logger.debug("[AUTH] get_current_user: resolved from API token (user_id=%s)", api_user.get("id"))
|
||||
return api_user
|
||||
return request.session.get("user")
|
||||
session_user = request.session.get("user")
|
||||
if session_user:
|
||||
logger.debug(
|
||||
"[AUTH] get_current_user: resolved from session (user=%s)",
|
||||
session_user.get("preferred_username") or session_user.get("email") or session_user.get("id"),
|
||||
)
|
||||
else:
|
||||
logger.debug("[AUTH] get_current_user: no user in session or API token")
|
||||
return session_user
|
||||
|
||||
|
||||
def _resolve_bearer_user(request: Request, db: Session) -> dict | None:
|
||||
@@ -141,10 +150,12 @@ def _resolve_bearer_user(request: Request, db: Session) -> dict | None:
|
||||
"""
|
||||
auth_header = request.headers.get("authorization", "")
|
||||
if not isinstance(auth_header, str) or not auth_header.startswith("Bearer "):
|
||||
logger.debug("[AUTH] _resolve_bearer_user: no Bearer token in Authorization header")
|
||||
return None
|
||||
|
||||
raw_token = auth_header[7:]
|
||||
if not raw_token or not isinstance(raw_token, str):
|
||||
logger.debug("[AUTH] _resolve_bearer_user: empty or invalid token after 'Bearer ' prefix")
|
||||
return None
|
||||
|
||||
from app.api.api_tokens import hash_token
|
||||
@@ -153,8 +164,15 @@ def _resolve_bearer_user(request: Request, db: Session) -> dict | None:
|
||||
token_hash = hash_token(raw_token)
|
||||
db_token = db.query(ApiToken).filter(ApiToken.token_hash == token_hash, ApiToken.is_active.is_(True)).first()
|
||||
if db_token is None:
|
||||
logger.debug("[AUTH] _resolve_bearer_user: no active API token matched the provided hash")
|
||||
return None
|
||||
|
||||
logger.debug(
|
||||
"[AUTH] _resolve_bearer_user: matched API token id=%s owner=%s",
|
||||
db_token.id,
|
||||
db_token.owner_id,
|
||||
)
|
||||
|
||||
# Update usage tracking
|
||||
try:
|
||||
db_token.last_used_at = datetime.now(timezone.utc)
|
||||
@@ -205,16 +223,18 @@ def require_login(func):
|
||||
|
||||
@wraps(func)
|
||||
async def wrapper(request: Request, *args, **kwargs):
|
||||
url_path = urlparse(str(request.url)).path
|
||||
# Check session auth first
|
||||
if request.session.get("user"):
|
||||
logger.debug("[AUTH] require_login: session auth OK for %s", url_path)
|
||||
if inspect.iscoroutinefunction(func):
|
||||
return await func(*args, request=request, **kwargs)
|
||||
else:
|
||||
return func(*args, request=request, **kwargs)
|
||||
|
||||
# Fall back to Bearer token auth for API endpoints
|
||||
url_path = urlparse(str(request.url)).path
|
||||
if url_path.startswith("/api/"):
|
||||
logger.debug("[AUTH] require_login: no session, trying Bearer token for %s", url_path)
|
||||
try:
|
||||
from app.database import SessionLocal
|
||||
|
||||
@@ -228,17 +248,22 @@ def require_login(func):
|
||||
|
||||
if api_user:
|
||||
request.state.api_token_user = api_user
|
||||
logger.debug(
|
||||
"[AUTH] require_login: Bearer token auth OK for %s (user=%s)", url_path, api_user.get("id")
|
||||
)
|
||||
if inspect.iscoroutinefunction(func):
|
||||
return await func(*args, request=request, **kwargs)
|
||||
else:
|
||||
return func(*args, request=request, **kwargs)
|
||||
|
||||
logger.debug("[AUTH] require_login: no valid auth for API endpoint %s — returning 401", url_path)
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
content={"error": "Not authenticated"},
|
||||
)
|
||||
|
||||
# Non-API endpoint with no session — redirect to login
|
||||
logger.debug("[AUTH] require_login: no session for %s — redirecting to /login", url_path)
|
||||
request.session["redirect_after_login"] = str(request.url)
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
@@ -307,9 +332,15 @@ async def login(request: Request):
|
||||
async def oauth_login(request: Request):
|
||||
"""Handle OAuth login flow"""
|
||||
if not OAUTH_CONFIGURED:
|
||||
logger.debug("[AUTH] oauth_login: OAuth not configured — redirecting to /login")
|
||||
return RedirectResponse(url="/login?error=OAuth+not+configured", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
redirect_uri = request.url_for("oauth_callback")
|
||||
logger.debug(
|
||||
"[AUTH] oauth_login: initiating Authentik OAuth redirect_uri=%s session_keys=%s",
|
||||
redirect_uri,
|
||||
list(request.session.keys()),
|
||||
)
|
||||
return await oauth.authentik.authorize_redirect(request, redirect_uri)
|
||||
|
||||
|
||||
@@ -324,13 +355,23 @@ async def social_login(request: Request, provider: str):
|
||||
A redirect to the provider's authorization page, or back to /login on error.
|
||||
"""
|
||||
if provider not in SOCIAL_PROVIDERS:
|
||||
logger.debug(
|
||||
"[AUTH] social_login: unknown provider=%r (registered=%s)", provider, list(SOCIAL_PROVIDERS.keys())
|
||||
)
|
||||
return RedirectResponse(url="/login?error=Unknown+social+provider", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
redirect_uri = request.url_for("social_callback", provider=provider)
|
||||
oauth_client = getattr(oauth, provider, None)
|
||||
if oauth_client is None:
|
||||
logger.debug("[AUTH] social_login: provider=%r registered but OAuth client not configured", provider)
|
||||
return RedirectResponse(url="/login?error=Provider+not+configured", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
logger.debug(
|
||||
"[AUTH] social_login: initiating %s OAuth, redirect_uri=%s session_keys=%s",
|
||||
provider,
|
||||
redirect_uri,
|
||||
list(request.session.keys()),
|
||||
)
|
||||
return await oauth_client.authorize_redirect(request, redirect_uri)
|
||||
|
||||
|
||||
@@ -389,27 +430,39 @@ async def social_callback(request: Request, provider: str, db: Session = Depends
|
||||
A redirect to the user's original destination or the upload page.
|
||||
"""
|
||||
if provider not in SOCIAL_PROVIDERS:
|
||||
logger.debug("[AUTH] social_callback: unknown provider=%r", provider)
|
||||
return RedirectResponse(url="/login?error=Unknown+social+provider", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
oauth_client = getattr(oauth, provider, None)
|
||||
if oauth_client is None:
|
||||
logger.debug("[AUTH] social_callback: provider=%r not configured", provider)
|
||||
return RedirectResponse(url="/login?error=Provider+not+configured", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
try:
|
||||
logger.debug("[AUTH] social_callback: exchanging auth code for provider=%s", provider)
|
||||
token = await oauth_client.authorize_access_token(request)
|
||||
|
||||
# Try standard OIDC userinfo first, fall back to token-embedded userinfo
|
||||
raw_userinfo = token.get("userinfo")
|
||||
if not raw_userinfo:
|
||||
logger.debug("[AUTH] social_callback: no userinfo in token, fetching from userinfo endpoint")
|
||||
try:
|
||||
resp = await oauth_client.userinfo(token=token)
|
||||
raw_userinfo = resp if isinstance(resp, dict) else resp.json() if hasattr(resp, "json") else {}
|
||||
except Exception:
|
||||
logger.debug("[AUTH] social_callback: userinfo endpoint failed, using empty dict", exc_info=True)
|
||||
raw_userinfo = {}
|
||||
|
||||
user_data = _normalize_social_userinfo(provider, token, raw_userinfo)
|
||||
logger.debug(
|
||||
"[AUTH] social_callback: normalized user_data email=%s sub=%s provider=%s",
|
||||
user_data.get("email"),
|
||||
user_data.get("sub"),
|
||||
provider,
|
||||
)
|
||||
|
||||
if not user_data.get("email"):
|
||||
logger.debug("[AUTH] social_callback: no email in user_data — aborting")
|
||||
return RedirectResponse(
|
||||
url="/login?error=Could+not+retrieve+email+from+provider",
|
||||
status_code=status.HTTP_302_FOUND,
|
||||
@@ -442,21 +495,29 @@ async def social_callback(request: Request, provider: str, db: Session = Depends
|
||||
)
|
||||
|
||||
# Mobile app flow: issue an inline API token and redirect back to the app.
|
||||
logger.debug(
|
||||
"[MOBILE] social_callback: checking for mobile redirect (session has mobile_redirect_uri=%s)",
|
||||
"mobile_redirect_uri" in request.session,
|
||||
)
|
||||
mobile_resp = _create_mobile_redirect(request, db)
|
||||
if mobile_resp:
|
||||
logger.info("[MOBILE] social_callback: returning mobile redirect response for provider=%s", provider)
|
||||
return mobile_resp
|
||||
|
||||
if user_id:
|
||||
profile = db.query(_UserProfile).filter(_UserProfile.user_id == user_id).first()
|
||||
if profile and not profile.onboarding_completed:
|
||||
logger.debug("[AUTH] social_callback: user=%s needs onboarding, redirecting", user_id)
|
||||
post_onboarding = request.session.pop("redirect_after_login", "/upload")
|
||||
request.session["post_onboarding_redirect"] = post_onboarding
|
||||
return RedirectResponse(url="/onboarding", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
redirect_url = request.session.pop("redirect_after_login", "/upload")
|
||||
logger.debug("[AUTH] social_callback: login complete, redirecting to %s", redirect_url)
|
||||
return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND)
|
||||
except Exception as e:
|
||||
logger.warning("[SECURITY] SOCIAL_LOGIN_FAILURE provider=%s error=%s", provider, type(e).__name__)
|
||||
logger.debug("[AUTH] social_callback: full exception for provider=%s", provider, exc_info=True)
|
||||
return RedirectResponse(
|
||||
url="/login?error=Social+login+failed.+Please+try+again.", status_code=status.HTTP_302_FOUND
|
||||
)
|
||||
@@ -561,15 +622,23 @@ def _ensure_user_profile(db: Session, user_data: dict, is_admin: bool = False) -
|
||||
async def oauth_callback(request: Request, db: Session = Depends(get_db)):
|
||||
"""Handle OAuth callback from provider"""
|
||||
try:
|
||||
logger.debug("[AUTH] oauth_callback: exchanging authorization code for token")
|
||||
token = await oauth.authentik.authorize_access_token(request)
|
||||
userinfo = token.get("userinfo")
|
||||
if not userinfo:
|
||||
logger.debug("[AUTH] oauth_callback: no userinfo in token response — aborting")
|
||||
return RedirectResponse(
|
||||
url="/login?error=Failed+to+retrieve+user+information", status_code=status.HTTP_302_FOUND
|
||||
)
|
||||
|
||||
# Store user info in session
|
||||
user_data = dict(userinfo)
|
||||
logger.debug(
|
||||
"[AUTH] oauth_callback: received userinfo email=%s sub=%s groups=%s",
|
||||
user_data.get("email"),
|
||||
user_data.get("sub"),
|
||||
user_data.get("groups", []),
|
||||
)
|
||||
|
||||
# Add Gravatar picture if no picture is provided
|
||||
if not user_data.get("picture") and user_data.get("email"):
|
||||
@@ -584,6 +653,12 @@ async def oauth_callback(request: Request, db: Session = Depends(get_db)):
|
||||
groups = user_data.get("groups", [])
|
||||
admin_group = (settings.admin_group_name or "admin").strip().lower()
|
||||
is_admin = admin_group in [group.lower() for group in groups]
|
||||
logger.debug(
|
||||
"[AUTH] oauth_callback: admin group check — looking for %r in %s → is_admin=%s",
|
||||
admin_group,
|
||||
[g.lower() for g in groups],
|
||||
is_admin,
|
||||
)
|
||||
|
||||
# Set is_admin flag (defaults to False for OAuth users unless they're in admin group)
|
||||
user_data["is_admin"] = is_admin
|
||||
@@ -623,15 +698,18 @@ async def oauth_callback(request: Request, db: Session = Depends(get_db)):
|
||||
if user_id:
|
||||
profile = db.query(_UserProfile).filter(_UserProfile.user_id == user_id).first()
|
||||
if profile and not profile.onboarding_completed:
|
||||
logger.debug("[AUTH] oauth_callback: user=%s needs onboarding, redirecting", user_id)
|
||||
post_onboarding = request.session.pop("redirect_after_login", "/upload")
|
||||
request.session["post_onboarding_redirect"] = post_onboarding
|
||||
return RedirectResponse(url="/onboarding", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
# Redirect to original destination or default
|
||||
redirect_url = request.session.pop("redirect_after_login", "/upload")
|
||||
logger.debug("[AUTH] oauth_callback: login complete, redirecting to %s", redirect_url)
|
||||
return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND)
|
||||
except Exception as e:
|
||||
logger.warning(f"[SECURITY] OAUTH_LOGIN_FAILURE error={type(e).__name__}")
|
||||
logger.debug("[AUTH] oauth_callback: full exception details", exc_info=True)
|
||||
return RedirectResponse(url=f"/login?error=Authentication+failed:+{str(e)}", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
|
||||
@@ -931,6 +1009,7 @@ async def logout(request: Request, db: Session = Depends(get_db)):
|
||||
username = "unknown"
|
||||
if isinstance(user, dict):
|
||||
username = user.get("preferred_username") or user.get("email") or "unknown"
|
||||
logger.debug("[AUTH] logout: clearing session for user=%s client_ip=%s", username, get_client_ip(request))
|
||||
logger.info(f"[SECURITY] LOGOUT user={username}")
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
@@ -48,6 +48,51 @@ class Settings(BaseSettings):
|
||||
workdir: str
|
||||
debug: bool = False # Default to False
|
||||
|
||||
# Logging level for the application. Accepts standard Python level names:
|
||||
# DEBUG, INFO, WARNING, ERROR, CRITICAL. When *debug* is True and
|
||||
# *log_level* has not been explicitly set, the effective level is forced to
|
||||
# DEBUG so that all ``logger.debug()`` calls produce output.
|
||||
log_level: str = Field(
|
||||
default="INFO",
|
||||
description=(
|
||||
"Python logging level for the application root logger. "
|
||||
"Accepts: DEBUG, INFO, WARNING, ERROR, CRITICAL. "
|
||||
"When DEBUG=True and LOG_LEVEL is not explicitly set, "
|
||||
"the effective level is automatically lowered to DEBUG."
|
||||
),
|
||||
)
|
||||
|
||||
# Log output format. ``text`` is the human-readable default.
|
||||
# ``json`` emits one JSON object per line, ideal for log collectors
|
||||
# (Promtail, Fluentd, Filebeat, Datadog agent) and SIEM ingestion.
|
||||
log_format: str = Field(
|
||||
default="text",
|
||||
description=(
|
||||
"Log output format: 'text' (human-readable, default) or "
|
||||
"'json' (structured JSON lines for SIEM / log aggregation)."
|
||||
),
|
||||
)
|
||||
|
||||
# Optional syslog forwarding for application logs (not just audit events).
|
||||
# When enabled, a Python SysLogHandler is added to the root logger so that
|
||||
# every log message is also sent to the configured syslog receiver.
|
||||
log_syslog_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Forward application logs to a syslog receiver in addition to stdout.",
|
||||
)
|
||||
log_syslog_host: str = Field(
|
||||
default="localhost",
|
||||
description="Hostname or IP of the syslog receiver for application logs.",
|
||||
)
|
||||
log_syslog_port: int = Field(
|
||||
default=514,
|
||||
description="Port of the syslog receiver for application logs.",
|
||||
)
|
||||
log_syslog_protocol: str = Field(
|
||||
default="udp",
|
||||
description="Protocol for syslog transport: 'udp' or 'tcp'.",
|
||||
)
|
||||
|
||||
# Making Dropbox optional
|
||||
dropbox_enabled: bool = Field(
|
||||
default=True,
|
||||
|
||||
@@ -271,6 +271,7 @@ def _ensure_indexes(engine: Any, inspector: Any) -> None:
|
||||
if table not in columns_by_table:
|
||||
columns_by_table[table] = {col["name"] for col in inspector.get_columns(table)}
|
||||
if column in columns_by_table[table]:
|
||||
# SECURITY: Quoted identifiers to prevent SQL injection during index creation
|
||||
quoted_idx = preparer.quote(idx_name)
|
||||
quoted_table = preparer.quote(table)
|
||||
quoted_col = preparer.quote(column)
|
||||
|
||||
+111
@@ -1,8 +1,11 @@
|
||||
#!/usr/bin/env python3
|
||||
import json as _json_mod
|
||||
import logging
|
||||
import os
|
||||
import pathlib
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime as _dt
|
||||
from datetime import timezone as _tz
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request, status
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
@@ -36,6 +39,114 @@ from app.views import router as frontend_router
|
||||
# Explicitly include the files router
|
||||
from app.views.files import router as files_router
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Configure Python root logging level early so that *all* loggers (including
|
||||
# those already created via ``logging.getLogger(__name__)`` in other modules)
|
||||
# respect the configured level.
|
||||
#
|
||||
# Standard behaviour (matches Django, Flask, 12-factor conventions):
|
||||
# • ``LOG_LEVEL`` env var takes precedence when explicitly set.
|
||||
# • When ``DEBUG=True`` and ``LOG_LEVEL`` is **not** set, the effective
|
||||
# level is automatically lowered to ``DEBUG``.
|
||||
# • Default (neither flag set): ``INFO``.
|
||||
#
|
||||
# ``LOG_FORMAT=json`` enables structured JSON lines on stdout, suitable for
|
||||
# Promtail, Fluentd, Filebeat, Datadog, Splunk UF, or any log collector.
|
||||
#
|
||||
# ``LOG_SYSLOG_ENABLED=true`` adds a Python SysLogHandler so that every log
|
||||
# message is also forwarded to the configured syslog receiver — useful for
|
||||
# traditional (non-container) deployments and centralised SIEM ingestion.
|
||||
#
|
||||
# Noisy third-party loggers (httpx, httpcore, authlib, etc.) are pinned to
|
||||
# WARNING when the app-level is DEBUG to keep output useful.
|
||||
# ---------------------------------------------------------------------------
|
||||
_explicit_log_level = os.environ.get("LOG_LEVEL")
|
||||
if settings.debug and _explicit_log_level is None:
|
||||
_effective_level = "DEBUG"
|
||||
else:
|
||||
_effective_level = settings.log_level.upper()
|
||||
|
||||
_effective_level_int = getattr(logging, _effective_level, logging.INFO)
|
||||
|
||||
|
||||
class _JsonFormatter(logging.Formatter):
|
||||
"""Emit one JSON object per log line for machine consumption.
|
||||
|
||||
Fields emitted: ``timestamp``, ``level``, ``logger``, ``message``,
|
||||
``module``, ``funcName``, ``lineno``, and — when present — ``exc_info``.
|
||||
Compatible with Grafana Loki, Splunk, ELK, Datadog, and most SIEM tools.
|
||||
"""
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
log_entry: dict = {
|
||||
"timestamp": _dt.fromtimestamp(record.created, tz=_tz.utc).isoformat(),
|
||||
"level": record.levelname,
|
||||
"logger": record.name,
|
||||
"message": record.getMessage(),
|
||||
"module": record.module,
|
||||
"funcName": record.funcName,
|
||||
"lineno": record.lineno,
|
||||
}
|
||||
if record.exc_info and record.exc_info[1] is not None:
|
||||
log_entry["exc_info"] = self.formatException(record.exc_info)
|
||||
return _json_mod.dumps(log_entry, default=str)
|
||||
|
||||
|
||||
# Choose formatter based on LOG_FORMAT setting
|
||||
if settings.log_format.lower() == "json":
|
||||
_handler = logging.StreamHandler()
|
||||
_handler.setFormatter(_JsonFormatter())
|
||||
logging.root.handlers = [_handler]
|
||||
logging.root.setLevel(_effective_level_int)
|
||||
else:
|
||||
logging.basicConfig(
|
||||
level=_effective_level_int,
|
||||
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
force=True,
|
||||
)
|
||||
|
||||
# Optional: forward application logs to a syslog receiver
|
||||
if settings.log_syslog_enabled:
|
||||
import logging.handlers as _lh
|
||||
import socket as _socket
|
||||
|
||||
_proto = settings.log_syslog_protocol.lower()
|
||||
_socktype = _socket.SOCK_STREAM if _proto == "tcp" else _socket.SOCK_DGRAM
|
||||
_syslog_handler = _lh.SysLogHandler(
|
||||
address=(settings.log_syslog_host, settings.log_syslog_port),
|
||||
socktype=_socktype,
|
||||
)
|
||||
_syslog_handler.setLevel(_effective_level_int)
|
||||
# Use the same formatter as stdout (text or JSON)
|
||||
if settings.log_format.lower() == "json":
|
||||
_syslog_handler.setFormatter(_JsonFormatter())
|
||||
else:
|
||||
_syslog_handler.setFormatter(logging.Formatter("%(name)s - %(levelname)s - %(message)s"))
|
||||
logging.root.addHandler(_syslog_handler)
|
||||
|
||||
# Keep noisy third-party loggers quiet at DEBUG level
|
||||
if _effective_level_int <= logging.DEBUG:
|
||||
for _noisy in (
|
||||
"httpx",
|
||||
"httpcore",
|
||||
"authlib",
|
||||
"urllib3",
|
||||
"hpack",
|
||||
"multipart",
|
||||
"watchfiles",
|
||||
):
|
||||
logging.getLogger(_noisy).setLevel(logging.WARNING)
|
||||
|
||||
_startup_logger = logging.getLogger(__name__)
|
||||
_startup_logger.info(
|
||||
"Root logging level set to %s (debug=%s, format=%s, syslog=%s)",
|
||||
_effective_level,
|
||||
settings.debug,
|
||||
settings.log_format,
|
||||
settings.log_syslog_enabled,
|
||||
)
|
||||
|
||||
# Load configuration from .env for the session key
|
||||
config = Config(".env")
|
||||
# Use settings.session_secret which has proper validation
|
||||
|
||||
@@ -11,6 +11,7 @@ from email.mime.image import MIMEImage
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from email.mime.text import MIMEText
|
||||
|
||||
import pypdf
|
||||
from jinja2 import Environment, FileSystemLoader, select_autoescape
|
||||
|
||||
from app.celery_app import celery
|
||||
@@ -80,8 +81,22 @@ def extract_metadata_from_file(file_path):
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load metadata from JSON file: {str(e)}")
|
||||
|
||||
# TODO: For PDF files, try to extract embedded metadata using PyPDF2
|
||||
# This would require additional dependencies, so for now we'll just check for external JSON
|
||||
# Try to extract embedded metadata from PDF
|
||||
if file_path.lower().endswith(".pdf") and os.path.exists(file_path):
|
||||
try:
|
||||
with open(file_path, "rb") as f:
|
||||
pdf_reader = pypdf.PdfReader(f)
|
||||
pdf_metadata = pdf_reader.metadata
|
||||
if pdf_metadata:
|
||||
# Convert metadata to a standard dictionary
|
||||
for key, value in pdf_metadata.items():
|
||||
# Remove the leading slash from PDF metadata keys (e.g., '/Title' -> 'Title')
|
||||
clean_key = key[1:] if key.startswith("/") else key
|
||||
metadata[clean_key] = str(value)
|
||||
|
||||
logger.info(f"Extracted embedded metadata from PDF: {file_path}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract metadata from PDF {file_path}: {str(e)}")
|
||||
|
||||
return metadata
|
||||
|
||||
|
||||
@@ -55,12 +55,12 @@ def upload_with_rclone(self, file_path: str, destination: str):
|
||||
|
||||
try:
|
||||
# Ensure the remote path exists (create folders if needed)
|
||||
mkdir_cmd = ["rclone", "mkdir", "--config", rclone_config_path, destination]
|
||||
mkdir_cmd = ["rclone", "mkdir", "--config", rclone_config_path, "--", destination]
|
||||
|
||||
subprocess.run(mkdir_cmd, check=True, capture_output=True) # noqa: S603
|
||||
|
||||
# Construct the upload command
|
||||
upload_cmd = ["rclone", "copy", "--config", rclone_config_path, file_path, destination, "--progress"]
|
||||
upload_cmd = ["rclone", "copy", "--config", rclone_config_path, "--progress", "--", file_path, destination]
|
||||
|
||||
log_task_progress(task_id, "rclone_upload", "in_progress", f"Executing rclone copy to {destination}")
|
||||
|
||||
@@ -71,7 +71,7 @@ def upload_with_rclone(self, file_path: str, destination: str):
|
||||
if result.returncode == 0:
|
||||
# Try to get a public link if possible
|
||||
try:
|
||||
link_cmd = ["rclone", "link", "--config", rclone_config_path, f"{destination}/{filename}"]
|
||||
link_cmd = ["rclone", "link", "--config", rclone_config_path, "--", f"{destination}/{filename}"]
|
||||
link_result = subprocess.run(link_cmd, capture_output=True, text=True, check=False) # noqa: S603
|
||||
public_url = link_result.stdout.strip() if link_result.returncode == 0 else None
|
||||
except (subprocess.SubprocessError, OSError) as e:
|
||||
|
||||
@@ -12,6 +12,7 @@ The utility:
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import MetaData, create_engine, inspect, text
|
||||
@@ -84,6 +85,9 @@ def preview_migration(source_url: str) -> dict[str, Any]:
|
||||
total = 0
|
||||
with src_engine.connect() as conn:
|
||||
for table_name in tables:
|
||||
if not re.match(r"^[a-zA-Z0-9_]+$", table_name):
|
||||
logger.warning(f"Skipping table with invalid name format: {table_name}")
|
||||
continue
|
||||
# table_name is safe — sourced from inspect().get_table_names(), not user input
|
||||
quoted_table = conn.dialect.identifier_preparer.quote(table_name)
|
||||
row = conn.execute(text(f"SELECT COUNT(*) FROM {quoted_table}")).fetchone() # noqa: S608
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
import logging
|
||||
import os
|
||||
from typing import Dict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def update_env_file(settings_to_update: Dict[str, str]) -> bool:
|
||||
"""
|
||||
Updates the .env file with the given settings (best-effort).
|
||||
Creates or modifies existing keys.
|
||||
|
||||
Args:
|
||||
settings_to_update: A dictionary mapping uppercase env var names to their new string values.
|
||||
|
||||
Returns:
|
||||
True if the file was successfully updated, False otherwise.
|
||||
"""
|
||||
try:
|
||||
env_path = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), ".env")
|
||||
if not os.path.exists(env_path):
|
||||
logger.warning(f".env file not found at {env_path}, skipping file write")
|
||||
return False
|
||||
|
||||
logger.info(f"Updating settings in {env_path}")
|
||||
|
||||
with open(env_path, "r") as f:
|
||||
env_lines = f.readlines()
|
||||
|
||||
updated = set()
|
||||
new_env_lines = []
|
||||
for line in env_lines:
|
||||
stripped_line = line.rstrip()
|
||||
is_updated = False
|
||||
for key, value in settings_to_update.items():
|
||||
if stripped_line.startswith(f"{key}=") or stripped_line.startswith(f"# {key}="):
|
||||
new_env_lines.append(f"{key}={value}")
|
||||
updated.add(key)
|
||||
is_updated = True
|
||||
break
|
||||
if not is_updated:
|
||||
new_env_lines.append(stripped_line)
|
||||
|
||||
for key, value in settings_to_update.items():
|
||||
if key not in updated:
|
||||
new_env_lines.append(f"{key}={value}")
|
||||
|
||||
with open(env_path, "w") as f:
|
||||
f.write("\n".join(new_env_lines) + "\n")
|
||||
|
||||
logger.info("Successfully updated settings in .env file")
|
||||
return True
|
||||
except Exception as env_err:
|
||||
logger.warning(f"Failed to write .env file (non-fatal): {env_err}")
|
||||
return False
|
||||
@@ -2628,6 +2628,63 @@ SETTING_METADATA = {
|
||||
"required": False,
|
||||
"restart_required": False,
|
||||
},
|
||||
# Logging
|
||||
"log_level": {
|
||||
"category": "Observability",
|
||||
"description": (
|
||||
"Python logging level for the application root logger. "
|
||||
"Accepts: DEBUG, INFO, WARNING, ERROR, CRITICAL. "
|
||||
"When DEBUG=True and LOG_LEVEL is not explicitly set, "
|
||||
"the effective level is automatically lowered to DEBUG."
|
||||
),
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"log_format": {
|
||||
"category": "Observability",
|
||||
"description": (
|
||||
"Log output format: 'text' (human-readable, default) or "
|
||||
"'json' (structured JSON lines for SIEM / log aggregation)."
|
||||
),
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"log_syslog_enabled": {
|
||||
"category": "Observability",
|
||||
"description": "Forward application logs to a syslog receiver in addition to stdout.",
|
||||
"type": "boolean",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"log_syslog_host": {
|
||||
"category": "Observability",
|
||||
"description": "Hostname or IP of the syslog receiver for application logs.",
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"log_syslog_port": {
|
||||
"category": "Observability",
|
||||
"description": "Port of the syslog receiver for application logs.",
|
||||
"type": "integer",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
"log_syslog_protocol": {
|
||||
"category": "Observability",
|
||||
"description": "Protocol for syslog transport: 'udp' or 'tcp'.",
|
||||
"type": "string",
|
||||
"sensitive": False,
|
||||
"required": False,
|
||||
"restart_required": True,
|
||||
},
|
||||
# Observability – Sentry
|
||||
"sentry_dsn": {
|
||||
"category": "Observability",
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
import time
|
||||
import os
|
||||
import sys
|
||||
import asyncio
|
||||
|
||||
# Mock settings before app imports to bypass validation
|
||||
os.environ["DATABASE_URL"] = "sqlite:///:memory:"
|
||||
os.environ["REDIS_URL"] = "redis://localhost:6379"
|
||||
os.environ["OPENAI_API_KEY"] = "mock_key"
|
||||
os.environ["WORKDIR"] = "/tmp/workdir"
|
||||
os.environ["AZURE_AI_KEY"] = "mock"
|
||||
os.environ["AZURE_REGION"] = "mock"
|
||||
os.environ["AZURE_ENDPOINT"] = "http://mock"
|
||||
os.environ["GOTENBERG_URL"] = "http://mock"
|
||||
os.environ["SESSION_SECRET"] = "mock_secret_mock_secret_mock_secret_mock_secret"
|
||||
|
||||
# Ensure app package is accessible
|
||||
sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from app.database import Base
|
||||
from app.models import FileRecord
|
||||
from app.api.duplicates import list_duplicate_groups
|
||||
|
||||
# Mocking Request object
|
||||
class MockRequest:
|
||||
def __init__(self):
|
||||
self.session = {"user": {"username": "testuser"}}
|
||||
self.state = type('State', (), {'user': {"username": "testuser"}})()
|
||||
|
||||
class MockURL:
|
||||
def include_query_params(self, **kwargs):
|
||||
return f"http://testserver/api/duplicates?page={kwargs.get('page')}"
|
||||
url = MockURL()
|
||||
|
||||
def setup_db():
|
||||
engine = create_engine('sqlite:///:memory:')
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine)
|
||||
db = Session()
|
||||
return db
|
||||
|
||||
def populate_data(db, num_groups, duplicates_per_group):
|
||||
for i in range(num_groups):
|
||||
filehash = f"hash_{i}"
|
||||
# Original
|
||||
original = FileRecord(
|
||||
filehash=filehash,
|
||||
local_filename=f"orig_{i}.txt",
|
||||
file_size=100,
|
||||
is_duplicate=False
|
||||
)
|
||||
db.add(original)
|
||||
# Duplicates
|
||||
for j in range(duplicates_per_group):
|
||||
dup = FileRecord(
|
||||
filehash=filehash,
|
||||
local_filename=f"dup_{i}_{j}.txt",
|
||||
file_size=100,
|
||||
is_duplicate=True
|
||||
)
|
||||
db.add(dup)
|
||||
db.commit()
|
||||
|
||||
async def run_benchmark(db):
|
||||
request = MockRequest()
|
||||
start_time = time.time()
|
||||
|
||||
# Run the function we want to benchmark
|
||||
result = list_duplicate_groups(request=request, db=db, page=1, per_page=500)
|
||||
if asyncio.iscoroutine(result):
|
||||
result = await result
|
||||
|
||||
end_time = time.time()
|
||||
return end_time - start_time, result
|
||||
|
||||
async def main():
|
||||
db = setup_db()
|
||||
# 500 groups, each with 20 duplicates = 10500 records total
|
||||
print("Populating data...")
|
||||
populate_data(db, 500, 20)
|
||||
print("Data populated. Running baseline benchmark...")
|
||||
|
||||
# Warmup
|
||||
result = list_duplicate_groups(request=MockRequest(), db=db, page=1, per_page=500)
|
||||
if asyncio.iscoroutine(result):
|
||||
await result
|
||||
|
||||
# Benchmark
|
||||
total_time = 0
|
||||
iterations = 10
|
||||
for _ in range(iterations):
|
||||
time_taken, _ = await run_benchmark(db)
|
||||
total_time += time_taken
|
||||
|
||||
avg_time = total_time / iterations
|
||||
print(f"Average time over {iterations} iterations: {avg_time:.4f} seconds")
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -0,0 +1,50 @@
|
||||
import json
|
||||
import time
|
||||
import pytest
|
||||
from app.database import get_db
|
||||
from app.models import UserNotificationTarget, UserNotificationPreference
|
||||
from app.main import app
|
||||
from tests.test_notifications_api import _make_client, _OWNER, _cleanup
|
||||
import statistics
|
||||
|
||||
def run_benchmark(notif_engine, notif_session, client, items_count, iterations=5):
|
||||
# Setup
|
||||
target = UserNotificationTarget(
|
||||
owner_id=_OWNER,
|
||||
channel_type="webhook",
|
||||
name="My Webhook",
|
||||
config=json.dumps({"url": "https://x.com"}),
|
||||
)
|
||||
notif_session.add(target)
|
||||
notif_session.commit()
|
||||
notif_session.refresh(target)
|
||||
|
||||
# Generate big payload
|
||||
preferences = []
|
||||
for i in range(items_count):
|
||||
preferences.append({
|
||||
"event_type": f"event.type.{i}",
|
||||
"channel_type": "webhook",
|
||||
"is_enabled": True,
|
||||
"target_id": target.id,
|
||||
})
|
||||
|
||||
payload = {"preferences": preferences}
|
||||
|
||||
# Warm up
|
||||
client.put("/api/user-notifications/preferences", json=payload)
|
||||
|
||||
times = []
|
||||
for _ in range(iterations):
|
||||
# Alter the values a bit so it's a real update
|
||||
for p in payload["preferences"]:
|
||||
p["is_enabled"] = not p["is_enabled"]
|
||||
|
||||
start = time.time()
|
||||
resp = client.put("/api/user-notifications/preferences", json=payload)
|
||||
end = time.time()
|
||||
|
||||
assert resp.status_code == 200
|
||||
times.append(end - start)
|
||||
|
||||
return statistics.mean(times)
|
||||
@@ -457,6 +457,87 @@ default overage buffer applied across all plans.
|
||||
|
||||
DocuElevate supports HTTP security headers to improve browser-side security. **These headers are disabled by default** since most deployments use a reverse proxy (Traefik, Nginx, etc.) that already adds them. Enable only if deploying directly without a reverse proxy. See [Deployment Guide - Security Headers](DeploymentGuide.md#security-headers) for detailed configuration examples.
|
||||
|
||||
### Application Logging
|
||||
|
||||
DocuElevate uses Python's standard `logging` module. Two environment variables control log verbosity:
|
||||
|
||||
| **Variable** | **Description** | **Default** |
|
||||
|-------------|----------------|-------------|
|
||||
| `LOG_LEVEL` | Root logger level. Accepts standard Python level names: `DEBUG`, `INFO`, `WARNING`, `ERROR`, `CRITICAL`. | `INFO` |
|
||||
| `DEBUG` | Enable debug mode. When `true` **and** `LOG_LEVEL` is **not** explicitly set, the effective log level is automatically lowered to `DEBUG`. | `false` |
|
||||
|
||||
**Precedence rules (standard behaviour):**
|
||||
|
||||
1. If `LOG_LEVEL` is explicitly set, it always wins — regardless of `DEBUG`.
|
||||
2. If only `DEBUG=true` is set (no `LOG_LEVEL`), the effective level becomes `DEBUG`.
|
||||
3. If neither is set, the default level is `INFO`.
|
||||
|
||||
```bash
|
||||
# Typical production (default)
|
||||
# LOG_LEVEL=INFO
|
||||
|
||||
# Quick debug mode — sets level to DEBUG automatically
|
||||
DEBUG=true
|
||||
|
||||
# Explicit level override (DEBUG flag is ignored for level selection)
|
||||
LOG_LEVEL=WARNING
|
||||
```
|
||||
|
||||
> **Tip:** At `DEBUG` level, noisy third-party libraries (httpx, authlib, urllib3, etc.) are automatically pinned to `WARNING` so that application debug output remains readable.
|
||||
|
||||
#### Structured JSON Logging
|
||||
|
||||
Set `LOG_FORMAT=json` to emit structured JSON lines on stdout — one JSON object per log message. This is the standard format for log collectors and SIEM tools:
|
||||
|
||||
| **Variable** | **Description** | **Default** |
|
||||
|-------------|----------------|-------------|
|
||||
| `LOG_FORMAT` | Log output format: `text` (human-readable) or `json` (structured JSON lines). | `text` |
|
||||
|
||||
Each JSON log line contains: `timestamp` (ISO 8601), `level`, `logger`, `message`, `module`, `funcName`, `lineno`, and `exc_info` (when an exception is logged).
|
||||
|
||||
```bash
|
||||
# Enable JSON logging for SIEM / log aggregation
|
||||
LOG_FORMAT=json
|
||||
```
|
||||
|
||||
**Example JSON output:**
|
||||
```json
|
||||
{"timestamp": "2025-03-16T09:18:05.192000+00:00", "level": "INFO", "logger": "app.auth", "message": "[SECURITY] OAUTH_LOGIN_SUCCESS user=alice@example.com admin=False", "module": "auth", "funcName": "oauth_callback", "lineno": 654}
|
||||
```
|
||||
|
||||
**Compatible with:**
|
||||
- **Grafana Loki** — Promtail scrapes JSON from Docker stdout
|
||||
- **Splunk** — Universal Forwarder or HEC with JSON sourcetype
|
||||
- **ELK / OpenSearch** — Filebeat with JSON codec
|
||||
- **Datadog** — Agent auto-parses JSON logs
|
||||
- **Fluentd / Vector** — JSON input plugin
|
||||
- **Docker log drivers** — `--log-driver=json-file` (default) preserves structure
|
||||
|
||||
#### Syslog Forwarding (Application Logs)
|
||||
|
||||
For traditional (non-container) deployments, application logs can be forwarded directly to a syslog receiver. This is **separate** from audit-log SIEM forwarding (see below) — it sends _every_ Python log message, not just audit events.
|
||||
|
||||
| **Variable** | **Description** | **Default** |
|
||||
|-------------|----------------|-------------|
|
||||
| `LOG_SYSLOG_ENABLED` | Forward application logs to a syslog receiver in addition to stdout. | `false` |
|
||||
| `LOG_SYSLOG_HOST` | Hostname or IP of the syslog receiver. | `localhost` |
|
||||
| `LOG_SYSLOG_PORT` | Port of the syslog receiver. | `514` |
|
||||
| `LOG_SYSLOG_PROTOCOL` | Protocol: `udp` or `tcp`. | `udp` |
|
||||
|
||||
```bash
|
||||
# Forward all application logs to syslog
|
||||
LOG_SYSLOG_ENABLED=true
|
||||
LOG_SYSLOG_HOST=syslog.internal.example.com
|
||||
LOG_SYSLOG_PORT=514
|
||||
LOG_SYSLOG_PROTOCOL=udp
|
||||
|
||||
# Combine with JSON format for structured syslog messages
|
||||
LOG_FORMAT=json
|
||||
LOG_SYSLOG_ENABLED=true
|
||||
```
|
||||
|
||||
> **Note:** When `LOG_FORMAT=json`, syslog messages are also sent as JSON. When `LOG_FORMAT=text`, syslog messages use the standard `name - level - message` format.
|
||||
|
||||
### Audit Logging
|
||||
|
||||
DocuElevate provides comprehensive audit logging that records significant actions (logins, document CRUD, settings changes) to an append-only database table. Every entry captures the timestamp, user, action, resource, client IP, and optional JSON details.
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
import re
|
||||
|
||||
with open("tests/test_api_saved_searches.py", "r") as f:
|
||||
content = f.read()
|
||||
|
||||
# We need to mock get_current_user in app.api.saved_searches (which is imported from app.auth)
|
||||
# because saved searches uses `_get_user_id` which calls `get_current_user(request)`.
|
||||
# But `_get_user_id` is NOT a dependency injected via `Depends`!
|
||||
# Let's verify `app/api/saved_searches.py` uses `Depends` or just calls it.
|
||||
|
||||
# In `app/api/saved_searches.py`:
|
||||
# def _get_user_id(request: Request) -> str:
|
||||
# user = get_current_user(request)
|
||||
# if user:
|
||||
# return user.get("preferred_username") ...
|
||||
# It's called directly inside the routes: `user_id = _get_user_id(request)`
|
||||
# It doesn't use `Depends(_get_user_id)`.
|
||||
# Ah! But earlier I saw `_get_user_id` wasn't mocked properly. Let's use patch to mock `_get_user_id`.
|
||||
|
||||
# Wait, `TestClient` can be given an active session, but `app.auth.get_current_user` uses `request.session.get("user")` or Bearer token.
|
||||
# Is `AUTH_ENABLED` false? The test env has `os.environ["AUTH_ENABLED"] = "False"` in `tests/conftest.py`.
|
||||
# If `AUTH_ENABLED` is false, `require_login` is a no-op, and `_get_user_id` falls back to "anonymous".
|
||||
# Actually, `_get_user_id` returns "anonymous" if `get_current_user(request)` is None.
|
||||
# If `_OWNER` is "test_user@example.com", we should probably just patch `_get_user_id`.
|
||||
|
||||
replacement = """def _make_client(int_engine, owner_id: str = _OWNER):
|
||||
\"\"\"Return a TestClient with *owner_id* injected as the authenticated user.\"\"\"
|
||||
from app.main import app
|
||||
from unittest.mock import patch
|
||||
|
||||
def override_db():
|
||||
Session = sessionmaker(bind=int_engine)
|
||||
session = Session()
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
app.dependency_overrides[get_db] = override_db
|
||||
with patch("app.api.saved_searches._get_user_id", return_value=owner_id):
|
||||
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
|
||||
yield client
|
||||
app.dependency_overrides.clear()"""
|
||||
|
||||
content = re.sub(
|
||||
r"def _make_client\(int_engine, owner_id: str = _OWNER\):.*?(?=@pytest\.fixture\(\)\ndef int_client\(int_engine\):)",
|
||||
replacement + "\n\n\n",
|
||||
content,
|
||||
flags=re.DOTALL
|
||||
)
|
||||
|
||||
with open("tests/test_api_saved_searches.py", "w") as f:
|
||||
f.write(content)
|
||||
@@ -23,6 +23,7 @@ djlint>=1.36.0 # HTML template linter for accessibility and best practices
|
||||
mkdocs-material>=9.5.0 # MkDocs Material theme – same package used in docs/requirements.txt
|
||||
|
||||
# Type stubs for mypy
|
||||
types-aiofiles>=24.1.0
|
||||
types-requests>=2.31.0
|
||||
types-paramiko>=3.0.0
|
||||
|
||||
@@ -37,4 +38,5 @@ pip-licenses==5.5.1 # For license compliance checking
|
||||
|
||||
# Release automation
|
||||
python-semantic-release>=9.0.0
|
||||
types-aiofiles>=23.2.0.20240106
|
||||
|
||||
types-aiofiles>=24.1.0.20240311 # Type stubs for aiofiles
|
||||
|
||||
+2
-1
@@ -58,4 +58,5 @@ sentry-sdk[fastapi,celery,sqlalchemy]>=2.20.0,<3.0.0
|
||||
|
||||
# GraphQL API
|
||||
strawberry-graphql[fastapi]>=0.243.0,<1.0.0
|
||||
aiofiles>=23.2.1
|
||||
|
||||
aiofiles>=24.1.0 # Asynchronous file I/O support
|
||||
|
||||
Executable
+2
@@ -0,0 +1,2 @@
|
||||
#!/bin/bash
|
||||
pytest tests/ -k "not test_e2e_full_stack and not test_upload_tasks and not test_slow" -m "not slow" -n 4
|
||||
@@ -269,6 +269,53 @@ class TestSavedSearchesCRUD:
|
||||
response2 = client.post("/api/saved-searches", json=payload)
|
||||
assert response2.status_code == 409
|
||||
|
||||
def test_create_saved_search_db_error(self, client: TestClient, monkeypatch):
|
||||
"""POST /api/saved-searches returns 500 on DB exception."""
|
||||
# Mock db.add or db.commit to raise an exception
|
||||
# We can monkeypatch the route's dependency or the models
|
||||
# It's easier to mock the SavedSearch model's __init__ or db's add
|
||||
# Since we use db: DbSession, it's an instance of sqlalchemy.orm.Session
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
original_commit = Session.commit
|
||||
|
||||
def mock_commit(*args, **kwargs):
|
||||
raise Exception("Simulated DB error")
|
||||
|
||||
monkeypatch.setattr(Session, "commit", mock_commit)
|
||||
|
||||
payload = {
|
||||
"name": "DB Error Search",
|
||||
"filters": {"status": "completed"},
|
||||
}
|
||||
response = client.post("/api/saved-searches", json=payload)
|
||||
assert response.status_code == 500
|
||||
assert "Failed to save search" in response.json()["detail"]
|
||||
|
||||
def test_create_saved_search_limit_reached(self, client: TestClient, monkeypatch):
|
||||
"""POST /api/saved-searches returns 409 if max limit is reached."""
|
||||
monkeypatch.setattr("app.api.saved_searches.MAX_SAVED_SEARCHES_PER_USER", 1)
|
||||
|
||||
# Create first one
|
||||
payload1 = {"name": "Search 1", "filters": {"status": "completed"}}
|
||||
response1 = client.post("/api/saved-searches", json=payload1)
|
||||
assert response1.status_code == 201
|
||||
|
||||
# Creating second one should fail due to limit
|
||||
payload2 = {"name": "Search 2", "filters": {"status": "pending"}}
|
||||
response2 = client.post("/api/saved-searches", json=payload2)
|
||||
assert response2.status_code == 409
|
||||
assert "Maximum of 1 saved searches reached" in response2.json()["detail"]
|
||||
|
||||
def test_create_saved_search_invalid_name_type(self, client: TestClient):
|
||||
"""POST /api/saved-searches with non-string name returns 422."""
|
||||
payload = {
|
||||
"name": 12345,
|
||||
"filters": {"status": "completed"},
|
||||
}
|
||||
response = client.post("/api/saved-searches", json=payload)
|
||||
assert response.status_code == 422
|
||||
|
||||
def test_update_saved_search(self, client: TestClient):
|
||||
"""PUT /api/saved-searches/{id} updates the saved search."""
|
||||
# Create
|
||||
@@ -296,6 +343,83 @@ class TestSavedSearchesCRUD:
|
||||
)
|
||||
assert response.status_code == 404
|
||||
|
||||
def test_update_saved_search_duplicate_name(self, client: TestClient):
|
||||
"""PUT /api/saved-searches/{id} with duplicate name returns 409."""
|
||||
# Create first search
|
||||
client.post(
|
||||
"/api/saved-searches",
|
||||
json={"name": "First Search", "filters": {"status": "pending"}},
|
||||
)
|
||||
# Create second search
|
||||
create_resp2 = client.post(
|
||||
"/api/saved-searches",
|
||||
json={"name": "Second Search", "filters": {"status": "completed"}},
|
||||
)
|
||||
search_id2 = create_resp2.json()["id"]
|
||||
|
||||
# Try to rename second search to "First Search"
|
||||
update_resp = client.put(
|
||||
f"/api/saved-searches/{search_id2}",
|
||||
json={"name": "First Search", "filters": {"status": "completed"}},
|
||||
)
|
||||
assert update_resp.status_code == 409
|
||||
|
||||
def test_update_saved_search_name_too_long(self, client: TestClient):
|
||||
"""PUT /api/saved-searches/{id} with name > 100 chars returns 422."""
|
||||
create_resp = client.post(
|
||||
"/api/saved-searches",
|
||||
json={"name": "Valid Name", "filters": {"status": "pending"}},
|
||||
)
|
||||
search_id = create_resp.json()["id"]
|
||||
|
||||
update_resp = client.put(
|
||||
f"/api/saved-searches/{search_id}",
|
||||
json={"name": "x" * 101, "filters": {"status": "completed"}},
|
||||
)
|
||||
assert update_resp.status_code == 422
|
||||
|
||||
def test_update_saved_search_empty_name(self, client: TestClient):
|
||||
"""PUT /api/saved-searches/{id} with empty name returns 422."""
|
||||
create_resp = client.post(
|
||||
"/api/saved-searches",
|
||||
json={"name": "Valid Name", "filters": {"status": "pending"}},
|
||||
)
|
||||
search_id = create_resp.json()["id"]
|
||||
|
||||
update_resp = client.put(
|
||||
f"/api/saved-searches/{search_id}",
|
||||
json={"name": "", "filters": {"status": "completed"}},
|
||||
)
|
||||
assert update_resp.status_code == 422
|
||||
|
||||
def test_update_saved_search_empty_filters(self, client: TestClient):
|
||||
"""PUT /api/saved-searches/{id} with empty filters returns 422."""
|
||||
create_resp = client.post(
|
||||
"/api/saved-searches",
|
||||
json={"name": "Valid Name", "filters": {"status": "pending"}},
|
||||
)
|
||||
search_id = create_resp.json()["id"]
|
||||
|
||||
update_resp = client.put(
|
||||
f"/api/saved-searches/{search_id}",
|
||||
json={"name": "Valid Name", "filters": {}},
|
||||
)
|
||||
assert update_resp.status_code == 422
|
||||
|
||||
def test_update_saved_search_invalid_filters(self, client: TestClient):
|
||||
"""PUT /api/saved-searches/{id} with only invalid filters returns 422."""
|
||||
create_resp = client.post(
|
||||
"/api/saved-searches",
|
||||
json={"name": "Valid Name", "filters": {"status": "pending"}},
|
||||
)
|
||||
search_id = create_resp.json()["id"]
|
||||
|
||||
update_resp = client.put(
|
||||
f"/api/saved-searches/{search_id}",
|
||||
json={"name": "Valid Name", "filters": {"invalid_key": "value"}},
|
||||
)
|
||||
assert update_resp.status_code == 422
|
||||
|
||||
def test_delete_saved_search(self, client: TestClient):
|
||||
"""DELETE /api/saved-searches/{id} removes the saved search."""
|
||||
# Create
|
||||
@@ -318,6 +442,22 @@ class TestSavedSearchesCRUD:
|
||||
response = client.delete("/api/saved-searches/999")
|
||||
assert response.status_code == 404
|
||||
|
||||
def test_delete_saved_search_db_error(self, client: TestClient):
|
||||
"""DELETE /api/saved-searches/{id} handles database errors (500)."""
|
||||
from unittest.mock import patch
|
||||
|
||||
# Create
|
||||
create_resp = client.post(
|
||||
"/api/saved-searches",
|
||||
json={"name": "To Delete DB Error", "filters": {"status": "failed"}},
|
||||
)
|
||||
search_id = create_resp.json()["id"]
|
||||
|
||||
with patch("sqlalchemy.orm.Session.delete", side_effect=Exception("DB Delete Error")):
|
||||
response = client.delete(f"/api/saved-searches/{search_id}")
|
||||
assert response.status_code == 500
|
||||
assert response.json()["detail"] == "Failed to delete saved search"
|
||||
|
||||
def test_create_name_too_long(self, client: TestClient):
|
||||
"""POST /api/saved-searches with name > 100 chars returns 422."""
|
||||
payload = {
|
||||
|
||||
@@ -7,7 +7,6 @@ Covers Dropbox OAuth endpoints, settings management, and token testing.
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
@@ -137,7 +136,7 @@ class TestTestDropboxToken:
|
||||
assert data["status"] == "error"
|
||||
assert "not fully configured" in data["message"]
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
@patch("app.api.dropbox.httpx.AsyncClient.post")
|
||||
@patch("app.api.dropbox.settings")
|
||||
def test_valid_token(self, mock_settings, mock_post, client):
|
||||
"""Test successful token validation."""
|
||||
@@ -162,7 +161,7 @@ class TestTestDropboxToken:
|
||||
assert data["account"] == "user@example.com"
|
||||
assert data["account_name"] == "Test User"
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
@patch("app.api.dropbox.httpx.AsyncClient.post")
|
||||
@patch("app.api.dropbox.settings")
|
||||
def test_expired_token_refreshed(self, mock_settings, mock_post, client):
|
||||
"""Test that expired token triggers refresh and retry."""
|
||||
@@ -194,7 +193,7 @@ class TestTestDropboxToken:
|
||||
data = response.json()
|
||||
assert data["status"] == "success"
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
@patch("app.api.dropbox.httpx.AsyncClient.post")
|
||||
@patch("app.api.dropbox.settings")
|
||||
def test_refresh_token_expired(self, mock_settings, mock_post, client):
|
||||
"""Test handling when refresh token itself is expired."""
|
||||
@@ -220,7 +219,7 @@ class TestTestDropboxToken:
|
||||
assert data["status"] == "error"
|
||||
assert data["needs_reauth"] is True
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
@patch("app.api.dropbox.httpx.AsyncClient.post")
|
||||
@patch("app.api.dropbox.settings")
|
||||
def test_token_validation_failure(self, mock_settings, mock_post, client):
|
||||
"""Test handling non-401, non-200 response."""
|
||||
@@ -240,16 +239,21 @@ class TestTestDropboxToken:
|
||||
data = response.json()
|
||||
assert data["status"] == "error"
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
@patch("app.api.dropbox.httpx.AsyncClient.post")
|
||||
@patch("app.api.dropbox.settings")
|
||||
def test_connection_error(self, mock_settings, mock_post, client):
|
||||
"""Test handling of connection exceptions."""
|
||||
import httpx
|
||||
|
||||
mock_settings.dropbox_refresh_token = "token"
|
||||
mock_settings.dropbox_app_key = "app-key"
|
||||
mock_settings.dropbox_app_secret = "app-secret"
|
||||
mock_settings.http_request_timeout = 30
|
||||
|
||||
mock_post.side_effect = requests.exceptions.ConnectionError("Connection refused")
|
||||
mock_post.side_effect = httpx.RequestError(
|
||||
"Connection refused",
|
||||
request=httpx.Request("POST", "https://api.dropboxapi.com/2/users/get_current_account"),
|
||||
)
|
||||
|
||||
response = client.get("/api/dropbox/test-token")
|
||||
|
||||
|
||||
@@ -84,8 +84,8 @@ class TestUpdateDropboxSettings:
|
||||
class TestTestDropboxToken:
|
||||
"""Tests for GET /dropbox/test-token endpoint."""
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
def test_test_token_success(self, mock_post):
|
||||
@patch("app.api.dropbox.httpx.AsyncClient")
|
||||
def test_test_token_success(self, mock_client_cls):
|
||||
"""Test successful token validation."""
|
||||
from app.config import settings
|
||||
|
||||
@@ -95,7 +95,7 @@ class TestTestDropboxToken:
|
||||
"email": "test@example.com",
|
||||
"name": {"display_name": "Test User"},
|
||||
}
|
||||
mock_post.return_value = mock_response
|
||||
mock_client_cls.return_value.__aenter__.return_value.post.return_value = mock_response
|
||||
|
||||
with patch.object(settings, "dropbox_refresh_token", "token"):
|
||||
with patch.object(settings, "dropbox_app_key", "key"):
|
||||
@@ -104,8 +104,8 @@ class TestTestDropboxToken:
|
||||
# Should include account email and name
|
||||
pass
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
def test_test_token_not_configured(self, mock_post):
|
||||
@patch("app.api.dropbox.httpx.AsyncClient")
|
||||
def test_test_token_not_configured(self, mock_client_cls):
|
||||
"""Test when credentials are not configured."""
|
||||
from app.config import settings
|
||||
|
||||
@@ -113,8 +113,8 @@ class TestTestDropboxToken:
|
||||
# Should return error indicating not configured
|
||||
pass
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
def test_test_token_partial_config(self, mock_post):
|
||||
@patch("app.api.dropbox.httpx.AsyncClient")
|
||||
def test_test_token_partial_config(self, mock_client_cls):
|
||||
"""Test with partial configuration (missing some credentials)."""
|
||||
from app.config import settings
|
||||
|
||||
@@ -123,8 +123,8 @@ class TestTestDropboxToken:
|
||||
# Should return error
|
||||
pass
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
def test_test_token_expired_requires_refresh(self, mock_post):
|
||||
@patch("app.api.dropbox.httpx.AsyncClient")
|
||||
def test_test_token_expired_requires_refresh(self, mock_client_cls):
|
||||
"""Test when access token is expired and needs refresh."""
|
||||
from app.config import settings
|
||||
|
||||
@@ -145,7 +145,9 @@ class TestTestDropboxToken:
|
||||
"name": {"display_name": "Test User"},
|
||||
}
|
||||
|
||||
mock_post.side_effect = [mock_response_401, mock_refresh_response, mock_success_response]
|
||||
mock_client = MagicMock()
|
||||
mock_client.post.side_effect = [mock_response_401, mock_refresh_response, mock_success_response]
|
||||
mock_client_cls.return_value.__aenter__.return_value = mock_client
|
||||
|
||||
with patch.object(settings, "dropbox_refresh_token", "token"):
|
||||
with patch.object(settings, "dropbox_app_key", "key"):
|
||||
@@ -153,8 +155,8 @@ class TestTestDropboxToken:
|
||||
# Should refresh and succeed
|
||||
pass
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
def test_test_token_refresh_failed(self, mock_post):
|
||||
@patch("app.api.dropbox.httpx.AsyncClient")
|
||||
def test_test_token_refresh_failed(self, mock_client_cls):
|
||||
"""Test when refresh token is invalid."""
|
||||
from app.config import settings
|
||||
|
||||
@@ -167,7 +169,9 @@ class TestTestDropboxToken:
|
||||
mock_refresh_response.status_code = 400
|
||||
mock_refresh_response.text = "Invalid refresh token"
|
||||
|
||||
mock_post.side_effect = [mock_response_401, mock_refresh_response]
|
||||
mock_client = MagicMock()
|
||||
mock_client.post.side_effect = [mock_response_401, mock_refresh_response]
|
||||
mock_client_cls.return_value.__aenter__.return_value = mock_client
|
||||
|
||||
with patch.object(settings, "dropbox_refresh_token", "token"):
|
||||
with patch.object(settings, "dropbox_app_key", "key"):
|
||||
@@ -175,8 +179,8 @@ class TestTestDropboxToken:
|
||||
# Should return error with needs_reauth: True
|
||||
pass
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
def test_test_token_perpetual_token_info(self, mock_post):
|
||||
@patch("app.api.dropbox.httpx.AsyncClient")
|
||||
def test_test_token_perpetual_token_info(self, mock_client_cls):
|
||||
"""Test that perpetual token info is returned."""
|
||||
from app.config import settings
|
||||
|
||||
@@ -186,7 +190,7 @@ class TestTestDropboxToken:
|
||||
"email": "test@example.com",
|
||||
"name": {"display_name": "Test User"},
|
||||
}
|
||||
mock_post.return_value = mock_response
|
||||
mock_client_cls.return_value.__aenter__.return_value.post.return_value = mock_response
|
||||
|
||||
with patch.object(settings, "dropbox_refresh_token", "token"):
|
||||
with patch.object(settings, "dropbox_app_key", "key"):
|
||||
@@ -194,12 +198,12 @@ class TestTestDropboxToken:
|
||||
# token_info should indicate never expires
|
||||
pass
|
||||
|
||||
@patch("app.api.dropbox.requests.post")
|
||||
def test_test_token_exception_handling(self, mock_post):
|
||||
@patch("app.api.dropbox.httpx.AsyncClient")
|
||||
def test_test_token_exception_handling(self, mock_client_cls):
|
||||
"""Test handling of exceptions."""
|
||||
from app.config import settings
|
||||
|
||||
mock_post.side_effect = Exception("Network error")
|
||||
mock_client_cls.return_value.__aenter__.return_value.post.side_effect = Exception("Network error")
|
||||
|
||||
with patch.object(settings, "dropbox_refresh_token", "token"):
|
||||
with patch.object(settings, "dropbox_app_key", "key"):
|
||||
|
||||
@@ -7,9 +7,9 @@ Targets the remaining uncovered branches from the 97.03% baseline:
|
||||
- 214 : test_google_drive_token — generic connection error (not token-related)
|
||||
- 302->306: get_google_drive_token_info — credentials already valid (no refresh)
|
||||
- 307->318: get_google_drive_token_info — credentials have no expiry
|
||||
- 395->397: save_dropbox_settings — refresh_token falsy inside use_oauth block
|
||||
- 449->451: save_dropbox_settings — refresh_token falsy in in-memory update
|
||||
- 468->470: save_dropbox_settings — folder_id falsy in db-persist block
|
||||
- 395->397: save_google_drive_settings — refresh_token falsy inside use_oauth block
|
||||
- 449->451: save_google_drive_settings — refresh_token falsy in in-memory update
|
||||
- 468->470: save_google_drive_settings — folder_id falsy in db-persist block
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
@@ -152,10 +152,10 @@ class TestGetTokenInfoCredentialsBranches:
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestSaveGoogleDriveSettingsFalsyFields:
|
||||
"""Cover branches 395->397, 449->451, 468->470 in save_dropbox_settings.
|
||||
"""Cover branches 395->397, 449->451, 468->470 in save_google_drive_settings.
|
||||
|
||||
Note: the Google Drive save endpoint is named save_dropbox_settings in the
|
||||
source (app/api/google_drive.py) due to an existing naming inconsistency.
|
||||
Note: the Google Drive save endpoint is named save_google_drive_settings in the
|
||||
source (app/api/google_drive.py).
|
||||
"""
|
||||
|
||||
@patch("app.api.google_drive.settings")
|
||||
@@ -167,7 +167,7 @@ class TestSaveGoogleDriveSettingsFalsyFields:
|
||||
|
||||
from starlette.requests import Request as StarletteRequest
|
||||
|
||||
from app.api.google_drive import save_dropbox_settings
|
||||
from app.api.google_drive import save_google_drive_settings
|
||||
|
||||
mock_request = MagicMock(spec=StarletteRequest)
|
||||
mock_request.session = {}
|
||||
@@ -175,7 +175,7 @@ class TestSaveGoogleDriveSettingsFalsyFields:
|
||||
|
||||
with patch("app.api.google_drive.save_setting_to_db"):
|
||||
with patch("app.api.google_drive.notify_settings_updated"):
|
||||
result = await save_dropbox_settings(
|
||||
result = await save_google_drive_settings(
|
||||
request=mock_request,
|
||||
refresh_token="", # falsy → branches 395->397 and 449->451
|
||||
client_id="cid",
|
||||
|
||||
@@ -75,8 +75,8 @@ class TestTestTokenRotation:
|
||||
patch.object(settings, "onedrive_refresh_token", "old_token"),
|
||||
patch.object(settings, "onedrive_client_id", "cid"),
|
||||
patch.object(settings, "onedrive_client_secret", "sec"),
|
||||
patch("app.api.onedrive.os.path.join", return_value=str(env_file)),
|
||||
patch("app.api.onedrive.os.path.exists", return_value=True),
|
||||
patch("app.utils.env_utils.os.path.join", return_value=str(env_file)),
|
||||
patch("app.utils.env_utils.os.path.exists", return_value=True),
|
||||
patch("app.database.SessionLocal") as mock_session_local,
|
||||
patch("app.api.onedrive.save_setting_to_db"),
|
||||
patch("app.api.onedrive.notify_settings_updated"),
|
||||
@@ -117,7 +117,7 @@ class TestTestTokenRotation:
|
||||
patch.object(settings, "onedrive_refresh_token", "old_token"),
|
||||
patch.object(settings, "onedrive_client_id", "cid"),
|
||||
patch.object(settings, "onedrive_client_secret", "sec"),
|
||||
patch("app.api.onedrive.os.path.exists", return_value=False),
|
||||
patch("app.utils.env_utils.os.path.exists", return_value=False),
|
||||
patch("app.database.SessionLocal") as mock_session_local,
|
||||
patch("app.api.onedrive.save_setting_to_db"),
|
||||
patch("app.api.onedrive.notify_settings_updated"),
|
||||
@@ -157,7 +157,7 @@ class TestTestTokenRotation:
|
||||
patch.object(settings, "onedrive_refresh_token", "old_token"),
|
||||
patch.object(settings, "onedrive_client_id", "cid"),
|
||||
patch.object(settings, "onedrive_client_secret", "sec"),
|
||||
patch("app.api.onedrive.os.path.exists", return_value=True),
|
||||
patch("app.utils.env_utils.os.path.exists", return_value=True),
|
||||
patch("builtins.open", side_effect=PermissionError("Permission denied")),
|
||||
patch("app.database.SessionLocal") as mock_session_local,
|
||||
patch("app.api.onedrive.save_setting_to_db"),
|
||||
@@ -198,7 +198,7 @@ class TestTestTokenRotation:
|
||||
patch.object(settings, "onedrive_refresh_token", "old_token"),
|
||||
patch.object(settings, "onedrive_client_id", "cid"),
|
||||
patch.object(settings, "onedrive_client_secret", "sec"),
|
||||
patch("app.api.onedrive.os.path.exists", return_value=False),
|
||||
patch("app.utils.env_utils.os.path.exists", return_value=False),
|
||||
patch("app.database.SessionLocal", side_effect=Exception("DB error")),
|
||||
):
|
||||
response = client.get("/api/onedrive/test-token")
|
||||
@@ -277,8 +277,8 @@ class TestTokenRotationEnvAppendLine:
|
||||
patch.object(settings, "onedrive_refresh_token", "old_token"),
|
||||
patch.object(settings, "onedrive_client_id", "cid"),
|
||||
patch.object(settings, "onedrive_client_secret", "sec"),
|
||||
patch("app.api.onedrive.os.path.join", return_value=str(env_file)),
|
||||
patch("app.api.onedrive.os.path.exists", return_value=True),
|
||||
patch("app.utils.env_utils.os.path.join", return_value=str(env_file)),
|
||||
patch("app.utils.env_utils.os.path.exists", return_value=True),
|
||||
patch("app.database.SessionLocal") as mock_sl,
|
||||
patch("app.api.onedrive.save_setting_to_db"),
|
||||
patch("app.api.onedrive.notify_settings_updated"),
|
||||
|
||||
@@ -0,0 +1,191 @@
|
||||
"""Tests for the saved searches API (app/api/saved_searches.py)."""
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from app.database import Base, get_db
|
||||
from app.models import SavedSearch
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test data constants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_OWNER = "test_user@example.com"
|
||||
_OTHER_OWNER = "other_user@example.com"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shared fixture helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def int_engine():
|
||||
"""In-memory SQLite engine for integration tests."""
|
||||
engine = create_engine(
|
||||
"sqlite:///:memory:",
|
||||
connect_args={"check_same_thread": False},
|
||||
poolclass=StaticPool,
|
||||
)
|
||||
Base.metadata.create_all(bind=engine)
|
||||
yield engine
|
||||
Base.metadata.drop_all(bind=engine)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def int_session(int_engine):
|
||||
"""DB session scoped to one test."""
|
||||
Session = sessionmaker(bind=int_engine)
|
||||
session = Session()
|
||||
yield session
|
||||
session.close()
|
||||
|
||||
|
||||
def _make_client(int_engine, owner_id: str = _OWNER):
|
||||
"""Return a TestClient with *owner_id* injected as the authenticated user."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.main import app
|
||||
|
||||
def override_db():
|
||||
Session = sessionmaker(bind=int_engine)
|
||||
session = Session()
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
app.dependency_overrides[get_db] = override_db
|
||||
with patch("app.api.saved_searches._get_user_id", return_value=owner_id):
|
||||
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
|
||||
yield client
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def int_client(int_engine):
|
||||
"""TestClient authenticated as _OWNER."""
|
||||
yield from _make_client(int_engine, _OWNER)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# CRUD tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
class TestSavedSearchesAPI:
|
||||
"""Tests for Saved Searches endpoints."""
|
||||
|
||||
def test_list_saved_searches_empty(self, int_client):
|
||||
"""No saved searches returns empty list."""
|
||||
resp = int_client.get("/api/saved-searches")
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == []
|
||||
|
||||
def test_create_saved_search(self, int_client):
|
||||
"""Create a saved search and verify the response."""
|
||||
payload = {"name": "My Invoices", "filters": {"tags": "invoice", "document_type": "Invoice"}}
|
||||
resp = int_client.post("/api/saved-searches", json=payload)
|
||||
assert resp.status_code == 201
|
||||
data = resp.json()
|
||||
assert data["name"] == "My Invoices"
|
||||
assert data["filters"] == {"tags": "invoice", "document_type": "Invoice"}
|
||||
assert "id" in data
|
||||
|
||||
def test_create_saved_search_invalid_filters(self, int_client):
|
||||
"""Creating with invalid filters returns 422."""
|
||||
# Missing filters parameter (or empty after sanitization)
|
||||
payload = {"name": "My Invoices", "filters": {}}
|
||||
resp = int_client.post("/api/saved-searches", json=payload)
|
||||
assert resp.status_code == 422
|
||||
|
||||
# Invalid filters format
|
||||
payload2 = {"name": "My Invoices", "filters": "not_a_dict"}
|
||||
resp2 = int_client.post("/api/saved-searches", json=payload2)
|
||||
assert resp2.status_code == 422
|
||||
|
||||
def test_create_saved_search_duplicate(self, int_client):
|
||||
"""Creating a duplicate named search returns 409."""
|
||||
payload = {"name": "Duplicate", "filters": {"q": "test"}}
|
||||
int_client.post("/api/saved-searches", json=payload)
|
||||
resp = int_client.post("/api/saved-searches", json=payload)
|
||||
assert resp.status_code == 409
|
||||
|
||||
def test_create_saved_search_limit(self, int_client, int_session):
|
||||
"""Exceeding MAX_SAVED_SEARCHES_PER_USER returns 409."""
|
||||
# Create 50 searches using the API to ensure they are visible
|
||||
for i in range(50):
|
||||
resp = int_client.post("/api/saved-searches", json={"name": f"Search LIMIT {i}", "filters": {"q": "test"}})
|
||||
assert resp.status_code == 201
|
||||
|
||||
payload = {"name": "One too many", "filters": {"q": "test"}}
|
||||
resp = int_client.post("/api/saved-searches", json=payload)
|
||||
assert resp.status_code == 409
|
||||
|
||||
def test_update_saved_search(self, int_client):
|
||||
"""Update an existing saved search."""
|
||||
payload = {"name": "Original Name", "filters": {"q": "test"}}
|
||||
created = int_client.post("/api/saved-searches", json=payload).json()
|
||||
search_id = created["id"]
|
||||
|
||||
update_payload = {"name": "Updated Name", "filters": {"tags": "new"}}
|
||||
resp = int_client.put(f"/api/saved-searches/{search_id}", json=update_payload)
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["name"] == "Updated Name"
|
||||
assert data["filters"] == {"tags": "new"}
|
||||
|
||||
def test_update_saved_search_not_found(self, int_client):
|
||||
"""Updating a non-existent search returns 404."""
|
||||
update_payload = {"name": "Updated Name"}
|
||||
resp = int_client.put("/api/saved-searches/999", json=update_payload)
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_update_saved_search_duplicate_name(self, int_client):
|
||||
"""Updating name to an existing search name returns 409."""
|
||||
payload1 = {"name": "Search 1", "filters": {"q": "a"}}
|
||||
payload2 = {"name": "Search 2", "filters": {"q": "b"}}
|
||||
int_client.post("/api/saved-searches", json=payload1)
|
||||
created2 = int_client.post("/api/saved-searches", json=payload2).json()
|
||||
search2_id = created2["id"]
|
||||
|
||||
update_payload = {"name": "Search 1"}
|
||||
resp = int_client.put(f"/api/saved-searches/{search2_id}", json=update_payload)
|
||||
assert resp.status_code == 409
|
||||
|
||||
def test_delete_saved_search(self, int_client, int_session):
|
||||
"""Delete an existing search."""
|
||||
payload = {"name": "To be deleted", "filters": {"q": "test"}}
|
||||
created = int_client.post("/api/saved-searches", json=payload).json()
|
||||
search_id = created["id"]
|
||||
|
||||
resp = int_client.delete(f"/api/saved-searches/{search_id}")
|
||||
assert resp.status_code == 204
|
||||
|
||||
assert int_session.query(SavedSearch).filter(SavedSearch.id == search_id).first() is None
|
||||
|
||||
def test_delete_saved_search_not_found(self, int_client):
|
||||
"""Deleting a non-existent search returns 404."""
|
||||
resp = int_client.delete("/api/saved-searches/999")
|
||||
assert resp.status_code == 404
|
||||
|
||||
def test_other_users_searches_isolated(self, int_engine, int_session):
|
||||
"""Users only see and can only modify their own saved searches."""
|
||||
int_session.add(SavedSearch(user_id=_OTHER_OWNER, name="Other Search", filters='{"q": "test"}'))
|
||||
int_session.commit()
|
||||
|
||||
client = next(_make_client(int_engine, _OWNER))
|
||||
resp = client.get("/api/saved-searches")
|
||||
assert resp.status_code == 200
|
||||
assert len(resp.json()) == 0
|
||||
|
||||
other_search = int_session.query(SavedSearch).first()
|
||||
resp = client.put(f"/api/saved-searches/{other_search.id}", json={"name": "Hacked"})
|
||||
assert resp.status_code == 404
|
||||
|
||||
resp = client.delete(f"/api/saved-searches/{other_search.id}")
|
||||
assert resp.status_code == 404
|
||||
+129
-6
@@ -101,6 +101,43 @@ def _cleanup(app):
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests – Auth Helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetOwnerId:
|
||||
"""Tests for the _get_owner_id dependency helper."""
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_get_owner_id_unauthenticated(self):
|
||||
"""_get_owner_id should raise a 401 if the user is not authenticated."""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from app.api.api_tokens import _get_owner_id
|
||||
|
||||
mock_request = MagicMock()
|
||||
with patch("app.api.api_tokens.get_current_owner_id", return_value=None):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_get_owner_id(mock_request)
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.detail == "Not authenticated"
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_get_owner_id_authenticated(self):
|
||||
"""_get_owner_id should return owner_id if user is authenticated."""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from app.api.api_tokens import _get_owner_id
|
||||
|
||||
mock_request = MagicMock()
|
||||
with patch("app.api.api_tokens.get_current_owner_id", return_value="owner123"):
|
||||
owner_id = _get_owner_id(mock_request)
|
||||
assert owner_id == "owner123"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests – Token CRUD
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -157,6 +194,46 @@ class TestTokenCreate:
|
||||
finally:
|
||||
_cleanup(app)
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_create_token_database_error(self, tok_engine, tok_session):
|
||||
"""Creating a token should rollback and raise 500 if database commit fails."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from sqlalchemy.orm import Session as SASession
|
||||
|
||||
from app.main import app
|
||||
|
||||
client = _make_client(tok_engine)
|
||||
try:
|
||||
# Wrap commit: flush first so changes are staged in the transaction,
|
||||
# then raise to simulate a commit failure after data has been written.
|
||||
def _fail_after_flush(self):
|
||||
self.flush() # stage changes inside the open transaction
|
||||
raise Exception("DB Failure")
|
||||
|
||||
# Spy on rollback so we can assert it is called.
|
||||
rollback_called = False
|
||||
real_rollback = SASession.rollback
|
||||
|
||||
def _spy_rollback(self):
|
||||
nonlocal rollback_called
|
||||
rollback_called = True
|
||||
real_rollback(self)
|
||||
|
||||
with patch.object(SASession, "commit", _fail_after_flush):
|
||||
with patch.object(SASession, "rollback", _spy_rollback):
|
||||
resp = client.post("/api/api-tokens/", json={"name": "DB Error Create Test"})
|
||||
assert resp.status_code == 500
|
||||
|
||||
# rollback() must have been called to undo the flushed changes.
|
||||
assert rollback_called, "db.rollback() was not called after commit failure in create_token"
|
||||
|
||||
# After rollback the token must not exist in the database.
|
||||
db_token = tok_session.query(ApiToken).filter(ApiToken.name == "DB Error Create Test").first()
|
||||
assert db_token is None
|
||||
finally:
|
||||
_cleanup(app)
|
||||
|
||||
|
||||
class TestTokenList:
|
||||
"""Tests for GET /api/api-tokens/."""
|
||||
@@ -326,12 +403,10 @@ class TestTokenRevoke:
|
||||
rollback_called = True
|
||||
real_rollback(self)
|
||||
|
||||
with (
|
||||
patch.object(SASession, "commit", _fail_after_flush),
|
||||
patch.object(SASession, "rollback", _spy_rollback),
|
||||
):
|
||||
resp = client.delete(f"/api/api-tokens/{token_id}")
|
||||
assert resp.status_code == 500
|
||||
with patch.object(SASession, "commit", _fail_after_flush):
|
||||
with patch.object(SASession, "rollback", _spy_rollback):
|
||||
resp = client.delete(f"/api/api-tokens/{token_id}")
|
||||
assert resp.status_code == 500
|
||||
|
||||
# rollback() must have been called to undo the flushed changes.
|
||||
assert rollback_called, "db.rollback() was not called after commit failure"
|
||||
@@ -534,6 +609,44 @@ class TestTokenUtils:
|
||||
tokens = {generate_api_token() for _ in range(100)}
|
||||
assert len(tokens) == 100
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_generate_api_token_length(self):
|
||||
"""Generated tokens should have the exact expected length based on TOKEN_BYTES."""
|
||||
import math
|
||||
|
||||
from app.api.api_tokens import TOKEN_BYTES, TOKEN_PREFIX, generate_api_token
|
||||
|
||||
# base64url encoding of N bytes without padding: ceil(N * 4 / 3) characters
|
||||
expected_b64_len = math.ceil(TOKEN_BYTES * 4 / 3)
|
||||
expected_total_len = len(TOKEN_PREFIX) + expected_b64_len
|
||||
|
||||
token = generate_api_token()
|
||||
assert len(token) == expected_total_len
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_generate_api_token_charset(self):
|
||||
"""Generated tokens should only contain URL-safe base64 characters and the prefix."""
|
||||
import re
|
||||
|
||||
from app.api.api_tokens import TOKEN_PREFIX, generate_api_token
|
||||
|
||||
token = generate_api_token()
|
||||
# Check it starts with prefix and the rest is base64url chars ([A-Za-z0-9_-])
|
||||
pattern = f"^{re.escape(TOKEN_PREFIX)}[A-Za-z0-9_\\-]+$"
|
||||
assert re.match(pattern, token) is not None
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_generate_api_token_uses_secrets(self):
|
||||
"""Generated tokens should use secrets.token_urlsafe with the correct number of bytes."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.api.api_tokens import TOKEN_BYTES, TOKEN_PREFIX, generate_api_token
|
||||
|
||||
with patch("app.api.api_tokens.secrets.token_urlsafe", return_value="mocked_token") as mock_secrets:
|
||||
token = generate_api_token()
|
||||
mock_secrets.assert_called_once_with(TOKEN_BYTES)
|
||||
assert token == f"{TOKEN_PREFIX}mocked_token"
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_hash_token_deterministic(self):
|
||||
"""Hashing the same token should always produce the same result."""
|
||||
@@ -554,3 +667,13 @@ class TestTokenUtils:
|
||||
# All characters should be valid lowercase hex digits.
|
||||
int(h, 16)
|
||||
assert h == h.lower()
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_hash_token_known_value(self):
|
||||
"""hash_token should return the exact expected PBKDF2 digest for a known input."""
|
||||
from app.api.api_tokens import hash_token
|
||||
|
||||
# PBKDF2-HMAC-SHA256 with 100,000 iterations and salt b"api-token-v1"
|
||||
token = "de_test_token_value"
|
||||
expected_hash = "9b89d9adf2f390c75bf2fd0ff2bb5622ef5a9dce438354cce6e39f2f5401129e"
|
||||
assert hash_token(token) == expected_hash
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import logging
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -38,6 +39,40 @@ class TestGetCurrentUser:
|
||||
result = get_current_user(mock_request)
|
||||
assert result is None
|
||||
|
||||
def test_logs_debug_when_session_user_found(self, caplog):
|
||||
"""Test that get_current_user emits a DEBUG log when session user is found."""
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.session = {"user": {"id": "u1", "preferred_username": "alice"}}
|
||||
mock_request.state = MagicMock(spec=[]) # no api_token_user attribute
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="app.auth"):
|
||||
get_current_user(mock_request)
|
||||
|
||||
assert any("[AUTH] get_current_user: resolved from session" in m for m in caplog.messages)
|
||||
|
||||
def test_logs_debug_when_no_user(self, caplog):
|
||||
"""Test that get_current_user emits a DEBUG log when no user is present."""
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.session = {}
|
||||
mock_request.state = MagicMock(spec=[])
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="app.auth"):
|
||||
get_current_user(mock_request)
|
||||
|
||||
assert any("[AUTH] get_current_user: no user in session or API token" in m for m in caplog.messages)
|
||||
|
||||
def test_logs_debug_when_api_token_user(self, caplog):
|
||||
"""Test that get_current_user emits a DEBUG log when resolved from API token."""
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.state.api_token_user = {"id": "tok_user"}
|
||||
mock_request.session = {}
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="app.auth"):
|
||||
result = get_current_user(mock_request)
|
||||
|
||||
assert result == {"id": "tok_user"}
|
||||
assert any("[AUTH] get_current_user: resolved from API token" in m for m in caplog.messages)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestGetGravatarUrl:
|
||||
|
||||
@@ -704,6 +704,18 @@ def _set_minimal_provider_settings(mock_settings):
|
||||
class TestSettingsSyncAdditional:
|
||||
"""Additional tests for settings_sync covering reload failure branch."""
|
||||
|
||||
def test_notify_settings_updated_redis_failure_logs_warning(self):
|
||||
"""Test that a Redis failure is logged, not raised (lines 58-59)."""
|
||||
from app.utils.settings_sync import notify_settings_updated
|
||||
|
||||
with patch("app.utils.settings_sync.redis") as mock_redis_module:
|
||||
mock_redis_module.from_url.side_effect = Exception("Redis connection failed")
|
||||
with patch("app.utils.settings_sync.logger") as mock_logger:
|
||||
notify_settings_updated()
|
||||
mock_logger.warning.assert_any_call(
|
||||
"Could not publish settings update to Redis: Redis connection failed"
|
||||
)
|
||||
|
||||
def test_reload_failure_is_logged_not_raised(self):
|
||||
"""Test that a reload failure is logged, not raised (lines 71-72)."""
|
||||
from app.utils.settings_sync import notify_settings_updated
|
||||
@@ -711,8 +723,66 @@ class TestSettingsSyncAdditional:
|
||||
with patch("app.utils.settings_sync.redis") as mock_redis_module:
|
||||
mock_redis_module.from_url.return_value = MagicMock() # Redis OK
|
||||
with patch("app.utils.config_loader.reload_settings_from_db", side_effect=Exception("reload failed")):
|
||||
# Should not raise despite reload failure
|
||||
notify_settings_updated()
|
||||
with patch("app.utils.settings_sync.logger") as mock_logger:
|
||||
# Should not raise despite reload failure
|
||||
notify_settings_updated()
|
||||
mock_logger.warning.assert_any_call("Could not reload in-process settings: reload failed")
|
||||
|
||||
def test_notify_settings_updated_ocr_failure_logs_warning(self):
|
||||
"""Test that OCR check failure is logged, not raised (lines 82-83)."""
|
||||
from app.utils.settings_sync import notify_settings_updated
|
||||
|
||||
with patch("app.utils.settings_sync.redis") as mock_redis_module:
|
||||
mock_redis_module.from_url.return_value = MagicMock() # Redis OK
|
||||
with patch("app.utils.config_loader.reload_settings_from_db"):
|
||||
with patch(
|
||||
"app.utils.ocr_language_manager.ensure_ocr_languages_async", side_effect=Exception("OCR failed")
|
||||
):
|
||||
with patch("app.utils.settings_sync.logger") as mock_logger:
|
||||
notify_settings_updated()
|
||||
mock_logger.warning.assert_any_call("Could not schedule OCR language check: OCR failed")
|
||||
|
||||
def test_signal_handler_ocr_check_failure_logs_warning(self):
|
||||
"""Test that signal handler logs warning if OCR language check fails on worker (lines 115-116)."""
|
||||
from app.utils.settings_sync import register_settings_reload_signal
|
||||
|
||||
handler_fn = None
|
||||
|
||||
def capture_connect(fn=None, weak=None, **kwargs):
|
||||
nonlocal handler_fn
|
||||
if fn is not None:
|
||||
handler_fn = fn
|
||||
return fn
|
||||
|
||||
def decorator(func):
|
||||
nonlocal handler_fn
|
||||
handler_fn = func
|
||||
return func
|
||||
|
||||
return decorator
|
||||
|
||||
with patch("app.utils.settings_sync.task_prerun") as mock_signal:
|
||||
mock_signal.connect = capture_connect
|
||||
register_settings_reload_signal()
|
||||
|
||||
assert handler_fn is not None
|
||||
|
||||
mock_redis = MagicMock()
|
||||
mock_redis.get.return_value = b"1234567890.0"
|
||||
|
||||
with patch("app.utils.settings_sync.redis") as mock_redis_mod:
|
||||
mock_redis_mod.from_url.return_value = mock_redis
|
||||
with patch("app.utils.config_loader.reload_settings_from_db"):
|
||||
with patch("app.utils.settings_sync._last_seen_version", ""):
|
||||
with patch(
|
||||
"app.utils.ocr_language_manager.ensure_ocr_languages_async",
|
||||
side_effect=Exception("Worker OCR fail"),
|
||||
):
|
||||
with patch("app.utils.settings_sync.logger") as mock_logger:
|
||||
handler_fn(sender=None)
|
||||
mock_logger.warning.assert_any_call(
|
||||
"Could not schedule OCR language check on worker: Worker OCR fail"
|
||||
)
|
||||
|
||||
def test_signal_handler_reloads_on_version_change(self):
|
||||
"""Test the task_prerun signal handler reloads settings when version changes (lines 95-98)."""
|
||||
@@ -820,6 +890,83 @@ class TestSettingsSyncAdditional:
|
||||
# Should not raise
|
||||
handler_fn(sender=None)
|
||||
|
||||
def test_signal_handler_no_version_returned(self):
|
||||
"""Test that handler does nothing if Redis returns None for version."""
|
||||
from app.utils.settings_sync import register_settings_reload_signal
|
||||
|
||||
handler_fn = None
|
||||
|
||||
def capture_connect(fn=None, weak=None, **kwargs):
|
||||
nonlocal handler_fn
|
||||
if fn is not None:
|
||||
handler_fn = fn
|
||||
return fn
|
||||
|
||||
def decorator(func):
|
||||
nonlocal handler_fn
|
||||
handler_fn = func
|
||||
return func
|
||||
|
||||
return decorator
|
||||
|
||||
with patch("app.utils.settings_sync.task_prerun") as mock_signal:
|
||||
mock_signal.connect = capture_connect
|
||||
register_settings_reload_signal()
|
||||
|
||||
assert handler_fn is not None
|
||||
|
||||
mock_redis = MagicMock()
|
||||
mock_redis.get.return_value = None # Return None for version
|
||||
|
||||
with patch("app.utils.settings_sync.redis") as mock_redis_mod:
|
||||
mock_redis_mod.from_url.return_value = mock_redis
|
||||
with patch("app.utils.config_loader.reload_settings_from_db") as mock_reload:
|
||||
handler_fn(sender=None)
|
||||
mock_reload.assert_not_called()
|
||||
|
||||
def test_signal_handler_ocr_language_manager_exception(self):
|
||||
"""Test that OCR language check exception inside handler is caught and logged."""
|
||||
from app.utils.settings_sync import register_settings_reload_signal
|
||||
|
||||
handler_fn = None
|
||||
|
||||
def capture_connect(fn=None, weak=None, **kwargs):
|
||||
nonlocal handler_fn
|
||||
if fn is not None:
|
||||
handler_fn = fn
|
||||
return fn
|
||||
|
||||
def decorator(func):
|
||||
nonlocal handler_fn
|
||||
handler_fn = func
|
||||
return func
|
||||
|
||||
return decorator
|
||||
|
||||
with patch("app.utils.settings_sync.task_prerun") as mock_signal:
|
||||
mock_signal.connect = capture_connect
|
||||
register_settings_reload_signal()
|
||||
|
||||
assert handler_fn is not None
|
||||
|
||||
mock_redis = MagicMock()
|
||||
mock_redis.get.return_value = b"9999999.0" # New version
|
||||
|
||||
with patch("app.utils.settings_sync.redis") as mock_redis_mod:
|
||||
mock_redis_mod.from_url.return_value = mock_redis
|
||||
with patch("app.utils.config_loader.reload_settings_from_db") as mock_reload:
|
||||
with patch("app.utils.settings_sync._last_seen_version", "111.0"):
|
||||
with patch("app.utils.settings_sync.logger") as mock_logger:
|
||||
with patch(
|
||||
"app.utils.ocr_language_manager.ensure_ocr_languages_async",
|
||||
side_effect=Exception("OCR failed"),
|
||||
):
|
||||
handler_fn(sender=None)
|
||||
mock_reload.assert_called_once()
|
||||
mock_logger.warning.assert_called_with(
|
||||
"Could not schedule OCR language check on worker: OCR failed"
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# app/api/logs.py – additional branches
|
||||
|
||||
@@ -345,7 +345,7 @@ class TestUploadErrorHandling:
|
||||
|
||||
def test_upload_disk_write_failure(self, client: TestClient):
|
||||
"""Test handling of disk write failures."""
|
||||
with patch("builtins.open", side_effect=IOError("Disk full")):
|
||||
with patch("aiofiles.open", side_effect=IOError("Disk full")):
|
||||
pdf_content = b"%PDF-1.4\n%EOF"
|
||||
|
||||
response = client.post(
|
||||
|
||||
@@ -0,0 +1,285 @@
|
||||
"""Tests for application logging configuration.
|
||||
|
||||
Validates that the LOG_LEVEL and DEBUG settings correctly control the
|
||||
Python root-logger level and that the standard precedence rules are respected:
|
||||
1. Explicit LOG_LEVEL always wins.
|
||||
2. DEBUG=True without LOG_LEVEL → effective DEBUG.
|
||||
3. Neither set → default INFO.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from app.config import Settings
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLogLevelSetting:
|
||||
"""Tests for the log_level config field."""
|
||||
|
||||
_BASE_KWARGS = {
|
||||
"database_url": "sqlite:///test.db",
|
||||
"redis_url": "redis://localhost:6379",
|
||||
"openai_api_key": "test",
|
||||
"azure_ai_key": "test",
|
||||
"azure_region": "test",
|
||||
"azure_endpoint": "https://test.example.com",
|
||||
"gotenberg_url": "http://localhost:3000",
|
||||
"workdir": "/tmp",
|
||||
"auth_enabled": False,
|
||||
"session_secret": None,
|
||||
}
|
||||
|
||||
def test_log_level_default_is_info(self):
|
||||
"""Test that log_level defaults to INFO."""
|
||||
config = Settings(**self._BASE_KWARGS)
|
||||
assert config.log_level.upper() == "INFO"
|
||||
|
||||
def test_log_level_accepts_debug(self):
|
||||
"""Test that log_level accepts DEBUG."""
|
||||
config = Settings(**self._BASE_KWARGS, log_level="DEBUG")
|
||||
assert config.log_level.upper() == "DEBUG"
|
||||
|
||||
def test_log_level_accepts_warning(self):
|
||||
"""Test that log_level accepts WARNING."""
|
||||
config = Settings(**self._BASE_KWARGS, log_level="WARNING")
|
||||
assert config.log_level.upper() == "WARNING"
|
||||
|
||||
def test_log_level_accepts_error(self):
|
||||
"""Test that log_level accepts ERROR."""
|
||||
config = Settings(**self._BASE_KWARGS, log_level="ERROR")
|
||||
assert config.log_level.upper() == "ERROR"
|
||||
|
||||
def test_log_level_case_insensitive(self):
|
||||
"""Test that log_level is case-insensitive in usage."""
|
||||
config = Settings(**self._BASE_KWARGS, log_level="debug")
|
||||
assert config.log_level.upper() == "DEBUG"
|
||||
|
||||
def test_debug_flag_defaults_to_false(self):
|
||||
"""Test that debug defaults to False."""
|
||||
config = Settings(**self._BASE_KWARGS)
|
||||
assert config.debug is False
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestEffectiveLogLevel:
|
||||
"""Tests for the effective log-level resolution logic in main.py."""
|
||||
|
||||
def test_debug_true_without_log_level_gives_debug(self):
|
||||
"""When DEBUG=True and LOG_LEVEL is not set, effective level is DEBUG."""
|
||||
with patch.dict(os.environ, {"DEBUG": "true"}, clear=False):
|
||||
# Remove LOG_LEVEL from env if present
|
||||
env = os.environ.copy()
|
||||
env.pop("LOG_LEVEL", None)
|
||||
with patch.dict(os.environ, env, clear=True):
|
||||
s = Settings(
|
||||
database_url="sqlite:///test.db",
|
||||
redis_url="redis://localhost:6379",
|
||||
openai_api_key="test",
|
||||
azure_ai_key="test",
|
||||
azure_region="test",
|
||||
azure_endpoint="https://test.example.com",
|
||||
gotenberg_url="http://localhost:3000",
|
||||
workdir="/tmp",
|
||||
auth_enabled=False,
|
||||
session_secret=None,
|
||||
debug=True,
|
||||
)
|
||||
explicit = os.environ.get("LOG_LEVEL")
|
||||
if s.debug and explicit is None:
|
||||
effective = "DEBUG"
|
||||
else:
|
||||
effective = s.log_level.upper()
|
||||
assert effective == "DEBUG"
|
||||
|
||||
def test_explicit_log_level_overrides_debug(self):
|
||||
"""When LOG_LEVEL is explicitly set, it takes precedence over DEBUG=True."""
|
||||
with patch.dict(os.environ, {"LOG_LEVEL": "WARNING", "DEBUG": "true"}, clear=False):
|
||||
s = Settings(
|
||||
database_url="sqlite:///test.db",
|
||||
redis_url="redis://localhost:6379",
|
||||
openai_api_key="test",
|
||||
azure_ai_key="test",
|
||||
azure_region="test",
|
||||
azure_endpoint="https://test.example.com",
|
||||
gotenberg_url="http://localhost:3000",
|
||||
workdir="/tmp",
|
||||
auth_enabled=False,
|
||||
session_secret=None,
|
||||
debug=True,
|
||||
log_level="WARNING",
|
||||
)
|
||||
explicit = os.environ.get("LOG_LEVEL")
|
||||
if s.debug and explicit is None:
|
||||
effective = "DEBUG"
|
||||
else:
|
||||
effective = s.log_level.upper()
|
||||
assert effective == "WARNING"
|
||||
|
||||
def test_default_no_flags_gives_info(self):
|
||||
"""When neither DEBUG nor LOG_LEVEL is set, effective level is INFO."""
|
||||
env = os.environ.copy()
|
||||
env.pop("LOG_LEVEL", None)
|
||||
env.pop("DEBUG", None)
|
||||
with patch.dict(os.environ, env, clear=True):
|
||||
s = Settings(
|
||||
database_url="sqlite:///test.db",
|
||||
redis_url="redis://localhost:6379",
|
||||
openai_api_key="test",
|
||||
azure_ai_key="test",
|
||||
azure_region="test",
|
||||
azure_endpoint="https://test.example.com",
|
||||
gotenberg_url="http://localhost:3000",
|
||||
workdir="/tmp",
|
||||
auth_enabled=False,
|
||||
session_secret=None,
|
||||
)
|
||||
explicit = os.environ.get("LOG_LEVEL")
|
||||
if s.debug and explicit is None:
|
||||
effective = "DEBUG"
|
||||
else:
|
||||
effective = s.log_level.upper()
|
||||
assert effective == "INFO"
|
||||
|
||||
def test_effective_level_maps_to_logging_constant(self):
|
||||
"""The effective level string maps to a valid logging constant."""
|
||||
for level_name in ("DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"):
|
||||
assert getattr(logging, level_name) is not None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLoggingConfiguredAtStartup:
|
||||
"""Tests that the main module configures the root logger on import."""
|
||||
|
||||
def test_root_logger_has_handler(self):
|
||||
"""Root logger should have at least one handler after app import."""
|
||||
root = logging.getLogger()
|
||||
assert len(root.handlers) > 0, "Root logger has no handlers after app startup"
|
||||
|
||||
def test_root_logger_level_is_not_warning_default(self):
|
||||
"""Root logger should not be at the unconfigured WARNING default.
|
||||
|
||||
Our basicConfig(force=True) should have set it to at least INFO.
|
||||
"""
|
||||
root = logging.getLogger()
|
||||
# The test env doesn't set DEBUG=True, so the level should be INFO (20)
|
||||
assert root.level <= logging.INFO
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestJsonFormatter:
|
||||
"""Tests for the _JsonFormatter used when LOG_FORMAT=json."""
|
||||
|
||||
def _make_formatter(self):
|
||||
"""Lazily import the JSON formatter from main module."""
|
||||
from app.main import _JsonFormatter
|
||||
|
||||
return _JsonFormatter()
|
||||
|
||||
def test_output_is_valid_json(self):
|
||||
"""JSON formatter output should be parseable JSON."""
|
||||
import json
|
||||
|
||||
fmt = self._make_formatter()
|
||||
record = logging.LogRecord(
|
||||
name="test.logger",
|
||||
level=logging.INFO,
|
||||
pathname="test.py",
|
||||
lineno=42,
|
||||
msg="Hello %s",
|
||||
args=("world",),
|
||||
exc_info=None,
|
||||
)
|
||||
result = fmt.format(record)
|
||||
parsed = json.loads(result)
|
||||
assert parsed["level"] == "INFO"
|
||||
assert parsed["logger"] == "test.logger"
|
||||
assert parsed["message"] == "Hello world"
|
||||
assert parsed["lineno"] == 42
|
||||
|
||||
def test_includes_timestamp_iso8601(self):
|
||||
"""JSON output should contain an ISO 8601 timestamp."""
|
||||
import json
|
||||
|
||||
fmt = self._make_formatter()
|
||||
record = logging.LogRecord(
|
||||
name="x",
|
||||
level=logging.DEBUG,
|
||||
pathname="x.py",
|
||||
lineno=1,
|
||||
msg="test",
|
||||
args=(),
|
||||
exc_info=None,
|
||||
)
|
||||
parsed = json.loads(fmt.format(record))
|
||||
assert "timestamp" in parsed
|
||||
# ISO 8601 timestamps contain "T" and "+00:00" (UTC)
|
||||
assert "T" in parsed["timestamp"]
|
||||
|
||||
def test_includes_exc_info_when_present(self):
|
||||
"""JSON output should include exc_info when an exception is logged."""
|
||||
import json
|
||||
|
||||
fmt = self._make_formatter()
|
||||
try:
|
||||
raise ValueError("boom") # noqa: TRY301
|
||||
except ValueError:
|
||||
import sys
|
||||
|
||||
record = logging.LogRecord(
|
||||
name="x",
|
||||
level=logging.ERROR,
|
||||
pathname="x.py",
|
||||
lineno=1,
|
||||
msg="error",
|
||||
args=(),
|
||||
exc_info=sys.exc_info(),
|
||||
)
|
||||
parsed = json.loads(fmt.format(record))
|
||||
assert "exc_info" in parsed
|
||||
assert "ValueError" in parsed["exc_info"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestLogFormatSetting:
|
||||
"""Tests for the log_format and log_syslog_* config fields."""
|
||||
|
||||
_BASE_KWARGS = {
|
||||
"database_url": "sqlite:///test.db",
|
||||
"redis_url": "redis://localhost:6379",
|
||||
"openai_api_key": "test",
|
||||
"azure_ai_key": "test",
|
||||
"azure_region": "test",
|
||||
"azure_endpoint": "https://test.example.com",
|
||||
"gotenberg_url": "http://localhost:3000",
|
||||
"workdir": "/tmp",
|
||||
"auth_enabled": False,
|
||||
"session_secret": None,
|
||||
}
|
||||
|
||||
def test_log_format_default_is_text(self):
|
||||
"""Test that log_format defaults to 'text'."""
|
||||
config = Settings(**self._BASE_KWARGS)
|
||||
assert config.log_format == "text"
|
||||
|
||||
def test_log_format_accepts_json(self):
|
||||
"""Test that log_format accepts 'json'."""
|
||||
config = Settings(**self._BASE_KWARGS, log_format="json")
|
||||
assert config.log_format == "json"
|
||||
|
||||
def test_log_syslog_defaults(self):
|
||||
"""Test syslog forwarding defaults."""
|
||||
config = Settings(**self._BASE_KWARGS)
|
||||
assert config.log_syslog_enabled is False
|
||||
assert config.log_syslog_host == "localhost"
|
||||
assert config.log_syslog_port == 514
|
||||
assert config.log_syslog_protocol == "udp"
|
||||
|
||||
def test_log_syslog_can_be_enabled(self):
|
||||
"""Test that syslog forwarding can be enabled."""
|
||||
config = Settings(**self._BASE_KWARGS, log_syslog_enabled=True, log_syslog_host="syslog.example.com")
|
||||
assert config.log_syslog_enabled is True
|
||||
assert config.log_syslog_host == "syslog.example.com"
|
||||
@@ -830,3 +830,58 @@ class TestUserNotificationService:
|
||||
|
||||
result = _send_email_notification({"smtp_host": "smtp.example.com"}, "Title", "Body")
|
||||
assert result is False
|
||||
|
||||
|
||||
class TestBenchmark:
|
||||
@pytest.mark.unit
|
||||
def test_update_preferences_benchmark(self, notif_engine, notif_session):
|
||||
import statistics
|
||||
import time
|
||||
|
||||
from app.main import app
|
||||
|
||||
target = UserNotificationTarget(
|
||||
owner_id=_OWNER,
|
||||
channel_type="webhook",
|
||||
name="My Webhook",
|
||||
config=json.dumps({"url": "https://x.com"}),
|
||||
)
|
||||
notif_session.add(target)
|
||||
notif_session.commit()
|
||||
notif_session.refresh(target)
|
||||
|
||||
client = _make_client(notif_engine, _OWNER)
|
||||
try:
|
||||
items_count = 100
|
||||
preferences = []
|
||||
for i in range(items_count):
|
||||
preferences.append(
|
||||
{
|
||||
"event_type": f"event.type.{i}",
|
||||
"channel_type": "webhook",
|
||||
"is_enabled": True,
|
||||
"target_id": target.id,
|
||||
}
|
||||
)
|
||||
|
||||
payload = {"preferences": preferences}
|
||||
|
||||
# Warm up
|
||||
client.put("/api/user-notifications/preferences", json=payload)
|
||||
|
||||
times = []
|
||||
for _ in range(5):
|
||||
# Alter the values a bit so it's a real update
|
||||
for p in payload["preferences"]:
|
||||
p["is_enabled"] = not p["is_enabled"]
|
||||
|
||||
start = time.time()
|
||||
resp = client.put("/api/user-notifications/preferences", json=payload)
|
||||
end = time.time()
|
||||
|
||||
assert resp.status_code == 200
|
||||
times.append(end - start)
|
||||
|
||||
print(f"\nAverage time: {statistics.mean(times):.4f}s")
|
||||
finally:
|
||||
_cleanup(app)
|
||||
|
||||
@@ -289,7 +289,8 @@ class TestNotifySettingsUpdated:
|
||||
call_args = mock_redis_instance.set.call_args[0]
|
||||
assert call_args[0] == SETTINGS_VERSION_KEY
|
||||
|
||||
def test_does_not_raise_on_redis_failure(self):
|
||||
@patch("app.utils.settings_sync.logger")
|
||||
def test_does_not_raise_on_redis_failure(self, mock_logger):
|
||||
"""notify_settings_updated must not propagate Redis errors."""
|
||||
from app.utils.settings_sync import notify_settings_updated
|
||||
|
||||
@@ -297,6 +298,35 @@ class TestNotifySettingsUpdated:
|
||||
mock_redis_module.from_url.side_effect = Exception("Redis down")
|
||||
# Should not raise
|
||||
notify_settings_updated()
|
||||
mock_logger.warning.assert_any_call("Could not publish settings update to Redis: Redis down")
|
||||
|
||||
@patch("app.utils.settings_sync.logger")
|
||||
def test_does_not_raise_on_reload_failure(self, mock_logger):
|
||||
"""notify_settings_updated must not propagate settings reload errors."""
|
||||
from app.utils.settings_sync import notify_settings_updated
|
||||
|
||||
with patch("app.utils.config_loader.reload_settings_from_db") as mock_reload:
|
||||
mock_reload.side_effect = Exception("Reload error")
|
||||
# We mock redis so that we skip over the redis block, and mock ensure_ocr_languages_async to prevent its side effects.
|
||||
with patch("app.utils.settings_sync.redis"):
|
||||
with patch("app.utils.ocr_language_manager.ensure_ocr_languages_async"):
|
||||
# Should not raise
|
||||
notify_settings_updated()
|
||||
mock_logger.warning.assert_any_call("Could not reload in-process settings: Reload error")
|
||||
|
||||
@patch("app.utils.settings_sync.logger")
|
||||
def test_does_not_raise_on_ocr_language_check_failure(self, mock_logger):
|
||||
"""notify_settings_updated must not propagate OCR language check errors."""
|
||||
from app.utils.settings_sync import notify_settings_updated
|
||||
|
||||
with patch("app.utils.ocr_language_manager.ensure_ocr_languages_async") as mock_ensure:
|
||||
mock_ensure.side_effect = Exception("OCR error")
|
||||
# We mock redis and reload_settings_from_db so we only test the OCR block failure.
|
||||
with patch("app.utils.settings_sync.redis"):
|
||||
with patch("app.utils.config_loader.reload_settings_from_db"):
|
||||
# Should not raise
|
||||
notify_settings_updated()
|
||||
mock_logger.warning.assert_any_call("Could not schedule OCR language check: OCR error")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
|
||||
@@ -0,0 +1,211 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import app.utils.settings_sync
|
||||
from app.utils.settings_sync import (
|
||||
SETTINGS_VERSION_KEY,
|
||||
notify_settings_updated,
|
||||
register_settings_reload_signal,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def reset_last_seen_version():
|
||||
"""Reset the global variable before and after tests."""
|
||||
app.utils.settings_sync._last_seen_version = ""
|
||||
yield
|
||||
app.utils.settings_sync._last_seen_version = ""
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.redis.from_url")
|
||||
@patch("app.utils.config_loader.reload_settings_from_db")
|
||||
@patch("app.utils.ocr_language_manager.ensure_ocr_languages_async")
|
||||
@patch("app.utils.settings_sync.time.time", return_value=12345.0)
|
||||
def test_notify_settings_updated_success(mock_time, mock_ensure_ocr, mock_reload, mock_redis):
|
||||
# Setup mock redis instance
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis.return_value = mock_redis_instance
|
||||
|
||||
notify_settings_updated()
|
||||
|
||||
# Verify redis calls
|
||||
mock_redis.assert_called_once()
|
||||
mock_redis_instance.set.assert_called_once_with(SETTINGS_VERSION_KEY, "12345.0")
|
||||
|
||||
# Verify other calls
|
||||
mock_reload.assert_called_once()
|
||||
mock_ensure_ocr.assert_called_once()
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.redis.from_url")
|
||||
@patch("app.utils.config_loader.reload_settings_from_db")
|
||||
@patch("app.utils.ocr_language_manager.ensure_ocr_languages_async")
|
||||
def test_notify_settings_updated_redis_failure(mock_ensure_ocr, mock_reload, mock_redis, caplog):
|
||||
# Setup mock redis to fail
|
||||
mock_redis.side_effect = Exception("Redis connection failed")
|
||||
|
||||
notify_settings_updated()
|
||||
|
||||
# Verification: should continue and call reload and ocr despite redis failure
|
||||
mock_reload.assert_called_once()
|
||||
mock_ensure_ocr.assert_called_once()
|
||||
assert "Could not publish settings update to Redis: Redis connection failed" in caplog.text
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.redis.from_url")
|
||||
@patch("app.utils.config_loader.reload_settings_from_db")
|
||||
@patch("app.utils.ocr_language_manager.ensure_ocr_languages_async")
|
||||
def test_notify_settings_updated_reload_failure(mock_ensure_ocr, mock_reload, mock_redis, caplog):
|
||||
# Setup reload to fail
|
||||
mock_reload.side_effect = Exception("Reload failed")
|
||||
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis.return_value = mock_redis_instance
|
||||
|
||||
notify_settings_updated()
|
||||
|
||||
# Verification: redis should be called, reload fails, ocr should still be called
|
||||
mock_redis_instance.set.assert_called_once()
|
||||
mock_ensure_ocr.assert_called_once()
|
||||
assert "Could not reload in-process settings: Reload failed" in caplog.text
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.redis.from_url")
|
||||
@patch("app.utils.config_loader.reload_settings_from_db")
|
||||
@patch("app.utils.ocr_language_manager.ensure_ocr_languages_async")
|
||||
def test_notify_settings_updated_ocr_failure(mock_ensure_ocr, mock_reload, mock_redis, caplog):
|
||||
# Setup ocr check to fail
|
||||
mock_ensure_ocr.side_effect = Exception("OCR check failed")
|
||||
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis.return_value = mock_redis_instance
|
||||
|
||||
notify_settings_updated()
|
||||
|
||||
# Verification: all should be called, ocr failure logged
|
||||
mock_redis_instance.set.assert_called_once()
|
||||
mock_reload.assert_called_once()
|
||||
assert "Could not schedule OCR language check: OCR check failed" in caplog.text
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.task_prerun.connect")
|
||||
def test_register_settings_reload_signal(mock_connect):
|
||||
register_settings_reload_signal()
|
||||
# It should register a signal with task_prerun
|
||||
mock_connect.assert_called_once_with(weak=False)
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.task_prerun.connect")
|
||||
@patch("app.utils.settings_sync.redis.from_url")
|
||||
@patch("app.utils.config_loader.reload_settings_from_db")
|
||||
@patch("app.utils.ocr_language_manager.ensure_ocr_languages_async")
|
||||
def test_reload_if_stale_new_version(mock_ensure_ocr, mock_reload, mock_redis, mock_connect, reset_last_seen_version):
|
||||
# Capture the registered callback
|
||||
mock_decorator = MagicMock()
|
||||
mock_connect.return_value = mock_decorator
|
||||
|
||||
register_settings_reload_signal()
|
||||
|
||||
mock_connect.assert_called_once_with(weak=False)
|
||||
# Get the callback function
|
||||
callback = mock_decorator.call_args[0][0]
|
||||
|
||||
# Setup redis to return a new version
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.get.return_value = b"new_version"
|
||||
mock_redis.return_value = mock_redis_instance
|
||||
|
||||
# Initial state check
|
||||
assert app.utils.settings_sync._last_seen_version == ""
|
||||
|
||||
# Call the callback
|
||||
callback(sender="test")
|
||||
|
||||
# Verification
|
||||
mock_redis_instance.get.assert_called_once_with(SETTINGS_VERSION_KEY)
|
||||
mock_reload.assert_called_once()
|
||||
mock_ensure_ocr.assert_called_once()
|
||||
assert app.utils.settings_sync._last_seen_version == "new_version"
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.task_prerun.connect")
|
||||
@patch("app.utils.settings_sync.redis.from_url")
|
||||
@patch("app.utils.config_loader.reload_settings_from_db")
|
||||
@patch("app.utils.ocr_language_manager.ensure_ocr_languages_async")
|
||||
def test_reload_if_stale_same_version(mock_ensure_ocr, mock_reload, mock_redis, mock_connect, reset_last_seen_version):
|
||||
# Set initial state
|
||||
app.utils.settings_sync._last_seen_version = "existing_version"
|
||||
|
||||
mock_decorator = MagicMock()
|
||||
mock_connect.return_value = mock_decorator
|
||||
register_settings_reload_signal()
|
||||
callback = mock_decorator.call_args[0][0]
|
||||
|
||||
# Setup redis to return the SAME version
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.get.return_value = b"existing_version"
|
||||
mock_redis.return_value = mock_redis_instance
|
||||
|
||||
# Call the callback
|
||||
callback(sender="test")
|
||||
|
||||
# Verification
|
||||
mock_redis_instance.get.assert_called_once_with(SETTINGS_VERSION_KEY)
|
||||
# Should NOT reload or check OCR
|
||||
mock_reload.assert_not_called()
|
||||
mock_ensure_ocr.assert_not_called()
|
||||
assert app.utils.settings_sync._last_seen_version == "existing_version"
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.task_prerun.connect")
|
||||
@patch("app.utils.settings_sync.redis.from_url")
|
||||
@patch("app.utils.config_loader.reload_settings_from_db")
|
||||
def test_reload_if_stale_redis_error(mock_reload, mock_redis, mock_connect, reset_last_seen_version, caplog):
|
||||
import logging
|
||||
|
||||
caplog.set_level(logging.DEBUG)
|
||||
mock_decorator = MagicMock()
|
||||
mock_connect.return_value = mock_decorator
|
||||
register_settings_reload_signal()
|
||||
callback = mock_decorator.call_args[0][0]
|
||||
|
||||
# Setup redis to fail
|
||||
mock_redis.side_effect = Exception("Redis error")
|
||||
|
||||
# Call the callback
|
||||
callback(sender="test")
|
||||
|
||||
# Verification
|
||||
mock_reload.assert_not_called()
|
||||
assert "Settings version check skipped: Redis error" in caplog.text
|
||||
|
||||
|
||||
@patch("app.utils.settings_sync.task_prerun.connect")
|
||||
@patch("app.utils.settings_sync.redis.from_url")
|
||||
@patch("app.utils.config_loader.reload_settings_from_db")
|
||||
@patch("app.utils.ocr_language_manager.ensure_ocr_languages_async")
|
||||
def test_reload_if_stale_ocr_error(
|
||||
mock_ensure_ocr, mock_reload, mock_redis, mock_connect, reset_last_seen_version, caplog
|
||||
):
|
||||
mock_decorator = MagicMock()
|
||||
mock_connect.return_value = mock_decorator
|
||||
register_settings_reload_signal()
|
||||
callback = mock_decorator.call_args[0][0]
|
||||
|
||||
# Setup redis to return a new version
|
||||
mock_redis_instance = MagicMock()
|
||||
mock_redis_instance.get.return_value = b"new_version"
|
||||
mock_redis.return_value = mock_redis_instance
|
||||
|
||||
# Setup OCR check to fail
|
||||
mock_ensure_ocr.side_effect = Exception("OCR error")
|
||||
|
||||
# Call the callback
|
||||
callback(sender="test")
|
||||
|
||||
# Verification
|
||||
mock_reload.assert_called_once()
|
||||
mock_ensure_ocr.assert_called_once()
|
||||
assert "Could not schedule OCR language check on worker: OCR error" in caplog.text
|
||||
assert app.utils.settings_sync._last_seen_version == "new_version"
|
||||
@@ -121,6 +121,34 @@ class TestExtractMetadataFromFile:
|
||||
|
||||
assert result == {}
|
||||
|
||||
def test_extract_metadata_from_pdf(self, tmp_path):
|
||||
"""Test extracting metadata from a PDF file using pypdf when JSON is missing."""
|
||||
import pypdf
|
||||
|
||||
file_path = tmp_path / "test.pdf"
|
||||
|
||||
# Create a test PDF with metadata
|
||||
writer = pypdf.PdfWriter()
|
||||
writer.add_blank_page(width=100, height=100)
|
||||
writer.add_metadata(
|
||||
{
|
||||
"/Title": "Test Title",
|
||||
"/Author": "Test Author",
|
||||
"/Subject": "Test Document",
|
||||
"/Keywords": "test, metadata, pypdf",
|
||||
}
|
||||
)
|
||||
with open(file_path, "wb") as f:
|
||||
writer.write(f)
|
||||
|
||||
result = extract_metadata_from_file(str(file_path))
|
||||
|
||||
# Check that the leading slash is stripped and keys/values match
|
||||
assert result.get("Title") == "Test Title"
|
||||
assert result.get("Author") == "Test Author"
|
||||
assert result.get("Subject") == "Test Document"
|
||||
assert result.get("Keywords") == "test, metadata, pypdf"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
class TestAttachLogo:
|
||||
|
||||
Reference in New Issue
Block a user