diff --git a/.env.example b/.env.example index 84763b5..e592935 100644 --- a/.env.example +++ b/.env.example @@ -17,7 +17,7 @@ POSTGRES_PASSWORD=changeme_super MINIO_ROOT_USER=minioadmin MINIO_ROOT_PASSWORD=changeme_minio_root MINIO_ENDPOINT=minio:9000 -# App-level access key — minimal permissions on docuvault bucket only +# App-level access key. docker-compose.yml provisions this user and a bucket-scoped policy. MINIO_ACCESS_KEY=docuvault_app MINIO_SECRET_KEY=changeme_minio_app MINIO_BUCKET=docuvault @@ -51,9 +51,9 @@ SMTP_PASSWORD= SMTP_FROM=noreply@docuvault.local # ── CORS (Phase 2 — D-09) ──────────────────────────────────────────────────── -# Comma-separated list of allowed origins. Default: http://localhost:5173 -# Example for production: https://app.docuvault.example.com -CORS_ORIGINS=http://localhost:5173 +# JSON list of allowed origins. Default: ["http://localhost:5173"] +# Example for production: ["https://app.docuvault.example.com"] +CORS_ORIGINS=["http://localhost:5173"] # ── Cloud Storage Backends (Phase 5) ───────────────────────────────────────── # Master key for HKDF per-user cloud credential encryption. diff --git a/README.md b/README.md index 34fafc3..0e19c9d 100644 --- a/README.md +++ b/README.md @@ -158,13 +158,14 @@ Copy `.env.example` to `.env`. Only the fields marked **Required** must be set b | `ADMIN_PASSWORD` | Bootstrap admin password (must pass strength check) | | `POSTGRES_PASSWORD` | PostgreSQL superuser password | | `MINIO_ROOT_PASSWORD` | MinIO root password | +| `MINIO_ACCESS_KEY` / `MINIO_SECRET_KEY` | App-level MinIO credentials; Docker Compose provisions this user and bucket policy | | `REDIS_PASSWORD` | Redis `requirepass` password | ### Optional (sensible defaults for local dev) | Variable | Default | Description | |----------|---------|-------------| -| `CORS_ORIGINS` | `http://localhost:5173` | Comma-separated allowed origins | +| `CORS_ORIGINS` | `["http://localhost:5173"]` | JSON list of allowed origins | | `SMTP_HOST` | *(unset)* | Leave empty to log password-reset links to stdout | | `LOG_JSON` | `false` | Set `true` in production for structured JSON logs | | `DEFAULT_AI_PROVIDER` | `ollama` | Used on first boot seed; overridable per-user in admin panel | diff --git a/backend/api/admin/quotas.py b/backend/api/admin/quotas.py index 906abb1..50eaf3f 100644 --- a/backend/api/admin/quotas.py +++ b/backend/api/admin/quotas.py @@ -11,13 +11,9 @@ from deps.auth import get_current_admin from deps.db import get_db from deps.utils import get_client_ip from services.audit import write_audit_log -from api.admin.shared import _user_to_dict - router = APIRouter() # NO prefix — parent __init__.py carries /api/admin (D-04) -# ── Request models ──────────────────────────────────────────────────────────── - class QuotaUpdate(BaseModel): limit_bytes: int @@ -28,9 +24,6 @@ class QuotaUpdate(BaseModel): raise ValueError("limit_bytes must be greater than 0") return v - -# ── Endpoints ───────────────────────────────────────────────────────────────── - @router.get("/users/{user_id}/quota") async def get_user_quota( user_id: uuid.UUID, diff --git a/backend/api/admin/users.py b/backend/api/admin/users.py index facddff..117ff82 100644 --- a/backend/api/admin/users.py +++ b/backend/api/admin/users.py @@ -21,12 +21,8 @@ from api.admin.shared import _user_to_dict router = APIRouter() # NO prefix — parent __init__.py carries /api/admin (D-04) -# ── Constants ───────────────────────────────────────────────────────────────── +_DEFAULT_QUOTA_BYTES = 104857600 -_DEFAULT_QUOTA_BYTES = 104857600 # 100 MB free-tier default (D-06) - - -# ── Request models ──────────────────────────────────────────────────────────── class UserCreate(BaseModel): handle: str @@ -51,20 +47,21 @@ class UserAiConfigUpdate(BaseModel): class SystemTopicCreate(BaseModel): - """Request model for admin system topic creation (D-09).""" - name: str description: str = "" color: str = "#6366f1" class UserDeleteConfirm(BaseModel): - """Admin password confirmation required before hard-deleting a user (ADMIN-02, T-05-11-01).""" - admin_password: str = Field(..., min_length=1) -# ── Endpoints ───────────────────────────────────────────────────────────────── +async def _get_user_or_404(session: AsyncSession, user_id: uuid.UUID) -> User: + user = await session.get(User, user_id) + if user is None: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found") + return user + @router.get("/users") async def list_users( @@ -111,7 +108,7 @@ async def create_user( role=body.role, is_active=True, totp_enabled=False, - password_must_change=True, # ADMIN-01: force password change on first login + password_must_change=True, ) session.add(new_user) @@ -121,8 +118,7 @@ async def create_user( used_bytes=0, ) session.add(quota) - await session.flush() # persist User + Quota before audit_log FK references them - # D-13: admin user created event + await session.flush() _ip_addr = get_client_ip(request) await write_audit_log( session, @@ -151,11 +147,8 @@ async def update_user_status( session: AsyncSession = Depends(get_db), _admin: User = Depends(get_current_admin), ) -> dict: - user = await session.get(User, user_id) - if user is None: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found") + user = await _get_user_or_404(session, user_id) - # Guard: cannot deactivate the only remaining active admin (T-02-29) if not body.is_active and user.role == "admin": count_result = await session.execute( select(func.count(User.id)).where( @@ -174,9 +167,7 @@ async def update_user_status( user.is_active = body.is_active if not body.is_active: - # Revoke all refresh tokens on deactivation await revoke_all_refresh_tokens(session, user.id) - # Revoke any pre-deactivation access tokens still within their TTL (T-7.2-01) await request.app.state.redis.set( f"user_nbf:{user.id}", int(time.time()), @@ -185,7 +176,6 @@ async def update_user_status( session.add(user) - # D-13: user deactivated/activated event _event = "admin.user_deactivated" if not body.is_active else "admin.user_activated" await write_audit_log( session, @@ -211,9 +201,7 @@ async def initiate_password_reset( session: AsyncSession = Depends(get_db), _admin: User = Depends(get_current_admin), ) -> dict: - user = await session.get(User, user_id) - if user is None: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found") + user = await _get_user_or_404(session, user_id) from services.auth import create_password_reset_token # noqa: PLC0415 from config import settings as _settings # noqa: PLC0415 @@ -236,16 +224,13 @@ async def update_ai_config( session: AsyncSession = Depends(get_db), _admin: User = Depends(get_current_admin), ) -> dict: - user = await session.get(User, user_id) - if user is None: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found") + user = await _get_user_or_404(session, user_id) _ip_addr = get_client_ip(request) user.ai_provider = body.ai_provider user.ai_model = body.ai_model session.add(user) - # D-13: AI provider assigned event await write_audit_log( session, event_type="admin.ai_provider_assigned", @@ -273,19 +258,14 @@ async def delete_user( session: AsyncSession = Depends(get_db), _admin: User = Depends(get_current_admin), ) -> None: - # T-05-11-01: Verify admin password before performing any destructive action. - # Fail fast — no DB reads for the target user until the admin is confirmed. if not verify_password(body.admin_password, _admin.password_hash): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="Invalid admin password", ) - user = await session.get(User, user_id) - if user is None: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found") + user = await _get_user_or_404(session, user_id) - # T-04-07-04: Cannot delete admin accounts if user.role == "admin": raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -294,15 +274,11 @@ async def delete_user( _ip_addr = get_client_ip(request) - # SEC-09 (cloud): purge cloud-stored documents and credentials BEFORE DB delete. - # Must run before MinIO cleanup so that credentials are still available to build - # the cloud backend instances for delete_object calls. cloud_conns_result = await session.execute( select(CloudConnection).where(CloudConnection.user_id == user_id) ) cloud_conns = cloud_conns_result.scalars().all() for conn in cloud_conns: - # Delete cloud objects stored in this provider for this user cloud_docs_result = await session.execute( select(Document).where( Document.user_id == user_id, @@ -314,12 +290,10 @@ async def delete_user( backend = await get_storage_backend_for_document(doc, user, session) await backend.delete_object(doc.object_key) except Exception: - pass # Best-effort cloud object cleanup; deletion proceeds regardless - # Purge the credentials row (FK cascade would also remove it, but explicit - # deletion here guarantees credentials_enc is gone before commit — SEC-09) + pass await session.delete(conn) if cloud_conns: - await session.flush() # Flush connection deletes before user delete + await session.flush() await write_audit_log( session, event_type="cloud.credentials_purged", @@ -330,7 +304,6 @@ async def delete_user( metadata_={"providers": [c.provider for c in cloud_conns]}, ) - # SEC-09 (minio): collect all user documents and delete MinIO objects BEFORE DB delete docs_result = await session.execute( select(Document).where(Document.user_id == user_id) ) @@ -341,9 +314,8 @@ async def delete_user( try: await storage.delete_object(doc.object_key) except Exception: - pass # Best-effort MinIO cleanup; DB deletion proceeds regardless + pass - # D-13: audit log BEFORE deleting the user row (user FK still valid at flush time) await write_audit_log( session, event_type="admin.user_deleted", @@ -354,7 +326,6 @@ async def delete_user( ) await session.flush() - # Delete user record (CASCADE removes quota, documents, refresh_tokens, etc.) await session.delete(user) await session.commit() diff --git a/backend/api/audit.py b/backend/api/audit.py index 6884050..d686ca5 100644 --- a/backend/api/audit.py +++ b/backend/api/audit.py @@ -1,23 +1,4 @@ -""" -Admin audit log API endpoints for DocuVault. - -All handlers require get_current_admin (ADMIN-06, SEC-07) — regular users -receive 403 Forbidden. - -Implements: - GET /api/admin/audit-log — paginated, filtered audit log viewer - GET /api/admin/audit-log/export — CSV streaming export with same filters - GET /api/admin/audit-log/daily-exports — list available Celery daily export files - GET /api/admin/audit-log/daily-exports/{date} — stream a specific daily export CSV - -Security invariants: - - All endpoints use Depends(get_current_admin) — verified by grep - - _audit_to_dict() is a pure whitelist: no filename, extracted_text, - password_hash, or credentials_enc can appear in responses (ADMIN-06, D-15) - - CSV export uses the same _audit_to_dict_with_handles() helper as the JSON viewer - - Date path parameter validated against YYYY-MM-DD regex before MinIO key - construction — prevents path traversal (T-06.2-04-01, Pitfall 6) -""" +"""Admin audit log API endpoints.""" from __future__ import annotations import asyncio @@ -44,17 +25,22 @@ from storage.minio_backend import MinIOBackend router = APIRouter(prefix="/api/admin", tags=["audit"]) _VALID_EVENT_PREFIXES = frozenset({"auth", "document", "folder", "share", "admin", "cloud"}) +_CSV_FIELDS = [ + "id", + "event_type", + "user_id", + "actor_id", + "user_handle", + "actor_handle", + "user_email", + "resource_id", + "ip_address", + "metadata_", + "created_at", +] -# ── Safe response helpers ───────────────────────────────────────────────────── - -def _audit_to_dict(entry: AuditLog) -> dict: - """Safe audit log serializer — never includes filename, extracted_text, or - document content (ADMIN-06, D-15). - - Whitelist: id, event_type, user_id, actor_id, resource_id, ip_address, - metadata_, created_at. No other keys are possible. - """ +def _audit_base_fields(entry: AuditLog) -> dict: return { "id": entry.id, "event_type": entry.event_type, @@ -67,38 +53,49 @@ def _audit_to_dict(entry: AuditLog) -> dict: } +def _audit_to_dict(entry: AuditLog) -> dict: + """Whitelisted audit serializer shared with the daily export task.""" + return _audit_base_fields(entry) + + def _audit_to_dict_with_handles( entry: AuditLog, user_handle: Optional[str], actor_handle: Optional[str], user_email: Optional[str] = None, ) -> dict: - """Extended audit log serializer that includes user_handle, actor_handle, and user_email. - - Returns the same fields as _audit_to_dict() plus: - - user_handle: str | None (the handle of the user who owns the entry) - - actor_handle: str | None (the handle of the actor who performed the event) - - user_email: str | None (the email of the user who owns the entry) - - Used by both the JSON viewer and CSV export endpoints (Pitfall 7 — both - endpoints must use the enriched function). - """ - return { - "id": entry.id, - "event_type": entry.event_type, - "user_id": str(entry.user_id) if entry.user_id else None, - "actor_id": str(entry.actor_id) if entry.actor_id else None, + data = _audit_base_fields(entry) + data.update({ "user_handle": user_handle or None, "actor_handle": actor_handle or None, "user_email": user_email or None, - "resource_id": str(entry.resource_id) if entry.resource_id else None, - "ip_address": str(entry.ip_address) if entry.ip_address else None, - "metadata_": entry.metadata_, - "created_at": entry.created_at.isoformat(), - } + }) + return data -# ── Query builder helpers ───────────────────────────────────────────────────── +def _validate_event_type(event_type: Optional[str]) -> None: + if event_type is not None and event_type not in _VALID_EVENT_PREFIXES: + raise HTTPException(status_code=422, detail="Invalid event_type prefix") + + +def _apply_audit_filters( + query, + start: Optional[datetime], + end: Optional[datetime], + user_uuid: Optional[uuid.UUID], + event_type: Optional[str], +): + _validate_event_type(event_type) + if start is not None: + query = query.where(AuditLog.created_at >= start) + if end is not None: + query = query.where(AuditLog.created_at <= end) + if user_uuid is not None: + query = query.where(AuditLog.user_id == user_uuid) + if event_type is not None: + query = query.where(AuditLog.event_type.like(f"{event_type}.%")) + return query + def _build_filtered_query( start: Optional[datetime], @@ -106,27 +103,13 @@ def _build_filtered_query( user_id: Optional[uuid.UUID], event_type: Optional[str], ): - """Return a SQLAlchemy Select for AuditLog with the given filters applied. - - Shared by count queries in both the paginated viewer and the CSV export - endpoints to ensure consistent filter semantics. - - NOTE: This function selects AuditLog only (no JOIN). It is used for COUNT - queries to avoid the subquery ambiguity that arises with multi-column JOINs - (Pitfall 4). Data queries use _build_filtered_query_with_handles() instead. - """ - q = select(AuditLog).order_by(AuditLog.created_at.desc()) - if start is not None: - q = q.where(AuditLog.created_at >= start) - if end is not None: - q = q.where(AuditLog.created_at <= end) - if user_id is not None: - q = q.where(AuditLog.user_id == user_id) - if event_type is not None: - if event_type not in _VALID_EVENT_PREFIXES: - raise HTTPException(status_code=422, detail="Invalid event_type prefix") - q = q.where(AuditLog.event_type.like(f"{event_type}.%")) - return q + return _apply_audit_filters( + select(AuditLog).order_by(AuditLog.created_at.desc()), + start, + end, + user_id, + event_type, + ) def _build_filtered_query_with_handles( @@ -135,15 +118,6 @@ def _build_filtered_query_with_handles( user_uuid: Optional[uuid.UUID], event_type: Optional[str], ): - """Return a multi-column Select that joins User twice for handle enrichment. - - Yields (AuditLog, user_handle: str|None, actor_handle: str|None) tuples. - Uses SQLAlchemy aliased() to join User twice without collision: - - UserSubject: resolves user_id FK → handle - - UserActor: resolves actor_id FK → handle - - outerjoin() ensures entries with NULL user_id or actor_id are still returned. - """ UserSubject = aliased(User) UserActor = aliased(User) @@ -158,37 +132,68 @@ def _build_filtered_query_with_handles( .outerjoin(UserActor, UserActor.id == AuditLog.actor_id) .order_by(AuditLog.created_at.desc()) ) - if start is not None: - q = q.where(AuditLog.created_at >= start) - if end is not None: - q = q.where(AuditLog.created_at <= end) - if user_uuid is not None: - q = q.where(AuditLog.user_id == user_uuid) - if event_type is not None: - if event_type not in _VALID_EVENT_PREFIXES: - raise HTTPException(status_code=422, detail="Invalid event_type prefix") - q = q.where(AuditLog.event_type.like(f"{event_type}.%")) - return q + return _apply_audit_filters(q, start, end, user_uuid, event_type) -# ── Endpoints ───────────────────────────────────────────────────────────────── -# IMPORTANT: daily-export routes are registered BEFORE /audit-log and -# /audit-log/export so FastAPI matches the more specific paths first. +async def _resolve_user_uuid(session: AsyncSession, user_handle: Optional[str]) -> uuid.UUID | None: + if not user_handle: + return None + + result = await session.execute(select(User.id).where(User.handle == user_handle)) + return result.scalar_one_or_none() + + +async def _count_audit_log( + session: AsyncSession, + start: Optional[datetime], + end: Optional[datetime], + user_uuid: Optional[uuid.UUID], + event_type: Optional[str], +) -> int: + count_q = _apply_audit_filters( + select(func.count(AuditLog.id)).where(True), + start, + end, + user_uuid, + event_type, + ) + result = await session.execute(count_q) + return result.scalar_one() + + +def _audit_rows_to_dicts(rows) -> list[dict]: + return [_audit_to_dict_with_handles(row[0], row[1], row[2], row[3]) for row in rows] + + +def _csv_response(csv_text: str, filename: str = "audit-export.csv") -> StreamingResponse: + return StreamingResponse( + iter([csv_text]), + media_type="text/csv", + headers={"Content-Disposition": f"attachment; filename={filename}"}, + ) + + +def _empty_csv_response() -> StreamingResponse: + output = io.StringIO() + csv.DictWriter(output, fieldnames=_CSV_FIELDS).writeheader() + return _csv_response(output.getvalue()) + + +def _audit_csv_response(rows) -> StreamingResponse: + output = io.StringIO() + writer = csv.DictWriter(output, fieldnames=_CSV_FIELDS) + writer.writeheader() + for record in _audit_rows_to_dicts(rows): + record["metadata_"] = json.dumps(record["metadata_"]) if record["metadata_"] is not None else "" + writer.writerow(record) + return _csv_response(output.getvalue()) @router.get("/audit-log/daily-exports") async def list_daily_exports( _admin: User = Depends(get_current_admin), ) -> dict: - """List available Celery daily audit export files from MinIO (D-15). - - Returns: { items: [{ date: "YYYY-MM-DD", key: "audit-logs/YYYY-MM-DD.csv" }] } - Items are sorted descending by date. - - Security: requires get_current_admin — regular users receive 403 (T-06.2-04-02). - Event loop safety: list_objects() is synchronous; wrapped in asyncio.to_thread - to avoid blocking the event loop (T-06.2-04-05). - """ + """List available Celery daily audit export files from MinIO.""" backend = get_storage_backend() if not isinstance(backend, MinIOBackend): return {"items": []} @@ -215,15 +220,7 @@ async def download_daily_export( date: str, _admin: User = Depends(get_current_admin), ) -> StreamingResponse: - """Stream a specific Celery daily audit export file from MinIO (D-16). - - The date path parameter is validated against YYYY-MM-DD regex before - MinIO key construction to prevent path traversal (T-06.2-04-01, Pitfall 6). - - Returns: StreamingResponse with Content-Type: text/csv. - - Security: requires get_current_admin — regular users receive 403 (T-06.2-04-02). - """ + """Stream a specific Celery daily audit export file from MinIO.""" if not re.fullmatch(r"\d{4}-\d{2}-\d{2}", date): raise HTTPException(status_code=404, detail="Invalid date format") @@ -263,56 +260,18 @@ async def list_audit_log( session: AsyncSession = Depends(get_db), _admin: User = Depends(get_current_admin), ) -> dict: - """Return paginated, filtered audit log entries (ADMIN-06). + """Return paginated, filtered audit log entries.""" + user_uuid = await _resolve_user_uuid(session, user_handle) + if user_handle and user_uuid is None: + return {"items": [], "total": 0, "page": page, "per_page": per_page} - Response: { items: [...], total: int, page: int, per_page: int } - Each item includes user_handle and actor_handle alongside UUID fields (D-11). - Entries never contain filename, extracted_text, or document content (D-15). - - user_handle filter: accepts a plain string handle and resolves to UUID - internally. Returns empty results (not 422) for unknown handles (D-12). - """ - # Handle-to-UUID resolution (D-12, Pattern 4) - user_uuid: Optional[uuid.UUID] = None - if user_handle: - handle_result = await session.execute( - select(User.id).where(User.handle == user_handle) - ) - uid = handle_result.scalar_one_or_none() - if uid is None: - # No user with that handle — return empty results (D-12) - return {"items": [], "total": 0, "page": page, "per_page": per_page} - user_uuid = uid - - # Count query: use the plain _build_filtered_query (no JOIN) to avoid - # COUNT ambiguity on multi-column subqueries (Pitfall 4) - count_q = select(func.count(AuditLog.id)).where(True) - if start is not None: - count_q = count_q.where(AuditLog.created_at >= start) - if end is not None: - count_q = count_q.where(AuditLog.created_at <= end) - if user_uuid is not None: - count_q = count_q.where(AuditLog.user_id == user_uuid) - if event_type is not None: - if event_type not in _VALID_EVENT_PREFIXES: - raise HTTPException(status_code=422, detail="Invalid event_type prefix") - count_q = count_q.where(AuditLog.event_type.like(f"{event_type}.%")) - count_result = await session.execute(count_q) - total = count_result.scalar_one() - - # Data query: use enriched JOIN for handle fields + total = await _count_audit_log(session, start, end, user_uuid, event_type) data_q = _build_filtered_query_with_handles(start, end, user_uuid, event_type) data_q = data_q.limit(per_page).offset((page - 1) * per_page) result = await session.execute(data_q) - rows = result.all() - - items = [] - for row in rows: - entry, user_handle_val, actor_handle_val, user_email_val = row[0], row[1], row[2], row[3] - items.append(_audit_to_dict_with_handles(entry, user_handle_val, actor_handle_val, user_email_val)) return { - "items": items, + "items": _audit_rows_to_dicts(result.all()), "total": total, "page": page, "per_page": per_page, @@ -329,68 +288,11 @@ async def export_audit_log( session: AsyncSession = Depends(get_db), _admin: User = Depends(get_current_admin), ) -> StreamingResponse: - """Stream a CSV export of filtered audit log entries (ADMIN-06). + """Stream a CSV export of filtered audit log entries.""" + user_uuid = await _resolve_user_uuid(session, user_handle) + if user_handle and user_uuid is None: + return _empty_csv_response() - Uses the same _audit_to_dict_with_handles() whitelist as the JSON viewer — - includes user_handle and actor_handle; no filename, extracted_text, or - document content appears in the export (D-15, T-04-06-02, Pitfall 7). - - Returns StreamingResponse with Content-Disposition: attachment; filename=audit-export.csv. - - user_handle filter: same handle-to-UUID resolution as the viewer (D-12). - """ - # Handle-to-UUID resolution (D-12) — same logic as list_audit_log - user_uuid: Optional[uuid.UUID] = None - if user_handle: - handle_result = await session.execute( - select(User.id).where(User.handle == user_handle) - ) - uid = handle_result.scalar_one_or_none() - if uid is None: - # Unknown handle — return empty CSV - empty_output = io.StringIO() - fields = [ - "id", "event_type", "user_id", "actor_id", "user_handle", "actor_handle", - "user_email", "resource_id", "ip_address", "metadata_", "created_at", - ] - writer = csv.DictWriter(empty_output, fieldnames=fields) - writer.writeheader() - return StreamingResponse( - iter([empty_output.getvalue()]), - media_type="text/csv", - headers={"Content-Disposition": "attachment; filename=audit-export.csv"}, - ) - user_uuid = uid - - # Data query with handle enrichment (Pitfall 7 — export must use enriched function) q = _build_filtered_query_with_handles(start, end, user_uuid, event_type) result = await session.execute(q) - rows = result.all() - - fields = [ - "id", - "event_type", - "user_id", - "actor_id", - "user_handle", - "actor_handle", - "user_email", - "resource_id", - "ip_address", - "metadata_", - "created_at", - ] - output = io.StringIO() - writer = csv.DictWriter(output, fieldnames=fields) - writer.writeheader() - for row in rows: - entry, user_handle_val, actor_handle_val, user_email_val = row[0], row[1], row[2], row[3] - record = _audit_to_dict_with_handles(entry, user_handle_val, actor_handle_val, user_email_val) - record["metadata_"] = json.dumps(record["metadata_"]) if record["metadata_"] is not None else "" - writer.writerow(record) - - return StreamingResponse( - iter([output.getvalue()]), - media_type="text/csv", - headers={"Content-Disposition": "attachment; filename=audit-export.csv"}, - ) + return _audit_csv_response(result.all()) diff --git a/backend/api/cloud.py b/backend/api/cloud.py index adcca7f..b2d1a33 100644 --- a/backend/api/cloud.py +++ b/backend/api/cloud.py @@ -1,23 +1,4 @@ -""" -Cloud storage connection management API endpoints for DocuVault. - -Endpoints: - GET /api/cloud/oauth/initiate/{provider} — start OAuth flow (Google Drive, OneDrive) - GET /api/cloud/oauth/callback/{provider} — OAuth callback: exchange code, store creds - POST /api/cloud/connections/webdav — connect WebDAV / Nextcloud with credentials - GET /api/cloud/connections — list all active cloud connections - DELETE /api/cloud/connections/{id} — disconnect and purge credentials_enc - GET /api/cloud/folders/{provider}/{folder_id} — lazy-load folder tree (TTL-cached 60s) - PATCH /api/users/me/default-storage — update user's default storage backend - -Security invariants: - - ALL endpoints use Depends(get_regular_user) — admin role gets 403 (D-18, D-19) - - credentials_enc column NEVER appears in any API response (CloudConnectionOut whitelist) - - DELETE /connections/{id} returns 404 for wrong-owner connections (prevents ID enumeration, D-19) - - OAuth state token: secrets.token_urlsafe(32), Redis key "oauth_state:{token}", TTL 1800, single-use - - write_audit_log metadata: only {"provider": provider} — no credentials, no tokens (T-05-05-08) - - WebDAV server_url validated via validate_cloud_url() before backend instantiation (T-05-05-05) -""" +"""Cloud storage connection management endpoints.""" from __future__ import annotations import asyncio @@ -27,7 +8,7 @@ import urllib.parse from typing import Optional from fastapi import APIRouter, Depends, HTTPException, Request, status -from fastapi.responses import RedirectResponse +from fastapi.responses import JSONResponse, RedirectResponse from pydantic import BaseModel from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession @@ -40,17 +21,16 @@ from deps.db import get_db from deps.utils import get_client_ip from services.audit import write_audit_log from services.rate_limiting import account_limiter -from storage.cloud_utils import encrypt_credentials, decrypt_credentials, validate_cloud_url - -# ── Router definitions ──────────────────────────────────────────────────────── +from storage.cloud_backend_factory import build_cloud_backend +from storage.cloud_utils import decrypt_credentials, encrypt_credentials, validate_cloud_url router = APIRouter(prefix="/api/cloud", tags=["cloud"]) users_router = APIRouter(prefix="/api/users", tags=["users"]) -# ── Provider constants ──────────────────────────────────────────────────────── - VALID_OAUTH_PROVIDERS = {"google_drive", "onedrive"} VALID_WEBDAV_PROVIDERS = {"nextcloud", "webdav"} +VALID_CLOUD_PROVIDERS = VALID_OAUTH_PROVIDERS | VALID_WEBDAV_PROVIDERS +_VALID_BACKENDS = frozenset({"minio", *VALID_CLOUD_PROVIDERS}) _DISPLAY_NAMES = { "google_drive": "Google Drive", @@ -59,12 +39,8 @@ _DISPLAY_NAMES = { "webdav": "WebDAV server", } -# ── Request models ──────────────────────────────────────────────────────────── - class WebDAVConnectRequest(BaseModel): - """Request body for POST /api/cloud/connections/webdav.""" - server_url: str username: str password: str @@ -72,197 +48,57 @@ class WebDAVConnectRequest(BaseModel): class DefaultStorageRequest(BaseModel): - """Request body for PATCH /api/users/me/default-storage.""" - backend: str -# ── _call_cloud_op helper ───────────────────────────────────────────────────── +def _master_key() -> bytes: + return settings.cloud_creds_key.encode() -async def _call_cloud_op(conn: CloudConnection, user: User, session: AsyncSession, op_fn): - """Wrap a cloud operation with transparent token refresh and invalid_grant handling. - - Design (B2, D-05, D-06): - 1. Call op_fn() — a zero-argument async callable that performs the cloud operation. - 2. On CloudConnectionError(reason="token_expired"): decrypt current credentials, - refresh the access token via the provider, encrypt the new credentials, update - conn.credentials_enc in the DB, rebuild the backend instance, and retry op_fn() once. - 3. On CloudConnectionError(reason="invalid_grant"): set conn.status="REQUIRES_REAUTH", - commit, write audit log, raise HTTPException(503). - 4. All other exceptions are propagated unchanged. - - Args: - conn: The CloudConnection ORM row for the current user + provider. - user: The authenticated User ORM instance. - session: An active AsyncSession (caller owns commit responsibility). - op_fn: An async callable with no arguments. Must capture the backend instance - in its closure. After token refresh, the caller should pass a new op_fn - that captures the rebuilt backend. - """ - from storage.google_drive_backend import CloudConnectionError # lazy import - - try: - return await op_fn() - except CloudConnectionError as exc: - if exc.reason == "token_expired": - # Refresh token and retry once - master_key = settings.cloud_creds_key.encode() - credentials = decrypt_credentials(master_key, str(user.id), conn.credentials_enc) - - try: - new_credentials = await _refresh_oauth_token(conn.provider, credentials) - except Exception: - # If refresh itself fails, treat as invalid_grant - conn.status = "REQUIRES_REAUTH" - await session.commit() - await write_audit_log( - session, - event_type="cloud.requires_reauth", - user_id=user.id, - actor_id=user.id, - resource_id=conn.id, - ip_address=None, - metadata_={"provider": conn.provider}, - ) - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Cloud connection requires re-authentication. Please reconnect in Settings.", - ) from exc - - # Persist new credentials - conn.credentials_enc = encrypt_credentials(master_key, str(user.id), new_credentials) - await session.commit() - - # Rebuild backend and retry - rebuilt_backend = _build_backend(conn.provider, new_credentials) - - async def retried_op(): - return await _invoke_backend_op(rebuilt_backend, op_fn) - - try: - return await retried_op() - except CloudConnectionError as retry_exc: - if retry_exc.reason == "invalid_grant": - conn.status = "REQUIRES_REAUTH" - await session.commit() - await write_audit_log( - session, - event_type="cloud.requires_reauth", - user_id=user.id, - actor_id=user.id, - resource_id=conn.id, - ip_address=None, - metadata_={"provider": conn.provider}, - ) - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Cloud connection requires re-authentication. Please reconnect in Settings.", - ) from retry_exc - raise - - elif exc.reason == "invalid_grant": - conn.status = "REQUIRES_REAUTH" - await session.commit() - await write_audit_log( - session, - event_type="cloud.requires_reauth", - user_id=user.id, - actor_id=user.id, - resource_id=conn.id, - ip_address=None, - metadata_={"provider": conn.provider}, - ) - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Cloud connection requires re-authentication. Please reconnect in Settings.", - ) from exc - raise +def _oauth_redirect_uri(provider: str) -> str: + return f"{settings.backend_url}/api/cloud/oauth/callback/{provider}" -async def _refresh_oauth_token(provider: str, credentials: dict) -> dict: - """Refresh the OAuth access token for google_drive or onedrive. +def _oauth_state_key(state_token: str) -> str: + return f"oauth_state:{state_token}" - Returns an updated credentials dict with new access_token and expiry. - Raises Exception on failure (including invalid_grant from the provider). - """ - if provider == "google_drive": - from google.oauth2.credentials import Credentials # lazy import - from google.auth.transport.requests import Request # lazy import - creds = Credentials( - token=credentials.get("access_token"), - refresh_token=credentials.get("refresh_token"), - token_uri=credentials.get("token_uri", "https://oauth2.googleapis.com/token"), - client_id=credentials.get("client_id"), - client_secret=credentials.get("client_secret"), +def _settings_redirect(**params: str) -> RedirectResponse: + query = urllib.parse.urlencode(params) + return RedirectResponse(url=f"{settings.frontend_url}/settings?{query}", status_code=302) + + +def _unsupported_provider_error(provider: str, valid_providers: set[str]) -> HTTPException: + return HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"Unsupported provider: {provider}. Valid providers: {sorted(valid_providers)}", + ) + + +def _ensure_oauth_configured(provider: str) -> None: + if provider == "google_drive" and (not settings.google_client_id or not settings.google_client_secret): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Google Drive OAuth is not configured on this server. Set GOOGLE_CLIENT_ID and GOOGLE_CLIENT_SECRET in your environment.", ) - await asyncio.to_thread(creds.refresh, Request()) - new_creds = dict(credentials) - new_creds["access_token"] = creds.token - if creds.expiry: - new_creds["expiry"] = creds.expiry.isoformat() - return new_creds - elif provider == "onedrive": - import msal # lazy import - - app = msal.ConfidentialClientApplication( - settings.onedrive_client_id, - client_credential=settings.onedrive_client_secret, - authority=f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}", + if provider == "onedrive" and ( + not settings.onedrive_client_id + or not settings.onedrive_client_secret + or not settings.onedrive_tenant_id + ): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="OneDrive OAuth is not configured on this server. Set ONEDRIVE_CLIENT_ID, ONEDRIVE_CLIENT_SECRET, and ONEDRIVE_TENANT_ID in your environment.", ) - result = await asyncio.to_thread( - app.acquire_token_by_refresh_token, - credentials.get("refresh_token", ""), - scopes=["Files.ReadWrite", "offline_access"], - ) - if "error" in result: - raise Exception(f"Token refresh failed: {result.get('error_description', result['error'])}") - new_creds = dict(credentials) - new_creds["access_token"] = result["access_token"] - if "refresh_token" in result: - new_creds["refresh_token"] = result["refresh_token"] - return new_creds - - raise ValueError(f"Cannot refresh token for provider: {provider}") -def _build_backend(provider: str, credentials: dict): - """Build a cloud backend instance from provider name and credentials dict.""" - if provider == "google_drive": - from storage.google_drive_backend import GoogleDriveBackend # lazy import - return GoogleDriveBackend(credentials) - elif provider == "onedrive": - from storage.onedrive_backend import OneDriveBackend # lazy import - return OneDriveBackend(credentials) - elif provider == "nextcloud": - from storage.nextcloud_backend import NextcloudBackend # lazy import - return NextcloudBackend( - credentials["server_url"], - credentials["username"], - credentials["password"], - ) - elif provider == "webdav": - from storage.webdav_backend import WebDAVBackend # lazy import - return WebDAVBackend( - credentials["server_url"], - credentials["username"], - credentials["password"], - ) - raise ValueError(f"Unknown provider: {provider}") - - -async def _invoke_backend_op(backend, op_fn): - """Call op_fn, but the backend in the closure needs to be the rebuilt one. - - Since op_fn is already a closure, we call it directly. This helper exists - so the retry path in _call_cloud_op can call op_fn with minimal friction. - """ - return await op_fn() - - -# ── Helper: upsert CloudConnection ─────────────────────────────────────────── +def _webdav_credentials(body: WebDAVConnectRequest) -> dict[str, str]: + return { + "server_url": body.server_url, + "username": body.username, + "password": body.password, + } async def _upsert_cloud_connection( @@ -271,20 +107,6 @@ async def _upsert_cloud_connection( provider: str, credentials_enc: str, ) -> CloudConnection: - """Insert or update a CloudConnection row for the given user + provider. - - If an existing row is found, updates credentials_enc and sets status=ACTIVE. - If no row exists, inserts a new one. - - Args: - session: Active async SQLAlchemy session. - user_id: The authenticated user's UUID. - provider: The cloud provider string (e.g. "google_drive"). - credentials_enc: The HKDF+Fernet-encrypted credentials string. - - Returns: - The CloudConnection ORM instance (added to session, not yet committed). - """ result = await session.execute( select(CloudConnection).where( CloudConnection.user_id == user_id, @@ -310,61 +132,43 @@ async def _upsert_cloud_connection( return conn -# ── GET /api/cloud/oauth/initiate/{provider} ────────────────────────────────── +async def _get_owned_connection( + session: AsyncSession, + connection_id: uuid.UUID, + user_id: uuid.UUID, +) -> CloudConnection: + conn = await session.get(CloudConnection, connection_id) + if conn is None or conn.user_id != user_id: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Connection not found") + return conn -@router.get("/oauth/initiate/{provider}") -@account_limiter.limit("100/minute") -async def oauth_initiate( +async def _get_active_connection( + session: AsyncSession, + user_id: uuid.UUID, provider: str, - request: Request, - current_user: User = Depends(get_regular_user), -) -> dict: - """Start the OAuth flow for Google Drive or OneDrive. - - Generates a CSRF state token, stores it in Redis with TTL 1800 (30 min), - and returns the provider's authorization URL as JSON so the frontend can - navigate using fetch() with the Bearer header (plan 05-10 fix). - - Returns: {"url": ""} - - Security: - - state token is secrets.token_urlsafe(32) — 256 bits of entropy (T-05-05-01) - - Redis key is single-use: deleted in the callback handler (T-05-05-02) - - Only google_drive and onedrive are accepted (T-05-05-06) - - Endpoint requires get_regular_user — no unauthenticated access (T-05-10-01) - """ - request.state.current_user = current_user - from fastapi.responses import JSONResponse # already available via fastapi - - if provider not in VALID_OAUTH_PROVIDERS: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=f"Unsupported OAuth provider: {provider}. Valid providers: {sorted(VALID_OAUTH_PROVIDERS)}", +) -> CloudConnection: + result = await session.execute( + select(CloudConnection).where( + CloudConnection.user_id == user_id, + CloudConnection.provider == provider, + CloudConnection.status == "ACTIVE", ) - - # Pre-flight: validate credentials are configured before allocating Redis state - if provider == "google_drive" and (not settings.google_client_id or not settings.google_client_secret): + ) + conn = result.scalar_one_or_none() + if conn is None: raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail="Google Drive OAuth is not configured on this server. Set GOOGLE_CLIENT_ID and GOOGLE_CLIENT_SECRET in your environment.", - ) - if provider == "onedrive" and ( - not settings.onedrive_client_id - or not settings.onedrive_client_secret - or not settings.onedrive_tenant_id - ): - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail="OneDrive OAuth is not configured on this server. Set ONEDRIVE_CLIENT_ID, ONEDRIVE_CLIENT_SECRET, and ONEDRIVE_TENANT_ID in your environment.", + status_code=status.HTTP_404_NOT_FOUND, + detail="No active cloud connection found for this provider", ) + return conn - state_token = secrets.token_urlsafe(32) - redis_client = request.app.state.redis - await redis_client.setex(f"oauth_state:{state_token}", 1800, str(current_user.id)) - redirect_uri = f"{settings.backend_url}/api/cloud/oauth/callback/{provider}" +def _decrypt_connection(conn: CloudConnection, user_id: uuid.UUID) -> dict: + return decrypt_credentials(_master_key(), str(user_id), conn.credentials_enc) + +async def _oauth_authorization_url(provider: str, redirect_uri: str, state_token: str) -> str: if provider == "google_drive": from google_auth_oauthlib.flow import Flow # lazy import @@ -385,9 +189,9 @@ async def oauth_initiate( prompt="consent", state=state_token, ) - return JSONResponse({"url": authorization_url}) + return authorization_url - elif provider == "onedrive": + if provider == "onedrive": import msal # lazy import app = msal.ConfidentialClientApplication( @@ -395,16 +199,186 @@ async def oauth_initiate( client_credential=settings.onedrive_client_secret, authority=f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}", ) - auth_url = await asyncio.to_thread( + return await asyncio.to_thread( app.get_authorization_request_url, scopes=["Files.ReadWrite", "offline_access"], redirect_uri=redirect_uri, state=state_token, ) - return JSONResponse({"url": auth_url}) + + raise ValueError(f"Unsupported OAuth provider: {provider}") -# ── GET /api/cloud/oauth/callback/{provider} ────────────────────────────────── +async def _oauth_credentials_from_code(provider: str, redirect_uri: str, code: Optional[str]) -> dict: + if provider == "google_drive": + from google_auth_oauthlib.flow import Flow # lazy import + + flow = Flow.from_client_config( + { + "web": { + "client_id": settings.google_client_id, + "client_secret": settings.google_client_secret, + "auth_uri": "https://accounts.google.com/o/oauth2/auth", + "token_uri": "https://oauth2.googleapis.com/token", + } + }, + scopes=["https://www.googleapis.com/auth/drive.file"], + redirect_uri=redirect_uri, + ) + await asyncio.to_thread(flow.fetch_token, code=code) + token_credentials = flow.credentials + return { + "access_token": token_credentials.token, + "refresh_token": token_credentials.refresh_token, + "token_uri": token_credentials.token_uri, + "client_id": token_credentials.client_id, + "client_secret": token_credentials.client_secret, + "expiry": token_credentials.expiry.isoformat() if token_credentials.expiry else None, + } + + if provider == "onedrive": + import msal # lazy import + + app = msal.ConfidentialClientApplication( + settings.onedrive_client_id, + client_credential=settings.onedrive_client_secret, + authority=f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}", + ) + result = await asyncio.to_thread( + app.acquire_token_by_authorization_code, + code, + scopes=["Files.ReadWrite", "offline_access"], + redirect_uri=redirect_uri, + ) + if "error" in result: + raise ValueError(f"Token exchange failed: {result.get('error_description', result['error'])}") + return { + "access_token": result["access_token"], + "refresh_token": result.get("refresh_token", ""), + "token_uri": f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}/oauth2/v2.0/token", + "client_id": settings.onedrive_client_id, + "client_secret": settings.onedrive_client_secret, + } + + raise ValueError(f"Unsupported OAuth provider: {provider}") + + +def _connection_response(conn: CloudConnection) -> dict: + data = CloudConnectionOut.model_validate(conn).model_dump() + if conn.provider not in VALID_WEBDAV_PROVIDERS: + return data + + try: + credentials = _decrypt_connection(conn, conn.user_id) + except Exception: + return data + + data["server_url"] = credentials.get("server_url") + data["connection_username"] = credentials.get("username") + return data + + +async def _fetch_google_drive_folders(credentials: dict, folder_id: str) -> list: + from google.oauth2.credentials import Credentials # lazy import + from googleapiclient.discovery import build # lazy import + + import datetime as dt + + expiry = None + if expiry_str := credentials.get("expiry"): + try: + expiry = dt.datetime.fromisoformat(expiry_str) + except ValueError: + pass + + creds = Credentials( + token=credentials.get("access_token"), + refresh_token=credentials.get("refresh_token"), + token_uri=credentials.get("token_uri", "https://oauth2.googleapis.com/token"), + client_id=credentials.get("client_id"), + client_secret=credentials.get("client_secret"), + expiry=expiry, + ) + + def _list_files() -> list: + service = build("drive", "v3", credentials=creds, cache_discovery=False) + response = service.files().list( + q=f"'{folder_id}' in parents and trashed=false", + fields="files(id,name,mimeType,size)", + pageSize=200, + ).execute() + items = [] + for item in response.get("files", []): + is_dir = item.get("mimeType") == "application/vnd.google-apps.folder" + items.append({ + "id": item["id"], + "name": item["name"], + "is_dir": is_dir, + "size": int(item.get("size", 0)) if not is_dir else 0, + }) + return items + + return await asyncio.to_thread(_list_files) + + +async def _fetch_onedrive_folders(credentials: dict, folder_id: str) -> list: + import httpx # lazy import + + if folder_id in ("root", ""): + url = "https://graph.microsoft.com/v1.0/me/drive/root/children" + else: + url = f"https://graph.microsoft.com/v1.0/me/drive/items/{folder_id}/children" + + async with httpx.AsyncClient() as client: + resp = await client.get( + url, + headers={"Authorization": f"Bearer {credentials.get('access_token', '')}"}, + timeout=30, + ) + resp.raise_for_status() + data = resp.json() + + items = [] + for item in data.get("value", []): + is_dir = "folder" in item + items.append({ + "id": item["id"], + "name": item["name"], + "is_dir": is_dir, + "size": item.get("size", 0) if not is_dir else 0, + }) + return items + + +async def _fetch_webdav_folders(provider: str, credentials: dict, folder_id: str) -> list: + webdav_path = "" if folder_id == "root" else folder_id + backend = build_cloud_backend(provider, credentials) + return await backend.list_folder(webdav_path) + + +@router.get("/oauth/initiate/{provider}") +@account_limiter.limit("100/minute") +async def oauth_initiate( + provider: str, + request: Request, + current_user: User = Depends(get_regular_user), +) -> dict: + """Start an OAuth flow and return the provider authorization URL.""" + request.state.current_user = current_user + + if provider not in VALID_OAUTH_PROVIDERS: + raise _unsupported_provider_error(provider, VALID_OAUTH_PROVIDERS) + _ensure_oauth_configured(provider) + + state_token = secrets.token_urlsafe(32) + await request.app.state.redis.setex(_oauth_state_key(state_token), 1800, str(current_user.id)) + + authorization_url = await _oauth_authorization_url( + provider, + redirect_uri=_oauth_redirect_uri(provider), + state_token=state_token, + ) + return JSONResponse({"url": authorization_url}) @router.get("/oauth/callback/{provider}", response_class=RedirectResponse) @@ -413,25 +387,11 @@ async def oauth_callback( request: Request, session: AsyncSession = Depends(get_db), ): - """Handle the OAuth callback from Google Drive or OneDrive. - - Validates state token (single-use, from Redis), exchanges the authorization - code for tokens, encrypts credentials, and upserts a CloudConnection row. - - On success: redirects to {settings.frontend_url}/settings?cloud_connected={provider} - On any error: redirects to {settings.frontend_url}/settings?cloud_error={message} - - Security: - - state token is looked up in Redis (T-05-05-01), deleted on first use (T-05-05-02) - - Missing state → 400 redirect (not silent failure) - - credentials_enc stored, never returned in response - """ + """Exchange OAuth code, store encrypted credentials, and redirect to settings.""" state = request.query_params.get("state") code = request.query_params.get("code") error_param = request.query_params.get("error") - frontend_settings = f"{settings.frontend_url}/settings" - try: if error_param: raise ValueError(f"OAuth provider returned error: {error_param}") @@ -443,86 +403,26 @@ async def oauth_callback( raise ValueError("Missing OAuth state parameter") redis_client = request.app.state.redis - stored_user_id = await redis_client.get(f"oauth_state:{state}") + state_key = _oauth_state_key(state) + stored_user_id = await redis_client.get(state_key) if not stored_user_id: - # state is missing or expired — single-use validation failed - return RedirectResponse( - url=f"{frontend_settings}?cloud_error={urllib.parse.quote('Invalid or expired OAuth state. Please try connecting again.')}", - status_code=302, + return _settings_redirect( + cloud_error="Invalid or expired OAuth state. Please try connecting again." ) - # Delete the Redis key immediately — single-use (T-05-05-02) - await redis_client.delete(f"oauth_state:{state}") + await redis_client.delete(state_key) - # Parse the user_id from the stored value if isinstance(stored_user_id, bytes): stored_user_id = stored_user_id.decode("utf-8") user_id = uuid.UUID(stored_user_id) - # Load user from DB (the OAuth callback is not authenticated via JWT — - # the state token binds this callback to the correct user session) user = await session.get(User, user_id) if user is None or not user.is_active: raise ValueError("User not found or inactive") - redirect_uri = f"{settings.backend_url}/api/cloud/oauth/callback/{provider}" - master_key = settings.cloud_creds_key.encode() - - if provider == "google_drive": - from google_auth_oauthlib.flow import Flow # lazy import - - flow = Flow.from_client_config( - { - "web": { - "client_id": settings.google_client_id, - "client_secret": settings.google_client_secret, - "auth_uri": "https://accounts.google.com/o/oauth2/auth", - "token_uri": "https://oauth2.googleapis.com/token", - } - }, - scopes=["https://www.googleapis.com/auth/drive.file"], - redirect_uri=redirect_uri, - ) - await asyncio.to_thread(flow.fetch_token, code=code) - token_credentials = flow.credentials - credentials = { - "access_token": token_credentials.token, - "refresh_token": token_credentials.refresh_token, - "token_uri": token_credentials.token_uri, - "client_id": token_credentials.client_id, - "client_secret": token_credentials.client_secret, - "expiry": token_credentials.expiry.isoformat() if token_credentials.expiry else None, - } - - elif provider == "onedrive": - import msal # lazy import - - app = msal.ConfidentialClientApplication( - settings.onedrive_client_id, - client_credential=settings.onedrive_client_secret, - authority=f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}", - ) - result = await asyncio.to_thread( - app.acquire_token_by_authorization_code, - code, - scopes=["Files.ReadWrite", "offline_access"], - redirect_uri=redirect_uri, - ) - if "error" in result: - raise ValueError(f"Token exchange failed: {result.get('error_description', result['error'])}") - credentials = { - "access_token": result["access_token"], - "refresh_token": result.get("refresh_token", ""), - "token_uri": f"https://login.microsoftonline.com/{settings.onedrive_tenant_id}/oauth2/v2.0/token", - "client_id": settings.onedrive_client_id, - "client_secret": settings.onedrive_client_secret, - } - - else: - raise ValueError(f"Unsupported OAuth provider: {provider}") - - credentials_enc = encrypt_credentials(master_key, str(user_id), credentials) + credentials = await _oauth_credentials_from_code(provider, _oauth_redirect_uri(provider), code) + credentials_enc = encrypt_credentials(_master_key(), str(user_id), credentials) conn = await _upsert_cloud_connection(session, user_id, provider, credentials_enc) await session.flush() @@ -537,20 +437,10 @@ async def oauth_callback( ) await session.commit() - return RedirectResponse( - url=f"{frontend_settings}?cloud_connected={provider}", - status_code=302, - ) + return _settings_redirect(cloud_connected=provider) except Exception as exc: - error_msg = str(exc) - return RedirectResponse( - url=f"{frontend_settings}?cloud_error={urllib.parse.quote(error_msg)}", - status_code=302, - ) - - -# ── POST /api/cloud/connections/webdav ──────────────────────────────────────── + return _settings_redirect(cloud_error=str(exc)) @router.post("/connections/webdav", status_code=status.HTTP_201_CREATED) @@ -561,16 +451,7 @@ async def connect_webdav( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> dict: - """Connect a WebDAV or Nextcloud server. - - Validates the URL for SSRF, tests the connection with health_check(), - encrypts credentials, and upserts a CloudConnection row. - - Security: - - validate_cloud_url() blocks RFC-1918, loopback, and link-local (T-05-05-05) - - health_check() requires a successful PROPFIND before storing credentials - - credentials_enc never returned in response (CloudConnectionOut whitelist) - """ + """Connect a WebDAV or Nextcloud server.""" request.state.current_user = current_user if body.provider not in VALID_WEBDAV_PROVIDERS: raise HTTPException( @@ -578,7 +459,6 @@ async def connect_webdav( detail=f"Unsupported WebDAV provider: {body.provider}. Valid values: {sorted(VALID_WEBDAV_PROVIDERS)}", ) - # SSRF validation (T-05-05-05) try: validate_cloud_url(body.server_url) except ValueError as exc: @@ -587,14 +467,9 @@ async def connect_webdav( detail=f"Invalid server URL: {exc}", ) from exc - # Instantiate backend and run health_check try: - if body.provider == "nextcloud": - from storage.nextcloud_backend import NextcloudBackend # lazy import - backend = NextcloudBackend(body.server_url, body.username, body.password) - else: - from storage.webdav_backend import WebDAVBackend # lazy import - backend = WebDAVBackend(body.server_url, body.username, body.password) + credentials = _webdav_credentials(body) + backend = build_cloud_backend(body.provider, credentials) except ValueError as exc: raise HTTPException( status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, @@ -616,13 +491,7 @@ async def connect_webdav( detail=f"Connection test failed — check server URL and credentials: {exc}", ) from exc - credentials = { - "server_url": body.server_url, - "username": body.username, - "password": body.password, - } - master_key = settings.cloud_creds_key.encode() - credentials_enc = encrypt_credentials(master_key, str(current_user.id), credentials) + credentials_enc = encrypt_credentials(_master_key(), str(current_user.id), credentials) conn = await _upsert_cloud_connection(session, current_user.id, body.provider, credentials_enc) await session.flush() @@ -643,9 +512,6 @@ async def connect_webdav( return CloudConnectionOut.model_validate(conn).model_dump() -# ── GET /api/cloud/connections ──────────────────────────────────────────────── - - @router.get("/connections") @account_limiter.limit("100/minute") async def list_connections( @@ -653,33 +519,12 @@ async def list_connections( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> dict: - """List all cloud connections for the current user. - - Security: - - Only connections owned by current_user.id are returned - - credentials_enc excluded by CloudConnectionOut whitelist (T-05-05-03) - """ + """List the current user's cloud connections.""" request.state.current_user = current_user result = await session.execute( select(CloudConnection).where(CloudConnection.user_id == current_user.id) ) - connections = result.scalars().all() - items = [] - master_key = settings.cloud_creds_key.encode() - for conn in connections: - d = CloudConnectionOut.model_validate(conn).model_dump() - if conn.provider in ("nextcloud", "webdav"): - try: - creds = decrypt_credentials(master_key, str(conn.user_id), conn.credentials_enc) - d["server_url"] = creds.get("server_url") - d["connection_username"] = creds.get("username") - except Exception: - pass - items.append(d) - return {"items": items} - - -# ── GET /api/cloud/connections/{connection_id}/config ──────────────────────── + return {"items": [_connection_response(conn) for conn in result.scalars().all()]} @router.get("/connections/{connection_id}/config") @@ -690,24 +535,9 @@ async def get_connection_config( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> dict: - """Return non-secret configuration fields for a WebDAV/Nextcloud connection. - - Returns server_url and connection_username (not password) so the frontend - can pre-populate the Edit modal without exposing credentials. - - Only applicable to WebDAV / Nextcloud connections (not OAuth providers). - Returns 404 for wrong-owner or unknown connections (prevents ID enumeration). - Returns 400 for OAuth providers (no non-secret config to return). - - Security: - - Only connection owned by current_user.id is returned (T-05-05-04) - - password is never included in the response (D-18) - - Returns 404 for wrong-owner connections (prevents ID enumeration) - """ + """Return non-secret WebDAV/Nextcloud connection fields.""" request.state.current_user = current_user - conn = await session.get(CloudConnection, connection_id) - if conn is None or conn.user_id != current_user.id: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Connection not found") + conn = await _get_owned_connection(session, connection_id, current_user.id) if conn.provider not in VALID_WEBDAV_PROVIDERS: raise HTTPException( @@ -715,16 +545,14 @@ async def get_connection_config( detail="Connection config is only available for WebDAV/Nextcloud connections", ) - master_key = settings.cloud_creds_key.encode() try: - credentials = decrypt_credentials(master_key, str(current_user.id), conn.credentials_enc) + credentials = _decrypt_connection(conn, current_user.id) except Exception: raise HTTPException( status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Failed to decrypt connection credentials", ) - # Return non-secret fields only — never expose the password return { "id": str(conn.id), "provider": conn.provider, @@ -733,9 +561,6 @@ async def get_connection_config( } -# ── DELETE /api/cloud/connections/{connection_id} ───────────────────────────── - - @router.delete("/connections/{connection_id}", status_code=status.HTTP_204_NO_CONTENT) @account_limiter.limit("100/minute") async def delete_connection( @@ -744,24 +569,14 @@ async def delete_connection( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> None: - """Disconnect a cloud connection and permanently purge credentials_enc. - - Returns 404 (not 403) for connections owned by other users to prevent - connection ID enumeration (T-05-05-04, D-19). - - On success: connection row is deleted, audit log written, cache invalidated. - """ + """Disconnect a cloud connection.""" request.state.current_user = current_user - conn = await session.get(CloudConnection, connection_id) - - # Return 404 for any access failure — prevents ID enumeration (T-05-05-04) - if conn is None or conn.user_id != current_user.id: - raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Connection not found") + conn = await _get_owned_connection(session, connection_id, current_user.id) provider = conn.provider - # Invalidate folder listing cache for this provider from services.cloud_cache import invalidate_provider_cache # lazy import + invalidate_provider_cache(str(current_user.id), provider) _ip = get_client_ip(request) @@ -780,9 +595,6 @@ async def delete_connection( await session.commit() -# ── GET /api/cloud/folders/{provider}/{folder_id} ───────────────────────────── - - @router.get("/folders/{provider}/{folder_id:path}") @account_limiter.limit("100/minute") async def list_cloud_folders( @@ -792,157 +604,30 @@ async def list_cloud_folders( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> dict: - """List folder contents for a cloud provider (TTL-cached 60 seconds, D-16). - - Loads the active CloudConnection for current_user + provider, decrypts - credentials, builds the appropriate backend, and calls list_folder via - the TTL cache. - - Returns 404 if no active connection found (prevents enumeration). - """ + """List folder contents for a connected cloud provider.""" request.state.current_user = current_user - all_providers = VALID_OAUTH_PROVIDERS | VALID_WEBDAV_PROVIDERS - if provider not in all_providers: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=f"Unsupported provider: {provider}", - ) + if provider not in VALID_CLOUD_PROVIDERS: + raise _unsupported_provider_error(provider, VALID_CLOUD_PROVIDERS) - result = await session.execute( - select(CloudConnection).where( - CloudConnection.user_id == current_user.id, - CloudConnection.provider == provider, - CloudConnection.status == "ACTIVE", - ) - ) - conn = result.scalar_one_or_none() - if conn is None: - raise HTTPException( - status_code=status.HTTP_404_NOT_FOUND, - detail="No active cloud connection found for this provider", - ) - - master_key = settings.cloud_creds_key.encode() - credentials = decrypt_credentials(master_key, str(current_user.id), conn.credentials_enc) + conn = await _get_active_connection(session, current_user.id, provider) + credentials = _decrypt_connection(conn, current_user.id) from services.cloud_cache import get_cloud_folders_cached # lazy import if provider == "google_drive": - from storage.google_drive_backend import GoogleDriveBackend # lazy import - - backend = GoogleDriveBackend(credentials) - - async def fetch_google_drive() -> list: - """Fetch Google Drive folder listing via Drive v3 files.list.""" - from googleapiclient.discovery import build # lazy import - from google.oauth2.credentials import Credentials # lazy import - - import datetime as _dt - - expiry_str = credentials.get("expiry") - expiry = None - if expiry_str: - try: - expiry = _dt.datetime.fromisoformat(expiry_str) - except ValueError: - pass - - creds = Credentials( - token=credentials.get("access_token"), - refresh_token=credentials.get("refresh_token"), - token_uri=credentials.get("token_uri", "https://oauth2.googleapis.com/token"), - client_id=credentials.get("client_id"), - client_secret=credentials.get("client_secret"), - expiry=expiry, - ) - - def _list_files(): - service = build("drive", "v3", credentials=creds, cache_discovery=False) - query = f"'{folder_id}' in parents and trashed=false" - response = service.files().list( - q=query, - fields="files(id,name,mimeType,size)", - pageSize=200, - ).execute() - files = response.get("files", []) - result = [] - for f in files: - is_dir = f.get("mimeType") == "application/vnd.google-apps.folder" - result.append({ - "id": f["id"], - "name": f["name"], - "is_dir": is_dir, - "size": int(f.get("size", 0)) if not is_dir else 0, - }) - return result - - return await asyncio.to_thread(_list_files) - - items = await get_cloud_folders_cached( - str(current_user.id), provider, folder_id, fetch_google_drive - ) + fetcher = lambda: _fetch_google_drive_folders(credentials, folder_id) elif provider == "onedrive": - async def fetch_onedrive() -> list: - """Fetch OneDrive folder listing via Microsoft Graph.""" - import httpx # lazy import - - access_token = credentials.get("access_token", "") - if folder_id in ("root", ""): - url = "https://graph.microsoft.com/v1.0/me/drive/root/children" - else: - url = f"https://graph.microsoft.com/v1.0/me/drive/items/{folder_id}/children" - - async with httpx.AsyncClient() as client: - resp = await client.get( - url, - headers={"Authorization": f"Bearer {access_token}"}, - timeout=30, - ) - resp.raise_for_status() - data = resp.json() - - result = [] - for item in data.get("value", []): - is_dir = "folder" in item - result.append({ - "id": item["id"], - "name": item["name"], - "is_dir": is_dir, - "size": item.get("size", 0) if not is_dir else 0, - }) - return result - - items = await get_cloud_folders_cached( - str(current_user.id), provider, folder_id, fetch_onedrive - ) - - elif provider in ("nextcloud", "webdav"): - backend = _build_backend(provider, credentials) - # "root" is a frontend sentinel meaning the WebDAV root; translate to "" - webdav_path = "" if folder_id == "root" else folder_id - - async def fetch_webdav() -> list: - return await backend.list_folder(webdav_path) - - items = await get_cloud_folders_cached( - str(current_user.id), provider, folder_id, fetch_webdav - ) + fetcher = lambda: _fetch_onedrive_folders(credentials, folder_id) else: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=f"Unsupported provider: {provider}", - ) + fetcher = lambda: _fetch_webdav_folders(provider, credentials, folder_id) + + items = await get_cloud_folders_cached(str(current_user.id), provider, folder_id, fetcher) return {"items": items} -# ── PATCH /api/users/me/default-storage ────────────────────────────────────── - -_VALID_BACKENDS = frozenset({"minio", "google_drive", "onedrive", "nextcloud", "webdav"}) - - @users_router.patch("/me/default-storage") @account_limiter.limit("100/minute") async def update_default_storage( diff --git a/backend/api/documents/content.py b/backend/api/documents/content.py index 8f5f9a8..156126b 100644 --- a/backend/api/documents/content.py +++ b/backend/api/documents/content.py @@ -15,14 +15,12 @@ Security: from __future__ import annotations import urllib.parse -import uuid - from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi.responses import StreamingResponse -from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession -from db.models import Document, Share, User +from db.models import User +from api.documents.shared import get_accessible_document from deps.auth import get_regular_user from deps.db import get_db from services.rate_limiting import account_limiter @@ -31,15 +29,12 @@ from storage.exceptions import CloudConnectionError router = APIRouter() -# ── Range header parsing helper ─────────────────────────────────────────────── - def _parse_range(range_header: str, file_size: int) -> tuple: """Parse a 'bytes=X-Y' Range header and return (start, end). Returns (start, end) where both are inclusive byte offsets. Raises HTTP 416 on any invalid or out-of-bounds range. - T-04-05-03: validates start <= end, start >= 0, end < file_size. """ try: h = range_header.replace("bytes=", "").split("-") @@ -52,8 +47,6 @@ def _parse_range(range_header: str, file_size: int) -> tuple: return start, end -# ── GET /api/documents/{doc_id}/content ────────────────────────────────────── - @router.get("/{doc_id}/content") @account_limiter.limit("100/minute") async def stream_document_content( @@ -63,26 +56,7 @@ async def stream_document_content( current_user: User = Depends(get_regular_user), ): request.state.current_user = current_user - try: - uid = uuid.UUID(doc_id) - except ValueError: - raise HTTPException(status_code=404, detail="Document not found") - - doc = await session.get(Document, uid) - if doc is None: - raise HTTPException(status_code=404, detail="Document not found") - - # Access control: owner OR share recipient (T-04-05-04) - if doc.user_id != current_user.id: - result = await session.execute( - select(Share).where( - Share.document_id == doc.id, - Share.recipient_id == current_user.id, - ) - ) - share = result.scalar_one_or_none() - if share is None: - raise HTTPException(status_code=404, detail="Document not found") + doc, _is_recipient = await get_accessible_document(session, doc_id, current_user.id) try: import api.documents as _doc_pkg # late import allows test monkeypatching via api.documents diff --git a/backend/api/documents/crud.py b/backend/api/documents/crud.py index 5b989d5..b6af280 100644 --- a/backend/api/documents/crud.py +++ b/backend/api/documents/crud.py @@ -1,22 +1,4 @@ -"""Document CRUD endpoints — list, get, patch, delete, and re-classify. - -Endpoints: - GET "" — list documents with sort, folder filter, and FTS (list_documents) - GET /{doc_id} — get document metadata (get_document) - PATCH /{doc_id} — update filename and/or folder_id (patch_document) - DELETE /{doc_id} — delete document, decrement quota atomically (delete_document) - POST /{doc_id}/classify — re-queue Celery classification (classify_document, D-08) - -Sub-router carries NO prefix — prefix="/api/documents" lives in __init__.py (D-04). - -Security: - T-03-11: ownership assertion on every resource endpoint — cross-user access returns 404. - T-05-09-01: get_regular_user dep rejects admins (403) and unauthenticated (401). - T-05-09-02: response uses storage.get_metadata() whitelist — no credentials_enc, no password_hash. - T-06.2-03-01: cloud documents skip MinIO quota decrement. - T-06.2-03-02: cloud delete failure returns {success: false, cloud_delete_failed: true} (HTTP 200). - T-07-10: classify endpoint — IDOR returns 404 per ownership assertion. -""" +"""Document CRUD endpoints.""" from __future__ import annotations import uuid @@ -41,14 +23,31 @@ from services.rate_limiting import account_limiter from storage import get_storage_backend_for_document as _get_storage_backend_for_document from tasks.document_tasks import extract_and_classify -from api.documents.shared import DocumentPatch, _CLOUD_PROVIDERS +from api.documents.shared import DocumentPatch, get_accessible_document, get_owned_document router = APIRouter() -# ── GET /api/documents ──────────────────────────────────────────────────────── -# Route registered on parent router in __init__.py (FastAPI 0.100+ disallows -# include_router when both the include prefix and route path are empty strings). +async def _shared_document_ids(session: AsyncSession, user_id: uuid.UUID) -> set[uuid.UUID]: + result = await session.execute(select(Share.document_id).where(Share.owner_id == user_id)) + return {row[0] for row in result.fetchall()} + + +async def _decorate_shared_flags( + session: AsyncSession, + user_id: uuid.UUID, + items: list[dict], +) -> list[dict]: + shared_ids = await _shared_document_ids(session, user_id) + for item in items: + try: + doc_id = uuid.UUID(item.get("id", "")) + except (TypeError, ValueError): + item["is_shared"] = False + else: + item["is_shared"] = doc_id in shared_ids + return items + @account_limiter.limit("100/minute") async def list_documents( @@ -63,40 +62,17 @@ async def list_documents( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """List documents with optional sort, folder filter, and full-text search. - - D-16: requires authenticated regular user (get_regular_user rejects admins). - Returns only documents belonging to the current user. - - FOLD-05: sort by name|date|size; order asc|desc; folder_id filter; - q full-text search via plainto_tsquery (PostgreSQL only — silently skipped - on SQLite when function is unavailable). FTS scope is always scoped to - current_user.id (T-04-03-02). - - Backward-compat: when sort/order/folder_id/q are not provided, behaviour - is identical to the pre-Phase-4 implementation. - """ + """List documents with optional sort, folder filter, and full-text search.""" request.state.current_user = current_user - # If no new params used, fall through to the legacy storage.list_metadata path - # to preserve full backward compatibility with topic filtering. if folder_id is None and q is None and sort == "date" and order == "desc": docs = await storage.list_metadata(session, user_id=current_user.id, topic=topic) total = len(docs) start = (page - 1) * per_page - # Add is_shared field (Phase 4 addition) - shared_result = await session.execute( - select(Share.document_id).where(Share.owner_id == current_user.id) + items = await _decorate_shared_flags( + session, + current_user.id, + docs[start : start + per_page], ) - shared_ids = {row[0] for row in shared_result.fetchall()} - items = [] - for d in docs[start : start + per_page]: - doc_id_str = d.get("id", "") - try: - doc_uuid = uuid.UUID(doc_id_str) - except (ValueError, AttributeError): - doc_uuid = None - d["is_shared"] = doc_uuid in shared_ids if doc_uuid else False - items.append(d) return {"items": items, "total": total, "page": page, "per_page": per_page} from db.models import DocumentTopic, Topic # noqa: PLC0415 (avoid circular at module top) @@ -126,8 +102,6 @@ async def list_documents( order_fn = sort_col.asc if order == "asc" else sort_col.desc stmt = stmt.order_by(order_fn()) - # Full-text search — plainto_tsquery on extracted_text (PostgreSQL only) - # Falls back to unfiltered if the DB dialect doesn't support @@ (e.g. SQLite in test env) fts_requested = q is not None and len(q) >= 2 if fts_requested: fts_stmt = stmt.where( @@ -143,14 +117,11 @@ async def list_documents( result = await session.execute(stmt) docs_orm = result.scalars().all() - shared_result = await session.execute( - select(Share.document_id).where(Share.owner_id == current_user.id) - ) - shared_ids = {row[0] for row in shared_result.fetchall()} - all_items = [] + shared_ids = await _shared_document_ids(session, current_user.id) for doc in docs_orm: from services.storage import _doc_to_dict, _load_topic_names # noqa: PLC0415 + topic_names = await _load_topic_names(session, doc.id) d = _doc_to_dict(doc, topic_names) d["is_shared"] = doc.id in shared_ids @@ -166,8 +137,6 @@ async def list_documents( } -# ── GET /api/documents/{doc_id} ─────────────────────────────────────────────── - @router.get("/{doc_id}") @account_limiter.limit("100/minute") async def get_document( @@ -176,44 +145,18 @@ async def get_document( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Return document metadata by ID. - - D-16: requires authenticated regular user. Asserts ownership — cross-user - access returns 404 (not 403) to avoid information leakage (T-03-11). - """ + """Return document metadata by ID.""" request.state.current_user = current_user - try: - uid = uuid.UUID(doc_id) - except ValueError: - raise HTTPException(404, "Document not found") - - doc = await session.get(Document, uid) - if doc is None: - raise HTTPException(404, "Document not found") - - is_recipient = False - if doc.user_id != current_user.id: - share_result = await session.execute( - select(Share).where( - Share.document_id == uid, - Share.recipient_id == current_user.id, - ) - ) - if share_result.scalar_one_or_none() is None: - raise HTTPException(404, "Document not found") - is_recipient = True + _, is_recipient = await get_accessible_document(session, doc_id, current_user.id) meta = await storage.get_metadata(session, doc_id) if meta is None: raise HTTPException(404, "Document not found") - # T-04-04-03: recipients get metadata only — extracted_text excluded (consistent with /shares/received) if is_recipient: meta.pop("extracted_text", None) return meta -# ── PATCH /api/documents/{doc_id} ──────────────────────────────────────────── - @router.patch("/{doc_id}") @account_limiter.limit("100/minute") async def patch_document( @@ -223,25 +166,9 @@ async def patch_document( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Update document metadata (filename and/or folder_id). - - T-05-09-01: get_regular_user dep rejects admins (403) and unauthenticated (401). - T-05-09-01: ownership check — non-owner gets 404 to avoid leaking document IDs (D-16). - T-05-09-02: response uses storage.get_metadata() which excludes credentials_enc and - password_hash via the _doc_to_dict whitelist. - - At least one field must be provided — empty body returns 422. - folder_id=null moves the document to the root (no folder). - """ + """Update document metadata.""" request.state.current_user = current_user - try: - uid = uuid.UUID(doc_id) - except ValueError: - raise HTTPException(404, "Document not found") - - doc = await session.get(Document, uid) - if doc is None or doc.user_id != current_user.id: - raise HTTPException(404, "Document not found") + doc = await get_owned_document(session, doc_id, current_user.id) if not body.model_fields_set: raise HTTPException(422, "At least one field (filename, folder_id) must be provided") @@ -250,7 +177,6 @@ async def patch_document( doc.filename = body.filename if "folder_id" in body.model_fields_set: - # folder_id=null → move to root (no folder); folder_id= → move to folder if body.folder_id is not None: target = await session.get(Folder, body.folder_id) if target is None or target.user_id != current_user.id: @@ -265,8 +191,6 @@ async def patch_document( return meta -# ── DELETE /api/documents/{doc_id} ─────────────────────────────────────────── - @router.delete("/{doc_id}") @account_limiter.limit("100/minute") async def delete_document( @@ -276,27 +200,9 @@ async def delete_document( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Delete a document and decrement quota atomically. - - For cloud-stored documents: - - Default path: attempt cloud provider delete first; on failure return - {success: false, cloud_delete_failed: true} (HTTP 200) so the frontend - can offer a "Remove from app" fallback (T-06.2-03-02). - - remove_only=true: skip cloud delete, remove DB row only, skip quota decrement. - - Cloud docs always use skip_quota=True (never charged MinIO quota, T-06.2-03-01). - - D-16: requires authenticated regular user. Asserts ownership — cross-user - delete returns 404 (not 403) to avoid information leakage (T-03-11). - """ + """Delete a document and decrement quota when appropriate.""" request.state.current_user = current_user - try: - uid = uuid.UUID(doc_id) - except ValueError: - raise HTTPException(404, "Document not found") - - doc = await session.get(Document, uid) - if doc is None or doc.user_id != current_user.id: - raise HTTPException(404, "Document not found") + doc = await get_owned_document(session, doc_id, current_user.id) is_cloud = doc.storage_backend != "minio" _doc_size = doc.size_bytes @@ -305,7 +211,8 @@ async def delete_document( if is_cloud and not remove_only: try: - import api.documents as _doc_pkg # late import allows test monkeypatching via api.documents + import api.documents as _doc_pkg + _gsb = _doc_pkg.get_storage_backend_for_document cloud_backend = await _gsb(doc, current_user, session) await cloud_backend.delete_object(doc.object_key) @@ -320,13 +227,10 @@ async def delete_document( }, ) - # auto_commit=False defers the commit so the audit log write below happens - # in the same transaction — avoids the split-transaction gap (WR-08). ok = await storage.delete_document(session, doc_id, skip_quota=is_cloud, auto_commit=False) if not ok: raise HTTPException(404, "Document not found") - # D-13: document deleted event — written in the same transaction as the delete (WR-08). await write_audit_log( session, event_type="document.deleted", @@ -341,8 +245,6 @@ async def delete_document( return {"success": True} -# ── POST /api/documents/{doc_id}/classify ──────────────────────────────────── - @router.post("/{doc_id}/classify") @account_limiter.limit("100/minute") async def classify_document( @@ -351,25 +253,9 @@ async def classify_document( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Re-queue a document for classification via Celery (D-11). - - Sets doc.status='processing', commits, dispatches extract_and_classify.delay(), - and returns {'document_id': str, 'status': 'processing'}. - - D-16: requires authenticated regular user. Asserts ownership — cross-user - classify returns 404 (not 403) to avoid information leakage (T-03-11). - T-07-10: ownership enforced here; IDOR returns 404 per STATE.md policy. - Placed in crud.py per D-08: same ownership-check pattern as get/patch/delete. - """ + """Re-queue a document for classification via Celery.""" request.state.current_user = current_user - try: - uid = uuid.UUID(doc_id) - except ValueError: - raise HTTPException(404, "Document not found") - - doc = await session.get(Document, uid) - if doc is None or doc.user_id != current_user.id: - raise HTTPException(404, "Document not found") + doc = await get_owned_document(session, doc_id, current_user.id) doc.status = "processing" await session.commit() diff --git a/backend/api/documents/shared.py b/backend/api/documents/shared.py index 5cd61b9..5c34de7 100644 --- a/backend/api/documents/shared.py +++ b/backend/api/documents/shared.py @@ -1,22 +1,16 @@ -"""Shared constants and Pydantic request models for the documents API package. - -CODE-08: Single definition of _CLOUD_PROVIDERS, UploadUrlRequest, and DocumentPatch. -These are imported by upload.py and crud.py — never duplicated. - -T-05-06-01: _CLOUD_PROVIDERS is an allowlist frozenset; target_backend validated - against it (never against user-supplied strings). -T-05-09-01: DocumentPatch fields declared explicitly — mass assignment prevented. -T-05-09-02: filename_no_path_separators validator preserved verbatim (path traversal - defense at the API boundary — D-11 analysis: stays in Pydantic model). -""" +"""Shared constants, request models, and access helpers for document routes.""" from __future__ import annotations import uuid from typing import Optional +from fastapi import HTTPException from pydantic import BaseModel, Field, field_validator +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from db.models import Document, Share -# Valid cloud backend slugs (T-05-06-01: validated against allowlist, not user-supplied string) _CLOUD_PROVIDERS = frozenset({"google_drive", "onedrive", "nextcloud", "webdav"}) @@ -26,14 +20,6 @@ class UploadUrlRequest(BaseModel): class DocumentPatch(BaseModel): - """Pydantic model for PATCH /api/documents/{doc_id}. - - Optional fields — model_fields_set distinguishes "not provided" from "set to null". - At least one field must be present in model_fields_set (enforced in the handler). - - T-05-09-01: explicit field declaration prevents mass assignment. - T-05-09-02: only filename and folder_id are accepted — no other fields can be set. - """ filename: Optional[str] = Field(None, min_length=1, max_length=255) folder_id: Optional[uuid.UUID] = None @@ -43,3 +29,50 @@ class DocumentPatch(BaseModel): if v is not None and ("/" in v or "\\" in v): raise ValueError("filename must not contain path separators") return v + + +def _document_not_found() -> HTTPException: + return HTTPException(status_code=404, detail="Document not found") + + +def parse_document_uuid(doc_id: str) -> uuid.UUID: + try: + return uuid.UUID(doc_id) + except ValueError: + raise _document_not_found() + + +async def get_owned_document( + session: AsyncSession, + doc_id: str, + user_id: uuid.UUID, +) -> Document: + uid = parse_document_uuid(doc_id) + doc = await session.get(Document, uid) + if doc is None or doc.user_id != user_id: + raise _document_not_found() + return doc + + +async def get_accessible_document( + session: AsyncSession, + doc_id: str, + user_id: uuid.UUID, +) -> tuple[Document, bool]: + uid = parse_document_uuid(doc_id) + doc = await session.get(Document, uid) + if doc is None: + raise _document_not_found() + + if doc.user_id == user_id: + return doc, False + + result = await session.execute( + select(Share).where( + Share.document_id == doc.id, + Share.recipient_id == user_id, + ) + ) + if result.scalar_one_or_none() is None: + raise _document_not_found() + return doc, True diff --git a/backend/api/documents/upload.py b/backend/api/documents/upload.py index 99060fe..f5785ed 100644 --- a/backend/api/documents/upload.py +++ b/backend/api/documents/upload.py @@ -1,32 +1,11 @@ -"""Document upload endpoints — presigned URL flow and direct cloud upload. - -Endpoints: - POST /upload-url — create pending Document row, return presigned PUT URL (D-05 step 1) - POST /upload — direct multipart upload supporting cloud backends (D-10, D-14, D-15) - POST /{doc_id}/confirm — stat MinIO for authoritative size, enforce quota atomically (D-05 step 3) - -Sub-router carries NO prefix — prefix="/api/documents" lives in __init__.py (D-04). - -Security: - T-03-04: object_key computed server-side using str(current_user.id) — never user-supplied. - T-03-05: size from backend.stat_object() — never from client. - T-03-06: atomic SQL UPDATE prevents concurrent over-quota uploads (STORE-03 SC2). - T-03-11: ownership assertion on confirm — cross-user access returns 404. - T-03-15: object_key prefix always the authenticated user's id. - T-05-06-01: target_backend validated against _CLOUD_PROVIDERS allowlist. - T-05-06-02: CloudConnectionError detail never includes provider error detail. -""" +"""Document upload endpoints.""" from __future__ import annotations import uuid from pathlib import Path -import structlog as _structlog - -_log = _structlog.get_logger(__name__) - -from fastapi import APIRouter, Depends, File, Form, HTTPException, Request, UploadFile, status -from sqlalchemy import text +from fastapi import APIRouter, Depends, File, Form, HTTPException, Request, UploadFile +from sqlalchemy import select, text from sqlalchemy.ext.asyncio import AsyncSession from config import settings @@ -36,24 +15,91 @@ from deps.db import get_db from deps.utils import get_client_ip from services.audit import write_audit_log from services.rate_limiting import account_limiter -from storage import get_storage_backend, get_storage_backend_for_document +from storage import get_storage_backend +from storage.cloud_backend_factory import build_cloud_backend from storage.cloud_utils import decrypt_credentials from storage.exceptions import CloudConnectionError from tasks.document_tasks import extract_and_classify -try: - from minio.error import S3Error -except ImportError: - S3Error = Exception # type: ignore[assignment,misc] - -from sqlalchemy import select - from api.documents.shared import UploadUrlRequest, _CLOUD_PROVIDERS router = APIRouter() -# ── POST /api/documents/upload-url ─────────────────────────────────────────── +def _new_minio_document(user_id: uuid.UUID, filename: str, content_type: str) -> Document: + doc_id = uuid.uuid4() + suffix = Path(filename).suffix.lower() + return Document( + id=doc_id, + user_id=user_id, + filename=filename, + content_type=content_type, + size_bytes=0, + storage_backend="minio", + status="pending", + object_key=f"{user_id}/{doc_id}/{uuid.uuid4()}{suffix}", + ) + + +async def _create_presigned_upload( + session: AsyncSession, + user_id: uuid.UUID, + filename: str, + content_type: str, +) -> dict: + doc = _new_minio_document(user_id, filename, content_type) + session.add(doc) + await session.commit() + + upload_url = await get_storage_backend().generate_presigned_put_url( + doc.object_key, expires_minutes=15 + ) + return {"upload_url": upload_url, "document_id": str(doc.id)} + + +async def _get_active_cloud_connection( + session: AsyncSession, + user_id: uuid.UUID, + provider: str, +) -> CloudConnection: + result = await session.execute( + select(CloudConnection).where( + CloudConnection.user_id == user_id, + CloudConnection.provider == provider, + CloudConnection.status == "ACTIVE", + ) + ) + conn = result.scalar_one_or_none() + if conn is None: + raise HTTPException( + status_code=404, + detail=f"No active {provider} connection found. Please connect in Settings.", + ) + return conn + + +def _decrypt_cloud_credentials(conn: CloudConnection, user_id: uuid.UUID) -> dict: + return decrypt_credentials(settings.cloud_creds_key.encode(), str(user_id), conn.credentials_enc) + + +async def _record_upload( + session: AsyncSession, + request: Request, + current_user: User, + doc: Document, + size_bytes: int, + storage_backend: str, +) -> None: + await write_audit_log( + session, + event_type="document.uploaded", + user_id=current_user.id, + actor_id=current_user.id, + resource_id=doc.id, + ip_address=get_client_ip(request) if request else None, + metadata_={"size_bytes": size_bytes, "storage_backend": storage_backend}, + ) + @router.post("/upload-url") @account_limiter.limit("100/minute") @@ -63,41 +109,15 @@ async def request_upload_url( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Create a pending Document row and return a presigned PUT URL. - - D-05 step 1: FastAPI creates a Document row (status='pending'), generates a - 15-minute presigned PUT URL, returns {upload_url, document_id}. - Quota is NOT reserved at this step — quota enforcement happens at /confirm. - - T-03-04: object_key is computed server-side using str(current_user.id); filename - stored in DB only (CLAUDE.md MinIO key schema). - T-03-15: object_key prefix is always the authenticated user's id — never user-supplied. - """ + """Create a pending Document row and return a presigned PUT URL.""" request.state.current_user = current_user - doc_id = uuid.uuid4() - suffix = Path(body.filename).suffix.lower() - object_key = f"{current_user.id}/{doc_id}/{uuid.uuid4()}{suffix}" - - doc = Document( - id=doc_id, - user_id=current_user.id, - filename=body.filename, - content_type=body.content_type, - size_bytes=0, - storage_backend="minio", - status="pending", - object_key=object_key, + return await _create_presigned_upload( + session, + current_user.id, + body.filename, + body.content_type, ) - session.add(doc) - await session.commit() - upload_url = await get_storage_backend().generate_presigned_put_url( - object_key, expires_minutes=15 - ) - return {"upload_url": upload_url, "document_id": str(doc_id)} - - -# ── POST /api/documents/upload ──────────────────────────────────────────────── @router.post("/upload") @account_limiter.limit("100/minute") @@ -109,47 +129,15 @@ async def upload_document( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Direct multipart upload endpoint supporting cloud backends (D-10, D-14, D-15). - - If target_backend == "minio": generates a presigned PUT URL (unchanged MinIO flow). - If target_backend in ("google_drive", "onedrive", "nextcloud", "webdav"): - 1. Reads file bytes from UploadFile - 2. Loads CloudConnection for current_user.id + target_backend; 404 if not found/not ACTIVE - 3. Decrypts credentials and instantiates the correct backend class - 4. Calls cloud_backend.put_object() to upload directly to the provider - 5. Creates Document with storage_backend=target_backend - 6. Returns {document_id, storage_backend} — no upload_url (cloud upload is synchronous) - - Cloud uploads do NOT use the atomic quota UPDATE — cloud files are not counted - against MinIO quota (D-11: separate backends; cloud storage quota is provider-side). - - Security: - T-05-06-01: target_backend validated against _CLOUD_PROVIDERS allowlist → 422 on invalid value - T-05-06-02: CloudConnectionError detail message never includes provider error detail - """ + """Direct multipart upload endpoint supporting cloud backends.""" request.state.current_user = current_user if target_backend == "minio": - doc_id = uuid.uuid4() - suffix = Path(file.filename or "file").suffix.lower() - object_key = f"{current_user.id}/{doc_id}/{uuid.uuid4()}{suffix}" - - doc = Document( - id=doc_id, - user_id=current_user.id, - filename=file.filename or "upload", - content_type=file.content_type or "application/octet-stream", - size_bytes=0, - storage_backend="minio", - status="pending", - object_key=object_key, + return await _create_presigned_upload( + session, + current_user.id, + file.filename or "upload", + file.content_type or "application/octet-stream", ) - session.add(doc) - await session.commit() - - upload_url = await get_storage_backend().generate_presigned_put_url( - object_key, expires_minutes=15 - ) - return {"upload_url": upload_url, "document_id": str(doc_id)} if target_backend not in _CLOUD_PROVIDERS: raise HTTPException( @@ -157,23 +145,8 @@ async def upload_document( detail=f"Invalid target_backend '{target_backend}'. Valid values: minio, {', '.join(sorted(_CLOUD_PROVIDERS))}", ) - # Load active CloudConnection for current user + provider (T-05-06-01: user-scoped query) - result = await session.execute( - select(CloudConnection).where( - CloudConnection.user_id == current_user.id, - CloudConnection.provider == target_backend, - CloudConnection.status == "ACTIVE", - ) - ) - conn = result.scalar_one_or_none() - if conn is None: - raise HTTPException( - status_code=404, - detail=f"No active {target_backend} connection found. Please connect in Settings.", - ) - - master_key = settings.cloud_creds_key.encode() - credentials = decrypt_credentials(master_key, str(current_user.id), conn.credentials_enc) + conn = await _get_active_cloud_connection(session, current_user.id, target_backend) + credentials = _decrypt_cloud_credentials(conn, current_user.id) file_bytes = await file.read() filename = file.filename or "upload" @@ -182,26 +155,7 @@ async def upload_document( doc_id = uuid.uuid4() - if target_backend == "google_drive": - from storage.google_drive_backend import GoogleDriveBackend # lazy import - cloud_backend = GoogleDriveBackend(credentials) - elif target_backend == "onedrive": - from storage.onedrive_backend import OneDriveBackend # lazy import - cloud_backend = OneDriveBackend(credentials) - elif target_backend == "nextcloud": - from storage.nextcloud_backend import NextcloudBackend # lazy import - cloud_backend = NextcloudBackend( - credentials["server_url"], - credentials["username"], - credentials["password"], - ) - elif target_backend == "webdav": - from storage.webdav_backend import WebDAVBackend # lazy import - cloud_backend = WebDAVBackend( - credentials["server_url"], - credentials["username"], - credentials["password"], - ) + cloud_backend = build_cloud_backend(target_backend, credentials) try: object_key = await cloud_backend.put_object( @@ -219,9 +173,9 @@ async def upload_document( detail="Cloud connection requires re-authentication. Please reconnect in Settings.", ) from exc - # Bust folder listing cache so the next GET /folders reflects the new file if cloud_folder_path: from services.cloud_cache import invalidate_provider_cache # lazy import + invalidate_provider_cache(str(current_user.id), target_backend) doc = Document( @@ -236,16 +190,7 @@ async def upload_document( ) session.add(doc) - _ip = get_client_ip(request) if request else None - await write_audit_log( - session, - event_type="document.uploaded", - user_id=current_user.id, - actor_id=current_user.id, - resource_id=doc.id, - ip_address=_ip, - metadata_={"size_bytes": len(file_bytes), "storage_backend": target_backend}, - ) + await _record_upload(session, request, current_user, doc, len(file_bytes), target_backend) await session.commit() extract_and_classify.delay(str(doc.id)) @@ -253,8 +198,6 @@ async def upload_document( return {"document_id": str(doc.id), "storage_backend": target_backend} -# ── POST /api/documents/{doc_id}/confirm ───────────────────────────────────── - @router.post("/{doc_id}/confirm") @account_limiter.limit("100/minute") async def confirm_upload( @@ -263,18 +206,7 @@ async def confirm_upload( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Confirm a presigned PUT upload: stat MinIO for size, enforce quota atomically. - - D-05 step 3: FastAPI reads authoritative file size from MinIO stat_object (never - from client), runs atomic quota UPDATE, sets status='uploaded', enqueues Celery task. - - Quota exceeded: HTTP 413 with {"used_bytes": N, "limit_bytes": M, "rejected_bytes": K} - Upload not found: HTTP 422 (presigned URL may have expired) - - T-03-05: size always comes from backend.stat_object(doc.object_key) — never client. - T-03-06: atomic SQL UPDATE prevents concurrent over-quota uploads (STORE-03 SC2). - T-03-11: ownership assertion — cross-user access returns 404 (D-16). - """ + """Confirm a presigned PUT upload and enforce quota.""" request.state.current_user = current_user try: uid = uuid.UUID(doc_id) @@ -285,7 +217,6 @@ async def confirm_upload( if doc is None or doc.user_id != current_user.id: raise HTTPException(status_code=404, detail="Document not found") - # Get authoritative file size from MinIO (T-03-05 — never trust client-supplied size) try: size = await get_storage_backend().stat_object(doc.object_key) except Exception as exc: @@ -300,7 +231,6 @@ async def confirm_upload( doc.size_bytes = size await session.flush() - # Atomic quota enforcement — user_id is always set post-migration (Plan 03-03+) result = await session.execute( text( "UPDATE quotas " @@ -323,7 +253,7 @@ async def confirm_upload( try: await get_storage_backend().delete_object(doc.object_key) except Exception: - pass # MinIO cleanup is best-effort; object TTL will eventually expire + pass await session.commit() raise HTTPException( status_code=413, @@ -337,17 +267,7 @@ async def confirm_upload( used_bytes = row.used_bytes doc.status = "uploaded" - # D-13: document uploaded event — size_bytes + storage_backend only, NO filename, NO extracted_text (T-04-07-02) - _ip = get_client_ip(request) - await write_audit_log( - session, - event_type="document.uploaded", - user_id=current_user.id, - actor_id=current_user.id, - resource_id=doc.id, - ip_address=_ip, - metadata_={"size_bytes": size, "storage_backend": "minio"}, - ) + await _record_upload(session, request, current_user, doc, size, "minio") await session.commit() extract_and_classify.delay(str(doc.id)) diff --git a/backend/api/folders.py b/backend/api/folders.py index e270bfa..19673b0 100644 --- a/backend/api/folders.py +++ b/backend/api/folders.py @@ -1,21 +1,4 @@ -""" -Folder API endpoints for DocuVault — Phase 4, Plan 03. - -Implements FOLD-01 through FOLD-05: - POST /api/folders — create folder (FOLD-01) - GET /api/folders — list top-level folders (FOLD-02) - GET /api/folders/{id} — get folder + breadcrumb (FOLD-02) - PATCH /api/folders/{id} — rename folder (FOLD-03) - DELETE /api/folders/{id} — delete folder (cascade) (FOLD-03) - PATCH /api/documents/{id}/folder — move document to folder (FOLD-04) - -Security invariants (all enforced): - T-04-03-01: get_regular_user on all endpoints (admin gets 403) - T-04-03-04: All folder IDOR paths return 404 not 403 - T-04-03-05: PATCH /api/documents/{id}/folder validates both doc and target folder ownership - T-04-03-06: IntegrityError (duplicate folder name) → 409 Conflict - T-04-03-03: Atomic quota decrement with CASE WHEN pattern (SQLite compat) -""" +"""Folder and document-folder organization endpoints.""" from __future__ import annotations import uuid @@ -23,11 +6,11 @@ from typing import Optional from fastapi import APIRouter, Depends, HTTPException, Request, status from pydantic import BaseModel -from sqlalchemy import select, text, func +from sqlalchemy import select, text from sqlalchemy.exc import IntegrityError, OperationalError from sqlalchemy.ext.asyncio import AsyncSession -from db.models import Document, Folder, Quota, Share, User +from db.models import Document, Folder, User from deps.auth import get_regular_user from deps.db import get_db from deps.utils import get_client_ip @@ -37,8 +20,6 @@ from storage import get_storage_backend router = APIRouter(prefix="/api/folders", tags=["folders"]) -# ── Request / response models ───────────────────────────────────────────────── - class FolderCreate(BaseModel): name: str parent_id: Optional[str] = None @@ -52,9 +33,6 @@ class DocumentMove(BaseModel): folder_id: Optional[str] = None - -# ── Helper: folder serialization ────────────────────────────────────────────── - def _folder_to_dict(folder: Folder) -> dict: return { "id": str(folder.id), @@ -65,8 +43,6 @@ def _folder_to_dict(folder: Folder) -> dict: } -# ── Helper: document serialization ──────────────────────────────────────────── - def _doc_to_dict(doc: Document) -> dict: return { "id": str(doc.id), @@ -81,7 +57,67 @@ def _doc_to_dict(doc: Document) -> dict: } -# ── POST /api/folders ───────────────────────────────────────────────────────── +def _parse_uuid(value: str, not_found_detail: str) -> uuid.UUID: + try: + return uuid.UUID(value) + except ValueError: + raise HTTPException(status_code=404, detail=not_found_detail) + + +async def _get_owned_folder( + session: AsyncSession, + folder_id: str | uuid.UUID, + user_id: uuid.UUID, + detail: str = "Folder not found", +) -> Folder: + uid = folder_id if isinstance(folder_id, uuid.UUID) else _parse_uuid(folder_id, detail) + folder = await session.get(Folder, uid) + if folder is None or folder.user_id != user_id: + raise HTTPException(status_code=404, detail=detail) + return folder + + +async def _get_owned_document( + session: AsyncSession, + doc_id: str, + user_id: uuid.UUID, +) -> Document: + uid = _parse_uuid(doc_id, "Document not found") + doc = await session.get(Document, uid) + if doc is None or doc.user_id != user_id: + raise HTTPException(status_code=404, detail="Document not found") + return doc + + +async def _ensure_unique_folder_name( + session: AsyncSession, + user_id: uuid.UUID, + name: str, + parent_id: Optional[uuid.UUID], + exclude_id: Optional[uuid.UUID] = None, +) -> None: + stmt = select(Folder).where( + Folder.user_id == user_id, + Folder.name == name, + Folder.parent_id == parent_id, + ) + if exclude_id is not None: + stmt = stmt.where(Folder.id != exclude_id) + + dup = await session.execute(stmt) + if dup.scalar_one_or_none() is not None: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail="A folder with that name already exists here", + ) + + +def _duplicate_folder_error() -> HTTPException: + return HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail="A folder with that name already exists here", + ) + @router.post("", status_code=status.HTTP_201_CREATED) async def create_folder( @@ -90,35 +126,14 @@ async def create_folder( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Create a new folder for the current user. - - FOLD-01: parent_id (if given) must belong to current_user — 404 otherwise. - Duplicate name under same parent returns 409 (T-04-03-06). - """ parent_uuid: Optional[uuid.UUID] = None if body.parent_id is not None: - try: - parent_uuid = uuid.UUID(body.parent_id) - except ValueError: - raise HTTPException(status_code=404, detail="Parent folder not found") - parent = await session.get(Folder, parent_uuid) - if parent is None or parent.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Parent folder not found") + parent = await _get_owned_folder( + session, body.parent_id, current_user.id, "Parent folder not found" + ) + parent_uuid = parent.id - # Explicit duplicate check — UniqueConstraint won't fire when parent_id IS NULL - # because SQL treats NULL as distinct from NULL in unique indexes. - dup = await session.execute( - select(Folder).where( - Folder.user_id == current_user.id, - Folder.name == body.name, - Folder.parent_id == parent_uuid, - ) - ) - if dup.scalar_one_or_none() is not None: - raise HTTPException( - status_code=status.HTTP_409_CONFLICT, - detail="A folder with that name already exists here", - ) + await _ensure_unique_folder_name(session, current_user.id, body.name, parent_uuid) folder = Folder( user_id=current_user.id, @@ -130,10 +145,7 @@ async def create_folder( await session.commit() except IntegrityError: await session.rollback() - raise HTTPException( - status_code=status.HTTP_409_CONFLICT, - detail="A folder with that name already exists here", - ) + raise _duplicate_folder_error() await write_audit_log( session, @@ -148,34 +160,20 @@ async def create_folder( return _folder_to_dict(folder) -# ── GET /api/folders ────────────────────────────────────────────────────────── - @router.get("") async def list_folders( parent_id: Optional[str] = None, session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """List the current user's folders at a given level. - - FOLD-02: when parent_id is omitted, returns root folders (parent_id IS NULL). - When parent_id is supplied, returns that folder's direct children (asserts ownership). - Each folder includes has_children so the frontend can hide expand arrows on leaf nodes. - """ parent_uuid: Optional[uuid.UUID] = None if parent_id is not None: - try: - parent_uuid = uuid.UUID(parent_id) - except ValueError: - raise HTTPException(status_code=404, detail="Parent folder not found") - parent_folder = await session.get(Folder, parent_uuid) - if parent_folder is None or parent_folder.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Parent folder not found") + parent_folder = await _get_owned_folder( + session, parent_id, current_user.id, "Parent folder not found" + ) + parent_uuid = parent_folder.id - if parent_uuid is None: - where_clause = Folder.parent_id.is_(None) - else: - where_clause = Folder.parent_id == parent_uuid + where_clause = Folder.parent_id.is_(None) if parent_uuid is None else Folder.parent_id == parent_uuid result = await session.execute( select(Folder) @@ -184,8 +182,6 @@ async def list_folders( ) folders = result.scalars().all() - # One extra query to know which of these folders have sub-folders. - # Allows the frontend to hide expand arrows on leaf nodes without extra round-trips. folder_ids = [f.id for f in folders] folders_with_children: set = set() if folder_ids: @@ -205,55 +201,34 @@ async def list_folders( } -# ── GET /api/folders/{folder_id} ────────────────────────────────────────────── - @router.get("/{folder_id}") async def get_folder( folder_id: str, session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Get folder metadata + breadcrumb array from root to this folder. + folder = await _get_owned_folder(session, folder_id, current_user.id) - FOLD-02 / FOLD-05: breadcrumb is built via iterative parent-walk in Python - (not WITH RECURSIVE) so it is compatible with both PostgreSQL and SQLite tests. - - Response: {id, name, parent_id, user_id, created_at, breadcrumb: [{id, name}, ...]} - The breadcrumb array is ordered root-first (root is breadcrumb[0]). - """ - try: - uid = uuid.UUID(folder_id) - except ValueError: - raise HTTPException(status_code=404, detail="Folder not found") - - folder = await session.get(Folder, uid) - if folder is None or folder.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Folder not found") - - # Build breadcrumb by walking up the parent chain iteratively. - # Ownership check on each ancestor ensures no cross-user traversal. crumbs = [{"id": str(folder.id), "name": folder.name}] current = folder visited: set = {current.id} while current.parent_id is not None: if current.parent_id in visited: - break # cycle guard (should not occur with proper constraints) + break parent = await session.get(Folder, current.parent_id) if parent is None or parent.user_id != current_user.id: - break # stop traversal if parent is inaccessible + break visited.add(parent.id) crumbs.append({"id": str(parent.id), "name": parent.name}) current = parent - crumbs.reverse() # root-first order + crumbs.reverse() response = _folder_to_dict(folder) response["breadcrumb"] = crumbs return response -# ── PATCH /api/folders/{folder_id} ─────────────────────────────────────────── - @router.patch("/{folder_id}") async def rename_folder( folder_id: str, @@ -262,46 +237,20 @@ async def rename_folder( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Rename a folder. - - FOLD-03: asserts ownership → 404 if not owner. - Duplicate name under same parent returns 409 (T-04-03-06). - """ - try: - uid = uuid.UUID(folder_id) - except ValueError: - raise HTTPException(status_code=404, detail="Folder not found") - - folder = await session.get(Folder, uid) - if folder is None or folder.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Folder not found") - + folder = await _get_owned_folder(session, folder_id, current_user.id) old_name = folder.name - # Explicit duplicate check — same NULL parent_id issue as create_folder. + if body.name != folder.name: - dup = await session.execute( - select(Folder).where( - Folder.user_id == current_user.id, - Folder.name == body.name, - Folder.parent_id == folder.parent_id, - Folder.id != folder.id, - ) + await _ensure_unique_folder_name( + session, current_user.id, body.name, folder.parent_id, folder.id ) - if dup.scalar_one_or_none() is not None: - raise HTTPException( - status_code=status.HTTP_409_CONFLICT, - detail="A folder with that name already exists here", - ) folder.name = body.name try: await session.commit() except IntegrityError: await session.rollback() - raise HTTPException( - status_code=status.HTTP_409_CONFLICT, - detail="A folder with that name already exists here", - ) + raise _duplicate_folder_error() await write_audit_log( session, @@ -316,8 +265,6 @@ async def rename_folder( return _folder_to_dict(folder) -# ── DELETE /api/folders/{folder_id} ────────────────────────────────────────── - @router.delete("/{folder_id}", status_code=status.HTTP_204_NO_CONTENT) async def delete_folder( folder_id: str, @@ -325,30 +272,9 @@ async def delete_folder( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Delete a folder and all of its contents (cascade). - - FOLD-03 + D-03: - - Collects all documents in the folder subtree using WITH RECURSIVE CTE - (wraps in try/except OperationalError for SQLite test compat; fallback - uses direct children only). - - Sums size_bytes, performs atomic quota decrement (CASE WHEN pattern for - SQLite compat — T-04-03-03). - - Deletes MinIO objects best-effort (per-object try/except — PATTERNS.md Pattern 2). - - Deletes all document rows and the folder row via ORM. - """ - try: - uid = uuid.UUID(folder_id) - except ValueError: - raise HTTPException(status_code=404, detail="Folder not found") - - folder = await session.get(Folder, uid) - if folder is None or folder.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Folder not found") - + folder = await _get_owned_folder(session, folder_id, current_user.id) folder_name = folder.name - # Collect all folder IDs in the subtree via WITH RECURSIVE CTE. - # Falls back to direct-children-only on SQLite (OperationalError on recursive CTE). subtree_folder_ids: list[str] = [] try: cte_result = await session.execute( @@ -361,17 +287,13 @@ async def delete_folder( " WHERE f.user_id = :uid" ") SELECT id FROM subtree" ), - # Use .hex (no dashes) — SQLite stores UUID as 32-char hex; PostgreSQL accepts both. {"root_id": folder.id.hex, "uid": current_user.id.hex}, ) subtree_folder_ids = [str(row[0]) for row in cte_result.fetchall()] except OperationalError: - # SQLite fallback: only direct children of this folder subtree_folder_ids = [str(folder.id)] - # Collect all documents in the subtree folder IDs if subtree_folder_ids: - # Build UUID list for IN query subtree_uuids = [] for fid in subtree_folder_ids: try: @@ -394,7 +316,6 @@ async def delete_folder( total_bytes = sum(d.size_bytes for d in docs) - # Atomic quota decrement (CASE WHEN for SQLite compat — never goes below 0) if total_bytes > 0: await session.execute( text( @@ -406,20 +327,16 @@ async def delete_folder( {"delta": total_bytes, "uid": current_user.id.hex}, ) - # Delete MinIO objects best-effort (per-object, never abort on failure) storage_backend = get_storage_backend() for doc in docs: try: await storage_backend.delete_object(doc.object_key) except Exception: - pass # best-effort; stale MinIO objects will be garbage-collected + pass - # Delete document rows for doc in docs: await session.delete(doc) - # Delete the folder (cascade will handle sub-folders in PostgreSQL; - # in SQLite test env we already collected and deleted all documents) await session.delete(folder) await session.commit() @@ -428,17 +345,12 @@ async def delete_folder( event_type="folder.deleted", user_id=current_user.id, actor_id=current_user.id, - resource_id=uid, + resource_id=folder.id, ip_address=get_client_ip(request), metadata_={"name": folder_name, "doc_count": len(docs)}, ) -# ── PATCH /api/documents/{doc_id}/folder ───────────────────────────────────── -# This endpoint lives in the folders router (not documents router) because it -# is logically a folder organisation operation. The URL prefix /api/documents -# is achieved via an explicit path on this APIRouter. FastAPI supports this. - document_move_router = APIRouter(prefix="/api/documents", tags=["folders"]) @@ -450,33 +362,12 @@ async def move_document( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ): - """Move a document to a different folder (or to root if folder_id is null). - - FOLD-04: - - Asserts document ownership → 404 if not owner. - - If folder_id given: asserts target folder ownership → 404 if not owner - (T-04-03-05: cross-user folder assignment blocked). - - Updates doc.folder_id and commits. - - Returns 200 with updated document dict. - """ - try: - doc_uid = uuid.UUID(doc_id) - except ValueError: - raise HTTPException(status_code=404, detail="Document not found") - - doc = await session.get(Document, doc_uid) - if doc is None or doc.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Document not found") + doc = await _get_owned_document(session, doc_id, current_user.id) target_folder_uuid: Optional[uuid.UUID] = None if body.folder_id is not None: - try: - target_folder_uuid = uuid.UUID(body.folder_id) - except ValueError: - raise HTTPException(status_code=404, detail="Folder not found") - target_folder = await session.get(Folder, target_folder_uuid) - if target_folder is None or target_folder.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Folder not found") + target_folder = await _get_owned_folder(session, body.folder_id, current_user.id) + target_folder_uuid = target_folder.id doc.folder_id = target_folder_uuid await session.commit() diff --git a/backend/api/shares.py b/backend/api/shares.py index 4bfe09f..72b058f 100644 --- a/backend/api/shares.py +++ b/backend/api/shares.py @@ -1,23 +1,7 @@ -""" -Sharing API for DocuVault — Phase 4, Plan 04-04. - -Implements SHARE-01 through SHARE-05: - POST /api/shares — grant share by recipient handle - GET /api/shares — list shares owned by current user for a document - GET /api/shares/received — virtual "Shared with me" folder (metadata only) - DELETE /api/shares/{share_id} — revoke share with IDOR protection - -Security invariants: - T-04-04-02: DELETE asserts share.owner_id == current_user.id → 404 on mismatch - T-04-04-03: GET /received returns metadata only — extracted_text is never included - T-04-04-04: No quota table is touched anywhere in this module - T-04-04-05: UniqueConstraint(document_id, recipient_id) → IntegrityError → 409 -""" +"""Document sharing endpoints.""" from __future__ import annotations import uuid -from typing import Optional - from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from pydantic import BaseModel, field_validator from sqlalchemy import select @@ -33,9 +17,6 @@ from services.audit import write_audit_log router = APIRouter(prefix="/api/shares", tags=["shares"]) -# ── Request models ──────────────────────────────────────────────────────────── - - class ShareCreate(BaseModel): document_id: str recipient_handle: str @@ -60,11 +41,47 @@ class SharePermissionPatch(BaseModel): return v -# ── Helpers ─────────────────────────────────────────────────────────────────── +def _parse_uuid(value: str, detail: str) -> uuid.UUID: + try: + return uuid.UUID(value) + except ValueError: + raise HTTPException(status_code=404, detail=detail) +async def _get_owned_document( + session: AsyncSession, + document_id: str, + owner_id: uuid.UUID, +) -> Document: + uid = _parse_uuid(document_id, "Document not found") + doc = await session.get(Document, uid) + if doc is None or doc.user_id != owner_id: + raise HTTPException(status_code=404, detail="Document not found") + return doc -# ── POST /api/shares ────────────────────────────────────────────────────────── + +async def _get_owned_share( + session: AsyncSession, + share_id: str, + owner_id: uuid.UUID, +) -> Share: + sid = _parse_uuid(share_id, "Share not found") + share = await session.get(Share, sid) + if share is None or share.owner_id != owner_id: + raise HTTPException(status_code=404, detail="Share not found") + return share + + +def _share_to_dict(share: Share, recipient: User) -> dict: + return { + "id": str(share.id), + "document_id": str(share.document_id), + "owner_id": str(share.owner_id), + "recipient_id": str(share.recipient_id), + "recipient_handle": recipient.handle, + "permission": share.permission, + "created_at": share.created_at.isoformat() if share.created_at else None, + } @router.post("", status_code=status.HTTP_201_CREATED) @@ -74,22 +91,7 @@ async def grant_share( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> dict: - """Grant document share to a user identified by their handle (SHARE-01, D-04). - - T-04-04-06: Only document owner can grant; 404 prevents ID enumeration. - T-04-04-01: get_regular_user ensures admins cannot invoke this endpoint. - T-04-04-05: Duplicate share → IntegrityError → 409 (no unbounded inserts). - """ - # Parse document_id as UUID (T-03-11 pattern) - try: - uid = uuid.UUID(body.document_id) - except ValueError: - raise HTTPException(status_code=404, detail="Document not found") - - # Ownership assertion — 404 prevents ID enumeration - doc = await session.get(Document, uid) - if doc is None or doc.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Document not found") + doc = await _get_owned_document(session, body.document_id, current_user.id) # Recipient lookup by exact handle (D-04) result = await session.execute( @@ -105,7 +107,7 @@ async def grant_share( # Create the share row share = Share( - document_id=uid, + document_id=doc.id, owner_id=current_user.id, recipient_id=recipient.id, permission=body.permission, @@ -126,25 +128,14 @@ async def grant_share( event_type="share.granted", user_id=current_user.id, actor_id=current_user.id, - resource_id=uid, + resource_id=doc.id, ip_address=get_client_ip(request), metadata_={"recipient_id": str(recipient.id)}, ) await session.commit() - return { - "id": str(share.id), - "document_id": str(share.document_id), - "owner_id": str(share.owner_id), - "recipient_id": str(share.recipient_id), - "recipient_handle": recipient.handle, - "permission": share.permission, - "created_at": share.created_at.isoformat() if share.created_at else None, - } - - -# ── GET /api/shares ─────────────────────────────────────────────────────────── + return _share_to_dict(share, recipient) @router.get("") @@ -153,46 +144,28 @@ async def list_shares( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> dict: - """List shares owned by current user for a specific document (SHARE-01, D-05). + doc = await _get_owned_document(session, document_id, current_user.id) - Only the document owner can list shares — 404 on mismatch or bad UUID. - """ - try: - uid = uuid.UUID(document_id) - except ValueError: - raise HTTPException(status_code=404, detail="Document not found") - - doc = await session.get(Document, uid) - if doc is None or doc.user_id != current_user.id: - raise HTTPException(status_code=404, detail="Document not found") - - # Join Share with User to get recipient handles stmt = ( select(Share, User) .join(User, User.id == Share.recipient_id) - .where(Share.document_id == uid) + .where(Share.document_id == doc.id) ) result = await session.execute(stmt) - rows = result.all() - - items = [] - for share, recipient in rows: - items.append( - { - "id": str(share.id), - "recipient_id": str(share.recipient_id), - "recipient_handle": recipient.handle, - "permission": share.permission, - "created_at": share.created_at.isoformat() - if share.created_at - else None, - } - ) + items = [ + { + "id": str(share.id), + "recipient_id": str(share.recipient_id), + "recipient_handle": recipient.handle, + "permission": share.permission, + "created_at": share.created_at.isoformat() if share.created_at else None, + } + for share, recipient in result.all() + ] return {"items": items} -# ── GET /api/shares/received ────────────────────────────────────────────────── # CRITICAL: This endpoint MUST be defined BEFORE DELETE /api/shares/{share_id}. # Defining it after would cause FastAPI to route GET /api/shares/received as # DELETE with share_id="received" (path parameter conflict). @@ -203,12 +176,6 @@ async def list_shared_with_me( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> dict: - """Return documents shared WITH the current user (virtual "Shared with me" folder — D-06). - - T-04-04-03: Only metadata is returned — extracted_text is never included. - T-04-04-04: No quota is modified. - Response: {items: [{id, filename, content_type, size_bytes, created_at, owner_handle}]} - """ stmt = ( select(Document, User) .join(Share, Share.document_id == Document.id) @@ -219,26 +186,21 @@ async def list_shared_with_me( result = await session.execute(stmt) rows = result.all() - items = [] - for doc, owner in rows: - # T-04-04-03: extracted_text is intentionally excluded here - items.append( - { - "id": str(doc.id), - "filename": doc.filename, - "content_type": doc.content_type, - "size_bytes": doc.size_bytes, - "created_at": doc.created_at.isoformat() if doc.created_at else None, - "owner_handle": owner.handle, - } - ) + items = [ + { + "id": str(doc.id), + "filename": doc.filename, + "content_type": doc.content_type, + "size_bytes": doc.size_bytes, + "created_at": doc.created_at.isoformat() if doc.created_at else None, + "owner_handle": owner.handle, + } + for doc, owner in rows + ] return {"items": items} -# ── PATCH /api/shares/{share_id} ───────────────────────────────────────────── - - @router.patch("/{share_id}", status_code=200) async def update_share_permission( share_id: str, @@ -247,20 +209,7 @@ async def update_share_permission( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> dict: - """Update the permission on an existing share (SHARE-03, D-09). - - T-06.2-02-01 IDOR protection: 404 on owner mismatch — mirrors revoke_share exactly. - T-06.2-02-02: SharePermissionPatch validator prevents arbitrary string passthrough. - """ - try: - sid = uuid.UUID(share_id) - except ValueError: - raise HTTPException(status_code=404, detail="Share not found") - - share = await session.get(Share, sid) - if share is None or share.owner_id != current_user.id: - raise HTTPException(status_code=404, detail="Share not found") - + share = await _get_owned_share(session, share_id, current_user.id) share.permission = body.permission await write_audit_log( @@ -277,9 +226,6 @@ async def update_share_permission( return {"id": str(share.id), "permission": share.permission} -# ── DELETE /api/shares/{share_id} ───────────────────────────────────────────── - - @router.delete("/{share_id}", status_code=status.HTTP_204_NO_CONTENT) async def revoke_share( share_id: str, @@ -287,27 +233,12 @@ async def revoke_share( session: AsyncSession = Depends(get_db), current_user: User = Depends(get_regular_user), ) -> None: - """Revoke a share. Only the share owner may revoke (SHARE-04, D-07). - - T-04-04-02 IDOR protection: asserts share.owner_id == current_user.id. - Returns 404 (not 403) on mismatch to prevent share ID enumeration. - """ - try: - sid = uuid.UUID(share_id) - except ValueError: - raise HTTPException(status_code=404, detail="Share not found") - - share = await session.get(Share, sid) - # CRITICAL IDOR check: 404 on mismatch (not 403) — prevents ID enumeration - if share is None or share.owner_id != current_user.id: - raise HTTPException(status_code=404, detail="Share not found") - + share = await _get_owned_share(session, share_id, current_user.id) document_id = share.document_id recipient_id = share.recipient_id await session.delete(share) - # Audit log before commit (D-14 — within the same transaction) await write_audit_log( session=session, event_type="share.revoked", diff --git a/backend/services/ai_config.py b/backend/services/ai_config.py index 652e947..fc46544 100644 --- a/backend/services/ai_config.py +++ b/backend/services/ai_config.py @@ -1,27 +1,7 @@ -""" -AI provider configuration service for DocuVault. - -Provides HKDF/Fernet encryption helpers, a provider config loader that reads -from the system_settings DB table, and a startup seed function that populates -the default provider row from env vars on first boot. - -Security design (D-05, T-07-02): - HKDF domain separation — the info bytes b"ai-provider-settings" differ from - b"cloud-credentials" used by storage/cloud_utils.py. Both use the same master - key (settings.cloud_creds_key) but produce DIFFERENT derived Fernet keys, so - a leaked cloud credential cannot decrypt an AI API key and vice versa. - -AlreadyFinalized warning (RESEARCH.md Pitfall 2 / .continue-here.md anti-pattern): - The cryptography library raises AlreadyFinalized if .derive() is called twice - on the same HKDF instance. _derive_ai_settings_key() creates a FRESH HKDF(...) - object on every call — never cache or reuse the HKDF object between calls. - -Pattern reference: storage/cloud_utils.py:_derive_fernet_key(). -""" +"""AI provider configuration and API-key encryption helpers.""" from __future__ import annotations import base64 -import logging from typing import Optional import structlog @@ -36,167 +16,78 @@ from config import settings logger = structlog.get_logger(__name__) +AI_SETTINGS_KEY_INFO = b"ai-provider-settings" -# ── HKDF key derivation ─────────────────────────────────────────────────────── def _derive_ai_settings_key(master_key: bytes, provider_id: str) -> Fernet: - """Derive a per-provider Fernet encryption key using HKDF-SHA256. - - Security notes: - - A FRESH HKDF instance is created on every call. The cryptography library - raises AlreadyFinalized if .derive() is called twice on the same instance. - Never cache or reuse the HKDF object (RESEARCH.md Pitfall 2). - - salt = provider_id.encode("utf-8") provides per-provider derivation - (deterministic: same provider → same key for encrypt/decrypt consistency). - - info = b"ai-provider-settings" provides domain separation from - b"cloud-credentials" — same master key, different derived keys. - A leaked cloud credential cannot decrypt an AI API key (T-07-02 mitigated). - - Args: - master_key: The CLOUD_CREDS_KEY env var as bytes. - provider_id: The provider slug, e.g. "openai", "anthropic" (used as HKDF salt). - - Returns: - A Fernet instance ready for encrypt/decrypt operations. - """ - # Create a FRESH HKDF instance — never cache (AlreadyFinalized guard) + """Derive a per-provider Fernet key with domain-separated HKDF-SHA256.""" hkdf = HKDF( algorithm=hashes.SHA256(), length=32, salt=provider_id.encode("utf-8"), - info=b"ai-provider-settings", # domain-separated from b"cloud-credentials" + info=AI_SETTINGS_KEY_INFO, ) raw_key: bytes = hkdf.derive(master_key) fernet_key = base64.urlsafe_b64encode(raw_key) return Fernet(fernet_key) -# ── Encryption helpers ──────────────────────────────────────────────────────── - def encrypt_api_key(master_key: bytes, provider_id: str, api_key: str) -> str: - """Encrypt a plaintext API key string to a Fernet token. - - The returned string is safe to store in system_settings.api_key_enc. - No JSON wrapping — the raw API key string is encrypted directly. - - Args: - master_key: The CLOUD_CREDS_KEY env var as bytes. - provider_id: The provider slug (used as HKDF salt for key derivation). - api_key: The plaintext API key, e.g. "sk-proj-...". - - Returns: - A URL-safe base64 Fernet token (str). - """ + """Encrypt a plaintext API key for storage in system_settings.api_key_enc.""" f = _derive_ai_settings_key(master_key, provider_id) return f.encrypt(api_key.encode("utf-8")).decode("utf-8") def decrypt_api_key(master_key: bytes, provider_id: str, api_key_enc: str) -> str: - """Decrypt a Fernet token back to the original plaintext API key. - - Args: - master_key: The CLOUD_CREDS_KEY env var as bytes. - provider_id: The provider slug (used as HKDF salt for key derivation). - api_key_enc: The Fernet token string from the database. - - Returns: - The original plaintext API key string. - """ + """Decrypt a stored API-key token.""" f = _derive_ai_settings_key(master_key, provider_id) return f.decrypt(api_key_enc.encode("utf-8")).decode("utf-8") -# ── Provider config loader ──────────────────────────────────────────────────── +def _config_from_settings_row(row) -> ProviderConfig: + api_key = "" + if row.api_key_enc: + master_key = settings.cloud_creds_key.encode("utf-8") + try: + api_key = decrypt_api_key(master_key, row.provider_id, row.api_key_enc) + except Exception: + logger.warning( + "ai_config: failed to decrypt api_key_enc", + provider_id=row.provider_id, + ) + + return ProviderConfig( + provider_id=row.provider_id, + api_key=api_key, + base_url=row.base_url, + model=row.model_name, + context_chars=row.context_chars, + ) + + +async def _load_settings_row(session: AsyncSession, *criteria): + from db.models import SystemSettings # local import avoids circular deps + + stmt = select(SystemSettings) + if criteria: + stmt = stmt.where(*criteria) + result = await session.execute(stmt) + return result.scalar_one_or_none() + async def load_provider_config(session: AsyncSession) -> Optional[ProviderConfig]: - """Load the active AI provider config from the system_settings table. - - Returns a ProviderConfig built from the row where is_active=True, - decrypting api_key_enc when present. Returns None when no active row exists. - - Args: - session: An open AsyncSession. - - Returns: - A ProviderConfig if an active row exists; None if the table is empty or - no row is marked active. - """ from db.models import SystemSettings # local import to avoid circular deps - stmt = select(SystemSettings).where(SystemSettings.is_active.is_(True)) - result = await session.execute(stmt) - row = result.scalar_one_or_none() + row = await _load_settings_row(session, SystemSettings.is_active.is_(True)) + return _config_from_settings_row(row) if row else None - if row is None: - return None - - # Decrypt API key if present - api_key = "" - if row.api_key_enc: - master_key = settings.cloud_creds_key.encode("utf-8") - try: - api_key = decrypt_api_key(master_key, row.provider_id, row.api_key_enc) - except Exception: - logger.warning( - "ai_config.load_provider_config: failed to decrypt api_key_enc", - provider_id=row.provider_id, - ) - - return ProviderConfig( - provider_id=row.provider_id, - api_key=api_key, - base_url=row.base_url, - model=row.model_name, - context_chars=row.context_chars, - ) - - -# ── Provider config loader by ID ───────────────────────────────────────────── async def load_provider_config_by_id(session: AsyncSession, provider_id: str) -> Optional[ProviderConfig]: - """Load an AI provider config from system_settings by provider_id. - - Unlike load_provider_config(), this function does NOT require is_active=True. - Used by the admin test-connection endpoint so admins can test inactive rows. - - Args: - session: An open AsyncSession. - provider_id: The provider slug to load (e.g. "openai", "lmstudio"). - - Returns: - A ProviderConfig if a row with the given provider_id exists; None otherwise. - """ from db.models import SystemSettings # local import to avoid circular deps - stmt = select(SystemSettings).where(SystemSettings.provider_id == provider_id) - result = await session.execute(stmt) - row = result.scalar_one_or_none() + row = await _load_settings_row(session, SystemSettings.provider_id == provider_id) + return _config_from_settings_row(row) if row else None - if row is None: - return None - - # Decrypt API key if present - api_key = "" - if row.api_key_enc: - master_key = settings.cloud_creds_key.encode("utf-8") - try: - api_key = decrypt_api_key(master_key, row.provider_id, row.api_key_enc) - except Exception: - logger.warning( - "ai_config.load_provider_config_by_id: failed to decrypt api_key_enc", - provider_id=row.provider_id, - ) - - return ProviderConfig( - provider_id=row.provider_id, - api_key=api_key, - base_url=row.base_url, - model=row.model_name, - context_chars=row.context_chars, - ) - - -# ── Provider ID validator (D-11 migration) ─────────────────────────────────── def validate_provider_id(v: str) -> str: """Service-layer provider_id validator; raises ValueError per CLAUDE.md service-vs-API rule.""" @@ -207,21 +98,8 @@ def validate_provider_id(v: str) -> str: return v -# ── Startup seed ────────────────────────────────────────────────────────────── - async def seed_system_settings_from_env(session: AsyncSession) -> None: - """Populate system_settings with a default provider row on first boot. - - Reads settings.default_ai_provider and settings.default_ai_model from config. - If no row exists for that provider_id, inserts one with is_active=True. - Never overwrites an existing row — idempotent across restarts (D-04). - - This function is called from the FastAPI lifespan in main.py after the - session factory is available. Caller is responsible for committing. - - Args: - session: An open AsyncSession. - """ + """Seed the default provider row from env settings if it does not exist.""" from db.models import SystemSettings # local import to avoid circular deps provider_id = settings.default_ai_provider @@ -232,13 +110,10 @@ async def seed_system_settings_from_env(session: AsyncSession) -> None: existing = result.scalar_one_or_none() if existing is not None: - # Row already exists — never overwrite (idempotent) return - # Use PROVIDER_DEFAULTS context_chars for the provider, fallback to 8000 context_chars = PROVIDER_DEFAULTS.get(provider_id, {}).get("context_chars", 8000) - # Insert default row with no API key (local providers like Ollama don't need one) row = SystemSettings( provider_id=provider_id, model_name=model_name, diff --git a/backend/services/auth.py b/backend/services/auth.py index 57c35f9..1738ac9 100644 --- a/backend/services/auth.py +++ b/backend/services/auth.py @@ -1,20 +1,4 @@ -""" -Auth service — pure Python, no FastAPI coupling. - -Handles: - - Password hashing (Argon2 via pwdlib) and constant-time verification (SEC-06) - - JWT access token creation/decode (PyJWT) - - Refresh token lifecycle with family revocation on reuse (AUTH-07, RFC 9700) - - TOTP provisioning and verification with replay prevention (AUTH-08) - - Backup code generation, storage, and constant-time verification (AUTH-02) - - HaveIBeenPwned k-anonymity check (SEC-03) - - Admin account bootstrap (D-04, D-05, D-06) - -Security invariants: - - All token/code comparisons use hmac.compare_digest (constant-time, SEC-06) - - No function raises HTTPException — callers (api/) map ValueError to HTTP errors - - refresh token family revocation enqueues send_security_alert_email.delay (AUTH-07) -""" +"""Authentication, token, TOTP, and account bootstrap services.""" from __future__ import annotations import base64 @@ -32,7 +16,7 @@ import jwt import pyotp from pwdlib import PasswordHash from pwdlib.hashers.argon2 import Argon2Hasher -from sqlalchemy import select, update +from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from config import settings @@ -45,8 +29,6 @@ _PASSWORD_DETAIL = ( logger = logging.getLogger(__name__) -# ── Password hashing ──────────────────────────────────────────────────────────── -# Single shared PasswordHash instance; Argon2 is the only enabled hasher. _pwd = PasswordHash([Argon2Hasher()]) @@ -82,10 +64,8 @@ def validate_password_strength(password: str) -> None: raise ValueError(_PASSWORD_DETAIL) -# ── JWT helpers ───────────────────────────────────────────────────────────────── - def _compute_fgp(user_agent: str, accept_lang: str) -> str: - """Return 16-char hex fingerprint binding a token to its client context (D-04).""" + """Return a short fingerprint binding a token to its client context.""" return hmac.new( settings.secret_key.encode(), (user_agent + "\x00" + accept_lang).encode(), @@ -93,6 +73,36 @@ def _compute_fgp(user_agent: str, accept_lang: str) -> str: ).hexdigest()[:16] +def _jwt_private_key() -> str: + return base64.b64decode(settings.jwt_private_key).decode() + + +def _jwt_public_key() -> str: + return base64.b64decode(settings.jwt_public_key).decode() + + +def _encode_jwt(payload: dict) -> str: + return jwt.encode(payload, _jwt_private_key(), algorithm="ES256") + + +def _decode_jwt( + token: str, + expected_type: str, + expired_message: str, + invalid_prefix: str, +) -> dict: + try: + payload = jwt.decode(token, _jwt_public_key(), algorithms=["ES256"]) + except jwt.ExpiredSignatureError as exc: + raise ValueError(expired_message) from exc + except jwt.PyJWTError as exc: + raise ValueError(f"{invalid_prefix}: {exc}") from exc + + if payload.get("typ") != expected_type: + raise ValueError(f"Token type mismatch: expected '{expected_type}'") + return payload + + def create_access_token( user_id: str, role: str, @@ -114,8 +124,7 @@ def create_access_token( "jti": str(uuid.uuid4()), "fgp": _compute_fgp(user_agent, accept_lang), } - private_pem = base64.b64decode(settings.jwt_private_key).decode() - return jwt.encode(payload, private_pem, algorithm="ES256") + return _encode_jwt(payload) def decode_access_token(token: str) -> dict: @@ -124,17 +133,7 @@ def decode_access_token(token: str) -> dict: Verifies: signature, expiry, and typ='access' (T-02-01 — prevents password-reset tokens from being used as access tokens). """ - try: - public_pem = base64.b64decode(settings.jwt_public_key).decode() - payload = jwt.decode(token, public_pem, algorithms=["ES256"]) - except jwt.ExpiredSignatureError as exc: - raise ValueError("Token has expired") from exc - except jwt.PyJWTError as exc: - raise ValueError(f"Invalid token: {exc}") from exc - - if payload.get("typ") != "access": - raise ValueError("Token type mismatch: expected 'access'") - return payload + return _decode_jwt(token, "access", "Token has expired", "Invalid token") def create_password_reset_token(user_id: str) -> str: @@ -149,8 +148,7 @@ def create_password_reset_token(user_id: str) -> str: "iat": now, "exp": now + timedelta(seconds=3600), } - private_pem = base64.b64decode(settings.jwt_private_key).decode() - return jwt.encode(payload, private_pem, algorithm="ES256") + return _encode_jwt(payload) def decode_password_reset_token(token: str) -> str: @@ -158,21 +156,15 @@ def decode_password_reset_token(token: str) -> str: Returns the user_id string. """ - try: - public_pem = base64.b64decode(settings.jwt_public_key).decode() - payload = jwt.decode(token, public_pem, algorithms=["ES256"]) - except jwt.ExpiredSignatureError as exc: - raise ValueError("Reset token has expired") from exc - except jwt.PyJWTError as exc: - raise ValueError(f"Invalid reset token: {exc}") from exc - - if payload.get("typ") != "password-reset": - raise ValueError("Token type mismatch: expected 'password-reset'") + payload = _decode_jwt( + token, + "password-reset", + "Reset token has expired", + "Invalid reset token", + ) return payload["sub"] -# ── Refresh token lifecycle ───────────────────────────────────────────────────── - async def create_refresh_token( session: AsyncSession, user_id: uuid.UUID, remember_me: bool = False ) -> str: @@ -181,8 +173,8 @@ async def create_refresh_token( The raw token is returned to the caller and set as an httpOnly cookie. Only the SHA-256 hash is stored in the database. - remember_me=False (default): TTL = refresh_token_expire_hours (16h short session, D-09, D-10) - remember_me=True: TTL = refresh_token_expire_days (30d extended session, D-11) + remember_me=False uses the short-session TTL; remember_me=True uses the + extended-session TTL. """ raw = secrets.token_urlsafe(32) token_hash = hashlib.sha256(raw.encode()).hexdigest() @@ -231,9 +223,7 @@ async def rotate_refresh_token( raise ValueError("Refresh token has expired") if row.revoked: - # T-02-02: reuse of revoked token — family revocation await revoke_all_refresh_tokens(session, row.user_id) - # Enqueue security alert email (deferred import to avoid circular dependency) from tasks.email_tasks import send_security_alert_email # noqa: PLC0415 send_security_alert_email.delay(str(row.user_id)) raise ValueError("token_family_revoked") @@ -273,8 +263,6 @@ async def revoke_all_refresh_tokens( return count -# ── TOTP provisioning ─────────────────────────────────────────────────────────── - async def provision_totp( session: AsyncSession, user_id: uuid.UUID ) -> tuple[str, str]: @@ -329,13 +317,10 @@ async def verify_totp( return True -# ── Backup codes ──────────────────────────────────────────────────────────────── - def generate_backup_codes(n: int = 10) -> list[str]: """Return *n* random 8-character uppercase alphanumeric backup codes.""" codes = [] for _ in range(n): - # secrets.token_hex(4) returns 8 hex chars; uppercase for readability code = secrets.token_hex(4).upper() codes.append(code) return codes @@ -349,7 +334,6 @@ async def store_backup_codes( Each code is stored as an Argon2 hash (never plaintext, T-02-03). Existing unused codes are deleted first to prevent accumulation. """ - # Delete existing unused codes result = await session.execute( select(BackupCode).where( BackupCode.user_id == user_id, @@ -360,7 +344,6 @@ async def store_backup_codes( await session.delete(row) await session.flush() - # Insert new hashed codes for code in codes: row = BackupCode( id=uuid.uuid4(), @@ -392,21 +375,17 @@ async def verify_backup_code( matched_row: Optional[BackupCode] = None for row in rows: - # Always call verify_password for ALL rows (constant-time: no early exit) if verify_password(code, row.code_hash): - matched_row = row # record match but keep iterating + matched_row = row # keep iterating for constant-time behavior if matched_row is None: return False - # Mark the matched code as used matched_row.used_at = datetime.now(timezone.utc) await session.commit() return True -# ── HaveIBeenPwned check ──────────────────────────────────────────────────────── - async def check_hibp(password: str) -> bool: """Check if password appears in HaveIBeenPwned using the k-anonymity model. @@ -442,19 +421,17 @@ async def check_hibp(password: str) -> bool: return False -# ── Admin bootstrap ───────────────────────────────────────────────────────────── - async def bootstrap_admin(session: AsyncSession) -> None: - """Idempotent admin account bootstrap (D-04, D-05, D-06). + """Idempotent admin account bootstrap. If the users table is empty AND settings.admin_email and settings.admin_password are both non-empty, creates an admin User row with a Quota row. - Logs a WARNING if env vars are missing (D-05) but never raises. + Logs a warning if env vars are missing but never raises. """ if not settings.admin_email or not settings.admin_password: logger.warning( - "Admin bootstrap skipped: ADMIN_EMAIL and/or ADMIN_PASSWORD not set (D-05). " + "Admin bootstrap skipped: ADMIN_EMAIL and/or ADMIN_PASSWORD not set. " "Set both env vars to seed the first admin account on startup." ) return @@ -462,7 +439,6 @@ async def bootstrap_admin(session: AsyncSession) -> None: # Check if any users exist result = await session.execute(select(User).limit(1)) if result.scalar_one_or_none() is not None: - # Users already exist — idempotent, skip (D-04) return admin_id = uuid.uuid4() @@ -477,7 +453,7 @@ async def bootstrap_admin(session: AsyncSession) -> None: ) quota = Quota( user_id=admin_id, - limit_bytes=104857600, # 100 MB default (D-06) + limit_bytes=104857600, used_bytes=0, ) session.add(admin_user) diff --git a/backend/services/storage.py b/backend/services/storage.py index f96d95f..017f87d 100644 --- a/backend/services/storage.py +++ b/backend/services/storage.py @@ -1,21 +1,4 @@ -""" -Async document/topic/settings storage service for DocuVault. - -This module replaces the legacy flat-file + filelock implementation with: - - Async SQLAlchemy ORM for document and topic persistence (PostgreSQL) - - MinIO SDK (via asyncio.to_thread) for binary object storage - -Public function names are PRESERVED from the old flat-file implementation so -that api/documents.py and api/topics.py can be updated in Plan 05 with minimal -changes (async def + await + session parameter). - -Phase 3 D-12: load_settings / save_settings / mask_api_key / settings_masked removed. -All AI config comes from DB (users.ai_provider / users.ai_model set by admin). - -D-05: Storage service layer switched to PostgreSQL + MinIO. -D-06: Object key schema: {user_id}/{document_id}/{uuid4()}{ext} — human filename in DB only. -D-03: documents.user_id is None (nullable) in Phase 1 — no auth system yet. -""" +"""Async document and topic storage helpers.""" from __future__ import annotations import sys @@ -31,26 +14,17 @@ from db.models import Document, DocumentTopic, Topic from storage import get_storage_backend -# ── Lazy singleton storage backend ──────────────────────────────────────────── - _storage = None def _backend(): - """Return the lazily-instantiated StorageBackend singleton. - - Mirrors the module-level singleton behaviour of the old filelock objects so - the MinIO client is created once per process, not once per request. - """ + """Return the lazily-instantiated StorageBackend singleton.""" global _storage _storage = _storage or get_storage_backend() return _storage -# ── Private helpers ──────────────────────────────────────────────────────────── - def _doc_to_dict(doc: Document, topic_names: list) -> dict: - """Convert a Document ORM row + resolved topic names to the legacy dict shape.""" return { "id": str(doc.id), "original_name": doc.filename, @@ -65,8 +39,22 @@ def _doc_to_dict(doc: Document, topic_names: list) -> dict: } +def _topic_to_dict(topic: Topic) -> dict: + return { + "id": str(topic.id), + "name": topic.name, + "description": topic.description, + "color": topic.color, + } + + +def _topic_namespace_filter(name: str, user_id: Optional[uuid.UUID]): + criteria = [sql_func.lower(Topic.name) == name.lower()] + criteria.append(Topic.user_id.is_(None) if user_id is None else Topic.user_id == user_id) + return criteria + + async def _load_topic_names(session: AsyncSession, doc_id: uuid.UUID) -> list: - """Return the list of topic names for a given document UUID.""" q = await session.execute( select(Topic.name) .join(DocumentTopic, DocumentTopic.topic_id == Topic.id) @@ -75,9 +63,6 @@ async def _load_topic_names(session: AsyncSession, doc_id: uuid.UUID) -> list: return [row[0] for row in q] -# ── Documents ───────────────────────────────────────────────────────────────── - - async def save_metadata(session: AsyncSession, meta: dict) -> None: """Update a Document row from the legacy metadata dict shape. @@ -120,10 +105,7 @@ async def get_metadata(session: AsyncSession, doc_id: str) -> Optional[dict]: async def list_metadata( session: AsyncSession, user_id: uuid.UUID, topic: Optional[str] = None ) -> list: - """Return a list of metadata dicts for a specific user, optionally filtered by topic name. - - D-16: always filters by user_id — a user can only see their own documents. - """ + """Return metadata dicts for a user's documents, optionally filtered by topic.""" stmt = select(Document).where(Document.user_id == user_id).order_by(Document.created_at.desc()) if topic is not None: stmt = ( @@ -176,10 +158,6 @@ async def delete_document( print(f"[storage] WARNING: MinIO delete_object failed for {doc.object_key!r}: {exc}", file=sys.stderr) if not skip_quota: - # Atomic quota decrement (STORE-06, D-07). - # user_id is always set post-migration (Plan 03-03+) — guard removed. - # Use CASE WHEN instead of GREATEST() for SQLite compatibility - # (PostgreSQL supports both; SQLite lacks the GREATEST scalar function). await session.execute( text( "UPDATE quotas " @@ -212,12 +190,10 @@ async def update_document_topics( if doc is None: return None - # Remove all existing associations await session.execute( delete(DocumentTopic).where(DocumentTopic.document_id == uid) ) - # Re-insert, deduplicating by name seen: set = set() for name in topics: if name in seen: @@ -257,15 +233,10 @@ async def remove_topic_from_all_documents( return result.rowcount -# ── Topics ──────────────────────────────────────────────────────────────────── - async def load_topics(session: AsyncSession) -> list: """Return all topics ordered by name.""" q = await session.execute(select(Topic).order_by(Topic.name)) - return [ - {"id": str(t.id), "name": t.name, "description": t.description, "color": t.color} - for t in q.scalars() - ] + return [_topic_to_dict(t) for t in q.scalars()] async def load_topics_for_user(session: AsyncSession, user_id: uuid.UUID) -> list: @@ -280,17 +251,11 @@ async def load_topics_for_user(session: AsyncSession, user_id: uuid.UUID) -> lis or_(Topic.user_id == user_id, Topic.user_id.is_(None)) ).order_by(Topic.name) ) - return [ - {"id": str(t.id), "name": t.name, "description": t.description, "color": t.color} - for t in q.scalars() - ] + return [_topic_to_dict(t) for t in q.scalars()] async def save_topics(session: AsyncSession, topics: list) -> None: - """Idempotent bulk replace — delete all Topic rows then insert the list. - - # legacy: not used by current endpoints; preserved for API compatibility. - """ + """Idempotent bulk replace; kept for compatibility with older callers.""" await session.execute(delete(Topic)) for t in topics: session.add( @@ -313,7 +278,7 @@ async def get_topic(session: AsyncSession, topic_id: str) -> Optional[dict]: t = await session.get(Topic, uid) if t is None: return None - return {"id": str(t.id), "name": t.name, "description": t.description, "color": t.color} + return _topic_to_dict(t) async def create_topic( @@ -323,51 +288,16 @@ async def create_topic( color: str = "#6366f1", user_id: Optional[uuid.UUID] = None, ) -> dict: - """Create a topic, or return the existing one (case-insensitive, namespace-scoped dedup). - - D-08: user_id=None creates a system topic (visible to all users). - D-08: user_id= creates a per-user topic (visible only to that user). - - Deduplication is scoped by user_id namespace: - - System topics (user_id=None) dedup against other system topics only - - Per-user topics dedup within that user's namespace only - This allows "Finance" to exist as both a system topic and a per-user topic. - - SQLite note: Uses a branching approach instead of IS NOT DISTINCT FROM - (SQLite doesn't support that PostgreSQL construct for NULL comparison). - """ - if user_id is None: - q = await session.execute( - select(Topic).where( - sql_func.lower(Topic.name) == name.lower(), - Topic.user_id.is_(None), - ) - ) - else: - q = await session.execute( - select(Topic).where( - sql_func.lower(Topic.name) == name.lower(), - Topic.user_id == user_id, - ) - ) + """Create a topic, or return an existing case-insensitive namespace match.""" + q = await session.execute(select(Topic).where(*_topic_namespace_filter(name, user_id))) existing = q.scalars().first() if existing is not None: - return { - "id": str(existing.id), - "name": existing.name, - "description": existing.description, - "color": existing.color, - } + return _topic_to_dict(existing) topic = Topic(name=name, description=description, color=color, user_id=user_id) session.add(topic) await session.commit() - return { - "id": str(topic.id), - "name": topic.name, - "description": topic.description, - "color": topic.color, - } + return _topic_to_dict(topic) async def update_topic( @@ -392,7 +322,7 @@ async def update_topic( if color is not None: t.color = color await session.commit() - return {"id": str(t.id), "name": t.name, "description": t.description, "color": t.color} + return _topic_to_dict(t) async def delete_topic(session: AsyncSession, topic_id: str) -> Optional[str]: @@ -408,7 +338,7 @@ async def delete_topic(session: AsyncSession, topic_id: str) -> Optional[str]: if t is None: return None name = t.name - await session.delete(t) # ondelete="CASCADE" removes DocumentTopic rows + await session.delete(t) await session.commit() return name @@ -437,8 +367,6 @@ async def topic_doc_counts( return {name: count for name, count in q} -# ── Public surface ───────────────────────────────────────────────────────────── - __all__ = [ "save_metadata", "get_metadata", diff --git a/backend/storage/cloud_backend_factory.py b/backend/storage/cloud_backend_factory.py new file mode 100644 index 0000000..e6458e0 --- /dev/null +++ b/backend/storage/cloud_backend_factory.py @@ -0,0 +1,33 @@ +"""Factory for user-scoped cloud storage backends.""" + + +def build_cloud_backend(provider: str, credentials: dict): + if provider == "google_drive": + from storage.google_drive_backend import GoogleDriveBackend # lazy import + + return GoogleDriveBackend(credentials) + + if provider == "onedrive": + from storage.onedrive_backend import OneDriveBackend # lazy import + + return OneDriveBackend(credentials) + + if provider == "nextcloud": + from storage.nextcloud_backend import NextcloudBackend # lazy import + + return NextcloudBackend( + credentials["server_url"], + credentials["username"], + credentials["password"], + ) + + if provider == "webdav": + from storage.webdav_backend import WebDAVBackend # lazy import + + return WebDAVBackend( + credentials["server_url"], + credentials["username"], + credentials["password"], + ) + + raise ValueError(f"Unknown provider: {provider}") diff --git a/backend/storage/exceptions.py b/backend/storage/exceptions.py index 18f0bf5..8fe8524 100644 --- a/backend/storage/exceptions.py +++ b/backend/storage/exceptions.py @@ -6,12 +6,8 @@ class CloudConnectionError(Exception): """Raised when a cloud provider signals a non-retryable connection problem. Attributes: - reason: "token_expired" — access token expired; API layer can refresh and retry. + reason: "token_expired" — access token expired. "invalid_grant" — refresh token revoked; user must reconnect. - - The backend never updates the DB. The API layer (_call_cloud_op in cloud.py) - catches this exception, performs the DB state transition, and decides whether - to retry or surface a 503 to the client (B2 design, D-05/D-06). """ def __init__(self, msg: str = "", *, reason: str = "") -> None: diff --git a/backend/storage/google_drive_backend.py b/backend/storage/google_drive_backend.py index eb9f13e..1419e05 100644 --- a/backend/storage/google_drive_backend.py +++ b/backend/storage/google_drive_backend.py @@ -11,10 +11,8 @@ Design notes: - D-14: presigned_get_url and generate_presigned_put_url raise NotImplementedError. The API upload endpoint detects cloud backends and uses the direct put_object() path instead. - - B2 design: This backend is STATELESS. It raises CloudConnectionError but - does NOT update the DB or CloudConnection objects. DB state transitions - (e.g., REQUIRES_REAUTH) are handled by the _call_cloud_op() helper in - cloud.py (Plan 05-05), which has the DB session. + - This backend is stateless. It raises CloudConnectionError but does not + update the DB or CloudConnection objects. - Token key format stored in credentials dict: access_token — current OAuth bearer token refresh_token — long-lived refresh token diff --git a/backend/storage/onedrive_backend.py b/backend/storage/onedrive_backend.py index 27465eb..01300e6 100644 --- a/backend/storage/onedrive_backend.py +++ b/backend/storage/onedrive_backend.py @@ -10,9 +10,8 @@ Design notes: already async and awaited directly. - CloudConnectionError is imported from google_drive_backend (shared type). This keeps the exception hierarchy unified across all cloud backends. - - B2 design: This backend is STATELESS. It raises CloudConnectionError but - does NOT update the DB or CloudConnection objects. DB state transitions - (e.g., REQUIRES_REAUTH) are handled by _call_cloud_op() in cloud.py. + - This backend is stateless. It raises CloudConnectionError but does not + update the DB or CloudConnection objects. - _ensure_valid_token() checks expiry before each API call and calls _refresh_token() if the token is within 60 seconds of expiry. If the refresh returns None (invalid_grant), CloudConnectionError is raised. diff --git a/backend/tests/test_cloud.py b/backend/tests/test_cloud.py index bee5036..ee279f7 100644 --- a/backend/tests/test_cloud.py +++ b/backend/tests/test_cloud.py @@ -496,15 +496,8 @@ async def test_invalid_grant_sets_requires_reauth( assert "re-authentication" in resp.json().get("detail", "").lower() or \ "reconnect" in resp.json().get("detail", "").lower() - # The test checks the 503 response — the DB REQUIRES_REAUTH state transition is - # handled by _call_cloud_op in cloud.py (which is invoked by the real backend flow). - # In this test we monkeypatched get_storage_backend_for_document directly, so - # documents.py's except CloudConnectionError block fires, returning 503. - # The REQUIRES_REAUTH DB state is written by _call_cloud_op, not by documents.py. - # For this test we verify: (1) 503 returned, and (2) the conn was not status="ACTIVE" - # after the call — since the monkeypatch bypasses _call_cloud_op, we re-check conn status. - # The 503 path in documents.py does NOT update conn.status — that is _call_cloud_op's job. - # We verify the HTTP contract here; the DB transition is covered by the cloud.py unit tests. + # This test verifies the document content endpoint's HTTP contract when the + # storage layer reports an invalid cloud grant. # ── CLOUD-06: Disconnect / credential deletion ──────────────────────────────── diff --git a/docker-compose.yml b/docker-compose.yml index 3ba7f4f..f8dcc43 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -37,6 +37,55 @@ services: retries: 5 start_period: 15s + minio-init: + image: minio/mc:latest + depends_on: + minio: + condition: service_healthy + environment: + MINIO_ROOT_USER: ${MINIO_ROOT_USER} + MINIO_ROOT_PASSWORD: ${MINIO_ROOT_PASSWORD} + MINIO_ACCESS_KEY: ${MINIO_ACCESS_KEY} + MINIO_SECRET_KEY: ${MINIO_SECRET_KEY} + MINIO_BUCKET: ${MINIO_BUCKET} + entrypoint: ["/bin/sh", "-c"] + command: + - | + set -eu + mc alias set local http://minio:9000 "$$MINIO_ROOT_USER" "$$MINIO_ROOT_PASSWORD" + mc mb --ignore-existing "local/$$MINIO_BUCKET" + mc admin user add local "$$MINIO_ACCESS_KEY" "$$MINIO_SECRET_KEY" || true + cat > /tmp/docuvault-policy.json < { + if (value) searchParams.set(key, value) + }) + return searchParams +} + +async function downloadCsv(url, filename, errorPrefix) { + const res = await fetchWithRetry(url) + if (!res.ok) throw new Error(`${errorPrefix}: ${res.status}`) + + const blob = new Blob([await res.text()], { type: 'text/csv' }) + const objectUrl = URL.createObjectURL(blob) + const a = document.createElement('a') + a.href = objectUrl + a.download = filename + document.body.appendChild(a) + a.click() + document.body.removeChild(a) + setTimeout(() => URL.revokeObjectURL(objectUrl), 1000) +} export function adminListUsers() { return request('/api/admin/users') } export function adminCreateUser(body) { - return request('/api/admin/users', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify(body), - }) + return jsonRequest('/api/admin/users', 'POST', body) } export function adminDeactivateUser(id) { - return request(`/api/admin/users/${id}/status`, { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ is_active: false }), - }) + return jsonRequest(`/api/admin/users/${id}/status`, 'PATCH', { is_active: false }) } export function adminReactivateUser(id) { - return request(`/api/admin/users/${id}/status`, { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ is_active: true }), - }) + return jsonRequest(`/api/admin/users/${id}/status`, 'PATCH', { is_active: true }) } export function adminResetUserPassword(id) { @@ -37,27 +48,18 @@ export function adminGetUserQuota(id) { } export function adminUpdateQuota(id, limitBytes) { - return request(`/api/admin/users/${id}/quota`, { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ limit_bytes: limitBytes }), - }) + return jsonRequest(`/api/admin/users/${id}/quota`, 'PATCH', { limit_bytes: limitBytes }) } export function adminUpdateAiConfig(id, provider, model) { - return request(`/api/admin/users/${id}/ai-config`, { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ ai_provider: provider, ai_model: model }), + return jsonRequest(`/api/admin/users/${id}/ai-config`, 'PATCH', { + ai_provider: provider, + ai_model: model, }) } export function adminDeleteUser(id, adminPassword) { - return request(`/api/admin/users/${id}`, { - method: 'DELETE', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ admin_password: adminPassword }), - }) + return jsonRequest(`/api/admin/users/${id}`, 'DELETE', { admin_password: adminPassword }) } export function getAiConfig() { @@ -65,11 +67,7 @@ export function getAiConfig() { } export function saveAiConfig(body) { - return request('/api/admin/ai-config', { - method: 'PUT', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify(body), - }) + return jsonRequest('/api/admin/ai-config', 'PUT', body) } export function testAiConnection(providerId) { @@ -87,38 +85,23 @@ export function getAiModels(providerId) { } export function adminListAuditLog({ start, end, user_handle, event_type, page = 1, per_page = 50 } = {}) { - const params = new URLSearchParams() - if (start) params.set('start', start) - if (end) params.set('end', end) - if (user_handle) params.set('user_handle', user_handle) - if (event_type) params.set('event_type', event_type) + const params = withPresentParams({}, { start, end, user_handle, event_type }) params.set('page', page) params.set('per_page', per_page) return request(`/api/admin/audit-log?${params}`) } -// Unlike window.location.href, fetchWithRetry sends the Authorization Bearer header so -// the endpoint can authenticate. Must NOT call res.json() — CSV is text/csv. export async function adminExportAuditLogCsv(params = {}) { - const searchParams = new URLSearchParams({ format: 'csv' }) - if (params.start) searchParams.set('start', params.start) - if (params.end) searchParams.set('end', params.end) - if (params.user_handle) searchParams.set('user_handle', params.user_handle) - if (params.event_type) searchParams.set('event_type', params.event_type) - - const res = await fetchWithRetry(`/api/admin/audit-log/export?${searchParams}`) - if (!res.ok) throw new Error(`Export failed: ${res.status}`) - - const text = await res.text() - const blob = new Blob([text], { type: 'text/csv' }) - const url = URL.createObjectURL(blob) - const a = document.createElement('a') - a.href = url - a.download = 'audit-export.csv' - document.body.appendChild(a) - a.click() - document.body.removeChild(a) - setTimeout(() => URL.revokeObjectURL(url), 1000) + const searchParams = withPresentParams( + { format: 'csv' }, + { + start: params.start, + end: params.end, + user_handle: params.user_handle, + event_type: params.event_type, + } + ) + return downloadCsv(`/api/admin/audit-log/export?${searchParams}`, 'audit-export.csv', 'Export failed') } export function adminListDailyExports() { @@ -129,19 +112,6 @@ export async function getAdminOverview() { return request('/api/admin/overview') } -// Uses fetchWithRetry() to send the Authorization Bearer header (D-17, T-06.2-04-03). export async function adminDownloadDailyExport(date) { - const res = await fetchWithRetry(`/api/admin/audit-log/daily-exports/${date}`) - if (!res.ok) throw new Error(`Download failed: ${res.status}`) - - const text = await res.text() - const blob = new Blob([text], { type: 'text/csv' }) - const url = URL.createObjectURL(blob) - const a = document.createElement('a') - a.href = url - a.download = `audit-${date}.csv` - document.body.appendChild(a) - a.click() - document.body.removeChild(a) - setTimeout(() => URL.revokeObjectURL(url), 1000) + return downloadCsv(`/api/admin/audit-log/daily-exports/${date}`, `audit-${date}.csv`, 'Download failed') } diff --git a/frontend/src/api/auth.js b/frontend/src/api/auth.js index 5d1e54d..c43c9a1 100644 --- a/frontend/src/api/auth.js +++ b/frontend/src/api/auth.js @@ -5,22 +5,14 @@ * TotpEnrollment.vue, PasswordResetView.vue */ -import { request } from './utils.js' +import { request, jsonRequest } from './utils.js' export function login(body) { - return request('/api/auth/login', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify(body), - }) + return jsonRequest('/api/auth/login', 'POST', body) } export function register(body) { - return request('/api/auth/register', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify(body), - }) + return jsonRequest('/api/auth/register', 'POST', body) } export function refreshToken() { @@ -41,11 +33,7 @@ export function getMe() { } export function changePassword(body) { - return request('/api/auth/change-password', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify(body), - }) + return jsonRequest('/api/auth/change-password', 'POST', body) } export function totpSetup() { @@ -53,11 +41,7 @@ export function totpSetup() { } export function totpEnable(code) { - return request('/api/auth/totp/enable', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ code }), - }) + return jsonRequest('/api/auth/totp/enable', 'POST', { code }) } export function totpDisable() { @@ -65,18 +49,13 @@ export function totpDisable() { } export function passwordResetRequest(email) { - return request('/api/auth/password-reset', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ email }), - }) + return jsonRequest('/api/auth/password-reset', 'POST', { email }) } export function passwordResetConfirm(token, newPassword) { - return request('/api/auth/password-reset/confirm', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ token, new_password: newPassword }), + return jsonRequest('/api/auth/password-reset/confirm', 'POST', { + token, + new_password: newPassword, }) } @@ -85,11 +64,7 @@ export function getMyPreferences() { } export function updateMyPreferences(payload) { - return request('/api/auth/me/preferences', { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify(payload), - }) + return jsonRequest('/api/auth/me/preferences', 'PATCH', payload) } export function getMyQuota() { diff --git a/frontend/src/api/cloud.js b/frontend/src/api/cloud.js index 1fc19de..7cf83d2 100644 --- a/frontend/src/api/cloud.js +++ b/frontend/src/api/cloud.js @@ -6,7 +6,7 @@ * CloudFolderTreeItem.vue */ -import { request } from './utils.js' +import { request, jsonRequest } from './utils.js' export function listCloudConnections() { return request('/api/cloud/connections') @@ -17,19 +17,16 @@ export function disconnectCloud(id) { } export function connectWebDav(provider, serverUrl, username, password) { - return request('/api/cloud/connections/webdav', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ provider, server_url: serverUrl, username, password }), + return jsonRequest('/api/cloud/connections/webdav', 'POST', { + provider, + server_url: serverUrl, + username, + password, }) } export function updateDefaultStorage(backend) { - return request('/api/users/me/default-storage', { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ backend }), - }) + return jsonRequest('/api/users/me/default-storage', 'PATCH', { backend }) } export function getCloudFolders(provider, folderId) { diff --git a/frontend/src/api/documents.js b/frontend/src/api/documents.js index 7226058..418cb91 100644 --- a/frontend/src/api/documents.js +++ b/frontend/src/api/documents.js @@ -5,7 +5,7 @@ * DocumentPreviewModal.vue, FileManagerView.vue, CloudFolderView.vue */ -import { request, fetchWithRetry } from './utils.js' +import { request, fetchWithRetry, jsonRequest } from './utils.js' export function listDocuments({ topic, page = 1, perPage = 20, folderId = null, q = null, sort = null, order = null } = {}) { const params = new URLSearchParams({ page, per_page: perPage }) @@ -31,18 +31,13 @@ export function deleteDocumentRemoveOnly(id) { } export function classifyDocument(id, topics = null) { - return request(`/api/documents/${id}/classify`, { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify(topics ? { topics } : {}), - }) + return jsonRequest(`/api/documents/${id}/classify`, 'POST', topics ? { topics } : {}) } export function getUploadUrl(filename, contentType) { - return request('/api/documents/upload-url', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ filename, content_type: contentType }), + return jsonRequest('/api/documents/upload-url', 'POST', { + filename, + content_type: contentType, }) } diff --git a/frontend/src/api/folders.js b/frontend/src/api/folders.js index 7f8906b..75da4d0 100644 --- a/frontend/src/api/folders.js +++ b/frontend/src/api/folders.js @@ -4,7 +4,7 @@ * Consumers: stores/folders.js, FolderTreeItem.vue, AppSidebar.vue */ -import { request } from './utils.js' +import { request, jsonRequest } from './utils.js' export function listFolders(parentId = null) { const params = new URLSearchParams() @@ -14,11 +14,7 @@ export function listFolders(parentId = null) { } export function createFolder(name, parentId = null) { - return request('/api/folders', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ name, parent_id: parentId || null }), - }) + return jsonRequest('/api/folders', 'POST', { name, parent_id: parentId || null }) } export function getFolder(folderId) { @@ -26,11 +22,7 @@ export function getFolder(folderId) { } export function renameFolder(folderId, name) { - return request(`/api/folders/${folderId}`, { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ name }), - }) + return jsonRequest(`/api/folders/${folderId}`, 'PATCH', { name }) } export function deleteFolder(folderId) { @@ -38,9 +30,5 @@ export function deleteFolder(folderId) { } export function moveDocument(docId, folderId) { - return request(`/api/documents/${docId}/folder`, { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ folder_id: folderId || null }), - }) + return jsonRequest(`/api/documents/${docId}/folder`, 'PATCH', { folder_id: folderId || null }) } diff --git a/frontend/src/api/shares.js b/frontend/src/api/shares.js index 14a3d42..6f902b1 100644 --- a/frontend/src/api/shares.js +++ b/frontend/src/api/shares.js @@ -4,22 +4,18 @@ * Consumers: SharedView.vue, ShareModal.vue */ -import { request } from './utils.js' +import { request, jsonRequest } from './utils.js' export function createShare(docId, recipientHandle, permission = 'view') { - return request('/api/shares', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ document_id: docId, recipient_handle: recipientHandle, permission }), + return jsonRequest('/api/shares', 'POST', { + document_id: docId, + recipient_handle: recipientHandle, + permission, }) } export function updateSharePermission(shareId, permission) { - return request(`/api/shares/${shareId}`, { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ permission }), - }) + return jsonRequest(`/api/shares/${shareId}`, 'PATCH', { permission }) } export function listShares(docId) { diff --git a/frontend/src/api/topics.js b/frontend/src/api/topics.js index 6996656..d439b0d 100644 --- a/frontend/src/api/topics.js +++ b/frontend/src/api/topics.js @@ -4,26 +4,18 @@ * Consumers: stores/topics.js */ -import { request } from './utils.js' +import { request, jsonRequest } from './utils.js' export function listTopics() { return request('/api/topics') } export function createTopic({ name, description = '', color = '#6366f1' }) { - return request('/api/topics', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ name, description, color }), - }) + return jsonRequest('/api/topics', 'POST', { name, description, color }) } export function updateTopic(id, patch) { - return request(`/api/topics/${id}`, { - method: 'PATCH', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify(patch), - }) + return jsonRequest(`/api/topics/${id}`, 'PATCH', patch) } export function deleteTopic(id) { @@ -31,9 +23,5 @@ export function deleteTopic(id) { } export function suggestTopics(documentId) { - return request('/api/topics/suggest', { - method: 'POST', - headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ document_id: documentId }), - }) + return jsonRequest('/api/topics/suggest', 'POST', { document_id: documentId }) } diff --git a/frontend/src/api/utils.js b/frontend/src/api/utils.js index fe564b5..80c3550 100644 --- a/frontend/src/api/utils.js +++ b/frontend/src/api/utils.js @@ -1,56 +1,42 @@ -/** - * HTTP transport + 401-retry consolidator. - * - * `request()` moved from client.js to break the circular dependency: - * domain modules import request from here; client.js re-exports from those - * same domain modules, so request cannot also live in client.js. - * - * `fetchWithRetry()` consolidates 3 blob-download patterns that previously - * duplicated identical auth-injection + 401-retry boilerplate: - * - adminExportAuditLogCsv (admin.js) - * - adminDownloadDailyExport (admin.js) - * - fetchDocumentContent (documents.js) - * - * Security: Bearer token injected from authStore (Pinia memory only — CLAUDE.md). - * Token is NEVER read from localStorage or sessionStorage. - */ +const NO_REFRESH_PATHS = new Set(['/api/auth/login', '/api/auth/register', '/api/auth/refresh']) -/** - * Core HTTP transport. All JSON-returning endpoints go through this function. - * - * On 401: attempts one token refresh via authStore.refresh() then retries. - * Skip-refresh guard: login/register return 401 for bad credentials (not expired - * tokens), and refresh itself must not retry (would cause infinite loop). - * - * @param {string} path — API path (e.g. '/api/documents') - * @param {RequestInit & {_retry?: boolean}} [options] — fetch options - * @returns {Promise} — parsed JSON response, or null for 204/empty body - */ -export async function request(path, options = {}) { - // Lazy import to avoid circular dependency (stores/auth.js → api/client.js → stores/auth.js) +async function getAuthStore() { const { useAuthStore } = await import('../stores/auth.js') - const authStore = useAuthStore() + return useAuthStore() +} +function withAuthHeaders(options, accessToken) { const headers = { ...(options.headers || {}) } - if (authStore.accessToken) { - headers['Authorization'] = `Bearer ${authStore.accessToken}` + if (accessToken) headers.Authorization = `Bearer ${accessToken}` + return { ...options, headers, credentials: 'include' } +} + +function clearSession(authStore) { + authStore.accessToken = null + authStore.user = null +} + +function fetchAuthenticated(url, options, authStore) { + return fetch(url, withAuthHeaders(options, authStore.accessToken)) +} + +async function refreshAndRetry(authStore, retryFn) { + try { + await authStore.refresh() + return retryFn() + } catch { + clearSession(authStore) + throw new Error('Session expired') } +} - const res = await fetch(path, { ...options, headers, credentials: 'include' }) +/** Core HTTP transport for JSON-returning API endpoints. */ +export async function request(path, options = {}) { + const authStore = await getAuthStore() + const res = await fetch(path, withAuthHeaders(options, authStore.accessToken)) - // 401 → attempt refresh → retry once - // Skip refresh for auth endpoints: login/register return 401 for bad credentials (not expired tokens), - // and refresh itself must not retry to avoid an infinite loop. - const noRefreshPaths = ['/api/auth/login', '/api/auth/register', '/api/auth/refresh'] - if (res.status === 401 && !options._retry && !noRefreshPaths.includes(path)) { - try { - await authStore.refresh() - return request(path, { ...options, _retry: true }) - } catch { - authStore.accessToken = null - authStore.user = null - throw new Error('Session expired') - } + if (res.status === 401 && !options._retry && !NO_REFRESH_PATHS.has(path)) { + return refreshAndRetry(authStore, () => request(path, { ...options, _retry: true })) } if (!res.ok) { @@ -74,45 +60,21 @@ export async function request(path, options = {}) { return res.json() } -/** - * Authenticated fetch with 401-retry for non-JSON responses (blobs, raw Response). - * - * Consolidates adminExportAuditLogCsv, adminDownloadDailyExport, fetchDocumentContent - * which share identical auth-injection + 401-retry boilerplate (CODE-08). - * - * Unlike request(), this does NOT parse the response — the raw Response is returned - * so callers can call .text(), .blob(), etc. as appropriate. - * - * Security: Bearer token from authStore.accessToken (memory only — CLAUDE.md). - * On 401-and-not-already-retried: calls authStore.refresh() and recurses with - * _retry=true. On refresh failure: clears in-memory token state and throws - * 'Session expired'. - * - * @param {string} url — full URL to fetch - * @param {RequestInit} [options] — fetch options (method, headers, etc.) - * @param {boolean} [_retry] — internal retry guard; callers must NOT pass this - * @returns {Promise} — raw Response; caller decides how to consume it - */ +export function jsonRequest(path, method, body) { + return request(path, { + method, + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(body), + }) +} + +/** Authenticated fetch with one refresh retry for raw/blob responses. */ export async function fetchWithRetry(url, options = {}, _retry = false) { - const { useAuthStore } = await import('../stores/auth.js') - const authStore = useAuthStore() - - const headers = { ...(options.headers || {}) } - if (authStore.accessToken) { - headers['Authorization'] = `Bearer ${authStore.accessToken}` - } - - const res = await fetch(url, { ...options, headers, credentials: 'include' }) + const authStore = await getAuthStore() + const res = await fetchAuthenticated(url, options, authStore) if (res.status === 401 && !_retry) { - try { - await authStore.refresh() - return fetchWithRetry(url, options, true) - } catch { - authStore.accessToken = null - authStore.user = null - throw new Error('Session expired') - } + return refreshAndRetry(authStore, () => fetchWithRetry(url, options, true)) } return res diff --git a/frontend/src/components/storage/StorageBrowser.vue b/frontend/src/components/storage/StorageBrowser.vue index d99f63c..799bcfe 100644 --- a/frontend/src/components/storage/StorageBrowser.vue +++ b/frontend/src/components/storage/StorageBrowser.vue @@ -1,7 +1,6 @@