diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f7d3659e8..558ea8a0b 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -21,7 +21,7 @@ jobs: runs-on: ubuntu-latest continue-on-error: true steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false @@ -73,10 +73,10 @@ jobs: name: Python syntax (compileall) runs-on: ubuntu-latest steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0 + - uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 with: python-version: "3.11" # Byte-compile sources — catches syntax errors without installing deps. @@ -86,10 +86,10 @@ jobs: name: JS syntax (node --check) runs-on: ubuntu-latest steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - - uses: actions/setup-node@48b55a011bda9f5d6aeb4c2d9c7362e8dae4041e # v6.4.0 + - uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7.0.0 with: node-version: "20" # Syntax-check our own JS (skip vendored libs in static/lib). @@ -108,7 +108,7 @@ jobs: # ROADMAP "fresh install smoke tests" item; make this required once green. continue-on-error: true steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 persist-credentials: false @@ -135,7 +135,7 @@ jobs: echo "docs_only=false" >> "$GITHUB_OUTPUT" fi - - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0 + - uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 if: steps.docs-check.outputs.docs_only != 'true' with: python-version: "3.11" diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index bb8a8c53e..290418194 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -27,15 +27,15 @@ jobs: language: [actions, javascript-typescript, python] steps: - name: Checkout - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - name: Initialize CodeQL - uses: github/codeql-action/init@8aad20d150bbac5944a9f9d289da16a4b0d87c1e # v4.36.2 + uses: github/codeql-action/init@f205ea1c3313d32999d8d6a48b4f6530d4437b38 # v4.37.4 with: languages: ${{ matrix.language }} build-mode: none - name: Perform CodeQL Analysis - uses: github/codeql-action/analyze@8aad20d150bbac5944a9f9d289da16a4b0d87c1e # v4.36.2 + uses: github/codeql-action/analyze@f205ea1c3313d32999d8d6a48b4f6530d4437b38 # v4.37.4 with: category: "/language:${{ matrix.language }}" diff --git a/.github/workflows/container-scan.yml b/.github/workflows/container-scan.yml index f1c4b5bfd..798d752d4 100644 --- a/.github/workflows/container-scan.yml +++ b/.github/workflows/container-scan.yml @@ -37,12 +37,12 @@ jobs: contents: read steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - name: Lint Dockerfile - uses: hadolint/hadolint-action@2332a7b74a6de0dda2e2221d575162eba76ba5e5 # v3.3.0 + uses: hadolint/hadolint-action@2a66e89f53d0771bb131a7fa31f3136336094aa6 # v3.4.0 with: dockerfile: Dockerfile # DL3008: pinning apt package versions is impractical on a -slim base diff --git a/.github/workflows/container-trivy.yml b/.github/workflows/container-trivy.yml index 2a482f067..a2d7a34a3 100644 --- a/.github/workflows/container-trivy.yml +++ b/.github/workflows/container-trivy.yml @@ -52,17 +52,17 @@ jobs: contents: read steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - name: Set up Buildx - uses: docker/setup-buildx-action@d7f5e7f509e45cec5c76c4d5afdd7de93d0b3df5 # v4.1.0 + uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0 # Build without pushing so a broken Dockerfile is caught here, and the # exact image we ship is what gets scanned. - name: Build image - uses: docker/build-push-action@f9f3042f7e2789586610d6e8b85c8f03e5195baf # v7.2.0 + uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7.3.0 with: context: . push: false @@ -93,15 +93,15 @@ jobs: security-events: write # upload SARIF to the Security tab steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - name: Set up Buildx - uses: docker/setup-buildx-action@d7f5e7f509e45cec5c76c4d5afdd7de93d0b3df5 # v4.1.0 + uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0 - name: Build image - uses: docker/build-push-action@f9f3042f7e2789586610d6e8b85c8f03e5195baf # v7.2.0 + uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7.3.0 with: context: . push: false @@ -119,7 +119,7 @@ jobs: TRIVY_DB_REPOSITORY: ghcr.io/aquasecurity/trivy-db:2 - name: Upload Trivy results - uses: github/codeql-action/upload-sarif@8aad20d150bbac5944a9f9d289da16a4b0d87c1e # v4.36.2 + uses: github/codeql-action/upload-sarif@f205ea1c3313d32999d8d6a48b4f6530d4437b38 # v4.37.4 with: sarif_file: trivy-results.sarif category: trivy-image diff --git a/.github/workflows/dependency-review.yml b/.github/workflows/dependency-review.yml index 0a587de19..0a5e30a4a 100644 --- a/.github/workflows/dependency-review.yml +++ b/.github/workflows/dependency-review.yml @@ -36,7 +36,7 @@ jobs: contents: read steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -55,12 +55,12 @@ jobs: contents: read steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - name: Set up Python - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0 + uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 with: python-version: '3.12' diff --git a/.github/workflows/docker-publish.yml b/.github/workflows/docker-publish.yml index d52c0c4e8..7db67e58c 100644 --- a/.github/workflows/docker-publish.yml +++ b/.github/workflows/docker-publish.yml @@ -45,20 +45,20 @@ jobs: arch: arm64 runner: ubuntu-24.04-arm steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - name: Set up Buildx - uses: docker/setup-buildx-action@d7f5e7f509e45cec5c76c4d5afdd7de93d0b3df5 # v4.1.0 + uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0 - name: Log in to GHCR - uses: docker/login-action@650006c6eb7dba73a995cc03b0b2d7f5ca915bee # v4.2.0 + uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0 with: registry: ${{ env.REGISTRY }} username: ${{ github.actor }} password: ${{ secrets.GITHUB_TOKEN }} - name: Build and push by digest id: build - uses: docker/build-push-action@f9f3042f7e2789586610d6e8b85c8f03e5195baf # v7.2.0 + uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7.3.0 with: context: . platforms: ${{ matrix.platform }} @@ -86,7 +86,7 @@ jobs: contents: read packages: write steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - name: Read APP_VERSION + short sha @@ -103,16 +103,16 @@ jobs: pattern: digest-* merge-multiple: true - name: Set up Buildx - uses: docker/setup-buildx-action@d7f5e7f509e45cec5c76c4d5afdd7de93d0b3df5 # v4.1.0 + uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0 - name: Log in to GHCR - uses: docker/login-action@650006c6eb7dba73a995cc03b0b2d7f5ca915bee # v4.2.0 + uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0 with: registry: ${{ env.REGISTRY }} username: ${{ github.actor }} password: ${{ secrets.GITHUB_TOKEN }} - name: Compute tags id: meta - uses: docker/metadata-action@80c7e94dd9b9319bd5eb7a0e0fe9291e23a2a2e9 # v6.1.0 + uses: docker/metadata-action@dc802804100637a589fabce1cb79ff13a1411302 # v6.2.0 with: images: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }} tags: | diff --git a/.github/workflows/issue-description-check.yml b/.github/workflows/issue-description-check.yml index 5ce6037f0..968f36c12 100644 --- a/.github/workflows/issue-description-check.yml +++ b/.github/workflows/issue-description-check.yml @@ -14,7 +14,7 @@ jobs: # Skip bots (Dependabot, release-drafter, etc.) if: ${{ github.event.issue.user.type != 'Bot' }} steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: sparse-checkout: .github/scripts persist-credentials: false diff --git a/.github/workflows/pr-description-check.yml b/.github/workflows/pr-description-check.yml index 53f0b5f50..4945b6fce 100644 --- a/.github/workflows/pr-description-check.yml +++ b/.github/workflows/pr-description-check.yml @@ -23,7 +23,7 @@ jobs: # Skip bots: they open PRs programmatically and have their own process. if: github.event.pull_request.user.type != 'Bot' steps: - - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: ref: ${{ github.base_ref }} sparse-checkout: .github/scripts diff --git a/.github/workflows/secret-scan.yml b/.github/workflows/secret-scan.yml index 02512204a..ec7b6092e 100644 --- a/.github/workflows/secret-scan.yml +++ b/.github/workflows/secret-scan.yml @@ -35,7 +35,7 @@ jobs: contents: read steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: # Full history so a secret committed in an earlier commit (and later # deleted) is still caught -- deletion does not remove it from Git. diff --git a/.github/workflows/workflow-security.yml b/.github/workflows/workflow-security.yml index ee345333b..b00cd03a4 100644 --- a/.github/workflows/workflow-security.yml +++ b/.github/workflows/workflow-security.yml @@ -36,7 +36,7 @@ jobs: contents: read steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false @@ -61,12 +61,12 @@ jobs: contents: read steps: - name: Checkout repository - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: persist-credentials: false - name: Set up Python - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0 + uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0 with: python-version: '3.12' diff --git a/app.py b/app.py index 8363ba4e9..2ae5ec761 100644 --- a/app.py +++ b/app.py @@ -739,7 +739,7 @@ app.include_router(setup_stt_routes(stt_service)) logger.info("STT service initialized (provider managed via settings)") # Documents (artifacts/canvas) -from routes.document.document_routes import setup_document_routes +from routes.document_routes import setup_document_routes document_router = setup_document_routes(session_manager, upload_handler) app.include_router(document_router) @@ -820,7 +820,7 @@ set_ai_rag_manager(rag_manager, personal_docs_mgr) logger.info("AI interaction tools initialized (session, memory, RAG, UI control)") # Webhooks -from routes.webhook.webhook_routes import setup_webhook_routes +from routes.webhook_routes import setup_webhook_routes app.include_router(setup_webhook_routes(webhook_manager, auth_manager, session_manager, api_key_manager)) # API Tokens @@ -852,7 +852,7 @@ app.include_router(setup_codex_routes( )) app.include_router(setup_claude_routes()) -from routes.vault.vault_routes import setup_vault_routes +from routes.vault_routes import setup_vault_routes app.include_router(setup_vault_routes()) # Contacts (CardDAV) diff --git a/mcp_servers/memory_server.py b/mcp_servers/memory_server.py index fd574fd1f..fafbcfc2b 100644 --- a/mcp_servers/memory_server.py +++ b/mcp_servers/memory_server.py @@ -17,8 +17,6 @@ from mcp.types import Tool, TextContent sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) -from src.memory import MemoryStoreUnreadable - server = Server("memory") # Late-initialized managers (set during first tool call) @@ -31,10 +29,6 @@ _OWNER_SCOPE_ERROR = ( "Error: Memory MCP owner is not configured for an owner-scoped memory store. " "Set ODYSSEUS_MCP_MEMORY_OWNER for this server or use the owner-aware native memory tool." ) -_UNREADABLE_STORE_ERROR = ( - "Error: Memory store is temporarily unreadable — nothing was saved. " - "Repair or restore memory.json, then retry." -) def _configured_owner() -> str | None: @@ -57,21 +51,9 @@ def _owner_scoped_store(entries: list[dict]) -> bool: return any(_entry_owner(entry) for entry in entries if isinstance(entry, dict)) -def _scope_entries(for_update: bool = False) -> tuple[str | None, list[dict], list[dict], str | None]: - """Return configured owner, all entries, visible entries, and optional error. - - ``for_update=True`` is for read-modify-write callers. They save the ``all - entries`` list back, so an unreadable store must be reported as an error - instead of degrading to ``[]`` — otherwise the save writes their one new - entry over the whole store (issue #5673). - """ - if for_update: - try: - entries = _memory_manager.load_all_for_update() - except MemoryStoreUnreadable as e: - return None, [], [], f"{_UNREADABLE_STORE_ERROR} ({e})" - else: - entries = _memory_manager.load_all() +def _scope_entries() -> tuple[str | None, list[dict], list[dict], str | None]: + """Return configured owner, all entries, visible entries, and optional error.""" + entries = _memory_manager.load_all() owner = _configured_owner() if owner is None and _owner_scoped_store(entries): return None, entries, [], _OWNER_SCOPE_ERROR @@ -179,7 +161,7 @@ async def call_tool(name: str, arguments: dict) -> list[TextContent]: category = arguments.get("category", "fact") if not text: return _text_result("Error: Memory text cannot be empty") - owner, memories, _visible, scope_error = _scope_entries(for_update=True) + owner, memories, _visible, scope_error = _scope_entries() if scope_error: return _text_result(scope_error) entry = _memory_manager.add_entry(text, source="ai_agent", category=category, owner=owner) diff --git a/routes/backup_routes.py b/routes/backup_routes.py index 4ecf4f165..313369370 100644 --- a/routes/backup_routes.py +++ b/routes/backup_routes.py @@ -6,7 +6,6 @@ from datetime import datetime from fastapi import APIRouter, HTTPException, Request, Response from core.middleware import require_admin -from services.memory import MemoryStoreUnreadable from src.auth_helpers import get_current_user from src.settings import load_settings, save_settings, load_features, save_features @@ -77,15 +76,7 @@ def setup_backup_routes(memory_manager, preset_manager, skills_manager) -> APIRo # ── Memories ── if "memories" in body and isinstance(body["memories"], list): - # Strict load: importing on top of an unreadable store would write - # only the incoming rows back and drop everything already saved. - try: - existing = memory_manager.load_all_for_update() - except MemoryStoreUnreadable as e: - logger.error("Refusing to import memories: %s", e) - raise HTTPException( - 503, "Memory store is temporarily unreadable — nothing was imported." - ) + existing = memory_manager.load_all() # Dedup against THIS user's own memories only. Using every tenant's # rows (load_all) meant a memory whose text matched any other # user's was silently skipped, so the importing user lost their own diff --git a/routes/document/__init__.py b/routes/document/__init__.py deleted file mode 100644 index 7f79ce1bb..000000000 --- a/routes/document/__init__.py +++ /dev/null @@ -1,6 +0,0 @@ -"""Document route domain package (slice 2m, #4082/#4071). - -Contains document_routes.py and document_helpers.py, migrated from the flat -routes/ directory. Backward-compat shims at routes/document_routes.py and -routes/document_helpers.py re-export from here. -""" diff --git a/routes/document/document_helpers.py b/routes/document/document_helpers.py deleted file mode 100644 index a0c2d08eb..000000000 --- a/routes/document/document_helpers.py +++ /dev/null @@ -1,243 +0,0 @@ -"""document_helpers.py — Pydantic models, doc serializers, owner gating, file-locator helpers shared with document_routes.py.""" - -"""Document routes — CRUD for living documents with version history.""" - -import logging -import os -import re -from typing import Any, Dict, Optional - -from fastapi import HTTPException, Request -from pydantic import BaseModel - -from core.database import Document, DocumentVersion -from core.database import Session as DbSession -from src.auth_helpers import _auth_disabled -from src.upload_handler import UploadHandler - -logger = logging.getLogger(__name__) - - -# ---- Request schemas ---- - -class DocumentCreate(BaseModel): - session_id: Optional[str] = None - title: str = "Untitled" - language: Optional[str] = None - content: str = "" - -class DocumentUpdate(BaseModel): - content: str - summary: Optional[str] = None - force_version: bool = False - -class DocumentPatch(BaseModel): - title: Optional[str] = None - language: Optional[str] = None - session_id: Optional[str] = None # link/unlink document to a session - - -# ---- Helpers ---- - -def _doc_to_dict(doc: Document) -> Dict[str, Any]: - return { - "id": doc.id, - "session_id": doc.session_id, - "title": doc.title, - "language": doc.language, - "current_content": doc.current_content, - "version_count": doc.version_count, - "is_active": doc.is_active, - "archived": bool(getattr(doc, "archived", False)), - "created_at": (doc.created_at.isoformat() + "Z") if doc.created_at else None, - "updated_at": (doc.updated_at.isoformat() + "Z") if doc.updated_at else None, - # Source-email provenance (set when doc was created from an email - # attachment) — drives the "Send signed reply" menu item. - "source_email_uid": getattr(doc, "source_email_uid", None), - "source_email_folder": getattr(doc, "source_email_folder", None), - "source_email_account_id": getattr(doc, "source_email_account_id", None), - "source_email_message_id": getattr(doc, "source_email_message_id", None), - } - -def _version_to_dict(v: DocumentVersion) -> Dict[str, Any]: - return { - "id": v.id, - "document_id": v.document_id, - "version_number": v.version_number, - "content": v.content, - "summary": v.summary, - "source": v.source, - "created_at": v.created_at.isoformat() if v.created_at else None, - } - - -def _verify_doc_owner(db, doc: Document, user: str): - """Verify `user` owns this document. Raise 404 if not. - - Documents now carry their own `owner` column, so a doc whose session - was deleted (session_id → NULL) can still prove ownership and stay - openable / cloneable. We trust that column first and only fall back to - the session join for any not-yet-backfilled legacy row. - """ - if user is None: - if _auth_disabled(): - return # Single-user / no-auth mode: allow access - raise HTTPException(403, "Authentication required") - if doc.owner is not None: - if doc.owner != user: - raise HTTPException(404, "Document not found") - return - # Legacy fallback: derive ownership from the linked session. - if not doc.session_id: - raise HTTPException(404, "Document not found") - session = db.query(DbSession).filter(DbSession.id == doc.session_id).first() - if not session or session.owner != user: - raise HTTPException(404, "Document not found") - - -def _owner_session_filter(q, user): - """Restrict a documents query to those owned by `user`. - - Documents now carry their own `owner` column (backfilled at boot from - the linked session, or assigned to the admin user for legacy/orphaned - docs). We filter on that directly rather than on a session join, so a - document whose session was deleted (session_id → NULL) still shows up - for its owner instead of silently vanishing from the Library + search. - - The owner backfill runs in init_db before the app serves requests, so - by the time this filter is live there are no NULL-owner rows to leak; - we therefore match the owner strictly for authenticated callers.""" - if not user: - if user == "" or _auth_disabled(): - return q - return q.filter(False) - return q.filter(Document.owner == user) - - - -def _slug(name: str) -> str: - """Filesystem-friendly version of a document title. - - Whitespace becomes underscores; other unsafe punctuation is dropped. - Preserves letters, digits, dot, hyphen, underscore. Idempotent. - """ - import re as _re - s = (name or "").strip() - # Drop the trailing extension if the title happens to include one - s = _re.sub(r'\.pdf$', '', s, flags=_re.IGNORECASE) - s = _re.sub(r'\s+', '_', s) - s = _re.sub(r'[^A-Za-z0-9._-]', '', s) - s = _re.sub(r'_+', '_', s).strip('_') - return s or "form" - - -# DPI scale for the interactive PDF view. ~150 DPI (2x of 72 PDF user-units). -_PDF_RENDER_SCALE = 2.0 - - -def _upload_path_inside(upload_dir: str, path: str) -> bool: - base = os.path.realpath(upload_dir) - p = os.path.realpath(path) - try: - return os.path.commonpath([base, p]) == base - except Exception: - return False - - -def _resolve_user_upload_path( - upload_handler: Any, - upload_id: str, - owner: Optional[str], - auth_manager=None, -) -> Optional[str]: - """Resolve an upload id to a filesystem path the caller may read.""" - if upload_handler is None: - return None - resolved = upload_handler.resolve_upload( - upload_id, - owner=owner, - auth_manager=auth_manager, - ) - if not isinstance(resolved, dict) or not resolved: - return None - path = resolved.get("path") - upload_dir = getattr(upload_handler, "upload_dir", None) - if path and upload_dir and not _upload_path_inside(upload_dir, path): - logger.warning("Upload path outside upload directory: %s", path) - return None - return path - - -def _locate_upload( - upload_dir: str, - file_id: str, - owner: Optional[str] = None, - auth_manager=None, - upload_handler: Any = None, -): - """Find an upload by its filename ID via UploadHandler.resolve_upload.""" - if upload_handler is None: - from src.upload_handler import UploadHandler - - base_dir = os.path.dirname(os.path.abspath(upload_dir)) - upload_handler = UploadHandler(base_dir, upload_dir) - return _resolve_user_upload_path(upload_handler, file_id, owner, auth_manager) - - -def _assert_pdf_marker_upload_owned( - request: Request, - content: str, - user: Optional[str], - upload_handler: Any, -) -> None: - """Reject document content whose pdf_source marker points at another user's upload.""" - if upload_handler is None: - return - from src.pdf_form_doc import find_source_upload_id - - upload_id = find_source_upload_id(content or "") - if not upload_id: - return - auth_manager = getattr(getattr(request.app, "state", None), "auth_manager", None) - if not _resolve_user_upload_path(upload_handler, upload_id, user, auth_manager): - raise HTTPException( - 400, - "Document PDF marker references an upload you do not own", - ) - - -def _derive_title(content: str) -> str: - """Derive a title from document content.""" - import re - if not isinstance(content, str): - return "Untitled" - text = content.strip() - if not text: - return "Untitled" - - # Markdown header - md = re.match(r'^#{1,3}\s+(.+)', text, re.MULTILINE) - if md: - title = md.group(1).strip() - if len(title) > 50: - title = title[:48] + "…" - return title - - # HTML heading - html = re.search(r']*>([^<]+)', text, re.IGNORECASE) - if html: - title = html.group(1).strip() - if len(title) > 50: - title = title[:48] + "…" - return title - - # First non-empty line (if short enough) - for line in text.split('\n'): - line = line.strip() - if line and 2 <= len(line) <= 60: - title = re.sub(r'[:#*`]+$', '', line).strip() - if title and len(title) > 50: - title = title[:48] + "…" - return title or "Untitled" - - return "Untitled" diff --git a/routes/document/document_routes.py b/routes/document/document_routes.py deleted file mode 100644 index dae8b09fa..000000000 --- a/routes/document/document_routes.py +++ /dev/null @@ -1,1810 +0,0 @@ -"""Document routes — CRUD for living documents with version history.""" - -import uuid -import logging -from datetime import datetime, timezone -from typing import Dict, Any, List, Optional - -from fastapi import APIRouter, HTTPException, Query, Request, UploadFile, File, Form - -from sqlalchemy import case, func, or_ -from core.database import SessionLocal, Document, DocumentVersion -from core.database import Session as DbSession -from src.auth_helpers import get_current_user, _auth_disabled -from src.constants import MAIL_ATTACHMENTS_DIR -from src.upload_handler import reserve_upload_references - -logger = logging.getLogger(__name__) - - -def _get_session_or_404(db, session_id: str, user: Optional[str]): - session = db.query(DbSession).filter(DbSession.id == session_id).first() - if not session: - raise HTTPException(404, "Session not found") - if user and session.owner != user: - raise HTTPException(404, "Session not found") - return session - - -def _aggregate_language_facets(lang_rows): - """Sum document counts per display language for the library facet. - - NULL-language and explicit "text" rows share the "text" bucket (the - language filter treats them as one), so they must be ADDED. The old dict - comprehension keyed both to "text", silently overwriting one group and - undercounting the facet versus what the filter actually returns. - """ - out = {} - for lang, cnt in lang_rows: - key = lang or "text" - out[key] = out.get(key, 0) + cnt - return out - - -def _library_language_for_document(doc: Document) -> str: - """Return the display language used by the document library. - - PDF documents are stored as markdown wrappers so the editor can preserve - extracted text, form fields, and annotations. The library should still - identify them as PDFs instead of exposing that internal wrapper format. - """ - from src.pdf_form_doc import find_source_upload_id - - if find_source_upload_id(doc.current_content or ""): - return "pdf" - return doc.language or "text" - - -def _email_source_key(content: str) -> tuple[str, str]: - """Return the source email identity embedded in an email draft document.""" - import re - - text = content or "" - uid_m = re.search(r"(?im)^X-Source-UID:\s*(.+?)\s*$", text) - folder_m = re.search(r"(?im)^X-Source-Folder:\s*(.+?)\s*$", text) - uid = (uid_m.group(1).strip() if uid_m else "") - folder = (folder_m.group(1).strip() if folder_m else "INBOX") - return uid, folder - - -from routes.document_helpers import ( - DocumentCreate, DocumentUpdate, DocumentPatch, - _doc_to_dict, _version_to_dict, - _verify_doc_owner, _owner_session_filter, - _slug, _resolve_user_upload_path, _assert_pdf_marker_upload_owned, _derive_title, - _PDF_RENDER_SCALE, -) - - -def setup_document_routes(session_manager, upload_handler=None) -> APIRouter: - router = APIRouter(tags=["documents"]) - - def _reserve_document_uploads(user: Optional[str], content: str) -> None: - missing_id = reserve_upload_references(upload_handler, user, content) - if missing_id: - raise HTTPException( - 409, - f"Referenced upload is no longer available: {missing_id}", - ) - - def _locate_current_user_upload(request: Request, upload_id: str, user: Optional[str]): - if upload_handler is None: - return None - auth_manager = getattr(getattr(request.app, "state", None), "auth_manager", None) - return _resolve_user_upload_path(upload_handler, upload_id, user, auth_manager) - - def _load_pdf_viewer_fitz(): - from src.pdf_runtime import load_pymupdf_for_pdf_viewer - - try: - return load_pymupdf_for_pdf_viewer() - except RuntimeError as exc: - raise HTTPException(503, str(exc)) from exc - - # ---- POST /api/document ---- - @router.post("/api/document") - async def create_document(request: Request, req: DocumentCreate) -> Dict[str, Any]: - from src.auth_helpers import require_privilege - user = require_privilege(request, "can_use_documents") - db = SessionLocal() - try: - # session_id is optional: a doc can be a session-less "library" doc - # (e.g. files imported from the library) — session_id is nullable and - # the doc is owner-stamped, so it lives in the library on its own. - session = None - if req.session_id: - # Match the lenient ownership model the rest of the app uses - # (see _owner_filter): only block when an AUTHENTICATED user is - # writing into a DIFFERENT user's session. In single-user / - # unconfigured / localhost-bypass mode, falsey users preserve - # the existing lenient path. - session = _get_session_or_404(db, req.session_id, user) - - # If no language was supplied (e.g. cloning a doc whose language - # was never set), detect it from the content rather than storing - # NULL — which made the editor fall back to plain text. Defaults - # to markdown for prose. - language = req.language - if not language: - from src.agent_tools.document_tools import _looks_like_email_document, _sniff_doc_language, _coerce_email_document_content - language = _sniff_doc_language(req.content) - else: - from src.agent_tools.document_tools import _looks_like_email_document, _coerce_email_document_content - if _looks_like_email_document(req.content, req.title): - language = "email" - - _reserve_document_uploads(user, req.content) - _assert_pdf_marker_upload_owned(request, req.content, user, upload_handler) - - # Reply drafts are keyed to the source email. If a UI/tool path tries - # to create a second draft for the same email in the same chat, - # update the existing draft instead so quoted thread history stays - # attached to the visible document. - if language == "email" and req.session_id: - source_uid, source_folder = _email_source_key(req.content) - if source_uid: - candidates = ( - db.query(Document) - .filter(Document.session_id == req.session_id) - .filter(Document.is_active == True) - .filter(Document.language == "email") - .order_by(Document.updated_at.desc()) - .limit(25) - .all() - ) - for existing in candidates: - old_uid, old_folder = _email_source_key(existing.current_content or "") - if old_uid != source_uid or old_folder != source_folder: - continue - merged = _coerce_email_document_content(existing.current_content or "", req.content) - if existing.current_content != merged: - new_ver = (existing.version_count or 1) + 1 - existing.current_content = merged - existing.title = req.title or existing.title - existing.version_count = new_ver - db.add(DocumentVersion( - id=str(uuid.uuid4()), - document_id=existing.id, - version_number=new_ver, - content=merged, - summary="Updated existing email draft", - source="user", - )) - db.commit() - db.refresh(existing) - return _doc_to_dict(existing) - - doc_id = str(uuid.uuid4()) - ver_id = str(uuid.uuid4()) - - doc = Document( - id=doc_id, - session_id=req.session_id, - title=req.title, - language=language, - current_content=req.content, - version_count=1, - is_active=True, - # Stamp ownership directly so the doc survives its session - # being deleted. Fall back to the session's owner when the - # request is unauthenticated (single-user / localhost bypass). - owner=user or (session.owner if session else None), - ) - ver = DocumentVersion( - id=ver_id, - document_id=doc_id, - version_number=1, - content=req.content, - summary="Initial version", - source="user", - ) - db.add(doc) - db.add(ver) - db.commit() - db.refresh(doc) - try: - from src.event_bus import fire_event - fire_event("document_created", doc.owner) - except Exception: - logger.debug("document_created event dispatch failed", exc_info=True) - return _doc_to_dict(doc) - except HTTPException: - raise - except Exception as e: - db.rollback() - logger.error(f"Failed to create document: {e}") - raise HTTPException(500, f"Failed to create document: {e}") - finally: - db.close() - - # ---- POST /api/documents/import-pdf ---- - @router.post("/api/documents/import-pdf") - async def import_pdf( - request: Request, - file: UploadFile = File(...), - session_id: Optional[str] = Form(None), - ) -> Dict[str, Any]: - """Upload a PDF and create the matching Document. - - Detects AcroForm fields — if any, creates a form-backed markdown doc - (clickable inputs in the PDF view). Otherwise creates a plain PDF doc - with a `pdf_source` marker so the viewer renders the pages without - overlays. - """ - from src.pdf_forms import has_form_fields, extract_fields - from src.pdf_form_doc import ( - save_field_sidecar, - create_form_markdown_document, - create_plain_pdf_document, - ) - from src.document_processor import _process_pdf, strip_pdf_content_marker - import os - - from src.auth_helpers import require_privilege - user = require_privilege(request, "can_use_documents") - - # session_id is optional — a library import isn't tied to a chat. When - # given, validate it; otherwise the PDF becomes a session-less library - # doc (the doc creators below already handle a missing session). - if session_id: - db = SessionLocal() - try: - _get_session_or_404(db, session_id, user) - finally: - db.close() - - if upload_handler is None: - raise HTTPException(500, "Upload handler not configured") - - client_ip = request.client.host if request.client else "unknown" - try: - meta = upload_handler.save_upload(file, client_ip, owner=user) - except HTTPException: - raise - except Exception as e: - logger.error(f"PDF import save_upload failed: {e}") - raise HTTPException(500, f"Upload failed: {e}") - - upload_id = meta["id"] - pdf_path = _locate_current_user_upload(request, upload_id, user) - if not pdf_path: - raise HTTPException(500, "Saved PDF could not be located") - - title = os.path.splitext(meta.get("original_name") or meta.get("name") or upload_id)[0] - try: - body_text = strip_pdf_content_marker(_process_pdf(pdf_path, owner=user)) - except Exception: - body_text = None - - is_form = False - try: - is_form = has_form_fields(pdf_path) - except Exception as e: - logger.warning(f"has_form_fields failed for {pdf_path}: {e}") - - if is_form: - fields = extract_fields(pdf_path) - save_field_sidecar(pdf_path, fields) - doc_id = create_form_markdown_document( - session_id=session_id, - fields=fields, - upload_id=upload_id, - title=title, - intro_text=body_text, - ) - else: - doc_id = create_plain_pdf_document( - session_id=session_id, - upload_id=upload_id, - title=title, - body_text=body_text, - ) - - if not doc_id: - raise HTTPException(500, "Failed to create document for PDF") - - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(500, "Created document not found") - # The PDF doc creators stamp owner from the session only; a - # session-less library import leaves owner NULL, which the Library's - # owner filter then hides. Stamp the requesting user so it shows. - if not doc.owner and user: - doc.owner = user - db.commit() - db.refresh(doc) - return _doc_to_dict(doc) - finally: - db.close() - - # ---- GET /api/documents/library ---- - @router.get("/api/documents/library") - async def documents_library( - request: Request, - search: Optional[str] = Query(None), - language: Optional[str] = Query(None), - sort: str = Query("recent"), - offset: int = Query(0, ge=0), - limit: int = Query(20, ge=1, le=50), - archived: bool = Query(False), - ) -> Dict[str, Any]: - user = get_current_user(request) - db = SessionLocal() - try: - from sqlalchemy import or_ - pdf_marker_cond = or_( - Document.current_content.like('%\s*\n+#[^\n]*\n+)', re.MULTILINE) - head_match = head_re.match(content) - head = head_match.group(1) if head_match else (content.splitlines()[0] + "\n\n# " + (doc.title or "PDF") + "\n\n") - doc.current_content = head + body_text.strip() + "\n" - doc.version_count = (doc.version_count or 1) + 1 - db.add(DocumentVersion( - id=str(__import__("uuid").uuid4()), - document_id=doc_id, - version_number=doc.version_count, - content=doc.current_content, - summary="PDF text re-extracted (OCR)", - source="ocr", - )) - db.commit() - return {"ok": True, "id": doc_id, "extracted": True, "chars": len(body_text)} - finally: - db.close() - - # ---- POST /api/documents/export-zip — bundle selected docs into a .zip ---- - @router.post("/api/documents/export-zip") - async def documents_export_zip(request: Request): - """Zip the selected documents (each as a text file with the right - extension) — mirrors the gallery's bulk download-zip so multi-export - is one file instead of a blocked flood of individual downloads.""" - user = get_current_user(request) - try: - data = await request.json() - except Exception as e: - logger.warning("Failed to parse export request body, defaulting to empty", exc_info=e) - data = {} - ids = data.get("ids") or [] - if not ids: - raise HTTPException(400, "No documents specified") - _ext = { - "javascript": ".js", "python": ".py", "html": ".html", "css": ".css", - "markdown": ".md", "json": ".json", "yaml": ".yml", "bash": ".sh", - "sql": ".sql", "rust": ".rs", "go": ".go", "java": ".java", "c": ".c", - "cpp": ".cpp", "typescript": ".ts", "ruby": ".rb", "php": ".php", - "text": ".txt", "xml": ".xml", "toml": ".toml", "ini": ".ini", - } - db = SessionLocal() - try: - import io - import re - import zipfile - from fastapi import Response - docs = db.query(Document).filter(Document.id.in_(ids)).all() - buf = io.BytesIO() - used = set() - wrote = 0 - with zipfile.ZipFile(buf, "w", zipfile.ZIP_DEFLATED) as zf: - for doc in docs: - try: - _verify_doc_owner(db, doc, user) - except HTTPException: - continue # skip docs the user doesn't own - ext = _ext.get(doc.language or "text", ".txt") - base = (doc.title or "document").strip() or "document" - base = re.sub(r"[^\w\-. ]+", "", base)[:60].strip() or doc.id - name = base if "." in base else base + ext - i = 1 - while name in used: - name = f"{base}-{i}" + ("" if "." in base else ext) - i += 1 - used.add(name) - zf.writestr(name, doc.current_content or "") - wrote += 1 - if not wrote: - raise HTTPException(404, "No documents found") - return Response( - content=buf.getvalue(), - media_type="application/zip", - headers={"Content-Disposition": 'attachment; filename="documents.zip"'}, - ) - finally: - db.close() - - # ---- PUT /api/document/{doc_id} — user manual edit ---- - # Coalesce window: if the last user version was saved within this many - # seconds, update it in-place (user is still actively editing). - # Once the gap exceeds this, the next save creates a new version. - VERSION_COALESCE_SECONDS = 60 - - @router.put("/api/document/{doc_id}") - async def update_document(request: Request, doc_id: str, req: DocumentUpdate) -> Dict[str, Any]: - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - - incoming_content = req.content - from src.agent_tools.document_tools import _coerce_email_document_content, _looks_like_email_document - is_email_doc = ( - (doc.language or "").lower() == "email" - or _looks_like_email_document(doc.current_content or "", doc.title or "") - or _looks_like_email_document(req.content or "", doc.title or "") - ) - if is_email_doc: - incoming_content = _coerce_email_document_content(doc.current_content or "", req.content) - doc.language = "email" - - # Skip if content is identical unless the caller explicitly wants - # a checkpoint version from the current editor state. - if doc.current_content == incoming_content and not req.force_version: - return _doc_to_dict(doc) - - _reserve_document_uploads(user, incoming_content) - _assert_pdf_marker_upload_owned(request, incoming_content, user, upload_handler) - - # Check if we can coalesce with the latest version - latest_ver = db.query(DocumentVersion).filter( - DocumentVersion.document_id == doc_id, - ).order_by(DocumentVersion.version_number.desc()).first() - - now = datetime.now(timezone.utc) - coalesced = False - if latest_ver and latest_ver.source == "user" and not req.force_version: - ver_time = latest_ver.created_at - if ver_time.tzinfo is None: - ver_time = ver_time.replace(tzinfo=timezone.utc) - age = (now - ver_time).total_seconds() - if age < VERSION_COALESCE_SECONDS: - # Update the existing version in-place - latest_ver.content = incoming_content - latest_ver.created_at = now - if req.summary: - latest_ver.summary = req.summary - coalesced = True - - if not coalesced: - new_ver = doc.version_count + 1 - ver = DocumentVersion( - id=str(uuid.uuid4()), - document_id=doc_id, - version_number=new_ver, - content=incoming_content, - summary=req.summary or "Manual edit", - source="user", - ) - doc.version_count = new_ver - db.add(ver) - - doc.current_content = incoming_content - db.commit() - db.refresh(doc) - return _doc_to_dict(doc) - except HTTPException: - raise - except Exception as e: - db.rollback() - raise HTTPException(500, f"Failed to update document: {e}") - finally: - db.close() - - # ---- PATCH /api/document/{doc_id} — metadata only ---- - @router.patch("/api/document/{doc_id}") - async def patch_document(request: Request, doc_id: str, req: DocumentPatch) -> Dict[str, Any]: - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - if req.title is not None: - doc.title = req.title - if req.language is not None: - doc.language = req.language - if req.session_id is not None: - # Empty string = unlink from session - if req.session_id: - _get_session_or_404(db, req.session_id, user) - doc.session_id = req.session_id if req.session_id else None - if not req.session_id: - # Tab closed / doc detached from its session — drop the - # in-memory active-doc pointer so the last-resort injection - # path doesn't re-surface this doc in a later chat (#1160). - try: - from src.agent_tools.document_tools import clear_active_document - clear_active_document(doc_id) - except Exception as e: - logger.warning("Failed to clear active document %r on detach", doc_id, exc_info=e) - db.commit() - db.refresh(doc) - return _doc_to_dict(doc) - except HTTPException: - raise - except Exception as e: - db.rollback() - raise HTTPException(500, str(e)) - finally: - db.close() - - # ---- DELETE /api/document/{doc_id} — soft delete ---- - @router.delete("/api/document/{doc_id}") - async def delete_document(request: Request, doc_id: str) -> Dict[str, str]: - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - doc.is_active = False - # Closed/deleted — drop the in-memory active-doc pointer so it isn't - # re-injected into a later, unrelated chat (#1160). - try: - from src.agent_tools.document_tools import clear_active_document - clear_active_document(doc_id) - except Exception: - pass - db.commit() - return {"status": "deleted", "id": doc_id} - except HTTPException: - raise - except Exception as e: - db.rollback() - raise HTTPException(500, str(e)) - finally: - db.close() - - # ---- GET /api/document/{doc_id}/versions ---- - @router.get("/api/document/{doc_id}/versions") - async def list_versions(request: Request, doc_id: str) -> List[Dict[str, Any]]: - user = get_current_user(request) - db = SessionLocal() - try: - # Verify ownership before listing versions - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - versions = db.query(DocumentVersion).filter( - DocumentVersion.document_id == doc_id - ).order_by(DocumentVersion.version_number.desc()).all() - return [{ - "id": v.id, - "version_number": v.version_number, - "content": v.content, - "summary": v.summary, - "source": v.source, - "created_at": v.created_at.isoformat() if v.created_at else None, - } for v in versions] - finally: - db.close() - - # ---- GET /api/document/{doc_id}/version/{num} ---- - @router.get("/api/document/{doc_id}/version/{num}") - async def get_version(request: Request, doc_id: str, num: int) -> Dict[str, Any]: - user = get_current_user(request) - db = SessionLocal() - try: - # Verify ownership - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - ver = db.query(DocumentVersion).filter( - DocumentVersion.document_id == doc_id, - DocumentVersion.version_number == num, - ).first() - if not ver: - raise HTTPException(404, "Version not found") - return _version_to_dict(ver) - finally: - db.close() - - # ---- POST /api/document/{doc_id}/restore/{num} ---- - @router.post("/api/document/{doc_id}/restore/{num}") - async def restore_version(request: Request, doc_id: str, num: int) -> Dict[str, Any]: - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - - old_ver = db.query(DocumentVersion).filter( - DocumentVersion.document_id == doc_id, - DocumentVersion.version_number == num, - ).first() - if not old_ver: - raise HTTPException(404, "Version not found") - - new_ver_num = doc.version_count + 1 - ver = DocumentVersion( - id=str(uuid.uuid4()), - document_id=doc_id, - version_number=new_ver_num, - content=old_ver.content, - summary=f"Restored from v{num}", - source="user", - ) - doc.current_content = old_ver.content - doc.version_count = new_ver_num - db.add(ver) - db.commit() - db.refresh(doc) - return _doc_to_dict(doc) - except HTTPException: - raise - except Exception as e: - db.rollback() - raise HTTPException(500, str(e)) - finally: - db.close() - - # ---- POST /api/documents/tidy — clean up broken/empty documents ---- - @router.post("/api/documents/tidy") - async def tidy_documents(request: Request) -> Dict[str, Any]: - """Fix empty titles and remove broken/empty documents (user's docs only).""" - user = get_current_user(request) - db = SessionLocal() - try: - q = ( - db.query(Document) - .outerjoin(DbSession, Document.session_id == DbSession.id) - .filter(Document.is_active == True) - .filter((Document.archived == False) | (Document.archived.is_(None))) - ) - q = _owner_session_filter(q, user) - docs = q.all() - fixed_titles = 0 - deleted = 0 - - # Same junk-detection logic as the scheduled tidy_documents - # action (src/document_actions.py). Keep these two in sync. - import re as _re - from src.document_actions import _JUNK_TITLES - - to_delete = [] - now = datetime.now(timezone.utc) - for doc in docs: - created = doc.created_at - if created and created.tzinfo is None: - created = created.replace(tzinfo=timezone.utc) - - # Skip freshly created documents to avoid deleting them while the user is actively editing - if created and (now - created).total_seconds() < 900: # 15 minutes - continue - - content = (doc.current_content or "").strip() - title_raw = (doc.title or "").strip() - title = title_raw.lower() - is_fresh_empty = ( - not content - and created is not None - and (now - created).total_seconds() < 1800 - ) - if is_fresh_empty: - continue - - # Strip markdown noise to get a "real" character count - stripped = _re.sub(r"^#{1,6}\s+", "", content, flags=_re.MULTILINE) - stripped = _re.sub(r"[*_`>\-=]+", "", stripped) - stripped = _re.sub(r"\s+", " ", stripped).strip() - real_len = len(stripped) - - # Detect email-scaffold stubs: "To: \nSubject: \n---\n" style - # bodies with nothing typed in. Stub = every meaningful line - # is a header label (To:/From:/Subject:/...) with no real - # value (blank, "empty", "(empty)", "-", "none", "n/a"). - _is_email_stub = False - _HEADER_RE = _re.compile(r"^(to|from|cc|bcc|subject|reply-to):\s*(.*)$", _re.I) - _PLACEHOLDER_VALS = {"", "empty", "(empty)", "-", "—", "none", "n/a", "na", "tbd"} - if title in ("new email", "new mail", "new message") or doc.language == "email": - body_lines = [ln.strip() for ln in content.split("\n") - if ln.strip() and ln.strip() != "---"] - def _is_filler(ln): - m = _HEADER_RE.match(ln) - if not m: - return False - val = (m.group(2) or "").strip().lower() - return val in _PLACEHOLDER_VALS - has_real_body = any(not _is_filler(ln) for ln in body_lines) - if body_lines and not has_real_body: - _is_email_stub = True - - # Hard-delete obviously empty / junk documents - if not content or content in ("", "# Untitled"): - to_delete.append(doc); deleted += 1; continue - if _is_email_stub: - to_delete.append(doc); deleted += 1; continue - if title in _JUNK_TITLES: - to_delete.append(doc); deleted += 1; continue - - # Fix empty or placeholder titles on survivors - if not title_raw or title_raw == "Untitled": - new_title = _derive_title(content) - if new_title and new_title != "Untitled": - doc.title = new_title - fixed_titles += 1 - - for doc in to_delete: - db.delete(doc) - - # Also clean up inactive empty docs from previous soft-deletes - inactive_q = ( - db.query(Document) - .outerjoin(DbSession, Document.session_id == DbSession.id) - .filter(Document.is_active == False) - .filter((Document.current_content == None) | (Document.current_content == "")) - ) - inactive_q = _owner_session_filter(inactive_q, user) - inactive_docs = inactive_q.all() - for doc in inactive_docs: - db.delete(doc) - deleted += len(inactive_docs) - - db.commit() - return { - "fixed_titles": fixed_titles, - "deleted": deleted, - "message": f"Fixed {fixed_titles} title{'s' if fixed_titles != 1 else ''}, removed {deleted} empty document{'s' if deleted != 1 else ''}", - } - except Exception as e: - db.rollback() - logger.error(f"Document tidy failed: {e}") - raise HTTPException(500, f"Tidy failed: {e}") - finally: - db.close() - - # ---- POST /api/documents/ai-tidy — AI-powered cleanup of junk/test documents ---- - @router.post("/api/documents/ai-tidy") - async def ai_tidy_documents(request: Request) -> Dict[str, Any]: - """Use AI to judge if documents are junk/test/accidental, then delete them. - Caches verdicts so previously-reviewed docs are skipped.""" - from src.task_endpoint import resolve_task_endpoint - from src.endpoint_resolver import resolve_endpoint - from src.llm_core import llm_call_async - - user = get_current_user(request) - url, model, headers = resolve_task_endpoint(owner=user or None) - if not url or not model: - # Fall back to default endpoint - url, model, headers = resolve_endpoint("default", owner=user or None) - if not url or not model: - raise HTTPException(500, "No endpoint configured for AI tidy") - - db = SessionLocal() - try: - q = ( - db.query(Document) - .outerjoin(DbSession, Document.session_id == DbSession.id) - .filter(Document.is_active == True) - .filter((Document.archived == False) | (Document.archived.is_(None))) - ) - q = _owner_session_filter(q, user) - docs = q.all() - - # Only review docs that haven't been reviewed yet - to_review = [d for d in docs if not d.tidy_verdict] - if not to_review: - return {"deleted": 0, "reviewed": 0, "message": "All documents already reviewed"} - - # Build a batch prompt — review up to 30 at a time - batch = to_review[:30] - doc_list = [] - for i, doc in enumerate(batch): - preview = (doc.current_content or "")[:300].strip() - doc_list.append(f"[{i}] title=\"{doc.title}\" lang={doc.language or 'text'} content_preview=\"{preview}\"") - - prompt = ( - "You are a document library cleaner. For each document below, decide if it is JUNK " - "(test, accidental, placeholder, empty-ish, tool-test, throwaway) or KEEP (real content worth saving).\n\n" - "Respond with ONLY a JSON array of verdicts, one per document, like: [\"junk\",\"keep\",\"junk\",...]\n" - "No explanation, no markdown, just the JSON array.\n\n" - + "\n".join(doc_list) - ) - - response = await llm_call_async( - url, model, - [{"role": "system", "content": "You classify documents as junk or keep. Respond only with a JSON array."}, - {"role": "user", "content": prompt}], - temperature=0.1, - max_tokens=200, - headers=headers, - timeout=30, - ) - - # Parse verdicts - import re - match = re.search(r'\[.*?\]', response, re.DOTALL) - if not match: - raise HTTPException(500, "AI returned invalid response") - - import json as _json - verdicts = _json.loads(match.group()) - - deleted = 0 - reviewed = 0 - for i, doc in enumerate(batch): - if i >= len(verdicts): - break - verdict = str(verdicts[i] or "").lower().strip() - if verdict == "junk": - doc.tidy_verdict = "junk" - db.delete(doc) - deleted += 1 - else: - doc.tidy_verdict = "keep" - reviewed += 1 - - db.commit() - return { - "deleted": deleted, - "reviewed": reviewed, - "remaining": len(to_review) - len(batch), - "message": f"Reviewed {reviewed}, removed {deleted} junk document{'s' if deleted != 1 else ''}", - } - except HTTPException: - raise - except Exception as e: - db.rollback() - logger.error(f"AI tidy failed: {e}") - raise HTTPException(500, f"AI tidy failed: {e}") - finally: - db.close() - - # ---- POST /api/document/{doc_id}/export-pdf/preview ---- - @router.post("/api/document/{doc_id}/export-pdf/preview") - async def export_pdf_preview(doc_id: str, request: Request) -> Dict[str, Any]: - """Return the field-value mapping that would be written to the PDF. - - Frontend shows this in a confirmation modal so the user can spot/fix - any wrong values before triggering the actual download. - """ - from src.pdf_form_doc import find_source_upload_id, parse_markdown_to_values, load_field_sidecar - - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - - upload_id = find_source_upload_id(doc.current_content or "") - if not upload_id: - raise HTTPException(400, "Document is not linked to a source PDF") - - pdf_path = _locate_current_user_upload(request, upload_id, user) - if not pdf_path: - raise HTTPException(404, f"Source PDF {upload_id} not found in uploads") - - fields = load_field_sidecar(pdf_path) - if not fields: - raise HTTPException(404, "Field schema sidecar missing for source PDF") - - values = parse_markdown_to_values(doc.current_content or "") - field_meta = {f["name"]: f for f in fields} - - preview = [] - for name, current in values.items(): - meta = field_meta.get(name) - if not meta: - continue - preview.append({ - "name": name, - "label": meta.get("label") or name, - "type": meta.get("type"), - "options": meta.get("options") or [], - "page": meta.get("page"), - "value": current, - }) - - unknown = [ - name for name in values - if name not in field_meta - ] - return { - "doc_id": doc_id, - "upload_id": upload_id, - "fields": preview, - "unknown_fields": unknown, - "total": len(fields), - "filled": sum(1 for p in preview if p["value"] not in ("", False, None)), - } - finally: - db.close() - - # ---- GET /api/document/{doc_id}/render-pages ---- - @router.get("/api/document/{doc_id}/render-pages") - async def render_pages(doc_id: str, request: Request) -> Dict[str, Any]: - """Return per-page metadata for the interactive PDF view. - - Each page entry has its rendered-image dimensions (matching what - /page/{n}.png returns at the same DPI) plus the list of form fields - on that page with their rects translated to image-pixel coordinates. - Frontend overlays HTML form controls at those positions. - """ - from src.pdf_form_doc import find_source_upload_id, parse_markdown_to_values, load_field_sidecar - - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - upload_id = find_source_upload_id(doc.current_content or "") - if not upload_id: - raise HTTPException(400, "Document is not linked to a source PDF") - pdf_path = _locate_current_user_upload(request, upload_id, user) - if not pdf_path: - raise HTTPException(404, f"Source PDF {upload_id} not found") - - fitz = _load_pdf_viewer_fitz() - schema = load_field_sidecar(pdf_path) or [] - values = parse_markdown_to_values(doc.current_content or "") - - # Group fields by page - by_page: Dict[int, list] = {} - for f in schema: - by_page.setdefault(f["page"], []).append(f) - - scale = _PDF_RENDER_SCALE - pdf_doc = fitz.open(pdf_path) - try: - pages_out = [] - for page_index in range(pdf_doc.page_count): - page = pdf_doc[page_index] - page_no = page_index + 1 - pw, ph = page.rect.width, page.rect.height - img_w = int(pw * scale) - img_h = int(ph * scale) - fields_out = [] - for f in by_page.get(page_no, []): - x0, y0, x1, y1 = f["rect"] - fields_out.append({ - "name": f["name"], - "type": f["type"], - "label": f.get("label") or "", - "options": f.get("options") or [], - "value": values.get(f["name"], f.get("value", "")), - "rect_px": [ - int(x0 * scale), int(y0 * scale), - int(x1 * scale), int(y1 * scale), - ], - }) - pages_out.append({ - "page": page_no, - "width": img_w, - "height": img_h, - "fields": fields_out, - }) - return {"doc_id": doc_id, "scale": scale, "pages": pages_out} - finally: - pdf_doc.close() - finally: - db.close() - - # ---- GET /api/document/{doc_id}/page/{n}.png ---- - @router.get("/api/document/{doc_id}/page/{page_no}.png") - async def render_page_png(doc_id: str, page_no: int, request: Request): - """Render one page of the source PDF as a PNG (no values stamped — the - frontend overlays HTML form inputs on top).""" - from fastapi.responses import Response - from src.pdf_form_doc import find_source_upload_id - - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - upload_id = find_source_upload_id(doc.current_content or "") - if not upload_id: - raise HTTPException(400, "Document is not linked to a source PDF") - pdf_path = _locate_current_user_upload(request, upload_id, user) - if not pdf_path: - raise HTTPException(404, "Source PDF not found") - finally: - db.close() - - fitz = _load_pdf_viewer_fitz() - pdf_doc = fitz.open(pdf_path) - try: - if page_no < 1 or page_no > pdf_doc.page_count: - raise HTTPException(404, "Page out of range") - page = pdf_doc[page_no - 1] - mat = fitz.Matrix(_PDF_RENDER_SCALE, _PDF_RENDER_SCALE) - pix = page.get_pixmap(matrix=mat, alpha=False) - png_bytes = pix.tobytes("png") - return Response( - content=png_bytes, - media_type="image/png", - headers={"Cache-Control": "public, max-age=3600"}, - ) - finally: - pdf_doc.close() - - # ---- POST /api/document/{doc_id}/ai-fill-annotations ---- - @router.post("/api/document/{doc_id}/ai-fill-annotations") - async def ai_fill_annotations(doc_id: str, request: Request) -> Dict[str, Any]: - """Ask a vision-capable LLM to locate fillable areas on a flat PDF and - propose annotation values for each, given a free-form user instruction. - - Returns a list of annotations: [{page, x, y, w, h, value}] where x/y/w/h - are page-percentages (0–100) — same coordinate system as the freeform - annotations the frontend already renders. - """ - import base64 - import json - import fitz - from src.pdf_form_doc import find_source_upload_id - from src.document_processor import _resolve_vl_model, _load_vl_settings - from src.llm_core import llm_call_async - - body = await request.json() if request.headers.get("content-type", "").startswith("application/json") else {} - instruction = (body or {}).get("instruction", "").strip() - if not instruction: - raise HTTPException(400, "instruction is required") - - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - upload_id = find_source_upload_id(doc.current_content or "") - if not upload_id: - raise HTTPException(400, "Document is not linked to a source PDF") - pdf_path = _locate_current_user_upload(request, upload_id, user) - if not pdf_path: - raise HTTPException(404, "Source PDF not found") - finally: - db.close() - - # Resolve VL model (admin-configured or auto-detected vision-capable) - settings = _load_vl_settings() - vl_model = settings.get("vision_model", "") - try: - url, model_id, headers = _resolve_vl_model(vl_model, owner=user) - except Exception as e: - raise HTTPException(503, f"No vision model available: {e}") - - system_prompt = ( - "You analyze rendered PDF page images and propose values to fill in. " - "For each blank line, box, underscore, or labeled space on the page that " - "should be filled given the user's instruction, output one annotation. " - "Coordinates are percentages (0-100) of the page width/height with the " - "origin at top-left. Width/height should match the visible blank box. " - "Return ONLY a JSON array, no prose, no markdown fences. Each entry: " - '{"x": number, "y": number, "w": number, "h": number, "value": string}. ' - "If a region should not be filled, omit it. If nothing should be filled, " - "return []." - ) - - all_annotations = [] - pdf_doc = fitz.open(pdf_path) - try: - for page_index in range(pdf_doc.page_count): - page = pdf_doc[page_index] - mat = fitz.Matrix(_PDF_RENDER_SCALE, _PDF_RENDER_SCALE) - pix = page.get_pixmap(matrix=mat, alpha=False) - png_bytes = pix.tobytes("png") - b64 = base64.b64encode(png_bytes).decode("ascii") - - messages = [ - {"role": "system", "content": system_prompt}, - { - "role": "user", - "content": [ - { - "type": "text", - "text": ( - f"User instruction:\n{instruction}\n\n" - f"This is page {page_index + 1} of {pdf_doc.page_count}. " - "Return JSON array of annotations to add to this page." - ), - }, - { - "type": "image_url", - "image_url": {"url": f"data:image/png;base64,{b64}"}, - }, - ], - }, - ] - try: - raw = await llm_call_async( - url, model_id, messages, - temperature=0.1, max_tokens=2000, headers=headers, - ) - except Exception as e: - logger.error(f"VL call failed on page {page_index + 1}: {e}") - continue - - raw = (raw or "").strip() - if raw.startswith("```"): - raw = raw.split("\n", 1)[-1].rsplit("```", 1)[0].strip() - try: - parsed = json.loads(raw) - except Exception: - logger.warning(f"AI fill: page {page_index + 1} returned non-JSON: {raw[:200]}") - continue - if not isinstance(parsed, list): - continue - for item in parsed: - if not isinstance(item, dict): - continue - try: - x = float(item.get("x", 0)) - y = float(item.get("y", 0)) - w = float(item.get("w", 0)) - h = float(item.get("h", 0)) - value = str(item.get("value", "") or "") - except Exception: - continue - # Clamp + reject zero-size entries - if w <= 0.5 or h <= 0.3: - continue - x = max(0.0, min(99.0, x)) - y = max(0.0, min(99.0, y)) - w = max(0.5, min(100.0 - x, w)) - h = max(0.3, min(100.0 - y, h)) - if not value.strip(): - continue - all_annotations.append({ - "page": page_index + 1, - "x": round(x, 2), - "y": round(y, 2), - "w": round(w, 2), - "h": round(h, 2), - "value": value, - }) - finally: - pdf_doc.close() - - return {"annotations": all_annotations} - - # ---- GET /api/document/{doc_id}/render-pdf ---- - @router.get("/api/document/{doc_id}/render-pdf") - async def render_pdf(doc_id: str, request: Request): - """Inline PDF preview filled with the current markdown values. - - Same plumbing as the export route, but no signature stamping and - served inline (Content-Disposition: inline) so the browser can - embed it in an iframe. Cache-busted by the caller via query string. - """ - import base64 - import os - import tempfile - from fastapi.responses import FileResponse - from starlette.background import BackgroundTask - from src.pdf_form_doc import find_source_upload_id, parse_markdown_to_values, parse_markdown_annotations - from src.pdf_forms import fill_fields, stamp_annotations - from core.database import Signature - - # Track temp files for this request so they get unlinked AFTER - # the response is fully sent (BackgroundTask runs post-send). - _to_unlink: list[str] = [] - def _cleanup_temps(): - for _p in _to_unlink: - try: - os.unlink(_p) - except FileNotFoundError: - pass - except Exception as _e: - logger.warning(f"Could not unlink temp PDF {_p}: {_e}") - - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - upload_id = find_source_upload_id(doc.current_content or "") - if not upload_id: - raise HTTPException(400, "Document is not linked to a source PDF") - pdf_path = _locate_current_user_upload(request, upload_id, user) - if not pdf_path: - raise HTTPException(404, f"Source PDF {upload_id} not found") - - # Fail fast with a clear 503 if the optional PyMuPDF dependency - # is missing — fill_fields/stamp_annotations will otherwise - # raise RuntimeError deep inside and bubble out as a 500. - # Mirrors the convention in _load_pdf_viewer_fitz above. - _load_pdf_viewer_fitz() - - values = parse_markdown_to_values(doc.current_content or "") - out_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name - _to_unlink.append(out_path) - try: - fill_fields(pdf_path, out_path, values) - except Exception as e: - logger.error(f"render_pdf fill_fields failed for {doc_id}: {e}") - _cleanup_temps() - raise HTTPException(500, f"PDF render failed: {e}") - - annotations = parse_markdown_annotations(doc.current_content or "") - if annotations: - ann_sig_ids = [ - a["value"][len("signature:"):].strip() - for a in annotations - if a.get("kind") == "signature" - and isinstance(a.get("value"), str) - and a["value"].startswith("signature:") - ] - ann_signature_pngs: dict[str, bytes] = {} - if ann_sig_ids: - # SECURITY: filter by owner so a caller can't reference - # someone else's signature ID from doc markdown and have - # it stamped/exported. - _sig_q = db.query(Signature).filter(Signature.id.in_(ann_sig_ids)) - if user: - _sig_q = _sig_q.filter(Signature.owner == user) - sig_rows = _sig_q.all() - for s in sig_rows: - try: - ann_signature_pngs[s.id] = base64.b64decode(s.data_png) - except Exception as e: - logger.warning(f"Bad annotation signature data for {s.id}: {e}") - annotated_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name - _to_unlink.append(annotated_path) - try: - stamp_annotations(out_path, annotated_path, annotations, ann_signature_pngs) - out_path = annotated_path - except Exception as e: - logger.error(f"stamp_annotations (render) failed for {doc_id}: {e}") - - return FileResponse( - out_path, - media_type="application/pdf", - headers={"Content-Disposition": "inline"}, - background=BackgroundTask(_cleanup_temps), - ) - finally: - db.close() - - # ---- GET /api/document/{doc_id}/export-pdf ---- - @router.get("/api/document/{doc_id}/export-pdf") - async def export_pdf(doc_id: str, request: Request): - """Stream the filled PDF for download. - - Reads field values and signature selections from the markdown — there - is no separate confirmation step. Signature fields contain their - chosen signature ID encoded as `signature:` in the value. - """ - import base64 - import os - import tempfile - from fastapi.responses import FileResponse - from starlette.background import BackgroundTask - from src.pdf_form_doc import find_source_upload_id, parse_markdown_to_values, load_field_sidecar, parse_markdown_annotations - from src.pdf_forms import fill_fields, stamp_signatures, stamp_annotations - from core.database import Signature - - _to_unlink: list[str] = [] - def _cleanup_temps(): - for _p in _to_unlink: - try: - os.unlink(_p) - except FileNotFoundError: - pass - except Exception as _e: - logger.warning(f"Could not unlink temp PDF {_p}: {_e}") - - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - - upload_id = find_source_upload_id(doc.current_content or "") - if not upload_id: - raise HTTPException(400, "Document is not linked to a source PDF") - - pdf_path = _locate_current_user_upload(request, upload_id, user) - if not pdf_path: - raise HTTPException(404, f"Source PDF {upload_id} not found in uploads") - - schema = load_field_sidecar(pdf_path) or [] - sig_field_names = {f["name"] for f in schema if f.get("type") == "signature"} - - all_values = parse_markdown_to_values(doc.current_content or "") - # Split: signature fields go to stamps, everything else to fill_fields - text_values: dict = {} - sig_ids: dict[str, str] = {} - for name, raw in all_values.items(): - if name in sig_field_names and isinstance(raw, str) and raw.startswith("signature:"): - sig_ids[name] = raw[len("signature:"):].strip() - elif name not in sig_field_names: - text_values[name] = raw - - stamps: dict = {} - if sig_ids: - # SECURITY: filter by owner — same reason as render_pdf. - _sig_q2 = db.query(Signature).filter(Signature.id.in_(list(sig_ids.values()))) - if user: - _sig_q2 = _sig_q2.filter(Signature.owner == user) - rows = _sig_q2.all() - by_id = {s.id: s for s in rows} - for field_name, sid in sig_ids.items(): - s = by_id.get(sid) - if not s: - continue - try: - stamps[field_name] = base64.b64decode(s.data_png) - except Exception as e: - logger.warning(f"Bad signature data for {sid}: {e}") - - filled_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name - _to_unlink.append(filled_path) - try: - fill_fields(pdf_path, filled_path, text_values) - except Exception as e: - logger.error(f"fill_fields failed for doc {doc_id}: {e}") - _cleanup_temps() - raise HTTPException(500, f"PDF fill failed: {e}") - - out_path = filled_path - if stamps: - stamped_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name - _to_unlink.append(stamped_path) - try: - stamp_signatures(filled_path, stamped_path, stamps) - out_path = stamped_path - except Exception as e: - logger.error(f"stamp_signatures failed for doc {doc_id}: {e}") - - # Burn freeform annotations (Text/Check/Sign drops) on top. - annotations = parse_markdown_annotations(doc.current_content or "") - if annotations: - # Resolve any signature annotations to their PNG bytes. - ann_sig_ids = [ - a["value"][len("signature:"):].strip() - for a in annotations - if a.get("kind") == "signature" - and isinstance(a.get("value"), str) - and a["value"].startswith("signature:") - ] - ann_signature_pngs: dict[str, bytes] = {} - if ann_sig_ids: - # SECURITY: filter by owner so a caller can't reference - # someone else's signature ID from doc markdown and have - # it stamped/exported. - _sig_q = db.query(Signature).filter(Signature.id.in_(ann_sig_ids)) - if user: - _sig_q = _sig_q.filter(Signature.owner == user) - sig_rows = _sig_q.all() - for s in sig_rows: - try: - ann_signature_pngs[s.id] = base64.b64decode(s.data_png) - except Exception as e: - logger.warning(f"Bad annotation signature data for {s.id}: {e}") - annotated_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name - _to_unlink.append(annotated_path) - try: - stamp_annotations(out_path, annotated_path, annotations, ann_signature_pngs) - out_path = annotated_path - except Exception as e: - logger.error(f"stamp_annotations failed for doc {doc_id}: {e}") - - download_name = _slug(doc.title or "form") + "_annotated.pdf" - return FileResponse( - out_path, - media_type="application/pdf", - filename=download_name, - background=BackgroundTask(_cleanup_temps), - ) - finally: - db.close() - - # ---- POST /api/document/{doc_id}/prepare-signed-reply ---- - @router.post("/api/document/{doc_id}/prepare-signed-reply") - async def prepare_signed_reply(doc_id: str, request: Request): - """Bake the current PDF state (form fields + signature stamps + - annotations) into a flattened PDF, drop it in COMPOSE_UPLOADS_DIR - and return the reply context (To/Subject/threading headers) so the - frontend can open a reply draft with this attachment pre-loaded. - - Requires the document to have source_email_* metadata (set when the - doc was created via /api/email/attachment-as-doc). Otherwise 400. - """ - import base64 - import tempfile - import shutil - import uuid as _uuid - import email as _email_mod - from src.pdf_form_doc import ( - find_source_upload_id, parse_markdown_to_values, - load_field_sidecar, parse_markdown_annotations, - ) - from src.pdf_forms import fill_fields, stamp_signatures, stamp_annotations - from core.database import Signature - # COMPOSE_UPLOADS_DIR lives in email_routes — re-derive here so we - # don't import from a routes file (cycle-prone). Same env override - # as email_routes (ODYSSEUS_MAIL_ATTACHMENTS_DIR). - from pathlib import Path as _Path - _COMPOSE_DIR = _Path(MAIL_ATTACHMENTS_DIR) / "_compose" - _COMPOSE_DIR.mkdir(parents=True, exist_ok=True) - - user = get_current_user(request) - db = SessionLocal() - try: - doc = db.query(Document).filter(Document.id == doc_id).first() - if not doc: - raise HTTPException(404, "Document not found") - _verify_doc_owner(db, doc, user) - - if not (doc.source_email_uid and doc.source_email_folder): - raise HTTPException(400, "Document has no source email — cannot reply") - - # 1) Build the flattened PDF (same pipeline as export_pdf) - upload_id = find_source_upload_id(doc.current_content or "") - if not upload_id: - raise HTTPException(400, "Document is not linked to a source PDF") - pdf_path = _locate_current_user_upload(request, upload_id, user) - if not pdf_path: - raise HTTPException(404, f"Source PDF {upload_id} not found") - - schema = load_field_sidecar(pdf_path) or [] - sig_field_names = {f["name"] for f in schema if f.get("type") == "signature"} - all_values = parse_markdown_to_values(doc.current_content or "") - text_values: dict = {} - sig_ids: dict[str, str] = {} - for name, raw in all_values.items(): - if name in sig_field_names and isinstance(raw, str) and raw.startswith("signature:"): - sig_ids[name] = raw[len("signature:"):].strip() - elif name not in sig_field_names: - text_values[name] = raw - - stamps: dict = {} - if sig_ids: - # SECURITY: filter by owner — same reason as render_pdf. - _sig_q2 = db.query(Signature).filter(Signature.id.in_(list(sig_ids.values()))) - if user: - _sig_q2 = _sig_q2.filter(Signature.owner == user) - rows = _sig_q2.all() - by_id = {s.id: s for s in rows} - for fname, sid in sig_ids.items(): - s = by_id.get(sid) - if not s: - continue - try: - stamps[fname] = base64.b64decode(s.data_png) - except Exception: - pass - - import os - _to_unlink: list[str] = [] - filled_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name - _to_unlink.append(filled_path) - fill_fields(pdf_path, filled_path, text_values) - out_path = filled_path - if stamps: - stamped_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name - _to_unlink.append(stamped_path) - try: - stamp_signatures(filled_path, stamped_path, stamps) - out_path = stamped_path - except Exception as e: - logger.warning(f"stamp_signatures failed for {doc_id}: {e}") - - annotations = parse_markdown_annotations(doc.current_content or "") - if annotations: - ann_sig_ids = [ - a["value"][len("signature:"):].strip() - for a in annotations - if a.get("kind") == "signature" - and isinstance(a.get("value"), str) - and a["value"].startswith("signature:") - ] - ann_signature_pngs: dict[str, bytes] = {} - if ann_sig_ids: - # SECURITY: filter by owner so a caller can't reference - # someone else's signature ID from doc markdown and have - # it stamped/exported. - _sig_q = db.query(Signature).filter(Signature.id.in_(ann_sig_ids)) - if user: - _sig_q = _sig_q.filter(Signature.owner == user) - sig_rows = _sig_q.all() - for s in sig_rows: - try: - ann_signature_pngs[s.id] = base64.b64decode(s.data_png) - except Exception: - pass - annotated_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name - _to_unlink.append(annotated_path) - try: - stamp_annotations(out_path, annotated_path, annotations, ann_signature_pngs) - out_path = annotated_path - except Exception as e: - logger.warning(f"stamp_annotations failed for {doc_id}: {e}") - - # 2) Move/copy into COMPOSE_UPLOADS_DIR with the token format - # `_` that /api/email/send expects. - filename = _slug(doc.title or "signed") + "_signed.pdf" - token = f"{_uuid.uuid4().hex}_{filename}" - dest = _COMPOSE_DIR / token - shutil.copyfile(out_path, str(dest)) - # Unlink the intermediate temp PDFs now that they've been - # copied into COMPOSE_UPLOADS_DIR. - for _p in _to_unlink: - try: - os.unlink(_p) - except FileNotFoundError: - pass - except Exception as _e: - logger.warning(f"Could not unlink temp PDF {_p}: {_e}") - - # 3) Fetch the source email's headers so we can build a clean reply - # context (To/Subject/In-Reply-To/References). - try: - from routes.email_routes import _imap, _decode_header - from routes.email_helpers import _q - except Exception: - _imap = None - _decode_header = lambda x: x or "" - _q = lambda x: x or "" - - to_addr = "" - from_name = "" - subject = "" - in_reply_to = doc.source_email_message_id or "" - references = in_reply_to - if _imap: - try: - with _imap(doc.source_email_account_id or None) as conn: - conn.select(_q(doc.source_email_folder), readonly=True) - status, data = conn.fetch(doc.source_email_uid.encode(), "(RFC822.HEADER)") - if status == "OK" and data and data[0]: - raw_hdr = data[0][1] - m = _email_mod.message_from_bytes(raw_hdr) - sender = _decode_header(m.get("From", "")) - from_name, to_addr = _email_mod.utils.parseaddr(sender) - if not to_addr: - to_addr = sender - subject = _decode_header(m.get("Subject", "") or "") - if subject and not subject.lower().startswith("re:"): - subject = "Re: " + subject - msg_refs = (m.get("References") or "").strip() - msg_in_reply = (m.get("Message-ID") or "").strip() or in_reply_to - in_reply_to = msg_in_reply - references = (msg_refs + " " + msg_in_reply).strip() if msg_refs else msg_in_reply - except Exception as e: - logger.warning(f"prepare-signed-reply header fetch failed: {e}") - - return { - "ok": True, - "attachment": { - "token": token, - "filename": filename, - "size": dest.stat().st_size, - }, - "reply": { - "to": to_addr, - "to_name": from_name, - "subject": subject, - "in_reply_to": in_reply_to, - "references": references, - "account_id": doc.source_email_account_id or None, - "source_uid": doc.source_email_uid, - "source_folder": doc.source_email_folder, - "source_message_id": doc.source_email_message_id, - }, - } - finally: - db.close() - - return router diff --git a/routes/document_helpers.py b/routes/document_helpers.py index c1f68ca51..a0c2d08eb 100644 --- a/routes/document_helpers.py +++ b/routes/document_helpers.py @@ -1,14 +1,243 @@ -"""Backward-compat shim — canonical location is routes/document/document_helpers.py. +"""document_helpers.py — Pydantic models, doc serializers, owner gating, file-locator helpers shared with document_routes.py.""" -This module is replaced in ``sys.modules`` by the canonical module object so -that ``import routes.document_helpers``, ``from routes.document_helpers import -X``, and the ``sys.modules.pop("routes.document_helpers")`` + re-import -pattern used by test_security_regressions.py all operate on the *same* object. -Keeps existing import paths working after slice 2m (#4082/#4071). -""" +"""Document routes — CRUD for living documents with version history.""" -import sys as _sys +import logging +import os +import re +from typing import Any, Dict, Optional -from routes.document import document_helpers as _canonical # noqa: F401 +from fastapi import HTTPException, Request +from pydantic import BaseModel -_sys.modules[__name__] = _canonical +from core.database import Document, DocumentVersion +from core.database import Session as DbSession +from src.auth_helpers import _auth_disabled +from src.upload_handler import UploadHandler + +logger = logging.getLogger(__name__) + + +# ---- Request schemas ---- + +class DocumentCreate(BaseModel): + session_id: Optional[str] = None + title: str = "Untitled" + language: Optional[str] = None + content: str = "" + +class DocumentUpdate(BaseModel): + content: str + summary: Optional[str] = None + force_version: bool = False + +class DocumentPatch(BaseModel): + title: Optional[str] = None + language: Optional[str] = None + session_id: Optional[str] = None # link/unlink document to a session + + +# ---- Helpers ---- + +def _doc_to_dict(doc: Document) -> Dict[str, Any]: + return { + "id": doc.id, + "session_id": doc.session_id, + "title": doc.title, + "language": doc.language, + "current_content": doc.current_content, + "version_count": doc.version_count, + "is_active": doc.is_active, + "archived": bool(getattr(doc, "archived", False)), + "created_at": (doc.created_at.isoformat() + "Z") if doc.created_at else None, + "updated_at": (doc.updated_at.isoformat() + "Z") if doc.updated_at else None, + # Source-email provenance (set when doc was created from an email + # attachment) — drives the "Send signed reply" menu item. + "source_email_uid": getattr(doc, "source_email_uid", None), + "source_email_folder": getattr(doc, "source_email_folder", None), + "source_email_account_id": getattr(doc, "source_email_account_id", None), + "source_email_message_id": getattr(doc, "source_email_message_id", None), + } + +def _version_to_dict(v: DocumentVersion) -> Dict[str, Any]: + return { + "id": v.id, + "document_id": v.document_id, + "version_number": v.version_number, + "content": v.content, + "summary": v.summary, + "source": v.source, + "created_at": v.created_at.isoformat() if v.created_at else None, + } + + +def _verify_doc_owner(db, doc: Document, user: str): + """Verify `user` owns this document. Raise 404 if not. + + Documents now carry their own `owner` column, so a doc whose session + was deleted (session_id → NULL) can still prove ownership and stay + openable / cloneable. We trust that column first and only fall back to + the session join for any not-yet-backfilled legacy row. + """ + if user is None: + if _auth_disabled(): + return # Single-user / no-auth mode: allow access + raise HTTPException(403, "Authentication required") + if doc.owner is not None: + if doc.owner != user: + raise HTTPException(404, "Document not found") + return + # Legacy fallback: derive ownership from the linked session. + if not doc.session_id: + raise HTTPException(404, "Document not found") + session = db.query(DbSession).filter(DbSession.id == doc.session_id).first() + if not session or session.owner != user: + raise HTTPException(404, "Document not found") + + +def _owner_session_filter(q, user): + """Restrict a documents query to those owned by `user`. + + Documents now carry their own `owner` column (backfilled at boot from + the linked session, or assigned to the admin user for legacy/orphaned + docs). We filter on that directly rather than on a session join, so a + document whose session was deleted (session_id → NULL) still shows up + for its owner instead of silently vanishing from the Library + search. + + The owner backfill runs in init_db before the app serves requests, so + by the time this filter is live there are no NULL-owner rows to leak; + we therefore match the owner strictly for authenticated callers.""" + if not user: + if user == "" or _auth_disabled(): + return q + return q.filter(False) + return q.filter(Document.owner == user) + + + +def _slug(name: str) -> str: + """Filesystem-friendly version of a document title. + + Whitespace becomes underscores; other unsafe punctuation is dropped. + Preserves letters, digits, dot, hyphen, underscore. Idempotent. + """ + import re as _re + s = (name or "").strip() + # Drop the trailing extension if the title happens to include one + s = _re.sub(r'\.pdf$', '', s, flags=_re.IGNORECASE) + s = _re.sub(r'\s+', '_', s) + s = _re.sub(r'[^A-Za-z0-9._-]', '', s) + s = _re.sub(r'_+', '_', s).strip('_') + return s or "form" + + +# DPI scale for the interactive PDF view. ~150 DPI (2x of 72 PDF user-units). +_PDF_RENDER_SCALE = 2.0 + + +def _upload_path_inside(upload_dir: str, path: str) -> bool: + base = os.path.realpath(upload_dir) + p = os.path.realpath(path) + try: + return os.path.commonpath([base, p]) == base + except Exception: + return False + + +def _resolve_user_upload_path( + upload_handler: Any, + upload_id: str, + owner: Optional[str], + auth_manager=None, +) -> Optional[str]: + """Resolve an upload id to a filesystem path the caller may read.""" + if upload_handler is None: + return None + resolved = upload_handler.resolve_upload( + upload_id, + owner=owner, + auth_manager=auth_manager, + ) + if not isinstance(resolved, dict) or not resolved: + return None + path = resolved.get("path") + upload_dir = getattr(upload_handler, "upload_dir", None) + if path and upload_dir and not _upload_path_inside(upload_dir, path): + logger.warning("Upload path outside upload directory: %s", path) + return None + return path + + +def _locate_upload( + upload_dir: str, + file_id: str, + owner: Optional[str] = None, + auth_manager=None, + upload_handler: Any = None, +): + """Find an upload by its filename ID via UploadHandler.resolve_upload.""" + if upload_handler is None: + from src.upload_handler import UploadHandler + + base_dir = os.path.dirname(os.path.abspath(upload_dir)) + upload_handler = UploadHandler(base_dir, upload_dir) + return _resolve_user_upload_path(upload_handler, file_id, owner, auth_manager) + + +def _assert_pdf_marker_upload_owned( + request: Request, + content: str, + user: Optional[str], + upload_handler: Any, +) -> None: + """Reject document content whose pdf_source marker points at another user's upload.""" + if upload_handler is None: + return + from src.pdf_form_doc import find_source_upload_id + + upload_id = find_source_upload_id(content or "") + if not upload_id: + return + auth_manager = getattr(getattr(request.app, "state", None), "auth_manager", None) + if not _resolve_user_upload_path(upload_handler, upload_id, user, auth_manager): + raise HTTPException( + 400, + "Document PDF marker references an upload you do not own", + ) + + +def _derive_title(content: str) -> str: + """Derive a title from document content.""" + import re + if not isinstance(content, str): + return "Untitled" + text = content.strip() + if not text: + return "Untitled" + + # Markdown header + md = re.match(r'^#{1,3}\s+(.+)', text, re.MULTILINE) + if md: + title = md.group(1).strip() + if len(title) > 50: + title = title[:48] + "…" + return title + + # HTML heading + html = re.search(r']*>([^<]+)', text, re.IGNORECASE) + if html: + title = html.group(1).strip() + if len(title) > 50: + title = title[:48] + "…" + return title + + # First non-empty line (if short enough) + for line in text.split('\n'): + line = line.strip() + if line and 2 <= len(line) <= 60: + title = re.sub(r'[:#*`]+$', '', line).strip() + if title and len(title) > 50: + title = title[:48] + "…" + return title or "Untitled" + + return "Untitled" diff --git a/routes/document_routes.py b/routes/document_routes.py index dd13e3c60..dae8b09fa 100644 --- a/routes/document_routes.py +++ b/routes/document_routes.py @@ -1,17 +1,1810 @@ -"""Backward-compat shim — canonical location is routes/document/document_routes.py. +"""Document routes — CRUD for living documents with version history.""" -This module is replaced in ``sys.modules`` by the canonical module object so -that ``import routes.document_routes``, ``from routes.document_routes import -X``, ``importlib.import_module("routes.document_routes")``, and the -``import ... as droutes`` + ``droutes.SessionLocal = ...`` / -``monkeypatch.setattr(droutes, ...)`` pattern used by multiple tests all -operate on the *same* object the application actually uses. Keeps existing -import paths working after slice 2m (#4082/#4071). Source-introspection tests -read the canonical file by path. -""" +import uuid +import logging +from datetime import datetime, timezone +from typing import Dict, Any, List, Optional -import sys as _sys +from fastapi import APIRouter, HTTPException, Query, Request, UploadFile, File, Form -from routes.document import document_routes as _canonical # noqa: F401 +from sqlalchemy import case, func, or_ +from core.database import SessionLocal, Document, DocumentVersion +from core.database import Session as DbSession +from src.auth_helpers import get_current_user, _auth_disabled +from src.constants import MAIL_ATTACHMENTS_DIR +from src.upload_handler import reserve_upload_references -_sys.modules[__name__] = _canonical +logger = logging.getLogger(__name__) + + +def _get_session_or_404(db, session_id: str, user: Optional[str]): + session = db.query(DbSession).filter(DbSession.id == session_id).first() + if not session: + raise HTTPException(404, "Session not found") + if user and session.owner != user: + raise HTTPException(404, "Session not found") + return session + + +def _aggregate_language_facets(lang_rows): + """Sum document counts per display language for the library facet. + + NULL-language and explicit "text" rows share the "text" bucket (the + language filter treats them as one), so they must be ADDED. The old dict + comprehension keyed both to "text", silently overwriting one group and + undercounting the facet versus what the filter actually returns. + """ + out = {} + for lang, cnt in lang_rows: + key = lang or "text" + out[key] = out.get(key, 0) + cnt + return out + + +def _library_language_for_document(doc: Document) -> str: + """Return the display language used by the document library. + + PDF documents are stored as markdown wrappers so the editor can preserve + extracted text, form fields, and annotations. The library should still + identify them as PDFs instead of exposing that internal wrapper format. + """ + from src.pdf_form_doc import find_source_upload_id + + if find_source_upload_id(doc.current_content or ""): + return "pdf" + return doc.language or "text" + + +def _email_source_key(content: str) -> tuple[str, str]: + """Return the source email identity embedded in an email draft document.""" + import re + + text = content or "" + uid_m = re.search(r"(?im)^X-Source-UID:\s*(.+?)\s*$", text) + folder_m = re.search(r"(?im)^X-Source-Folder:\s*(.+?)\s*$", text) + uid = (uid_m.group(1).strip() if uid_m else "") + folder = (folder_m.group(1).strip() if folder_m else "INBOX") + return uid, folder + + +from routes.document_helpers import ( + DocumentCreate, DocumentUpdate, DocumentPatch, + _doc_to_dict, _version_to_dict, + _verify_doc_owner, _owner_session_filter, + _slug, _resolve_user_upload_path, _assert_pdf_marker_upload_owned, _derive_title, + _PDF_RENDER_SCALE, +) + + +def setup_document_routes(session_manager, upload_handler=None) -> APIRouter: + router = APIRouter(tags=["documents"]) + + def _reserve_document_uploads(user: Optional[str], content: str) -> None: + missing_id = reserve_upload_references(upload_handler, user, content) + if missing_id: + raise HTTPException( + 409, + f"Referenced upload is no longer available: {missing_id}", + ) + + def _locate_current_user_upload(request: Request, upload_id: str, user: Optional[str]): + if upload_handler is None: + return None + auth_manager = getattr(getattr(request.app, "state", None), "auth_manager", None) + return _resolve_user_upload_path(upload_handler, upload_id, user, auth_manager) + + def _load_pdf_viewer_fitz(): + from src.pdf_runtime import load_pymupdf_for_pdf_viewer + + try: + return load_pymupdf_for_pdf_viewer() + except RuntimeError as exc: + raise HTTPException(503, str(exc)) from exc + + # ---- POST /api/document ---- + @router.post("/api/document") + async def create_document(request: Request, req: DocumentCreate) -> Dict[str, Any]: + from src.auth_helpers import require_privilege + user = require_privilege(request, "can_use_documents") + db = SessionLocal() + try: + # session_id is optional: a doc can be a session-less "library" doc + # (e.g. files imported from the library) — session_id is nullable and + # the doc is owner-stamped, so it lives in the library on its own. + session = None + if req.session_id: + # Match the lenient ownership model the rest of the app uses + # (see _owner_filter): only block when an AUTHENTICATED user is + # writing into a DIFFERENT user's session. In single-user / + # unconfigured / localhost-bypass mode, falsey users preserve + # the existing lenient path. + session = _get_session_or_404(db, req.session_id, user) + + # If no language was supplied (e.g. cloning a doc whose language + # was never set), detect it from the content rather than storing + # NULL — which made the editor fall back to plain text. Defaults + # to markdown for prose. + language = req.language + if not language: + from src.agent_tools.document_tools import _looks_like_email_document, _sniff_doc_language, _coerce_email_document_content + language = _sniff_doc_language(req.content) + else: + from src.agent_tools.document_tools import _looks_like_email_document, _coerce_email_document_content + if _looks_like_email_document(req.content, req.title): + language = "email" + + _reserve_document_uploads(user, req.content) + _assert_pdf_marker_upload_owned(request, req.content, user, upload_handler) + + # Reply drafts are keyed to the source email. If a UI/tool path tries + # to create a second draft for the same email in the same chat, + # update the existing draft instead so quoted thread history stays + # attached to the visible document. + if language == "email" and req.session_id: + source_uid, source_folder = _email_source_key(req.content) + if source_uid: + candidates = ( + db.query(Document) + .filter(Document.session_id == req.session_id) + .filter(Document.is_active == True) + .filter(Document.language == "email") + .order_by(Document.updated_at.desc()) + .limit(25) + .all() + ) + for existing in candidates: + old_uid, old_folder = _email_source_key(existing.current_content or "") + if old_uid != source_uid or old_folder != source_folder: + continue + merged = _coerce_email_document_content(existing.current_content or "", req.content) + if existing.current_content != merged: + new_ver = (existing.version_count or 1) + 1 + existing.current_content = merged + existing.title = req.title or existing.title + existing.version_count = new_ver + db.add(DocumentVersion( + id=str(uuid.uuid4()), + document_id=existing.id, + version_number=new_ver, + content=merged, + summary="Updated existing email draft", + source="user", + )) + db.commit() + db.refresh(existing) + return _doc_to_dict(existing) + + doc_id = str(uuid.uuid4()) + ver_id = str(uuid.uuid4()) + + doc = Document( + id=doc_id, + session_id=req.session_id, + title=req.title, + language=language, + current_content=req.content, + version_count=1, + is_active=True, + # Stamp ownership directly so the doc survives its session + # being deleted. Fall back to the session's owner when the + # request is unauthenticated (single-user / localhost bypass). + owner=user or (session.owner if session else None), + ) + ver = DocumentVersion( + id=ver_id, + document_id=doc_id, + version_number=1, + content=req.content, + summary="Initial version", + source="user", + ) + db.add(doc) + db.add(ver) + db.commit() + db.refresh(doc) + try: + from src.event_bus import fire_event + fire_event("document_created", doc.owner) + except Exception: + logger.debug("document_created event dispatch failed", exc_info=True) + return _doc_to_dict(doc) + except HTTPException: + raise + except Exception as e: + db.rollback() + logger.error(f"Failed to create document: {e}") + raise HTTPException(500, f"Failed to create document: {e}") + finally: + db.close() + + # ---- POST /api/documents/import-pdf ---- + @router.post("/api/documents/import-pdf") + async def import_pdf( + request: Request, + file: UploadFile = File(...), + session_id: Optional[str] = Form(None), + ) -> Dict[str, Any]: + """Upload a PDF and create the matching Document. + + Detects AcroForm fields — if any, creates a form-backed markdown doc + (clickable inputs in the PDF view). Otherwise creates a plain PDF doc + with a `pdf_source` marker so the viewer renders the pages without + overlays. + """ + from src.pdf_forms import has_form_fields, extract_fields + from src.pdf_form_doc import ( + save_field_sidecar, + create_form_markdown_document, + create_plain_pdf_document, + ) + from src.document_processor import _process_pdf, strip_pdf_content_marker + import os + + from src.auth_helpers import require_privilege + user = require_privilege(request, "can_use_documents") + + # session_id is optional — a library import isn't tied to a chat. When + # given, validate it; otherwise the PDF becomes a session-less library + # doc (the doc creators below already handle a missing session). + if session_id: + db = SessionLocal() + try: + _get_session_or_404(db, session_id, user) + finally: + db.close() + + if upload_handler is None: + raise HTTPException(500, "Upload handler not configured") + + client_ip = request.client.host if request.client else "unknown" + try: + meta = upload_handler.save_upload(file, client_ip, owner=user) + except HTTPException: + raise + except Exception as e: + logger.error(f"PDF import save_upload failed: {e}") + raise HTTPException(500, f"Upload failed: {e}") + + upload_id = meta["id"] + pdf_path = _locate_current_user_upload(request, upload_id, user) + if not pdf_path: + raise HTTPException(500, "Saved PDF could not be located") + + title = os.path.splitext(meta.get("original_name") or meta.get("name") or upload_id)[0] + try: + body_text = strip_pdf_content_marker(_process_pdf(pdf_path, owner=user)) + except Exception: + body_text = None + + is_form = False + try: + is_form = has_form_fields(pdf_path) + except Exception as e: + logger.warning(f"has_form_fields failed for {pdf_path}: {e}") + + if is_form: + fields = extract_fields(pdf_path) + save_field_sidecar(pdf_path, fields) + doc_id = create_form_markdown_document( + session_id=session_id, + fields=fields, + upload_id=upload_id, + title=title, + intro_text=body_text, + ) + else: + doc_id = create_plain_pdf_document( + session_id=session_id, + upload_id=upload_id, + title=title, + body_text=body_text, + ) + + if not doc_id: + raise HTTPException(500, "Failed to create document for PDF") + + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(500, "Created document not found") + # The PDF doc creators stamp owner from the session only; a + # session-less library import leaves owner NULL, which the Library's + # owner filter then hides. Stamp the requesting user so it shows. + if not doc.owner and user: + doc.owner = user + db.commit() + db.refresh(doc) + return _doc_to_dict(doc) + finally: + db.close() + + # ---- GET /api/documents/library ---- + @router.get("/api/documents/library") + async def documents_library( + request: Request, + search: Optional[str] = Query(None), + language: Optional[str] = Query(None), + sort: str = Query("recent"), + offset: int = Query(0, ge=0), + limit: int = Query(20, ge=1, le=50), + archived: bool = Query(False), + ) -> Dict[str, Any]: + user = get_current_user(request) + db = SessionLocal() + try: + from sqlalchemy import or_ + pdf_marker_cond = or_( + Document.current_content.like('%\s*\n+#[^\n]*\n+)', re.MULTILINE) + head_match = head_re.match(content) + head = head_match.group(1) if head_match else (content.splitlines()[0] + "\n\n# " + (doc.title or "PDF") + "\n\n") + doc.current_content = head + body_text.strip() + "\n" + doc.version_count = (doc.version_count or 1) + 1 + db.add(DocumentVersion( + id=str(__import__("uuid").uuid4()), + document_id=doc_id, + version_number=doc.version_count, + content=doc.current_content, + summary="PDF text re-extracted (OCR)", + source="ocr", + )) + db.commit() + return {"ok": True, "id": doc_id, "extracted": True, "chars": len(body_text)} + finally: + db.close() + + # ---- POST /api/documents/export-zip — bundle selected docs into a .zip ---- + @router.post("/api/documents/export-zip") + async def documents_export_zip(request: Request): + """Zip the selected documents (each as a text file with the right + extension) — mirrors the gallery's bulk download-zip so multi-export + is one file instead of a blocked flood of individual downloads.""" + user = get_current_user(request) + try: + data = await request.json() + except Exception as e: + logger.warning("Failed to parse export request body, defaulting to empty", exc_info=e) + data = {} + ids = data.get("ids") or [] + if not ids: + raise HTTPException(400, "No documents specified") + _ext = { + "javascript": ".js", "python": ".py", "html": ".html", "css": ".css", + "markdown": ".md", "json": ".json", "yaml": ".yml", "bash": ".sh", + "sql": ".sql", "rust": ".rs", "go": ".go", "java": ".java", "c": ".c", + "cpp": ".cpp", "typescript": ".ts", "ruby": ".rb", "php": ".php", + "text": ".txt", "xml": ".xml", "toml": ".toml", "ini": ".ini", + } + db = SessionLocal() + try: + import io + import re + import zipfile + from fastapi import Response + docs = db.query(Document).filter(Document.id.in_(ids)).all() + buf = io.BytesIO() + used = set() + wrote = 0 + with zipfile.ZipFile(buf, "w", zipfile.ZIP_DEFLATED) as zf: + for doc in docs: + try: + _verify_doc_owner(db, doc, user) + except HTTPException: + continue # skip docs the user doesn't own + ext = _ext.get(doc.language or "text", ".txt") + base = (doc.title or "document").strip() or "document" + base = re.sub(r"[^\w\-. ]+", "", base)[:60].strip() or doc.id + name = base if "." in base else base + ext + i = 1 + while name in used: + name = f"{base}-{i}" + ("" if "." in base else ext) + i += 1 + used.add(name) + zf.writestr(name, doc.current_content or "") + wrote += 1 + if not wrote: + raise HTTPException(404, "No documents found") + return Response( + content=buf.getvalue(), + media_type="application/zip", + headers={"Content-Disposition": 'attachment; filename="documents.zip"'}, + ) + finally: + db.close() + + # ---- PUT /api/document/{doc_id} — user manual edit ---- + # Coalesce window: if the last user version was saved within this many + # seconds, update it in-place (user is still actively editing). + # Once the gap exceeds this, the next save creates a new version. + VERSION_COALESCE_SECONDS = 60 + + @router.put("/api/document/{doc_id}") + async def update_document(request: Request, doc_id: str, req: DocumentUpdate) -> Dict[str, Any]: + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + + incoming_content = req.content + from src.agent_tools.document_tools import _coerce_email_document_content, _looks_like_email_document + is_email_doc = ( + (doc.language or "").lower() == "email" + or _looks_like_email_document(doc.current_content or "", doc.title or "") + or _looks_like_email_document(req.content or "", doc.title or "") + ) + if is_email_doc: + incoming_content = _coerce_email_document_content(doc.current_content or "", req.content) + doc.language = "email" + + # Skip if content is identical unless the caller explicitly wants + # a checkpoint version from the current editor state. + if doc.current_content == incoming_content and not req.force_version: + return _doc_to_dict(doc) + + _reserve_document_uploads(user, incoming_content) + _assert_pdf_marker_upload_owned(request, incoming_content, user, upload_handler) + + # Check if we can coalesce with the latest version + latest_ver = db.query(DocumentVersion).filter( + DocumentVersion.document_id == doc_id, + ).order_by(DocumentVersion.version_number.desc()).first() + + now = datetime.now(timezone.utc) + coalesced = False + if latest_ver and latest_ver.source == "user" and not req.force_version: + ver_time = latest_ver.created_at + if ver_time.tzinfo is None: + ver_time = ver_time.replace(tzinfo=timezone.utc) + age = (now - ver_time).total_seconds() + if age < VERSION_COALESCE_SECONDS: + # Update the existing version in-place + latest_ver.content = incoming_content + latest_ver.created_at = now + if req.summary: + latest_ver.summary = req.summary + coalesced = True + + if not coalesced: + new_ver = doc.version_count + 1 + ver = DocumentVersion( + id=str(uuid.uuid4()), + document_id=doc_id, + version_number=new_ver, + content=incoming_content, + summary=req.summary or "Manual edit", + source="user", + ) + doc.version_count = new_ver + db.add(ver) + + doc.current_content = incoming_content + db.commit() + db.refresh(doc) + return _doc_to_dict(doc) + except HTTPException: + raise + except Exception as e: + db.rollback() + raise HTTPException(500, f"Failed to update document: {e}") + finally: + db.close() + + # ---- PATCH /api/document/{doc_id} — metadata only ---- + @router.patch("/api/document/{doc_id}") + async def patch_document(request: Request, doc_id: str, req: DocumentPatch) -> Dict[str, Any]: + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + if req.title is not None: + doc.title = req.title + if req.language is not None: + doc.language = req.language + if req.session_id is not None: + # Empty string = unlink from session + if req.session_id: + _get_session_or_404(db, req.session_id, user) + doc.session_id = req.session_id if req.session_id else None + if not req.session_id: + # Tab closed / doc detached from its session — drop the + # in-memory active-doc pointer so the last-resort injection + # path doesn't re-surface this doc in a later chat (#1160). + try: + from src.agent_tools.document_tools import clear_active_document + clear_active_document(doc_id) + except Exception as e: + logger.warning("Failed to clear active document %r on detach", doc_id, exc_info=e) + db.commit() + db.refresh(doc) + return _doc_to_dict(doc) + except HTTPException: + raise + except Exception as e: + db.rollback() + raise HTTPException(500, str(e)) + finally: + db.close() + + # ---- DELETE /api/document/{doc_id} — soft delete ---- + @router.delete("/api/document/{doc_id}") + async def delete_document(request: Request, doc_id: str) -> Dict[str, str]: + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + doc.is_active = False + # Closed/deleted — drop the in-memory active-doc pointer so it isn't + # re-injected into a later, unrelated chat (#1160). + try: + from src.agent_tools.document_tools import clear_active_document + clear_active_document(doc_id) + except Exception: + pass + db.commit() + return {"status": "deleted", "id": doc_id} + except HTTPException: + raise + except Exception as e: + db.rollback() + raise HTTPException(500, str(e)) + finally: + db.close() + + # ---- GET /api/document/{doc_id}/versions ---- + @router.get("/api/document/{doc_id}/versions") + async def list_versions(request: Request, doc_id: str) -> List[Dict[str, Any]]: + user = get_current_user(request) + db = SessionLocal() + try: + # Verify ownership before listing versions + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + versions = db.query(DocumentVersion).filter( + DocumentVersion.document_id == doc_id + ).order_by(DocumentVersion.version_number.desc()).all() + return [{ + "id": v.id, + "version_number": v.version_number, + "content": v.content, + "summary": v.summary, + "source": v.source, + "created_at": v.created_at.isoformat() if v.created_at else None, + } for v in versions] + finally: + db.close() + + # ---- GET /api/document/{doc_id}/version/{num} ---- + @router.get("/api/document/{doc_id}/version/{num}") + async def get_version(request: Request, doc_id: str, num: int) -> Dict[str, Any]: + user = get_current_user(request) + db = SessionLocal() + try: + # Verify ownership + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + ver = db.query(DocumentVersion).filter( + DocumentVersion.document_id == doc_id, + DocumentVersion.version_number == num, + ).first() + if not ver: + raise HTTPException(404, "Version not found") + return _version_to_dict(ver) + finally: + db.close() + + # ---- POST /api/document/{doc_id}/restore/{num} ---- + @router.post("/api/document/{doc_id}/restore/{num}") + async def restore_version(request: Request, doc_id: str, num: int) -> Dict[str, Any]: + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + + old_ver = db.query(DocumentVersion).filter( + DocumentVersion.document_id == doc_id, + DocumentVersion.version_number == num, + ).first() + if not old_ver: + raise HTTPException(404, "Version not found") + + new_ver_num = doc.version_count + 1 + ver = DocumentVersion( + id=str(uuid.uuid4()), + document_id=doc_id, + version_number=new_ver_num, + content=old_ver.content, + summary=f"Restored from v{num}", + source="user", + ) + doc.current_content = old_ver.content + doc.version_count = new_ver_num + db.add(ver) + db.commit() + db.refresh(doc) + return _doc_to_dict(doc) + except HTTPException: + raise + except Exception as e: + db.rollback() + raise HTTPException(500, str(e)) + finally: + db.close() + + # ---- POST /api/documents/tidy — clean up broken/empty documents ---- + @router.post("/api/documents/tidy") + async def tidy_documents(request: Request) -> Dict[str, Any]: + """Fix empty titles and remove broken/empty documents (user's docs only).""" + user = get_current_user(request) + db = SessionLocal() + try: + q = ( + db.query(Document) + .outerjoin(DbSession, Document.session_id == DbSession.id) + .filter(Document.is_active == True) + .filter((Document.archived == False) | (Document.archived.is_(None))) + ) + q = _owner_session_filter(q, user) + docs = q.all() + fixed_titles = 0 + deleted = 0 + + # Same junk-detection logic as the scheduled tidy_documents + # action (src/document_actions.py). Keep these two in sync. + import re as _re + from src.document_actions import _JUNK_TITLES + + to_delete = [] + now = datetime.now(timezone.utc) + for doc in docs: + created = doc.created_at + if created and created.tzinfo is None: + created = created.replace(tzinfo=timezone.utc) + + # Skip freshly created documents to avoid deleting them while the user is actively editing + if created and (now - created).total_seconds() < 900: # 15 minutes + continue + + content = (doc.current_content or "").strip() + title_raw = (doc.title or "").strip() + title = title_raw.lower() + is_fresh_empty = ( + not content + and created is not None + and (now - created).total_seconds() < 1800 + ) + if is_fresh_empty: + continue + + # Strip markdown noise to get a "real" character count + stripped = _re.sub(r"^#{1,6}\s+", "", content, flags=_re.MULTILINE) + stripped = _re.sub(r"[*_`>\-=]+", "", stripped) + stripped = _re.sub(r"\s+", " ", stripped).strip() + real_len = len(stripped) + + # Detect email-scaffold stubs: "To: \nSubject: \n---\n" style + # bodies with nothing typed in. Stub = every meaningful line + # is a header label (To:/From:/Subject:/...) with no real + # value (blank, "empty", "(empty)", "-", "none", "n/a"). + _is_email_stub = False + _HEADER_RE = _re.compile(r"^(to|from|cc|bcc|subject|reply-to):\s*(.*)$", _re.I) + _PLACEHOLDER_VALS = {"", "empty", "(empty)", "-", "—", "none", "n/a", "na", "tbd"} + if title in ("new email", "new mail", "new message") or doc.language == "email": + body_lines = [ln.strip() for ln in content.split("\n") + if ln.strip() and ln.strip() != "---"] + def _is_filler(ln): + m = _HEADER_RE.match(ln) + if not m: + return False + val = (m.group(2) or "").strip().lower() + return val in _PLACEHOLDER_VALS + has_real_body = any(not _is_filler(ln) for ln in body_lines) + if body_lines and not has_real_body: + _is_email_stub = True + + # Hard-delete obviously empty / junk documents + if not content or content in ("", "# Untitled"): + to_delete.append(doc); deleted += 1; continue + if _is_email_stub: + to_delete.append(doc); deleted += 1; continue + if title in _JUNK_TITLES: + to_delete.append(doc); deleted += 1; continue + + # Fix empty or placeholder titles on survivors + if not title_raw or title_raw == "Untitled": + new_title = _derive_title(content) + if new_title and new_title != "Untitled": + doc.title = new_title + fixed_titles += 1 + + for doc in to_delete: + db.delete(doc) + + # Also clean up inactive empty docs from previous soft-deletes + inactive_q = ( + db.query(Document) + .outerjoin(DbSession, Document.session_id == DbSession.id) + .filter(Document.is_active == False) + .filter((Document.current_content == None) | (Document.current_content == "")) + ) + inactive_q = _owner_session_filter(inactive_q, user) + inactive_docs = inactive_q.all() + for doc in inactive_docs: + db.delete(doc) + deleted += len(inactive_docs) + + db.commit() + return { + "fixed_titles": fixed_titles, + "deleted": deleted, + "message": f"Fixed {fixed_titles} title{'s' if fixed_titles != 1 else ''}, removed {deleted} empty document{'s' if deleted != 1 else ''}", + } + except Exception as e: + db.rollback() + logger.error(f"Document tidy failed: {e}") + raise HTTPException(500, f"Tidy failed: {e}") + finally: + db.close() + + # ---- POST /api/documents/ai-tidy — AI-powered cleanup of junk/test documents ---- + @router.post("/api/documents/ai-tidy") + async def ai_tidy_documents(request: Request) -> Dict[str, Any]: + """Use AI to judge if documents are junk/test/accidental, then delete them. + Caches verdicts so previously-reviewed docs are skipped.""" + from src.task_endpoint import resolve_task_endpoint + from src.endpoint_resolver import resolve_endpoint + from src.llm_core import llm_call_async + + user = get_current_user(request) + url, model, headers = resolve_task_endpoint(owner=user or None) + if not url or not model: + # Fall back to default endpoint + url, model, headers = resolve_endpoint("default", owner=user or None) + if not url or not model: + raise HTTPException(500, "No endpoint configured for AI tidy") + + db = SessionLocal() + try: + q = ( + db.query(Document) + .outerjoin(DbSession, Document.session_id == DbSession.id) + .filter(Document.is_active == True) + .filter((Document.archived == False) | (Document.archived.is_(None))) + ) + q = _owner_session_filter(q, user) + docs = q.all() + + # Only review docs that haven't been reviewed yet + to_review = [d for d in docs if not d.tidy_verdict] + if not to_review: + return {"deleted": 0, "reviewed": 0, "message": "All documents already reviewed"} + + # Build a batch prompt — review up to 30 at a time + batch = to_review[:30] + doc_list = [] + for i, doc in enumerate(batch): + preview = (doc.current_content or "")[:300].strip() + doc_list.append(f"[{i}] title=\"{doc.title}\" lang={doc.language or 'text'} content_preview=\"{preview}\"") + + prompt = ( + "You are a document library cleaner. For each document below, decide if it is JUNK " + "(test, accidental, placeholder, empty-ish, tool-test, throwaway) or KEEP (real content worth saving).\n\n" + "Respond with ONLY a JSON array of verdicts, one per document, like: [\"junk\",\"keep\",\"junk\",...]\n" + "No explanation, no markdown, just the JSON array.\n\n" + + "\n".join(doc_list) + ) + + response = await llm_call_async( + url, model, + [{"role": "system", "content": "You classify documents as junk or keep. Respond only with a JSON array."}, + {"role": "user", "content": prompt}], + temperature=0.1, + max_tokens=200, + headers=headers, + timeout=30, + ) + + # Parse verdicts + import re + match = re.search(r'\[.*?\]', response, re.DOTALL) + if not match: + raise HTTPException(500, "AI returned invalid response") + + import json as _json + verdicts = _json.loads(match.group()) + + deleted = 0 + reviewed = 0 + for i, doc in enumerate(batch): + if i >= len(verdicts): + break + verdict = str(verdicts[i] or "").lower().strip() + if verdict == "junk": + doc.tidy_verdict = "junk" + db.delete(doc) + deleted += 1 + else: + doc.tidy_verdict = "keep" + reviewed += 1 + + db.commit() + return { + "deleted": deleted, + "reviewed": reviewed, + "remaining": len(to_review) - len(batch), + "message": f"Reviewed {reviewed}, removed {deleted} junk document{'s' if deleted != 1 else ''}", + } + except HTTPException: + raise + except Exception as e: + db.rollback() + logger.error(f"AI tidy failed: {e}") + raise HTTPException(500, f"AI tidy failed: {e}") + finally: + db.close() + + # ---- POST /api/document/{doc_id}/export-pdf/preview ---- + @router.post("/api/document/{doc_id}/export-pdf/preview") + async def export_pdf_preview(doc_id: str, request: Request) -> Dict[str, Any]: + """Return the field-value mapping that would be written to the PDF. + + Frontend shows this in a confirmation modal so the user can spot/fix + any wrong values before triggering the actual download. + """ + from src.pdf_form_doc import find_source_upload_id, parse_markdown_to_values, load_field_sidecar + + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + + upload_id = find_source_upload_id(doc.current_content or "") + if not upload_id: + raise HTTPException(400, "Document is not linked to a source PDF") + + pdf_path = _locate_current_user_upload(request, upload_id, user) + if not pdf_path: + raise HTTPException(404, f"Source PDF {upload_id} not found in uploads") + + fields = load_field_sidecar(pdf_path) + if not fields: + raise HTTPException(404, "Field schema sidecar missing for source PDF") + + values = parse_markdown_to_values(doc.current_content or "") + field_meta = {f["name"]: f for f in fields} + + preview = [] + for name, current in values.items(): + meta = field_meta.get(name) + if not meta: + continue + preview.append({ + "name": name, + "label": meta.get("label") or name, + "type": meta.get("type"), + "options": meta.get("options") or [], + "page": meta.get("page"), + "value": current, + }) + + unknown = [ + name for name in values + if name not in field_meta + ] + return { + "doc_id": doc_id, + "upload_id": upload_id, + "fields": preview, + "unknown_fields": unknown, + "total": len(fields), + "filled": sum(1 for p in preview if p["value"] not in ("", False, None)), + } + finally: + db.close() + + # ---- GET /api/document/{doc_id}/render-pages ---- + @router.get("/api/document/{doc_id}/render-pages") + async def render_pages(doc_id: str, request: Request) -> Dict[str, Any]: + """Return per-page metadata for the interactive PDF view. + + Each page entry has its rendered-image dimensions (matching what + /page/{n}.png returns at the same DPI) plus the list of form fields + on that page with their rects translated to image-pixel coordinates. + Frontend overlays HTML form controls at those positions. + """ + from src.pdf_form_doc import find_source_upload_id, parse_markdown_to_values, load_field_sidecar + + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + upload_id = find_source_upload_id(doc.current_content or "") + if not upload_id: + raise HTTPException(400, "Document is not linked to a source PDF") + pdf_path = _locate_current_user_upload(request, upload_id, user) + if not pdf_path: + raise HTTPException(404, f"Source PDF {upload_id} not found") + + fitz = _load_pdf_viewer_fitz() + schema = load_field_sidecar(pdf_path) or [] + values = parse_markdown_to_values(doc.current_content or "") + + # Group fields by page + by_page: Dict[int, list] = {} + for f in schema: + by_page.setdefault(f["page"], []).append(f) + + scale = _PDF_RENDER_SCALE + pdf_doc = fitz.open(pdf_path) + try: + pages_out = [] + for page_index in range(pdf_doc.page_count): + page = pdf_doc[page_index] + page_no = page_index + 1 + pw, ph = page.rect.width, page.rect.height + img_w = int(pw * scale) + img_h = int(ph * scale) + fields_out = [] + for f in by_page.get(page_no, []): + x0, y0, x1, y1 = f["rect"] + fields_out.append({ + "name": f["name"], + "type": f["type"], + "label": f.get("label") or "", + "options": f.get("options") or [], + "value": values.get(f["name"], f.get("value", "")), + "rect_px": [ + int(x0 * scale), int(y0 * scale), + int(x1 * scale), int(y1 * scale), + ], + }) + pages_out.append({ + "page": page_no, + "width": img_w, + "height": img_h, + "fields": fields_out, + }) + return {"doc_id": doc_id, "scale": scale, "pages": pages_out} + finally: + pdf_doc.close() + finally: + db.close() + + # ---- GET /api/document/{doc_id}/page/{n}.png ---- + @router.get("/api/document/{doc_id}/page/{page_no}.png") + async def render_page_png(doc_id: str, page_no: int, request: Request): + """Render one page of the source PDF as a PNG (no values stamped — the + frontend overlays HTML form inputs on top).""" + from fastapi.responses import Response + from src.pdf_form_doc import find_source_upload_id + + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + upload_id = find_source_upload_id(doc.current_content or "") + if not upload_id: + raise HTTPException(400, "Document is not linked to a source PDF") + pdf_path = _locate_current_user_upload(request, upload_id, user) + if not pdf_path: + raise HTTPException(404, "Source PDF not found") + finally: + db.close() + + fitz = _load_pdf_viewer_fitz() + pdf_doc = fitz.open(pdf_path) + try: + if page_no < 1 or page_no > pdf_doc.page_count: + raise HTTPException(404, "Page out of range") + page = pdf_doc[page_no - 1] + mat = fitz.Matrix(_PDF_RENDER_SCALE, _PDF_RENDER_SCALE) + pix = page.get_pixmap(matrix=mat, alpha=False) + png_bytes = pix.tobytes("png") + return Response( + content=png_bytes, + media_type="image/png", + headers={"Cache-Control": "public, max-age=3600"}, + ) + finally: + pdf_doc.close() + + # ---- POST /api/document/{doc_id}/ai-fill-annotations ---- + @router.post("/api/document/{doc_id}/ai-fill-annotations") + async def ai_fill_annotations(doc_id: str, request: Request) -> Dict[str, Any]: + """Ask a vision-capable LLM to locate fillable areas on a flat PDF and + propose annotation values for each, given a free-form user instruction. + + Returns a list of annotations: [{page, x, y, w, h, value}] where x/y/w/h + are page-percentages (0–100) — same coordinate system as the freeform + annotations the frontend already renders. + """ + import base64 + import json + import fitz + from src.pdf_form_doc import find_source_upload_id + from src.document_processor import _resolve_vl_model, _load_vl_settings + from src.llm_core import llm_call_async + + body = await request.json() if request.headers.get("content-type", "").startswith("application/json") else {} + instruction = (body or {}).get("instruction", "").strip() + if not instruction: + raise HTTPException(400, "instruction is required") + + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + upload_id = find_source_upload_id(doc.current_content or "") + if not upload_id: + raise HTTPException(400, "Document is not linked to a source PDF") + pdf_path = _locate_current_user_upload(request, upload_id, user) + if not pdf_path: + raise HTTPException(404, "Source PDF not found") + finally: + db.close() + + # Resolve VL model (admin-configured or auto-detected vision-capable) + settings = _load_vl_settings() + vl_model = settings.get("vision_model", "") + try: + url, model_id, headers = _resolve_vl_model(vl_model, owner=user) + except Exception as e: + raise HTTPException(503, f"No vision model available: {e}") + + system_prompt = ( + "You analyze rendered PDF page images and propose values to fill in. " + "For each blank line, box, underscore, or labeled space on the page that " + "should be filled given the user's instruction, output one annotation. " + "Coordinates are percentages (0-100) of the page width/height with the " + "origin at top-left. Width/height should match the visible blank box. " + "Return ONLY a JSON array, no prose, no markdown fences. Each entry: " + '{"x": number, "y": number, "w": number, "h": number, "value": string}. ' + "If a region should not be filled, omit it. If nothing should be filled, " + "return []." + ) + + all_annotations = [] + pdf_doc = fitz.open(pdf_path) + try: + for page_index in range(pdf_doc.page_count): + page = pdf_doc[page_index] + mat = fitz.Matrix(_PDF_RENDER_SCALE, _PDF_RENDER_SCALE) + pix = page.get_pixmap(matrix=mat, alpha=False) + png_bytes = pix.tobytes("png") + b64 = base64.b64encode(png_bytes).decode("ascii") + + messages = [ + {"role": "system", "content": system_prompt}, + { + "role": "user", + "content": [ + { + "type": "text", + "text": ( + f"User instruction:\n{instruction}\n\n" + f"This is page {page_index + 1} of {pdf_doc.page_count}. " + "Return JSON array of annotations to add to this page." + ), + }, + { + "type": "image_url", + "image_url": {"url": f"data:image/png;base64,{b64}"}, + }, + ], + }, + ] + try: + raw = await llm_call_async( + url, model_id, messages, + temperature=0.1, max_tokens=2000, headers=headers, + ) + except Exception as e: + logger.error(f"VL call failed on page {page_index + 1}: {e}") + continue + + raw = (raw or "").strip() + if raw.startswith("```"): + raw = raw.split("\n", 1)[-1].rsplit("```", 1)[0].strip() + try: + parsed = json.loads(raw) + except Exception: + logger.warning(f"AI fill: page {page_index + 1} returned non-JSON: {raw[:200]}") + continue + if not isinstance(parsed, list): + continue + for item in parsed: + if not isinstance(item, dict): + continue + try: + x = float(item.get("x", 0)) + y = float(item.get("y", 0)) + w = float(item.get("w", 0)) + h = float(item.get("h", 0)) + value = str(item.get("value", "") or "") + except Exception: + continue + # Clamp + reject zero-size entries + if w <= 0.5 or h <= 0.3: + continue + x = max(0.0, min(99.0, x)) + y = max(0.0, min(99.0, y)) + w = max(0.5, min(100.0 - x, w)) + h = max(0.3, min(100.0 - y, h)) + if not value.strip(): + continue + all_annotations.append({ + "page": page_index + 1, + "x": round(x, 2), + "y": round(y, 2), + "w": round(w, 2), + "h": round(h, 2), + "value": value, + }) + finally: + pdf_doc.close() + + return {"annotations": all_annotations} + + # ---- GET /api/document/{doc_id}/render-pdf ---- + @router.get("/api/document/{doc_id}/render-pdf") + async def render_pdf(doc_id: str, request: Request): + """Inline PDF preview filled with the current markdown values. + + Same plumbing as the export route, but no signature stamping and + served inline (Content-Disposition: inline) so the browser can + embed it in an iframe. Cache-busted by the caller via query string. + """ + import base64 + import os + import tempfile + from fastapi.responses import FileResponse + from starlette.background import BackgroundTask + from src.pdf_form_doc import find_source_upload_id, parse_markdown_to_values, parse_markdown_annotations + from src.pdf_forms import fill_fields, stamp_annotations + from core.database import Signature + + # Track temp files for this request so they get unlinked AFTER + # the response is fully sent (BackgroundTask runs post-send). + _to_unlink: list[str] = [] + def _cleanup_temps(): + for _p in _to_unlink: + try: + os.unlink(_p) + except FileNotFoundError: + pass + except Exception as _e: + logger.warning(f"Could not unlink temp PDF {_p}: {_e}") + + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + upload_id = find_source_upload_id(doc.current_content or "") + if not upload_id: + raise HTTPException(400, "Document is not linked to a source PDF") + pdf_path = _locate_current_user_upload(request, upload_id, user) + if not pdf_path: + raise HTTPException(404, f"Source PDF {upload_id} not found") + + # Fail fast with a clear 503 if the optional PyMuPDF dependency + # is missing — fill_fields/stamp_annotations will otherwise + # raise RuntimeError deep inside and bubble out as a 500. + # Mirrors the convention in _load_pdf_viewer_fitz above. + _load_pdf_viewer_fitz() + + values = parse_markdown_to_values(doc.current_content or "") + out_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name + _to_unlink.append(out_path) + try: + fill_fields(pdf_path, out_path, values) + except Exception as e: + logger.error(f"render_pdf fill_fields failed for {doc_id}: {e}") + _cleanup_temps() + raise HTTPException(500, f"PDF render failed: {e}") + + annotations = parse_markdown_annotations(doc.current_content or "") + if annotations: + ann_sig_ids = [ + a["value"][len("signature:"):].strip() + for a in annotations + if a.get("kind") == "signature" + and isinstance(a.get("value"), str) + and a["value"].startswith("signature:") + ] + ann_signature_pngs: dict[str, bytes] = {} + if ann_sig_ids: + # SECURITY: filter by owner so a caller can't reference + # someone else's signature ID from doc markdown and have + # it stamped/exported. + _sig_q = db.query(Signature).filter(Signature.id.in_(ann_sig_ids)) + if user: + _sig_q = _sig_q.filter(Signature.owner == user) + sig_rows = _sig_q.all() + for s in sig_rows: + try: + ann_signature_pngs[s.id] = base64.b64decode(s.data_png) + except Exception as e: + logger.warning(f"Bad annotation signature data for {s.id}: {e}") + annotated_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name + _to_unlink.append(annotated_path) + try: + stamp_annotations(out_path, annotated_path, annotations, ann_signature_pngs) + out_path = annotated_path + except Exception as e: + logger.error(f"stamp_annotations (render) failed for {doc_id}: {e}") + + return FileResponse( + out_path, + media_type="application/pdf", + headers={"Content-Disposition": "inline"}, + background=BackgroundTask(_cleanup_temps), + ) + finally: + db.close() + + # ---- GET /api/document/{doc_id}/export-pdf ---- + @router.get("/api/document/{doc_id}/export-pdf") + async def export_pdf(doc_id: str, request: Request): + """Stream the filled PDF for download. + + Reads field values and signature selections from the markdown — there + is no separate confirmation step. Signature fields contain their + chosen signature ID encoded as `signature:` in the value. + """ + import base64 + import os + import tempfile + from fastapi.responses import FileResponse + from starlette.background import BackgroundTask + from src.pdf_form_doc import find_source_upload_id, parse_markdown_to_values, load_field_sidecar, parse_markdown_annotations + from src.pdf_forms import fill_fields, stamp_signatures, stamp_annotations + from core.database import Signature + + _to_unlink: list[str] = [] + def _cleanup_temps(): + for _p in _to_unlink: + try: + os.unlink(_p) + except FileNotFoundError: + pass + except Exception as _e: + logger.warning(f"Could not unlink temp PDF {_p}: {_e}") + + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + + upload_id = find_source_upload_id(doc.current_content or "") + if not upload_id: + raise HTTPException(400, "Document is not linked to a source PDF") + + pdf_path = _locate_current_user_upload(request, upload_id, user) + if not pdf_path: + raise HTTPException(404, f"Source PDF {upload_id} not found in uploads") + + schema = load_field_sidecar(pdf_path) or [] + sig_field_names = {f["name"] for f in schema if f.get("type") == "signature"} + + all_values = parse_markdown_to_values(doc.current_content or "") + # Split: signature fields go to stamps, everything else to fill_fields + text_values: dict = {} + sig_ids: dict[str, str] = {} + for name, raw in all_values.items(): + if name in sig_field_names and isinstance(raw, str) and raw.startswith("signature:"): + sig_ids[name] = raw[len("signature:"):].strip() + elif name not in sig_field_names: + text_values[name] = raw + + stamps: dict = {} + if sig_ids: + # SECURITY: filter by owner — same reason as render_pdf. + _sig_q2 = db.query(Signature).filter(Signature.id.in_(list(sig_ids.values()))) + if user: + _sig_q2 = _sig_q2.filter(Signature.owner == user) + rows = _sig_q2.all() + by_id = {s.id: s for s in rows} + for field_name, sid in sig_ids.items(): + s = by_id.get(sid) + if not s: + continue + try: + stamps[field_name] = base64.b64decode(s.data_png) + except Exception as e: + logger.warning(f"Bad signature data for {sid}: {e}") + + filled_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name + _to_unlink.append(filled_path) + try: + fill_fields(pdf_path, filled_path, text_values) + except Exception as e: + logger.error(f"fill_fields failed for doc {doc_id}: {e}") + _cleanup_temps() + raise HTTPException(500, f"PDF fill failed: {e}") + + out_path = filled_path + if stamps: + stamped_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name + _to_unlink.append(stamped_path) + try: + stamp_signatures(filled_path, stamped_path, stamps) + out_path = stamped_path + except Exception as e: + logger.error(f"stamp_signatures failed for doc {doc_id}: {e}") + + # Burn freeform annotations (Text/Check/Sign drops) on top. + annotations = parse_markdown_annotations(doc.current_content or "") + if annotations: + # Resolve any signature annotations to their PNG bytes. + ann_sig_ids = [ + a["value"][len("signature:"):].strip() + for a in annotations + if a.get("kind") == "signature" + and isinstance(a.get("value"), str) + and a["value"].startswith("signature:") + ] + ann_signature_pngs: dict[str, bytes] = {} + if ann_sig_ids: + # SECURITY: filter by owner so a caller can't reference + # someone else's signature ID from doc markdown and have + # it stamped/exported. + _sig_q = db.query(Signature).filter(Signature.id.in_(ann_sig_ids)) + if user: + _sig_q = _sig_q.filter(Signature.owner == user) + sig_rows = _sig_q.all() + for s in sig_rows: + try: + ann_signature_pngs[s.id] = base64.b64decode(s.data_png) + except Exception as e: + logger.warning(f"Bad annotation signature data for {s.id}: {e}") + annotated_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name + _to_unlink.append(annotated_path) + try: + stamp_annotations(out_path, annotated_path, annotations, ann_signature_pngs) + out_path = annotated_path + except Exception as e: + logger.error(f"stamp_annotations failed for doc {doc_id}: {e}") + + download_name = _slug(doc.title or "form") + "_annotated.pdf" + return FileResponse( + out_path, + media_type="application/pdf", + filename=download_name, + background=BackgroundTask(_cleanup_temps), + ) + finally: + db.close() + + # ---- POST /api/document/{doc_id}/prepare-signed-reply ---- + @router.post("/api/document/{doc_id}/prepare-signed-reply") + async def prepare_signed_reply(doc_id: str, request: Request): + """Bake the current PDF state (form fields + signature stamps + + annotations) into a flattened PDF, drop it in COMPOSE_UPLOADS_DIR + and return the reply context (To/Subject/threading headers) so the + frontend can open a reply draft with this attachment pre-loaded. + + Requires the document to have source_email_* metadata (set when the + doc was created via /api/email/attachment-as-doc). Otherwise 400. + """ + import base64 + import tempfile + import shutil + import uuid as _uuid + import email as _email_mod + from src.pdf_form_doc import ( + find_source_upload_id, parse_markdown_to_values, + load_field_sidecar, parse_markdown_annotations, + ) + from src.pdf_forms import fill_fields, stamp_signatures, stamp_annotations + from core.database import Signature + # COMPOSE_UPLOADS_DIR lives in email_routes — re-derive here so we + # don't import from a routes file (cycle-prone). Same env override + # as email_routes (ODYSSEUS_MAIL_ATTACHMENTS_DIR). + from pathlib import Path as _Path + _COMPOSE_DIR = _Path(MAIL_ATTACHMENTS_DIR) / "_compose" + _COMPOSE_DIR.mkdir(parents=True, exist_ok=True) + + user = get_current_user(request) + db = SessionLocal() + try: + doc = db.query(Document).filter(Document.id == doc_id).first() + if not doc: + raise HTTPException(404, "Document not found") + _verify_doc_owner(db, doc, user) + + if not (doc.source_email_uid and doc.source_email_folder): + raise HTTPException(400, "Document has no source email — cannot reply") + + # 1) Build the flattened PDF (same pipeline as export_pdf) + upload_id = find_source_upload_id(doc.current_content or "") + if not upload_id: + raise HTTPException(400, "Document is not linked to a source PDF") + pdf_path = _locate_current_user_upload(request, upload_id, user) + if not pdf_path: + raise HTTPException(404, f"Source PDF {upload_id} not found") + + schema = load_field_sidecar(pdf_path) or [] + sig_field_names = {f["name"] for f in schema if f.get("type") == "signature"} + all_values = parse_markdown_to_values(doc.current_content or "") + text_values: dict = {} + sig_ids: dict[str, str] = {} + for name, raw in all_values.items(): + if name in sig_field_names and isinstance(raw, str) and raw.startswith("signature:"): + sig_ids[name] = raw[len("signature:"):].strip() + elif name not in sig_field_names: + text_values[name] = raw + + stamps: dict = {} + if sig_ids: + # SECURITY: filter by owner — same reason as render_pdf. + _sig_q2 = db.query(Signature).filter(Signature.id.in_(list(sig_ids.values()))) + if user: + _sig_q2 = _sig_q2.filter(Signature.owner == user) + rows = _sig_q2.all() + by_id = {s.id: s for s in rows} + for fname, sid in sig_ids.items(): + s = by_id.get(sid) + if not s: + continue + try: + stamps[fname] = base64.b64decode(s.data_png) + except Exception: + pass + + import os + _to_unlink: list[str] = [] + filled_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name + _to_unlink.append(filled_path) + fill_fields(pdf_path, filled_path, text_values) + out_path = filled_path + if stamps: + stamped_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name + _to_unlink.append(stamped_path) + try: + stamp_signatures(filled_path, stamped_path, stamps) + out_path = stamped_path + except Exception as e: + logger.warning(f"stamp_signatures failed for {doc_id}: {e}") + + annotations = parse_markdown_annotations(doc.current_content or "") + if annotations: + ann_sig_ids = [ + a["value"][len("signature:"):].strip() + for a in annotations + if a.get("kind") == "signature" + and isinstance(a.get("value"), str) + and a["value"].startswith("signature:") + ] + ann_signature_pngs: dict[str, bytes] = {} + if ann_sig_ids: + # SECURITY: filter by owner so a caller can't reference + # someone else's signature ID from doc markdown and have + # it stamped/exported. + _sig_q = db.query(Signature).filter(Signature.id.in_(ann_sig_ids)) + if user: + _sig_q = _sig_q.filter(Signature.owner == user) + sig_rows = _sig_q.all() + for s in sig_rows: + try: + ann_signature_pngs[s.id] = base64.b64decode(s.data_png) + except Exception: + pass + annotated_path = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False).name + _to_unlink.append(annotated_path) + try: + stamp_annotations(out_path, annotated_path, annotations, ann_signature_pngs) + out_path = annotated_path + except Exception as e: + logger.warning(f"stamp_annotations failed for {doc_id}: {e}") + + # 2) Move/copy into COMPOSE_UPLOADS_DIR with the token format + # `_` that /api/email/send expects. + filename = _slug(doc.title or "signed") + "_signed.pdf" + token = f"{_uuid.uuid4().hex}_{filename}" + dest = _COMPOSE_DIR / token + shutil.copyfile(out_path, str(dest)) + # Unlink the intermediate temp PDFs now that they've been + # copied into COMPOSE_UPLOADS_DIR. + for _p in _to_unlink: + try: + os.unlink(_p) + except FileNotFoundError: + pass + except Exception as _e: + logger.warning(f"Could not unlink temp PDF {_p}: {_e}") + + # 3) Fetch the source email's headers so we can build a clean reply + # context (To/Subject/In-Reply-To/References). + try: + from routes.email_routes import _imap, _decode_header + from routes.email_helpers import _q + except Exception: + _imap = None + _decode_header = lambda x: x or "" + _q = lambda x: x or "" + + to_addr = "" + from_name = "" + subject = "" + in_reply_to = doc.source_email_message_id or "" + references = in_reply_to + if _imap: + try: + with _imap(doc.source_email_account_id or None) as conn: + conn.select(_q(doc.source_email_folder), readonly=True) + status, data = conn.fetch(doc.source_email_uid.encode(), "(RFC822.HEADER)") + if status == "OK" and data and data[0]: + raw_hdr = data[0][1] + m = _email_mod.message_from_bytes(raw_hdr) + sender = _decode_header(m.get("From", "")) + from_name, to_addr = _email_mod.utils.parseaddr(sender) + if not to_addr: + to_addr = sender + subject = _decode_header(m.get("Subject", "") or "") + if subject and not subject.lower().startswith("re:"): + subject = "Re: " + subject + msg_refs = (m.get("References") or "").strip() + msg_in_reply = (m.get("Message-ID") or "").strip() or in_reply_to + in_reply_to = msg_in_reply + references = (msg_refs + " " + msg_in_reply).strip() if msg_refs else msg_in_reply + except Exception as e: + logger.warning(f"prepare-signed-reply header fetch failed: {e}") + + return { + "ok": True, + "attachment": { + "token": token, + "filename": filename, + "size": dest.stat().st_size, + }, + "reply": { + "to": to_addr, + "to_name": from_name, + "subject": subject, + "in_reply_to": in_reply_to, + "references": references, + "account_id": doc.source_email_account_id or None, + "source_uid": doc.source_email_uid, + "source_folder": doc.source_email_folder, + "source_message_id": doc.source_email_message_id, + }, + } + finally: + db.close() + + return router diff --git a/routes/email_helpers.py b/routes/email_helpers.py index 257f5f921..c8639e1c7 100644 --- a/routes/email_helpers.py +++ b/routes/email_helpers.py @@ -247,7 +247,6 @@ import re as _re_reply _REPLY_OPEN_RE = _re_reply.compile(r"<<<\s*(?:REPLY|SUMMARY|OUTPUT)\s*>>+", _re_reply.I) _REPLY_CLOSE_RE = _re_reply.compile(r"<<<\s*END\s*>>+", _re_reply.I) _REPLY_ROLE_MARKER_RE = _re_reply.compile(r"?|?", _re_reply.I) -_SUMMARY_BULLET_RE = _re_reply.compile(r"^(?:[-*\u2022]\s+|\d+[.)]\s+)") def _extract_reply(text: str) -> str: @@ -278,125 +277,6 @@ def _extract_reply(text: str) -> str: return _strip_think(t).strip() -def _build_email_summary_messages(sender: str, subject: str, body_for_llm: str) -> list[dict[str, str]]: - return [ - { - "role": "system", - "content": ( - "You are an email summarizer. Format: 1-3 short bullet points " - "(use '- '). Cover: main point, action items, deadlines. If the " - "email has attachments (marked '--- ATTACHMENTS ---'), USE THEIR " - "CONTENTS - pull invoice totals, deadlines, key clauses, concrete " - "numbers/dates from PDFs/docs into the bullets. Be terse.\n\n" - "OUTPUT FORMAT: Put ONLY the bullet points between these exact " - "markers, each on its own line:\n" - "<<>>\n" - "- ...\n" - "<<>>\n" - "Any reasoning must come BEFORE <<>> (ideally inside " - "...). Only the text between the markers is kept." - ), - }, - { - "role": "user", - "content": ( - f"From: {sender}\nSubject: {subject}\n\n{body_for_llm[:12000]}" - "\n\n---\n\nSummarize the email. Output the bullets between " - "<<>> and <<>>." - ), - }, - ] - - -async def _generate_email_summary( - url: str, - model: str, - sender: str, - subject: str, - body_for_llm: str, - *, - headers: dict | None = None, - max_tokens: int = 8192, - timeout: int = 180, -) -> str: - """Generate an interactive email summary through the shared LLM adapter.""" - from src.llm_core import llm_call_async - - raw = await llm_call_async( - url=url, - model=model, - messages=_build_email_summary_messages(sender, subject, body_for_llm), - temperature=0.3, - max_tokens=max_tokens, - headers=headers, - timeout=timeout, - workload="foreground", - ) - return _normalize_email_summary(raw) - - -async def _generate_scheduled_email_summary( - url: str, - model: str, - sender: str, - subject: str, - body_for_llm: str, - *, - headers: dict | None = None, - owner: str | None = None, - max_tokens: int = 8192, - timeout: int = 180, -) -> str: - """Generate a scheduled summary through the background task candidate chain.""" - from src.task_endpoint import task_llm_call_async - - raw = await task_llm_call_async( - messages=_build_email_summary_messages(sender, subject, body_for_llm), - fallback_url=url, - fallback_model=model, - fallback_headers=headers, - owner=owner, - temperature=0.3, - max_tokens=max_tokens, - timeout=timeout, - ) - return _normalize_email_summary(raw) - - -def _normalize_email_summary(raw) -> str: - """Extract a stable cache/UI summary from provider output.""" - raw_text = raw or "" - if _REPLY_OPEN_RE.search(raw_text): - summary = _extract_reply(raw_text) - if summary: - return summary - - cleaned = _strip_think(raw_text).strip() - bullets = [ - line.strip() - for line in cleaned.splitlines() - if _SUMMARY_BULLET_RE.match(line.strip()) - ] - if bullets: - return "\n".join(bullets) - return cleaned.strip() - - -EMAIL_SUMMARY_ERROR_CODE = "email_summary_unavailable" -EMAIL_SUMMARY_ERROR_MESSAGE = "Failed to summarize" - - -def _email_summary_failure_log_detail(exc: BaseException) -> str: - """Return useful provider-failure metadata without echoing exception text.""" - detail = f"type={type(exc).__name__}" - status = getattr(exc, "status_code", None) - if status is None: - status = getattr(getattr(exc, "response", None), "status_code", None) - if isinstance(status, int): - detail += f" status={status}" - return detail - - def _apply_email_style_mechanics(text: str) -> str: """Enforce deterministic writing-style mechanics that models often miss.""" if not text: diff --git a/routes/email_pollers.py b/routes/email_pollers.py index a2507989d..5d96bd0f9 100644 --- a/routes/email_pollers.py +++ b/routes/email_pollers.py @@ -40,7 +40,6 @@ from routes.email_helpers import ( _pre_retrieve_context, _attach_compose_uploads, _cleanup_compose_uploads, _q, SCHEDULED_DB, _EMAIL_REPLY_SYS_PROMPT_BASE, _email_cache_owner_clause, - _generate_scheduled_email_summary, _email_summary_failure_log_detail, ) logger = logging.getLogger(__name__) @@ -654,7 +653,6 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None no_msgid = 0 examined = 0 _summaries_created = 0 - _summary_failed = 0 _events_created = 0 _replies_drafted = 0 _reply_failed = 0 @@ -787,17 +785,16 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None if need_sum: try: - summary = await _generate_scheduled_email_summary( - url=url, - model=model, - sender=sender, - subject=subject, - body_for_llm=body_for_llm, - headers=req_headers, + summary = await task_llm_call_async( + messages=[ + {"role": "system", "content": "You are an email summarizer. Format: 1-3 short bullet points (use '- '). Cover: main point, action items, deadlines. If the email has attachments (marked '--- ATTACHMENTS ---'), USE THEIR CONTENTS — pull out invoice totals, deadlines, key clauses, any concrete numbers/dates in PDFs/docs, and reflect them in the bullets. Be terse.\n\nOUTPUT FORMAT: Put ONLY the bullet points between these exact markers, each on its own line:\n<<>>\n- ...\n<<>>\nAny reasoning or planning must come BEFORE <<>> (ideally inside ...). Only the text between the markers is kept."}, + {"role": "user", "content": f"From: {sender}\nSubject: {subject}\n\n{body_for_llm[:12000]}\n\n---\n\nSummarize the email. Output the bullets between <<>> and <<>>."}, + ], + fallback_url=url, fallback_model=model, fallback_headers=headers, owner=account_owner or None, - max_tokens=16384, - timeout=240, + temperature=0.3, max_tokens=16384, timeout=240, ) + summary = _extract_reply((summary or "").strip()) if summary: _c = _sql3.connect(SCHEDULED_DB) _c.execute(""" @@ -811,19 +808,10 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None _summaries_created += 1 _uid_text = uid.decode() if isinstance(uid, bytes) else str(uid) _detail_lines.append(f"summary · {_folder}#{_uid_text} · {subject or '(no subject)'} — {sender or '(unknown sender)'}") - else: - _summary_failed += 1 - _uid_text = uid.decode() if isinstance(uid, bytes) else str(uid) - _detail_lines.append(f"summary empty · {_folder}#{_uid_text} · {subject or '(no subject)'} — {sender or '(unknown sender)'}") except Exception as e: - _summary_failed += 1 _uid_text = uid.decode() if isinstance(uid, bytes) else str(uid) _detail_lines.append(f"summary failed · {_folder}#{_uid_text} · {subject or '(no subject)'} — {sender or '(unknown sender)'}") - logger.warning( - "Auto-summary uid=%s failed %s", - _uid_text, - _email_summary_failure_log_detail(e), - ) + logger.warning(f"Auto-summary {uid} failed: {e}") if need_reply: await _emit_progress(progress_cb, f"Drafting reply {processed + 1}/{_max_process} · checked {examined}/{len(uid_list)}") @@ -1332,8 +1320,6 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None parts.append(f"processed {processed} new") if auto_sum: parts.append(f"summarized {_summaries_created}") - if _summary_failed: - parts.append(f"{_summary_failed} summary failed") if auto_reply_draft: parts.append(f"drafted {_replies_drafted} repl" + ("y" if _replies_drafted == 1 else "ies")) if _reply_failed: diff --git a/routes/email_routes.py b/routes/email_routes.py index 76a744ce1..3c8e407bd 100644 --- a/routes/email_routes.py +++ b/routes/email_routes.py @@ -57,8 +57,7 @@ from routes.email_helpers import ( _extract_attachment_to_disk, _extract_html, _extract_text, _fetch_sender_thread_context, _pre_retrieve_context, _EMAIL_REPLY_SYS_PROMPT_BASE, _POOL_HOOKS, - _friendly_email_auth_error, _email_summary_failure_log_detail, - _generate_email_summary, EMAIL_SUMMARY_ERROR_CODE, EMAIL_SUMMARY_ERROR_MESSAGE, + _friendly_email_auth_error, SendEmailRequest, ExtractStyleRequest, ATTACHMENTS_DIR, COMPOSE_UPLOADS_DIR, SCHEDULED_DB, attachment_extract_dir, _email_cache_owner_clause, email_translation_body_hash, @@ -4767,6 +4766,8 @@ def setup_email_routes(): """Generate a quick AI summary of an email body.""" try: from src.endpoint_resolver import resolve_endpoint + from src.llm_core import _uses_max_completion_tokens, _restricts_temperature + import requests as _req body = data.get("body", "") subject = data.get("subject", "") @@ -4777,11 +4778,7 @@ def setup_email_routes(): if account_id: _assert_owns_account(account_id, owner) if not body: - return { - "success": False, - "error": "No body provided", - "error_code": "email_summary_missing_body", - } + return {"success": False, "error": "No body provided"} # If we know which UID this is, fetch the raw message and pull # attachment text so the summary can reference invoice totals, @@ -4810,43 +4807,53 @@ def setup_email_routes(): if not url: url, model, headers = resolve_endpoint("default", owner=owner) if not url or not model: - return { - "success": False, - "error": "No model configured for email summaries", - "error_code": "email_summary_not_configured", - } + return {"success": False, "error": "No LLM endpoint configured"} req_headers = {"Content-Type": "application/json"} if headers: req_headers.update(headers) - try: - content = await _generate_email_summary( - url=url, - model=model, - sender=sender, - subject=subject, - body_for_llm=body_for_llm, - headers=req_headers, - max_tokens=8192, - timeout=180, - ) - except Exception as e: - logger.warning( - "Email summary LLM call failed %s", - _email_summary_failure_log_detail(e), - ) - return { - "success": False, - "error": EMAIL_SUMMARY_ERROR_MESSAGE, - "error_code": EMAIL_SUMMARY_ERROR_CODE, - } + tok_key = "max_completion_tokens" if _uses_max_completion_tokens(model) else "max_tokens" + payload = { + "model": model, + "messages": [ + {"role": "system", "content": "You are an email summarizer. Format: 1-3 short bullet points (use '- '). Cover: main point, action items, deadlines. If the email has attachments (marked '--- ATTACHMENTS ---'), USE THEIR CONTENTS — pull invoice totals, deadlines, key clauses, concrete numbers/dates from PDFs/docs into the bullets. Be terse.\n\nOUTPUT FORMAT: Put ONLY the bullet points between these exact markers, each on its own line:\n<<>>\n- ...\n<<>>\nAny reasoning must come BEFORE <<>> (ideally inside ...). Only the text between the markers is kept."}, + {"role": "user", "content": f"From: {sender}\nSubject: {subject}\n\n{body_for_llm[:12000]}\n\n---\n\nSummarize the email. Output the bullets between <<>> and <<>>."}, + ], + tok_key: 8192, + "temperature": 0.3, + "stream": False, + } + # Reasoning models (o1/o3/o4/gpt-5) reject an explicit temperature. + if _restricts_temperature(model): + payload.pop("temperature", None) + resp = await asyncio.to_thread( + _req.post, url, json=payload, headers=req_headers, timeout=180 + ) + if not resp.ok: + return {"success": False, "error": f"LLM HTTP {resp.status_code}"} + rdata = resp.json() + msg = (rdata.get("choices") or [{}])[0].get("message", {}) + content = (msg.get("content") or "").strip() + content = _extract_reply(content) if not content: - return { - "success": False, - "error": "The model returned an empty summary", - "error_code": "email_summary_empty", - } + # Model put everything in reasoning_content — extract bullet points + rc = (msg.get("reasoning_content") or "").strip() + # Find bullet-point style output (lines starting with -, •, *, or numbered) + bullet_lines = [] + for line in rc.split("\n"): + stripped = line.strip() + if re.match(r"^[-•*]\s+|^\d+[.)]\s+", stripped): + bullet_lines.append(stripped) + if bullet_lines: + content = "\n".join(bullet_lines) + else: + # Last resort: take the last paragraph + paragraphs = [p.strip() for p in rc.split("\n\n") if p.strip()] + content = paragraphs[-1] if paragraphs else rc[:500] + + if not content: + return {"success": False, "error": "Empty response from model"} # Cache the summary if we have a message_id mid = data.get("message_id", "") @@ -4869,15 +4876,8 @@ def setup_email_routes(): return {"success": True, "summary": content, "model_used": model} except Exception as e: - logger.error( - "Email summary route failed %s", - _email_summary_failure_log_detail(e), - ) - return { - "success": False, - "error": EMAIL_SUMMARY_ERROR_MESSAGE, - "error_code": EMAIL_SUMMARY_ERROR_CODE, - } + logger.error(f"Failed to summarize: {e}") + return {"success": False, "error": "Mail operation failed"} @router.post("/translate") async def translate_email(data: dict, owner: str = Depends(require_owner)): diff --git a/routes/memory/memory_routes.py b/routes/memory/memory_routes.py index c4232bec4..d290046ec 100644 --- a/routes/memory/memory_routes.py +++ b/routes/memory/memory_routes.py @@ -21,7 +21,7 @@ def _strip_list_prefix(text: str) -> str: return text return _LIST_PREFIX_RE.sub("", text, count=1).strip() -from services.memory import MemoryManager, MemoryStoreUnreadable +from services.memory import MemoryManager from core.session_manager import SessionManager from src.request_models import MemoryAddRequest from core.database import SessionLocal @@ -35,22 +35,6 @@ from src.upload_limits import read_upload_limited, MEMORY_IMPORT_MAX_BYTES logger = logging.getLogger(__name__) -def _load_for_update(memory_manager) -> List[Dict[str, Any]]: - """Load the whole store for a read-modify-write cycle. - - A transient read failure must not look like an empty store: the caller - would append to ``[]`` and save that back, atomically destroying every - existing memory (issue #5673). Surface it as a 503 and change nothing. - """ - try: - return memory_manager.load_all_for_update() - except MemoryStoreUnreadable as e: - logger.error("Refusing to rewrite the memory store: %s", e) - raise HTTPException( - 503, "Memory store is temporarily unreadable — no changes were made." - ) - - def setup_memory_routes(memory_manager: MemoryManager, session_manager: SessionManager, memory_vector=None): """Set up memory-related routes.""" router = APIRouter(prefix="/api/memory", tags=["memory"]) @@ -132,7 +116,7 @@ def setup_memory_routes(memory_manager: MemoryManager, session_manager: SessionM new_entry = memory_manager.add_entry(text, memory_data.source, memory_data.category, owner=user) if memory_data.session_id: new_entry["session_id"] = memory_data.session_id - all_mem = _load_for_update(memory_manager) + all_mem = memory_manager.load_all() all_mem.append(new_entry) memory_manager.save(all_mem) # Sync vector index @@ -503,7 +487,7 @@ def setup_memory_routes(memory_manager: MemoryManager, session_manager: SessionM def pin_memory(request: Request, memory_id: str, pinned: bool = Form(True)): """Pin or unpin a memory. Pinned memories are always included in context.""" user = _owner(request) - all_mem = _load_for_update(memory_manager) + all_mem = memory_manager.load_all() for i, memory in enumerate(all_mem): if memory["id"] == memory_id: _verify_memory_owner(memory, user) @@ -528,7 +512,7 @@ def setup_memory_routes(memory_manager: MemoryManager, session_manager: SessionM def update_memory(request: Request, memory_id: str, text: str = Form(...), category: str = Form(None)): """Update an existing memory item with new text and optional category.""" user = _owner(request) - all_mem = _load_for_update(memory_manager) + all_mem = memory_manager.load_all() for i, memory in enumerate(all_mem): if memory["id"] == memory_id: _verify_memory_owner(memory, user) @@ -550,7 +534,7 @@ def setup_memory_routes(memory_manager: MemoryManager, session_manager: SessionM def delete_memory(request: Request, memory_id: str): """Delete a memory item by its ID.""" user = _owner(request) - all_mem = _load_for_update(memory_manager) + all_mem = memory_manager.load_all() # Find and verify ownership before deleting target = next((m for m in all_mem if m["id"] == memory_id), None) diff --git a/routes/vault/__init__.py b/routes/vault/__init__.py deleted file mode 100644 index 8aa82701d..000000000 --- a/routes/vault/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -"""Vault route domain package (slice 2k, #4082/#4071). - -Contains vault_routes.py, migrated from the flat routes/ directory. -Backward-compat shim at routes/vault_routes.py re-exports from here. -""" diff --git a/routes/vault/vault_routes.py b/routes/vault/vault_routes.py deleted file mode 100644 index 7e97500f0..000000000 --- a/routes/vault/vault_routes.py +++ /dev/null @@ -1,242 +0,0 @@ -""" -vault_routes.py - -Vaultwarden / Bitwarden CLI integration — config and unlock endpoints. -Stores the BW_SESSION key in data/vault.json with restrictive permissions. -""" - -import json -import logging -import os -import shutil -import asyncio -from pathlib import Path -from datetime import datetime -from fastapi import APIRouter, Request -from pydantic import BaseModel - -from core.middleware import require_admin -from core.platform_compat import IS_WINDOWS, safe_chmod, which_tool -from src.constants import VAULT_FILE as _VAULT_FILE - -logger = logging.getLogger(__name__) - -VAULT_FILE = Path(_VAULT_FILE) - - -def _find_bw() -> str: - """Locate the bw binary, checking PATH and common npm-global locations. - - On Windows the Bitwarden CLI shim is `bw.cmd`/`bw.exe`, resolved by - which_tool via PATHEXT. - """ - p = which_tool("bw") - if p: - return p - if IS_WINDOWS: - appdata = os.environ.get("APPDATA", os.path.expanduser("~")) - for candidate in ( - os.path.join(appdata, "npm", "bw.cmd"), - os.path.join(appdata, "npm", "bw.exe"), - ): - if os.path.isfile(candidate): - return candidate - return "bw" - home = os.path.expanduser("~") - for candidate in ( - f"{home}/.npm-global/bin/bw", - f"{home}/.nvm/versions/node/*/bin/bw", - "/usr/local/bin/bw", - "/opt/homebrew/bin/bw", - ): - if "*" in candidate: - import glob - for m in glob.glob(candidate): - if os.path.isfile(m) and os.access(m, os.X_OK): - return m - elif os.path.isfile(candidate) and os.access(candidate, os.X_OK): - return candidate - return "bw" # fall back to PATH lookup (will FileNotFoundError, handled below) - - -def _load_config() -> dict: - if VAULT_FILE.exists(): - try: - data = json.loads(VAULT_FILE.read_text(encoding="utf-8")) - return data if isinstance(data, dict) else {} - except Exception: - pass - return {} - - -def _save_config(cfg: dict): - VAULT_FILE.parent.mkdir(parents=True, exist_ok=True) - VAULT_FILE.write_text(json.dumps(cfg, indent=2), encoding="utf-8") - # POSIX: restrict the BW_SESSION store to 0o600. Windows: no-op (profile dir - # is ACL-restricted already). - safe_chmod(str(VAULT_FILE), 0o600) - - -async def _run_bw(args: list, session: str = None, input_text: str = None, - bw_password: str = None) -> tuple: - env = {} - env.update(os.environ) - if session: - env["BW_SESSION"] = session - # Secrets must never be passed as argv — process arguments are world-readable - # via `ps` / `/proc//cmdline` to any local user. Keep --passwordenv - # support for bw commands that need it; unlock/login callers should prefer - # stdin so the master password is not left in the child environment either. - if bw_password is not None: - env["BW_PASSWORD"] = bw_password - bw_path = _find_bw() - try: - proc = await asyncio.create_subprocess_exec( - bw_path, *args, - stdin=asyncio.subprocess.PIPE if input_text else None, - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - env=env, - ) - except FileNotFoundError: - return "", "bw CLI not installed (install `nodejs-bitwarden-cli` or `bitwarden-cli`)", 127 - except Exception as e: - return "", f"Failed to launch bw: {e}", 1 - try: - stdout, stderr = await proc.communicate(input=input_text.encode() if input_text else None) - except Exception as e: - return "", f"bw subprocess error: {e}", 1 - return stdout.decode(errors="replace").strip(), stderr.decode(errors="replace").strip(), proc.returncode - - -class VaultConfig(BaseModel): - server_url: str = "" - email: str = "" - - -class VaultUnlockRequest(BaseModel): - master_password: str - - -class VaultLoginRequest(BaseModel): - email: str - master_password: str - - -def setup_vault_routes(): - router = APIRouter(prefix="/api/vault", tags=["vault"]) - - @router.get("/config") - async def get_config(request: Request): - """Return vault config (no sensitive fields).""" - require_admin(request) - cfg = _load_config() - return { - "server_url": cfg.get("server_url", ""), - "email": cfg.get("email", ""), - "unlocked": bool(cfg.get("session")), - "unlocked_at": cfg.get("unlocked_at", ""), - "bw_installed": await _check_bw_installed(), - } - - @router.post("/config") - async def save_config(req: VaultConfig, request: Request): - """Save vault URL + email. Runs 'bw config server' to point at Vaultwarden.""" - require_admin(request) - cfg = _load_config() - cfg["server_url"] = req.server_url.strip().rstrip("/") - cfg["email"] = req.email.strip() - - if cfg["server_url"]: - _, stderr, rc = await _run_bw(["config", "server", cfg["server_url"]]) - if rc != 0: - return {"ok": False, "error": f"bw config failed: {stderr[:300]}"} - - _save_config(cfg) - return {"ok": True} - - @router.post("/login") - async def login(req: VaultLoginRequest, request: Request): - """Log in to Vaultwarden (required once per account).""" - require_admin(request) - cfg = _load_config() - # Update email - cfg["email"] = req.email - _save_config(cfg) - - stdout, stderr, rc = await _run_bw( - ["login", req.email, "--raw"], - input_text=req.master_password + "\n", - ) - if rc != 0: - # Already logged in is OK - if "already logged in" in stderr.lower(): - return {"ok": True, "already": True} - return {"ok": False, "error": f"Login failed: {stderr[:300]}"} - # bw login --raw prints session key on success (when 2FA disabled) - if stdout: - cfg["session"] = stdout - cfg["unlocked_at"] = datetime.utcnow().isoformat() - _save_config(cfg) - return {"ok": True} - - @router.post("/unlock") - async def unlock(req: VaultUnlockRequest, request: Request): - """Unlock the vault and save the session key.""" - require_admin(request) - # Pass the master password on stdin, not argv. argv is visible through - # `ps` / /proc//cmdline; stdin also avoids leaving the secret in - # the child process environment. - stdout, stderr, rc = await _run_bw( - ["unlock", "--raw"], - input_text=req.master_password + "\n", - ) - if rc != 0: - return {"ok": False, "error": f"Unlock failed: {stderr[:300]}"} - session = stdout.strip() - if not session: - return {"ok": False, "error": "bw returned empty session"} - cfg = _load_config() - cfg["session"] = session - cfg["unlocked_at"] = datetime.utcnow().isoformat() - _save_config(cfg) - return {"ok": True, "message": "Vault unlocked"} - - @router.post("/lock") - async def lock(request: Request): - """Lock the vault (clear session from config).""" - require_admin(request) - cfg = _load_config() - cfg.pop("session", None) - cfg.pop("unlocked_at", None) - _save_config(cfg) - # Also tell bw to lock - await _run_bw(["lock"]) - return {"ok": True, "message": "Vault locked"} - - @router.post("/logout") - async def logout(request: Request): - """Log out of the Bitwarden CLI completely.""" - require_admin(request) - await _run_bw(["logout"]) - cfg = _load_config() - cfg.pop("session", None) - cfg.pop("email", None) - cfg.pop("unlocked_at", None) - _save_config(cfg) - return {"ok": True} - - return router - - -async def _check_bw_installed() -> bool: - try: - proc = await asyncio.create_subprocess_exec( - _find_bw(), "--version", - stdout=asyncio.subprocess.PIPE, - stderr=asyncio.subprocess.PIPE, - ) - await proc.communicate() - return proc.returncode == 0 - except Exception: - return False diff --git a/routes/vault_routes.py b/routes/vault_routes.py index cfed2ba39..7e97500f0 100644 --- a/routes/vault_routes.py +++ b/routes/vault_routes.py @@ -1,14 +1,242 @@ -"""Backward-compat shim — canonical location is routes/vault/vault_routes.py. +""" +vault_routes.py -This module is replaced in ``sys.modules`` by the canonical module object so -that ``import routes.vault_routes``, ``from routes.vault_routes import X``, -and the ``import ... as vr`` + ``monkeypatch.setattr(vr, ...)`` pattern used -by test_vault_password_not_in_argv.py all operate on the *same* object. -Keeps existing import paths working after slice 2k (#4082/#4071). +Vaultwarden / Bitwarden CLI integration — config and unlock endpoints. +Stores the BW_SESSION key in data/vault.json with restrictive permissions. """ -import sys as _sys +import json +import logging +import os +import shutil +import asyncio +from pathlib import Path +from datetime import datetime +from fastapi import APIRouter, Request +from pydantic import BaseModel -from routes.vault import vault_routes as _canonical # noqa: F401 +from core.middleware import require_admin +from core.platform_compat import IS_WINDOWS, safe_chmod, which_tool +from src.constants import VAULT_FILE as _VAULT_FILE -_sys.modules[__name__] = _canonical +logger = logging.getLogger(__name__) + +VAULT_FILE = Path(_VAULT_FILE) + + +def _find_bw() -> str: + """Locate the bw binary, checking PATH and common npm-global locations. + + On Windows the Bitwarden CLI shim is `bw.cmd`/`bw.exe`, resolved by + which_tool via PATHEXT. + """ + p = which_tool("bw") + if p: + return p + if IS_WINDOWS: + appdata = os.environ.get("APPDATA", os.path.expanduser("~")) + for candidate in ( + os.path.join(appdata, "npm", "bw.cmd"), + os.path.join(appdata, "npm", "bw.exe"), + ): + if os.path.isfile(candidate): + return candidate + return "bw" + home = os.path.expanduser("~") + for candidate in ( + f"{home}/.npm-global/bin/bw", + f"{home}/.nvm/versions/node/*/bin/bw", + "/usr/local/bin/bw", + "/opt/homebrew/bin/bw", + ): + if "*" in candidate: + import glob + for m in glob.glob(candidate): + if os.path.isfile(m) and os.access(m, os.X_OK): + return m + elif os.path.isfile(candidate) and os.access(candidate, os.X_OK): + return candidate + return "bw" # fall back to PATH lookup (will FileNotFoundError, handled below) + + +def _load_config() -> dict: + if VAULT_FILE.exists(): + try: + data = json.loads(VAULT_FILE.read_text(encoding="utf-8")) + return data if isinstance(data, dict) else {} + except Exception: + pass + return {} + + +def _save_config(cfg: dict): + VAULT_FILE.parent.mkdir(parents=True, exist_ok=True) + VAULT_FILE.write_text(json.dumps(cfg, indent=2), encoding="utf-8") + # POSIX: restrict the BW_SESSION store to 0o600. Windows: no-op (profile dir + # is ACL-restricted already). + safe_chmod(str(VAULT_FILE), 0o600) + + +async def _run_bw(args: list, session: str = None, input_text: str = None, + bw_password: str = None) -> tuple: + env = {} + env.update(os.environ) + if session: + env["BW_SESSION"] = session + # Secrets must never be passed as argv — process arguments are world-readable + # via `ps` / `/proc//cmdline` to any local user. Keep --passwordenv + # support for bw commands that need it; unlock/login callers should prefer + # stdin so the master password is not left in the child environment either. + if bw_password is not None: + env["BW_PASSWORD"] = bw_password + bw_path = _find_bw() + try: + proc = await asyncio.create_subprocess_exec( + bw_path, *args, + stdin=asyncio.subprocess.PIPE if input_text else None, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + env=env, + ) + except FileNotFoundError: + return "", "bw CLI not installed (install `nodejs-bitwarden-cli` or `bitwarden-cli`)", 127 + except Exception as e: + return "", f"Failed to launch bw: {e}", 1 + try: + stdout, stderr = await proc.communicate(input=input_text.encode() if input_text else None) + except Exception as e: + return "", f"bw subprocess error: {e}", 1 + return stdout.decode(errors="replace").strip(), stderr.decode(errors="replace").strip(), proc.returncode + + +class VaultConfig(BaseModel): + server_url: str = "" + email: str = "" + + +class VaultUnlockRequest(BaseModel): + master_password: str + + +class VaultLoginRequest(BaseModel): + email: str + master_password: str + + +def setup_vault_routes(): + router = APIRouter(prefix="/api/vault", tags=["vault"]) + + @router.get("/config") + async def get_config(request: Request): + """Return vault config (no sensitive fields).""" + require_admin(request) + cfg = _load_config() + return { + "server_url": cfg.get("server_url", ""), + "email": cfg.get("email", ""), + "unlocked": bool(cfg.get("session")), + "unlocked_at": cfg.get("unlocked_at", ""), + "bw_installed": await _check_bw_installed(), + } + + @router.post("/config") + async def save_config(req: VaultConfig, request: Request): + """Save vault URL + email. Runs 'bw config server' to point at Vaultwarden.""" + require_admin(request) + cfg = _load_config() + cfg["server_url"] = req.server_url.strip().rstrip("/") + cfg["email"] = req.email.strip() + + if cfg["server_url"]: + _, stderr, rc = await _run_bw(["config", "server", cfg["server_url"]]) + if rc != 0: + return {"ok": False, "error": f"bw config failed: {stderr[:300]}"} + + _save_config(cfg) + return {"ok": True} + + @router.post("/login") + async def login(req: VaultLoginRequest, request: Request): + """Log in to Vaultwarden (required once per account).""" + require_admin(request) + cfg = _load_config() + # Update email + cfg["email"] = req.email + _save_config(cfg) + + stdout, stderr, rc = await _run_bw( + ["login", req.email, "--raw"], + input_text=req.master_password + "\n", + ) + if rc != 0: + # Already logged in is OK + if "already logged in" in stderr.lower(): + return {"ok": True, "already": True} + return {"ok": False, "error": f"Login failed: {stderr[:300]}"} + # bw login --raw prints session key on success (when 2FA disabled) + if stdout: + cfg["session"] = stdout + cfg["unlocked_at"] = datetime.utcnow().isoformat() + _save_config(cfg) + return {"ok": True} + + @router.post("/unlock") + async def unlock(req: VaultUnlockRequest, request: Request): + """Unlock the vault and save the session key.""" + require_admin(request) + # Pass the master password on stdin, not argv. argv is visible through + # `ps` / /proc//cmdline; stdin also avoids leaving the secret in + # the child process environment. + stdout, stderr, rc = await _run_bw( + ["unlock", "--raw"], + input_text=req.master_password + "\n", + ) + if rc != 0: + return {"ok": False, "error": f"Unlock failed: {stderr[:300]}"} + session = stdout.strip() + if not session: + return {"ok": False, "error": "bw returned empty session"} + cfg = _load_config() + cfg["session"] = session + cfg["unlocked_at"] = datetime.utcnow().isoformat() + _save_config(cfg) + return {"ok": True, "message": "Vault unlocked"} + + @router.post("/lock") + async def lock(request: Request): + """Lock the vault (clear session from config).""" + require_admin(request) + cfg = _load_config() + cfg.pop("session", None) + cfg.pop("unlocked_at", None) + _save_config(cfg) + # Also tell bw to lock + await _run_bw(["lock"]) + return {"ok": True, "message": "Vault locked"} + + @router.post("/logout") + async def logout(request: Request): + """Log out of the Bitwarden CLI completely.""" + require_admin(request) + await _run_bw(["logout"]) + cfg = _load_config() + cfg.pop("session", None) + cfg.pop("email", None) + cfg.pop("unlocked_at", None) + _save_config(cfg) + return {"ok": True} + + return router + + +async def _check_bw_installed() -> bool: + try: + proc = await asyncio.create_subprocess_exec( + _find_bw(), "--version", + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + await proc.communicate() + return proc.returncode == 0 + except Exception: + return False diff --git a/routes/webhook/__init__.py b/routes/webhook/__init__.py deleted file mode 100644 index e51389e3a..000000000 --- a/routes/webhook/__init__.py +++ /dev/null @@ -1,5 +0,0 @@ -"""Webhook route domain package (slice 2l, #4082/#4071). - -Contains webhook_routes.py, migrated from the flat routes/ directory. -Backward-compat shim at routes/webhook_routes.py re-exports from here. -""" diff --git a/routes/webhook/webhook_routes.py b/routes/webhook/webhook_routes.py deleted file mode 100644 index 8d3a704c6..000000000 --- a/routes/webhook/webhook_routes.py +++ /dev/null @@ -1,395 +0,0 @@ -"""Webhook, API Token, and sync chat routes.""" - -import uuid -import logging -from typing import Optional - -import httpx -from fastapi import APIRouter, HTTPException, Request, Form -from pydantic import BaseModel, Field - -from core.database import SessionLocal, Webhook, ModelEndpoint -from src.auth_helpers import owner_filter -from src.url_security import validate_public_http_url -from src.webhook_manager import WebhookManager, validate_webhook_url, validate_events - -logger = logging.getLogger(__name__) - -router = APIRouter(prefix="/api", tags=["webhooks"]) - -# Input limits -MAX_NAME_LEN = 100 -MAX_URL_LEN = 2048 -MAX_SECRET_LEN = 256 -MAX_MESSAGE_LEN = 32_000 - - -from core.middleware import require_admin as _require_admin - - -def _select_api_chat_fallback_endpoint(db, token_owner: Optional[str]): - """First enabled ModelEndpoint visible to token_owner — their own rows plus - legacy null-owner ("shared") rows. Owner-scoped: an unscoped .first() would - let a chat-scoped token fall back onto another user's private endpoint and - silently spend that owner's API key/quota. Prefer owner rows before shared - rows. Fails closed to null-owner rows only when token_owner is absent. - Does not validate base_url — admin-configured local/LAN endpoints remain allowed. - """ - query = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True) # noqa: E712 - if token_owner: - query = owner_filter(query, ModelEndpoint, token_owner) - return query.order_by(ModelEndpoint.owner.desc(), ModelEndpoint.created_at).first() - return query.filter(ModelEndpoint.owner == None).order_by(ModelEndpoint.created_at).first() # noqa: E711 - - -def _caller_owns_session(sess_owner, caller) -> bool: - """Strict session-ownership gate for the token-authenticated sync-chat - endpoint (`POST /api/v1/chat`). - - Mirrors ``_verify_session_owner`` in session_routes.py and the null-owner - gates in notes/calendar/gallery: a caller may resume a session ONLY when - its owner matches them exactly. A null/empty session owner (legacy or - migrated rows) is deliberately NOT resumable by an arbitrary token — the - old ``sess_owner and sess_owner != caller`` form skipped the check whenever - ``sess_owner`` was falsy, so any chat-scoped token (e.g. a paired mobile - device) could resume such a session, inject a message, and read back its - history and reuse the owner's endpoint credentials. Fail closed: an - unresolvable caller also returns False. - """ - if not caller: - return False - return sess_owner == caller - - -def setup_webhook_routes( - webhook_manager: WebhookManager, - auth_manager, - session_manager=None, - api_key_manager=None, -) -> APIRouter: - - @router.get("/webhooks") - def list_webhooks(request: Request): - _require_admin(request) - db = SessionLocal() - try: - hooks = db.query(Webhook).all() - return [ - { - "id": w.id, - "name": w.name, - "url": w.url, - "has_secret": bool(w.secret), - "events": w.events.split(",") if w.events else [], - "is_active": w.is_active, - "last_triggered_at": w.last_triggered_at.isoformat() if w.last_triggered_at else None, - "last_status_code": w.last_status_code, - "last_error": w.last_error, - "created_at": w.created_at.isoformat() if w.created_at else None, - } - for w in hooks - ] - finally: - db.close() - - @router.post("/webhooks") - def create_webhook( - request: Request, - name: str = Form(""), - url: str = Form(""), - secret: str = Form(""), - events: str = Form(""), - ): - _require_admin(request) - name = name.strip()[:MAX_NAME_LEN] - if not name: - raise HTTPException(400, "Webhook name is required") - try: - url = validate_webhook_url(url) - except ValueError as e: - raise HTTPException(400, str(e)) - try: - events = validate_events(events) - except ValueError as e: - raise HTTPException(400, str(e)) - - secret_val = secret.strip()[:MAX_SECRET_LEN] or None - # Encrypt the secret at rest using the same Fernet key as API keys - encrypted_secret = None - if secret_val and api_key_manager: - encrypted_secret = api_key_manager.encrypt_api_key(secret_val) - elif secret_val: - encrypted_secret = secret_val # Fallback if no encryption available - - webhook_id = str(uuid.uuid4())[:8] - db = SessionLocal() - try: - db.add(Webhook( - id=webhook_id, - name=name, - url=url, - secret=encrypted_secret, - events=events, - is_active=True, - )) - db.commit() - finally: - db.close() - - return {"id": webhook_id, "name": name} - - @router.post("/webhooks/{webhook_id}/test") - async def test_webhook(request: Request, webhook_id: str): - _require_admin(request) - db = SessionLocal() - try: - wh = db.query(Webhook).filter(Webhook.id == webhook_id).first() - if not wh: - raise HTTPException(404, "Webhook not found") - url, secret = wh.url, wh.secret - finally: - db.close() - - await webhook_manager.deliver_test(webhook_id, url, secret) - return {"status": "sent"} - - @router.patch("/webhooks/{webhook_id}") - def toggle_webhook(request: Request, webhook_id: str): - _require_admin(request) - db = SessionLocal() - try: - wh = db.query(Webhook).filter(Webhook.id == webhook_id).first() - if not wh: - raise HTTPException(404, "Webhook not found") - wh.is_active = not wh.is_active - db.commit() - return {"id": webhook_id, "is_active": wh.is_active} - finally: - db.close() - - @router.delete("/webhooks/{webhook_id}") - def delete_webhook(request: Request, webhook_id: str): - _require_admin(request) - db = SessionLocal() - try: - deleted = db.query(Webhook).filter(Webhook.id == webhook_id).delete() - db.commit() - if not deleted: - raise HTTPException(404, "Webhook not found") - finally: - db.close() - return {"status": "deleted"} - - # ================================================================ - # Sync Chat Endpoint (for n8n / Make / Activepieces) - # ================================================================ - - # Known provider base URLs — auto-resolved from api_key prefix or model name - KNOWN_PROVIDERS = { - "deepseek": "https://api.deepseek.com/v1", - "openai": "https://api.openai.com/v1", - "mistral": "https://api.mistral.ai/v1", - "groq": "https://api.groq.com/openai/v1", - "together": "https://api.together.xyz/v1", - "openrouter": "https://openrouter.ai/api/v1", - "ollama": "https://ollama.com/api", - "opencode-zen": "https://opencode.ai/zen/v1", - "opencode-go": "https://opencode.ai/zen/go/v1", - "fireworks": "https://api.fireworks.ai/inference/v1", - "venice": "https://api.venice.ai/api/v1", - "kimi-code": "https://api.kimi.com/coding/v1", - "kimicode": "https://api.kimi.com/coding/v1", - } - - # Model prefix → provider mapping for auto-detection - MODEL_PROVIDER_MAP = { - "deepseek": "deepseek", - "gpt-": "openai", - "o1": "openai", - "o3": "openai", - "o4": "openai", - "mistral": "mistral", - "llama": "groq", - "mixtral": "groq", - "kimi-for-coding": "kimi-code", - "kimi": "kimi-code", - } - - def _resolve_base_url(model: Optional[str], provider: Optional[str]) -> Optional[str]: - """Try to auto-resolve a base URL from provider name or model prefix.""" - if provider and provider.lower() in KNOWN_PROVIDERS: - return KNOWN_PROVIDERS[provider.lower()] - if model: - model_lower = model.lower() - for prefix, prov in MODEL_PROVIDER_MAP.items(): - if model_lower.startswith(prefix): - return KNOWN_PROVIDERS[prov] - return None - - class SyncChatRequest(BaseModel): - message: str = Field(..., max_length=MAX_MESSAGE_LEN) - model: Optional[str] = Field(None, max_length=200) - session: Optional[str] = Field(None, max_length=100) - api_key: Optional[str] = Field(None, max_length=256) - base_url: Optional[str] = Field(None, max_length=MAX_URL_LEN) - provider: Optional[str] = Field(None, max_length=50) - - @router.post("/v1/chat") - async def sync_chat(request: Request, body: SyncChatRequest): - if not getattr(request.state, "api_token", False): - raise HTTPException(403, "This endpoint requires an API token") - scopes = set(getattr(request.state, "api_token_scopes", []) or []) - if "chat" not in scopes: - raise HTTPException(403, "API token is not scoped for chat") - token_owner = getattr(request.state, "api_token_owner", None) - - from core.models import ChatMessage - from src.llm_core import llm_call_async - from src.endpoint_resolver import build_chat_url, build_headers, build_models_url, normalize_base - - message = body.message.strip() - if not message: - raise HTTPException(400, "Message is required") - - session_id = body.session - sess = None - - # --- Case 1: Resume an existing session --- - if session_id and session_manager: - try: - sess = session_manager.get_session(session_id) - except (KeyError, Exception): - raise HTTPException(404, "Session not found") - # SECURITY: verify the API-token's user owns this session — without - # this any token holder could resume any user's chat by passing its - # ID. The token's user is on request.state.user (set by API-token - # middleware); fall back to require_user if not present. - try: - from src.auth_helpers import get_current_user as _gcu - _tok_user = token_owner or getattr(request.state, "user", None) or _gcu(request) - except Exception: - _tok_user = None - # Strict ownership (see _caller_owns_session): fail closed so a - # null-owner / cross-owner session can't be resumed by an arbitrary - # chat-scoped token. - _sess_owner = getattr(sess, "owner", None) - if not _caller_owns_session(_sess_owner, _tok_user): - raise HTTPException(404, "Session not found") - - # --- Case 2: Direct API key + model (no pre-configured endpoint needed) --- - if not sess and body.api_key: - api_key = body.api_key.strip() - model = body.model or "deepseek-chat" - - # Validate only token-supplied direct base_url; auto-resolved known-provider - # URLs are not subject to extra local/LAN blocking beyond existing provider logic. - direct_base_url = body.base_url.strip().rstrip("/") if body.base_url else None - if direct_base_url: - try: - base_url = validate_public_http_url(direct_base_url) - except ValueError as e: - detail = str(e).replace("URL", "base_url", 1) - raise HTTPException(400, detail) - else: - base_url = _resolve_base_url(model, body.provider) - if not base_url: - raise HTTPException(400, - "Could not auto-detect provider. Pass base_url (e.g. 'https://api.deepseek.com/v1') " - "or provider ('deepseek', 'openai', 'groq', etc.)") - base_url = normalize_base(base_url) - endpoint_url = build_chat_url(base_url) - - if not session_manager: - raise HTTPException(500, "Session manager not available") - - sid = str(uuid.uuid4()) - sess = session_manager.create_session( - session_id=sid, name="API Chat", endpoint_url=endpoint_url, - model=model, owner=token_owner, - ) - sess.headers = build_headers(api_key, base_url) - session_manager.save_sessions() - session_id = sid - - # --- Case 3: Fall back to first configured ModelEndpoint --- - if not sess: - db = SessionLocal() - try: - ep = _select_api_chat_fallback_endpoint(db, token_owner) - finally: - db.close() - - if not ep: - raise HTTPException(400, - "No session, api_key, or configured endpoints. " - "Pass api_key + model, or configure an endpoint in Admin.") - - base_url = normalize_base(ep.base_url) - endpoint_url = build_chat_url(base_url) - model = body.model or "auto" - api_key = ep.api_key - if getattr(ep, "provider_auth_id", None): - try: - from src.endpoint_resolver import resolve_endpoint_runtime - base_url, api_key = resolve_endpoint_runtime(ep, owner=token_owner) - endpoint_url = build_chat_url(base_url) - except Exception: - raise HTTPException(500, "Could not resolve endpoint credentials") - - if model == "auto": - try: - async with httpx.AsyncClient(timeout=5) as client: - models_url = build_models_url(base_url) - hdrs = build_headers(api_key, base_url) - if models_url: - resp = await client.get(models_url, headers=hdrs) - resp.raise_for_status() - data = resp.json() - items = data if isinstance(data, list) else (data.get("data") or []) - ids = [m.get("id") for m in items if isinstance(m, dict) and m.get("id")] - if not ids and isinstance(data, dict): - ids = [ - m.get("name") or m.get("model") - for m in (data.get("models") or []) - if m.get("name") or m.get("model") - ] - else: - import json as _json - ids = _json.loads(ep.cached_models or "[]") - model = ids[0] if ids else "auto" - except Exception: - raise HTTPException(500, "Could not discover models from endpoint") - - if not session_manager: - raise HTTPException(500, "Session manager not available") - - sid = str(uuid.uuid4()) - sess = session_manager.create_session( - session_id=sid, name="API Chat", endpoint_url=endpoint_url, - model=model, owner=token_owner, - ) - if api_key: - sess.headers = build_headers(api_key, base_url) - session_manager.save_sessions() - session_id = sid - - # --- Send message and get response --- - sess.add_message(ChatMessage("user", message)) - - messages = [{"role": m.role, "content": m.content} for m in sess.history] - - reply = await llm_call_async( - sess.endpoint_url, sess.model, messages, - headers=sess.headers, timeout=120, - ) - sess.add_message(ChatMessage("assistant", reply)) - session_manager.save_sessions() - - webhook_manager.fire_and_forget("chat.completed", { - "session_id": session_id, "model": sess.model, - "user_message": message[:2000], "response": reply[:2000], - }) - - return {"response": reply, "session_id": session_id, "model": sess.model} - - return router diff --git a/routes/webhook_routes.py b/routes/webhook_routes.py index 7c5e0453e..8d3a704c6 100644 --- a/routes/webhook_routes.py +++ b/routes/webhook_routes.py @@ -1,16 +1,395 @@ -"""Backward-compat shim — canonical location is routes/webhook/webhook_routes.py. +"""Webhook, API Token, and sync chat routes.""" -This module is replaced in ``sys.modules`` by the canonical module object so -that ``import routes.webhook_routes``, ``from routes.webhook_routes import X``, -``importlib.import_module("routes.webhook_routes")``, and the -``__import__("routes.webhook_routes", fromlist=[...])`` + ``setattr(wh_mod, -...)`` pattern used by test_null_owner_gates.py all operate on the *same* -object. Keeps existing import paths working after slice 2l (#4082/#4071). -Source-introspection tests read the canonical file by path. -""" +import uuid +import logging +from typing import Optional -import sys as _sys +import httpx +from fastapi import APIRouter, HTTPException, Request, Form +from pydantic import BaseModel, Field -from routes.webhook import webhook_routes as _canonical # noqa: F401 +from core.database import SessionLocal, Webhook, ModelEndpoint +from src.auth_helpers import owner_filter +from src.url_security import validate_public_http_url +from src.webhook_manager import WebhookManager, validate_webhook_url, validate_events -_sys.modules[__name__] = _canonical +logger = logging.getLogger(__name__) + +router = APIRouter(prefix="/api", tags=["webhooks"]) + +# Input limits +MAX_NAME_LEN = 100 +MAX_URL_LEN = 2048 +MAX_SECRET_LEN = 256 +MAX_MESSAGE_LEN = 32_000 + + +from core.middleware import require_admin as _require_admin + + +def _select_api_chat_fallback_endpoint(db, token_owner: Optional[str]): + """First enabled ModelEndpoint visible to token_owner — their own rows plus + legacy null-owner ("shared") rows. Owner-scoped: an unscoped .first() would + let a chat-scoped token fall back onto another user's private endpoint and + silently spend that owner's API key/quota. Prefer owner rows before shared + rows. Fails closed to null-owner rows only when token_owner is absent. + Does not validate base_url — admin-configured local/LAN endpoints remain allowed. + """ + query = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True) # noqa: E712 + if token_owner: + query = owner_filter(query, ModelEndpoint, token_owner) + return query.order_by(ModelEndpoint.owner.desc(), ModelEndpoint.created_at).first() + return query.filter(ModelEndpoint.owner == None).order_by(ModelEndpoint.created_at).first() # noqa: E711 + + +def _caller_owns_session(sess_owner, caller) -> bool: + """Strict session-ownership gate for the token-authenticated sync-chat + endpoint (`POST /api/v1/chat`). + + Mirrors ``_verify_session_owner`` in session_routes.py and the null-owner + gates in notes/calendar/gallery: a caller may resume a session ONLY when + its owner matches them exactly. A null/empty session owner (legacy or + migrated rows) is deliberately NOT resumable by an arbitrary token — the + old ``sess_owner and sess_owner != caller`` form skipped the check whenever + ``sess_owner`` was falsy, so any chat-scoped token (e.g. a paired mobile + device) could resume such a session, inject a message, and read back its + history and reuse the owner's endpoint credentials. Fail closed: an + unresolvable caller also returns False. + """ + if not caller: + return False + return sess_owner == caller + + +def setup_webhook_routes( + webhook_manager: WebhookManager, + auth_manager, + session_manager=None, + api_key_manager=None, +) -> APIRouter: + + @router.get("/webhooks") + def list_webhooks(request: Request): + _require_admin(request) + db = SessionLocal() + try: + hooks = db.query(Webhook).all() + return [ + { + "id": w.id, + "name": w.name, + "url": w.url, + "has_secret": bool(w.secret), + "events": w.events.split(",") if w.events else [], + "is_active": w.is_active, + "last_triggered_at": w.last_triggered_at.isoformat() if w.last_triggered_at else None, + "last_status_code": w.last_status_code, + "last_error": w.last_error, + "created_at": w.created_at.isoformat() if w.created_at else None, + } + for w in hooks + ] + finally: + db.close() + + @router.post("/webhooks") + def create_webhook( + request: Request, + name: str = Form(""), + url: str = Form(""), + secret: str = Form(""), + events: str = Form(""), + ): + _require_admin(request) + name = name.strip()[:MAX_NAME_LEN] + if not name: + raise HTTPException(400, "Webhook name is required") + try: + url = validate_webhook_url(url) + except ValueError as e: + raise HTTPException(400, str(e)) + try: + events = validate_events(events) + except ValueError as e: + raise HTTPException(400, str(e)) + + secret_val = secret.strip()[:MAX_SECRET_LEN] or None + # Encrypt the secret at rest using the same Fernet key as API keys + encrypted_secret = None + if secret_val and api_key_manager: + encrypted_secret = api_key_manager.encrypt_api_key(secret_val) + elif secret_val: + encrypted_secret = secret_val # Fallback if no encryption available + + webhook_id = str(uuid.uuid4())[:8] + db = SessionLocal() + try: + db.add(Webhook( + id=webhook_id, + name=name, + url=url, + secret=encrypted_secret, + events=events, + is_active=True, + )) + db.commit() + finally: + db.close() + + return {"id": webhook_id, "name": name} + + @router.post("/webhooks/{webhook_id}/test") + async def test_webhook(request: Request, webhook_id: str): + _require_admin(request) + db = SessionLocal() + try: + wh = db.query(Webhook).filter(Webhook.id == webhook_id).first() + if not wh: + raise HTTPException(404, "Webhook not found") + url, secret = wh.url, wh.secret + finally: + db.close() + + await webhook_manager.deliver_test(webhook_id, url, secret) + return {"status": "sent"} + + @router.patch("/webhooks/{webhook_id}") + def toggle_webhook(request: Request, webhook_id: str): + _require_admin(request) + db = SessionLocal() + try: + wh = db.query(Webhook).filter(Webhook.id == webhook_id).first() + if not wh: + raise HTTPException(404, "Webhook not found") + wh.is_active = not wh.is_active + db.commit() + return {"id": webhook_id, "is_active": wh.is_active} + finally: + db.close() + + @router.delete("/webhooks/{webhook_id}") + def delete_webhook(request: Request, webhook_id: str): + _require_admin(request) + db = SessionLocal() + try: + deleted = db.query(Webhook).filter(Webhook.id == webhook_id).delete() + db.commit() + if not deleted: + raise HTTPException(404, "Webhook not found") + finally: + db.close() + return {"status": "deleted"} + + # ================================================================ + # Sync Chat Endpoint (for n8n / Make / Activepieces) + # ================================================================ + + # Known provider base URLs — auto-resolved from api_key prefix or model name + KNOWN_PROVIDERS = { + "deepseek": "https://api.deepseek.com/v1", + "openai": "https://api.openai.com/v1", + "mistral": "https://api.mistral.ai/v1", + "groq": "https://api.groq.com/openai/v1", + "together": "https://api.together.xyz/v1", + "openrouter": "https://openrouter.ai/api/v1", + "ollama": "https://ollama.com/api", + "opencode-zen": "https://opencode.ai/zen/v1", + "opencode-go": "https://opencode.ai/zen/go/v1", + "fireworks": "https://api.fireworks.ai/inference/v1", + "venice": "https://api.venice.ai/api/v1", + "kimi-code": "https://api.kimi.com/coding/v1", + "kimicode": "https://api.kimi.com/coding/v1", + } + + # Model prefix → provider mapping for auto-detection + MODEL_PROVIDER_MAP = { + "deepseek": "deepseek", + "gpt-": "openai", + "o1": "openai", + "o3": "openai", + "o4": "openai", + "mistral": "mistral", + "llama": "groq", + "mixtral": "groq", + "kimi-for-coding": "kimi-code", + "kimi": "kimi-code", + } + + def _resolve_base_url(model: Optional[str], provider: Optional[str]) -> Optional[str]: + """Try to auto-resolve a base URL from provider name or model prefix.""" + if provider and provider.lower() in KNOWN_PROVIDERS: + return KNOWN_PROVIDERS[provider.lower()] + if model: + model_lower = model.lower() + for prefix, prov in MODEL_PROVIDER_MAP.items(): + if model_lower.startswith(prefix): + return KNOWN_PROVIDERS[prov] + return None + + class SyncChatRequest(BaseModel): + message: str = Field(..., max_length=MAX_MESSAGE_LEN) + model: Optional[str] = Field(None, max_length=200) + session: Optional[str] = Field(None, max_length=100) + api_key: Optional[str] = Field(None, max_length=256) + base_url: Optional[str] = Field(None, max_length=MAX_URL_LEN) + provider: Optional[str] = Field(None, max_length=50) + + @router.post("/v1/chat") + async def sync_chat(request: Request, body: SyncChatRequest): + if not getattr(request.state, "api_token", False): + raise HTTPException(403, "This endpoint requires an API token") + scopes = set(getattr(request.state, "api_token_scopes", []) or []) + if "chat" not in scopes: + raise HTTPException(403, "API token is not scoped for chat") + token_owner = getattr(request.state, "api_token_owner", None) + + from core.models import ChatMessage + from src.llm_core import llm_call_async + from src.endpoint_resolver import build_chat_url, build_headers, build_models_url, normalize_base + + message = body.message.strip() + if not message: + raise HTTPException(400, "Message is required") + + session_id = body.session + sess = None + + # --- Case 1: Resume an existing session --- + if session_id and session_manager: + try: + sess = session_manager.get_session(session_id) + except (KeyError, Exception): + raise HTTPException(404, "Session not found") + # SECURITY: verify the API-token's user owns this session — without + # this any token holder could resume any user's chat by passing its + # ID. The token's user is on request.state.user (set by API-token + # middleware); fall back to require_user if not present. + try: + from src.auth_helpers import get_current_user as _gcu + _tok_user = token_owner or getattr(request.state, "user", None) or _gcu(request) + except Exception: + _tok_user = None + # Strict ownership (see _caller_owns_session): fail closed so a + # null-owner / cross-owner session can't be resumed by an arbitrary + # chat-scoped token. + _sess_owner = getattr(sess, "owner", None) + if not _caller_owns_session(_sess_owner, _tok_user): + raise HTTPException(404, "Session not found") + + # --- Case 2: Direct API key + model (no pre-configured endpoint needed) --- + if not sess and body.api_key: + api_key = body.api_key.strip() + model = body.model or "deepseek-chat" + + # Validate only token-supplied direct base_url; auto-resolved known-provider + # URLs are not subject to extra local/LAN blocking beyond existing provider logic. + direct_base_url = body.base_url.strip().rstrip("/") if body.base_url else None + if direct_base_url: + try: + base_url = validate_public_http_url(direct_base_url) + except ValueError as e: + detail = str(e).replace("URL", "base_url", 1) + raise HTTPException(400, detail) + else: + base_url = _resolve_base_url(model, body.provider) + if not base_url: + raise HTTPException(400, + "Could not auto-detect provider. Pass base_url (e.g. 'https://api.deepseek.com/v1') " + "or provider ('deepseek', 'openai', 'groq', etc.)") + base_url = normalize_base(base_url) + endpoint_url = build_chat_url(base_url) + + if not session_manager: + raise HTTPException(500, "Session manager not available") + + sid = str(uuid.uuid4()) + sess = session_manager.create_session( + session_id=sid, name="API Chat", endpoint_url=endpoint_url, + model=model, owner=token_owner, + ) + sess.headers = build_headers(api_key, base_url) + session_manager.save_sessions() + session_id = sid + + # --- Case 3: Fall back to first configured ModelEndpoint --- + if not sess: + db = SessionLocal() + try: + ep = _select_api_chat_fallback_endpoint(db, token_owner) + finally: + db.close() + + if not ep: + raise HTTPException(400, + "No session, api_key, or configured endpoints. " + "Pass api_key + model, or configure an endpoint in Admin.") + + base_url = normalize_base(ep.base_url) + endpoint_url = build_chat_url(base_url) + model = body.model or "auto" + api_key = ep.api_key + if getattr(ep, "provider_auth_id", None): + try: + from src.endpoint_resolver import resolve_endpoint_runtime + base_url, api_key = resolve_endpoint_runtime(ep, owner=token_owner) + endpoint_url = build_chat_url(base_url) + except Exception: + raise HTTPException(500, "Could not resolve endpoint credentials") + + if model == "auto": + try: + async with httpx.AsyncClient(timeout=5) as client: + models_url = build_models_url(base_url) + hdrs = build_headers(api_key, base_url) + if models_url: + resp = await client.get(models_url, headers=hdrs) + resp.raise_for_status() + data = resp.json() + items = data if isinstance(data, list) else (data.get("data") or []) + ids = [m.get("id") for m in items if isinstance(m, dict) and m.get("id")] + if not ids and isinstance(data, dict): + ids = [ + m.get("name") or m.get("model") + for m in (data.get("models") or []) + if m.get("name") or m.get("model") + ] + else: + import json as _json + ids = _json.loads(ep.cached_models or "[]") + model = ids[0] if ids else "auto" + except Exception: + raise HTTPException(500, "Could not discover models from endpoint") + + if not session_manager: + raise HTTPException(500, "Session manager not available") + + sid = str(uuid.uuid4()) + sess = session_manager.create_session( + session_id=sid, name="API Chat", endpoint_url=endpoint_url, + model=model, owner=token_owner, + ) + if api_key: + sess.headers = build_headers(api_key, base_url) + session_manager.save_sessions() + session_id = sid + + # --- Send message and get response --- + sess.add_message(ChatMessage("user", message)) + + messages = [{"role": m.role, "content": m.content} for m in sess.history] + + reply = await llm_call_async( + sess.endpoint_url, sess.model, messages, + headers=sess.headers, timeout=120, + ) + sess.add_message(ChatMessage("assistant", reply)) + session_manager.save_sessions() + + webhook_manager.fire_and_forget("chat.completed", { + "session_id": session_id, "model": sess.model, + "user_message": message[:2000], "response": reply[:2000], + }) + + return {"response": reply, "session_id": session_id, "model": sess.model} + + return router diff --git a/services/memory/__init__.py b/services/memory/__init__.py index 31fa1d5fa..53fc80bd8 100644 --- a/services/memory/__init__.py +++ b/services/memory/__init__.py @@ -2,7 +2,7 @@ """Memory service — persistent memory storage and retrieval.""" from .service import MemoryService, Memory, MemorySearchResult -from .memory import MemoryManager, MemoryStoreUnreadable +from .memory import MemoryManager from .memory_vector import MemoryVectorStore __all__ = [ @@ -10,6 +10,5 @@ __all__ = [ "Memory", "MemorySearchResult", "MemoryManager", - "MemoryStoreUnreadable", "MemoryVectorStore", ] diff --git a/services/memory/memory.py b/services/memory/memory.py index b9aaaa2a8..031c13ac4 100644 --- a/services/memory/memory.py +++ b/services/memory/memory.py @@ -5,16 +5,6 @@ application runtime instantiates ``src.memory.MemoryManager``, so keeping a parallel implementation here risks silent drift between import paths. """ -from src.memory import ( - MemoryManager, - MemoryStoreUnreadable, - get_text_similarity, - tokenize, -) +from src.memory import MemoryManager, get_text_similarity, tokenize -__all__ = [ - "MemoryManager", - "MemoryStoreUnreadable", - "get_text_similarity", - "tokenize", -] +__all__ = ["MemoryManager", "get_text_similarity", "tokenize"] diff --git a/services/memory/memory_extractor.py b/services/memory/memory_extractor.py index 11539263b..e5f609250 100644 --- a/services/memory/memory_extractor.py +++ b/services/memory/memory_extractor.py @@ -17,8 +17,6 @@ import os import re from typing import Optional -from src.memory import MemoryStoreUnreadable - logger = logging.getLogger(__name__) @@ -389,13 +387,7 @@ async def extract_and_store( # Get owner from session _owner = getattr(session, 'owner', None) - # Strict load: this is a read-modify-write. Degrading to [] here would - # save only the newly extracted facts and drop the entire store. - try: - existing = memory_manager.load_all_for_update() - except MemoryStoreUnreadable as e: - logger.error("Skipping auto memory extraction, store unreadable: %s", e) - return + existing = memory_manager.load_all() added = 0 for fact in facts: @@ -634,18 +626,7 @@ async def audit_memories( # Merge audited entries back with other users' entries if owner: - # Strict load: the merge below reconstructs the whole file. If this - # degraded to [] we would save only this owner's audited slice and - # destroy every other tenant's memories. - try: - all_entries = memory_manager.load_all_for_update() - except MemoryStoreUnreadable as e: - logger.error("Aborting memory audit save, store unreadable: %s", e) - return { - "before": before_count, - "after": before_count, - "error": "store_unreadable", - } + all_entries = memory_manager.load_all() audited_ids = {e["id"] for e in final_entries} other_entries = [e for e in all_entries if e.get("owner") != owner and (e.get("owner") is not None)] # Also keep legacy entries that weren't part of this audit diff --git a/src/ai_interaction.py b/src/ai_interaction.py index e777ca32a..9ee97368f 100644 --- a/src/ai_interaction.py +++ b/src/ai_interaction.py @@ -22,7 +22,6 @@ import time from typing import Any, Awaitable, Callable, Dict, Optional, Tuple from src.constants import GENERATED_IMAGES_DIR -from src.memory import MemoryStoreUnreadable logger = logging.getLogger(__name__) @@ -385,15 +384,7 @@ async def do_manage_memory(content: str, session_id: Optional[str] = None, owner return {"error": "Memory text cannot be empty"} entry = _memory_manager.add_entry(text, source="ai_agent", category=category, owner=owner) - # Strict load: this is a read-modify-write, and it is the path an - # ordinary "remember that I prefer X" takes. Degrading to [] here would - # save just this one entry over a store we only failed to read, - # atomically destroying every memory in it (issue #5673). - try: - memories = _memory_manager.load_all_for_update() - except MemoryStoreUnreadable as e: - logger.error("Refusing to add memory, store unreadable: %s", e) - return {"error": "Memory store is temporarily unreadable — nothing was saved."} + memories = _memory_manager.load_all() memories.append(entry) _memory_manager.save(memories) diff --git a/src/integrations.py b/src/integrations.py index 52dd4b2d1..aa6c4982e 100644 --- a/src/integrations.py +++ b/src/integrations.py @@ -1,14 +1,11 @@ -import ipaddress import json import os -import time import uuid import logging import re from typing import Dict, List, Optional, Any from urllib.parse import urljoin, urlparse, urlunparse -import httpcore import httpx from fastapi import HTTPException @@ -357,152 +354,6 @@ def _find_integration(identifier: str) -> Optional[Dict[str, Any]]: return None -# httpcore raises its own exception hierarchy; map the ones a simple request can -# surface back to their httpx equivalents so the caller's `except httpx.*` blocks -# below behave exactly as they did with the default transport. -_HTTPCORE_TO_HTTPX_EXC = { - httpcore.ConnectError: httpx.ConnectError, - httpcore.ConnectTimeout: httpx.ConnectTimeout, - httpcore.NetworkError: httpx.NetworkError, - httpcore.PoolTimeout: httpx.PoolTimeout, - httpcore.ProtocolError: httpx.ProtocolError, - httpcore.ReadError: httpx.ReadError, - httpcore.ReadTimeout: httpx.ReadTimeout, - httpcore.RemoteProtocolError: httpx.RemoteProtocolError, - httpcore.TimeoutException: httpx.TimeoutException, - httpcore.WriteError: httpx.WriteError, - httpcore.WriteTimeout: httpx.WriteTimeout, -} - - -class _PinnedAsyncBackend(httpcore.AsyncNetworkBackend): - """Network backend that connects only to the pre-validated IPs, in order. - - Every address here came out of the single SSRF resolution, so moving to the - next one after a connect failure is not re-resolution — it's ordinary - multi-address fallback restricted to the set the guard already approved. - httpcore takes TLS SNI and the ``Host`` header from the request URL rather - than the connect host, so pinning the socket destination leaves certificate - validation and vhost routing pointed at the original hostname. - """ - - def __init__(self, ips: List[ipaddress._BaseAddress]): - self._ips = [str(ip) for ip in ips] - self._real = httpcore.AnyIOBackend() - - async def connect_tcp(self, host, port, timeout=None, local_address=None, - socket_options=None): - # One shared connect budget: each attempt gets the time left until the - # original deadline, so N dead addresses can't stretch the connect - # phase to N * timeout. - deadline = None if timeout is None else time.monotonic() + timeout - last_exc: Optional[Exception] = None - for ip in self._ips: - remaining = None if deadline is None else max(0.0, deadline - time.monotonic()) - try: - return await self._real.connect_tcp( - ip, port, remaining, local_address, socket_options - ) - except (httpcore.ConnectError, httpcore.ConnectTimeout) as exc: - last_exc = exc - if deadline is not None and time.monotonic() >= deadline: - break - raise last_exc - - async def connect_unix_socket(self, path, timeout=None, socket_options=None): - return await self._real.connect_unix_socket(path, timeout, socket_options) - - async def sleep(self, seconds: float) -> None: - return await self._real.sleep(seconds) - - -class _PinnedAsyncTransport(httpx.AsyncBaseTransport): - """httpx transport that pins the TCP connect to the pre-resolved IP(s). - - Kept local, mirroring the per-module pinned transports web fetch and - webhook delivery already carry, rather than coupling api_call to the - webhook subsystem. The request URL passes through unchanged, so SNI and the - ``Host`` header stay the original hostname; only the socket destination is - pinned, which is what closes the rebinding window. - """ - - def __init__(self, ips: List[ipaddress._BaseAddress]): - self._pinned_ips = list(ips) - self._pool = httpcore.AsyncConnectionPool( - # Reuse the CA trust the default httpx client would build (certifi - # plus SSL_CERT_FILE / SSL_CERT_DIR when trust_env is set) so - # swapping in this transport doesn't quietly change which chains - # verify. ssl.create_default_context() would use system roots. - ssl_context=httpx.create_ssl_context(), - http1=True, - http2=False, - network_backend=_PinnedAsyncBackend(ips), - ) - - async def handle_async_request(self, request: httpx.Request) -> httpx.Response: - core_req = httpcore.Request( - method=request.method, - url=httpcore.URL( - scheme=request.url.raw_scheme, - host=request.url.raw_host, - port=request.url.port, - target=request.url.raw_path, - ), - headers=request.headers.raw, - content=request.stream, - extensions=request.extensions, - ) - try: - core_resp = await self._pool.handle_async_request(core_req) - content = b"".join([chunk async for chunk in core_resp.aiter_stream()]) - await core_resp.aclose() - except Exception as exc: - mapped = _HTTPCORE_TO_HTTPX_EXC.get(type(exc)) - if mapped is not None: - raise mapped(str(exc)) from exc - raise - return httpx.Response( - status_code=core_resp.status, - headers=core_resp.headers, - content=content, - extensions=core_resp.extensions, - ) - - async def aclose(self) -> None: - await self._pool.aclose() - - -def _validated_ips(raw_ips: List[str]) -> List[ipaddress._BaseAddress]: - """Return every entry that parses as an IP address, de-duplicated, order - preserved. - - check_outbound_url only reports ok when *all* of these classify as safe, so - the whole list is guard-approved and any of them is a legitimate connect - target. Skipping unparseable entries mirrors how the guard walks the same - resolver output. - - De-duplication matters because the resolver is getaddrinfo(host, None) with - no socktype filter, so glibc reports the same address once per socktype - (SOCK_STREAM/SOCK_DGRAM/SOCK_RAW) — a single-homed host comes back three - times. Without this, the connect fallback would spend the shared deadline - retrying one dead address instead of moving on to a genuinely different one. - """ - ips: List[ipaddress._BaseAddress] = [] - seen = set() - for raw in raw_ips: - if not isinstance(raw, str): - continue - try: - ip = ipaddress.ip_address(raw.split("%")[0]) # strip IPv6 zone id - except ValueError: - continue - if ip in seen: - continue - seen.add(ip) - ips.append(ip) - return ips - - async def execute_api_call( integration_id: str, method: str, @@ -558,31 +409,13 @@ async def execute_api_call( # loopback for locked-down deployments. Private stays allowed by default # because LAN integrations (Home Assistant, Miniflux, ntfy) are the # primary use case. - from src.url_safety import check_outbound_url, _default_resolver + from src.url_safety import check_outbound_url block_private = os.getenv( "INTEGRATION_API_BLOCK_PRIVATE_IPS", "false" ).lower() == "true" - # Resolve the host exactly once and remember the IPs the guard validated so - # the request below can be pinned to them. check_outbound_url only reports - # (ok, reason); a plain httpx client re-resolves the host at connect time, - # which reopens a DNS-rebinding TOCTOU — a base_url host that answers with a - # public IP for the guard and then flips to 169.254.169.254 for the connect - # would reach cloud metadata with the integration's auth headers attached. - resolved_ips: List[str] = [] - - def _recording_resolver(host: str) -> List[str]: - ips = _default_resolver(host) - resolved_ips[:] = ips - return ips - - ok, reason = check_outbound_url( - url, block_private=block_private, resolver=_recording_resolver - ) + ok, reason = check_outbound_url(url, block_private=block_private) if not ok: return {"error": f"URL rejected: {reason}", "exit_code": 1} - pinned_ips = _validated_ips(resolved_ips) - if not pinned_ips: - return {"error": "URL rejected: host did not resolve to a usable address", "exit_code": 1} method = method.upper() @@ -622,9 +455,7 @@ async def execute_api_call( auth = httpx.BasicAuth(parts[0], parts[1]) try: - async with httpx.AsyncClient( - timeout=30.0, transport=_PinnedAsyncTransport(pinned_ips) - ) as client: + async with httpx.AsyncClient(timeout=30.0) as client: response = await client.request( method, url, diff --git a/src/memory.py b/src/memory.py index 92efbf5b2..1d8cdbc1e 100644 --- a/src/memory.py +++ b/src/memory.py @@ -10,18 +10,6 @@ from datetime import datetime logger = logging.getLogger(__name__) - -class MemoryStoreUnreadable(RuntimeError): - """memory.json exists on disk but could not be read or parsed. - - "The contents are unknown" is categorically different from "there are no - memories". A read-modify-write caller that conflates the two appends to an - empty view and then persists it, destroying the whole store — the writes - are atomic, so the loss is durable. Raised by - :meth:`MemoryManager.load_all_for_update` so those callers fail closed. - """ - - def tokenize(text: str) -> List[str]: """Simple tokenizer that splits on whitespace and removes punctuation.""" return [word.strip('.,!?";') for word in text.split()] @@ -122,69 +110,21 @@ class MemoryManager: with open(self.memory_file, 'w', encoding='utf-8') as f: json.dump([], f, ensure_ascii=False, indent=2) - def _read_entries(self) -> List[Dict]: - """Parse the store, or raise :class:`MemoryStoreUnreadable`. - - Returns ``[]`` only when the file genuinely does not exist. Every other - failure mode raises, so callers can tell "no memories" apart from - "couldn't read the memories". - """ + def load_all(self) -> List[Dict]: + """Load all memory entries from JSON file (unfiltered).""" if not os.path.exists(self.memory_file): return [] try: with open(self.memory_file, "r", encoding="utf-8") as f: data = json.load(f) - except OSError as e: - # PermissionError is an OSError (a scanner holding the file, a - # permissions problem, bad media). - raise MemoryStoreUnreadable( - f"cannot read {self.memory_file}: {e}" - ) from e - except json.JSONDecodeError as e: - # This is the branch that actually destroyed stores: the file reads - # back fine, so nothing stops the save that follows. A truncated - # memory.json is reachable because core/database.py rewrites it with - # a plain open(..,"w") + json.dump during migration. - # - # Preserved behaviour: a corrupt store still gets one shot at the - # pre-JSON memory.txt migration. Only raise when that finds nothing, - # so we never report "empty" for a store we simply failed to parse. - legacy = self._migrate_from_legacy() - if legacy: - return legacy - raise MemoryStoreUnreadable( - f"{self.memory_file} is not valid JSON: {e}" - ) from e - - if not isinstance(data, list): - raise MemoryStoreUnreadable( - f"{self.memory_file} is not a JSON array (got {type(data).__name__})" - ) - return self._validate_entries(data) - - def load_all(self) -> List[Dict]: - """Load all memory entries from JSON file (unfiltered). - - Lenient by design: this feeds display, search, and context-injection - paths, so an unreadable store degrades to an empty list rather than - breaking chat. Never build a value from this that you intend to save - back — use :meth:`load_all_for_update` for that. - """ - try: - return self._read_entries() - except MemoryStoreUnreadable as e: + if isinstance(data, list): + return self._validate_entries(data) + except (json.JSONDecodeError, PermissionError) as e: logger.error("Error loading memory.json: %s", e) - return [] + return self._migrate_from_legacy() - def load_all_for_update(self) -> List[Dict]: - """Load for a read-modify-write cycle. - - Propagates :class:`MemoryStoreUnreadable` instead of degrading to ``[]`` - so a caller can never append to an empty view and persist it over a - store that was only temporarily unreadable (issue #5673). - """ - return self._read_entries() + return [] def load(self, owner: str = None) -> List[Dict]: """Load memory entries, optionally filtered by owner.""" @@ -195,12 +135,7 @@ class MemoryManager: def claim_ownerless(self, owner: str): """Assign all ownerless memory entries to the given owner.""" - try: - entries = self.load_all_for_update() - except MemoryStoreUnreadable as e: - # Skip the sweep rather than rewrite the store from an unknown view. - logger.error("Skipping ownerless claim, memory store unreadable: %s", e) - return + entries = self.load_all() changed = False claimed = 0 for entry in entries: @@ -300,12 +235,7 @@ class MemoryManager: if not ids: return id_set = set(ids) - try: - entries = self.load_all_for_update() - except MemoryStoreUnreadable as e: - # Best-effort counter; never worth rewriting the store blind. - logger.error("Skipping uses bump, memory store unreadable: %s", e) - return + entries = self.load_all() changed = False for e in entries: if e.get("id") in id_set: diff --git a/src/memory_provider.py b/src/memory_provider.py index 8974a6e84..925c59192 100644 --- a/src/memory_provider.py +++ b/src/memory_provider.py @@ -157,11 +157,7 @@ class NativeMemoryProvider(MemoryProvider): if metadata: entry["metadata"] = dict(metadata) - # Strict load: read-modify-write. `load_all` degrades an unreadable - # store to [], which would save this single entry over everything - # already stored (issue #5673). The provider API has no error channel, - # so MemoryStoreUnreadable propagates to the caller. - memories = self.memory_manager.load_all_for_update() + memories = self.memory_manager.load_all() memories.append(entry) self.memory_manager.save(memories) @@ -227,10 +223,7 @@ class NativeMemoryProvider(MemoryProvider): ] async def delete(self, memory_id: str, *, owner: Optional[str] = None) -> bool: - # Strict load for the same reason: `remaining` is derived from this - # list and saved back, so it must never be built from a store we - # failed to read. - memories = self.memory_manager.load_all_for_update() + memories = self.memory_manager.load_all() remaining = [] deleted_id = None diff --git a/src/tool_parsing.py b/src/tool_parsing.py index 98dc1b5f6..2885cc00f 100644 --- a/src/tool_parsing.py +++ b/src/tool_parsing.py @@ -187,12 +187,8 @@ _FUNCTION_MODEL_NAME_RE = re.compile( _FUNCTION_MODEL_PARAMS_OPEN_RE = re.compile(r"\s*", re.IGNORECASE) _FUNCTION_MODEL_PARAMS_CLOSE_RE = re.compile(r"", re.IGNORECASE) _QWEN_ROLE_MARKER_RE = re.compile(r"?|?", re.IGNORECASE) -# At least one pipe is required around `end`. Both pipes used to be optional -# (`\|?end\|?`), which also matched a bare `end` on its own line and deleted it -# from ordinary prose and from Ruby/Lua/shell snippets that close blocks with -# one; see #5547. `|end`, `end|`, `|end|` and `/|end|` still strip as before. _QWEN_BARE_MARKER_RE = re.compile( - r"(?:^|[\t\r\n ])(?:/?\|end\||\|end|end\|)(?=[\t\r\n ]|$)|" + r"(?:^|[\t\r\n ])(?:\|?end\|?|/?\|end\|)(?=[\t\r\n ]|$)|" r"(?:^|[\t\r\n ])assistan(?:t)?(?=[\t\r\n ]|$)", re.IGNORECASE, ) diff --git a/src/tools/system.py b/src/tools/system.py index c2eb9ceab..813d57df2 100644 --- a/src/tools/system.py +++ b/src/tools/system.py @@ -46,9 +46,7 @@ async def do_manage_skills(content: str, owner: Optional[str] = None) -> Dict: except ValueError: return {"error": "Invalid JSON arguments", "exit_code": 1} - action = (args.get("action") or "").strip().lower() - if not action: - return {"error": "action is required (list|view|view_ref|add|edit|patch|publish|delete|search)", "exit_code": 1} + action = (args.get("action") or "").lower() from services.memory.skills import SkillsManager from services.memory.skill_format import Skill, slugify from src.constants import DATA_DIR @@ -57,7 +55,7 @@ async def do_manage_skills(content: str, owner: Optional[str] = None) -> Dict: # Accept legacy `skill_id` as an alias for `name`. name = (args.get("name") or args.get("skill_id") or "").strip() - if action in ("list", "index"): + if action in ("list", "index", ""): all_skills = sm.load(owner=owner) if not all_skills: return {"results": "No skills yet. Create one with action='add'."} diff --git a/static/app.js b/static/app.js index 2f1e8d4bf..97f0ae77e 100644 --- a/static/app.js +++ b/static/app.js @@ -10,14 +10,14 @@ import modelsModule from './js/models.js?v=20260715startupcalm2'; import ragModule from './js/rag.js'; import presetsModule from './js/presets.js'; import searchModule from './js/search.js'; -import chatModule from './js/chat.js?v=20260801fix1'; +import chatModule from './js/chat.js?v=20260722ctxheader4'; import compareModule from './js/compare/index.js?v=20260723compareicon2'; import documentModule from './js/document.js?v=20260722emailfastindex1'; import searchChatModule from './js/search-chat.js'; import { makeWindowDraggable } from './js/windowDrag.js'; import markdownModule from './js/markdown.js'; import chatRenderer from './js/chatRenderer.js?v=20260722emailfastindex1'; -import sessionModule from './js/sessions.js'; +import sessionModule from './js/sessions.js?v=20260722ctxheader4'; import memoryModule from './js/memory.js?v=20260722memoryloading1'; import voiceRecorderModule from './js/voiceRecorder.js'; import censorModule from './js/censor.js'; @@ -1689,20 +1689,12 @@ function initializeEventListeners() { const newMemoryInput = el('new-memory-input'); if (newMemoryInput) { - // keydown, not the deprecated keypress: keypress is not guaranteed to - // fire for Enter everywhere, which left the Add Memory form with no - // working submit path (#5828). - newMemoryInput.addEventListener('keydown', (e) => { - if (e.key === 'Enter' && !e.isComposing) { - e.preventDefault(); + newMemoryInput.addEventListener('keypress', (e) => { + if (e.key === 'Enter') { memoryModule.addNewMemory(); } }); } - const newMemoryAddBtn = el('new-memory-add-btn'); - if (newMemoryAddBtn) { - newMemoryAddBtn.addEventListener('click', () => memoryModule.addNewMemory()); - } // Voice recording is handled by the dual-purpose send/mic button (see below) @@ -3916,10 +3908,85 @@ function startOdysseusApp() { const messageInput = el('message'); const modelPickerWrap = document.getElementById('model-picker-wrap'); - // ArrowUp/ArrowDown prompt recall on #message lives in - // static/js/composerArrowUpRecall.js (wired from chat.js). Do not re-add a - // copy here: two capture-phase listeners on the same textarea meant the one - // without the draft guard won and ate unsent multi-line prompts (#5862). + function _readComposerPromptHistory() { + const chatBox = document.getElementById('chat-history'); + if (!chatBox) return []; + return Array.from(chatBox.querySelectorAll('.msg-user')) + .reverse() + .map(msg => { + const body = msg.querySelector('.body'); + return msg.dataset?.raw || (body ? body.textContent : '') || ''; + }) + .filter(Boolean); + } + + if (messageInput && !messageInput._odysseusPromptRecallCapture) { + messageInput._odysseusPromptRecallCapture = true; + let recallHistory = []; + let recallIndex = -1; + let lastRecalled = ''; + const norm = (v) => String(v || '').replace(/\r\n/g, '\n').trimEnd(); + messageInput.addEventListener('input', () => { + if (norm(messageInput.value) === norm(lastRecalled)) return; + recallHistory = []; + recallIndex = -1; + lastRecalled = ''; + try { delete messageInput.dataset.odysseusRecallIndex; } catch {} + }, true); + messageInput.addEventListener('keydown', (e) => { + if (e.key !== 'ArrowUp' && e.key !== 'ArrowDown') return; + if (e.shiftKey || e.altKey || e.ctrlKey || e.metaKey || e.isComposing) return; + if (window._ghostAutocomplete?.isActive?.()) return; + const fresh = _readComposerPromptHistory(); + const history = fresh.length ? fresh : recallHistory; + if (!history.length) return; + const current = norm(messageInput.value); + let currentIndex = current ? history.findIndex(item => norm(item) === current) : -1; + if (current && currentIndex < 0 && current === norm(lastRecalled)) currentIndex = recallIndex; + if (current && currentIndex < 0) { + const markedIndex = Number(messageInput.dataset.odysseusRecallIndex); + if (Number.isInteger(markedIndex) && markedIndex >= 0 && markedIndex < history.length) { + currentIndex = markedIndex; + } + } + e.preventDefault(); + e.stopPropagation(); + e.stopImmediatePropagation(); + if (e.key === 'ArrowDown') { + if (currentIndex < 0) return; + const nextIndex = currentIndex - 1; + if (nextIndex < 0) { + recallHistory = history; + recallIndex = -1; + lastRecalled = ''; + try { delete messageInput.dataset.odysseusRecallIndex; } catch {} + messageInput.value = ''; + try { messageInput.selectionStart = messageInput.selectionEnd = 0; } catch {} + try { uiModule.autoResize(messageInput); } catch {} + return; + } + const recalled = history[nextIndex]; + recallHistory = history; + recallIndex = nextIndex; + lastRecalled = recalled; + try { messageInput.dataset.odysseusRecallIndex = String(nextIndex); } catch {} + messageInput.value = recalled; + try { messageInput.selectionStart = messageInput.selectionEnd = recalled.length; } catch {} + try { uiModule.autoResize(messageInput); } catch {} + return; + } + const nextIndex = currentIndex >= 0 ? Math.min(currentIndex + 1, history.length - 1) : 0; + const recalled = history[nextIndex]; + if (!recalled) return; + recallHistory = history; + recallIndex = nextIndex; + lastRecalled = recalled; + try { messageInput.dataset.odysseusRecallIndex = String(nextIndex); } catch {} + messageInput.value = recalled; + try { messageInput.selectionStart = messageInput.selectionEnd = recalled.length; } catch {} + try { uiModule.autoResize(messageInput); } catch {} + }, true); + } const _sendIcon = ''; const _micIcon = ''; diff --git a/static/index.html b/static/index.html index fea4e20ac..8257660fe 100644 --- a/static/index.html +++ b/static/index.html @@ -250,9 +250,9 @@ - + - + @@ -365,7 +365,6 @@ Add a memory — e.g. 'I prefer concise replies' -
@@ -1006,7 +1005,7 @@ var tips = mobile ? phone : desktop; var el = document.getElementById('welcome-tip'); if (el) { - el.textContent = tips[Math.floor(Math.random() * tips.length)]; + el.textContent = 'Pick a model if you want, or just type.'; } fetch('/api/version').then(function(r){return r.json()}).then(function(d){ if (d.version) window._appVersion = d.version; @@ -2505,7 +2504,7 @@ - + @@ -2523,7 +2522,7 @@ - + diff --git a/static/js/chat.js b/static/js/chat.js index 3c8bbe850..ea2d8c1bb 100644 --- a/static/js/chat.js +++ b/static/js/chat.js @@ -349,9 +349,6 @@ import { wireArrowUpRecall, getUserMessagesFromChatHistory } from './composerArr async function _adoptOpenedSessionBeforeAutoCreate() { if (!sessionModule || !sessionModule.getCurrentSessionId || sessionModule.getCurrentSessionId()) return true; - // Don't adopt a stale session when the user explicitly started a New Chat - // (pending state set) — the send path must materialize the pending session. - if (sessionModule.hasPendingChat && sessionModule.hasPendingChat()) return false; const activeRowId = document.querySelector('.list-item.active-session[data-session-id], .session-item.active[data-session-id]')?.dataset?.sessionId || ''; const hashId = _hashSessionCandidate(); const lastSelectedId = String(window.__odysseusLastSelectedSessionId || '').trim(); @@ -1406,8 +1403,6 @@ import { wireArrowUpRecall, getUserMessagesFromChatHistory } from './composerArr currentAccumulated = ''; currentHolder = null; - let abortCtrl = null; - let streamingTTS = false; try { // Re-enable auto-scroll when user sends a message uiModule.setAutoScroll(true); @@ -1721,7 +1716,7 @@ import { wireArrowUpRecall, getUserMessagesFromChatHistory } from './composerArr } - abortCtrl = new AbortController(); + const abortCtrl = new AbortController(); abortCtrl._reason = ''; currentAbort = abortCtrl; @@ -1902,7 +1897,7 @@ import { wireArrowUpRecall, getUserMessagesFromChatHistory } from './composerArr let isThinking = false; let thinkingStartTime = null; // Streaming TTS: synthesize sentence-by-sentence during streaming - streamingTTS = !!(window.aiTTSManager && window.aiTTSManager.autoPlay && window.aiTTSManager.available); + const streamingTTS = !!(window.aiTTSManager && window.aiTTSManager.autoPlay && window.aiTTSManager.available); if (streamingTTS) window.aiTTSManager.streamingStart(); // Multi-bubble agent tracking let roundHolder = holder; // Current AI text bubble (changes per round) @@ -4792,8 +4787,7 @@ import { wireArrowUpRecall, getUserMessagesFromChatHistory } from './composerArr if (msgIndex < 0) return; const bodyEl = userMsgElement.querySelector('.body'); - let currentText = (userMsgElement.dataset.raw || (bodyEl ? bodyEl.textContent : '') || '').trim(); - currentText = currentText.replace(/\s*\[\d+ attachment\(s\)\]$/, ''); + const currentText = bodyEl ? bodyEl.textContent.trim().replace(/\s*\[\d+ attachment\(s\)\]$/, '') : ''; // Replace body with an editable textarea const editor = document.createElement('textarea'); diff --git a/static/js/chatRenderer.js b/static/js/chatRenderer.js index 1d6e2e4a9..10709679d 100644 --- a/static/js/chatRenderer.js +++ b/static/js/chatRenderer.js @@ -478,10 +478,7 @@ const DSML_STRAY_RE = /<\s*\/?\s*[||]+\s*DSML\s*[||]+[^>]*>/gi; const DSML_INVOKE_RE = /<\s*[||]+\s*DSML\s*[||]+\s*invoke\b[^>]*>[\s\S]*?(?:<\s*\/\s*[||]+\s*DSML\s*[||]+\s*invoke\s*>|$)/gi; const RAW_OPENAI_TOOL_JSON_RE = /(?:\[\s*)?\{\s*"function"\s*:\s*\{[\s\S]*?\}\s*,\s*"id"\s*:\s*"[^"]*"\s*,\s*"type"\s*:\s*"function"\s*\}\s*\]?/gi; const QWEN_ROLE_MARKER_RE = /<\/?\|(?:assistant|assistan|user|system|tool)\|>?|<\/\|end\|>?/gi; -// Keep in sync with _QWEN_BARE_MARKER_RE in src/tool_parsing.py. At least one -// pipe is required around `end`: with both optional (`\|?end\|?`) this also ate -// a bare `end` on its own line, breaking Ruby/Lua/shell snippets (#5547). -const QWEN_BARE_MARKER_RE = /(?:^|[\t\r\n ])(?:\/?\|end\||\|end|end\|)(?=[\t\r\n ]|$)|(?:^|[\t\r\n ])assistan(?:t)?(?=[\t\r\n ]|$)/gi; +const QWEN_BARE_MARKER_RE = /(?:^|[\t\r\n ])(?:\|?end\|?|\/?\|end\|)(?=[\t\r\n ]|$)|(?:^|[\t\r\n ])assistan(?:t)?(?=[\t\r\n ]|$)/gi; // Self-narration about tool results (model echoing stdout/exit_code) const TOOL_NARRATION_RE = /(?:The (?:result|output) shows?:?\s*)?-?\s*(?:stdout|stderr|exit_code):\s*.+/gi; diff --git a/static/js/composerArrowUpRecall.js b/static/js/composerArrowUpRecall.js index 83141bfe9..e0b20d6b4 100644 --- a/static/js/composerArrowUpRecall.js +++ b/static/js/composerArrowUpRecall.js @@ -143,9 +143,9 @@ export function wireArrowUpRecall(composer, getUserMessages, options = {}) { return; } - // ArrowUp walks older prompts. An unmatched draft already returned above, - // so reaching here means the composer is empty or holds a recalled prompt - // — the caret-navigation case is never hijacked. + // ArrowUp owns prompt history in the chat composer. If the current text + // is not already a recalled prompt, start from newest instead of letting + // the browser move the caret inside the textarea. const nextIndex = currentIndex >= 0 ? Math.min(currentIndex + 1, history.length - 1) : 0; const recalled = history[nextIndex]; if (!recalled) { diff --git a/static/js/emailLibrary.js b/static/js/emailLibrary.js index 32b906ddc..6a0d3e294 100644 --- a/static/js/emailLibrary.js +++ b/static/js/emailLibrary.js @@ -13,7 +13,7 @@ import { makeWindowDraggable } from './windowDrag.js'; import { _esc, _escLinkify, _extractName, _parseTurnMeta, _formatBubbleDate, _formatRecipients, _senderColor, _initials, - _sanitizeHtml, _renderEmailSummaryError, + _sanitizeHtml, _TALON_WROTE, _TALON_FROM, _TALON_SENT, _TALON_SUBJ, _TALON_TO, _TALON_ORIG_RE, _SIG_BLOAT_MIN_CHARS, } from './emailLibrary/utils.js'; @@ -7259,11 +7259,12 @@ async function _generateSummary(reader, data, btn) { if (label) label.textContent = 'Summary'; } } else { - _renderEmailSummaryError(content, result); + content.innerHTML = `${_esc(result.error || 'Failed to summarize')}`; + panel.remove(); } } catch (e) { sp.destroy(); - _renderEmailSummaryError(content, null); + panel.remove(); if (uiModule) uiModule.showError?.('Failed to summarize'); } finally { if (btn) btn.disabled = false; diff --git a/static/js/emailLibrary/utils.js b/static/js/emailLibrary/utils.js index f634c9949..82a5c86ec 100644 --- a/static/js/emailLibrary/utils.js +++ b/static/js/emailLibrary/utils.js @@ -30,25 +30,6 @@ export function _esc(text) { return div.innerHTML; } -const _EMAIL_SUMMARY_ERROR_MESSAGES = Object.freeze({ - email_summary_missing_body: 'No email body to summarize', - email_summary_not_configured: 'No model configured for email summaries', - email_summary_empty: 'The model returned an empty summary', - email_summary_unavailable: 'Failed to summarize', -}); - -export function _emailSummaryErrorMessage(result) { - const code = String(result?.error_code || ''); - return _EMAIL_SUMMARY_ERROR_MESSAGES[code] || 'Failed to summarize'; -} - -export function _renderEmailSummaryError(container, result) { - const message = container.ownerDocument.createElement('span'); - message.style.color = 'var(--red)'; - message.textContent = _emailSummaryErrorMessage(result); - container.replaceChildren(message); -} - function _attrEsc(text) { return String(text ?? '') .replace(/"/g, '"') diff --git a/static/js/sessions.js b/static/js/sessions.js index edf83c8a4..cf59d478c 100644 --- a/static/js/sessions.js +++ b/static/js/sessions.js @@ -1847,10 +1847,6 @@ export async function selectSession(id, { keepSidebar = false, showLoading = tru const _isTransientChat = !!_meta && (_meta.folder === 'Assistant' || _meta.folder === 'Tasks'); if (!_isTransientChat) { Storage.set('lastSessionId', id); - // Update URL hash without triggering hashchange handler - if (window.location.hash !== '#' + id) { - history.replaceState(null, '', '#' + id); - } } // Restore character preset for persistent chats try { @@ -2317,7 +2313,6 @@ export async function materializePendingSession() { currentSessionId = payload.id; if (!isIncognito) { Storage.set('lastSessionId', payload.id); - history.replaceState(null, '', '#' + payload.id); } // Reload the sidebar in the background. Awaiting this used to block the first diff --git a/static/js/skills.js b/static/js/skills.js index b45403570..84974d446 100644 --- a/static/js/skills.js +++ b/static/js/skills.js @@ -83,9 +83,11 @@ export async function loadSkills(cascade = false) { // Play the domino-in entrance on this load (set when the tab is opened, // not for the silent re-loads after an edit/delete). if (cascade) _cascadeNext = true; - // Always re-fetch when the tab is explicitly opened — the cascade - // animation is handled inside renderSkillsList() via _cascadeNext. - // Skipping the fetch here caused stale data on panel close/reopen (#5870). + if (cascade && loaded && !_loadPromise && _playSkillsCascade()) { + _cascadeNext = false; + updateCount(); + return; + } if (_loadPromise) return _loadPromise; _loadPromise = (async () => { try { diff --git a/tests/test_api_chat_security.py b/tests/test_api_chat_security.py index d92a31620..7dcec324e 100644 --- a/tests/test_api_chat_security.py +++ b/tests/test_api_chat_security.py @@ -76,7 +76,7 @@ def _load_webhook_routes_for_test(monkeypatch): module_name = "routes.webhook_routes_under_test" spec = importlib.util.spec_from_file_location( module_name, - Path(__file__).resolve().parent.parent / "routes" / "webhook" / "webhook_routes.py", + Path(__file__).resolve().parent.parent / "routes" / "webhook_routes.py", ) module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) diff --git a/tests/test_backup_import_cross_user_dedup.py b/tests/test_backup_import_cross_user_dedup.py index 135be78ee..2df5936ef 100644 --- a/tests/test_backup_import_cross_user_dedup.py +++ b/tests/test_backup_import_cross_user_dedup.py @@ -27,9 +27,6 @@ def _setup(monkeypatch, store, user="alice"): mem = MagicMock() mem.load_all.return_value = list(store) - # import_data reads through the strict loader so a store it cannot read is - # never overwritten (#5673); the double has to offer the same entry point. - mem.load_all_for_update.return_value = list(store) saved = {} mem.save.side_effect = lambda entries: saved.__setitem__("entries", entries) diff --git a/tests/test_composer_arrow_up_recall_js.py b/tests/test_composer_arrow_up_recall_js.py index 022fcbc02..eadc3bc94 100644 --- a/tests/test_composer_arrow_up_recall_js.py +++ b/tests/test_composer_arrow_up_recall_js.py @@ -306,24 +306,3 @@ def test_integration_recalls_from_chat_history_dom(): ) assert proc.returncode == 0, proc.stderr assert json.loads(proc.stdout.strip()) == {"value": "stored prompt", "prevented": True} - - -def test_prompt_recall_is_not_duplicated_in_app_js(): - """Only composerArrowUpRecall.js may own ArrowUp on #message (issue #5862). - - static/app.js once carried a near-verbatim copy of this recall logic, wired - as a second capture-phase listener on the same textarea. That copy lacked - the draft guard here, and because it called stopImmediatePropagation it won - regardless of registration order — so a typed multi-line prompt was replaced - by the last sent one instead of the caret moving up a line. - """ - app_js = (_REPO / "static" / "app.js").read_text(encoding="utf-8") - for marker in ( - "_odysseusPromptRecallCapture", - "_readComposerPromptHistory", - "odysseusRecallIndex", - ): - assert marker not in app_js, ( - f"static/app.js reintroduces prompt recall ({marker!r}); " - "it belongs to static/js/composerArrowUpRecall.js alone" - ) diff --git a/tests/test_document_routes_shim.py b/tests/test_document_routes_shim.py deleted file mode 100644 index 68d049a62..000000000 --- a/tests/test_document_routes_shim.py +++ /dev/null @@ -1,29 +0,0 @@ -"""Regression test for the document route shim (slice 2m, #4082/#4071). - -The backward-compat shims at ``routes/document_routes.py`` and -``routes/document_helpers.py`` use ``sys.modules`` replacement so the legacy -import paths and the canonical ``routes.document.*`` paths resolve to the -*same* module objects. This is required because multiple tests do -``import routes.document_routes as droutes`` followed by -``droutes.SessionLocal = ...`` / ``monkeypatch.setattr(droutes, ...)`` and -``sys.modules.pop("routes.document_helpers")`` + re-import — for those to -take effect at runtime, the legacy and canonical module objects must be -identical. -""" - -import importlib - -import routes.document_routes as _shim_routes # noqa: F401 -import routes.document_helpers as _shim_helpers # noqa: F401 - - -def test_legacy_and_canonical_routes_are_same_object(): - legacy = importlib.import_module("routes.document_routes") - canonical = importlib.import_module("routes.document.document_routes") - assert legacy is canonical - - -def test_legacy_and_canonical_helpers_are_same_object(): - legacy = importlib.import_module("routes.document_helpers") - canonical = importlib.import_module("routes.document.document_helpers") - assert legacy is canonical diff --git a/tests/test_email_summary_error_ui_js.py b/tests/test_email_summary_error_ui_js.py deleted file mode 100644 index 1afc3bec9..000000000 --- a/tests/test_email_summary_error_ui_js.py +++ /dev/null @@ -1,52 +0,0 @@ -import json -import shutil -import subprocess -from pathlib import Path - -import pytest - - -_REPO = Path(__file__).resolve().parent.parent -_UTILS = (_REPO / "static" / "js" / "emailLibrary" / "utils.js").as_posix() -_HAS_NODE = shutil.which("node") is not None - -pytestmark = pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH") - - -def test_email_summary_renderer_ignores_untrusted_provider_error_text(): - secret = ( - "endpoint=https://private.example.internal/v1 provider=ollama " - "model=private-model response_body=private-response " - "Authorization: Bearer token-secret-value" - ) - script = f""" - import {{ _renderEmailSummaryError }} from '{_UTILS}'; - const host = {{ - ownerDocument: {{ - createElement() {{ return {{ style: {{}}, textContent: '' }}; }}, - }}, - replaceChildren(node) {{ this.child = node; }}, - }}; - _renderEmailSummaryError(host, {{ - error_code: 'email_summary_unavailable', - error: {json.dumps(secret)}, - }}); - console.log(JSON.stringify({{ - text: host.child.textContent, - color: host.child.style.color, - }})); - """ - - proc = subprocess.run( - ["node", "--input-type=module"], - input=script, - capture_output=True, - text=True, - cwd=str(_REPO), - timeout=30, - ) - - assert proc.returncode == 0, proc.stderr - rendered = json.loads(proc.stdout) - assert rendered == {"text": "Failed to summarize", "color": "var(--red)"} - assert secret not in proc.stdout diff --git a/tests/test_email_summary_llm.py b/tests/test_email_summary_llm.py deleted file mode 100644 index b0ab7b3be..000000000 --- a/tests/test_email_summary_llm.py +++ /dev/null @@ -1,406 +0,0 @@ -import asyncio -import json -import logging -import os -import sqlite3 -import sys -import tempfile -from pathlib import Path - -import pytest - - -_TMP_DATA = Path(tempfile.mkdtemp(prefix="odysseus-email-summary-")) -os.environ.setdefault("DATA_DIR", str(_TMP_DATA)) -os.environ.setdefault("DATABASE_URL", f"sqlite:///{_TMP_DATA / 'app.db'}") - -PROJECT_ROOT = Path(__file__).resolve().parent.parent -if str(PROJECT_ROOT) not in sys.path: - sys.path.insert(0, str(PROJECT_ROOT)) - - -def _route_endpoint(router, path: str, method: str): - method = method.upper() - for route in router.routes: - if route.path == path and method in getattr(route, "methods", set()): - return route.endpoint - raise AssertionError(f"route not found: {method} {path}") - - -@pytest.mark.asyncio -async def test_generate_email_summary_uses_shared_llm_adapter(monkeypatch): - import routes.email_helpers as email_helpers - import src.llm_core as llm_core - - calls = {} - - async def fake_llm_call_async(url, model, messages, **kwargs): - calls["url"] = url - calls["model"] = model - calls["messages"] = messages - calls["kwargs"] = kwargs - return "thinking before marker\n<<>>\n- Pay the invoice by Friday.\n<<>>" - - monkeypatch.setattr(llm_core, "llm_call_async", fake_llm_call_async) - - summary = await email_helpers._generate_email_summary( - url="https://chatgpt.com/backend-api/codex/responses", - model="gpt-5.5", - sender="Billing ", - subject="Invoice due", - body_for_llm="Please pay invoice 123 by Friday.", - headers={"Authorization": "Bearer test"}, - max_tokens=1234, - timeout=45, - ) - - assert summary == "- Pay the invoice by Friday." - assert calls["url"] == "https://chatgpt.com/backend-api/codex/responses" - assert calls["model"] == "gpt-5.5" - assert calls["kwargs"]["headers"] == {"Authorization": "Bearer test"} - assert calls["kwargs"]["temperature"] == 0.3 - assert calls["kwargs"]["max_tokens"] == 1234 - assert calls["kwargs"]["timeout"] == 45 - assert calls["kwargs"]["workload"] == "foreground" - assert calls["messages"][0]["role"] == "system" - assert calls["messages"][1]["role"] == "user" - - -@pytest.mark.asyncio -async def test_scheduled_email_summary_uses_background_fallback_chain(monkeypatch): - import routes.email_helpers as email_helpers - import src.llm_core as llm_core - import src.task_endpoint as task_endpoint - - candidates = [ - ("http://primary.invalid/v1", "primary-model", {"X-Candidate": "primary"}), - ("http://fallback.invalid/v1", "fallback-model", {"X-Candidate": "fallback"}), - ] - resolve_calls = [] - wait_calls = [] - llm_calls = [] - - def fake_resolve_task_candidates(**kwargs): - resolve_calls.append(kwargs) - return candidates - - async def fake_wait_for_interactive_quiet(label): - wait_calls.append(label) - return False - - async def fake_llm_call_async(url, model, messages, **kwargs): - llm_calls.append((url, model, messages, kwargs)) - if model == "primary-model": - raise RuntimeError("primary unavailable") - return "<<>>\n- Used the fallback model.\n<<>>" - - monkeypatch.setattr(task_endpoint, "resolve_task_candidates", fake_resolve_task_candidates) - monkeypatch.setattr(task_endpoint, "wait_for_interactive_quiet", fake_wait_for_interactive_quiet) - monkeypatch.setattr(llm_core, "llm_call_async", fake_llm_call_async) - - summary = await email_helpers._generate_scheduled_email_summary( - url="http://caller-fallback.invalid/v1", - model="caller-fallback-model", - sender="Sender ", - subject="Scheduled subject", - body_for_llm="Please summarize this scheduled email.", - headers={"Authorization": "Bearer test"}, - owner="alice", - max_tokens=321, - timeout=54, - ) - - assert summary == "- Used the fallback model." - assert resolve_calls == [{ - "fallback_url": "http://caller-fallback.invalid/v1", - "fallback_model": "caller-fallback-model", - "fallback_headers": {"Authorization": "Bearer test"}, - "owner": "alice", - }] - assert wait_calls == ["background task LLM"] - assert [call[1] for call in llm_calls] == ["primary-model", "fallback-model"] - assert all(call[3]["workload"] == "background" for call in llm_calls) - assert all(call[3]["max_tokens"] == 321 for call in llm_calls) - assert all(call[3]["timeout"] == 54 for call in llm_calls) - - -@pytest.mark.asyncio -async def test_scheduled_local_summary_is_preempted_by_foreground_call(monkeypatch): - import routes.email_helpers as email_helpers - import src.llm_core as llm_core - import src.task_endpoint as task_endpoint - - local_url = "http://127.0.0.1:11434/v1/chat/completions" - background_started = asyncio.Event() - never_release = asyncio.Event() - observed_workloads = [] - - monkeypatch.setenv("ODYSSEUS_LOCAL_MODEL_GATE", "true") - monkeypatch.setenv("BACKGROUND_TASK_FOREGROUND_GATE", "false") - monkeypatch.setattr(llm_core, "_LOCAL_MODEL_LOCK", asyncio.Lock()) - monkeypatch.setattr(llm_core, "_LOCAL_MODEL_CURRENT", {}) - monkeypatch.setattr(llm_core, "_LOCAL_MODEL_WAITING_FOREGROUND", 0) - monkeypatch.setattr( - task_endpoint, - "resolve_task_candidates", - lambda **_kwargs: [(local_url, "scheduled-model", {})], - ) - - async def fake_wait_for_interactive_quiet(_label): - return False - - async def gated_llm_call(url, model, messages, **kwargs): - assert messages - workload = kwargs.get("workload") - observed_workloads.append(workload) - async with llm_core._local_model_slot(url, model, workload=workload): - background_started.set() - await never_release.wait() - return "unreachable" - - monkeypatch.setattr(task_endpoint, "wait_for_interactive_quiet", fake_wait_for_interactive_quiet) - monkeypatch.setattr(llm_core, "llm_call_async", gated_llm_call) - - background_task = asyncio.create_task(email_helpers._generate_scheduled_email_summary( - url=local_url, - model="scheduled-model", - sender="Sender", - subject="Scheduled", - body_for_llm="Scheduled body", - owner="alice", - )) - foreground_task = None - try: - await asyncio.wait_for(background_started.wait(), timeout=1) - - async def run_foreground(): - async with llm_core._local_model_slot( - local_url, - "interactive-model", - workload="foreground", - ): - return True - - foreground_task = asyncio.create_task(run_foreground()) - with pytest.raises(asyncio.CancelledError): - await asyncio.wait_for(background_task, timeout=1) - assert await asyncio.wait_for(foreground_task, timeout=1) is True - assert observed_workloads == ["background"] - finally: - for task in (background_task, foreground_task): - if task is not None and not task.done(): - task.cancel() - - -@pytest.mark.asyncio -async def test_manual_email_summary_uses_shared_helper_and_caches(tmp_path, monkeypatch): - import routes.email_helpers as email_helpers - import routes.email_routes as email_routes - import src.endpoint_resolver as endpoint_resolver - - db_path = tmp_path / "scheduled_emails.db" - monkeypatch.setattr(email_helpers, "SCHEDULED_DB", db_path) - monkeypatch.setattr(email_routes, "SCHEDULED_DB", db_path) - email_helpers._init_scheduled_db() - - resolve_calls = [] - - def fake_resolve_endpoint(kind, owner=None): - resolve_calls.append((kind, owner)) - assert kind == "utility" - assert owner == "alice" - return ( - "https://chatgpt.com/backend-api/codex/responses", - "gpt-5.5", - {"Authorization": "Bearer test"}, - ) - - helper_calls = {} - - async def fake_generate_email_summary(**kwargs): - helper_calls.update(kwargs) - return "- Manual summary" - - monkeypatch.setattr(endpoint_resolver, "resolve_endpoint", fake_resolve_endpoint) - monkeypatch.setattr(email_routes, "_generate_email_summary", fake_generate_email_summary) - - router = email_routes.setup_email_routes() - summarize = _route_endpoint(router, "/api/email/summarize", "POST") - - result = await summarize( - { - "body": "This is a long enough email body for manual summary.", - "subject": "Manual subject", - "from": "Sender ", - "message_id": "", - "folder": "INBOX", - }, - owner="alice", - ) - - assert result == { - "success": True, - "summary": "- Manual summary", - "model_used": "gpt-5.5", - } - assert resolve_calls == [("utility", "alice")] - assert helper_calls["url"] == "https://chatgpt.com/backend-api/codex/responses" - assert helper_calls["model"] == "gpt-5.5" - assert helper_calls["headers"]["Authorization"] == "Bearer test" - assert helper_calls["headers"]["Content-Type"] == "application/json" - - conn = sqlite3.connect(db_path) - try: - row = conn.execute( - "SELECT owner, summary, model_used FROM email_summaries WHERE message_id=?", - ("",), - ).fetchone() - finally: - conn.close() - assert row == ("alice", "- Manual summary", "gpt-5.5") - - -@pytest.mark.asyncio -@pytest.mark.parametrize("exception_kind", ["http", "runtime"]) -async def test_manual_email_summary_never_exposes_provider_exception( - monkeypatch, - caplog, - exception_kind, -): - from fastapi import HTTPException - import routes.email_routes as email_routes - import src.endpoint_resolver as endpoint_resolver - - secret_detail = ( - "endpoint=https://private.example.internal/v1 provider=ollama " - "model=private-model response_body=private-response " - "Authorization: Bearer token-secret-value" - ) - - def fake_resolve_endpoint(kind, owner=None): - assert kind == "utility" - assert owner == "alice" - return ( - "https://private.example.internal/v1", - "private-model", - {"Authorization": "Bearer token-secret-value"}, - ) - - async def fail_summary(**_kwargs): - if exception_kind == "http": - raise HTTPException(status_code=502, detail=secret_detail) - raise RuntimeError(secret_detail) - - monkeypatch.setattr(endpoint_resolver, "resolve_endpoint", fake_resolve_endpoint) - monkeypatch.setattr(email_routes, "_generate_email_summary", fail_summary) - caplog.set_level(logging.WARNING, logger=email_routes.__name__) - - router = email_routes.setup_email_routes() - summarize = _route_endpoint(router, "/api/email/summarize", "POST") - result = await summarize( - { - "body": "This email body is long enough to summarize.", - "subject": "Sensitive provider failure", - "from": "Sender ", - }, - owner="alice", - ) - - assert result == { - "success": False, - "error": "Failed to summarize", - "error_code": "email_summary_unavailable", - } - exposed = json.dumps(result) + caplog.text - for marker in ( - "private.example.internal", - "ollama", - "private-model", - "private-response", - "token-secret-value", - ): - assert marker not in exposed - assert f"type={'HTTPException' if exception_kind == 'http' else 'RuntimeError'}" in caplog.text - - -@pytest.mark.asyncio -async def test_scheduled_email_summary_uses_shared_helper_and_caches(tmp_path, monkeypatch): - import routes.email_helpers as email_helpers - import routes.email_pollers as email_pollers - - db_path = tmp_path / "scheduled_emails.db" - monkeypatch.setattr(email_helpers, "SCHEDULED_DB", db_path) - monkeypatch.setattr(email_pollers, "SCHEDULED_DB", db_path) - email_helpers._init_scheduled_db() - - raw_email = ( - b"From: Sender \r\n" - b"To: Alice \r\n" - b"Subject: Scheduled subject\r\n" - b"Message-ID: \r\n" - b"Date: Tue, 01 Jan 2026 12:00:00 +0000\r\n" - b"Content-Type: text/plain; charset=utf-8\r\n" - b"\r\n" - + (b"Please review this scheduled summary email. " * 8) - ) - - class FakeImap: - def __init__(self): - self.logout_calls = 0 - - def select(self, _folder, readonly=True): - return "OK", [] - - def uid(self, command, *args): - if command == "SEARCH": - return "OK", [b"1"] - if command == "FETCH": - return "OK", [(b"1 (RFC822)", raw_email)] - raise AssertionError(f"unexpected uid command: {command!r} {args!r}") - - def logout(self): - self.logout_calls += 1 - - fake_conn = FakeImap() - - def fake_resolve_task_candidates(owner=None): - assert owner == "alice" - return [( - "https://chatgpt.com/backend-api/codex/responses", - "gpt-5.5", - {"Authorization": "Bearer test"}, - )] - - helper_calls = {} - - async def fake_generate_email_summary(**kwargs): - helper_calls.update(kwargs) - return "- Scheduled summary" - - monkeypatch.setattr(email_pollers, "_load_settings", lambda: {"email_auto_summarize": True}) - monkeypatch.setattr(email_pollers, "_owner_for_email_account", lambda _account_id: "alice") - monkeypatch.setattr(email_pollers, "_imap_connect", lambda account_id=None, owner="": fake_conn) - monkeypatch.setattr(email_pollers, "_get_email_config", lambda account_id=None, owner="": {"from_address": "alice@example.com"}) - monkeypatch.setattr(email_pollers, "resolve_task_candidates", fake_resolve_task_candidates) - monkeypatch.setattr(email_pollers, "_generate_scheduled_email_summary", fake_generate_email_summary) - - result = await email_pollers._auto_summarize_pass_single(account_id="acct-alice") - - assert "summarized 1" in result - assert "summary failed" not in result - assert helper_calls["url"] == "https://chatgpt.com/backend-api/codex/responses" - assert helper_calls["model"] == "gpt-5.5" - assert helper_calls["headers"]["Authorization"] == "Bearer test" - assert helper_calls["headers"]["Content-Type"] == "application/json" - assert helper_calls["owner"] == "alice" - assert fake_conn.logout_calls == 1 - - conn = sqlite3.connect(db_path) - try: - row = conn.execute( - "SELECT owner, summary, model_used FROM email_summaries WHERE message_id=?", - ("",), - ).fetchone() - finally: - conn.close() - assert row == ("alice", "- Scheduled summary", "gpt-5.5") diff --git a/tests/test_imap_mailbox_quoting.py b/tests/test_imap_mailbox_quoting.py index 636270a56..7c5bb1645 100644 --- a/tests/test_imap_mailbox_quoting.py +++ b/tests/test_imap_mailbox_quoting.py @@ -87,7 +87,7 @@ def test_known_imap_mailbox_call_sites_are_quoted(): assert "conn.select(sent_name" not in pollers assert "imap.append(sent_folder" not in pollers - document_routes = Path("routes/document/document_routes.py").read_text() + document_routes = Path("routes/document_routes.py").read_text() assert "conn.select(doc.source_email_folder" not in document_routes diff --git a/tests/test_integration_api_call_ssrf.py b/tests/test_integration_api_call_ssrf.py index f23cc40de..53dc671c5 100644 --- a/tests/test_integration_api_call_ssrf.py +++ b/tests/test_integration_api_call_ssrf.py @@ -9,13 +9,8 @@ link-local/metadata is always rejected; RFC-1918/loopback only when INTEGRATION_API_BLOCK_PRIVATE_IPS=true (LAN integrations are the primary use case, so private stays allowed by default). """ -import asyncio -import ipaddress -import ssl from unittest.mock import AsyncMock, MagicMock, patch -import httpcore -import httpx import pytest from src import integrations @@ -102,238 +97,3 @@ async def test_private_base_url_allowed_by_default_blocked_with_knob(monkeypatch assert result["exit_code"] == 1 assert "rejected" in result["error"].lower() client.request.assert_not_called() - - -async def _call_capturing_transport(base_url, path="/items"): - """Drive execute_api_call and return (result, transport) where transport is - the object passed to httpx.AsyncClient(transport=...).""" - resp = MagicMock() - resp.status_code = 200 - resp.headers = {"content-type": "application/json"} - resp.json.return_value = {"ok": True} - resp.text = '{"ok": true}' - - client = AsyncMock() - client.__aenter__ = AsyncMock(return_value=client) - client.__aexit__ = AsyncMock(return_value=None) - client.request = AsyncMock(return_value=resp) - - captured = {} - - def _fake_async_client(*args, **kwargs): - captured.update(kwargs) - return client - - with ( - patch.object(integrations, "_find_integration", - return_value=_integration(base_url)), - patch("httpx.AsyncClient", side_effect=_fake_async_client), - ): - result = await integrations.execute_api_call("test_integ", "GET", path) - return result, captured.get("transport"), client - - -@pytest.mark.asyncio -async def test_connection_is_pinned_to_the_validated_ip(monkeypatch): - """DNS-rebinding defense: the guard resolves the host once to a benign - public IP, and the request must be pinned to *that* IP so a host that - rebinds to the metadata range at connect time can't be reached with the - integration's auth headers. Static resolution passing the guard is not - enough — a plain client would re-resolve at connect.""" - monkeypatch.setattr("src.url_safety._default_resolver", - lambda host: ["93.184.216.34"]) - result, transport, client = await _call_capturing_transport( - "http://rebinding.attacker.example") - - assert result.get("exit_code") == 0 - client.request.assert_called_once() - assert isinstance(transport, integrations._PinnedAsyncTransport) - assert [str(ip) for ip in transport._pinned_ips] == ["93.184.216.34"] - - -@pytest.mark.asyncio -async def test_pin_carries_the_whole_validated_ip_set(monkeypatch): - """When a host resolves to several records the transport keeps all of them - (check_outbound_url validated every one), in resolver order, so it can fall - back past a dead first address instead of failing the whole call.""" - monkeypatch.setattr("src.url_safety._default_resolver", - lambda host: ["93.184.216.34", "198.51.100.7"]) - result, transport, _ = await _call_capturing_transport("http://multi.example") - - assert result.get("exit_code") == 0 - assert [str(ip) for ip in transport._pinned_ips] == ["93.184.216.34", "198.51.100.7"] - - -class _FakeStream: - """Stand-in for the connected socket the real backend returns.""" - - -class _RecordingBackend: - """Fake httpcore backend: connect_tcp fails for the addresses in `dead` - and succeeds for the rest, recording the order it was asked to connect.""" - - def __init__(self, dead): - self.dead = set(dead) - self.attempts = [] - - async def connect_tcp(self, host, port, timeout=None, local_address=None, - socket_options=None): - self.attempts.append((host, timeout)) - if host in self.dead: - raise httpcore.ConnectError(f"connection refused: {host}") - return _FakeStream() - - -def _pinned_backend(ips, dead): - """A _PinnedAsyncBackend whose underlying connect is the recording fake.""" - backend = integrations._PinnedAsyncBackend(ips) - backend._real = _RecordingBackend(dead) - return backend - - -@pytest.mark.asyncio -async def test_connect_falls_back_from_dead_first_to_live_second(): - """first-dead / second-live: the pinned backend must try the next validated - address when the first refuses, rather than surfacing the failure. It also - ignores the `host` httpcore passes (the original hostname) and connects to - the pinned IPs, which is what keeps TLS SNI / Host on the real hostname.""" - ips = [ipaddress.ip_address("203.0.113.10"), ipaddress.ip_address("198.51.100.7")] - backend = _pinned_backend(ips, dead={"203.0.113.10"}) - - stream = await backend.connect_tcp("original.hostname.example", 443, timeout=5.0) - - assert isinstance(stream, _FakeStream) - # Tried the dead address first, then the live one — never the hostname. - assert [host for host, _ in backend._real.attempts] == ["203.0.113.10", "198.51.100.7"] - # Fallback shared one budget: the second attempt got the time left, not a fresh 5s. - assert backend._real.attempts[1][1] <= 5.0 - - -@pytest.mark.asyncio -async def test_connect_raises_when_every_validated_address_is_dead(): - ips = [ipaddress.ip_address("203.0.113.10"), ipaddress.ip_address("198.51.100.7")] - backend = _pinned_backend(ips, dead={"203.0.113.10", "198.51.100.7"}) - - with pytest.raises(httpcore.ConnectError): - await backend.connect_tcp("original.hostname.example", 443, timeout=5.0) - assert [host for host, _ in backend._real.attempts] == ["203.0.113.10", "198.51.100.7"] - - -@pytest.mark.asyncio -async def test_pinned_transport_reuses_httpx_ca_trust(monkeypatch): - """TLS trust must come from the same builder the default httpx client uses - (certifi + SSL_CERT_FILE / SSL_CERT_DIR via trust_env), not from - ssl.create_default_context()'s system roots — otherwise chains that verified - under the old default client can silently stop verifying.""" - sentinel = ssl.create_default_context() - calls = [] - - def _fake_create(*args, **kwargs): - calls.append(kwargs) - return sentinel - - monkeypatch.setattr(httpx, "create_ssl_context", _fake_create) - transport = integrations._PinnedAsyncTransport([ipaddress.ip_address("93.184.216.34")]) - try: - assert calls, "transport did not build its context via httpx.create_ssl_context" - assert transport._pool._ssl_context is sentinel - finally: - await transport.aclose() - - -@pytest.mark.asyncio -async def test_real_socket_falls_back_from_dead_first_to_live_second(): - """End-to-end over real loopback sockets: pin [127.0.0.2 (nothing - listening), 127.0.0.1 (live)], and the request must succeed by falling back - to the second address while the Host header stays the original hostname — - i.e. only the socket destination moved, vhost/SNI routing did not.""" - captured = {} - - async def handle(reader, writer): - request = await reader.read(4096) - for line in request.split(b"\r\n"): - if line.lower().startswith(b"host:"): - captured["host"] = line.split(b":", 1)[1].strip().decode() - writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nhi") - await writer.drain() - writer.close() - - server = await asyncio.start_server(handle, "127.0.0.1", 0) - port = server.sockets[0].getsockname()[1] - async with server: - await server.start_serving() - transport = integrations._PinnedAsyncTransport( - [ipaddress.ip_address("127.0.0.2"), ipaddress.ip_address("127.0.0.1")] - ) - try: - async with httpx.AsyncClient(transport=transport) as client: - resp = await client.get(f"http://pinned.example:{port}/health") - finally: - await transport.aclose() - - assert resp.status_code == 200 - assert resp.text == "hi" - assert captured.get("host") == f"pinned.example:{port}" - - -@pytest.mark.asyncio -async def test_ip_literal_base_url_still_pins_and_is_not_rejected(): - """A base_url that is already an IP has nothing to rebind, but it must not - trip the "did not resolve" guard either. - - check_outbound_url resolves even a literal (getaddrinfo returns the address - itself), so the captured list is populated and the pin is a no-op rather - than a rejection. Uses the real resolver on purpose — no monkeypatch — so - this would catch the fail-closed branch firing on a literal. - """ - result, transport, client = await _call_capturing_transport( - "http://93.184.216.34") - - assert result.get("exit_code") == 0 - assert isinstance(transport, integrations._PinnedAsyncTransport) - assert [str(ip) for ip in transport._pinned_ips] == ["93.184.216.34"] - - -@pytest.mark.asyncio -async def test_ipv6_base_url_pins_every_validated_address(monkeypatch): - """IPv6 goes down the same path as v4. - - Resolution is stubbed rather than using a literal so this doesn't depend on - the runner having IPv6 configured. - """ - v6 = "2606:2800:220:1:248:1893:25c8:1946" - monkeypatch.setattr("src.url_safety._default_resolver", lambda host: [v6]) - result, transport, client = await _call_capturing_transport("http://v6.example") - - assert result.get("exit_code") == 0 - assert isinstance(transport, integrations._PinnedAsyncTransport) - assert [str(ip) for ip in transport._pinned_ips] == [v6] - - -def test_validated_ips_strips_zone_id_and_drops_junk(): - """getaddrinfo can hand back a scoped v6 address like 'fe80::1%eth0'.""" - got = integrations._validated_ips( - ["93.184.216.34", "fe80::1%eth0", "not-an-ip", None, "2001:db8::5"] - ) - assert [str(ip) for ip in got] == ["93.184.216.34", "fe80::1", "2001:db8::5"] - - -def test_validated_ips_deduplicates_repeated_addresses(): - """The resolver is getaddrinfo(host, None) with no socktype filter, so glibc - returns one record per socktype and a single-homed host arrives three times - over. Duplicates must collapse (first-seen order kept) or the connect - fallback wastes its shared deadline retrying one dead address.""" - got = integrations._validated_ips( - ["93.184.216.34", "93.184.216.34", "93.184.216.34"] - ) - assert [str(ip) for ip in got] == ["93.184.216.34"] - - # Order is first-seen, and distinct addresses all survive. - got = integrations._validated_ips( - ["198.51.100.7", "93.184.216.34", "198.51.100.7", "2001:db8::5"] - ) - assert [str(ip) for ip in got] == ["198.51.100.7", "93.184.216.34", "2001:db8::5"] - - # A zone-id variant is the same address once stripped, so it collapses too. - got = integrations._validated_ips(["fe80::1%eth0", "fe80::1%eth1", "fe80::1"]) - assert [str(ip) for ip in got] == ["fe80::1"] diff --git a/tests/test_integrations_api_call_truncation.py b/tests/test_integrations_api_call_truncation.py index a0ad61b4a..bf1ec7d05 100644 --- a/tests/test_integrations_api_call_truncation.py +++ b/tests/test_integrations_api_call_truncation.py @@ -83,10 +83,9 @@ async def _call(json_data, status=200): with ( patch.object(integrations, "_find_integration", return_value=DUMMY_INTEGRATION), patch("httpx.AsyncClient", return_value=mock_client), - # api.example.com doesn't resolve. Point the resolver at a public - # address instead of stubbing the guard open, so the real check (and - # the connect-IP pinning that reads its result) still runs. - patch("src.url_safety._default_resolver", lambda host: ["93.184.216.34"]), + # api.example.com doesn't resolve; the SSRF guard would fail closed. + # These tests are about truncation, so stub the guard open. + patch("src.url_safety.check_outbound_url", return_value=(True, "ok")), ): return await integrations.execute_api_call("test_integ", "GET", "/items") @@ -102,10 +101,9 @@ async def _call_with_integration(integration, path="/items"): with ( patch.object(integrations, "_find_integration", return_value=integration), patch("httpx.AsyncClient", return_value=mock_client), - # api.example.com doesn't resolve. Point the resolver at a public - # address instead of stubbing the guard open, so the real check (and - # the connect-IP pinning that reads its result) still runs. - patch("src.url_safety._default_resolver", lambda host: ["93.184.216.34"]), + # api.example.com doesn't resolve; the SSRF guard would fail closed. + # These tests are about URL joining, so stub the guard open. + patch("src.url_safety.check_outbound_url", return_value=(True, "ok")), ): result = await integrations.execute_api_call("test_integ", "GET", path) return result, mock_client diff --git a/tests/test_manage_skills_action_required.py b/tests/test_manage_skills_action_required.py deleted file mode 100644 index 4efae8026..000000000 --- a/tests/test_manage_skills_action_required.py +++ /dev/null @@ -1,24 +0,0 @@ -import json - -import pytest - -from src.tools.system import do_manage_skills - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "payload", - [ - {}, - {"action": ""}, - {"action": " "}, - {"name": "demo", "description": "x", "procedure": ["step"]}, - ], -) -async def test_manage_skills_requires_action(payload): - result = await do_manage_skills(json.dumps(payload), owner="test") - - assert result == { - "error": "action is required (list|view|view_ref|add|edit|patch|publish|delete|search)", - "exit_code": 1, - } diff --git a/tests/test_memory_add_submit_regression.py b/tests/test_memory_add_submit_regression.py deleted file mode 100644 index 450d63003..000000000 --- a/tests/test_memory_add_submit_regression.py +++ /dev/null @@ -1,54 +0,0 @@ -"""The Brain > Add Memory form must be submittable (#5828). - -The form previously had no submit button and relied on a deprecated -``keypress`` listener for Enter, which is not guaranteed to fire on all -platforms — leaving the form with no working submit path. Pins: - -- a visible, keyboard-accessible submit button next to the category select; -- the button wired to ``memoryModule.addNewMemory()``; -- Enter handled via ``keydown`` with ``preventDefault()`` (and no lingering - ``keypress`` handler on the input). -""" -from pathlib import Path - -APP_JS = Path("static/app.js") -INDEX_HTML = Path("static/index.html") - - -def _add_memory_row(html): - start = html.index('id="new-memory-input"') - end = html.index("
", html.index('id="new-memory-add-btn"', start)) - return html[start:end] - - -def test_add_memory_form_renders_a_submit_button(): - html = INDEX_HTML.read_text() - row = _add_memory_row(html) - - assert 'id="new-memory-category"' in row, "button must sit in the same row as the form fields" - btn_start = row.index('id="new-memory-add-btn"') - btn_tag = row[row.rindex("", btn_start)] - assert 'type="button"' in btn_tag, "must not rely on implicit submit semantics" - - -def _new_memory_wiring_block(source): - start = source.index("const newMemoryInput = el('new-memory-input');") - end = source.index("// Voice recording", start) - return source[start:end] - - -def test_submit_button_is_wired_to_add_new_memory(): - block = _new_memory_wiring_block(APP_JS.read_text()) - - assert "el('new-memory-add-btn')" in block - assert "addEventListener('click', () => memoryModule.addNewMemory())" in block - - -def test_enter_uses_keydown_with_prevent_default(): - block = _new_memory_wiring_block(APP_JS.read_text()) - - assert "addEventListener('keydown'" in block - assert "addEventListener('keypress'" not in block, "keypress is deprecated and unreliable for Enter" - assert "e.preventDefault();" in block - assert "!e.isComposing" in block, "IME composition must not submit the form" - assert "memoryModule.addNewMemory();" in block diff --git a/tests/test_memory_extractor_vector_cross_tenant.py b/tests/test_memory_extractor_vector_cross_tenant.py index 06ca31667..49702c17f 100644 --- a/tests/test_memory_extractor_vector_cross_tenant.py +++ b/tests/test_memory_extractor_vector_cross_tenant.py @@ -67,12 +67,6 @@ class FakeMemoryManager: def load_all(self): return list(self.rows) - def load_all_for_update(self): - # Mirrors the real MemoryManager: extraction is a read-modify-write and - # goes through the strict loader (#5673). A healthy store behaves the - # same as load_all. - return list(self.rows) - def load(self, owner=None): return [r for r in self.rows if r.get("owner") == owner] diff --git a/tests/test_memory_store_unreadable_no_wipe.py b/tests/test_memory_store_unreadable_no_wipe.py deleted file mode 100644 index 4b9076065..000000000 --- a/tests/test_memory_store_unreadable_no_wipe.py +++ /dev/null @@ -1,255 +0,0 @@ -"""A memory store that cannot be READ must never be overwritten (issue #5673). - -`MemoryManager.save` is atomic, and the add/import/extract paths are all -read-modify-write: load the whole store, append, save it back. `load_all` -used to answer a *failed read* with `[]` — indistinguishable from "no -memories" — so a failed read turned into - - load_all() -> [] -> [].append(new) -> save([new]) - -which atomically replaced the entire store with one entry. - -The trigger that actually bites is a store that is **readable but not -parseable** — a truncated file, or one holding `{}` instead of `[]`. Nothing -obstructs the write, so the request succeeds with HTTP 200 and every existing -memory is destroyed silently. Truncation is reachable: `core/database.py` -rewrites memory.json during migration with a plain `open(..., "w")` + -`json.dump`, which is not atomic. - -A live exclusive lock is NOT the dangerous case: it blocks the read and the -`os.replace` alike, so the save fails too and the store survives (verified -end-to-end — clean dev returns 500 there and loses nothing). - -`load_all_for_update` is the strict loader those callers now use: it raises -`MemoryStoreUnreadable` rather than reporting an empty store. -""" - -import asyncio -import builtins -import json -import os - -import pytest - -from src.memory import MemoryManager, MemoryStoreUnreadable - -_SEED = [ - {"id": "m1", "text": "user prefers dark mode", "owner": "alice"}, - {"id": "m2", "text": "user lives in Berlin", "owner": "alice"}, - {"id": "m3", "text": "bob's cat is called Mila", "owner": "bob"}, -] - - -def _seeded(tmp_path): - m = MemoryManager(str(tmp_path)) - m.save([dict(e) for e in _SEED]) - return m - - -def _break_reads_of(monkeypatch, target, exc): - """Make open() raise `exc` for `target` only, leaving every other path alone.""" - real_open = builtins.open - - def fake_open(file, mode="r", *args, **kwargs): - if os.path.abspath(str(file)) == os.path.abspath(target) and "r" in mode: - raise exc - return real_open(file, mode, *args, **kwargs) - - monkeypatch.setattr(builtins, "open", fake_open) - - -# ── the strict loader signals, rather than reporting "empty" ────────────── - -def test_strict_load_raises_on_permission_error(tmp_path, monkeypatch): - m = _seeded(tmp_path) - _break_reads_of(monkeypatch, m.memory_file, PermissionError(13, "locked")) - with pytest.raises(MemoryStoreUnreadable): - m.load_all_for_update() - - -def test_strict_load_raises_on_corrupt_json(tmp_path): - m = _seeded(tmp_path) - with open(m.memory_file, "w", encoding="utf-8") as f: - f.write('[{"id": "m1", "text": "truncated mid-writ') - with pytest.raises(MemoryStoreUnreadable): - m.load_all_for_update() - - -def test_strict_load_raises_when_store_is_not_a_list(tmp_path): - # A file holding `{}` or `null` is not an empty store, it is a broken one. - m = _seeded(tmp_path) - with open(m.memory_file, "w", encoding="utf-8") as f: - json.dump({}, f) - with pytest.raises(MemoryStoreUnreadable): - m.load_all_for_update() - - -def test_strict_load_returns_entries_when_healthy(tmp_path): - m = _seeded(tmp_path) - assert {e["id"] for e in m.load_all_for_update()} == {"m1", "m2", "m3"} - - -def test_strict_load_returns_empty_when_file_genuinely_absent(tmp_path): - m = _seeded(tmp_path) - os.remove(m.memory_file) - # Absent is the one case that legitimately means "no memories yet". - assert m.load_all_for_update() == [] - - -# ── read paths stay lenient, so an unreadable store can't break chat ────── - -def test_read_path_still_degrades_to_empty(tmp_path, monkeypatch): - m = _seeded(tmp_path) - _break_reads_of(monkeypatch, m.memory_file, PermissionError(13, "locked")) - # Context injection / search must not raise; they just see nothing. - assert m.load_all() == [] - assert m.load(owner="alice") == [] - - -# ── the actual #5673 regression: the store survives ─────────────────────── - -def test_add_cycle_under_transient_read_error_does_not_wipe(tmp_path, monkeypatch): - """Mirrors routes/memory/memory_routes.py api_add_memory exactly.""" - m = _seeded(tmp_path) - new_entry = m.add_entry("a brand new fact", owner="alice") - - with monkeypatch.context() as mp: - _break_reads_of(mp, m.memory_file, PermissionError(13, "locked")) - with pytest.raises(MemoryStoreUnreadable): - all_mem = m.load_all_for_update() - all_mem.append(new_entry) - m.save(all_mem) - - # Reads work again; every original memory is still there and the file was - # never replaced by the single new entry. - assert {e["id"] for e in m.load_all()} == {"m1", "m2", "m3"} - - -def test_audit_merge_cannot_drop_other_tenants(tmp_path, monkeypatch): - """The audit path rebuilds the whole file from load_all + one owner's slice. - - Reading [] there would save only the audited owner's entries and destroy - every other tenant's memories, so it has to fail closed too. - """ - m = _seeded(tmp_path) - alice_slice = [e for e in _SEED if e["owner"] == "alice"] - - with monkeypatch.context() as mp: - _break_reads_of(mp, m.memory_file, PermissionError(13, "locked")) - with pytest.raises(MemoryStoreUnreadable): - all_entries = m.load_all_for_update() - others = [e for e in all_entries if e.get("owner") != "alice"] - m.save(alice_slice + others) - - assert any(e["id"] == "m3" for e in m.load_all()), "bob's memory was destroyed" - - -def test_uses_bump_skips_write_when_unreadable(tmp_path, monkeypatch): - m = _seeded(tmp_path) - with monkeypatch.context() as mp: - _break_reads_of(mp, m.memory_file, PermissionError(13, "locked")) - m.increment_uses(["m1"]) # must not raise, must not write - assert {e["id"] for e in m.load_all()} == {"m1", "m2", "m3"} - - -def test_claim_ownerless_skips_write_when_unreadable(tmp_path, monkeypatch): - m = _seeded(tmp_path) - with monkeypatch.context() as mp: - _break_reads_of(mp, m.memory_file, PermissionError(13, "locked")) - m.claim_ownerless("alice") - assert {e["id"] for e in m.load_all()} == {"m1", "m2", "m3"} - - -# ── the add sinks users actually reach ──────────────────────────────────── -# -# The tests above replay the read-modify-write shape. These drive the real -# entry points end to end, because those are what #5673 reports: "remember -# that I prefer X" in ordinary chat (src/ai_interaction.py do_manage_memory, -# routed from src/tool_execution.py) and the built-in memory MCP server -# (mcp_servers/memory_server.py, registered in src/builtin_mcp.py). -# -# They use a truncated store rather than a read error on purpose: it reads -# fine, so nothing stops the save, which is the case that silently destroyed -# stores. The assertion is that the file is left byte-identical — still broken, -# but still holding the user's memories, so it can be repaired by hand. - - -def _truncated_store(tmp_path): - """Seed a store that reads back fine but no longer parses.""" - m = _seeded(tmp_path) - good = json.dumps([dict(e) for e in _SEED], indent=2) - with open(m.memory_file, "w", encoding="utf-8") as f: - f.write(good[:good.rindex("]")]) # drop the closing bracket only - with open(m.memory_file, "rb") as f: - return m, f.read() - - -def _on_disk(manager) -> bytes: - with open(manager.memory_file, "rb") as f: - return f.read() - - -def test_agent_memory_add_does_not_overwrite_unreadable_store(tmp_path, monkeypatch): - """src/ai_interaction.py do_manage_memory, action "add".""" - from src import ai_interaction - - manager, before = _truncated_store(tmp_path) - monkeypatch.setattr(ai_interaction, "_memory_manager", manager) - monkeypatch.setattr(ai_interaction, "_memory_vector", None) - - result = asyncio.run(ai_interaction.do_manage_memory("add\nuser prefers tabs")) - - assert _on_disk(manager) == before, "the unreadable store was overwritten" - assert b"m3" in _on_disk(manager) - assert "error" in result, "the add reported success over an unreadable store" - - -def test_mcp_memory_add_does_not_overwrite_unreadable_store(tmp_path, monkeypatch): - """mcp_servers/memory_server.py, action "add".""" - import mcp_servers.memory_server as memory_server - - manager, before = _truncated_store(tmp_path) - monkeypatch.setattr(memory_server, "_memory_manager", manager) - monkeypatch.setattr(memory_server, "_memory_vector", None) - monkeypatch.setattr(memory_server, "_initialized", True) - for key in memory_server._OWNER_ENV_KEYS: - monkeypatch.delenv(key, raising=False) - - result = asyncio.run(memory_server.call_tool( - "manage_memory", {"action": "add", "text": "user prefers tabs"} - )) - - assert _on_disk(manager) == before, "the unreadable store was overwritten" - assert b"m3" in _on_disk(manager) - assert result[0].text.startswith("Error:") - - -def test_native_provider_remember_does_not_overwrite_unreadable_store(tmp_path): - """src/memory_provider.py NativeMemoryProvider.remember. - - Registered into app state in src/app_initializer.py but not yet consumed - outside tests, so this is the pattern held in place before it goes live. - """ - from src.memory_provider import NativeMemoryProvider - - manager, before = _truncated_store(tmp_path) - provider = NativeMemoryProvider(manager) - - with pytest.raises(MemoryStoreUnreadable): - asyncio.run(provider.remember("user prefers tabs", owner="alice")) - - assert _on_disk(manager) == before - - -# ── the legacy memory.txt migration is preserved ────────────────────────── - -def test_corrupt_store_still_migrates_from_legacy_txt(tmp_path): - m = _seeded(tmp_path) - with open(m.memory_file, "w", encoding="utf-8") as f: - f.write("{ not json") - legacy = os.path.join(str(tmp_path), "memory.txt") - with open(legacy, "w", encoding="utf-8") as f: - f.write("recovered fact one\nrecovered fact two\n") - - entries = m.load_all_for_update() - assert [e["text"] for e in entries] == ["recovered fact one", "recovered fact two"] diff --git a/tests/test_model_helper_owner_scope.py b/tests/test_model_helper_owner_scope.py index f48a1f7e2..dafbad594 100644 --- a/tests/test_model_helper_owner_scope.py +++ b/tests/test_model_helper_owner_scope.py @@ -14,7 +14,7 @@ def _function_source(path: str, name: str) -> str: def test_document_ai_tidy_resolves_with_owner_scope(): - body = _function_source("routes/document/document_routes.py", "ai_tidy_documents") + body = _function_source("routes/document_routes.py", "ai_tidy_documents") assert "resolve_task_endpoint(owner=user or None)" in body assert 'resolve_endpoint("default", owner=user or None)' in body diff --git a/tests/test_tool_parsing_bare_end_marker.py b/tests/test_tool_parsing_bare_end_marker.py deleted file mode 100644 index 6167c8dde..000000000 --- a/tests/test_tool_parsing_bare_end_marker.py +++ /dev/null @@ -1,96 +0,0 @@ -"""Regression: the Qwen bare-marker scrub must not eat a lone `end` (#5547). - -`_QWEN_BARE_MARKER_RE` cleans Qwen turn markers that leak into content. Its -`end` branch was `\\|?end\\|?` — both pipes optional — so it also matched a bare -`end` surrounded by whitespace and replaced it with a space. Any message -containing Ruby, Lua or shell code that closes a block with a lone `end` had -those lines silently deleted, in the stored text and in the rendered message. - -Requiring at least one pipe keeps every real marker (`|end`, `end|`, `|end|`, -`/|end|`) stripping as before. The same pattern is duplicated in -static/js/chatRenderer.js, so the JS copy is checked here too — the two must -not drift. -""" -import json -import re -import shutil -import subprocess -from pathlib import Path - -import pytest - -import src.agent_tools # noqa: F401 (break agent_tools<->tool_parsing import cycle) -from src.tool_parsing import strip_tool_blocks - -_REPO = Path(__file__).resolve().parent.parent -_CHAT_RENDERER = _REPO / "static" / "js" / "chatRenderer.js" - -# Inputs that must survive untouched, and the substring that proves they did. -KEPT = [ - ("loop do\n puts \"yo\"\nend\n", "\nend"), # the reported Ruby case - ("if x then\nend", "\nend"), - ("function f()\nend\n", "\nend"), - ("a end b", "a end b"), - ("append end", "append end"), - ("END", "END"), - ("\nEnd\n", "End"), -] - -# Real markers — at least one pipe, plus the role word — with the exact output -# they must still produce. Asserted as equality rather than "marker not in out" -# so narrowing the pattern can't pass by deleting more than it should. -STRIPPED = [ - ("a |end| b", "a b"), - ("a /|end| b", "a b"), - ("a |end b", "a b"), - ("a end| b", "a b"), - ("x assistant y", "x y"), -] - - -@pytest.mark.parametrize("text,kept", KEPT) -def test_bare_end_survives_stripping(text, kept): - assert kept in strip_tool_blocks(text) - - -@pytest.mark.parametrize("text,expected", STRIPPED) -def test_piped_end_markers_are_still_stripped(text, expected): - assert strip_tool_blocks(text) == expected - - -def test_bare_end_inside_a_fenced_block_survives(): - """The scrub runs over the whole message, fenced regions included.""" - out = strip_tool_blocks("Here:\n```ruby\nloop do\n puts 1\nend\n```\nDone.") - assert "\nend\n" in out - - -def _js_bare_marker_regex_source(): - src = _CHAT_RENDERER.read_text(encoding="utf-8") - m = re.search(r"^const QWEN_BARE_MARKER_RE = (/.*/[gimsuy]*);$", src, re.MULTILINE) - assert m, "QWEN_BARE_MARKER_RE literal not found in chatRenderer.js" - return m.group(1) - - -def test_js_copy_of_the_pattern_matches_the_python_one(): - """Guard the duplication: the JS branch must require a pipe too.""" - if shutil.which("node") is None: - pytest.skip("node binary not on PATH") - - cases = [text for text, _ in KEPT] + [text for text, _ in STRIPPED] - script = ( - "const RE = %s;\n" - "const cases = JSON.parse(process.argv[1]);\n" - "console.log(JSON.stringify(cases.map(c => c.replace(RE, ' '))));" - % _js_bare_marker_regex_source() - ) - result = subprocess.run( - ["node", "--input-type=module", "-e", script, json.dumps(cases)], - cwd=_REPO, capture_output=True, timeout=15, text=True, - ) - assert result.returncode == 0, f"node failed:\n{result.stderr}" - got = json.loads(result.stdout.splitlines()[-1]) - - for (text, kept), out in zip(KEPT, got): - assert kept in out, f"JS regex dropped {kept!r} from {text!r}" - for (text, expected), out in zip(STRIPPED, got[len(KEPT):]): - assert out == expected, f"JS regex: {text!r} -> {out!r}, expected {expected!r}" diff --git a/tests/test_vault_routes_shim.py b/tests/test_vault_routes_shim.py deleted file mode 100644 index 9577395f7..000000000 --- a/tests/test_vault_routes_shim.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Regression test for the vault route shim (slice 2k, #4082/#4071).""" - -import importlib - -import routes.vault_routes as _shim_vault # noqa: F401 - - -def test_legacy_and_canonical_vault_module_are_same_object(): - legacy = importlib.import_module("routes.vault_routes") - canonical = importlib.import_module("routes.vault.vault_routes") - assert legacy is canonical diff --git a/tests/test_vision_owner_scope.py b/tests/test_vision_owner_scope.py index 29de101a3..f0d3a184d 100644 --- a/tests/test_vision_owner_scope.py +++ b/tests/test_vision_owner_scope.py @@ -88,7 +88,7 @@ def test_request_vision_call_sites_pass_owner(): chat_source = (ROOT / "src" / "chat_handler.py").read_text() processor_source = (ROOT / "src" / "document_processor.py").read_text() upload_source = (ROOT / "routes" / "upload_routes.py").read_text() - document_source = (ROOT / "routes" / "document" / "document_routes.py").read_text() + document_source = (ROOT / "routes" / "document_routes.py").read_text() gallery_source = (ROOT / "routes" / "gallery" / "gallery_routes.py").read_text() memory_source = (ROOT / "routes" / "memory" / "memory_routes.py").read_text() diff --git a/tests/test_webhook_routes_shim.py b/tests/test_webhook_routes_shim.py deleted file mode 100644 index f6312e8e6..000000000 --- a/tests/test_webhook_routes_shim.py +++ /dev/null @@ -1,11 +0,0 @@ -"""Regression test for the webhook route shim (slice 2l, #4082/#4071).""" - -import importlib - -import routes.webhook_routes as _shim_webhook # noqa: F401 - - -def test_legacy_and_canonical_webhook_module_are_same_object(): - legacy = importlib.import_module("routes.webhook_routes") - canonical = importlib.import_module("routes.webhook.webhook_routes") - assert legacy is canonical