From d96c7af3df769508de01900b2264520b649caa4c Mon Sep 17 00:00:00 2001 From: Dividesbyzer0 <54127744+zoomdbz@users.noreply.github.com> Date: Mon, 27 Jul 2026 11:29:29 -0400 Subject: [PATCH 01/32] fix(agent): import Any for tool event helper (#5735) --- src/agent_loop.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 ( From 01790c2f08233f1b8d15457fbe7dd0d089d7b8ef Mon Sep 17 00:00:00 2001 From: RaresKeY <158580472+RaresKeY@users.noreply.github.com> Date: Tue, 28 Jul 2026 18:11:34 +0100 Subject: [PATCH 02/32] fix(mcp): keep built-in servers on SDK v1 (#5820) --- requirements.txt | 5 ++++- tests/test_mcp_dependency_compatibility.py | 15 +++++++++++++++ 2 files changed, 19 insertions(+), 1 deletion(-) create mode 100644 tests/test_mcp_dependency_compatibility.py 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/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 From 5104a9a96710f682d641a1e10c4977ae196ab912 Mon Sep 17 00:00:00 2001 From: Boody Date: Tue, 28 Jul 2026 21:34:03 +0300 Subject: [PATCH 03/32] feat(tts): implement TTS cache size limit and eviction policy --- .env.example | 1 + services/tts/tts_service.py | 35 +++++++++ tests/test_tts_service_enforce_cache_limit.py | 73 +++++++++++++++++++ 3 files changed, 109 insertions(+) create mode 100644 tests/test_tts_service_enforce_cache_limit.py diff --git a/.env.example b/.env.example index d23276eb8..4eb4695e0 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=52428800 # TTS cache (500 MB) # ============================================================ # Host Docker access (explicit opt-in) diff --git a/services/tts/tts_service.py b/services/tts/tts_service.py index 2120d7720..e1c67d4da 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,35 @@ 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 + + files = [f for f in self.cache_dir.glob("*.*") if f.is_file()] + total_size = sum(f.stat().st_size for f in files) + + 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) + files.sort(key=lambda f: f.stat().st_mtime) + + # Trim down to 80% of max capacity so we aren't constantly triggering this on every new generation + 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 FileNotFoundError: + # File was deleted by another process + continue + def clear_cache(self): count = 0 for f in self.cache_dir.glob("*.*"): 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..f4a31c75b --- /dev/null +++ b/tests/test_tts_service_enforce_cache_limit.py @@ -0,0 +1,73 @@ +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("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("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("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 \ No newline at end of file From 98e4d8451bdcdf2f682bb39a54cccfeb0299ab9a Mon Sep 17 00:00:00 2001 From: Boody Date: Tue, 28 Jul 2026 22:00:54 +0300 Subject: [PATCH 04/32] fix(tests): update environment variable for TTS cache limit to include ODYSSEUS prefix --- tests/test_tts_service_enforce_cache_limit.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/test_tts_service_enforce_cache_limit.py b/tests/test_tts_service_enforce_cache_limit.py index f4a31c75b..8a726557b 100644 --- a/tests/test_tts_service_enforce_cache_limit.py +++ b/tests/test_tts_service_enforce_cache_limit.py @@ -9,7 +9,7 @@ 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("TTS_CACHE_MAX_BYTES", "100") + monkeypatch.setenv("ODYSSEUS_TTS_CACHE_MAX_BYTES", "100") # Initialize service with pytest's temporary directory service = TTSService(cache_dir=str(tmp_path)) @@ -25,7 +25,7 @@ def test_cache_under_limit(tmp_path, monkeypatch): 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("TTS_CACHE_MAX_BYTES", "100") + 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) @@ -59,7 +59,7 @@ def test_cache_exceeds_limit_triggers_eviction(tmp_path, monkeypatch): def test_cache_limit_disabled(tmp_path, monkeypatch): """Test that setting max bytes to 0 disables eviction.""" - monkeypatch.setenv("TTS_CACHE_MAX_BYTES", "0") + monkeypatch.setenv("ODYSSEUS_TTS_CACHE_MAX_BYTES", "0") service = TTSService(cache_dir=str(tmp_path)) # Write 3 large files that would normally trigger eviction From 25a4d134b1611b5329675ce212070bedcb8daad1 Mon Sep 17 00:00:00 2001 From: "Tal.Yuan" Date: Wed, 29 Jul 2026 04:26:29 +0800 Subject: [PATCH 05/32] refactor(routes): move search domain into routes/search/ subpackage (#5779) Slice 2j of the route-domain reorganization (#4082/#4071). Moves search_routes.py into routes/search/, leaving a backward-compat sys.modules shim. Pure file reorganization, no behavior change. --- app.py | 2 +- routes/search/__init__.py | 5 ++ routes/search/search_routes.py | 111 +++++++++++++++++++++++++++++ routes/search_routes.py | 116 +++---------------------------- tests/test_search_routes_shim.py | 11 +++ 5 files changed, 137 insertions(+), 108 deletions(-) create mode 100644 routes/search/__init__.py create mode 100644 routes/search/search_routes.py create mode 100644 tests/test_search_routes_shim.py diff --git a/app.py b/app.py index e740ad518..2ae5ec761 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 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/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 From 61c138d9e7f50eaaf73185d1249f8debe809094b Mon Sep 17 00:00:00 2001 From: Boody Date: Wed, 29 Jul 2026 12:26:09 +0300 Subject: [PATCH 06/32] fixed .env.example ODYSSEUS_TTS_CACHE_MAX_BYTES into correct 500 MBs --- .env.example | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.env.example b/.env.example index 4eb4695e0..2d1be3373 100644 --- a/.env.example +++ b/.env.example @@ -189,7 +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=52428800 # TTS cache (500 MB) +# ODYSSEUS_TTS_CACHE_MAX_BYTES=524288000 # TTS cache (500 MB) # ============================================================ # Host Docker access (explicit opt-in) From 46905ab9b0f4a2ecf85996b23f47f678501792e4 Mon Sep 17 00:00:00 2001 From: Boody Date: Wed, 29 Jul 2026 12:32:43 +0300 Subject: [PATCH 07/32] added ODYSSEUS_TTS_CACHE_MAX_BYTES env variable to docker compose files --- docker-compose.gpu-amd.yml | 1 + docker-compose.gpu-nvidia.yml | 1 + 2 files changed, 2 insertions(+) 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:-} From 9914651cc9f1756fd1610952af42de8848b6015d Mon Sep 17 00:00:00 2001 From: Boody Date: Wed, 29 Jul 2026 12:41:31 +0300 Subject: [PATCH 08/32] improve cache eviction logic to handle file access errors and ensure stability --- services/tts/tts_service.py | 64 ++++++++++++++++++++++++------------- 1 file changed, 41 insertions(+), 23 deletions(-) diff --git a/services/tts/tts_service.py b/services/tts/tts_service.py index e1c67d4da..c7f787954 100644 --- a/services/tts/tts_service.py +++ b/services/tts/tts_service.py @@ -98,31 +98,49 @@ class TTSService: 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 + """Evicts oldest files if the cache exceeds the configured byte limit.""" + if self.max_cache_bytes <= 0: + return - files = [f for f in self.cache_dir.glob("*.*") if f.is_file()] - total_size = sum(f.stat().st_size for f in files) + try: + files = [] + total_size = 0 - 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) - files.sort(key=lambda f: f.stat().st_mtime) - - # Trim down to 80% of max capacity so we aren't constantly triggering this on every new generation - 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 FileNotFoundError: - # File was deleted by another process - continue + # Safely scan files and sum sizes, ignoring files deleted mid-scan + for f in self.cache_dir.glob("*.*"): + try: + if f.is_file(): + 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 From d183fe545b25766badbf5ada731ccd0fe7b0c594 Mon Sep 17 00:00:00 2001 From: Boody Date: Wed, 29 Jul 2026 12:42:20 +0300 Subject: [PATCH 09/32] add test for cache eviction handling unlink errors gracefully --- tests/test_tts_service_enforce_cache_limit.py | 26 ++++++++++++++++++- 1 file changed, 25 insertions(+), 1 deletion(-) diff --git a/tests/test_tts_service_enforce_cache_limit.py b/tests/test_tts_service_enforce_cache_limit.py index 8a726557b..1da9d16c0 100644 --- a/tests/test_tts_service_enforce_cache_limit.py +++ b/tests/test_tts_service_enforce_cache_limit.py @@ -70,4 +70,28 @@ def test_cache_limit_disabled(tmp_path, monkeypatch): # Ensure nothing was deleted files = list(tmp_path.glob("*.*")) assert len(files) == 3 - assert sum(f.stat().st_size for f in files) == 3000 \ No newline at end of file + 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 From 2e631ad8160c76c42f757471d372cffa59ac88e4 Mon Sep 17 00:00:00 2001 From: Boody Date: Wed, 29 Jul 2026 12:47:55 +0300 Subject: [PATCH 10/32] improve cache size calculation by filtering file types --- services/tts/tts_service.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/services/tts/tts_service.py b/services/tts/tts_service.py index c7f787954..dd37865a7 100644 --- a/services/tts/tts_service.py +++ b/services/tts/tts_service.py @@ -107,9 +107,9 @@ class TTSService: total_size = 0 # Safely scan files and sum sizes, ignoring files deleted mid-scan - for f in self.cache_dir.glob("*.*"): + for f in self.cache_dir.iterdir(): try: - if f.is_file(): + if f.is_file() and f.suffix.lower() in (".mp3", ".wav"): files.append(f) total_size += f.stat().st_size except OSError: From 9297bed5b9574ea5ba13039821a51d77a2afaccd Mon Sep 17 00:00:00 2001 From: Boody Date: Wed, 29 Jul 2026 12:54:55 +0300 Subject: [PATCH 11/32] add ODYSSEUS_TTS_CACHE_MAX_BYTES environment variable to docker-compose --- docker-compose.yml | 1 + 1 file changed, 1 insertion(+) 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:-} From 3250a4ce68ffb05e6d5023ca55c6c2189d548686 Mon Sep 17 00:00:00 2001 From: RaresKeY <158580472+RaresKeY@users.noreply.github.com> Date: Wed, 29 Jul 2026 22:04:28 +0100 Subject: [PATCH 12/32] fix(ci): clear review label when issues close (#5813) The issue-close lifecycle change is narrowly scoped and correct. Closed issues remove the stale \`ready for review\` label and return before normal validation can restore it. Focused regressions cover closure and subsequent edits to a closed issue. The branch was updated onto current \`dev\`. The focused test, merged-result validation, diff checks, and GitHub CI passed. No blocking review threads remain. --- .github/scripts/check-issue-description.js | 13 ++- .github/workflows/issue-description-check.yml | 2 +- tests/test_issue_description_check.py | 86 +++++++++++++++++++ 3 files changed, 97 insertions(+), 4 deletions(-) create mode 100644 tests/test_issue_description_check.py 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/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", + }, + } + ] From 6a84398e75835e899d792953e0b4e55ac40d9ccf Mon Sep 17 00:00:00 2001 From: holden093 Date: Thu, 30 Jul 2026 10:06:31 +0200 Subject: [PATCH 13/32] fix(skills): use utility model for skill tests instead of chat default (#5746) Skill tests are background automation tasks (like auto-naming and memory audit) and should use the configured utility model. Previously they resolved via resolve_endpoint("default") which returned the chat model, bypassing the utility model entirely. This completes the sweep started in PR #4027 which fixed auto-naming and memory audit but missed skill tests. --- routes/skills_routes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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) From f23221420fea7df82d9ca532cf6b16a4868b1136 Mon Sep 17 00:00:00 2001 From: Husam Date: Thu, 30 Jul 2026 11:54:59 +0300 Subject: [PATCH 14/32] fix(skills): replace deprecated utcnow in skill timestamp helper (#5777) * fix(skills): replace deprecated utcnow in skill timestamp helper _now_iso() builds the 'created' value in skill frontmatter. datetime.utcnow() returns a naive datetime and has been deprecated since Python 3.12, scheduled for removal. Switch to the timezone-aware datetime.now(timezone.utc), keeping the serialized YYYY-MM-DDTHH:MM:SSZ shape unchanged so existing skill files keep parsing. timezone.utc is used rather than the datetime.UTC alias, which is 3.11+ only. Adds regression tests covering the deprecation, the serialized shape, and UTC correctness under a non-UTC local timezone -- the last guards against a bare datetime.now(), which yields the same shape but local wall time. Fixes #5697 * test(skills): skip timezone mutation where unsupported --------- Co-authored-by: Alexandre Teixeira --- services/memory/skill_format.py | 4 +- tests/test_skill_format_timestamp.py | 58 ++++++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 2 deletions(-) create mode 100644 tests/test_skill_format_timestamp.py 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/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() From 578312200ac1828bc3bae039eac3c827036fdc69 Mon Sep 17 00:00:00 2001 From: Husam Date: Thu, 30 Jul 2026 12:48:31 +0300 Subject: [PATCH 15/32] fix(markdown): restore extracted blocks verbatim so $& and $$ survive (#5768) The placeholder-restore pass in mdToHtml put code, math, mermaid and allowed-HTML blocks back with a string replacement, so String.replace read `$&`, `` $` ``, `$'` and `$$` in the *replacement* as substitution patterns. A fenced block containing them rendered corrupted: `$&` re-inserted the placeholder (`perl -pe 's/world/$& again/'` became `s/world/___CODE_BLOCK_0___amp; again/`), `` $` `` and `$'` spliced in the surrounding document, and `$$` collapsed to a single `$`. Pass a function replacer at all four sites, matching the inline-code site below them, which was already fixed this way. A function's return value is inserted verbatim with no `$` interpretation. The inline-code comment claimed `echo $1` would be read as a back-reference; with a string search value there are no capture groups, so `$1` is already literal. Reworded to name the four sequences that do corrupt. Fixes #5663 --- static/js/markdown.js | 20 ++++++++----- tests/test_markdown_rendering_js.py | 44 +++++++++++++++++++++++++++++ 2 files changed, 57 insertions(+), 7 deletions(-) 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/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 From 84709a00d979c396dfb34f5d60b2992efb5a9284 Mon Sep 17 00:00:00 2001 From: Husam Date: Thu, 30 Jul 2026 13:30:00 +0300 Subject: [PATCH 16/32] fix(llm): omit temperature for major-only Opus ids (claude-opus-5) (#5761) The version pattern in _anthropic_rejects_temperature() required a minor component, so major-only ids like `claude-opus-5` never matched and the guard reported that the model accepts `temperature`. Anthropic rejects the field outright on Opus 4.7+, so every such call returned HTTP 400 and the stream aborted with zero tokens ("the model returned an empty response"). Make the minor optional and read a missing minor as `.0`. The major is also capped at 1-2 digits with a no-trailing-digit lookahead, mirroring the minor: once the minor is optional, a greedy major would swallow the date in `claude-3-opus-20240229` and read it as version 20240229, dropping temperature from a model that accepts it. Fixes #5753 Co-authored-by: Alexandre Teixeira <111787685+alteixeira20@users.noreply.github.com> --- src/llm_core.py | 26 ++++++++++++++++------ tests/test_llm_core_anthropic_temp_omit.py | 26 +++++++++++++++++++++- 2 files changed, 44 insertions(+), 8 deletions(-) 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/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 From 28c333e64780ae2fbb7f0757d752e8bb5fb74990 Mon Sep 17 00:00:00 2001 From: RaresKeY <158580472+RaresKeY@users.noreply.github.com> Date: Thu, 30 Jul 2026 12:24:39 +0100 Subject: [PATCH 17/32] fix(email): preserve OAuth SMTP security (#5802) --- static/js/settings.js | 2 ++ tests/test_email_oauth_connect_smtp_security.py | 15 +++++++++++++++ 2 files changed, 17 insertions(+) create mode 100644 tests/test_email_oauth_connect_smtp_security.py diff --git a/static/js/settings.js b/static/js/settings.js index 72936adee..3819c83fa 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; } 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 From 25c9e735ef5ce605f47f8f666ac6689056d2c10c Mon Sep 17 00:00:00 2001 From: RaresKeY <158580472+RaresKeY@users.noreply.github.com> Date: Thu, 30 Jul 2026 14:57:07 +0100 Subject: [PATCH 18/32] fix(email): open settings after OAuth callback (#5803) --- static/js/settings.js | 45 +++++++++++---------- tests/test_email_oauth_settings_redirect.py | 19 +++++++++ 2 files changed, 42 insertions(+), 22 deletions(-) create mode 100644 tests/test_email_oauth_settings_redirect.py diff --git a/static/js/settings.js b/static/js/settings.js index 3819c83fa..540acff00 100644 --- a/static/js/settings.js +++ b/static/js/settings.js @@ -5790,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/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 From 0de76c4056a0f0fb8f24b3bf2972fa5a185d7b0c Mon Sep 17 00:00:00 2001 From: "Tal.Yuan" Date: Tue, 4 Aug 2026 02:44:00 +0800 Subject: [PATCH 19/32] refactor(routes): move vault domain into routes/vault/ subpackage (#5780) Slice 2k of the route-domain reorganization (#4082/#4071). Moves vault_routes.py into routes/vault/, leaving a backward-compat sys.modules shim. Pure file reorganization, no behavior change. --- app.py | 2 +- routes/vault/__init__.py | 5 + routes/vault/vault_routes.py | 242 +++++++++++++++++++++++++++++++ routes/vault_routes.py | 246 ++------------------------------ tests/test_vault_routes_shim.py | 11 ++ 5 files changed, 268 insertions(+), 238 deletions(-) create mode 100644 routes/vault/__init__.py create mode 100644 routes/vault/vault_routes.py create mode 100644 tests/test_vault_routes_shim.py diff --git a/app.py b/app.py index 2ae5ec761..5fb2da54d 100644 --- a/app.py +++ b/app.py @@ -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/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/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 From fb8c391a8893be254a9ce4e2e954ea297c573554 Mon Sep 17 00:00:00 2001 From: "Tal.Yuan" Date: Tue, 4 Aug 2026 02:44:31 +0800 Subject: [PATCH 20/32] refactor(routes): move webhook domain into routes/webhook/ subpackage (#5781) Slice 2l of the route-domain reorganization (#4082/#4071). Moves webhook_routes.py into routes/webhook/, leaving a backward-compat sys.modules shim. Pure file reorganization, no behavior change. One source-introspection test repointed (test_api_chat_security.py). --- app.py | 2 +- routes/webhook/__init__.py | 5 + routes/webhook/webhook_routes.py | 395 +++++++++++++++++++++++++++++ routes/webhook_routes.py | 403 +----------------------------- tests/test_api_chat_security.py | 2 +- tests/test_webhook_routes_shim.py | 11 + 6 files changed, 425 insertions(+), 393 deletions(-) create mode 100644 routes/webhook/__init__.py create mode 100644 routes/webhook/webhook_routes.py create mode 100644 tests/test_webhook_routes_shim.py diff --git a/app.py b/app.py index 5fb2da54d..c85d425fb 100644 --- a/app.py +++ b/app.py @@ -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 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/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_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 From bb719f217a77b19d89d26f96168cf463ab73b6ba Mon Sep 17 00:00:00 2001 From: "Tal.Yuan" Date: Tue, 4 Aug 2026 17:54:55 +0800 Subject: [PATCH 21/32] refactor(routes): move document domain into routes/document/ subpackage (#5885) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Slice 2m of the route-domain reorganization (#4082/#4071, per specs/architecture-runtime-inventory.md §6.3). Moves document_routes.py (1810 lines) and document_helpers.py (243 lines) into routes/document/, leaving backward-compat sys.modules shims at the old paths. Pure file reorganization, no behavior change. Both shims use sys.modules replacement so the `import ... as droutes` + `droutes.SessionLocal = ...` / `monkeypatch.setattr(droutes, ...)` pattern in multiple tests, and the `sys.modules.pop("routes.document_helpers")` + re-import pattern in test_security_regressions.py, all reach the canonical modules. The canonical document_routes.py imports helpers from the canonical path (routes.document.document_helpers), not the legacy shim. Three source-introspection test sites repointed to the new canonical path: - test_imap_mailbox_quoting.py - test_model_helper_owner_scope.py - test_vision_owner_scope.py (shared with other domains; document entry repointed) Adds tests/test_document_routes_shim.py to pin the sys.modules shim contract for both modules. Verified: compileall clean; full suite 4789 passed, 3 skipped. --- app.py | 2 +- routes/document/__init__.py | 6 + routes/document/document_helpers.py | 243 ++++ routes/document/document_routes.py | 1810 +++++++++++++++++++++++ routes/document_helpers.py | 249 +--- routes/document_routes.py | 1819 +----------------------- tests/test_document_routes_shim.py | 29 + tests/test_imap_mailbox_quoting.py | 2 +- tests/test_model_helper_owner_scope.py | 2 +- tests/test_vision_owner_scope.py | 2 +- 10 files changed, 2115 insertions(+), 2049 deletions(-) create mode 100644 routes/document/__init__.py create mode 100644 routes/document/document_helpers.py create mode 100644 routes/document/document_routes.py create mode 100644 tests/test_document_routes_shim.py diff --git a/app.py b/app.py index c85d425fb..8363ba4e9 100644 --- a/app.py +++ b/app.py @@ -739,7 +739,7 @@ app.include_router(setup_stt_routes(stt_service)) logger.info("STT service initialized (provider managed via settings)") # Documents (artifacts/canvas) -from routes.document_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) 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/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_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_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_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() From 9d686180dd20e6ef842f0c23f9ef2f0ce39cee4f Mon Sep 17 00:00:00 2001 From: Ashvin <76151462+ashvinctrl@users.noreply.github.com> Date: Tue, 4 Aug 2026 15:47:41 +0530 Subject: [PATCH 22/32] fix(integrations): pin api_call to the SSRF-validated IP (#5727) * fix(integrations): pin api_call to the SSRF-validated IP execute_api_call runs check_outbound_url on the target, but that guard only resolves the host to answer (ok, reason) and hands back no address. The request right after it opened a plain httpx.AsyncClient, which resolves the host again at connect time. A base_url host on a low TTL can pass the guard as a public IP and then flip to 169.254.169.254 for the connect, so the call lands on cloud metadata with the integration's stored auth headers attached. Resolve once, remember the IPs the guard actually validated, and pin the client's socket to that set through a small AnyIO-backed transport. SNI and the Host header still come from the URL, so TLS and vhost routing are unchanged; connect-time fallback stays inside the approved address set over one shared deadline. This is the same pinning the webhook sender and web-fetch paths already do -- api_call was the last outbound path that skipped it. Fixes #5513 * fix(integrations): de-duplicate the pinned IP list _default_resolver calls getaddrinfo(host, None) with no socktype filter, so glibc returns one record per socktype and a single-homed host comes back three times over. _validated_ips kept every entry, so the transport pinned the same address repeatedly and the connect fallback could spend its shared deadline retrying one dead address instead of moving on to a genuinely different one. Windows getaddrinfo collapses those duplicate records, which is why the ip-literal pin test only failed on CI and not locally. --- src/integrations.py | 175 ++++++++++++- tests/test_integration_api_call_ssrf.py | 240 ++++++++++++++++++ .../test_integrations_api_call_truncation.py | 14 +- 3 files changed, 420 insertions(+), 9 deletions(-) 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/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 From 20e7fc0164286e1521569d9edc17a4ae4d0d2e22 Mon Sep 17 00:00:00 2001 From: adabarbulescu <94562950+adabarbulescu@users.noreply.github.com> Date: Tue, 4 Aug 2026 13:17:45 +0300 Subject: [PATCH 23/32] fix(skills): require manage_skills action (#5856) --- src/tools/system.py | 6 ++++-- tests/test_manage_skills_action_required.py | 24 +++++++++++++++++++++ 2 files changed, 28 insertions(+), 2 deletions(-) create mode 100644 tests/test_manage_skills_action_required.py 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/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, + } From c8a012d4d2db27142196a9a7b290687323979b5b Mon Sep 17 00:00:00 2001 From: Ashvin <76151462+ashvinctrl@users.noreply.github.com> Date: Thu, 6 Aug 2026 14:03:50 +0530 Subject: [PATCH 24/32] fix(memory): don't let an unreadable store get overwritten with an empty one (#5831) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(memory): don't let an unreadable store get overwritten with an empty one load_all() answered a failed read the same way it answered an empty store: with []. Every mutation path is a read-modify-write (load the whole file, change it, save it back), so a failed read became load_all() -> [] -> [].append(new) -> save([new]) and save() is atomic, so the replacement stuck. The case that actually destroys data is a store that is READABLE but not parseable - a truncated file, or one holding {} instead of []. Nothing obstructs the write, so adding a memory returns HTTP 200 and every memory already stored is gone. Verified end-to-end against a running instance: on the current code a truncated memory.json plus one add leaves the file holding only the new entry. 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. That path currently 500s and loses nothing. _read_entries() now returns [] only when the file genuinely does not exist and raises MemoryStoreUnreadable for every other failure, including a store that parses but is not a JSON array. load_all() keeps the old lenient behaviour so display, search and context injection still degrade quietly instead of breaking chat. The read-modify-write callers switch to load_all_for_update(), which propagates the error: the memory routes turn it into a 503 and change nothing, backup import refuses rather than saving only the incoming rows, and auto-extraction and the audit merge skip the write. The audit merge mattered most - it rebuilds the whole file from one owner's slice plus everyone else's rows, so an empty read there dropped every other tenant's memories. The corrupt-JSON path still gets its one shot at the legacy memory.txt migration before raising, so that recovery is unchanged. The two updated fakes gained load_all_for_update because the real class has it; MagicMock would otherwise hand the import path a Mock instead of the seeded list. Fixes #5673 * fix(memory): fail closed on the remaining read-modify-write add paths The strict loader landed with the routes, the backup import and the extractor converted, but three read-modify-write sinks still called load_all(), which degrades an unreadable store to []. Two of them are the paths users actually reach, so the data loss in #5673 stayed reproducible: - src/ai_interaction.py do_manage_memory, action "add" — reached from ordinary chat via src/tool_execution.py:793 -> dispatch_ai_tool. "Remember that I prefer X" against an unreadable store wrote a one-entry file over it and reported success. - mcp_servers/memory_server.py, action "add" — the same shape through _scope_entries(), registered as a built-in in src/builtin_mcp.py. - src/memory_provider.py NativeMemoryProvider.remember and .delete — wired into app state in src/app_initializer.py but not consumed outside tests yet, converted here so the pattern is uniform before it goes live. The MCP server takes _scope_entries(for_update=True) so list keeps the lenient read. The edit and delete branches on both tool paths were already fail-closed by accident — an empty view matches nothing and returns before the save — so they are left alone. The three new tests drive the real entry points rather than replaying the shape, and use a truncated store, which is the case that reads back fine so nothing stops the save. Each asserts memory.json is byte-identical afterwards; all three fail on the previous commit with the store overwritten. --- mcp_servers/memory_server.py | 26 +- routes/backup_routes.py | 11 +- routes/memory/memory_routes.py | 26 +- services/memory/__init__.py | 3 +- services/memory/memory.py | 14 +- services/memory/memory_extractor.py | 23 +- src/ai_interaction.py | 11 +- src/memory.py | 90 ++++++- src/memory_provider.py | 11 +- tests/test_backup_import_cross_user_dedup.py | 3 + ...st_memory_extractor_vector_cross_tenant.py | 6 + tests/test_memory_store_unreadable_no_wipe.py | 255 ++++++++++++++++++ 12 files changed, 451 insertions(+), 28 deletions(-) create mode 100644 tests/test_memory_store_unreadable_no_wipe.py 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/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/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/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/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/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/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_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"] From 5ddef23d949b0ebb5ed90d6ae4904798d1619bca Mon Sep 17 00:00:00 2001 From: adabarbulescu <94562950+adabarbulescu@users.noreply.github.com> Date: Fri, 7 Aug 2026 20:12:21 +0300 Subject: [PATCH 25/32] fix(welcome): rotate startup tips (#5871) --- static/index.html | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/static/index.html b/static/index.html index 8257660fe..0136f2316 100644 --- a/static/index.html +++ b/static/index.html @@ -1005,7 +1005,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; From 36d409842177e18017dac2fa4bbc5266bb451ac7 Mon Sep 17 00:00:00 2001 From: Jakub Grula Date: Fri, 7 Aug 2026 19:15:50 +0200 Subject: [PATCH 26/32] fix: Edit box formatting was removing triple tick boxes (#5737) --- static/js/chat.js | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/static/js/chat.js b/static/js/chat.js index ea2d8c1bb..ca583c5fc 100644 --- a/static/js/chat.js +++ b/static/js/chat.js @@ -4787,7 +4787,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'); From f1e96d102e5692fca3a91b1a40f46e22028e7ecb Mon Sep 17 00:00:00 2001 From: Husam Date: Fri, 7 Aug 2026 20:33:14 +0300 Subject: [PATCH 27/32] fix(tool_parsing): require a pipe on the Qwen bare end marker (#5829) The `end` branch of _QWEN_BARE_MARKER_RE had both pipes optional (`\|?end\|?`), so it also matched a bare `end` between whitespace and replaced it with a space. Messages containing Ruby, Lua or shell code that closes a block with a lone `end` had those lines deleted, and ordinary prose lost the word too. Require at least one pipe so only real turn markers match; `|end`, `end|`, `|end|` and `/|end|` strip exactly as before. Applied to the duplicated pattern in static/js/chatRenderer.js as well. Fixes #5547 --- src/tool_parsing.py | 6 +- static/js/chatRenderer.js | 5 +- tests/test_tool_parsing_bare_end_marker.py | 96 ++++++++++++++++++++++ 3 files changed, 105 insertions(+), 2 deletions(-) create mode 100644 tests/test_tool_parsing_bare_end_marker.py 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/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/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}" From 99566d28b53cdd53efaecff05a61ef3b7b5daa06 Mon Sep 17 00:00:00 2001 From: Husam Date: Fri, 7 Aug 2026 20:34:50 +0300 Subject: [PATCH 28/32] fix(chat): stop ArrowUp from eating an unsent multi-line prompt (#5875) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit static/app.js carried a near-verbatim copy of the prompt-recall logic in static/js/composerArrowUpRecall.js, wired as a second capture-phase keydown listener on the same #message textarea. The copy omitted the draft guard the module has: it called preventDefault() and stopImmediatePropagation() unconditionally, then recalled history[0] over whatever the user had typed. Because it stopped immediate propagation, the copy won regardless of registration order — if it ran first the module never saw the event, and if it ran second the module had already declined to stop propagation on an unmatched draft. The guard at composerArrowUpRecall.js:109 was unreachable on the real page, so ArrowUp on a multi-line draft replaced it with the last sent prompt instead of moving the caret up a line. Delete the duplicate. The module keeps ownership of ArrowUp/ArrowDown recall, which is the behavior MODULE_SUMMARY.md documents ("on an empty composer") and the behavior tests/test_composer_arrow_up_recall_js.py already pins via test_non_empty_composer_does_not_recall and test_multiline_caret_navigation_preserved. Also correct a stale comment in the module that described the deleted behavior and contradicted the guard 35 lines above it, and add a regression test asserting app.js does not reintroduce a second handler. Fixes #5862 --- static/app.js | 83 ++--------------------- static/js/composerArrowUpRecall.js | 6 +- tests/test_composer_arrow_up_recall_js.py | 21 ++++++ 3 files changed, 28 insertions(+), 82 deletions(-) diff --git a/static/app.js b/static/app.js index 97f0ae77e..c9e3a567f 100644 --- a/static/app.js +++ b/static/app.js @@ -3908,85 +3908,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/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/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" + ) From f06a0a30a80e6739a7f2d7f9b23ed38a8ffb21fd Mon Sep 17 00:00:00 2001 From: Samy <12219635+touzenesmy@users.noreply.github.com> Date: Fri, 7 Aug 2026 16:04:53 -0400 Subject: [PATCH 29/32] fix(session): restore session URL hash writes (removed in cf4e240a) (#5872) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Fix: restore session URL hash writes (removed in cf4e240a) Restores history.replaceState() calls in selectSession() and materializePendingSession() that were dropped during the July 23 merge. Without these, chat URLs never update the address bar hash, making sessions unshareable and causing bare-URL reloads to land on the welcome screen instead of restoring the last active chat. Root cause: selectSession() had its hash-write deliberately removed; materializePendingSession() lost its during a larger refactor that added the stale-response and incognito guards. Fixes #5870 (upstream) * fix: session URL hash lost when sending message mid-stream Two independent bugs caused the session hash to disappear from the URL: Bug 1 — ReferenceError in catch block silently killed error recovery In handleChatSubmit, two const variables (streamingTTS at line 1922 and abortCtrl at line 1741) were declared inside the try block but referenced in the catch block. Since const is block-scoped in JavaScript, they were undefined in catch, causing a ReferenceError that silently aborted the error handler. This prevented materializePendingSession() from ever being called, so no hash was written to the URL. Fix: Hoisted both as let declarations before the try { block. Bug 2 — Dual sessions.js ES module instances with mismatched state app.js imported sessions.js with a version query string (?v=20260722ctxheader4) while every other module imported ./sessions.js without one. The browser treated them as different URLs, creating two separate module instances with independent _pendingChat and currentSessionId state. createDirectChat() set pending on one instance while handleChatSubmit() checked hasPendingChat() on the other — so the pending session never materialized. Fix: Removed the version query string from the sessions.js import in app.js and from the modulepreload + script tags in index.html. All modules now share a single sessions.js instance. Bonus guard: _adoptOpenedSessionBeforeAutoCreate() now checks hasPendingChat() before adopting a stale DOM-active session, preventing the send path from landing in the wrong session when a New Chat is pending. --------- Co-authored-by: samy --- static/app.js | 4 ++-- static/index.html | 8 ++++---- static/js/chat.js | 9 +++++++-- static/js/sessions.js | 5 +++++ 4 files changed, 18 insertions(+), 8 deletions(-) diff --git a/static/app.js b/static/app.js index c9e3a567f..9d05d8991 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'; diff --git a/static/index.html b/static/index.html index 0136f2316..0fea836ba 100644 --- a/static/index.html +++ b/static/index.html @@ -250,9 +250,9 @@ - + - + @@ -2504,7 +2504,7 @@ - + @@ -2522,7 +2522,7 @@ - + diff --git a/static/js/chat.js b/static/js/chat.js index ca583c5fc..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) 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 From 378518f6dfb994481a8ec5fbb0032d4a8c23c4e1 Mon Sep 17 00:00:00 2001 From: Samy <12219635+touzenesmy@users.noreply.github.com> Date: Fri, 7 Aug 2026 16:06:17 -0400 Subject: [PATCH 30/32] Fix #5870: stale skills panel data on tab reopen (#5876) Remove early-return guard in loadSkills() that skipped both API re-fetch and renderSkillsList() when the Skills tab was reopened after first load. The cascade entrance animation is already handled inside renderSkillsList() via _cascadeNext, so the guard was unnecessary and caused deleted/edited skills to remain visible until a full page reload. Co-authored-by: samy --- static/js/skills.js | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) 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 { From e4fa4ae5dd1d709ce4168397bd1d200fec1b2494 Mon Sep 17 00:00:00 2001 From: Wes Huber Date: Fri, 7 Aug 2026 13:07:07 -0700 Subject: [PATCH 31/32] fix(brain): give the Add Memory form a submit button and reliable Enter handling (#5830) The Brain > Add tab rendered only a text input and category select with no submit control, and Enter submission relied on a deprecated keypress listener that is not guaranteed to fire, so the form could not be submitted at all (#5828). Add a labelled submit button styled like the neighbouring Skill Import button (theme-io-btn, inline SVG icon), switch the Enter handler to keydown with preventDefault, ignore IME composition, and pin both submit paths with a source-level regression test. Fixes #5828 Co-authored-by: Claude Fable 5 --- static/app.js | 12 ++++- static/index.html | 1 + tests/test_memory_add_submit_regression.py | 54 ++++++++++++++++++++++ 3 files changed, 65 insertions(+), 2 deletions(-) create mode 100644 tests/test_memory_add_submit_regression.py diff --git a/static/app.js b/static/app.js index 9d05d8991..2f1e8d4bf 100644 --- a/static/app.js +++ b/static/app.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) diff --git a/static/index.html b/static/index.html index 0fea836ba..fea4e20ac 100644 --- a/static/index.html +++ b/static/index.html @@ -365,6 +365,7 @@ Add a memory — e.g. 'I prefer concise replies' +

", 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 From 42da399b4d9d814400b8e19365dffedad5b66ab3 Mon Sep 17 00:00:00 2001 From: Matyas Gosztonyi Date: Sat, 8 Aug 2026 23:06:41 +0200 Subject: [PATCH 32/32] fix(email): route summaries through shared LLM adapter (#5841) * fix(email): route summaries through shared llm adapter * chore(ci): refresh PR checks * fix(email): preserve scheduled summary safeguards --------- Co-authored-by: Matyas Fenyves <16389204+uhhgoat@users.noreply.github.com> --- routes/email_helpers.py | 120 +++++++ routes/email_pollers.py | 32 +- routes/email_routes.py | 94 +++--- static/js/emailLibrary.js | 7 +- static/js/emailLibrary/utils.js | 19 ++ tests/test_email_summary_error_ui_js.py | 52 +++ tests/test_email_summary_llm.py | 406 ++++++++++++++++++++++++ 7 files changed, 670 insertions(+), 60 deletions(-) create mode 100644 tests/test_email_summary_error_ui_js.py create mode 100644 tests/test_email_summary_llm.py 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/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/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")