diff --git a/.env.example b/.env.example index d23276eb8..2d1be3373 100644 --- a/.env.example +++ b/.env.example @@ -189,6 +189,7 @@ SEARXNG_INSTANCE=http://localhost:8080 # ODYSSEUS_EMAIL_COMPOSE_UPLOAD_MAX_BYTES=26214400 # email compose attachment (25 MB) # ODYSSEUS_STT_MAX_AUDIO_BYTES=26214400 # speech-to-text audio (25 MB) # ODYSSEUS_ICS_MAX_BYTES=10485760 # calendar .ics import (10 MB) +# ODYSSEUS_TTS_CACHE_MAX_BYTES=524288000 # TTS cache (500 MB) # ============================================================ # Host Docker access (explicit opt-in) diff --git a/.github/scripts/check-issue-description.js b/.github/scripts/check-issue-description.js index a76ca29ab..63162b0d7 100644 --- a/.github/scripts/check-issue-description.js +++ b/.github/scripts/check-issue-description.js @@ -153,6 +153,16 @@ module.exports = async ({ github, context, core }) => { } } + const LABEL_BAD = 'needs more info'; + const LABEL_GOOD = 'ready for review'; + + // Closed issues are no longer awaiting review. + // This also prevents later edits to closed issues from restoring the label. + if (issue.state === 'closed') { + await dropLabel(LABEL_GOOD); + return; + } + // ── Find existing bot comment to update in-place ────────────────────────── const MARKER = ''; const { data: comments } = await github.rest.issues.listComments({ @@ -160,9 +170,6 @@ module.exports = async ({ github, context, core }) => { }); const existing = comments.find(c => c.user.type === 'Bot' && c.body.includes(MARKER)); - const LABEL_BAD = 'needs more info'; - const LABEL_GOOD = 'ready for review'; - if (failures.length === 0) { if (existing) { await github.rest.issues.deleteComment({ owner, repo, comment_id: existing.id }); diff --git a/.github/workflows/issue-description-check.yml b/.github/workflows/issue-description-check.yml index 52e9dddae..5ce6037f0 100644 --- a/.github/workflows/issue-description-check.yml +++ b/.github/workflows/issue-description-check.yml @@ -2,7 +2,7 @@ name: ci / issue description check on: issues: - types: [opened, edited, reopened] + types: [opened, edited, reopened, closed] permissions: issues: write diff --git a/app.py b/app.py index e740ad518..8363ba4e9 100644 --- a/app.py +++ b/app.py @@ -692,7 +692,7 @@ from routes.history.history_routes import setup_history_routes app.include_router(setup_history_routes(session_manager, upload_handler=upload_handler)) # Search -from routes.search_routes import setup_search_routes +from routes.search.search_routes import setup_search_routes app.include_router(setup_search_routes(config)) # Presets @@ -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_routes import setup_document_routes +from routes.document.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_routes import setup_webhook_routes +from routes.webhook.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_routes import setup_vault_routes +from routes.vault.vault_routes import setup_vault_routes app.include_router(setup_vault_routes()) # Contacts (CardDAV) diff --git a/docker-compose.gpu-amd.yml b/docker-compose.gpu-amd.yml index 91e223e05..9699fc038 100644 --- a/docker-compose.gpu-amd.yml +++ b/docker-compose.gpu-amd.yml @@ -67,6 +67,7 @@ services: - ODYSSEUS_EMAIL_COMPOSE_UPLOAD_MAX_BYTES=${ODYSSEUS_EMAIL_COMPOSE_UPLOAD_MAX_BYTES:-26214400} - ODYSSEUS_STT_MAX_AUDIO_BYTES=${ODYSSEUS_STT_MAX_AUDIO_BYTES:-26214400} - ODYSSEUS_ICS_MAX_BYTES=${ODYSSEUS_ICS_MAX_BYTES:-10485760} + - ODYSSEUS_TTS_CACHE_MAX_BYTES=${ODYSSEUS_TTS_CACHE_MAX_BYTES} - DATA_BRAVE_API_KEY=${DATA_BRAVE_API_KEY:-} - GOOGLE_API_KEY=${GOOGLE_API_KEY:-} - GOOGLE_PSE_CX=${GOOGLE_PSE_CX:-} diff --git a/docker-compose.gpu-nvidia.yml b/docker-compose.gpu-nvidia.yml index e8c2fd032..804a0a14e 100644 --- a/docker-compose.gpu-nvidia.yml +++ b/docker-compose.gpu-nvidia.yml @@ -66,6 +66,7 @@ services: - ODYSSEUS_EMAIL_COMPOSE_UPLOAD_MAX_BYTES=${ODYSSEUS_EMAIL_COMPOSE_UPLOAD_MAX_BYTES:-26214400} - ODYSSEUS_STT_MAX_AUDIO_BYTES=${ODYSSEUS_STT_MAX_AUDIO_BYTES:-26214400} - ODYSSEUS_ICS_MAX_BYTES=${ODYSSEUS_ICS_MAX_BYTES:-10485760} + - ODYSSEUS_TTS_CACHE_MAX_BYTES=${ODYSSEUS_TTS_CACHE_MAX_BYTES} - DATA_BRAVE_API_KEY=${DATA_BRAVE_API_KEY:-} - GOOGLE_API_KEY=${GOOGLE_API_KEY:-} - GOOGLE_PSE_CX=${GOOGLE_PSE_CX:-} diff --git a/docker-compose.yml b/docker-compose.yml index b1f2c37ee..b0efb4439 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -55,6 +55,7 @@ services: - ODYSSEUS_EMAIL_COMPOSE_UPLOAD_MAX_BYTES=${ODYSSEUS_EMAIL_COMPOSE_UPLOAD_MAX_BYTES:-26214400} - ODYSSEUS_STT_MAX_AUDIO_BYTES=${ODYSSEUS_STT_MAX_AUDIO_BYTES:-26214400} - ODYSSEUS_ICS_MAX_BYTES=${ODYSSEUS_ICS_MAX_BYTES:-10485760} + - ODYSSEUS_TTS_CACHE_MAX_BYTES=${ODYSSEUS_TTS_CACHE_MAX_BYTES} - DATA_BRAVE_API_KEY=${DATA_BRAVE_API_KEY:-} - GOOGLE_API_KEY=${GOOGLE_API_KEY:-} - GOOGLE_PSE_CX=${GOOGLE_PSE_CX:-} diff --git a/mcp_servers/memory_server.py b/mcp_servers/memory_server.py index fafbcfc2b..fd574fd1f 100644 --- a/mcp_servers/memory_server.py +++ b/mcp_servers/memory_server.py @@ -17,6 +17,8 @@ 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) @@ -29,6 +31,10 @@ _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: @@ -51,9 +57,21 @@ def _owner_scoped_store(entries: list[dict]) -> bool: return any(_entry_owner(entry) for entry in entries if isinstance(entry, dict)) -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() +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() owner = _configured_owner() if owner is None and _owner_scoped_store(entries): return None, entries, [], _OWNER_SCOPE_ERROR @@ -161,7 +179,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() + owner, memories, _visible, scope_error = _scope_entries(for_update=True) if scope_error: return _text_result(scope_error) entry = _memory_manager.add_entry(text, source="ai_agent", category=category, owner=owner) diff --git a/requirements.txt b/requirements.txt index be5f5d450..3c5114f53 100644 --- a/requirements.txt +++ b/requirements.txt @@ -38,7 +38,10 @@ python-dateutil caldav cryptography bcrypt -mcp +# Built-in servers use the v1 low-level Server decorator API. MCP SDK v2 is a +# breaking rewrite, so keep fresh installs on the maintained v1 line until the +# servers are migrated together. +mcp<2 pyotp qrcode[pil] croniter diff --git a/routes/backup_routes.py b/routes/backup_routes.py index 313369370..4ecf4f165 100644 --- a/routes/backup_routes.py +++ b/routes/backup_routes.py @@ -6,6 +6,7 @@ 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 @@ -76,7 +77,15 @@ def setup_backup_routes(memory_manager, preset_manager, skills_manager) -> APIRo # ── Memories ── if "memories" in body and isinstance(body["memories"], list): - existing = memory_manager.load_all() + # 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." + ) # 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 new file mode 100644 index 000000000..7f79ce1bb --- /dev/null +++ b/routes/document/__init__.py @@ -0,0 +1,6 @@ +"""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 new file mode 100644 index 000000000..a0c2d08eb --- /dev/null +++ b/routes/document/document_helpers.py @@ -0,0 +1,243 @@ +"""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 new file mode 100644 index 000000000..dae8b09fa --- /dev/null +++ b/routes/document/document_routes.py @@ -0,0 +1,1810 @@ +"""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 a0c2d08eb..c1f68ca51 100644 --- a/routes/document_helpers.py +++ b/routes/document_helpers.py @@ -1,243 +1,14 @@ -"""document_helpers.py — Pydantic models, doc serializers, owner gating, file-locator helpers shared with document_routes.py.""" +"""Backward-compat shim — canonical location is routes/document/document_helpers.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_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). +""" -import logging -import os -import re -from typing import Any, Dict, Optional +import sys as _sys -from fastapi import HTTPException, Request -from pydantic import BaseModel +from routes.document import document_helpers as _canonical # noqa: F401 -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" +_sys.modules[__name__] = _canonical diff --git a/routes/document_routes.py b/routes/document_routes.py index dae8b09fa..dd13e3c60 100644 --- a/routes/document_routes.py +++ b/routes/document_routes.py @@ -1,1810 +1,17 @@ -"""Document routes — CRUD for living documents with version history.""" +"""Backward-compat shim — canonical location is routes/document/document_routes.py. -import uuid -import logging -from datetime import datetime, timezone -from typing import Dict, Any, List, Optional +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. +""" -from fastapi import APIRouter, HTTPException, Query, Request, UploadFile, File, Form +import sys as _sys -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 +from routes.document import document_routes as _canonical # noqa: F401 -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 +_sys.modules[__name__] = _canonical diff --git a/routes/email_helpers.py b/routes/email_helpers.py index c8639e1c7..257f5f921 100644 --- a/routes/email_helpers.py +++ b/routes/email_helpers.py @@ -247,6 +247,7 @@ 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: @@ -277,6 +278,125 @@ 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 5d96bd0f9..a2507989d 100644 --- a/routes/email_pollers.py +++ b/routes/email_pollers.py @@ -40,6 +40,7 @@ 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__) @@ -653,6 +654,7 @@ 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 @@ -785,16 +787,17 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None if need_sum: try: - 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, + summary = await _generate_scheduled_email_summary( + url=url, + model=model, + sender=sender, + subject=subject, + body_for_llm=body_for_llm, + headers=req_headers, owner=account_owner or None, - temperature=0.3, max_tokens=16384, timeout=240, + max_tokens=16384, + timeout=240, ) - summary = _extract_reply((summary or "").strip()) if summary: _c = _sql3.connect(SCHEDULED_DB) _c.execute(""" @@ -808,10 +811,19 @@ 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(f"Auto-summary {uid} failed: {e}") + logger.warning( + "Auto-summary uid=%s failed %s", + _uid_text, + _email_summary_failure_log_detail(e), + ) if need_reply: await _emit_progress(progress_cb, f"Drafting reply {processed + 1}/{_max_process} · checked {examined}/{len(uid_list)}") @@ -1320,6 +1332,8 @@ 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 3c8e407bd..76a744ce1 100644 --- a/routes/email_routes.py +++ b/routes/email_routes.py @@ -57,7 +57,8 @@ 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, + _friendly_email_auth_error, _email_summary_failure_log_detail, + _generate_email_summary, EMAIL_SUMMARY_ERROR_CODE, EMAIL_SUMMARY_ERROR_MESSAGE, SendEmailRequest, ExtractStyleRequest, ATTACHMENTS_DIR, COMPOSE_UPLOADS_DIR, SCHEDULED_DB, attachment_extract_dir, _email_cache_owner_clause, email_translation_body_hash, @@ -4766,8 +4767,6 @@ 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", "") @@ -4778,7 +4777,11 @@ def setup_email_routes(): if account_id: _assert_owns_account(account_id, owner) if not body: - return {"success": False, "error": "No body provided"} + return { + "success": False, + "error": "No body provided", + "error_code": "email_summary_missing_body", + } # If we know which UID this is, fetch the raw message and pull # attachment text so the summary can reference invoice totals, @@ -4807,53 +4810,43 @@ 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 LLM endpoint configured"} + return { + "success": False, + "error": "No model configured for email summaries", + "error_code": "email_summary_not_configured", + } req_headers = {"Content-Type": "application/json"} if headers: req_headers.update(headers) - 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) + 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, + } if not content: - # 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"} + return { + "success": False, + "error": "The model returned an empty summary", + "error_code": "email_summary_empty", + } # Cache the summary if we have a message_id mid = data.get("message_id", "") @@ -4876,8 +4869,15 @@ def setup_email_routes(): return {"success": True, "summary": content, "model_used": model} except Exception as e: - logger.error(f"Failed to summarize: {e}") - return {"success": False, "error": "Mail operation failed"} + 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, + } @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 d290046ec..c4232bec4 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 +from services.memory import MemoryManager, MemoryStoreUnreadable from core.session_manager import SessionManager from src.request_models import MemoryAddRequest from core.database import SessionLocal @@ -35,6 +35,22 @@ 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"]) @@ -116,7 +132,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 = memory_manager.load_all() + all_mem = _load_for_update(memory_manager) all_mem.append(new_entry) memory_manager.save(all_mem) # Sync vector index @@ -487,7 +503,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 = memory_manager.load_all() + all_mem = _load_for_update(memory_manager) for i, memory in enumerate(all_mem): if memory["id"] == memory_id: _verify_memory_owner(memory, user) @@ -512,7 +528,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 = memory_manager.load_all() + all_mem = _load_for_update(memory_manager) for i, memory in enumerate(all_mem): if memory["id"] == memory_id: _verify_memory_owner(memory, user) @@ -534,7 +550,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 = memory_manager.load_all() + all_mem = _load_for_update(memory_manager) # Find and verify ownership before deleting target = next((m for m in all_mem if m["id"] == memory_id), None) diff --git a/routes/search/__init__.py b/routes/search/__init__.py new file mode 100644 index 000000000..ea051bbe0 --- /dev/null +++ b/routes/search/__init__.py @@ -0,0 +1,5 @@ +"""Search route domain package (slice 2j, #4082/#4071). + +Contains search_routes.py, migrated from the flat routes/ directory. +Backward-compat shim at routes/search_routes.py re-exports from here. +""" diff --git a/routes/search/search_routes.py b/routes/search/search_routes.py new file mode 100644 index 000000000..1effb7b8f --- /dev/null +++ b/routes/search/search_routes.py @@ -0,0 +1,111 @@ +"""Search routes — /api/search/config GET, /api/search POST.""" + +import logging +from typing import Dict, Any + +from fastapi import APIRouter, Request + +import time + +from services.search import get_search_config, comprehensive_web_search, PROVIDER_INFO +from services.search.core import _call_provider +from services.search.providers import _get_provider_key, _get_search_instance + +logger = logging.getLogger(__name__) + + +async def _request_values(request: Request) -> Dict[str, Any]: + """Accept JSON, form data, or query params for search endpoints. + + The browser UI posts FormData, while the agent's generic app_api tool + posts JSON. FastAPI Form(...) rejects JSON with a 422 before our handler + runs, which made the model think SearXNG was broken. + """ + values: Dict[str, Any] = dict(request.query_params) + content_type = (request.headers.get("content-type") or "").lower() + try: + if "application/json" in content_type: + body = await request.json() + if isinstance(body, dict): + values.update(body) + else: + form = await request.form() + values.update(dict(form)) + except Exception: + pass + return values + + +def setup_search_routes(config) -> APIRouter: + router = APIRouter(tags=["search"]) + + @router.get("/api/search/config") + async def get_search_settings() -> Dict[str, Any]: + return get_search_config() + + @router.post("/api/search") + async def do_web_search(request: Request) -> Dict[str, Any]: + """Standalone web search — returns context string + source list. + + Used by Compare mode to pre-search once and share results across panes. + """ + values = await _request_values(request) + query = str(values.get("query") or values.get("q") or "").strip() + if not query: + return {"context": "", "sources": [], "error": "query is required"} + time_filter = values.get("time_filter") or values.get("freshness") + if time_filter is not None: + time_filter = str(time_filter).strip() or None + try: + context, sources = comprehensive_web_search( + query, return_sources=True, time_filter=time_filter, + ) + return {"context": context, "sources": sources} + except Exception as e: + logger.error(f"Standalone web search failed: {e}") + return {"context": "", "sources": [], "error": str(e)} + + @router.get("/api/search/providers") + async def list_search_providers(): + """Return available search providers with config status.""" + providers = [] + for pid, (label, needs_key, needs_url) in PROVIDER_INFO.items(): + if pid == "disabled": + continue + available = True + if needs_key and not _get_provider_key(pid): + available = False + if needs_url and pid == "searxng" and not _get_search_instance(): + available = False + providers.append({ + "id": pid, + "label": label, + "available": available, + }) + return providers + + @router.post("/api/search/query") + async def search_with_provider(request: Request) -> Dict[str, Any]: + """Search using a specific provider. Used by compare search mode.""" + values = await _request_values(request) + query = str(values.get("query") or values.get("q") or "").strip() + provider = str(values.get("provider") or "").strip() + try: + count = int(values.get("count") or values.get("limit") or 10) + except Exception: + count = 10 + if not query: + return {"results": [], "provider": provider, "error": "query is required"} + if provider not in PROVIDER_INFO or provider == "disabled": + return {"results": [], "provider": provider, "error": "Unknown provider"} + t0 = time.time() + try: + results = _call_provider(provider, query, min(count, 20)) + elapsed = round(time.time() - t0, 2) + return {"results": results, "provider": provider, "time": elapsed} + except Exception as e: + elapsed = round(time.time() - t0, 2) + logger.error(f"Search provider {provider} failed: {e}") + return {"results": [], "provider": provider, "time": elapsed, "error": str(e)} + + return router diff --git a/routes/search_routes.py b/routes/search_routes.py index 1effb7b8f..03b94438b 100644 --- a/routes/search_routes.py +++ b/routes/search_routes.py @@ -1,111 +1,13 @@ -"""Search routes — /api/search/config GET, /api/search POST.""" +"""Backward-compat shim — canonical location is routes/search/search_routes.py. -import logging -from typing import Dict, Any +This module is replaced in ``sys.modules`` by the canonical module object so +that ``import routes.search_routes`` and ``from routes.search_routes import X`` +keep resolving to the canonical module. Keeps existing import paths working +after slice 2j (#4082/#4071). +""" -from fastapi import APIRouter, Request +import sys as _sys -import time +from routes.search import search_routes as _canonical # noqa: F401 -from services.search import get_search_config, comprehensive_web_search, PROVIDER_INFO -from services.search.core import _call_provider -from services.search.providers import _get_provider_key, _get_search_instance - -logger = logging.getLogger(__name__) - - -async def _request_values(request: Request) -> Dict[str, Any]: - """Accept JSON, form data, or query params for search endpoints. - - The browser UI posts FormData, while the agent's generic app_api tool - posts JSON. FastAPI Form(...) rejects JSON with a 422 before our handler - runs, which made the model think SearXNG was broken. - """ - values: Dict[str, Any] = dict(request.query_params) - content_type = (request.headers.get("content-type") or "").lower() - try: - if "application/json" in content_type: - body = await request.json() - if isinstance(body, dict): - values.update(body) - else: - form = await request.form() - values.update(dict(form)) - except Exception: - pass - return values - - -def setup_search_routes(config) -> APIRouter: - router = APIRouter(tags=["search"]) - - @router.get("/api/search/config") - async def get_search_settings() -> Dict[str, Any]: - return get_search_config() - - @router.post("/api/search") - async def do_web_search(request: Request) -> Dict[str, Any]: - """Standalone web search — returns context string + source list. - - Used by Compare mode to pre-search once and share results across panes. - """ - values = await _request_values(request) - query = str(values.get("query") or values.get("q") or "").strip() - if not query: - return {"context": "", "sources": [], "error": "query is required"} - time_filter = values.get("time_filter") or values.get("freshness") - if time_filter is not None: - time_filter = str(time_filter).strip() or None - try: - context, sources = comprehensive_web_search( - query, return_sources=True, time_filter=time_filter, - ) - return {"context": context, "sources": sources} - except Exception as e: - logger.error(f"Standalone web search failed: {e}") - return {"context": "", "sources": [], "error": str(e)} - - @router.get("/api/search/providers") - async def list_search_providers(): - """Return available search providers with config status.""" - providers = [] - for pid, (label, needs_key, needs_url) in PROVIDER_INFO.items(): - if pid == "disabled": - continue - available = True - if needs_key and not _get_provider_key(pid): - available = False - if needs_url and pid == "searxng" and not _get_search_instance(): - available = False - providers.append({ - "id": pid, - "label": label, - "available": available, - }) - return providers - - @router.post("/api/search/query") - async def search_with_provider(request: Request) -> Dict[str, Any]: - """Search using a specific provider. Used by compare search mode.""" - values = await _request_values(request) - query = str(values.get("query") or values.get("q") or "").strip() - provider = str(values.get("provider") or "").strip() - try: - count = int(values.get("count") or values.get("limit") or 10) - except Exception: - count = 10 - if not query: - return {"results": [], "provider": provider, "error": "query is required"} - if provider not in PROVIDER_INFO or provider == "disabled": - return {"results": [], "provider": provider, "error": "Unknown provider"} - t0 = time.time() - try: - results = _call_provider(provider, query, min(count, 20)) - elapsed = round(time.time() - t0, 2) - return {"results": results, "provider": provider, "time": elapsed} - except Exception as e: - elapsed = round(time.time() - t0, 2) - logger.error(f"Search provider {provider} failed: {e}") - return {"results": [], "provider": provider, "time": elapsed, "error": str(e)} - - return router +_sys.modules[__name__] = _canonical diff --git a/routes/skills_routes.py b/routes/skills_routes.py index 711baa2e5..00bef589f 100644 --- a/routes/skills_routes.py +++ b/routes/skills_routes.py @@ -1409,7 +1409,7 @@ def setup_skills_routes(skills_manager: SkillsManager) -> APIRouter: # Prefer the configured DEFAULT (→ Utility) model — not the current chat # session's model. Fall back to the caller's session model only if unset. - url, model, headers = resolve_endpoint("default", owner=user) + url, model, headers = resolve_endpoint("utility", owner=user) if not url or not model: url = url or ((body.get("endpoint_url") or "").strip() or None) model = model or ((body.get("model") or "").strip() or None) diff --git a/routes/vault/__init__.py b/routes/vault/__init__.py new file mode 100644 index 000000000..8aa82701d --- /dev/null +++ b/routes/vault/__init__.py @@ -0,0 +1,5 @@ +"""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 new file mode 100644 index 000000000..7e97500f0 --- /dev/null +++ b/routes/vault/vault_routes.py @@ -0,0 +1,242 @@ +""" +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 7e97500f0..cfed2ba39 100644 --- a/routes/vault_routes.py +++ b/routes/vault_routes.py @@ -1,242 +1,14 @@ -""" -vault_routes.py +"""Backward-compat shim — canonical location is routes/vault/vault_routes.py. -Vaultwarden / Bitwarden CLI integration — config and unlock endpoints. -Stores the BW_SESSION key in data/vault.json with restrictive permissions. +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). """ -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 +import sys as _sys -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 +from routes.vault import vault_routes as _canonical # noqa: F401 -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 +_sys.modules[__name__] = _canonical diff --git a/routes/webhook/__init__.py b/routes/webhook/__init__.py new file mode 100644 index 000000000..e51389e3a --- /dev/null +++ b/routes/webhook/__init__.py @@ -0,0 +1,5 @@ +"""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 new file mode 100644 index 000000000..8d3a704c6 --- /dev/null +++ b/routes/webhook/webhook_routes.py @@ -0,0 +1,395 @@ +"""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 8d3a704c6..7c5e0453e 100644 --- a/routes/webhook_routes.py +++ b/routes/webhook_routes.py @@ -1,395 +1,16 @@ -"""Webhook, API Token, and sync chat routes.""" +"""Backward-compat shim — canonical location is routes/webhook/webhook_routes.py. -import uuid -import logging -from typing import Optional +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 httpx -from fastapi import APIRouter, HTTPException, Request, Form -from pydantic import BaseModel, Field +import sys as _sys -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 +from routes.webhook import webhook_routes as _canonical # noqa: F401 -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 +_sys.modules[__name__] = _canonical diff --git a/services/memory/__init__.py b/services/memory/__init__.py index 53fc80bd8..31fa1d5fa 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 +from .memory import MemoryManager, MemoryStoreUnreadable from .memory_vector import MemoryVectorStore __all__ = [ @@ -10,5 +10,6 @@ __all__ = [ "Memory", "MemorySearchResult", "MemoryManager", + "MemoryStoreUnreadable", "MemoryVectorStore", ] diff --git a/services/memory/memory.py b/services/memory/memory.py index 031c13ac4..b9aaaa2a8 100644 --- a/services/memory/memory.py +++ b/services/memory/memory.py @@ -5,6 +5,16 @@ application runtime instantiates ``src.memory.MemoryManager``, so keeping a parallel implementation here risks silent drift between import paths. """ -from src.memory import MemoryManager, get_text_similarity, tokenize +from src.memory import ( + MemoryManager, + MemoryStoreUnreadable, + get_text_similarity, + tokenize, +) -__all__ = ["MemoryManager", "get_text_similarity", "tokenize"] +__all__ = [ + "MemoryManager", + "MemoryStoreUnreadable", + "get_text_similarity", + "tokenize", +] diff --git a/services/memory/memory_extractor.py b/services/memory/memory_extractor.py index e5f609250..11539263b 100644 --- a/services/memory/memory_extractor.py +++ b/services/memory/memory_extractor.py @@ -17,6 +17,8 @@ import os import re from typing import Optional +from src.memory import MemoryStoreUnreadable + logger = logging.getLogger(__name__) @@ -387,7 +389,13 @@ async def extract_and_store( # Get owner from session _owner = getattr(session, 'owner', None) - existing = memory_manager.load_all() + # 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 added = 0 for fact in facts: @@ -626,7 +634,18 @@ async def audit_memories( # Merge audited entries back with other users' entries if owner: - all_entries = memory_manager.load_all() + # 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", + } 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/services/memory/skill_format.py b/services/memory/skill_format.py index 2b2dfb1b3..628474b04 100644 --- a/services/memory/skill_format.py +++ b/services/memory/skill_format.py @@ -50,7 +50,7 @@ import json import logging import re from dataclasses import dataclass, field -from datetime import datetime +from datetime import datetime, timezone from typing import Any, Dict, List, Optional logger = logging.getLogger(__name__) @@ -441,4 +441,4 @@ class Skill: def _now_iso() -> str: - return datetime.utcnow().strftime("%Y-%m-%dT%H:%M:%SZ") + return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ") diff --git a/services/tts/tts_service.py b/services/tts/tts_service.py index 2120d7720..dd37865a7 100644 --- a/services/tts/tts_service.py +++ b/services/tts/tts_service.py @@ -2,6 +2,7 @@ """Multi-provider TTS service — dispatches to local Kokoro, OpenAI-compatible API, or browser.""" import io +import os import wave import logging import hashlib @@ -41,6 +42,11 @@ class TTSService: self.cache_dir = Path(cache_dir) self.cache_dir.mkdir(parents=True, exist_ok=True) self._kokoro = None # lazy-init + + try: + self.max_cache_bytes = int(os.getenv("ODYSSEUS_TTS_CACHE_MAX_BYTES", 500 * 1024 * 1024)) + except ValueError: + self.max_cache_bytes = 500 * 1024 * 1024 # ── Settings ── @@ -89,6 +95,53 @@ class TTSService: ext = ".mp3" if (len(data) >= 3 and (data[:3] == b'ID3' or (data[0] == 0xff and (data[1] & 0xe0) == 0xe0))) else ".wav" (self.cache_dir / f"{key}{ext}").write_bytes(data) + self._enforce_cache_limit() + + def _enforce_cache_limit(self): + """Evicts oldest files if the cache exceeds the configured byte limit.""" + if self.max_cache_bytes <= 0: + return + + try: + files = [] + total_size = 0 + + # Safely scan files and sum sizes, ignoring files deleted mid-scan + for f in self.cache_dir.iterdir(): + try: + if f.is_file() and f.suffix.lower() in (".mp3", ".wav"): + files.append(f) + total_size += f.stat().st_size + except OSError: + continue + + if total_size > self.max_cache_bytes: + logger.info( + f"TTS cache ({total_size} bytes) exceeded limit ({self.max_cache_bytes} bytes). Evicting oldest files." + ) + + # Sort files by modification time (oldest first) + try: + files.sort(key=lambda f: f.stat().st_mtime) + except OSError as e: + logger.warning(f"Failed to sort cache files by mtime: {e}") + + # Trim down to 80% of max capacity + target_size = self.max_cache_bytes * 0.8 + + while files and total_size > target_size: + f = files.pop(0) + try: + size = f.stat().st_size + f.unlink() + total_size -= size + except OSError as e: + logger.warning(f"Failed to evict cache file {f}: {e}") + continue + + except Exception as e: + logger.warning(f"Error enforcing TTS cache limit: {e}", exc_info=True) + def clear_cache(self): count = 0 for f in self.cache_dir.glob("*.*"): diff --git a/src/agent_loop.py b/src/agent_loop.py index 592ebaec1..cca93fe56 100644 --- a/src/agent_loop.py +++ b/src/agent_loop.py @@ -12,7 +12,7 @@ import json import re import time import logging -from typing import AsyncGenerator, List, Dict, Optional, Set +from typing import Any, AsyncGenerator, List, Dict, Optional, Set from urllib.parse import urlparse from src.llm_core import ( diff --git a/src/ai_interaction.py b/src/ai_interaction.py index 9ee97368f..e777ca32a 100644 --- a/src/ai_interaction.py +++ b/src/ai_interaction.py @@ -22,6 +22,7 @@ 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__) @@ -384,7 +385,15 @@ 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) - memories = _memory_manager.load_all() + # 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.append(entry) _memory_manager.save(memories) diff --git a/src/integrations.py b/src/integrations.py index aa6c4982e..52dd4b2d1 100644 --- a/src/integrations.py +++ b/src/integrations.py @@ -1,11 +1,14 @@ +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 @@ -354,6 +357,152 @@ 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, @@ -409,13 +558,31 @@ 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 + from src.url_safety import check_outbound_url, _default_resolver block_private = os.getenv( "INTEGRATION_API_BLOCK_PRIVATE_IPS", "false" ).lower() == "true" - ok, reason = check_outbound_url(url, block_private=block_private) + # 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 + ) 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() @@ -455,7 +622,9 @@ async def execute_api_call( auth = httpx.BasicAuth(parts[0], parts[1]) try: - async with httpx.AsyncClient(timeout=30.0) as client: + async with httpx.AsyncClient( + timeout=30.0, transport=_PinnedAsyncTransport(pinned_ips) + ) as client: response = await client.request( method, url, diff --git a/src/llm_core.py b/src/llm_core.py index 4dec32376..3e84c1060 100644 --- a/src/llm_core.py +++ b/src/llm_core.py @@ -1237,15 +1237,27 @@ def _anthropic_rejects_temperature(model: str) -> bool: return False # `(?= 4.7. Dated 4.7+ snapshots (`claude-opus-4-7- - # 20260201`) keep their explicit minor and are still matched. - match = re.search(r"(?= 4.7 (issue #5753). Without + # this, every Opus 5 call kept `temperature` and failed with HTTP 400 — visible + # only on paths that pass a temperature, e.g. scheduled tasks inheriting + # `stream_agent_loop`'s 0.3 default, which returned empty responses. + match = re.search( + r"(?= (4, 7) + major = int(match.group(1)) + minor = int(match.group(2)) if match.group(2) else 0 + return (major, minor) >= (4, 7) # Reasoning effort level sent to Mistral thinking-capable models. Mistral's # API accepts "high", "medium", "low", "none" — see diff --git a/src/memory.py b/src/memory.py index 1d8cdbc1e..92efbf5b2 100644 --- a/src/memory.py +++ b/src/memory.py @@ -10,6 +10,18 @@ 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()] @@ -110,21 +122,69 @@ class MemoryManager: with open(self.memory_file, 'w', encoding='utf-8') as f: json.dump([], f, ensure_ascii=False, indent=2) - def load_all(self) -> List[Dict]: - """Load all memory entries from JSON file (unfiltered).""" + 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". + """ 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) - if isinstance(data, list): - return self._validate_entries(data) - except (json.JSONDecodeError, PermissionError) as e: - logger.error("Error loading memory.json: %s", e) - return self._migrate_from_legacy() + 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 - return [] + 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: + logger.error("Error loading memory.json: %s", e) + return [] + + 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() def load(self, owner: str = None) -> List[Dict]: """Load memory entries, optionally filtered by owner.""" @@ -135,7 +195,12 @@ class MemoryManager: def claim_ownerless(self, owner: str): """Assign all ownerless memory entries to the given owner.""" - entries = self.load_all() + 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 changed = False claimed = 0 for entry in entries: @@ -235,7 +300,12 @@ class MemoryManager: if not ids: return id_set = set(ids) - entries = self.load_all() + 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 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 925c59192..8974a6e84 100644 --- a/src/memory_provider.py +++ b/src/memory_provider.py @@ -157,7 +157,11 @@ class NativeMemoryProvider(MemoryProvider): if metadata: entry["metadata"] = dict(metadata) - memories = self.memory_manager.load_all() + # 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.append(entry) self.memory_manager.save(memories) @@ -223,7 +227,10 @@ class NativeMemoryProvider(MemoryProvider): ] async def delete(self, memory_id: str, *, owner: Optional[str] = None) -> bool: - memories = self.memory_manager.load_all() + # 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() remaining = [] deleted_id = None diff --git a/src/tool_parsing.py b/src/tool_parsing.py index 2885cc00f..98dc1b5f6 100644 --- a/src/tool_parsing.py +++ b/src/tool_parsing.py @@ -187,8 +187,12 @@ _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\|)(?=[\t\r\n ]|$)|" + r"(?:^|[\t\r\n ])(?:/?\|end\||\|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 813d57df2..c2eb9ceab 100644 --- a/src/tools/system.py +++ b/src/tools/system.py @@ -46,7 +46,9 @@ 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 "").lower() + 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} from services.memory.skills import SkillsManager from services.memory.skill_format import Skill, slugify from src.constants import DATA_DIR @@ -55,7 +57,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 97f0ae77e..2f1e8d4bf 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=20260722ctxheader4'; +import chatModule from './js/chat.js?v=20260801fix1'; 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?v=20260722ctxheader4'; +import sessionModule from './js/sessions.js'; import memoryModule from './js/memory.js?v=20260722memoryloading1'; import voiceRecorderModule from './js/voiceRecorder.js'; import censorModule from './js/censor.js'; @@ -1689,12 +1689,20 @@ function initializeEventListeners() { const newMemoryInput = el('new-memory-input'); if (newMemoryInput) { - newMemoryInput.addEventListener('keypress', (e) => { - if (e.key === 'Enter') { + // 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(); 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) @@ -3908,85 +3916,10 @@ function startOdysseusApp() { const messageInput = el('message'); const modelPickerWrap = document.getElementById('model-picker-wrap'); - 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); - } + // 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). const _sendIcon = ''; const _micIcon = ''; diff --git a/static/index.html b/static/index.html index 8257660fe..fea4e20ac 100644 --- a/static/index.html +++ b/static/index.html @@ -250,9 +250,9 @@ - + - + @@ -365,6 +365,7 @@ Add a memory — e.g. 'I prefer concise replies' +
@@ -1005,7 +1006,7 @@ var tips = mobile ? phone : desktop; var el = document.getElementById('welcome-tip'); if (el) { - el.textContent = 'Pick a model if you want, or just type.'; + el.textContent = tips[Math.floor(Math.random() * tips.length)]; } fetch('/api/version').then(function(r){return r.json()}).then(function(d){ if (d.version) window._appVersion = d.version; @@ -2504,7 +2505,7 @@ - + @@ -2522,7 +2523,7 @@ - + diff --git a/static/js/chat.js b/static/js/chat.js index ea2d8c1bb..3c8bbe850 100644 --- a/static/js/chat.js +++ b/static/js/chat.js @@ -349,6 +349,9 @@ 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(); @@ -1403,6 +1406,8 @@ 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); @@ -1716,7 +1721,7 @@ import { wireArrowUpRecall, getUserMessagesFromChatHistory } from './composerArr } - const abortCtrl = new AbortController(); + abortCtrl = new AbortController(); abortCtrl._reason = ''; currentAbort = abortCtrl; @@ -1897,7 +1902,7 @@ import { wireArrowUpRecall, getUserMessagesFromChatHistory } from './composerArr let isThinking = false; let thinkingStartTime = null; // Streaming TTS: synthesize sentence-by-sentence during streaming - const streamingTTS = !!(window.aiTTSManager && window.aiTTSManager.autoPlay && window.aiTTSManager.available); + 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) @@ -4787,7 +4792,8 @@ import { wireArrowUpRecall, getUserMessagesFromChatHistory } from './composerArr if (msgIndex < 0) return; const bodyEl = userMsgElement.querySelector('.body'); - const currentText = bodyEl ? bodyEl.textContent.trim().replace(/\s*\[\d+ attachment\(s\)\]$/, '') : ''; + let currentText = (userMsgElement.dataset.raw || (bodyEl ? bodyEl.textContent : '') || '').trim(); + currentText = currentText.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 10709679d..1d6e2e4a9 100644 --- a/static/js/chatRenderer.js +++ b/static/js/chatRenderer.js @@ -478,7 +478,10 @@ 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; -const QWEN_BARE_MARKER_RE = /(?:^|[\t\r\n ])(?:\|?end\|?|\/?\|end\|)(?=[\t\r\n ]|$)|(?:^|[\t\r\n ])assistan(?:t)?(?=[\t\r\n ]|$)/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; // 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 e0b20d6b4..83141bfe9 100644 --- a/static/js/composerArrowUpRecall.js +++ b/static/js/composerArrowUpRecall.js @@ -143,9 +143,9 @@ export function wireArrowUpRecall(composer, getUserMessages, options = {}) { return; } - // 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. + // 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. 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 6a0d3e294..32b906ddc 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, + _sanitizeHtml, _renderEmailSummaryError, _TALON_WROTE, _TALON_FROM, _TALON_SENT, _TALON_SUBJ, _TALON_TO, _TALON_ORIG_RE, _SIG_BLOAT_MIN_CHARS, } from './emailLibrary/utils.js'; @@ -7259,12 +7259,11 @@ async function _generateSummary(reader, data, btn) { if (label) label.textContent = 'Summary'; } } else { - content.innerHTML = `${_esc(result.error || 'Failed to summarize')}`; - panel.remove(); + _renderEmailSummaryError(content, result); } } catch (e) { sp.destroy(); - panel.remove(); + _renderEmailSummaryError(content, null); 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 82a5c86ec..f634c9949 100644 --- a/static/js/emailLibrary/utils.js +++ b/static/js/emailLibrary/utils.js @@ -30,6 +30,25 @@ 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/markdown.js b/static/js/markdown.js index 8735b83e7..f249facc9 100644 --- a/static/js/markdown.js +++ b/static/js/markdown.js @@ -758,30 +758,36 @@ export function mdToHtml(src, opts) { // Remove empty paragraphs s = s.replace(/

<\/p>/g, ''); + // Every restore below passes a function replacer rather than the block string + // itself. With a string replacement, `String.replace` reads `$&`, `` $` ``, + // `$'` and `$$` in the *replacement* as substitution patterns, so a restored + // block containing them is corrupted: `$&` re-inserts the placeholder, `` $` `` + // and `$'` splice in the surrounding document, and `$$` collapses to `$`. Those + // sequences are ordinary content in fenced code (`perl -pe 's/x/$& y/'`, + // `echo "$$USD"`). A function replacer inserts its return value verbatim. + // CRITICAL: Restore allowed HTML blocks first allowedHtmlBlocks.forEach((block, index) => { - s = s.replace(`___ALLOWED_HTML_${index}___`, block); + s = s.replace(`___ALLOWED_HTML_${index}___`, () => block); }); // Restore math blocks mathBlocks.forEach((block, index) => { - s = s.replace(`___MATH_BLOCK_${index}___`, block); + s = s.replace(`___MATH_BLOCK_${index}___`, () => block); }); // Restore mermaid diagram blocks mermaidBlocks.forEach((block, index) => { - s = s.replace(`___MERMAID_BLOCK_${index}___`, block); + s = s.replace(`___MERMAID_BLOCK_${index}___`, () => block); }); // CRITICAL: Restore code blocks at the end codeBlocks.forEach((block, index) => { - s = s.replace(`___CODE_BLOCK_${index}___`, block); + s = s.replace(`___CODE_BLOCK_${index}___`, () => block); }); // Restore inline code spans last, so placeholders carried inside restored - // /allowed-HTML blocks are resolved too. The function replacer keeps the - // escaped code literal — e.g. a shell snippet like `echo $1` is not treated - // as a regex back-reference. + // /allowed-HTML blocks are resolved too. inlineCodeBlocks.forEach((block, index) => { s = s.replace(`___INLINE_CODE_${index}___`, () => block); }); diff --git a/static/js/sessions.js b/static/js/sessions.js index cf59d478c..edf83c8a4 100644 --- a/static/js/sessions.js +++ b/static/js/sessions.js @@ -1847,6 +1847,10 @@ 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 { @@ -2313,6 +2317,7 @@ 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/settings.js b/static/js/settings.js index 72936adee..540acff00 100644 --- a/static/js/settings.js +++ b/static/js/settings.js @@ -3031,12 +3031,14 @@ async function initEmailAccountsSettings() { const body = { name: el('eaf-name').value.trim() || el('eaf-from').value.trim(), from_address: el('eaf-from').value.trim(), + display_name: el('eaf-display-name').value.trim(), imap_host: el('eaf-imap-host').value.trim(), imap_port: parseInt(el('eaf-imap-port').value) || 993, imap_user: el('eaf-imap-user').value.trim(), imap_starttls: el('eaf-imap-starttls').checked, smtp_host: el('eaf-smtp-host').value.trim(), smtp_port: parseInt(el('eaf-smtp-port').value) || 587, + smtp_security: el('eaf-smtp-security').value, smtp_user: el('eaf-imap-user').value.trim(), }; if (!body.name) { el('eaf-msg').textContent = 'Enter a Name or Email first'; el('eaf-msg').style.color = 'var(--red)'; return; } @@ -5788,29 +5790,30 @@ export function close() { window.history.replaceState(null, '', clean); const success = sp.has('email_oauth_success'); const errMsg = sp.get('email_oauth_error') || ''; - // Open settings → integrations after the app has initialised. - function _tryOpen() { - if (window.settingsModule && typeof window.settingsModule.open === 'function') { - window.settingsModule.open('integrations'); - // Brief toast-style banner. - const banner = document.createElement('div'); - banner.textContent = success - ? '✓ Google account connected — email is ready' - : `Google OAuth failed: ${errMsg || 'unknown error'}`; - Object.assign(banner.style, { - position: 'fixed', bottom: '24px', left: '50%', transform: 'translateX(-50%)', - background: success ? 'var(--accent, #50fa7b)' : 'var(--red, #ff5555)', - color: '#000', padding: '8px 18px', borderRadius: '6px', fontSize: '12px', - fontWeight: '600', zIndex: '99999', pointerEvents: 'none', - boxShadow: '0 2px 12px rgba(0,0,0,0.3)', - }); - document.body.appendChild(banner); - setTimeout(() => banner.remove(), 4000); - } else { - setTimeout(_tryOpen, 100); - } + // Open settings → integrations once the document is ready. This module owns + // the open() API, so it does not need to wait for a window-level alias. + function _showResult() { + open('integrations'); + // Brief toast-style banner. + const banner = document.createElement('div'); + banner.textContent = success + ? 'Google account connected — email is ready' + : `Google OAuth failed: ${errMsg || 'unknown error'}`; + Object.assign(banner.style, { + position: 'fixed', bottom: '24px', left: '50%', transform: 'translateX(-50%)', + background: success ? 'var(--accent, #50fa7b)' : 'var(--red, #ff5555)', + color: '#000', padding: '8px 18px', borderRadius: '6px', fontSize: '12px', + fontWeight: '600', zIndex: '99999', pointerEvents: 'none', + boxShadow: '0 2px 12px rgba(0,0,0,0.3)', + }); + document.body.appendChild(banner); + setTimeout(() => banner.remove(), 4000); + } + if (document.readyState === 'loading') { + document.addEventListener('DOMContentLoaded', _showResult, { once: true }); + } else { + _showResult(); } - _tryOpen(); })(); const settingsModule = { open, close, initIntegrations, initUnifiedIntegrations, syncAdminVisibility, refreshAiModelEndpoints }; diff --git a/static/js/skills.js b/static/js/skills.js index 84974d446..b45403570 100644 --- a/static/js/skills.js +++ b/static/js/skills.js @@ -83,11 +83,9 @@ 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; - if (cascade && loaded && !_loadPromise && _playSkillsCascade()) { - _cascadeNext = false; - updateCount(); - return; - } + // 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 (_loadPromise) return _loadPromise; _loadPromise = (async () => { try { diff --git a/tests/test_api_chat_security.py b/tests/test_api_chat_security.py index 7dcec324e..d92a31620 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_routes.py", + Path(__file__).resolve().parent.parent / "routes" / "webhook" / "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 2df5936ef..135be78ee 100644 --- a/tests/test_backup_import_cross_user_dedup.py +++ b/tests/test_backup_import_cross_user_dedup.py @@ -27,6 +27,9 @@ 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 eadc3bc94..022fcbc02 100644 --- a/tests/test_composer_arrow_up_recall_js.py +++ b/tests/test_composer_arrow_up_recall_js.py @@ -306,3 +306,24 @@ 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 new file mode 100644 index 000000000..68d049a62 --- /dev/null +++ b/tests/test_document_routes_shim.py @@ -0,0 +1,29 @@ +"""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_oauth_connect_smtp_security.py b/tests/test_email_oauth_connect_smtp_security.py new file mode 100644 index 000000000..21c4224a6 --- /dev/null +++ b/tests/test_email_oauth_connect_smtp_security.py @@ -0,0 +1,15 @@ +"""Regression coverage for SMTP security saved before Google OAuth.""" + +from pathlib import Path + + +_REPO = Path(__file__).resolve().parents[1] + + +def test_email_tab_oauth_connect_persists_selected_smtp_security(): + source = (_REPO / "static" / "js" / "settings.js").read_text(encoding="utf-8") + start = source.index("el('eaf-oauth-btn').addEventListener") + handler_body = source[start:source.index("if (!body.name)", start)] + + assert "smtp_security: el('eaf-smtp-security').value" in handler_body + assert "display_name: el('eaf-display-name').value.trim()" in handler_body diff --git a/tests/test_email_oauth_settings_redirect.py b/tests/test_email_oauth_settings_redirect.py new file mode 100644 index 000000000..f7d588132 --- /dev/null +++ b/tests/test_email_oauth_settings_redirect.py @@ -0,0 +1,19 @@ +"""Regression coverage for the settings UI after Google OAuth redirects.""" + +from pathlib import Path + + +_REPO = Path(__file__).resolve().parents[1] + + +def test_oauth_redirect_uses_the_module_local_settings_api(): + source = (_REPO / "static" / "js" / "settings.js").read_text(encoding="utf-8") + handler = source[ + source.index("(function _handleOauthRedirect"): + source.index("const settingsModule =") + ] + + assert "open('integrations');" in handler + assert "window.settingsModule" not in handler + assert "window.__odysseusAppStarted" not in handler + assert "document.addEventListener('DOMContentLoaded', _showResult, { once: true })" in handler diff --git a/tests/test_email_summary_error_ui_js.py b/tests/test_email_summary_error_ui_js.py new file mode 100644 index 000000000..1afc3bec9 --- /dev/null +++ b/tests/test_email_summary_error_ui_js.py @@ -0,0 +1,52 @@ +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 new file mode 100644 index 000000000..b0ab7b3be --- /dev/null +++ b/tests/test_email_summary_llm.py @@ -0,0 +1,406 @@ +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 7c5bb1645..636270a56 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_routes.py").read_text() + document_routes = Path("routes/document/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 53dc671c5..f23cc40de 100644 --- a/tests/test_integration_api_call_ssrf.py +++ b/tests/test_integration_api_call_ssrf.py @@ -9,8 +9,13 @@ 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 @@ -97,3 +102,238 @@ 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 bf1ec7d05..a0ad61b4a 100644 --- a/tests/test_integrations_api_call_truncation.py +++ b/tests/test_integrations_api_call_truncation.py @@ -83,9 +83,10 @@ 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; 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")), + # 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"]), ): return await integrations.execute_api_call("test_integ", "GET", "/items") @@ -101,9 +102,10 @@ 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; 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")), + # 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"]), ): result = await integrations.execute_api_call("test_integ", "GET", path) return result, mock_client diff --git a/tests/test_issue_description_check.py b/tests/test_issue_description_check.py new file mode 100644 index 000000000..196f21cfc --- /dev/null +++ b/tests/test_issue_description_check.py @@ -0,0 +1,86 @@ +"""Regression coverage for issue-description label lifecycle events.""" + +import json +import shutil +import subprocess +from pathlib import Path + +import pytest + + +_REPO = Path(__file__).resolve().parent.parent +_CHECKER = _REPO / ".github" / "scripts" / "check-issue-description.js" +_WORKFLOW = _REPO / ".github" / "workflows" / "issue-description-check.yml" +pytestmark = pytest.mark.skipif(not shutil.which("node"), reason="node not on PATH") + + +def _run_closed_issue(action): + harness = r""" +const checkIssueDescription = require(process.argv[1]); +const action = process.argv[2]; +const calls = []; +const unexpected = (name) => async () => { + throw new Error(`${name} should not be called for a closed issue`); +}; + +const github = { + rest: { + issues: { + removeLabel: async (params) => calls.push({ method: 'removeLabel', params }), + getLabel: unexpected('getLabel'), + addLabels: unexpected('addLabels'), + listComments: unexpected('listComments'), + createComment: unexpected('createComment'), + updateComment: unexpected('updateComment'), + deleteComment: unexpected('deleteComment'), + }, + }, +}; +const context = { + payload: { + action, + issue: { number: 42, state: 'closed', body: '', labels: [] }, + }, + repo: { owner: 'odysseus-dev', repo: 'odysseus' }, +}; +const core = { + warning: unexpected('core.warning'), + setFailed: unexpected('core.setFailed'), +}; + +checkIssueDescription({ github, context, core }) + .then(() => process.stdout.write(JSON.stringify(calls))) + .catch((error) => { + console.error(error); + process.exitCode = 1; + }); +""" + proc = subprocess.run( + ["node", "-e", harness, str(_CHECKER), action], + capture_output=True, + text=True, + cwd=str(_REPO), + timeout=30, + ) + assert proc.returncode == 0, proc.stderr + return json.loads(proc.stdout) + + +def test_workflow_handles_issue_closures(): + workflow = _WORKFLOW.read_text() + assert "types: [opened, edited, reopened, closed]" in workflow + + +@pytest.mark.parametrize("action", ["closed", "edited"]) +def test_closed_issue_only_drops_ready_for_review(action): + assert _run_closed_issue(action) == [ + { + "method": "removeLabel", + "params": { + "owner": "odysseus-dev", + "repo": "odysseus", + "issue_number": 42, + "name": "ready for review", + }, + } + ] diff --git a/tests/test_llm_core_anthropic_temp_omit.py b/tests/test_llm_core_anthropic_temp_omit.py index 2274f1dc9..f7d26aef0 100644 --- a/tests/test_llm_core_anthropic_temp_omit.py +++ b/tests/test_llm_core_anthropic_temp_omit.py @@ -29,6 +29,13 @@ from src.llm_core import _anthropic_rejects_temperature, _build_anthropic_payloa "anthropic/claude-opus-4-7", # tolerate a provider-prefixed id "claude-opus-4-10", # future minor still >= 4.7 "claude-opus-5-0", # future major + # Major-only ids: a missing minor reads as `.0`, so these are >= 4.7 too + # (issue #5753). Before the fix the version pattern required a minor, so + # these fell through to "accepts temperature" and every call 400'd. + "claude-opus-5", + "claude-opus-5-20260101", # major-only + dated snapshot + "anthropic/claude-opus-5", # major-only behind a provider prefix + "claude-opus-6", # future major-only ], ) def test_opus_47_plus_rejects_temperature(model): @@ -48,7 +55,10 @@ def test_opus_47_plus_rejects_temperature(model): "claude-opus-4-6-20251201", # dated 4.6 snapshot — older, still keeps temperature "claude-sonnet-4-6", "claude-3-5-sonnet", - "claude-3-opus-20240229", # legacy Claude 3 Opus — no opus-N-M pattern, kept + "claude-3-opus-20240229", # legacy Claude 3 Opus — date directly after + # "opus-", so the major must not swallow it as version 20240229 (that is + # what makes capping the major at 1-2 digits necessary once the minor + # became optional in #5753). "claude-haiku-4-5", "claude-x", "octopus-4-8", # "opus" only as a substring of another word — must not match @@ -87,6 +97,20 @@ def test_payload_keeps_temperature_for_older_models(): assert _payload("claude-3-5-sonnet", 1.2)["temperature"] == 1.0 +def test_payload_omits_temperature_for_major_only_opus_5(): + # Issue #5753: the scheduled-task path calls stream_agent_loop() without a + # temperature and inherits its 0.3 default, so `claude-opus-5` 400'd on every + # run and surfaced as "the model returned an empty response". Interactive chat + # leaves temperature None and never hit it. + assert "temperature" not in _payload("claude-opus-5", 0.3) + + +def test_payload_keeps_temperature_for_legacy_claude_3_opus(): + # Guards the major-digit cap: `opus-20240229` must not parse as version + # 20240229, or Claude 3 Opus would silently lose the caller's temperature. + assert _payload("claude-3-opus-20240229", 0.5)["temperature"] == 0.5 + + def test_payload_keeps_temperature_for_dated_opus_4_0(): # Anthropic's dated id for Opus 4.0 (claude-opus-4-20250514) is in this repo's # ANTHROPIC_MODELS list. The date must not be misread as a >= 4.7 minor, or the diff --git a/tests/test_manage_skills_action_required.py b/tests/test_manage_skills_action_required.py new file mode 100644 index 000000000..4efae8026 --- /dev/null +++ b/tests/test_manage_skills_action_required.py @@ -0,0 +1,24 @@ +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_markdown_rendering_js.py b/tests/test_markdown_rendering_js.py index 2ffe8914f..536789b89 100644 --- a/tests/test_markdown_rendering_js.py +++ b/tests/test_markdown_rendering_js.py @@ -214,6 +214,50 @@ def test_inline_code_content_is_html_escaped(node_available): assert "" not in html +def test_fenced_code_keeps_dollar_ampersand(node_available): + # Issue #5663: the block-restore pass used a string replacement, so `$&` in a + # restored block was read as "the matched text" and re-inserted the + # placeholder. `perl -pe 's/world/$& again/'` rendered as + # "s/world/___CODE_BLOCK_0___amp; again/" — the trailing "amp;" is the orphan + # left behind after `$&` consumed the `$&` of the escaped `$&`. + html = _run_markdown_case( + "```sh\necho \"hello world\" | perl -pe 's/world/$& again/'\n```" + ) + + assert "___CODE_BLOCK_" not in html + assert "s/world/$& again/" in html + assert "amp; again" not in html.replace("$& again", "") + + +def test_fenced_code_keeps_dollar_backtick_and_quote(node_available): + # `` $` `` and `$'` splice the text before/after the placeholder into the + # block. Unlike `$&` these leave no placeholder behind — the characters just + # vanish — so assert the content survives verbatim. + html = _run_markdown_case("```sh\nsed \"s/$`/x/\" && sed \"s/$'/y/\"\n```") + + assert "___CODE_BLOCK_" not in html + assert "s/$`/x/" in html + assert "s/$'/y/" in html + + +def test_fenced_code_keeps_double_dollar(node_available): + # `$$` collapsed to a single `$` in the restored block. + html = _run_markdown_case('```sh\necho "$$USD and $$"\n```') + + assert "$$USD and $$" in html + + +def test_mermaid_block_keeps_dollar_ampersand(node_available): + # The mermaid restore site had the same hazard: a node label containing `$&` + # re-inserted the ___MERMAID_BLOCK_n___ placeholder into the diagram source, + # which then fails to parse. The math and allowed-HTML sites are fixed the + # same way; they need KaTeX/sanitizer conditions this harness doesn't set up. + html = _run_markdown_case('```mermaid\ngraph TD; A["$&"] --> B;\n```') + + assert "___MERMAID_BLOCK_" not in html + assert "$&" in html + + def test_currency_dollar_amounts_are_not_rendered_as_math(node_available): # "$5 to $10" used to pair the two dollar signs as inline-math delimiters # and render "5 to" through KaTeX. Pandoc-style rules now reject it: the diff --git a/tests/test_mcp_dependency_compatibility.py b/tests/test_mcp_dependency_compatibility.py new file mode 100644 index 000000000..9efefe4fe --- /dev/null +++ b/tests/test_mcp_dependency_compatibility.py @@ -0,0 +1,15 @@ +"""Regression coverage for the built-in MCP servers' SDK compatibility line.""" + +from pathlib import Path + + +REQUIREMENTS = Path(__file__).resolve().parents[1] / "requirements.txt" + + +def test_mcp_requirement_excludes_breaking_v2_sdk(): + requirements = [ + line.split("#", 1)[0].strip().replace(" ", "") + for line in REQUIREMENTS.read_text(encoding="utf-8").splitlines() + ] + + assert "mcp<2" in requirements diff --git a/tests/test_memory_add_submit_regression.py b/tests/test_memory_add_submit_regression.py new file mode 100644 index 000000000..450d63003 --- /dev/null +++ b/tests/test_memory_add_submit_regression.py @@ -0,0 +1,54 @@ +"""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 49702c17f..06ca31667 100644 --- a/tests/test_memory_extractor_vector_cross_tenant.py +++ b/tests/test_memory_extractor_vector_cross_tenant.py @@ -67,6 +67,12 @@ 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 new file mode 100644 index 000000000..4b9076065 --- /dev/null +++ b/tests/test_memory_store_unreadable_no_wipe.py @@ -0,0 +1,255 @@ +"""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 dafbad594..f48a1f7e2 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_routes.py", "ai_tidy_documents") + body = _function_source("routes/document/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_search_routes_shim.py b/tests/test_search_routes_shim.py new file mode 100644 index 000000000..a8b278488 --- /dev/null +++ b/tests/test_search_routes_shim.py @@ -0,0 +1,11 @@ +"""Regression test for the search route shim (slice 2j, #4082/#4071).""" + +import importlib + +import routes.search_routes as _shim_search # noqa: F401 + + +def test_legacy_and_canonical_search_module_are_same_object(): + legacy = importlib.import_module("routes.search_routes") + canonical = importlib.import_module("routes.search.search_routes") + assert legacy is canonical diff --git a/tests/test_skill_format_timestamp.py b/tests/test_skill_format_timestamp.py new file mode 100644 index 000000000..a9309bdc1 --- /dev/null +++ b/tests/test_skill_format_timestamp.py @@ -0,0 +1,58 @@ +"""Regression for issue #5697 — skill timestamps must not use ``datetime.utcnow()``. + +``_now_iso()`` builds the ``created`` value in skill frontmatter. ``utcnow()`` +returns a *naive* datetime and has been deprecated since Python 3.12, scheduled +for removal. The replacement must stay timezone-aware while keeping the +serialized ``YYYY-MM-DDTHH:MM:SSZ`` shape, so skill files written by older +versions keep parsing. + +The UTC check matters on its own: a bare ``datetime.now()`` also produces the +right shape, but emits local wall time, which would silently backdate or +postdate skills for every user outside UTC. +""" + +import os +import re +import time +import warnings +from datetime import datetime, timezone + +import pytest + +from services.memory.skill_format import _now_iso + +_ISO_Z = re.compile(r"^\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}Z$") + + +def test_now_iso_keeps_serialized_shape(): + assert _ISO_Z.match(_now_iso()) + + +def test_now_iso_emits_no_deprecation_warning(): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + _now_iso() + assert not [w for w in caught if issubclass(w.category, DeprecationWarning)] + + +@pytest.mark.skipif( + not hasattr(time, "tzset"), + reason="time.tzset is unavailable on this platform", +) +def test_now_iso_is_utc_not_local_time(): + """Pin UTC under a non-UTC local timezone, where the two visibly diverge.""" + original_tz = os.environ.get("TZ") + os.environ["TZ"] = "Asia/Amman" # UTC+3, never UTC + time.tzset() + try: + emitted = datetime.strptime(_now_iso(), "%Y-%m-%dT%H:%M:%SZ").replace( + tzinfo=timezone.utc + ) + drift = abs((emitted - datetime.now(timezone.utc)).total_seconds()) + assert drift < 60, f"timestamp is {drift}s off UTC — local time leaked in" + finally: + if original_tz is None: + os.environ.pop("TZ", None) + else: + os.environ["TZ"] = original_tz + time.tzset() diff --git a/tests/test_tool_parsing_bare_end_marker.py b/tests/test_tool_parsing_bare_end_marker.py new file mode 100644 index 000000000..6167c8dde --- /dev/null +++ b/tests/test_tool_parsing_bare_end_marker.py @@ -0,0 +1,96 @@ +"""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_tts_service_enforce_cache_limit.py b/tests/test_tts_service_enforce_cache_limit.py new file mode 100644 index 000000000..1da9d16c0 --- /dev/null +++ b/tests/test_tts_service_enforce_cache_limit.py @@ -0,0 +1,97 @@ +import os +import time +from pathlib import Path +import pytest + +# Adjust the import path if your file is directly in ./services instead of ./services/tts +from services.tts.tts_service import TTSService + +def test_cache_under_limit(tmp_path, monkeypatch): + """Test that writing a file under the size limit does not trigger eviction.""" + # Set a tiny limit: 100 bytes + monkeypatch.setenv("ODYSSEUS_TTS_CACHE_MAX_BYTES", "100") + + # Initialize service with pytest's temporary directory + service = TTSService(cache_dir=str(tmp_path)) + + # Write a 40-byte file (under the 100-byte limit) + service._put_cache("test_key", b"x" * 40) + + # Verify the file was written and nothing was deleted + files = list(tmp_path.glob("*.*")) + assert len(files) == 1 + assert sum(f.stat().st_size for f in files) == 40 + +def test_cache_exceeds_limit_triggers_eviction(tmp_path, monkeypatch): + """Test that exceeding the limit evicts the oldest files down to 80% capacity.""" + # Set limit to 100 bytes. 80% target capacity will be 80 bytes. + monkeypatch.setenv("ODYSSEUS_TTS_CACHE_MAX_BYTES", "100") + service = TTSService(cache_dir=str(tmp_path)) + + # 1. Setup: Manually create two older files (40 bytes each) + file1 = tmp_path / "oldest.wav" + file2 = tmp_path / "middle.wav" + + file1.write_bytes(b"a" * 40) + file2.write_bytes(b"b" * 40) + + # Spoof timestamps so file1 is explicitly older than file2 + now = time.time() + os.utime(file1, (now - 100, now - 100)) # 100 seconds ago + os.utime(file2, (now - 50, now - 50)) # 50 seconds ago + + # 2. Action: Write a 3rd file using the service method (40 bytes) + # Total cache is now 120 bytes, which exceeds 100. + # It should delete oldest (file1) to drop to 80 bytes (which matches the 80% target). + service._put_cache("newest", b"c" * 40) + + # 3. Assertions + # The newest file should exist (saved as .wav because it lacks MP3 magic bytes) + newest_file = tmp_path / "newest.wav" + + assert not file1.exists(), "The oldest file should have been evicted." + assert file2.exists(), "The middle file should still exist." + assert newest_file.exists(), "The newest file should have been saved." + + # Verify the final directory size is <= 80 bytes + total_size = sum(f.stat().st_size for f in tmp_path.glob("*.*")) + assert total_size <= 80 + +def test_cache_limit_disabled(tmp_path, monkeypatch): + """Test that setting max bytes to 0 disables eviction.""" + monkeypatch.setenv("ODYSSEUS_TTS_CACHE_MAX_BYTES", "0") + service = TTSService(cache_dir=str(tmp_path)) + + # Write 3 large files that would normally trigger eviction + service._put_cache("file1", b"x" * 1000) + service._put_cache("file2", b"x" * 1000) + service._put_cache("file3", b"x" * 1000) + + # Ensure nothing was deleted + files = list(tmp_path.glob("*.*")) + assert len(files) == 3 + assert sum(f.stat().st_size for f in files) == 3000 + +def test_cache_eviction_handles_unlink_error_gracefully(tmp_path, monkeypatch): + """Test that if unlinking a file fails, _put_cache still succeeds without raising.""" + service = TTSService(cache_dir=str(tmp_path)) + service.max_cache_bytes = 50 + + # Create a file to evict + old_file = tmp_path / "old.wav" + old_file.write_bytes(b"x" * 40) + + # Monkeypatch unlink on Path objects to simulate a PermissionError / file-lock failure + def mock_unlink(self_path): + raise OSError("Permission denied / file locked") + + monkeypatch.setattr(Path, "unlink", mock_unlink) + + # Writing a new file triggers eviction which encounters the mocked unlink error + try: + service._put_cache("new_key", b"y" * 40) + except Exception as e: + pytest.fail(f"_put_cache raised an exception during failed eviction: {e}") + + # The new file should still be written successfully + assert (tmp_path / "new_key.wav").exists() \ No newline at end of file diff --git a/tests/test_vault_routes_shim.py b/tests/test_vault_routes_shim.py new file mode 100644 index 000000000..9577395f7 --- /dev/null +++ b/tests/test_vault_routes_shim.py @@ -0,0 +1,11 @@ +"""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 f0d3a184d..29de101a3 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_routes.py").read_text() + document_source = (ROOT / "routes" / "document" / "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 new file mode 100644 index 000000000..f6312e8e6 --- /dev/null +++ b/tests/test_webhook_routes_shim.py @@ -0,0 +1,11 @@ +"""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