mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-10-11 17:32:20 +02:00
Merge verified Odysseus fixes
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
"""Admin wipe route domain package (slice 2h, #4082/#4071).
|
||||
|
||||
Contains admin_wipe_routes.py, migrated from the flat routes/ directory.
|
||||
Backward-compat shim at routes/admin_wipe_routes.py re-exports from here.
|
||||
"""
|
||||
@@ -0,0 +1,176 @@
|
||||
"""Admin Danger Zone — per-category wipes.
|
||||
|
||||
Each endpoint is admin-only and truncates exactly one domain so the
|
||||
user can selectively reset memory / skills / notes / etc. without
|
||||
nuking everything. The catch-all `chats` endpoint mirrors the
|
||||
existing /api/sessions/all so the Danger Zone speaks one URL pattern.
|
||||
|
||||
URL shape: DELETE /api/admin/wipe/{kind}
|
||||
Kinds: chats, memory, skills, notes, tasks, documents, gallery, calendar.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
|
||||
from core.middleware import require_admin
|
||||
from core.database import (
|
||||
SessionLocal,
|
||||
Session as DbSession,
|
||||
ChatMessage as DbChatMessage,
|
||||
Memory,
|
||||
Note,
|
||||
ScheduledTask,
|
||||
TaskRun,
|
||||
Document,
|
||||
DocumentVersion,
|
||||
GalleryImage,
|
||||
GalleryAlbum,
|
||||
CalendarEvent,
|
||||
CalendarCal,
|
||||
)
|
||||
from src.constants import DATA_DIR, SKILLS_DIR, SKILLS_FILE, GALLERY_DIR, GALLERY_UPLOADS_DIR
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _wipe_memory_files():
|
||||
"""Blank memory.json + drop the per-owner tidy-state sidecar so the
|
||||
next audit doesn't try to diff against gone memories."""
|
||||
for name in ("memory.json", "memory_tidy_state.json"):
|
||||
p = os.path.join(DATA_DIR, name)
|
||||
if not os.path.exists(p):
|
||||
continue
|
||||
try:
|
||||
if name == "memory.json":
|
||||
with open(p, "w", encoding="utf-8") as f:
|
||||
json.dump([], f)
|
||||
else:
|
||||
os.remove(p)
|
||||
except OSError as e:
|
||||
logger.warning(f"Could not reset {name}: {e}")
|
||||
|
||||
|
||||
def _rmtree_quiet(path: str):
|
||||
"""rmtree that doesn't crash if the path doesn't exist."""
|
||||
if os.path.isdir(path):
|
||||
try:
|
||||
shutil.rmtree(path)
|
||||
except OSError as e:
|
||||
logger.warning(f"Could not remove {path}: {e}")
|
||||
|
||||
|
||||
def setup_admin_wipe_routes(session_manager):
|
||||
"""The session_manager is passed in so we can also clear its
|
||||
in-memory cache when wiping chats — without it the DB is empty
|
||||
but the next /api/sessions returns stale entries."""
|
||||
router = APIRouter(prefix="/api/admin")
|
||||
|
||||
@router.delete("/wipe/{kind}")
|
||||
def wipe(kind: str, request: Request):
|
||||
require_admin(request)
|
||||
kind = (kind or "").strip().lower()
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
if kind == "chats":
|
||||
count = db.query(DbSession).count()
|
||||
db.query(DbChatMessage).delete()
|
||||
db.query(DbSession).delete()
|
||||
db.commit()
|
||||
try:
|
||||
session_manager.sessions.clear()
|
||||
except Exception:
|
||||
pass
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "memory":
|
||||
count = db.query(Memory).count()
|
||||
db.query(Memory).delete()
|
||||
db.commit()
|
||||
_wipe_memory_files()
|
||||
# Drop the vector store too so semantic search doesn't
|
||||
# return ghosts. Lazy import — chromadb may not be
|
||||
# initialised in every deployment.
|
||||
try:
|
||||
from src.memory_vector import get_memory_vector_store
|
||||
mv = get_memory_vector_store()
|
||||
if mv and hasattr(mv, "clear"):
|
||||
mv.clear()
|
||||
except Exception as e:
|
||||
logger.info(f"Memory vector clear skipped: {e}")
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "skills":
|
||||
# Skills live as SKILL.md files under data/skills/. Drop
|
||||
# the entire directory; the SkillsManager re-creates the
|
||||
# tree on next write.
|
||||
skills_dir = SKILLS_DIR
|
||||
count = 0
|
||||
if os.path.isdir(skills_dir):
|
||||
# Count SKILL.md files for the response — quick walk.
|
||||
for _, _, files in os.walk(skills_dir):
|
||||
count += sum(1 for f in files if f == "SKILL.md")
|
||||
_rmtree_quiet(skills_dir)
|
||||
# Legacy fallback file
|
||||
legacy = SKILLS_FILE
|
||||
if os.path.exists(legacy):
|
||||
try:
|
||||
os.remove(legacy)
|
||||
except OSError:
|
||||
pass
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "notes":
|
||||
count = db.query(Note).count()
|
||||
db.query(Note).delete()
|
||||
db.commit()
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "tasks":
|
||||
# TaskRun rows reference tasks via FK — clear them first.
|
||||
db.query(TaskRun).delete()
|
||||
count = db.query(ScheduledTask).count()
|
||||
db.query(ScheduledTask).delete()
|
||||
db.commit()
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "documents":
|
||||
# DocumentVersion FKs Document — clear children first.
|
||||
db.query(DocumentVersion).delete()
|
||||
count = db.query(Document).count()
|
||||
db.query(Document).delete()
|
||||
db.commit()
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "gallery":
|
||||
count = db.query(GalleryImage).count() + db.query(GalleryAlbum).count()
|
||||
db.query(GalleryImage).delete()
|
||||
db.query(GalleryAlbum).delete()
|
||||
db.commit()
|
||||
# Also drop the upload dir so disk doesn't keep orphans.
|
||||
_rmtree_quiet(GALLERY_DIR)
|
||||
_rmtree_quiet(GALLERY_UPLOADS_DIR)
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "calendar":
|
||||
# Events FK calendars — clear children first, then both.
|
||||
db.query(CalendarEvent).delete()
|
||||
count = db.query(CalendarCal).count()
|
||||
db.query(CalendarCal).delete()
|
||||
db.commit()
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
raise HTTPException(400, f"Unknown wipe kind: {kind!r}")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.exception(f"Wipe {kind} failed")
|
||||
raise HTTPException(500, f"Wipe {kind} failed: {e}")
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return router
|
||||
+12
-171
@@ -1,176 +1,17 @@
|
||||
"""Admin Danger Zone — per-category wipes.
|
||||
"""Backward-compat shim — canonical location is routes/admin_wipe/admin_wipe_routes.py.
|
||||
|
||||
Each endpoint is admin-only and truncates exactly one domain so the
|
||||
user can selectively reset memory / skills / notes / etc. without
|
||||
nuking everything. The catch-all `chats` endpoint mirrors the
|
||||
existing /api/sessions/all so the Danger Zone speaks one URL pattern.
|
||||
|
||||
URL shape: DELETE /api/admin/wipe/{kind}
|
||||
Kinds: chats, memory, skills, notes, tasks, documents, gallery, calendar.
|
||||
This module is replaced in ``sys.modules`` by the canonical module object so
|
||||
that ``import routes.admin_wipe_routes``, ``from routes.admin_wipe_routes
|
||||
import X``, ``importlib.import_module("routes.admin_wipe_routes")``, and the
|
||||
``import ... as admin_wipe_routes`` + ``monkeypatch.setattr(admin_wipe_routes,
|
||||
"SessionLocal", ...)`` / ``"require_admin"`` pattern used by
|
||||
test_admin_wipe_gallery.py all operate on the *same* object the application
|
||||
actually uses. Keeps existing import paths working after slice 2h
|
||||
(#4082/#4071).
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
import sys as _sys
|
||||
|
||||
from core.middleware import require_admin
|
||||
from core.database import (
|
||||
SessionLocal,
|
||||
Session as DbSession,
|
||||
ChatMessage as DbChatMessage,
|
||||
Memory,
|
||||
Note,
|
||||
ScheduledTask,
|
||||
TaskRun,
|
||||
Document,
|
||||
DocumentVersion,
|
||||
GalleryImage,
|
||||
GalleryAlbum,
|
||||
CalendarEvent,
|
||||
CalendarCal,
|
||||
)
|
||||
from src.constants import DATA_DIR, SKILLS_DIR, SKILLS_FILE, GALLERY_DIR, GALLERY_UPLOADS_DIR
|
||||
from routes.admin_wipe import admin_wipe_routes as _canonical # noqa: F401
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _wipe_memory_files():
|
||||
"""Blank memory.json + drop the per-owner tidy-state sidecar so the
|
||||
next audit doesn't try to diff against gone memories."""
|
||||
for name in ("memory.json", "memory_tidy_state.json"):
|
||||
p = os.path.join(DATA_DIR, name)
|
||||
if not os.path.exists(p):
|
||||
continue
|
||||
try:
|
||||
if name == "memory.json":
|
||||
with open(p, "w", encoding="utf-8") as f:
|
||||
json.dump([], f)
|
||||
else:
|
||||
os.remove(p)
|
||||
except OSError as e:
|
||||
logger.warning(f"Could not reset {name}: {e}")
|
||||
|
||||
|
||||
def _rmtree_quiet(path: str):
|
||||
"""rmtree that doesn't crash if the path doesn't exist."""
|
||||
if os.path.isdir(path):
|
||||
try:
|
||||
shutil.rmtree(path)
|
||||
except OSError as e:
|
||||
logger.warning(f"Could not remove {path}: {e}")
|
||||
|
||||
|
||||
def setup_admin_wipe_routes(session_manager):
|
||||
"""The session_manager is passed in so we can also clear its
|
||||
in-memory cache when wiping chats — without it the DB is empty
|
||||
but the next /api/sessions returns stale entries."""
|
||||
router = APIRouter(prefix="/api/admin")
|
||||
|
||||
@router.delete("/wipe/{kind}")
|
||||
def wipe(kind: str, request: Request):
|
||||
require_admin(request)
|
||||
kind = (kind or "").strip().lower()
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
if kind == "chats":
|
||||
count = db.query(DbSession).count()
|
||||
db.query(DbChatMessage).delete()
|
||||
db.query(DbSession).delete()
|
||||
db.commit()
|
||||
try:
|
||||
session_manager.sessions.clear()
|
||||
except Exception:
|
||||
pass
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "memory":
|
||||
count = db.query(Memory).count()
|
||||
db.query(Memory).delete()
|
||||
db.commit()
|
||||
_wipe_memory_files()
|
||||
# Drop the vector store too so semantic search doesn't
|
||||
# return ghosts. Lazy import — chromadb may not be
|
||||
# initialised in every deployment.
|
||||
try:
|
||||
from src.memory_vector import get_memory_vector_store
|
||||
mv = get_memory_vector_store()
|
||||
if mv and hasattr(mv, "clear"):
|
||||
mv.clear()
|
||||
except Exception as e:
|
||||
logger.info(f"Memory vector clear skipped: {e}")
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "skills":
|
||||
# Skills live as SKILL.md files under data/skills/. Drop
|
||||
# the entire directory; the SkillsManager re-creates the
|
||||
# tree on next write.
|
||||
skills_dir = SKILLS_DIR
|
||||
count = 0
|
||||
if os.path.isdir(skills_dir):
|
||||
# Count SKILL.md files for the response — quick walk.
|
||||
for _, _, files in os.walk(skills_dir):
|
||||
count += sum(1 for f in files if f == "SKILL.md")
|
||||
_rmtree_quiet(skills_dir)
|
||||
# Legacy fallback file
|
||||
legacy = SKILLS_FILE
|
||||
if os.path.exists(legacy):
|
||||
try:
|
||||
os.remove(legacy)
|
||||
except OSError:
|
||||
pass
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "notes":
|
||||
count = db.query(Note).count()
|
||||
db.query(Note).delete()
|
||||
db.commit()
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "tasks":
|
||||
# TaskRun rows reference tasks via FK — clear them first.
|
||||
db.query(TaskRun).delete()
|
||||
count = db.query(ScheduledTask).count()
|
||||
db.query(ScheduledTask).delete()
|
||||
db.commit()
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "documents":
|
||||
# DocumentVersion FKs Document — clear children first.
|
||||
db.query(DocumentVersion).delete()
|
||||
count = db.query(Document).count()
|
||||
db.query(Document).delete()
|
||||
db.commit()
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "gallery":
|
||||
count = db.query(GalleryImage).count() + db.query(GalleryAlbum).count()
|
||||
db.query(GalleryImage).delete()
|
||||
db.query(GalleryAlbum).delete()
|
||||
db.commit()
|
||||
# Also drop the upload dir so disk doesn't keep orphans.
|
||||
_rmtree_quiet(GALLERY_DIR)
|
||||
_rmtree_quiet(GALLERY_UPLOADS_DIR)
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
if kind == "calendar":
|
||||
# Events FK calendars — clear children first, then both.
|
||||
db.query(CalendarEvent).delete()
|
||||
count = db.query(CalendarCal).count()
|
||||
db.query(CalendarCal).delete()
|
||||
db.commit()
|
||||
return {"status": "deleted", "kind": kind, "count": count}
|
||||
|
||||
raise HTTPException(400, f"Unknown wipe kind: {kind!r}")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.exception(f"Wipe {kind} failed")
|
||||
raise HTTPException(500, f"Wipe {kind} failed: {e}")
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return router
|
||||
_sys.modules[__name__] = _canonical
|
||||
|
||||
@@ -13,8 +13,9 @@ from sqlalchemy import or_, and_
|
||||
from dateutil.rrule import rrulestr
|
||||
|
||||
from core.database import SessionLocal, CalendarCal, CalendarDeletedEvent, CalendarEvent
|
||||
from src.auth_helpers import require_user
|
||||
from src.auth_helpers import effective_user, require_user
|
||||
from src.upload_limits import read_upload_limited, ICS_MAX_BYTES
|
||||
from src.upload_handler import reserve_upload_references
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -697,9 +698,18 @@ def _expand_rrule(
|
||||
|
||||
# ── Routes ──
|
||||
|
||||
def setup_calendar_routes() -> APIRouter:
|
||||
def setup_calendar_routes(upload_handler=None) -> APIRouter:
|
||||
router = APIRouter(prefix="/api/calendar", tags=["calendar"])
|
||||
|
||||
def _reserve_calendar_uploads(request: Request, *values) -> None:
|
||||
missing_id = reserve_upload_references(
|
||||
upload_handler,
|
||||
effective_user(request),
|
||||
*values,
|
||||
)
|
||||
if missing_id:
|
||||
raise HTTPException(409, f"Referenced upload is no longer available: {missing_id}")
|
||||
|
||||
# ── CalDAV multi-account helpers ─────────────────────────────────────────
|
||||
|
||||
def _get_caldav_accounts(owner: str) -> list:
|
||||
@@ -913,7 +923,24 @@ def setup_calendar_routes() -> APIRouter:
|
||||
'</d:prop></d:propfind>'
|
||||
)
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=8.0, follow_redirects=False, trust_env=False) as cx:
|
||||
# Build an SSL context that trusts the operator's custom CA bundle
|
||||
# (SSL_CERT_FILE / REQUESTS_CA_BUNDLE) so self-signed CalDAV servers
|
||||
# pass the pre-flight the same way they pass the real sync.
|
||||
# trust_env=False is kept to block proxy/auth env leakage; the CA
|
||||
# bundle is loaded explicitly instead.
|
||||
import ssl as _ssl
|
||||
_ssl_ctx = _ssl.create_default_context()
|
||||
# Disable VERIFY_X509_STRICT so certs without a keyUsage extension
|
||||
# (common in self-signed setups) are accepted, matching the
|
||||
# requests/urllib3 behavior used by the CalDAV sync path.
|
||||
_ssl_ctx.verify_flags &= ~_ssl.VERIFY_X509_STRICT
|
||||
_ca_bundle = _os.environ.get("SSL_CERT_FILE") or _os.environ.get("REQUESTS_CA_BUNDLE")
|
||||
if _ca_bundle:
|
||||
if _os.path.isfile(_ca_bundle):
|
||||
_ssl_ctx.load_verify_locations(_ca_bundle)
|
||||
else:
|
||||
logger.warning("CalDAV test: CA bundle %s not found, using system CAs", _ca_bundle)
|
||||
async with httpx.AsyncClient(timeout=8.0, follow_redirects=False, trust_env=False, verify=_ssl_ctx) as cx:
|
||||
r = await cx.request(
|
||||
"PROPFIND", url,
|
||||
auth=(user, pw),
|
||||
@@ -1070,6 +1097,7 @@ def setup_calendar_routes() -> APIRouter:
|
||||
@router.post("/events")
|
||||
async def create_event(request: Request, data: EventCreate):
|
||||
owner = _require_user(request)
|
||||
_reserve_calendar_uploads(request, data.color, data.description, data.location)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
cal = None
|
||||
@@ -1131,6 +1159,7 @@ def setup_calendar_routes() -> APIRouter:
|
||||
@router.put("/events/{uid}")
|
||||
async def update_event(request: Request, uid: str, data: EventUpdate):
|
||||
owner = _require_user(request)
|
||||
_reserve_calendar_uploads(request, data.color, data.description, data.location)
|
||||
try:
|
||||
base_uid = _resolve_base_uid(uid)
|
||||
except ValueError as e:
|
||||
@@ -1224,6 +1253,7 @@ def setup_calendar_routes() -> APIRouter:
|
||||
@router.post("/calendars")
|
||||
async def create_calendar(request: Request, name: str = "Imported", color: str = "#5b8abf"):
|
||||
owner = _require_user(request)
|
||||
_reserve_calendar_uploads(request, color)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
cal = CalendarCal(
|
||||
@@ -1246,6 +1276,7 @@ def setup_calendar_routes() -> APIRouter:
|
||||
@router.put("/calendars/{cal_id}")
|
||||
async def update_calendar(request: Request, cal_id: str, name: str = None, color: str = None):
|
||||
owner = _require_user(request)
|
||||
_reserve_calendar_uploads(request, color)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
cal = _get_or_404_calendar(db, cal_id, owner)
|
||||
|
||||
+77
-22
@@ -5,6 +5,7 @@ import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Optional
|
||||
|
||||
@@ -17,6 +18,7 @@ from src.context_compactor import maybe_compact, trim_for_context
|
||||
from src.model_context import estimate_tokens
|
||||
from src.auth_helpers import effective_user
|
||||
from src.prompt_security import untrusted_context_message
|
||||
from src.attachment_refs import attachment_ref
|
||||
from routes.prefs_routes import _load_for_user as load_prefs_for_user
|
||||
|
||||
from fastapi import HTTPException
|
||||
@@ -55,6 +57,9 @@ def _is_casual_low_signal(text: str) -> bool:
|
||||
# the background work (extraction, auto-naming) silently never runs.
|
||||
# Mirrors WebhookManager._spawn_tracked from src/webhook_manager.py.
|
||||
_BG_TASKS: set[asyncio.Task] = set()
|
||||
_INCOGNITO_CONTEXTS: dict[str, dict[str, Any]] = {}
|
||||
_INCOGNITO_CONTEXT_TTL_SECONDS = 6 * 60 * 60
|
||||
_INCOGNITO_CONTEXT_MAX_MESSAGES = 80
|
||||
|
||||
|
||||
def _spawn_bg(coro) -> asyncio.Task:
|
||||
@@ -65,6 +70,40 @@ def _spawn_bg(coro) -> asyncio.Task:
|
||||
return task
|
||||
|
||||
|
||||
def _prune_incognito_contexts(now: float | None = None):
|
||||
now = now or time.time()
|
||||
stale = [
|
||||
sid for sid, bundle in _INCOGNITO_CONTEXTS.items()
|
||||
if now - float(bundle.get("updated_at") or 0) > _INCOGNITO_CONTEXT_TTL_SECONDS
|
||||
]
|
||||
for sid in stale:
|
||||
_INCOGNITO_CONTEXTS.pop(sid, None)
|
||||
|
||||
|
||||
def _incognito_messages(session_id: str) -> list[dict[str, Any]]:
|
||||
_prune_incognito_contexts()
|
||||
bundle = _INCOGNITO_CONTEXTS.get(str(session_id or ""))
|
||||
if not bundle:
|
||||
return []
|
||||
return [dict(m) for m in bundle.get("messages", []) if isinstance(m, dict)]
|
||||
|
||||
|
||||
def _append_incognito_message(session_id: str, role: str, content: Any, metadata: dict | None = None):
|
||||
sid = str(session_id or "").strip()
|
||||
if not sid:
|
||||
return
|
||||
_prune_incognito_contexts()
|
||||
bundle = _INCOGNITO_CONTEXTS.setdefault(sid, {"messages": [], "updated_at": time.time()})
|
||||
msg: dict[str, Any] = {"role": role, "content": content}
|
||||
if metadata:
|
||||
msg["metadata"] = dict(metadata)
|
||||
messages = bundle.setdefault("messages", [])
|
||||
messages.append(msg)
|
||||
if len(messages) > _INCOGNITO_CONTEXT_MAX_MESSAGES:
|
||||
del messages[:-_INCOGNITO_CONTEXT_MAX_MESSAGES]
|
||||
bundle["updated_at"] = time.time()
|
||||
|
||||
|
||||
# ── Data containers ────────────────────────────────────────────────────── #
|
||||
|
||||
@dataclass
|
||||
@@ -418,24 +457,28 @@ def build_uploaded_file_manifest(att_ids: list, upload_handler, owner: Optional[
|
||||
except Exception:
|
||||
path = None
|
||||
|
||||
manifest.append({
|
||||
"id": info.get("id") or str(att_id),
|
||||
"name": info.get("name") or info.get("original_name") or str(att_id),
|
||||
"mime": info.get("mime", ""),
|
||||
"size": info.get("size", 0),
|
||||
ref = attachment_ref({**info, "id": info.get("id") or str(att_id)})
|
||||
ref.update({
|
||||
"id": ref["attachment_id"],
|
||||
"uri": f"odysseus://attachment/{ref['attachment_id']}",
|
||||
"read_policy": "owner_checked_upload",
|
||||
# Transitional compatibility: existing built-in tools can still use
|
||||
# this path, but only after owner, upload-root, and tool-root checks.
|
||||
"path": path,
|
||||
})
|
||||
manifest.append(ref)
|
||||
return manifest
|
||||
|
||||
|
||||
def add_user_message(sess, chat_handler, preprocessed: PreprocessedMessage, incognito: bool = False):
|
||||
"""Add user message to session history and update session name.
|
||||
In incognito mode, still add to in-memory history (for conversation context)
|
||||
but skip session name update (which would persist)."""
|
||||
Incognito messages must not mutate persistent session history, even in
|
||||
memory, because a later normal turn can persist the same session object."""
|
||||
if incognito:
|
||||
return
|
||||
user_meta = {"attachments": preprocessed.attachment_meta} if preprocessed.attachment_meta else None
|
||||
sess.add_message(ChatMessage("user", preprocessed.user_content, metadata=user_meta))
|
||||
if not incognito:
|
||||
chat_handler.update_session_name_if_needed(sess, preprocessed.text_for_context)
|
||||
chat_handler.update_session_name_if_needed(sess, preprocessed.text_for_context)
|
||||
|
||||
|
||||
def fire_message_event(request, webhook_manager, session_id: str, sess, message: str, compare_mode: bool = False):
|
||||
@@ -664,8 +707,14 @@ async def build_chat_context(
|
||||
allow_tool_preprocessing=allow_tool_preprocessing,
|
||||
)
|
||||
|
||||
# Add user message to history
|
||||
add_user_message(sess, chat_handler, preprocessed, incognito=incognito)
|
||||
# Add user message to history. Nobody/incognito uses a request-local
|
||||
# transcript store instead of session history so stale saved chats cannot
|
||||
# bleed into context and the turn is not persisted.
|
||||
if incognito:
|
||||
user_meta = {"attachments": preprocessed.attachment_meta} if preprocessed.attachment_meta else None
|
||||
_append_incognito_message(session_id, "user", preprocessed.user_content, user_meta)
|
||||
else:
|
||||
add_user_message(sess, chat_handler, preprocessed, incognito=False)
|
||||
|
||||
# Fire events
|
||||
if not incognito:
|
||||
@@ -756,8 +805,10 @@ async def build_chat_context(
|
||||
if norm:
|
||||
sess.model = norm
|
||||
|
||||
# Build messages
|
||||
messages = preface + sess.get_context_messages()
|
||||
# Build messages. In Nobody/incognito mode, never read saved session
|
||||
# history: the session id may be a temporary wrapper or, in buggy clients, a
|
||||
# stale normal session id. Only the ephemeral incognito transcript is safe.
|
||||
messages = preface + (_incognito_messages(session_id) if incognito else sess.get_context_messages())
|
||||
|
||||
# Current date/time — injected as a standalone *user*-role context message
|
||||
# placed immediately before the latest user turn, NOT folded into the
|
||||
@@ -1023,7 +1074,12 @@ def save_assistant_response(
|
||||
tool_events: list = None,
|
||||
incognito: bool = False,
|
||||
):
|
||||
"""Add assistant response to session history. In incognito mode, keeps in-memory context but skips DB persistence."""
|
||||
"""Add assistant response to session history.
|
||||
|
||||
Incognito responses are intentionally not added to the session object. The
|
||||
session may later be saved by a normal turn, so "in-memory only" is not
|
||||
private enough.
|
||||
"""
|
||||
md = dict(last_metrics) if last_metrics else {}
|
||||
def _model_value(value) -> str:
|
||||
if value is None:
|
||||
@@ -1063,19 +1119,18 @@ def save_assistant_response(
|
||||
_content = _think_info["reply"]
|
||||
else:
|
||||
_content = full_response
|
||||
if incognito:
|
||||
_append_incognito_message(session_id, "assistant", _content, md)
|
||||
return None
|
||||
sess.add_message(ChatMessage("assistant", _content, metadata=md))
|
||||
|
||||
if not incognito:
|
||||
from core.database import update_session_last_accessed
|
||||
update_session_last_accessed(session_id)
|
||||
session_manager.save_sessions()
|
||||
from core.database import update_session_last_accessed
|
||||
update_session_last_accessed(session_id)
|
||||
session_manager.save_sessions()
|
||||
|
||||
# Return the persisted message's DB id so the stream can wire it onto the
|
||||
# freshly-rendered bubble — lets the user edit/delete a just-streamed reply
|
||||
# without reloading. Incognito returns None: those messages are ephemeral,
|
||||
# so we don't hand out an edit/delete handle for them.
|
||||
if incognito:
|
||||
return None
|
||||
# without reloading.
|
||||
try:
|
||||
_last = sess.history[-1]
|
||||
_meta = getattr(_last, "metadata", None)
|
||||
|
||||
+314
-26
@@ -42,6 +42,7 @@ from routes.chat_helpers import (
|
||||
_enforce_chat_privileges,
|
||||
)
|
||||
from src.action_intents import ToolIntent, classify_tool_intent as _classify_tool_intent
|
||||
from src.image_model_ids import looks_like_image_generation_model
|
||||
from src.tool_policy import (
|
||||
WEB_TOOL_NAMES,
|
||||
build_effective_tool_policy,
|
||||
@@ -53,7 +54,6 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# Track active streams for partial-save safety net
|
||||
_active_streams: Dict[str, dict] = {}
|
||||
_IMAGE_MODEL_PREFIXES = ("gpt-image", "dall-e", "chatgpt-image")
|
||||
|
||||
|
||||
def _stream_set(session_id: str, **fields) -> None:
|
||||
@@ -111,7 +111,8 @@ def _ensure_current_request_is_latest_user(messages: List[Dict[str, Any]], curre
|
||||
_WEB_FOLLOWUP_RE = re.compile(
|
||||
r"^\s*(?:(?:can|could|would|will)\s+you\s+)?"
|
||||
r"(?:check|try\s+again|look(?:\s+now|\s+it\s+up)?|search(?:\s+now|\s+online|\s+it)?|"
|
||||
r"do\s+it|again)\??\s*$",
|
||||
r"do\s+it|again|approved|approve(?:d)?|yes|ok(?:ay)?|proceed|go\s+ahead|"
|
||||
r"send(?:\s+it)?|submit(?:\s+it)?|email(?:\s+them|\s+it)?)\??\s*$",
|
||||
re.I,
|
||||
)
|
||||
_RECENT_WEB_CONTEXT_RE = re.compile(
|
||||
@@ -119,6 +120,26 @@ _RECENT_WEB_CONTEXT_RE = re.compile(
|
||||
r"price|current|latest|search|look\s+up|online)\b",
|
||||
re.I,
|
||||
)
|
||||
_RECENT_BROWSER_CONTEXT_RE = re.compile(
|
||||
r"\b(?:browser|browse|open\s+(?:the\s+)?(?:site|page|url|link)|click|"
|
||||
r"fill(?:\s+out)?|submit|send\s+(?:the\s+)?form|contact\s+form|web\s*form|"
|
||||
r"form\s+submission|playwright|automation)\b",
|
||||
re.I,
|
||||
)
|
||||
_BROWSER_MCP_TOOLS = {
|
||||
"mcp__builtin_browser__browser_navigate",
|
||||
"mcp__builtin_browser__browser_snapshot",
|
||||
"mcp__builtin_browser__browser_click",
|
||||
"mcp__builtin_browser__browser_type",
|
||||
"mcp__builtin_browser__browser_fill_form",
|
||||
"mcp__builtin_browser__browser_select_option",
|
||||
"mcp__builtin_browser__browser_press_key",
|
||||
"mcp__builtin_browser__browser_wait_for",
|
||||
"mcp__builtin_browser__browser_take_screenshot",
|
||||
"mcp__builtin_browser__browser_drag",
|
||||
"mcp__builtin_browser__browser_navigate_back",
|
||||
"mcp__builtin_browser__browser_close",
|
||||
}
|
||||
|
||||
|
||||
def _recent_session_text(sess, limit: int = 8, max_chars: int = 2000) -> str:
|
||||
@@ -141,6 +162,13 @@ def _is_contextual_web_followup(message: str, sess) -> bool:
|
||||
return bool(_RECENT_WEB_CONTEXT_RE.search(_recent_session_text(sess)))
|
||||
|
||||
|
||||
def _is_contextual_browser_followup(message: str, sess) -> bool:
|
||||
"""Treat short retry replies as browser tasks when recent context was forms/browser automation."""
|
||||
if not message or not _WEB_FOLLOWUP_RE.search(message):
|
||||
return False
|
||||
return bool(_RECENT_BROWSER_CONTEXT_RE.search(_recent_session_text(sess, limit=12, max_chars=4000)))
|
||||
|
||||
|
||||
def _resolve_request_workspace(request, raw_value) -> tuple:
|
||||
"""Resolve the posted workspace for this request: (workspace, rejected).
|
||||
|
||||
@@ -168,6 +196,46 @@ def _resolve_request_workspace(request, raw_value) -> tuple:
|
||||
return workspace, (requested if not workspace else "")
|
||||
|
||||
|
||||
_ABS_PATH_RE = re.compile(r"(?<!\S)(~?/[^\"'\s`<>]+)")
|
||||
_LOCAL_FILE_TASK_RE = re.compile(
|
||||
r"\b(?:file|folder|directory|path|workspace|repo|project|movie|video|"
|
||||
r"subtitle|subtitles|srt|vtt|ass|download|save|rename|move|copy|extract|"
|
||||
r"convert|ffmpeg|run|execute|open|read|inspect|fix|debug|test|build)\b",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _resolve_workspace_from_message_path(request, message: str) -> tuple[str, str]:
|
||||
"""Auto-bind a workspace only when the user names an explicit safe path.
|
||||
|
||||
This is intentionally deterministic rather than LLM/RAG-driven: RAG can
|
||||
choose the tool family, but filesystem binding must not let a prompt infer
|
||||
or probe arbitrary host paths. For a file path, bind its parent directory.
|
||||
For a directory path, bind that directory.
|
||||
"""
|
||||
text = str(message or "")
|
||||
if not text or not _LOCAL_FILE_TASK_RE.search(text):
|
||||
return "", ""
|
||||
|
||||
from src.tool_security import owner_is_admin_or_single_user
|
||||
if not owner_is_admin_or_single_user(get_current_user(request)):
|
||||
return "", ""
|
||||
|
||||
from src.tool_execution import vet_workspace
|
||||
|
||||
for match in _ABS_PATH_RE.finditer(text):
|
||||
raw = match.group(1).rstrip(".,;:)]}")
|
||||
expanded = os.path.realpath(os.path.expanduser(raw))
|
||||
candidates = [expanded]
|
||||
if os.path.isfile(expanded):
|
||||
candidates.insert(0, os.path.dirname(expanded))
|
||||
for candidate in candidates:
|
||||
workspace = vet_workspace(candidate) or ""
|
||||
if workspace:
|
||||
return workspace, ""
|
||||
return "", ""
|
||||
|
||||
|
||||
def _session_url_matches_endpoint(session_url: str, endpoint_base: str) -> bool:
|
||||
if not session_url or not endpoint_base:
|
||||
return False
|
||||
@@ -243,7 +311,7 @@ def _is_image_generation_session(sess, owner: str | None = None) -> bool:
|
||||
models into the image-generation path.
|
||||
"""
|
||||
model = (getattr(sess, "model", "") or "").strip()
|
||||
if any(model.lower().startswith(prefix) for prefix in _IMAGE_MODEL_PREFIXES):
|
||||
if looks_like_image_generation_model(model):
|
||||
return True
|
||||
|
||||
endpoint_url = (getattr(sess, "endpoint_url", "") or "").strip()
|
||||
@@ -271,6 +339,29 @@ def _is_image_generation_session(sess, owner: str | None = None) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _first_image_attachment(chat_handler, att_ids: List[str], owner: str | None = None) -> Optional[Dict[str, Any]]:
|
||||
"""Return the first attached image file that this owner can read."""
|
||||
upload_handler = getattr(chat_handler, "upload_handler", None)
|
||||
if not upload_handler:
|
||||
return None
|
||||
for att_id in att_ids or []:
|
||||
try:
|
||||
info = upload_handler.resolve_upload(att_id, owner=owner)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to resolve image edit upload %s", att_id, exc_info=e)
|
||||
continue
|
||||
if not info:
|
||||
continue
|
||||
name = info.get("name") or info.get("original_name") or info.get("id") or ""
|
||||
mime = info.get("mime", "")
|
||||
try:
|
||||
if upload_handler.is_image_file(name, mime):
|
||||
return info
|
||||
except Exception:
|
||||
continue
|
||||
return None
|
||||
|
||||
|
||||
def _recover_empty_session_model(sess, session_id: str, owner: str | None = None) -> bool:
|
||||
"""Re-populate sess.model from the matching endpoint's cached models.
|
||||
|
||||
@@ -381,9 +472,85 @@ def _recover_empty_session_model(sess, session_id: str, owner: str | None = None
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.warning("Failed to recover empty session model for %s: %s", session_id, e)
|
||||
return False
|
||||
|
||||
|
||||
def _reconcile_selected_route_from_request(
|
||||
request: Request,
|
||||
sess,
|
||||
session_id: str,
|
||||
form_data,
|
||||
owner: str | None = None,
|
||||
) -> bool:
|
||||
"""Apply the model route the browser selected before streaming.
|
||||
|
||||
The frontend creates a pending chat first and only materializes it on first
|
||||
send. Startup/default-model refreshes can race with that UI state, so the
|
||||
stream request includes the route that was selected at click/send time.
|
||||
Trust only registered endpoint ids, or the session's existing endpoint URL.
|
||||
"""
|
||||
selected_model = str(form_data.get("selected_model") or "").strip()
|
||||
selected_endpoint_id = str(form_data.get("selected_endpoint_id") or "").strip()
|
||||
selected_endpoint_url = str(form_data.get("selected_endpoint_url") or "").strip()
|
||||
if not selected_model:
|
||||
return False
|
||||
|
||||
endpoint_url = ""
|
||||
headers = None
|
||||
if selected_endpoint_id or selected_endpoint_url:
|
||||
try:
|
||||
from src.auth_helpers import owner_filter
|
||||
from src.endpoint_resolver import build_headers, normalize_base
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True)
|
||||
if selected_endpoint_id:
|
||||
q = q.filter(ModelEndpoint.id == selected_endpoint_id)
|
||||
if owner:
|
||||
q = owner_filter(q, ModelEndpoint, owner)
|
||||
candidates = q.all() if selected_endpoint_url and not selected_endpoint_id else [q.first()]
|
||||
ep = None
|
||||
for cand in candidates:
|
||||
if not cand:
|
||||
continue
|
||||
if selected_endpoint_id or _session_url_matches_endpoint(selected_endpoint_url, cand.base_url or ""):
|
||||
ep = cand
|
||||
break
|
||||
if not ep:
|
||||
return False
|
||||
endpoint_url = build_chat_url(normalize_base(ep.base_url or ""))
|
||||
headers = build_headers(ep.api_key or "", ep.base_url or "") if ep.api_key else {}
|
||||
finally:
|
||||
db.close()
|
||||
except Exception as e:
|
||||
logger.warning("Failed to resolve selected endpoint %s/%s for %s: %s", selected_endpoint_id, selected_endpoint_url, session_id, e)
|
||||
return False
|
||||
|
||||
if not endpoint_url:
|
||||
return False
|
||||
|
||||
if (
|
||||
selected_model == (getattr(sess, "model", "") or "")
|
||||
and endpoint_url == (getattr(sess, "endpoint_url", "") or "")
|
||||
):
|
||||
return False
|
||||
|
||||
sess.model = selected_model
|
||||
sess.endpoint_url = endpoint_url
|
||||
sess.headers = headers or {}
|
||||
db = SessionLocal()
|
||||
try:
|
||||
db_session = db.query(DBSession).filter(DBSession.id == session_id).first()
|
||||
if db_session:
|
||||
db_session.model = selected_model
|
||||
db_session.endpoint_url = endpoint_url
|
||||
db_session.headers = sess.headers or {}
|
||||
db_session.updated_at = datetime.utcnow()
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
logger.info("Reconciled selected route for %s: model=%r endpoint=%s", session_id, selected_model, redact_url(endpoint_url))
|
||||
return True
|
||||
|
||||
|
||||
def _set_user_time_from_request(request: Request) -> None:
|
||||
@@ -565,9 +732,7 @@ def setup_chat_routes(
|
||||
search_context = form_data.get("search_context") # pre-fetched web search results (compare mode)
|
||||
compare_mode = str(form_data.get("compare_mode", "")).lower() == "true"
|
||||
incognito = str(form_data.get("incognito", "")).lower() == "true"
|
||||
# Plan mode is not part of the merge-ready UI. Ignore stale clients or
|
||||
# manual form posts that still send plan_mode=true.
|
||||
plan_mode = False
|
||||
plan_mode = str(form_data.get("plan_mode") or (body or {}).get("plan_mode") or "").lower() == "true"
|
||||
chat_mode = str(form_data.get("mode", "")).lower() # 'chat' or 'agent'
|
||||
# Workspace: confine the agent's file/shell tools to this folder.
|
||||
workspace, workspace_rejected = _resolve_request_workspace(
|
||||
@@ -589,6 +754,25 @@ def setup_chat_routes(
|
||||
# not chats we quietly promoted for a notes/calendar intent.
|
||||
user_requested_agent = (chat_mode == "agent")
|
||||
_search_enabled = web_search_enabled_for_turn(allow_web_search, use_web)
|
||||
_explicit_web_intent = False
|
||||
_explicit_browser_intent = False
|
||||
if isinstance(message, str):
|
||||
_msg_l = message.lower()
|
||||
_explicit_web_intent = bool(re.search(
|
||||
r"\b(search|look\s*up|lookup|google|browse|web|online|latest|current|today|news|weather|forecast|rate|exchange\s+rate)\b",
|
||||
_msg_l,
|
||||
))
|
||||
_explicit_browser_intent = bool(re.search(
|
||||
r"\b(browser|browse|open\s+(?:the\s+)?(?:site|page|url|link)|"
|
||||
r"click|fill(?:\s+out)?|submit|send\s+(?:the\s+)?form|"
|
||||
r"contact\s+form|web\s*form|form\s+submission)\b",
|
||||
_msg_l,
|
||||
))
|
||||
_allow_browser_for_web_turn = bool(
|
||||
_explicit_browser_intent
|
||||
or _explicit_web_intent
|
||||
or _search_enabled
|
||||
)
|
||||
# Intent auto-escalation: if the user is clearly asking the assistant
|
||||
# to create a todo, reminder, or calendar event, promote chat → agent
|
||||
# for this turn so the LLM has access to manage_notes / manage_calendar.
|
||||
@@ -598,9 +782,13 @@ def setup_chat_routes(
|
||||
# shell disabled).
|
||||
auto_escalated = False
|
||||
_tool_intent = _classify_tool_intent(message) if isinstance(message, str) else None
|
||||
_workspace_agent_intent = False
|
||||
if chat_mode == "chat" and _tool_intent and _tool_intent.needs_tools:
|
||||
chat_mode = "agent"
|
||||
auto_escalated = True
|
||||
_workspace_agent_intent = _tool_intent.category in {"shell", "workspace"}
|
||||
if _workspace_agent_intent:
|
||||
allow_bash = "true"
|
||||
logger.info(
|
||||
"chat→agent auto-escalation: category=%s reason=%s",
|
||||
_tool_intent.category,
|
||||
@@ -610,6 +798,10 @@ def setup_chat_routes(
|
||||
chat_mode = "agent"
|
||||
auto_escalated = True
|
||||
logger.info("chat→agent auto-escalation: search enabled")
|
||||
elif chat_mode == "chat" and _explicit_web_intent:
|
||||
chat_mode = "agent"
|
||||
auto_escalated = True
|
||||
logger.info("chat→agent auto-escalation: explicit web intent")
|
||||
active_doc_id = form_data.get("active_doc_id", "").strip()
|
||||
logger.info(f"[doc-inject] chat_mode={chat_mode}, active_doc_id={active_doc_id!r}")
|
||||
|
||||
@@ -688,6 +880,7 @@ def setup_chat_routes(
|
||||
_verify_session_owner(request, session)
|
||||
sess = session_manager.get_session(session)
|
||||
owner = effective_user(request)
|
||||
_reconcile_selected_route_from_request(request, sess, session, form_data, owner=owner)
|
||||
if _clear_orphaned_session_endpoint(sess, owner=owner):
|
||||
raise HTTPException(400, "Selected model endpoint was removed. Pick another model in Settings.")
|
||||
# Issue #587: picker shows a model from the endpoint cache but
|
||||
@@ -711,11 +904,28 @@ def setup_chat_routes(
|
||||
_tool_intent = ToolIntent(True, "web", "contextual web lookup follow-up")
|
||||
chat_mode = "agent"
|
||||
auto_escalated = True
|
||||
_workspace_agent_intent = False
|
||||
logger.info(
|
||||
"chat→agent auto-escalation: category=%s reason=%s",
|
||||
_tool_intent.category,
|
||||
_tool_intent.reason,
|
||||
)
|
||||
if isinstance(message, str) and _is_contextual_browser_followup(message, sess):
|
||||
_explicit_browser_intent = True
|
||||
if chat_mode == "chat":
|
||||
chat_mode = "agent"
|
||||
auto_escalated = True
|
||||
_workspace_agent_intent = False
|
||||
logger.info("chat→agent auto-escalation: contextual browser/form follow-up")
|
||||
if not workspace and isinstance(message, str):
|
||||
_auto_workspace, _ = _resolve_workspace_from_message_path(request, message)
|
||||
if _auto_workspace:
|
||||
workspace = _auto_workspace
|
||||
chat_mode = "agent"
|
||||
auto_escalated = True
|
||||
_workspace_agent_intent = True
|
||||
allow_bash = "true"
|
||||
logger.info("chat→agent auto-escalation: explicit path workspace=%s", workspace)
|
||||
except SessionNotFoundError as e:
|
||||
raise HTTPException(404, str(e))
|
||||
except (ValueError, ValidationError):
|
||||
@@ -750,7 +960,12 @@ def setup_chat_routes(
|
||||
except Exception as e:
|
||||
logger.warning("Failed to parse attachments JSON, ignoring attachments", exc_info=e)
|
||||
|
||||
image_generation_session = _is_image_generation_session(sess, owner=effective_user(request))
|
||||
no_memory = str(form_data.get("no_memory", "")).lower() == "true"
|
||||
if image_generation_session:
|
||||
no_memory = True
|
||||
use_rag = "false"
|
||||
search_context = None
|
||||
pre_context_tool_policy = build_effective_tool_policy(
|
||||
last_user_message=message,
|
||||
)
|
||||
@@ -879,7 +1094,7 @@ def setup_chat_routes(
|
||||
# explicitly enable it.
|
||||
if allow_bash is not None and str(allow_bash).lower() != "true":
|
||||
disabled_tools.add("bash")
|
||||
_explicit_web_intent = bool(_tool_intent and _tool_intent.category == "web")
|
||||
_explicit_web_intent = _explicit_web_intent or bool(_tool_intent and _tool_intent.category == "web")
|
||||
if is_web_search_explicitly_denied(allow_web_search) or not _search_enabled:
|
||||
disabled_tools.update(WEB_TOOL_NAMES)
|
||||
if _explicit_web_intent:
|
||||
@@ -893,7 +1108,7 @@ def setup_chat_routes(
|
||||
"create_document", "edit_document", "update_document",
|
||||
"send_email", "reply_to_email",
|
||||
"manage_notes", "manage_calendar", "manage_tasks",
|
||||
"api_call", "builtin_browser",
|
||||
"api_call",
|
||||
})
|
||||
if _search_enabled:
|
||||
disabled_tools.difference_update(WEB_TOOL_NAMES)
|
||||
@@ -909,6 +1124,11 @@ def setup_chat_routes(
|
||||
"manage_memory", # persistent memory store
|
||||
"search_chats", # past chat history
|
||||
"manage_skills", # skill presets tied to user
|
||||
"create_session",
|
||||
"list_sessions",
|
||||
"manage_session",
|
||||
"send_to_session",
|
||||
"chat_with_model",
|
||||
})
|
||||
|
||||
# Active email reader open → strip the tools that let the agent drift
|
||||
@@ -935,7 +1155,7 @@ def setup_chat_routes(
|
||||
if not _privs.get("can_use_bash", True):
|
||||
disabled_tools.update({"bash", "python", "read_file", "write_file"})
|
||||
if not _privs.get("can_use_browser", True):
|
||||
disabled_tools.add("builtin_browser")
|
||||
disabled_tools.update(_BROWSER_MCP_TOOLS)
|
||||
if not _privs.get("can_use_documents", True):
|
||||
disabled_tools.update({"create_document", "edit_document", "update_document", "suggest_document"})
|
||||
if not _privs.get("can_generate_images", True):
|
||||
@@ -958,10 +1178,12 @@ def setup_chat_routes(
|
||||
# the heavy "do things on the computer" tools — otherwise the model
|
||||
# tries to shell out for a request that never needed it, then fails
|
||||
# (and looks broken when the shell is disabled).
|
||||
if auto_escalated:
|
||||
if auto_escalated and not _workspace_agent_intent:
|
||||
disabled_tools.update({
|
||||
"bash", "python", "read_file", "write_file", "builtin_browser",
|
||||
"bash", "python", "read_file", "write_file",
|
||||
})
|
||||
if not _allow_browser_for_web_turn:
|
||||
disabled_tools.update(_BROWSER_MCP_TOOLS)
|
||||
|
||||
# Disable document tools in compare sessions — they break the pane UI
|
||||
if sess.name and sess.name.startswith("[CMP]"):
|
||||
@@ -1195,7 +1417,7 @@ def setup_chat_routes(
|
||||
_model_info["character_name"] = ctx.preset.character_name
|
||||
yield f'data: {json.dumps(_model_info)}\n\n'
|
||||
|
||||
if _is_image_generation_session(sess, owner=_user):
|
||||
if image_generation_session:
|
||||
from src.settings import get_setting
|
||||
if tool_policy.blocks("generate_image"):
|
||||
_blocked_msg = tool_policy.reason_for("generate_image")
|
||||
@@ -1208,26 +1430,85 @@ def setup_chat_routes(
|
||||
yield "data: [DONE]\n\n"
|
||||
_active_streams.pop(session, None)
|
||||
return
|
||||
from src.ai_interaction import do_generate_image
|
||||
from src.ai_interaction import do_edit_image, do_generate_image
|
||||
_user_msg = message or ""
|
||||
yield f'data: {json.dumps({"type": "tool_start", "tool": "generate_image", "command": _user_msg[:100]})}\n\n'
|
||||
_image_upload = _first_image_attachment(chat_handler, att_ids, owner=_user)
|
||||
_image_tool_name = "edit_image" if _image_upload else "generate_image"
|
||||
yield f'data: {json.dumps({"type": "tool_start", "tool": _image_tool_name, "command": _user_msg[:100]})}\n\n'
|
||||
yield ": heartbeat\n\n"
|
||||
_img_result = await do_generate_image(f"{_user_msg}\n{sess.model}", session, owner=_user)
|
||||
_progress_queue: asyncio.Queue = asyncio.Queue()
|
||||
|
||||
async def _image_progress_callback(progress: Dict[str, Any]):
|
||||
try:
|
||||
_progress_queue.put_nowait(progress)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if _image_upload:
|
||||
_img_task = asyncio.create_task(do_edit_image(
|
||||
_user_msg,
|
||||
_image_upload.get("path", ""),
|
||||
model_spec=sess.model,
|
||||
session_id=session,
|
||||
owner=_user,
|
||||
size="1024x1024",
|
||||
progress_callback=_image_progress_callback,
|
||||
))
|
||||
else:
|
||||
_img_task = asyncio.create_task(do_generate_image(f"{_user_msg}\n{sess.model}\n512x512", session, owner=_user))
|
||||
_img_started = time.time()
|
||||
_img_tick = 0
|
||||
while not _img_task.done():
|
||||
try:
|
||||
_progress = await asyncio.wait_for(_progress_queue.get(), timeout=2.0)
|
||||
except asyncio.TimeoutError:
|
||||
_progress = None
|
||||
_img_tick += 1
|
||||
_elapsed = int(time.time() - _img_started)
|
||||
_label = "Editing image" if _image_upload else "Generating image"
|
||||
yield ": image generation still running\n\n"
|
||||
_progress_data = {"type": "tool_progress", "tool": _image_tool_name, "message": f"{_label}… {_elapsed}s", "elapsed": _elapsed, "tick": _img_tick}
|
||||
if isinstance(_progress, dict) and _progress.get("total"):
|
||||
_step = int(_progress.get("step") or 0)
|
||||
_total = int(_progress.get("total") or 0)
|
||||
_percent = _progress.get("percent")
|
||||
_progress_data.update({
|
||||
"step": _step,
|
||||
"total": _total,
|
||||
"percent": _percent,
|
||||
"message": f"{_label}… {_step}/{_total}",
|
||||
})
|
||||
yield f'data: {json.dumps(_progress_data)}\n\n'
|
||||
_img_result = await _img_task
|
||||
_img_output = _img_result.get("results", _img_result.get("error", ""))
|
||||
_img_tool_data = {"type": "tool_output", "tool": "generate_image", "command": _user_msg[:100], "output": _img_output, "exit_code": 0 if "error" not in _img_result else 1}
|
||||
_img_tool_data = {"type": "tool_output", "tool": _image_tool_name, "command": _user_msg[:100], "output": _img_output, "exit_code": 0 if "error" not in _img_result else 1}
|
||||
for _k in ("image_url", "image_id", "image_prompt", "image_model", "image_size", "image_quality"):
|
||||
if _k in _img_result:
|
||||
_img_tool_data[_k] = _img_result[_k]
|
||||
if _image_upload:
|
||||
_img_tool_data["source_image"] = {
|
||||
"id": _image_upload.get("id"),
|
||||
"name": _image_upload.get("name") or _image_upload.get("original_name"),
|
||||
}
|
||||
yield f'data: {json.dumps(_img_tool_data)}\n\n'
|
||||
if _img_result.get("image_url"):
|
||||
_img_event = {"type": "generated_image", "url": _img_result.get("image_url")}
|
||||
for _k in ("image_url", "image_id", "image_prompt", "image_model", "image_size", "image_quality"):
|
||||
if _img_result.get(_k):
|
||||
_img_event[_k] = _img_result[_k]
|
||||
yield f'data: {json.dumps(_img_event)}\n\n'
|
||||
_desc = _img_result.get("results", _img_result.get("error", "Image generation complete"))
|
||||
full_response = _desc
|
||||
yield f'data: {json.dumps({"delta": _desc})}\n\n'
|
||||
# Save to session history
|
||||
if not incognito:
|
||||
_ev = {"round": 1, "tool": "generate_image", "command": _user_msg[:100], "output": _img_output, "exit_code": 0 if "error" not in _img_result else 1}
|
||||
_ev = {"round": 1, "tool": _image_tool_name, "command": _user_msg[:100], "output": _img_output, "exit_code": 0 if "error" not in _img_result else 1}
|
||||
for _ek in ("image_url", "image_id", "image_prompt", "image_model", "image_size", "image_quality"):
|
||||
if _img_result.get(_ek):
|
||||
_ev[_ek] = _img_result[_ek]
|
||||
if _image_upload:
|
||||
_ev["source_image_id"] = _image_upload.get("id")
|
||||
_ev["source_image_name"] = _image_upload.get("name") or _image_upload.get("original_name")
|
||||
sess.add_message(ChatMessage("assistant", full_response, metadata={"tool_events": [_ev], "model": sess.model}))
|
||||
session_manager.save_sessions()
|
||||
yield f'data: {json.dumps({"type": "metrics", "data": {"total_time": 0}})}\n\n'
|
||||
@@ -1292,8 +1573,10 @@ def setup_chat_routes(
|
||||
last_metrics["context_messages_after_trim"] = ctx.context_messages_after_trim
|
||||
last_metrics["context_tokens_before_trim"] = ctx.context_tokens_before_trim
|
||||
last_metrics["context_tokens_after_trim"] = ctx.context_tokens_after_trim
|
||||
if ctx.context_length and last_metrics.get("input_tokens"):
|
||||
pct = min(round((last_metrics["input_tokens"] / ctx.context_length) * 100, 1), 100.0)
|
||||
request_context_tokens = ctx.context_tokens_after_trim or estimate_tokens(messages)
|
||||
last_metrics["request_context_tokens"] = request_context_tokens
|
||||
if ctx.context_length and request_context_tokens:
|
||||
pct = min(round((request_context_tokens / ctx.context_length) * 100, 1), 100.0)
|
||||
last_metrics["context_percent"] = pct
|
||||
last_metrics["context_length"] = ctx.context_length
|
||||
# The frontend reads `tokens_per_second`; the raw usage event
|
||||
@@ -1326,6 +1609,7 @@ def setup_chat_routes(
|
||||
"input_tokens": _est_in,
|
||||
"output_tokens": _est_out,
|
||||
"tokens_per_second": _tps,
|
||||
"request_context_tokens": _est_in,
|
||||
"context_percent": _ctx_pct,
|
||||
"context_length": ctx.context_length,
|
||||
"model": _actual_model or _answered_by or _requested_model,
|
||||
@@ -1360,7 +1644,7 @@ def setup_chat_routes(
|
||||
_stream_set(session, status="done")
|
||||
yield chunk
|
||||
except (asyncio.CancelledError, GeneratorExit):
|
||||
if full_response:
|
||||
if full_response and not incognito:
|
||||
logger.info("Client disconnected mid-stream (chat mode) for session %s, saving partial (%d chars)", session, len(full_response))
|
||||
_stopped_content, _stopped_md = clean_thinking_for_save(
|
||||
full_response,
|
||||
@@ -1371,8 +1655,7 @@ def setup_chat_routes(
|
||||
},
|
||||
)
|
||||
sess.add_message(ChatMessage("assistant", _stopped_content, metadata=_stopped_md))
|
||||
if not incognito:
|
||||
session_manager.save_sessions()
|
||||
session_manager.save_sessions()
|
||||
raise
|
||||
finally:
|
||||
_active_streams.pop(session, None)
|
||||
@@ -1405,6 +1688,10 @@ def setup_chat_routes(
|
||||
_forced_tools = None
|
||||
if _search_enabled:
|
||||
_forced_tools = set(WEB_TOOL_NAMES)
|
||||
if _explicit_browser_intent:
|
||||
_forced_tools |= set(_BROWSER_MCP_TOOLS)
|
||||
elif _explicit_browser_intent:
|
||||
_forced_tools = set(_BROWSER_MCP_TOOLS)
|
||||
|
||||
async for chunk in stream_agent_loop(
|
||||
sess.endpoint_url,
|
||||
@@ -1450,7 +1737,9 @@ def setup_chat_routes(
|
||||
"tool_start", "tool_output", "agent_step",
|
||||
"doc_stream_open", "doc_stream_delta",
|
||||
"doc_update", "doc_suggestions", "ui_control",
|
||||
"rounds_exhausted",
|
||||
"rounds_exhausted", "budget_exceeded",
|
||||
"loop_breaker_triggered",
|
||||
"intent_nudge_exhausted",
|
||||
"ask_user",
|
||||
"plan_update",
|
||||
):
|
||||
@@ -1527,7 +1816,7 @@ def setup_chat_routes(
|
||||
# outer finally from running and left _active_streams
|
||||
# with a stale entry).
|
||||
try:
|
||||
if full_response:
|
||||
if full_response and not incognito:
|
||||
logger.info("Client disconnected mid-stream for session %s, saving partial response (%d chars)", session, len(full_response))
|
||||
_stopped_content2, _stopped_md2 = clean_thinking_for_save(
|
||||
full_response,
|
||||
@@ -1538,8 +1827,7 @@ def setup_chat_routes(
|
||||
},
|
||||
)
|
||||
sess.add_message(ChatMessage("assistant", _stopped_content2, metadata=_stopped_md2))
|
||||
if not incognito:
|
||||
session_manager.save_sessions()
|
||||
session_manager.save_sessions()
|
||||
except Exception:
|
||||
logger.exception("Failed to save partial response on disconnect (session %s)", session)
|
||||
raise
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Cleanup route domain package (slice 2g, #4082/#4071).
|
||||
|
||||
Contains cleanup_routes.py, migrated from the flat routes/ directory.
|
||||
Backward-compat shim at routes/cleanup_routes.py re-exports from here.
|
||||
"""
|
||||
@@ -0,0 +1,60 @@
|
||||
# routes/cleanup_routes.py
|
||||
"""Routes for cleanup operations."""
|
||||
import logging
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from src.cleanup_service import get_cleanup_preview, cleanup_sessions
|
||||
from src.auth_helpers import get_current_user
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
def setup_cleanup_routes(session_manager):
|
||||
"""
|
||||
Setup cleanup-related routes.
|
||||
|
||||
Args:
|
||||
session_manager: SessionManager instance
|
||||
|
||||
Returns:
|
||||
APIRouter instance with cleanup routes
|
||||
"""
|
||||
router = APIRouter(prefix="/api/cleanup")
|
||||
|
||||
@router.get("/preview")
|
||||
async def cleanup_preview(request: Request):
|
||||
"""
|
||||
Preview what would be cleaned up without making any changes.
|
||||
|
||||
Returns:
|
||||
JSON response with lists of sessions that would be archived/deleted and estimated space savings
|
||||
"""
|
||||
user = get_current_user(request)
|
||||
try:
|
||||
preview = await get_cleanup_preview(owner=user)
|
||||
return preview
|
||||
except Exception as e:
|
||||
logger.error(f"Cleanup preview failed: {e}")
|
||||
raise HTTPException(500, "Cleanup preview generation failed")
|
||||
|
||||
@router.post("")
|
||||
async def cleanup_endpoint(request: Request):
|
||||
"""
|
||||
Perform cleanup operations:
|
||||
1. Archive inactive sessions (not accessed for 7 days)
|
||||
2. Delete old sessions (archived, not important, not accessed for 14+ days, with fewer than 10 messages)
|
||||
|
||||
Returns:
|
||||
JSON response with counts of deleted and archived sessions, and space freed
|
||||
"""
|
||||
user = get_current_user(request)
|
||||
try:
|
||||
archived_count, deleted_count, space_freed_mb = await cleanup_sessions(session_manager, owner=user)
|
||||
return {
|
||||
"archived_count": archived_count,
|
||||
"deleted_count": deleted_count,
|
||||
"space_freed_mb": round(space_freed_mb, 2)
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Cleanup failed: {e}")
|
||||
raise HTTPException(500, "Cleanup operation failed")
|
||||
|
||||
return router
|
||||
+13
-56
@@ -1,60 +1,17 @@
|
||||
# routes/cleanup_routes.py
|
||||
"""Routes for cleanup operations."""
|
||||
import logging
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from src.cleanup_service import get_cleanup_preview, cleanup_sessions
|
||||
from src.auth_helpers import get_current_user
|
||||
"""Backward-compat shim — canonical location is routes/cleanup/cleanup_routes.py.
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
This module is replaced in ``sys.modules`` by the canonical module object so
|
||||
that ``import routes.cleanup_routes``, ``from routes.cleanup_routes import X``,
|
||||
``importlib.import_module("routes.cleanup_routes")``, and the string-targeted
|
||||
``monkeypatch.setattr("routes.cleanup_routes.get_cleanup_preview", ...)`` /
|
||||
``"routes.cleanup_routes.get_current_user"`` / ``"routes.cleanup_routes.
|
||||
cleanup_sessions"`` pattern used by test_cleanup_owner_scope.py all operate
|
||||
on the *same* object the application actually uses. Keeps existing import
|
||||
paths working after slice 2g (#4082/#4071).
|
||||
"""
|
||||
|
||||
def setup_cleanup_routes(session_manager):
|
||||
"""
|
||||
Setup cleanup-related routes.
|
||||
import sys as _sys
|
||||
|
||||
Args:
|
||||
session_manager: SessionManager instance
|
||||
from routes.cleanup import cleanup_routes as _canonical # noqa: F401
|
||||
|
||||
Returns:
|
||||
APIRouter instance with cleanup routes
|
||||
"""
|
||||
router = APIRouter(prefix="/api/cleanup")
|
||||
|
||||
@router.get("/preview")
|
||||
async def cleanup_preview(request: Request):
|
||||
"""
|
||||
Preview what would be cleaned up without making any changes.
|
||||
|
||||
Returns:
|
||||
JSON response with lists of sessions that would be archived/deleted and estimated space savings
|
||||
"""
|
||||
user = get_current_user(request)
|
||||
try:
|
||||
preview = await get_cleanup_preview(owner=user)
|
||||
return preview
|
||||
except Exception as e:
|
||||
logger.error(f"Cleanup preview failed: {e}")
|
||||
raise HTTPException(500, "Cleanup preview generation failed")
|
||||
|
||||
@router.post("")
|
||||
async def cleanup_endpoint(request: Request):
|
||||
"""
|
||||
Perform cleanup operations:
|
||||
1. Archive inactive sessions (not accessed for 7 days)
|
||||
2. Delete old sessions (archived, not important, not accessed for 14+ days, with fewer than 10 messages)
|
||||
|
||||
Returns:
|
||||
JSON response with counts of deleted and archived sessions, and space freed
|
||||
"""
|
||||
user = get_current_user(request)
|
||||
try:
|
||||
archived_count, deleted_count, space_freed_mb = await cleanup_sessions(session_manager, owner=user)
|
||||
return {
|
||||
"archived_count": archived_count,
|
||||
"deleted_count": deleted_count,
|
||||
"space_freed_mb": round(space_freed_mb, 2)
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Cleanup failed: {e}")
|
||||
raise HTTPException(500, "Cleanup operation failed")
|
||||
|
||||
return router
|
||||
_sys.modules[__name__] = _canonical
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Compare route domain package (slice 2i, #4082/#4071).
|
||||
|
||||
Contains compare_routes.py, migrated from the flat routes/ directory.
|
||||
Backward-compat shim at routes/compare_routes.py re-exports from here.
|
||||
"""
|
||||
@@ -0,0 +1,365 @@
|
||||
# routes/compare_routes.py
|
||||
"""Model A/B comparison routes."""
|
||||
import json
|
||||
import uuid
|
||||
import random
|
||||
from datetime import datetime
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from typing import List
|
||||
from pydantic import BaseModel
|
||||
import logging
|
||||
|
||||
from core.database import Comparison, SessionLocal
|
||||
from core.session_manager import SessionManager
|
||||
from src.auth_helpers import get_current_user
|
||||
from routes.session_routes import _reject_raw_endpoint_url_for_non_admin
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/compare", tags=["compare"])
|
||||
|
||||
|
||||
def _owned_endpoint_by_url(db, base_url, owner):
|
||||
"""ModelEndpoint whose base_url == `base_url` and is VISIBLE to `owner`
|
||||
(their own rows + legacy null-owner "shared" rows); None otherwise.
|
||||
|
||||
Owner-scoped on purpose. ModelEndpoint is per-user (core/database.py: non-null
|
||||
owner = private, "the model picker only shows the endpoint to that user") and
|
||||
holds a decrypted `api_key`. start_comparison copies the matched row's api_key
|
||||
into the caller-owned [CMP] session's headers, which then drives that session's
|
||||
/api/chat_stream calls — so an UNSCOPED base_url match would let a user mint a
|
||||
comparison bound to ANOTHER user's private endpoint and spend that owner's
|
||||
api_key / reach whatever base_url they configured. Mirrors
|
||||
session_routes._owned_endpoint. A null/empty owner is a no-op (single-user /
|
||||
legacy mode).
|
||||
"""
|
||||
from core.database import ModelEndpoint
|
||||
from src.auth_helpers import owner_filter
|
||||
q = db.query(ModelEndpoint).filter(ModelEndpoint.base_url == base_url)
|
||||
return owner_filter(q, ModelEndpoint, owner).first()
|
||||
|
||||
|
||||
def _owned_endpoint_by_id(db, endpoint_id, owner):
|
||||
"""ModelEndpoint whose id == `endpoint_id` and is VISIBLE to `owner` (their
|
||||
own rows + legacy null-owner "shared" rows); None otherwise.
|
||||
|
||||
Preferred over _owned_endpoint_by_url for credential resolution: two visible
|
||||
endpoints can share the same base_url but hold DIFFERENT api_keys (e.g. two
|
||||
accounts on the same provider). A base_url-only match returns whichever row
|
||||
sorts first, so it can copy the WRONG owner-scoped key into the [CMP] session.
|
||||
An id pins the exact registered endpoint, so /api/compare/start prefers it and
|
||||
only falls back to URL matching for legacy / admin raw-URL callers. Owner
|
||||
scoping is identical to _owned_endpoint_by_url (a null/empty owner is a no-op).
|
||||
"""
|
||||
from core.database import ModelEndpoint
|
||||
from src.auth_helpers import owner_filter
|
||||
q = db.query(ModelEndpoint).filter(ModelEndpoint.id == endpoint_id)
|
||||
return owner_filter(q, ModelEndpoint, owner).first()
|
||||
|
||||
|
||||
class RecordVoteRequest(BaseModel):
|
||||
prompt: str
|
||||
models: List[str]
|
||||
winner: str # model name or "tie"
|
||||
is_blind: bool = True
|
||||
|
||||
|
||||
def setup_compare_routes(session_manager: SessionManager):
|
||||
"""Setup comparison routes."""
|
||||
|
||||
@router.post("/start")
|
||||
def start_comparison(
|
||||
request: Request,
|
||||
prompt: str = Form(...),
|
||||
model_a: str = Form(...),
|
||||
model_b: str = Form(...),
|
||||
endpoint_a: str = Form(""),
|
||||
endpoint_b: str = Form(""),
|
||||
endpoint_a_id: str = Form(""),
|
||||
endpoint_b_id: str = Form(""),
|
||||
is_blind: str = Form("true"),
|
||||
):
|
||||
"""Create two ephemeral sessions and a comparison record.
|
||||
|
||||
Returns the comparison ID and the two session IDs so the client
|
||||
can fire two independent SSE streams to /api/chat_stream.
|
||||
"""
|
||||
user = getattr(request.state, 'current_user', None)
|
||||
comp_id = str(uuid.uuid4())
|
||||
sid_a = str(uuid.uuid4())
|
||||
sid_b = str(uuid.uuid4())
|
||||
|
||||
# Blind mapping: randomly assign left/right
|
||||
blind = str(is_blind).lower() == "true"
|
||||
if blind:
|
||||
mapping = {"left": "a", "right": "b"}
|
||||
if random.random() > 0.5:
|
||||
mapping = {"left": "b", "right": "a"}
|
||||
else:
|
||||
mapping = {"left": "a", "right": "b"}
|
||||
|
||||
# Map session IDs to left/right based on blind mapping
|
||||
session_left = sid_a if mapping["left"] == "a" else sid_b
|
||||
session_right = sid_a if mapping["right"] == "a" else sid_b
|
||||
|
||||
# In blind mode, name the helper sessions by their neutral slot
|
||||
# ("Model A" / "Model B") instead of the real model. Otherwise the
|
||||
# session name leaks the model in the sidebar and GET /api/sessions,
|
||||
# de-anonymizing the comparison before the user votes (issue #1285).
|
||||
slot_name = {session_left: "Model A", session_right: "Model B"}
|
||||
|
||||
# SECURITY: resolve and validate BOTH endpoints before creating any
|
||||
# session. Compare copies a registered endpoint's Authorization header
|
||||
# into the [CMP] session, so validating one endpoint while creating its
|
||||
# session, then rejecting the other, would leave a partial compare
|
||||
# session behind with that header attached. Doing all the owner-scope
|
||||
# resolution + raw-URL rejection up front means a 403 on either endpoint
|
||||
# aborts the whole request with nothing created and no header copied.
|
||||
from src.endpoint_resolver import build_chat_url, build_headers, normalize_base
|
||||
resolved = []
|
||||
db = SessionLocal()
|
||||
try:
|
||||
for sid, model, endpoint, endpoint_id in [
|
||||
(sid_a, model_a, endpoint_a, endpoint_a_id),
|
||||
(sid_b, model_b, endpoint_b, endpoint_b_id),
|
||||
]:
|
||||
# Prefer an explicit endpoint id: it pins the EXACT registered
|
||||
# endpoint (and its api_key), even when two endpoints visible to
|
||||
# the caller share a base_url with different keys — a URL-only
|
||||
# match would copy whichever row sorts first, i.e. possibly the
|
||||
# wrong key. Fall back to URL resolution only for legacy / admin
|
||||
# raw-URL callers that don't send an id.
|
||||
eid = endpoint_id.strip() if isinstance(endpoint_id, str) else ""
|
||||
if eid:
|
||||
ep = _owned_endpoint_by_id(db, eid, user)
|
||||
if ep is None:
|
||||
# An id the caller can't see (wrong owner / deleted) must
|
||||
# NOT silently fall back to a same-URL row with a different
|
||||
# key — that's exactly the mix-up ids exist to prevent.
|
||||
raise HTTPException(404, "Model endpoint not found")
|
||||
# The id already resolved the endpoint; ignore any raw URL the
|
||||
# caller also sent and dial the stored config instead.
|
||||
endpoint = ep.base_url
|
||||
elif not endpoint:
|
||||
raise HTTPException(
|
||||
422, "endpoint_a/endpoint_b or endpoint_a_id/endpoint_b_id is required"
|
||||
)
|
||||
else:
|
||||
# Resolve the supplied URL to a ModelEndpoint the caller owns
|
||||
# (their own rows + legacy null-owner shared rows), scoped so a
|
||||
# comparison can't borrow another user's private endpoint key.
|
||||
base = normalize_base(endpoint)
|
||||
ep = _owned_endpoint_by_url(db, base, user)
|
||||
# Reject *unregistered* raw URLs for signed-in non-admins; a
|
||||
# matched registered endpoint supplies an id so the caller can
|
||||
# still compare endpoints they own. Blanket-rejecting here (the
|
||||
# earlier `endpoint_id=None` call) locked non-admins out of
|
||||
# compare entirely, since compare resolves endpoints by URL with
|
||||
# no endpoint_id. Mirrors the gallery inpaint/harmonize checks.
|
||||
# Raised here (phase 1), before any session exists.
|
||||
_reject_raw_endpoint_url_for_non_admin(
|
||||
request, user, str(ep.id) if ep is not None else None, endpoint
|
||||
)
|
||||
# Bind the [CMP] session to the RESOLVED endpoint, not the raw
|
||||
# caller-supplied string. When the URL matches a registered
|
||||
# endpoint visible to the caller, use that row's own normalized
|
||||
# base URL (the same value owner scoping + endpoint validation
|
||||
# already vetted) so the session dials exactly where the stored
|
||||
# config points. The raw `endpoint` only survives for callers
|
||||
# allowed to pass one — admins / single-user mode, where
|
||||
# `_reject_raw_endpoint_url_for_non_admin` is a no-op and `ep`
|
||||
# is None. Mirrors the registered-endpoint path in session_routes.
|
||||
session_endpoint_url = (
|
||||
build_chat_url(normalize_base(ep.base_url)) if ep is not None else endpoint
|
||||
)
|
||||
# Headers come only from a matched endpoint's key; None when
|
||||
# `ep` is None (raw admin URL or no match), so a comparison can
|
||||
# never inherit another user's key/headers.
|
||||
headers = build_headers(ep.api_key, ep.base_url) if (ep and ep.api_key) else None
|
||||
resolved.append((sid, model, session_endpoint_url, headers))
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# Both endpoints validated — only now create the ephemeral [CMP]
|
||||
# sessions and copy any resolved headers.
|
||||
for sid, model, session_endpoint_url, headers in resolved:
|
||||
name = f"[CMP] {slot_name[sid]}" if blind else f"[CMP] {model.split('/')[-1]}"
|
||||
session_manager.create_session(
|
||||
session_id=sid,
|
||||
name=name,
|
||||
endpoint_url=session_endpoint_url,
|
||||
model=model,
|
||||
rag=False,
|
||||
owner=user,
|
||||
)
|
||||
if headers:
|
||||
s = session_manager.sessions.get(sid)
|
||||
if s:
|
||||
s.headers = headers
|
||||
|
||||
# Store comparison record
|
||||
db = SessionLocal()
|
||||
try:
|
||||
comp = Comparison(
|
||||
id=comp_id,
|
||||
prompt=prompt,
|
||||
model_a=model_a,
|
||||
model_b=model_b,
|
||||
# Record the URL the session actually dials. For URL callers this
|
||||
# is their raw input; for id-only callers (empty endpoint_a/_b)
|
||||
# fall back to the resolved endpoint URL so the column stays
|
||||
# meaningful and non-null. resolved is in [a, b] order.
|
||||
endpoint_a=endpoint_a or resolved[0][2],
|
||||
endpoint_b=endpoint_b or resolved[1][2],
|
||||
is_blind=blind,
|
||||
blind_mapping=json.dumps(mapping),
|
||||
owner=user,
|
||||
)
|
||||
db.add(comp)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# In blind mode, withhold the model identities AND the left/right
|
||||
# mapping from the response. The client already knows model_a/model_b
|
||||
# (it sent them), so returning either would defeat blind mode. They are
|
||||
# revealed by POST /api/compare/{id}/vote once the user has voted (#1285).
|
||||
return {
|
||||
"id": comp_id,
|
||||
"session_left": session_left,
|
||||
"session_right": session_right,
|
||||
"model_left": None if blind else (model_a if mapping["left"] == "a" else model_b),
|
||||
"model_right": None if blind else (model_a if mapping["right"] == "a" else model_b),
|
||||
"is_blind": blind,
|
||||
"mapping": None if blind else mapping,
|
||||
}
|
||||
|
||||
@router.post("/{comp_id}/vote")
|
||||
def vote_comparison(
|
||||
request: Request,
|
||||
comp_id: str,
|
||||
winner: str = Form(...), # "left", "right", or "tie"
|
||||
):
|
||||
"""Record the user's vote and reveal model names if blind."""
|
||||
user = get_current_user(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
comp = db.query(Comparison).filter(Comparison.id == comp_id).first()
|
||||
if not comp:
|
||||
raise HTTPException(404, "Comparison not found")
|
||||
# SECURITY: strict ownership — null-owner Comparisons were
|
||||
# accessible to every user.
|
||||
if user and comp.owner != user:
|
||||
raise HTTPException(404, "Comparison not found")
|
||||
if comp.winner:
|
||||
raise HTTPException(400, "Already voted")
|
||||
|
||||
mapping = json.loads(comp.blind_mapping) if comp.blind_mapping else {"left": "a", "right": "b"}
|
||||
|
||||
if winner == "tie":
|
||||
comp.winner = "tie"
|
||||
elif winner == "left":
|
||||
comp.winner = mapping["left"]
|
||||
elif winner == "right":
|
||||
comp.winner = mapping["right"]
|
||||
else:
|
||||
raise HTTPException(400, "winner must be 'left', 'right', or 'tie'")
|
||||
|
||||
comp.voted_at = datetime.utcnow()
|
||||
db.commit()
|
||||
|
||||
return {
|
||||
"winner": comp.winner,
|
||||
"model_a": comp.model_a,
|
||||
"model_b": comp.model_b,
|
||||
"revealed": {
|
||||
"left": comp.model_a if mapping["left"] == "a" else comp.model_b,
|
||||
"right": comp.model_a if mapping["right"] == "a" else comp.model_b,
|
||||
},
|
||||
}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.post("/record")
|
||||
def record_comparison(request: Request, body: RecordVoteRequest):
|
||||
"""Lightweight endpoint to record a comparison vote from the frontend."""
|
||||
user = get_current_user(request)
|
||||
comp_id = str(uuid.uuid4())
|
||||
|
||||
model_a = body.models[0] if len(body.models) > 0 else ""
|
||||
model_b = body.models[1] if len(body.models) > 1 else ""
|
||||
|
||||
# For N>2 models, store the full list as JSON in blind_mapping
|
||||
if len(body.models) > 2:
|
||||
blind_mapping = json.dumps({"models": body.models})
|
||||
else:
|
||||
blind_mapping = None
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
comp = Comparison(
|
||||
id=comp_id,
|
||||
prompt=body.prompt[:500],
|
||||
model_a=model_a,
|
||||
model_b=model_b,
|
||||
endpoint_a="",
|
||||
endpoint_b="",
|
||||
winner=body.winner,
|
||||
is_blind=body.is_blind,
|
||||
blind_mapping=blind_mapping,
|
||||
voted_at=datetime.utcnow(),
|
||||
owner=user,
|
||||
)
|
||||
db.add(comp)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return {"status": "ok", "id": comp_id}
|
||||
|
||||
@router.get("/history")
|
||||
def list_comparisons(request: Request):
|
||||
"""List past comparisons."""
|
||||
user = get_current_user(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(Comparison)
|
||||
if user:
|
||||
q = q.filter(Comparison.owner == user)
|
||||
comps = q.order_by(Comparison.created_at.desc()).limit(50).all()
|
||||
return [
|
||||
{
|
||||
"id": c.id,
|
||||
"prompt": c.prompt[:100],
|
||||
"model_a": c.model_a,
|
||||
"model_b": c.model_b,
|
||||
"winner": c.winner,
|
||||
"is_blind": c.is_blind,
|
||||
"voted_at": c.voted_at.isoformat() if c.voted_at else None,
|
||||
"created_at": c.created_at.isoformat() if c.created_at else None,
|
||||
}
|
||||
for c in comps
|
||||
]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.delete("/{comp_id}")
|
||||
def delete_comparison(request: Request, comp_id: str):
|
||||
"""Delete a comparison and its ephemeral sessions."""
|
||||
user = get_current_user(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
comp = db.query(Comparison).filter(Comparison.id == comp_id).first()
|
||||
if not comp:
|
||||
raise HTTPException(404, "Comparison not found")
|
||||
# SECURITY: strict ownership — null-owner Comparisons were
|
||||
# accessible to every user.
|
||||
if user and comp.owner != user:
|
||||
raise HTTPException(404, "Comparison not found")
|
||||
db.delete(comp)
|
||||
db.commit()
|
||||
return {"status": "deleted"}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return router
|
||||
+14
-361
@@ -1,365 +1,18 @@
|
||||
# routes/compare_routes.py
|
||||
"""Model A/B comparison routes."""
|
||||
import json
|
||||
import uuid
|
||||
import random
|
||||
from datetime import datetime
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from typing import List
|
||||
from pydantic import BaseModel
|
||||
import logging
|
||||
"""Backward-compat shim — canonical location is routes/compare/compare_routes.py.
|
||||
|
||||
from core.database import Comparison, SessionLocal
|
||||
from core.session_manager import SessionManager
|
||||
from src.auth_helpers import get_current_user
|
||||
from routes.session_routes import _reject_raw_endpoint_url_for_non_admin
|
||||
This module is replaced in ``sys.modules`` by the canonical module object so
|
||||
that ``import routes.compare_routes``, ``from routes.compare_routes import X``,
|
||||
``importlib.import_module("routes.compare_routes")``, and the
|
||||
``import ... as cr`` + ``monkeypatch.setattr(cr, "SessionLocal", ...)`` /
|
||||
``"_owned_endpoint_by_url"`` / ``"_owned_endpoint_by_id"`` pattern used by
|
||||
test_endpoint_owner_scope_followup.py all operate on the *same* object the
|
||||
application actually uses. Keeps existing import paths working after
|
||||
slice 2i (#4082/#4071). Source-introspection tests read the canonical file
|
||||
by path.
|
||||
"""
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
import sys as _sys
|
||||
|
||||
router = APIRouter(prefix="/api/compare", tags=["compare"])
|
||||
from routes.compare import compare_routes as _canonical # noqa: F401
|
||||
|
||||
|
||||
def _owned_endpoint_by_url(db, base_url, owner):
|
||||
"""ModelEndpoint whose base_url == `base_url` and is VISIBLE to `owner`
|
||||
(their own rows + legacy null-owner "shared" rows); None otherwise.
|
||||
|
||||
Owner-scoped on purpose. ModelEndpoint is per-user (core/database.py: non-null
|
||||
owner = private, "the model picker only shows the endpoint to that user") and
|
||||
holds a decrypted `api_key`. start_comparison copies the matched row's api_key
|
||||
into the caller-owned [CMP] session's headers, which then drives that session's
|
||||
/api/chat_stream calls — so an UNSCOPED base_url match would let a user mint a
|
||||
comparison bound to ANOTHER user's private endpoint and spend that owner's
|
||||
api_key / reach whatever base_url they configured. Mirrors
|
||||
session_routes._owned_endpoint. A null/empty owner is a no-op (single-user /
|
||||
legacy mode).
|
||||
"""
|
||||
from core.database import ModelEndpoint
|
||||
from src.auth_helpers import owner_filter
|
||||
q = db.query(ModelEndpoint).filter(ModelEndpoint.base_url == base_url)
|
||||
return owner_filter(q, ModelEndpoint, owner).first()
|
||||
|
||||
|
||||
def _owned_endpoint_by_id(db, endpoint_id, owner):
|
||||
"""ModelEndpoint whose id == `endpoint_id` and is VISIBLE to `owner` (their
|
||||
own rows + legacy null-owner "shared" rows); None otherwise.
|
||||
|
||||
Preferred over _owned_endpoint_by_url for credential resolution: two visible
|
||||
endpoints can share the same base_url but hold DIFFERENT api_keys (e.g. two
|
||||
accounts on the same provider). A base_url-only match returns whichever row
|
||||
sorts first, so it can copy the WRONG owner-scoped key into the [CMP] session.
|
||||
An id pins the exact registered endpoint, so /api/compare/start prefers it and
|
||||
only falls back to URL matching for legacy / admin raw-URL callers. Owner
|
||||
scoping is identical to _owned_endpoint_by_url (a null/empty owner is a no-op).
|
||||
"""
|
||||
from core.database import ModelEndpoint
|
||||
from src.auth_helpers import owner_filter
|
||||
q = db.query(ModelEndpoint).filter(ModelEndpoint.id == endpoint_id)
|
||||
return owner_filter(q, ModelEndpoint, owner).first()
|
||||
|
||||
|
||||
class RecordVoteRequest(BaseModel):
|
||||
prompt: str
|
||||
models: List[str]
|
||||
winner: str # model name or "tie"
|
||||
is_blind: bool = True
|
||||
|
||||
|
||||
def setup_compare_routes(session_manager: SessionManager):
|
||||
"""Setup comparison routes."""
|
||||
|
||||
@router.post("/start")
|
||||
def start_comparison(
|
||||
request: Request,
|
||||
prompt: str = Form(...),
|
||||
model_a: str = Form(...),
|
||||
model_b: str = Form(...),
|
||||
endpoint_a: str = Form(""),
|
||||
endpoint_b: str = Form(""),
|
||||
endpoint_a_id: str = Form(""),
|
||||
endpoint_b_id: str = Form(""),
|
||||
is_blind: str = Form("true"),
|
||||
):
|
||||
"""Create two ephemeral sessions and a comparison record.
|
||||
|
||||
Returns the comparison ID and the two session IDs so the client
|
||||
can fire two independent SSE streams to /api/chat_stream.
|
||||
"""
|
||||
user = getattr(request.state, 'current_user', None)
|
||||
comp_id = str(uuid.uuid4())
|
||||
sid_a = str(uuid.uuid4())
|
||||
sid_b = str(uuid.uuid4())
|
||||
|
||||
# Blind mapping: randomly assign left/right
|
||||
blind = str(is_blind).lower() == "true"
|
||||
if blind:
|
||||
mapping = {"left": "a", "right": "b"}
|
||||
if random.random() > 0.5:
|
||||
mapping = {"left": "b", "right": "a"}
|
||||
else:
|
||||
mapping = {"left": "a", "right": "b"}
|
||||
|
||||
# Map session IDs to left/right based on blind mapping
|
||||
session_left = sid_a if mapping["left"] == "a" else sid_b
|
||||
session_right = sid_a if mapping["right"] == "a" else sid_b
|
||||
|
||||
# In blind mode, name the helper sessions by their neutral slot
|
||||
# ("Model A" / "Model B") instead of the real model. Otherwise the
|
||||
# session name leaks the model in the sidebar and GET /api/sessions,
|
||||
# de-anonymizing the comparison before the user votes (issue #1285).
|
||||
slot_name = {session_left: "Model A", session_right: "Model B"}
|
||||
|
||||
# SECURITY: resolve and validate BOTH endpoints before creating any
|
||||
# session. Compare copies a registered endpoint's Authorization header
|
||||
# into the [CMP] session, so validating one endpoint while creating its
|
||||
# session, then rejecting the other, would leave a partial compare
|
||||
# session behind with that header attached. Doing all the owner-scope
|
||||
# resolution + raw-URL rejection up front means a 403 on either endpoint
|
||||
# aborts the whole request with nothing created and no header copied.
|
||||
from src.endpoint_resolver import build_chat_url, build_headers, normalize_base
|
||||
resolved = []
|
||||
db = SessionLocal()
|
||||
try:
|
||||
for sid, model, endpoint, endpoint_id in [
|
||||
(sid_a, model_a, endpoint_a, endpoint_a_id),
|
||||
(sid_b, model_b, endpoint_b, endpoint_b_id),
|
||||
]:
|
||||
# Prefer an explicit endpoint id: it pins the EXACT registered
|
||||
# endpoint (and its api_key), even when two endpoints visible to
|
||||
# the caller share a base_url with different keys — a URL-only
|
||||
# match would copy whichever row sorts first, i.e. possibly the
|
||||
# wrong key. Fall back to URL resolution only for legacy / admin
|
||||
# raw-URL callers that don't send an id.
|
||||
eid = endpoint_id.strip() if isinstance(endpoint_id, str) else ""
|
||||
if eid:
|
||||
ep = _owned_endpoint_by_id(db, eid, user)
|
||||
if ep is None:
|
||||
# An id the caller can't see (wrong owner / deleted) must
|
||||
# NOT silently fall back to a same-URL row with a different
|
||||
# key — that's exactly the mix-up ids exist to prevent.
|
||||
raise HTTPException(404, "Model endpoint not found")
|
||||
# The id already resolved the endpoint; ignore any raw URL the
|
||||
# caller also sent and dial the stored config instead.
|
||||
endpoint = ep.base_url
|
||||
elif not endpoint:
|
||||
raise HTTPException(
|
||||
422, "endpoint_a/endpoint_b or endpoint_a_id/endpoint_b_id is required"
|
||||
)
|
||||
else:
|
||||
# Resolve the supplied URL to a ModelEndpoint the caller owns
|
||||
# (their own rows + legacy null-owner shared rows), scoped so a
|
||||
# comparison can't borrow another user's private endpoint key.
|
||||
base = normalize_base(endpoint)
|
||||
ep = _owned_endpoint_by_url(db, base, user)
|
||||
# Reject *unregistered* raw URLs for signed-in non-admins; a
|
||||
# matched registered endpoint supplies an id so the caller can
|
||||
# still compare endpoints they own. Blanket-rejecting here (the
|
||||
# earlier `endpoint_id=None` call) locked non-admins out of
|
||||
# compare entirely, since compare resolves endpoints by URL with
|
||||
# no endpoint_id. Mirrors the gallery inpaint/harmonize checks.
|
||||
# Raised here (phase 1), before any session exists.
|
||||
_reject_raw_endpoint_url_for_non_admin(
|
||||
request, user, str(ep.id) if ep is not None else None, endpoint
|
||||
)
|
||||
# Bind the [CMP] session to the RESOLVED endpoint, not the raw
|
||||
# caller-supplied string. When the URL matches a registered
|
||||
# endpoint visible to the caller, use that row's own normalized
|
||||
# base URL (the same value owner scoping + endpoint validation
|
||||
# already vetted) so the session dials exactly where the stored
|
||||
# config points. The raw `endpoint` only survives for callers
|
||||
# allowed to pass one — admins / single-user mode, where
|
||||
# `_reject_raw_endpoint_url_for_non_admin` is a no-op and `ep`
|
||||
# is None. Mirrors the registered-endpoint path in session_routes.
|
||||
session_endpoint_url = (
|
||||
build_chat_url(normalize_base(ep.base_url)) if ep is not None else endpoint
|
||||
)
|
||||
# Headers come only from a matched endpoint's key; None when
|
||||
# `ep` is None (raw admin URL or no match), so a comparison can
|
||||
# never inherit another user's key/headers.
|
||||
headers = build_headers(ep.api_key, ep.base_url) if (ep and ep.api_key) else None
|
||||
resolved.append((sid, model, session_endpoint_url, headers))
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# Both endpoints validated — only now create the ephemeral [CMP]
|
||||
# sessions and copy any resolved headers.
|
||||
for sid, model, session_endpoint_url, headers in resolved:
|
||||
name = f"[CMP] {slot_name[sid]}" if blind else f"[CMP] {model.split('/')[-1]}"
|
||||
session_manager.create_session(
|
||||
session_id=sid,
|
||||
name=name,
|
||||
endpoint_url=session_endpoint_url,
|
||||
model=model,
|
||||
rag=False,
|
||||
owner=user,
|
||||
)
|
||||
if headers:
|
||||
s = session_manager.sessions.get(sid)
|
||||
if s:
|
||||
s.headers = headers
|
||||
|
||||
# Store comparison record
|
||||
db = SessionLocal()
|
||||
try:
|
||||
comp = Comparison(
|
||||
id=comp_id,
|
||||
prompt=prompt,
|
||||
model_a=model_a,
|
||||
model_b=model_b,
|
||||
# Record the URL the session actually dials. For URL callers this
|
||||
# is their raw input; for id-only callers (empty endpoint_a/_b)
|
||||
# fall back to the resolved endpoint URL so the column stays
|
||||
# meaningful and non-null. resolved is in [a, b] order.
|
||||
endpoint_a=endpoint_a or resolved[0][2],
|
||||
endpoint_b=endpoint_b or resolved[1][2],
|
||||
is_blind=blind,
|
||||
blind_mapping=json.dumps(mapping),
|
||||
owner=user,
|
||||
)
|
||||
db.add(comp)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# In blind mode, withhold the model identities AND the left/right
|
||||
# mapping from the response. The client already knows model_a/model_b
|
||||
# (it sent them), so returning either would defeat blind mode. They are
|
||||
# revealed by POST /api/compare/{id}/vote once the user has voted (#1285).
|
||||
return {
|
||||
"id": comp_id,
|
||||
"session_left": session_left,
|
||||
"session_right": session_right,
|
||||
"model_left": None if blind else (model_a if mapping["left"] == "a" else model_b),
|
||||
"model_right": None if blind else (model_a if mapping["right"] == "a" else model_b),
|
||||
"is_blind": blind,
|
||||
"mapping": None if blind else mapping,
|
||||
}
|
||||
|
||||
@router.post("/{comp_id}/vote")
|
||||
def vote_comparison(
|
||||
request: Request,
|
||||
comp_id: str,
|
||||
winner: str = Form(...), # "left", "right", or "tie"
|
||||
):
|
||||
"""Record the user's vote and reveal model names if blind."""
|
||||
user = get_current_user(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
comp = db.query(Comparison).filter(Comparison.id == comp_id).first()
|
||||
if not comp:
|
||||
raise HTTPException(404, "Comparison not found")
|
||||
# SECURITY: strict ownership — null-owner Comparisons were
|
||||
# accessible to every user.
|
||||
if user and comp.owner != user:
|
||||
raise HTTPException(404, "Comparison not found")
|
||||
if comp.winner:
|
||||
raise HTTPException(400, "Already voted")
|
||||
|
||||
mapping = json.loads(comp.blind_mapping) if comp.blind_mapping else {"left": "a", "right": "b"}
|
||||
|
||||
if winner == "tie":
|
||||
comp.winner = "tie"
|
||||
elif winner == "left":
|
||||
comp.winner = mapping["left"]
|
||||
elif winner == "right":
|
||||
comp.winner = mapping["right"]
|
||||
else:
|
||||
raise HTTPException(400, "winner must be 'left', 'right', or 'tie'")
|
||||
|
||||
comp.voted_at = datetime.utcnow()
|
||||
db.commit()
|
||||
|
||||
return {
|
||||
"winner": comp.winner,
|
||||
"model_a": comp.model_a,
|
||||
"model_b": comp.model_b,
|
||||
"revealed": {
|
||||
"left": comp.model_a if mapping["left"] == "a" else comp.model_b,
|
||||
"right": comp.model_a if mapping["right"] == "a" else comp.model_b,
|
||||
},
|
||||
}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.post("/record")
|
||||
def record_comparison(request: Request, body: RecordVoteRequest):
|
||||
"""Lightweight endpoint to record a comparison vote from the frontend."""
|
||||
user = get_current_user(request)
|
||||
comp_id = str(uuid.uuid4())
|
||||
|
||||
model_a = body.models[0] if len(body.models) > 0 else ""
|
||||
model_b = body.models[1] if len(body.models) > 1 else ""
|
||||
|
||||
# For N>2 models, store the full list as JSON in blind_mapping
|
||||
if len(body.models) > 2:
|
||||
blind_mapping = json.dumps({"models": body.models})
|
||||
else:
|
||||
blind_mapping = None
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
comp = Comparison(
|
||||
id=comp_id,
|
||||
prompt=body.prompt[:500],
|
||||
model_a=model_a,
|
||||
model_b=model_b,
|
||||
endpoint_a="",
|
||||
endpoint_b="",
|
||||
winner=body.winner,
|
||||
is_blind=body.is_blind,
|
||||
blind_mapping=blind_mapping,
|
||||
voted_at=datetime.utcnow(),
|
||||
owner=user,
|
||||
)
|
||||
db.add(comp)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return {"status": "ok", "id": comp_id}
|
||||
|
||||
@router.get("/history")
|
||||
def list_comparisons(request: Request):
|
||||
"""List past comparisons."""
|
||||
user = get_current_user(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(Comparison)
|
||||
if user:
|
||||
q = q.filter(Comparison.owner == user)
|
||||
comps = q.order_by(Comparison.created_at.desc()).limit(50).all()
|
||||
return [
|
||||
{
|
||||
"id": c.id,
|
||||
"prompt": c.prompt[:100],
|
||||
"model_a": c.model_a,
|
||||
"model_b": c.model_b,
|
||||
"winner": c.winner,
|
||||
"is_blind": c.is_blind,
|
||||
"voted_at": c.voted_at.isoformat() if c.voted_at else None,
|
||||
"created_at": c.created_at.isoformat() if c.created_at else None,
|
||||
}
|
||||
for c in comps
|
||||
]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.delete("/{comp_id}")
|
||||
def delete_comparison(request: Request, comp_id: str):
|
||||
"""Delete a comparison and its ephemeral sessions."""
|
||||
user = get_current_user(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
comp = db.query(Comparison).filter(Comparison.id == comp_id).first()
|
||||
if not comp:
|
||||
raise HTTPException(404, "Comparison not found")
|
||||
# SECURITY: strict ownership — null-owner Comparisons were
|
||||
# accessible to every user.
|
||||
if user and comp.owner != user:
|
||||
raise HTTPException(404, "Comparison not found")
|
||||
db.delete(comp)
|
||||
db.commit()
|
||||
return {"status": "deleted"}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return router
|
||||
_sys.modules[__name__] = _canonical
|
||||
|
||||
@@ -463,14 +463,22 @@ def _cached_model_scan_script(model_dirs: list[str] | None = None, add_hf_cache:
|
||||
" if sz == 0 and os.path.isdir(snap):",
|
||||
" sz2, nf2, ic2 = snapshot_size()",
|
||||
" sz, nf, ic = sz2, nf2, ic or ic2",
|
||||
" is_diffusion = False; gguf_files = []",
|
||||
" is_video = bool(re.search(r'(?i)(^|/)Lightricks/LTX-|(^|/)LTX[-_/]|video|text-to-video|image-to-video', rid))",
|
||||
" is_diffusion = is_video; is_adapter = bool(re.search(r'(?i)(lora|adapter|peft|qlora|control[-_]?lora|diffusion[-_]?lora)', rid)); gguf_files = []",
|
||||
" if os.path.isdir(snap):",
|
||||
" for sd in os.listdir(snap):",
|
||||
" sf = os.path.join(snap, sd)",
|
||||
" if not os.path.isdir(sf): continue",
|
||||
" if os.path.exists(os.path.join(sf, 'model_index.json')): is_diffusion = True",
|
||||
" if os.path.exists(os.path.join(sf, 'adapter_config.json')) or os.path.exists(os.path.join(sf, 'adapter_model.safetensors')): is_adapter = True",
|
||||
" for _root, _dirs, _fns in safe_walk(sf):",
|
||||
" for _fn in _fns:",
|
||||
" _lfn = _fn.lower()",
|
||||
" if _lfn.endswith('.safetensors') and re.search(r'(?i)(ltx|video|upscaler)', _lfn): is_video = True; is_diffusion = True",
|
||||
" if _lfn in ('adapter_config.json','adapter_model.safetensors','pytorch_lora_weights.safetensors') or 'lora' in _lfn:",
|
||||
" is_adapter = True",
|
||||
" for f in collect_ggufs(sf): f['rel_path'] = sd + '/' + f['rel_path']; gguf_files.append(f)",
|
||||
" models.append({'repo_id':rid,'size_bytes':sz,'nb_files':nf,'has_incomplete':ic,'path':cache,'is_diffusion':is_diffusion,'is_gguf':bool(gguf_files),'gguf_files':gguf_files})",
|
||||
" models.append({'repo_id':rid,'size_bytes':sz,'nb_files':nf,'has_incomplete':ic,'path':cache,'is_diffusion':is_diffusion,'is_video':is_video,'is_adapter':is_adapter,'is_gguf':bool(gguf_files),'gguf_files':gguf_files})",
|
||||
"def hf_cache_paths():",
|
||||
" candidates = []",
|
||||
" def add(p):",
|
||||
@@ -505,11 +513,12 @@ def _cached_model_scan_script(model_dirs: list[str] | None = None, add_hf_cache:
|
||||
" fp = os.path.join(p, d)",
|
||||
" if not os.path.isdir(fp) or os.path.islink(fp) or not safe_path(fp): continue",
|
||||
" if d in seen: continue",
|
||||
" is_model = False; gguf_files = []",
|
||||
" is_model = False; is_adapter = bool(re.search(r'(?i)(lora|adapter|peft|qlora|control[-_]?lora|diffusion[-_]?lora)', d)); gguf_files = []",
|
||||
" for root, dirs, fns in safe_walk(fp):",
|
||||
" for fn in fns:",
|
||||
" if fn.lower().endswith('.gguf'): is_model = True",
|
||||
" elif fn == 'config.json' or fn.endswith('.safetensors') or fn.endswith('.bin'): is_model = True",
|
||||
" if fn in ('adapter_config.json','adapter_model.safetensors','pytorch_lora_weights.safetensors') or 'lora' in fn.lower(): is_adapter = True",
|
||||
" if is_model: break",
|
||||
" if not is_model: continue",
|
||||
" gguf_files = collect_ggufs(fp)",
|
||||
@@ -520,7 +529,7 @@ def _cached_model_scan_script(model_dirs: list[str] | None = None, add_hf_cache:
|
||||
" try: nf += 1; sz += os.path.getsize(os.path.join(dp, fn))",
|
||||
" except Exception: pass",
|
||||
" is_diff = os.path.exists(os.path.join(fp, 'model_index.json'))",
|
||||
" models.append({'repo_id':d,'size_bytes':sz,'nb_files':nf,'has_incomplete':False,'path':p,'is_local_dir':True,'is_diffusion':is_diff,'is_gguf':bool(gguf_files),'gguf_files':gguf_files})",
|
||||
" models.append({'repo_id':d,'size_bytes':sz,'nb_files':nf,'has_incomplete':False,'path':p,'is_local_dir':True,'is_diffusion':is_diff,'is_adapter':is_adapter,'is_gguf':bool(gguf_files),'gguf_files':gguf_files})",
|
||||
"def parse_size(num, unit):",
|
||||
" try: n = float(num)",
|
||||
" except Exception: return 0",
|
||||
@@ -1320,6 +1329,26 @@ def _diagnose_serve_output(text: str) -> dict | None:
|
||||
"MLX LM is not installed on this server.",
|
||||
[{"label": "install mlx-lm in Cookbook Dependencies", "op": "dependency", "package": "mlx-lm"}],
|
||||
),
|
||||
(
|
||||
r"OmniGen2Pipeline|module diffusers has no attribute .*Pipeline|custom_pipeline=.*failed",
|
||||
"This image model uses a custom Diffusers pipeline that the launch environment does not know yet.",
|
||||
[{"label": "update Diffusers image dependencies", "op": "dependency", "package": "diffusers transformers accelerate"}],
|
||||
),
|
||||
(
|
||||
r"mflux-generate-qwen.*not found|mflux-generate.*not found|MLX image serving requires mflux|No module named ['\"]?mflux",
|
||||
"MLX image serving requires mflux on this Apple Silicon server.",
|
||||
[{"label": "install mflux in Cookbook Dependencies", "op": "dependency", "package": "mflux"}],
|
||||
),
|
||||
(
|
||||
r"mlx-lama-swift|odysseus-mlx-inpaint|mlx-lama-serve|LaMa / MI-GAN MLX inpainting models require",
|
||||
"LaMa / MI-GAN MLX inpainting requires an Odysseus-compatible mlx-lama-swift bridge on this Apple Silicon server.",
|
||||
[{"label": "build mlx-lama-swift bridge and put odysseus-mlx-inpaint or mlx-lama-serve on PATH", "op": "dependency", "package": "mlx_lama_swift"}],
|
||||
),
|
||||
(
|
||||
r"mlx-ddcolor-swift|odysseus-mlx-colorize|mlx-ddcolor-serve|DDColor MLX models require",
|
||||
"DDColor MLX colorization requires an Odysseus-compatible mlx-ddcolor-swift bridge on this Apple Silicon server.",
|
||||
[{"label": "build mlx-ddcolor-swift bridge and put odysseus-mlx-colorize or mlx-ddcolor-serve on PATH", "op": "dependency", "package": "mlx_ddcolor_swift"}],
|
||||
),
|
||||
(
|
||||
r"Unable to quantize model of type <class ['\"]mlx_lm\.models\.switch_layers\.QuantizedSwitchLinear['\"]>|QuantizedSwitchLinear",
|
||||
"MLX-LM tried to quantize an already-quantized DeepSeek switch layer.",
|
||||
@@ -1358,9 +1387,9 @@ def _diagnose_serve_output(text: str) -> dict | None:
|
||||
[{"label": "download a GGUF build of this model (repo name usually ends in -GGUF, file like Q4_K_M.gguf)", "op": "manual"}],
|
||||
),
|
||||
(
|
||||
r"No module named 'torch'|No module named torch|No module named 'diffusers'|No module named diffusers",
|
||||
"Diffusion serving requires PyTorch and diffusers.",
|
||||
[{"label": "install diffusers[torch] in Cookbook Dependencies", "op": "dependency", "package": "diffusers[torch]"}],
|
||||
r"No module named 'torch'|No module named torch|No module named 'torchvision'|No module named torchvision|No module named 'diffusers'|No module named diffusers|No module named 'scipy'|No module named scipy|install scipy if you want to use beta sigmas|requires the Torchvision library",
|
||||
"Diffusion serving requires PyTorch, Torchvision, Diffusers, Accelerate, and SciPy.",
|
||||
[{"label": "install Diffusers image deps in Cookbook Dependencies", "op": "dependency", "package": "diffusers[torch] torchvision accelerate scipy python-multipart"}],
|
||||
),
|
||||
(
|
||||
r"403 Forbidden|401 Unauthorized|Access to model.*is restricted|gated repo|not in the authorized list|awaiting a review",
|
||||
|
||||
+187
-28
@@ -73,6 +73,23 @@ _HF_TOKEN_STATUS_SNIPPET = (
|
||||
)
|
||||
|
||||
|
||||
def _append_mlx_image_server_script(runner_lines: list[str]) -> None:
|
||||
"""Write the MLX image API helper next to the tmux runner on remote hosts."""
|
||||
script_path = Path(__file__).resolve().parents[1] / "scripts" / "mlx_image_server.py"
|
||||
try:
|
||||
script = script_path.read_text(encoding="utf-8")
|
||||
except Exception as e:
|
||||
logger.warning("Failed to read mlx_image_server.py: %s", e)
|
||||
runner_lines.append('echo "ERROR: Odysseus could not prepare the MLX image server helper."')
|
||||
runner_lines.append('ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
return
|
||||
runner_lines.append('mkdir -p scripts')
|
||||
runner_lines.append("cat > scripts/mlx_image_server.py <<'PY'")
|
||||
runner_lines.extend(script.splitlines())
|
||||
runner_lines.append("PY")
|
||||
runner_lines.append('chmod +x scripts/mlx_image_server.py 2>/dev/null || true')
|
||||
|
||||
|
||||
def _venv_root_from_serve_cmd(cmd: str) -> str:
|
||||
"""Best-effort venv root from an absolute venv python in a serve command."""
|
||||
try:
|
||||
@@ -492,6 +509,11 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
"MLX LM is not installed on this server.",
|
||||
[{"label": "install mlx-lm in Cookbook Dependencies", "op": "dependency", "package": "mlx-lm"}],
|
||||
),
|
||||
(
|
||||
r"OmniGen2Pipeline|module diffusers has no attribute .*Pipeline|custom_pipeline=.*failed",
|
||||
"This image model uses a custom Diffusers pipeline that the launch environment does not know yet.",
|
||||
[{"label": "update Diffusers image dependencies", "op": "dependency", "package": "diffusers transformers accelerate"}],
|
||||
),
|
||||
(
|
||||
r"Unable to quantize model of type <class ['\"]mlx_lm\.models\.switch_layers\.QuantizedSwitchLinear['\"]>|QuantizedSwitchLinear",
|
||||
"MLX-LM tried to quantize an already-quantized DeepSeek switch layer.",
|
||||
@@ -530,9 +552,9 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
[{"label": "download a GGUF build of this model (repo name usually ends in -GGUF, file like Q4_K_M.gguf)", "op": "manual"}],
|
||||
),
|
||||
(
|
||||
r"No module named 'torch'|No module named torch|No module named 'diffusers'|No module named diffusers",
|
||||
"Diffusion serving requires PyTorch and diffusers.",
|
||||
[{"label": "install diffusers[torch] in Cookbook Dependencies", "op": "dependency", "package": "diffusers[torch]"}],
|
||||
r"No module named 'torch'|No module named torch|No module named 'torchvision'|No module named torchvision|No module named 'diffusers'|No module named diffusers|No module named 'scipy'|No module named scipy|install scipy if you want to use beta sigmas|requires the Torchvision library",
|
||||
"Diffusion serving requires PyTorch, Torchvision, Diffusers, Accelerate, and SciPy.",
|
||||
[{"label": "install Diffusers image deps in Cookbook Dependencies", "op": "dependency", "package": "diffusers[torch] torchvision accelerate scipy python-multipart"}],
|
||||
),
|
||||
(
|
||||
r"403 Forbidden|401 Unauthorized|Access to model.*is restricted|gated repo|not in the authorized list|awaiting a review",
|
||||
@@ -1433,9 +1455,11 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
"nb_files": m["nb_files"],
|
||||
"has_incomplete": m["has_incomplete"],
|
||||
"status": "downloading" if m["has_incomplete"] else "ready",
|
||||
"path": m.get("path", ""),
|
||||
"is_diffusion": m.get("is_diffusion", False),
|
||||
}
|
||||
"path": m.get("path", ""),
|
||||
"is_diffusion": m.get("is_diffusion", False),
|
||||
"is_video": m.get("is_video", False),
|
||||
"is_adapter": m.get("is_adapter", False),
|
||||
}
|
||||
if m.get("is_local_dir"):
|
||||
entry["is_local_dir"] = True
|
||||
if m.get("is_gguf"):
|
||||
@@ -1460,6 +1484,7 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
"""Register a diffusion model as an image endpoint so it appears in the model selector."""
|
||||
import re
|
||||
from core.database import SessionLocal, ModelEndpoint
|
||||
from src.settings import load_settings, save_settings
|
||||
|
||||
# Parse port from command (--port NNNN), default 8100 for diffusion_server
|
||||
port_match = re.search(r'--port\s+(\d+)', req.cmd)
|
||||
@@ -1477,6 +1502,7 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
# Friendly display name from repo_id
|
||||
short_name = req.repo_id.split("/")[-1] if "/" in req.repo_id else req.repo_id
|
||||
display_name = f"{short_name} (image)"
|
||||
pinned_models = [req.repo_id] if req.repo_id else []
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
@@ -1486,7 +1512,16 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
existing.is_enabled = True
|
||||
existing.model_type = "image"
|
||||
existing.name = display_name
|
||||
existing.endpoint_kind = "local"
|
||||
existing.model_refresh_mode = "manual"
|
||||
if pinned_models:
|
||||
existing.cached_models = json.dumps(pinned_models)
|
||||
existing.pinned_models = json.dumps(pinned_models)
|
||||
db.commit()
|
||||
settings = load_settings()
|
||||
if settings.get("image_gen_enabled") is not True:
|
||||
settings["image_gen_enabled"] = True
|
||||
save_settings(settings)
|
||||
logger.info(f"Updated existing image endpoint: {base_url}")
|
||||
return existing.id
|
||||
|
||||
@@ -1498,9 +1533,18 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
api_key=None,
|
||||
is_enabled=True,
|
||||
model_type="image",
|
||||
endpoint_kind="local",
|
||||
model_refresh_mode="manual",
|
||||
cached_models=json.dumps(pinned_models) if pinned_models else None,
|
||||
pinned_models=json.dumps(pinned_models) if pinned_models else None,
|
||||
)
|
||||
db.add(ep)
|
||||
db.commit()
|
||||
settings = load_settings()
|
||||
settings["image_gen_enabled"] = True
|
||||
if not settings.get("image_model"):
|
||||
settings["image_model"] = req.repo_id
|
||||
save_settings(settings)
|
||||
logger.info(f"Auto-registered image endpoint: {display_name} @ {base_url}")
|
||||
return ep_id
|
||||
except Exception as e:
|
||||
@@ -2356,24 +2400,8 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
runner_lines.append('fi')
|
||||
elif "sglang.launch_server" in req.cmd:
|
||||
runner_lines.append('export PATH="$HOME/.local/bin:$PATH"')
|
||||
runner_lines.append('if ! command -v sglang &>/dev/null; then')
|
||||
runner_lines.append(' echo "ERROR: SGLang is not installed."')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append('elif ! ODYSSEUS_SGLANG_IMPORT_ERROR="$(python3 -c "import sglang" 2>&1)"; then')
|
||||
runner_lines.append(' echo "ERROR: SGLang is installed but failed to import."')
|
||||
runner_lines.append(' printf "%s\\n" "$ODYSSEUS_SGLANG_IMPORT_ERROR"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append('fi')
|
||||
elif "mlx_lm.server" in req.cmd:
|
||||
runner_lines.append('export PATH="$HOME/.local/bin:/opt/homebrew/bin:/usr/local/bin:$PATH"')
|
||||
runner_lines.append('if ! ODYSSEUS_MLX_IMPORT_ERROR="$(python3 -c "import mlx_lm" 2>&1)"; then')
|
||||
runner_lines.append(' echo "ERROR: MLX LM is not installed."')
|
||||
runner_lines.append(' printf "%s\\n" "$ODYSSEUS_MLX_IMPORT_ERROR"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append('fi')
|
||||
runner_lines.append(f"ODYSSEUS_SERVE_CMD='{_bash_squote(req.cmd)}'")
|
||||
runner_lines.append('if [ -z "$ODYSSEUS_PREFLIGHT_EXIT" ]; then')
|
||||
runner_lines.append(' ODYSSEUS_MLX_CMD_PY="$(python3 - "$ODYSSEUS_SERVE_CMD" <<\'PY\'')
|
||||
runner_lines.append('ODYSSEUS_SGLANG_CMD_PY="$(python3 - "$ODYSSEUS_SERVE_CMD" <<\'PY\'')
|
||||
runner_lines.append('import shlex, sys')
|
||||
runner_lines.append('parts = shlex.split(sys.argv[1])')
|
||||
runner_lines.append('py = "python3"')
|
||||
@@ -2384,6 +2412,36 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
runner_lines.append('print(py)')
|
||||
runner_lines.append('PY')
|
||||
runner_lines.append(')"')
|
||||
runner_lines.append('if ! "$ODYSSEUS_SGLANG_CMD_PY" -c "import sglang" &>/dev/null; then')
|
||||
runner_lines.append(' if ! command -v sglang &>/dev/null; then')
|
||||
runner_lines.append(' echo "ERROR: SGLang is not installed."')
|
||||
runner_lines.append(' else')
|
||||
runner_lines.append(' echo "ERROR: SGLang is installed but failed to import in the launch Python."')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append(' ODYSSEUS_SGLANG_IMPORT_ERROR="$("$ODYSSEUS_SGLANG_CMD_PY" -c "import sglang" 2>&1)"')
|
||||
runner_lines.append(' printf "%s\\n" "$ODYSSEUS_SGLANG_IMPORT_ERROR"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append('fi')
|
||||
elif "mlx_lm.server" in req.cmd:
|
||||
runner_lines.append('export PATH="$HOME/.local/bin:/opt/homebrew/bin:/usr/local/bin:$PATH"')
|
||||
runner_lines.append(f"ODYSSEUS_SERVE_CMD='{_bash_squote(req.cmd)}'")
|
||||
runner_lines.append('ODYSSEUS_MLX_CMD_PY="$(python3 - "$ODYSSEUS_SERVE_CMD" <<\'PY\'')
|
||||
runner_lines.append('import shlex, sys')
|
||||
runner_lines.append('parts = shlex.split(sys.argv[1])')
|
||||
runner_lines.append('py = "python3"')
|
||||
runner_lines.append('for i, part in enumerate(parts):')
|
||||
runner_lines.append(' if part.endswith("/bin/python") or part.endswith("/bin/python3") or "/bin/python3." in part:')
|
||||
runner_lines.append(' py = part')
|
||||
runner_lines.append(' break')
|
||||
runner_lines.append('print(py)')
|
||||
runner_lines.append('PY')
|
||||
runner_lines.append(')"')
|
||||
runner_lines.append('if ! ODYSSEUS_MLX_IMPORT_ERROR="$("$ODYSSEUS_MLX_CMD_PY" -c "import mlx_lm" 2>&1)"; then')
|
||||
runner_lines.append(' echo "ERROR: MLX LM is not installed in the launch Python: $ODYSSEUS_MLX_CMD_PY"')
|
||||
runner_lines.append(' printf "%s\\n" "$ODYSSEUS_MLX_IMPORT_ERROR"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append('fi')
|
||||
runner_lines.append('if [ -z "$ODYSSEUS_PREFLIGHT_EXIT" ]; then')
|
||||
runner_lines.append(' ODYSSEUS_SERVE_CMD="$("$ODYSSEUS_MLX_CMD_PY" - "$ODYSSEUS_SERVE_CMD" <<\'PY\'')
|
||||
runner_lines.append('import json, os, shlex, sys')
|
||||
runner_lines.append('from pathlib import Path')
|
||||
@@ -2474,10 +2532,111 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
runner_lines.append('PY')
|
||||
runner_lines.append(')"')
|
||||
runner_lines.append('fi')
|
||||
elif "scripts/mlx_image_server.py" in req.cmd or ".mlx_image_server.py" in req.cmd:
|
||||
_append_mlx_image_server_script(runner_lines)
|
||||
runner_lines.append('export PATH="$HOME/.local/bin:/opt/homebrew/bin:/usr/local/bin:$PATH"')
|
||||
runner_lines.append(f"ODYSSEUS_SERVE_CMD='{_bash_squote(req.cmd)}'")
|
||||
runner_lines.append('ODYSSEUS_MLX_IMAGE_CMD_PY="$(python3 - "$ODYSSEUS_SERVE_CMD" <<\'PY\'')
|
||||
runner_lines.append('import shlex, sys')
|
||||
runner_lines.append('parts = shlex.split(sys.argv[1])')
|
||||
runner_lines.append('py = "python3"')
|
||||
runner_lines.append('for part in parts:')
|
||||
runner_lines.append(' if part.endswith("/bin/python") or part.endswith("/bin/python3") or "/bin/python3." in part:')
|
||||
runner_lines.append(' py = part')
|
||||
runner_lines.append(' break')
|
||||
runner_lines.append('print(py)')
|
||||
runner_lines.append('PY')
|
||||
runner_lines.append(')"')
|
||||
runner_lines.append('ODYSSEUS_MLX_IMAGE_BIN_DIR="$(dirname "$ODYSSEUS_MLX_IMAGE_CMD_PY" 2>/dev/null || true)"')
|
||||
runner_lines.append('if [ -n "$ODYSSEUS_MLX_IMAGE_BIN_DIR" ]; then export PATH="$ODYSSEUS_MLX_IMAGE_BIN_DIR:$PATH"; fi')
|
||||
runner_lines.append('if ! "$ODYSSEUS_MLX_IMAGE_CMD_PY" -c "import fastapi, uvicorn, multipart" >/dev/null 2>&1; then')
|
||||
runner_lines.append(' echo "ERROR: MLX image serving requires FastAPI + uvicorn + python-multipart in the launch Python: $ODYSSEUS_MLX_IMAGE_CMD_PY. Install the MLX image dependencies in Cookbook Dependencies."')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append('fi')
|
||||
runner_lines.append('ODYSSEUS_MLX_IMAGE_MODEL="$(python3 - "$ODYSSEUS_SERVE_CMD" <<\'PY\'')
|
||||
runner_lines.append('import shlex, sys')
|
||||
runner_lines.append('parts = shlex.split(sys.argv[1])')
|
||||
runner_lines.append('model = ""')
|
||||
runner_lines.append('for i, part in enumerate(parts):')
|
||||
runner_lines.append(' if part == "--model" and i + 1 < len(parts):')
|
||||
runner_lines.append(' model = parts[i + 1]')
|
||||
runner_lines.append(' break')
|
||||
runner_lines.append('print(model)')
|
||||
runner_lines.append('PY')
|
||||
runner_lines.append(')"')
|
||||
runner_lines.append('if printf "%s" "$ODYSSEUS_MLX_IMAGE_MODEL" | grep -qi hidream; then')
|
||||
runner_lines.append(' if ! "$ODYSSEUS_MLX_IMAGE_CMD_PY" -c "import mlx, mlx_vlm, transformers, huggingface_hub, safetensors, numpy, PIL" >/dev/null 2>&1; then')
|
||||
runner_lines.append(' echo "ERROR: HiDream MLX serving needs the model requirements in the launch Python: $ODYSSEUS_MLX_IMAGE_CMD_PY."')
|
||||
runner_lines.append(' echo "Install with: $ODYSSEUS_MLX_IMAGE_CMD_PY -m pip install -U fastapi uvicorn python-multipart mlx mlx-vlm \'transformers>=4.57.0,<6.0\' huggingface_hub safetensors numpy pillow tqdm sentencepiece hf_transfer"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append('elif printf "%s" "$ODYSSEUS_MLX_IMAGE_MODEL" | grep -qi boogu; then')
|
||||
runner_lines.append(' if ! "$ODYSSEUS_MLX_IMAGE_CMD_PY" -c "import boogu_image_mlx, mlx, huggingface_hub, safetensors, numpy, PIL" >/dev/null 2>&1; then')
|
||||
runner_lines.append(' echo "ERROR: Boogu MLX serving needs boogu-image-mlx in the launch Python: $ODYSSEUS_MLX_IMAGE_CMD_PY."')
|
||||
runner_lines.append(' echo "Install with: $ODYSSEUS_MLX_IMAGE_CMD_PY -m pip install -U git+https://github.com/xocialize/boogu-image-mlx.git fastapi uvicorn python-multipart pillow"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append('elif printf "%s" "$ODYSSEUS_MLX_IMAGE_MODEL" | grep -Eqi "ddcolor"; then')
|
||||
runner_lines.append(' if ! "$ODYSSEUS_MLX_IMAGE_CMD_PY" -c "import PIL" >/dev/null 2>&1; then')
|
||||
runner_lines.append(' echo "ERROR: DDColor MLX serving needs Pillow in the launch Python: $ODYSSEUS_MLX_IMAGE_CMD_PY."')
|
||||
runner_lines.append(' echo "Install with: $ODYSSEUS_MLX_IMAGE_CMD_PY -m pip install -U fastapi uvicorn python-multipart pillow huggingface_hub"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append(' if ! command -v odysseus-mlx-colorize >/dev/null 2>&1 && ! command -v mlx-ddcolor-serve >/dev/null 2>&1; then')
|
||||
runner_lines.append(' echo "ERROR: DDColor MLX serving requires the Odysseus mlx-ddcolor-swift bridge on PATH: odysseus-mlx-colorize or mlx-ddcolor-serve."')
|
||||
runner_lines.append(' echo "Build it from swift/odysseus-mlx-image-bridge in Cookbook Dependencies."')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append(' ODYSSEUS_DDCOLOR_BIN="$(command -v odysseus-mlx-colorize 2>/dev/null || command -v mlx-ddcolor-serve 2>/dev/null || true)"')
|
||||
runner_lines.append(' if [ -n "$ODYSSEUS_DDCOLOR_BIN" ]; then')
|
||||
runner_lines.append(' ODYSSEUS_DDCOLOR_DIR="$(dirname "$ODYSSEUS_DDCOLOR_BIN")"')
|
||||
runner_lines.append(' if [ ! -f "$ODYSSEUS_DDCOLOR_DIR/mlx.metallib" ] && [ ! -f "$ODYSSEUS_DDCOLOR_DIR/default.metallib" ]; then')
|
||||
runner_lines.append(' echo "ERROR: DDColor MLX serving found the Swift runner, but mlx.metallib/default.metallib is missing next to it."')
|
||||
runner_lines.append(' echo "Run the DDColor MLX image editing dependency install again; it copies mlx.metallib from the launch Python MLX package."')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append('elif printf "%s" "$ODYSSEUS_MLX_IMAGE_MODEL" | grep -Eqi "mi-gan|migan|lama"; then')
|
||||
runner_lines.append(' if ! "$ODYSSEUS_MLX_IMAGE_CMD_PY" -c "import PIL" >/dev/null 2>&1; then')
|
||||
runner_lines.append(' echo "ERROR: LaMa / MI-GAN MLX serving needs Pillow in the launch Python: $ODYSSEUS_MLX_IMAGE_CMD_PY."')
|
||||
runner_lines.append(' echo "Install with: $ODYSSEUS_MLX_IMAGE_CMD_PY -m pip install -U fastapi uvicorn python-multipart pillow huggingface_hub"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append(' if ! command -v odysseus-mlx-inpaint >/dev/null 2>&1 && ! command -v mlx-lama-serve >/dev/null 2>&1; then')
|
||||
runner_lines.append(' echo "ERROR: LaMa / MI-GAN MLX serving requires the Odysseus mlx-lama-swift bridge on PATH: odysseus-mlx-inpaint or mlx-lama-serve."')
|
||||
runner_lines.append(' echo "Build it from swift/odysseus-mlx-image-bridge in Cookbook Dependencies."')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append(' ODYSSEUS_INPAINT_BIN="$(command -v odysseus-mlx-inpaint 2>/dev/null || command -v mlx-lama-serve 2>/dev/null || true)"')
|
||||
runner_lines.append(' if [ -n "$ODYSSEUS_INPAINT_BIN" ]; then')
|
||||
runner_lines.append(' ODYSSEUS_INPAINT_DIR="$(dirname "$ODYSSEUS_INPAINT_BIN")"')
|
||||
runner_lines.append(' if [ ! -f "$ODYSSEUS_INPAINT_DIR/mlx.metallib" ] && [ ! -f "$ODYSSEUS_INPAINT_DIR/default.metallib" ]; then')
|
||||
runner_lines.append(' echo "ERROR: LaMa / MI-GAN MLX serving found the Swift runner, but mlx.metallib/default.metallib is missing next to it."')
|
||||
runner_lines.append(' echo "Run the LaMa / MI-GAN MLX image editing dependency install again; it copies mlx.metallib from the launch Python MLX package."')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append(' fi')
|
||||
runner_lines.append('elif ! command -v mflux-generate >/dev/null 2>&1 && ! command -v mflux-generate-qwen >/dev/null 2>&1; then')
|
||||
runner_lines.append(' echo "ERROR: mflux-compatible MLX image serving requires mflux-generate or mflux-generate-qwen in PATH for launch Python: $ODYSSEUS_MLX_IMAGE_CMD_PY."')
|
||||
runner_lines.append(' echo "Install with: $ODYSSEUS_MLX_IMAGE_CMD_PY -m pip install -U mflux fastapi uvicorn python-multipart"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append('fi')
|
||||
elif "scripts/diffusion_server.py" in req.cmd or ".diffusion_server.py" in req.cmd:
|
||||
runner_lines.append('export PATH="$HOME/.local/bin:$PATH"')
|
||||
runner_lines.append('if ! ODYSSEUS_DIFFUSION_IMPORT_ERROR="$(python3 -c "import torch, diffusers" 2>&1)"; then')
|
||||
runner_lines.append(' echo "ERROR: Diffusion serving requires PyTorch + diffusers."')
|
||||
runner_lines.append(f"ODYSSEUS_SERVE_CMD='{_bash_squote(req.cmd)}'")
|
||||
runner_lines.append('ODYSSEUS_DIFFUSION_CMD_PY="$(python3 - "$ODYSSEUS_SERVE_CMD" <<\'PY\'')
|
||||
runner_lines.append('import shlex, sys')
|
||||
runner_lines.append('parts = shlex.split(sys.argv[1])')
|
||||
runner_lines.append('py = "python3"')
|
||||
runner_lines.append('for part in parts:')
|
||||
runner_lines.append(' if part.endswith("/bin/python") or part.endswith("/bin/python3") or "/bin/python3." in part:')
|
||||
runner_lines.append(' py = part')
|
||||
runner_lines.append(' break')
|
||||
runner_lines.append('print(py)')
|
||||
runner_lines.append('PY')
|
||||
runner_lines.append(')"')
|
||||
runner_lines.append('if ! ODYSSEUS_DIFFUSION_IMPORT_ERROR="$("$ODYSSEUS_DIFFUSION_CMD_PY" -c "import torch, torchvision, diffusers" 2>&1)"; then')
|
||||
runner_lines.append(' echo "ERROR: Diffusion serving requires PyTorch + Torchvision + diffusers in the launch Python: $ODYSSEUS_DIFFUSION_CMD_PY."')
|
||||
runner_lines.append(' printf "%s\\n" "$ODYSSEUS_DIFFUSION_IMPORT_ERROR"')
|
||||
runner_lines.append(' ODYSSEUS_PREFLIGHT_EXIT=127')
|
||||
runner_lines.append('fi')
|
||||
@@ -2588,8 +2747,8 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
# endpoint; any other real model serve (i.e. not a pip-install task) gets
|
||||
# a local LLM endpoint pointed at its /v1.
|
||||
endpoint_id = None
|
||||
is_diffusion = "diffusion_server.py" in req.cmd
|
||||
if is_diffusion:
|
||||
is_image_endpoint = "diffusion_server.py" in req.cmd or "mlx_image_server.py" in req.cmd
|
||||
if is_image_endpoint:
|
||||
endpoint_id = _auto_register_image_endpoint(req, remote)
|
||||
elif not is_pip_install:
|
||||
endpoint_id = _auto_register_llm_endpoint(req, remote)
|
||||
@@ -2605,7 +2764,7 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
# if N != 0 within the watch window, delete the endpoint we just
|
||||
# created. Skipped for diffusion (different image-endpoint cleanup
|
||||
# path) and pip-install tasks (no endpoint to drop).
|
||||
if endpoint_id and not is_diffusion and not is_pip_install:
|
||||
if endpoint_id and not is_image_endpoint and not is_pip_install:
|
||||
asyncio.create_task(_serve_crash_watchdog(
|
||||
endpoint_id=endpoint_id,
|
||||
session_id=session_id,
|
||||
|
||||
@@ -12,6 +12,7 @@ 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__)
|
||||
|
||||
@@ -78,6 +79,14 @@ from routes.document_helpers import (
|
||||
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
|
||||
@@ -124,6 +133,7 @@ def setup_document_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
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
|
||||
@@ -636,6 +646,7 @@ def setup_document_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
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
|
||||
|
||||
+16
-3
@@ -246,6 +246,7 @@ import re as _re_reply
|
||||
# serves replies and summaries (any fenced final-output block).
|
||||
_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"</?\|(?:assistant|assistan|user|system|tool)\|>?|</\|end\|>?", _re_reply.I)
|
||||
|
||||
|
||||
def _extract_reply(text: str) -> str:
|
||||
@@ -272,6 +273,7 @@ def _extract_reply(text: str) -> str:
|
||||
# Drop any stray/duplicate marker tokens, then strip think markup.
|
||||
t = _REPLY_OPEN_RE.sub("", t)
|
||||
t = _REPLY_CLOSE_RE.sub("", t)
|
||||
t = _REPLY_ROLE_MARKER_RE.sub("", t)
|
||||
return _strip_think(t).strip()
|
||||
|
||||
|
||||
@@ -1035,13 +1037,23 @@ def _coerce_imap_timeout_seconds(raw: str | None) -> int:
|
||||
_IMAP_TIMEOUT_SECONDS = _coerce_imap_timeout_seconds(os.environ.get("ODYSSEUS_IMAP_TIMEOUT_SECONDS"))
|
||||
|
||||
|
||||
def _open_imap_connection(host: str, port: int, *, starttls: bool, timeout: int = _IMAP_TIMEOUT_SECONDS):
|
||||
def _open_imap_connection(
|
||||
host: str,
|
||||
port: int,
|
||||
*,
|
||||
starttls: bool,
|
||||
timeout: int = _IMAP_TIMEOUT_SECONDS,
|
||||
ssl_context=None,
|
||||
):
|
||||
"""Open an IMAP connection using the configured security mode."""
|
||||
port = int(port or 993)
|
||||
if starttls:
|
||||
conn = imaplib.IMAP4(host, port, timeout=timeout)
|
||||
try:
|
||||
conn.starttls()
|
||||
if ssl_context:
|
||||
conn.starttls(ssl_context=ssl_context)
|
||||
else:
|
||||
conn.starttls()
|
||||
except Exception:
|
||||
# Don't leak the open plain socket if the STARTTLS upgrade is
|
||||
# rejected; close it before propagating. (#3174)
|
||||
@@ -1051,7 +1063,8 @@ def _open_imap_connection(host: str, port: int, *, starttls: bool, timeout: int
|
||||
pass
|
||||
raise
|
||||
elif port == 993:
|
||||
conn = imaplib.IMAP4_SSL(host, port, timeout=timeout)
|
||||
kwargs = {"ssl_context": ssl_context} if ssl_context else {}
|
||||
conn = imaplib.IMAP4_SSL(host, port, timeout=timeout, **kwargs)
|
||||
else:
|
||||
conn = imaplib.IMAP4(host, port, timeout=timeout)
|
||||
try:
|
||||
|
||||
+370
-25
@@ -100,6 +100,272 @@ def _owner_for_email_account(account_id: str | None) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
def _email_date_only(value: str | None):
|
||||
value = (value or "").strip()
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
return datetime.strptime(value[:10], "%Y-%m-%d").date()
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
_AUTO_REPLY_KEYS = {
|
||||
"email_auto_reply",
|
||||
"email_auto_reply_start",
|
||||
"email_auto_reply_end",
|
||||
"email_auto_reply_subject",
|
||||
"email_auto_reply_message",
|
||||
"email_auto_reply_cooldown",
|
||||
"email_auto_reply_scope",
|
||||
"email_auto_reply_account_id",
|
||||
"email_auto_reply_exclude_automated",
|
||||
"email_auto_reply_pause_notifications",
|
||||
"email_auto_reply_enabled_at",
|
||||
}
|
||||
|
||||
|
||||
def _effective_settings_for_email_account(settings: dict, account_id: str | None) -> dict:
|
||||
"""Overlay per-account auto-reply settings onto global settings.
|
||||
|
||||
Other automation toggles remain global. This lets each mailbox have its own
|
||||
away reply while preserving existing installs that only have global keys.
|
||||
"""
|
||||
effective = dict(settings or {})
|
||||
key = str(account_id or "").strip()
|
||||
by_account = effective.get("email_auto_reply_by_account") or {}
|
||||
account_cfg = by_account.get(key) if key and isinstance(by_account, dict) else None
|
||||
if isinstance(account_cfg, dict):
|
||||
for k in _AUTO_REPLY_KEYS:
|
||||
if k in account_cfg:
|
||||
effective[k] = account_cfg[k]
|
||||
return effective
|
||||
|
||||
|
||||
def _away_reply_active(settings: dict, account_id: str | None) -> bool:
|
||||
if not settings.get("email_auto_reply", False):
|
||||
return False
|
||||
|
||||
scope = str(settings.get("email_auto_reply_scope") or "all").strip().lower()
|
||||
if scope == "account":
|
||||
selected = str(settings.get("email_auto_reply_account_id") or "").strip()
|
||||
if selected and selected != str(account_id or ""):
|
||||
return False
|
||||
|
||||
today = datetime.utcnow().date()
|
||||
start = _email_date_only(settings.get("email_auto_reply_start"))
|
||||
end = _email_date_only(settings.get("email_auto_reply_end"))
|
||||
if start and today < start:
|
||||
return False
|
||||
if end and today > end:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _message_after_away_enabled(settings: dict, msg) -> bool:
|
||||
enabled_at = (settings.get("email_auto_reply_enabled_at") or "").strip()
|
||||
if not enabled_at:
|
||||
# Existing installs may already have the toggle on before this feature
|
||||
# existed. Do not back-reply old mail until the user saves/toggles it.
|
||||
return False
|
||||
try:
|
||||
enabled_dt = datetime.fromisoformat(enabled_at.replace("Z", "+00:00"))
|
||||
except Exception:
|
||||
return False
|
||||
try:
|
||||
msg_dt = email.utils.parsedate_to_datetime(msg.get("Date", ""))
|
||||
except Exception:
|
||||
return False
|
||||
try:
|
||||
if enabled_dt.tzinfo and not msg_dt.tzinfo:
|
||||
msg_dt = msg_dt.replace(tzinfo=enabled_dt.tzinfo)
|
||||
elif msg_dt.tzinfo and not enabled_dt.tzinfo:
|
||||
enabled_dt = enabled_dt.replace(tzinfo=msg_dt.tzinfo)
|
||||
except Exception:
|
||||
pass
|
||||
return msg_dt >= enabled_dt
|
||||
|
||||
|
||||
def _away_reply_period_key(settings: dict) -> str:
|
||||
start = (settings.get("email_auto_reply_start") or "").strip()
|
||||
end = (settings.get("email_auto_reply_end") or "").strip()
|
||||
return f"{start or '*'}..{end or '*'}"
|
||||
|
||||
|
||||
def _away_reply_cooldown_seconds(settings: dict) -> int | None:
|
||||
raw = str(settings.get("email_auto_reply_cooldown") or "period").strip().lower()
|
||||
if raw == "1d":
|
||||
return 24 * 60 * 60
|
||||
if raw == "3d":
|
||||
return 3 * 24 * 60 * 60
|
||||
if raw == "7d":
|
||||
return 7 * 24 * 60 * 60
|
||||
return None
|
||||
|
||||
|
||||
def _ensure_away_reply_table():
|
||||
import sqlite3 as _sql3
|
||||
conn = _sql3.connect(SCHEDULED_DB)
|
||||
try:
|
||||
conn.execute("""
|
||||
CREATE TABLE IF NOT EXISTS email_away_replies (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
owner TEXT DEFAULT '',
|
||||
account_id TEXT DEFAULT '',
|
||||
message_id TEXT DEFAULT '',
|
||||
sender_addr TEXT DEFAULT '',
|
||||
subject TEXT DEFAULT '',
|
||||
period_key TEXT DEFAULT '',
|
||||
sent_at TEXT DEFAULT ''
|
||||
)
|
||||
""")
|
||||
conn.execute("CREATE INDEX IF NOT EXISTS idx_email_away_msg ON email_away_replies(owner, account_id, message_id)")
|
||||
conn.execute("CREATE INDEX IF NOT EXISTS idx_email_away_sender ON email_away_replies(owner, account_id, sender_addr, sent_at)")
|
||||
conn.commit()
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def _sender_is_automated(msg, sender_addr: str) -> bool:
|
||||
auto_submitted = (msg.get("Auto-Submitted") or "").strip().lower()
|
||||
if auto_submitted and auto_submitted != "no":
|
||||
return True
|
||||
precedence = (msg.get("Precedence") or "").strip().lower()
|
||||
if precedence in {"bulk", "junk", "list"}:
|
||||
return True
|
||||
if msg.get("List-Id") or msg.get("List-Unsubscribe"):
|
||||
return True
|
||||
local = (sender_addr or "").split("@", 1)[0].lower()
|
||||
return local in {
|
||||
"no-reply", "noreply", "do-not-reply", "donotreply",
|
||||
"notification", "notifications", "automated", "mailer-daemon",
|
||||
"postmaster",
|
||||
}
|
||||
|
||||
|
||||
def _away_reply_already_sent(settings: dict, account_owner: str, account_id: str | None,
|
||||
message_id: str, sender_addr: str) -> bool:
|
||||
import sqlite3 as _sql3
|
||||
_ensure_away_reply_table()
|
||||
owner = account_owner or ""
|
||||
aid = account_id or ""
|
||||
sender = (sender_addr or "").strip().lower()
|
||||
conn = _sql3.connect(SCHEDULED_DB)
|
||||
try:
|
||||
row = conn.execute(
|
||||
"SELECT 1 FROM email_away_replies WHERE owner=? AND account_id=? AND message_id=? LIMIT 1",
|
||||
(owner, aid, message_id),
|
||||
).fetchone()
|
||||
if row:
|
||||
return True
|
||||
|
||||
cooldown = _away_reply_cooldown_seconds(settings)
|
||||
if cooldown is None:
|
||||
period_key = _away_reply_period_key(settings)
|
||||
row = conn.execute(
|
||||
"SELECT 1 FROM email_away_replies WHERE owner=? AND account_id=? AND sender_addr=? AND period_key=? LIMIT 1",
|
||||
(owner, aid, sender, period_key),
|
||||
).fetchone()
|
||||
return bool(row)
|
||||
|
||||
since = datetime.utcnow().timestamp() - cooldown
|
||||
rows = conn.execute(
|
||||
"SELECT sent_at FROM email_away_replies WHERE owner=? AND account_id=? AND sender_addr=? ORDER BY sent_at DESC LIMIT 5",
|
||||
(owner, aid, sender),
|
||||
).fetchall()
|
||||
for (sent_at,) in rows:
|
||||
try:
|
||||
if datetime.fromisoformat(sent_at).timestamp() >= since:
|
||||
return True
|
||||
except Exception:
|
||||
continue
|
||||
return False
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def _record_away_reply(settings: dict, account_owner: str, account_id: str | None,
|
||||
message_id: str, sender_addr: str, subject: str):
|
||||
import sqlite3 as _sql3
|
||||
_ensure_away_reply_table()
|
||||
conn = _sql3.connect(SCHEDULED_DB)
|
||||
try:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO email_away_replies
|
||||
(owner, account_id, message_id, sender_addr, subject, period_key, sent_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
account_owner or "",
|
||||
account_id or "",
|
||||
message_id,
|
||||
(sender_addr or "").strip().lower(),
|
||||
subject or "",
|
||||
_away_reply_period_key(settings),
|
||||
datetime.utcnow().isoformat(),
|
||||
),
|
||||
)
|
||||
conn.commit()
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def _send_away_reply(settings: dict, account_owner: str, account_id: str | None,
|
||||
msg, message_id: str, sender: str, subject: str):
|
||||
sender_name, sender_addr = email.utils.parseaddr(sender or "")
|
||||
sender_addr = (sender_addr or "").strip()
|
||||
if not sender_addr:
|
||||
return False, "missing sender"
|
||||
|
||||
cfg = _get_email_config(account_id, owner=account_owner)
|
||||
from_addr = (cfg.get("from_address") or cfg.get("smtp_user") or "").strip()
|
||||
if not from_addr:
|
||||
return False, "missing from address"
|
||||
if sender_addr.lower() == from_addr.lower():
|
||||
return False, "self mail"
|
||||
if settings.get("email_auto_reply_exclude_automated", True) and _sender_is_automated(msg, sender_addr):
|
||||
return False, "automated sender"
|
||||
if _away_reply_already_sent(settings, account_owner, account_id, message_id, sender_addr):
|
||||
return False, "already sent"
|
||||
|
||||
body = (settings.get("email_auto_reply_message") or "").strip()
|
||||
if not body:
|
||||
body = "Thanks for your email. I'm away and may be slower to reply."
|
||||
|
||||
subject_template = (settings.get("email_auto_reply_subject") or "(Away) {subject}").strip()
|
||||
if subject_template:
|
||||
original_subject = subject or ""
|
||||
reply_subject = (
|
||||
subject_template
|
||||
.replace("{subject}", original_subject)
|
||||
.replace("{original_subject}", original_subject)
|
||||
).strip() or "Re:"
|
||||
else:
|
||||
reply_subject = subject or ""
|
||||
if not reply_subject.lower().lstrip().startswith("re:"):
|
||||
reply_subject = f"Re: {reply_subject}" if reply_subject else "Re:"
|
||||
|
||||
outer = MIMEMultipart("alternative")
|
||||
display = cfg.get("display_name") or ""
|
||||
outer["From"] = email.utils.formataddr((display, from_addr)) if display else from_addr
|
||||
outer["To"] = email.utils.formataddr((sender_name, sender_addr)) if sender_name else sender_addr
|
||||
outer["Subject"] = reply_subject
|
||||
outer["Date"] = email.utils.formatdate(localtime=False)
|
||||
outer["Message-ID"] = email.utils.make_msgid()
|
||||
outer["Auto-Submitted"] = "auto-replied"
|
||||
outer["X-Auto-Response-Suppress"] = "All"
|
||||
if message_id:
|
||||
outer["In-Reply-To"] = message_id
|
||||
refs = (msg.get("References") or "").strip()
|
||||
outer["References"] = f"{refs} {message_id}".strip()
|
||||
outer.attach(MIMEText(body, "plain", "utf-8"))
|
||||
|
||||
_send_smtp_message(cfg, from_addr, [sender_addr], outer.as_string())
|
||||
_record_away_reply(settings, account_owner, account_id, message_id, sender_addr, subject)
|
||||
return True, sender_addr
|
||||
|
||||
|
||||
# ── Routes ──
|
||||
|
||||
async def _emit_progress(progress_cb, message: str):
|
||||
@@ -125,9 +391,10 @@ async def _run_auto_summarize_once(do_summary: bool = True, do_reply: bool = Tru
|
||||
settings = _load_settings()
|
||||
prev = {k: settings.get(k, False) for k in
|
||||
("email_auto_summarize", "email_auto_reply", "email_auto_tag",
|
||||
"email_auto_spam", "email_auto_calendar")}
|
||||
"email_auto_spam", "email_auto_calendar", "_email_auto_reply_draft_only")}
|
||||
settings["email_auto_summarize"] = bool(do_summary)
|
||||
settings["email_auto_reply"] = bool(do_reply)
|
||||
settings["_email_auto_reply_draft_only"] = bool(do_reply)
|
||||
settings["email_auto_tag"] = bool(do_tag)
|
||||
settings["email_auto_spam"] = bool(do_spam)
|
||||
settings["email_auto_calendar"] = bool(do_calendar)
|
||||
@@ -142,7 +409,10 @@ async def _run_auto_summarize_once(do_summary: bool = True, do_reply: bool = Tru
|
||||
finally:
|
||||
s2 = _load_settings()
|
||||
for k, v in prev.items():
|
||||
s2[k] = v
|
||||
if v is None and k.startswith("_"):
|
||||
s2.pop(k, None)
|
||||
else:
|
||||
s2[k] = v
|
||||
_save_settings(s2)
|
||||
|
||||
|
||||
@@ -176,7 +446,7 @@ def _latest_inbox_fallback_uids(conn, reconnect):
|
||||
return [], reconnect()
|
||||
|
||||
|
||||
async def _auto_summarize_pass(days_back: int = 1, account_id: str | None = None, max_process: int | None = None, progress_cb=None) -> str:
|
||||
async def _auto_summarize_pass(days_back: int = 1, account_id: str | None = None, max_process: int | None = None, progress_cb=None, away_only: bool = False) -> str:
|
||||
"""Single pass of the auto-summarize/reply scan.
|
||||
|
||||
When account_id is None, iterates over every enabled account in
|
||||
@@ -208,6 +478,7 @@ async def _auto_summarize_pass(days_back: int = 1, account_id: str | None = None
|
||||
account_id=(ids[0] if ids else None),
|
||||
max_process=max_process,
|
||||
progress_cb=progress_cb,
|
||||
away_only=away_only,
|
||||
)
|
||||
outs = []
|
||||
for idx, aid in enumerate(ids, start=1):
|
||||
@@ -218,6 +489,7 @@ async def _auto_summarize_pass(days_back: int = 1, account_id: str | None = None
|
||||
account_id=aid,
|
||||
max_process=max_process,
|
||||
progress_cb=progress_cb,
|
||||
away_only=away_only,
|
||||
)
|
||||
outs.append(f"[{names.get(aid, aid[:8])}] {result}")
|
||||
except Exception as e:
|
||||
@@ -229,23 +501,32 @@ async def _auto_summarize_pass(days_back: int = 1, account_id: str | None = None
|
||||
account_id=account_id,
|
||||
max_process=max_process,
|
||||
progress_cb=progress_cb,
|
||||
away_only=away_only,
|
||||
)
|
||||
|
||||
|
||||
async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None = None, max_process: int | None = None, progress_cb=None) -> str:
|
||||
async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None = None, max_process: int | None = None, progress_cb=None, away_only: bool = False) -> str:
|
||||
"""Single pass of the auto-summarize/reply scan for ONE account.
|
||||
Reads current settings flags."""
|
||||
import asyncio
|
||||
import sqlite3 as _sql3
|
||||
from src.llm_core import _uses_max_completion_tokens
|
||||
|
||||
settings = _load_settings()
|
||||
settings = _effective_settings_for_email_account(_load_settings(), account_id)
|
||||
auto_sum = settings.get("email_auto_summarize", False)
|
||||
auto_reply = settings.get("email_auto_reply", False)
|
||||
auto_reply_draft = bool(auto_reply and settings.get("_email_auto_reply_draft_only", False))
|
||||
auto_reply_away = bool(auto_reply and not auto_reply_draft and _away_reply_active(settings, account_id))
|
||||
auto_tag = settings.get("email_auto_tag", False)
|
||||
auto_spam = settings.get("email_auto_spam", False)
|
||||
auto_cal = settings.get("email_auto_calendar", False)
|
||||
if not auto_sum and not auto_reply and not auto_tag and not auto_spam and not auto_cal:
|
||||
if away_only:
|
||||
auto_sum = False
|
||||
auto_reply_draft = False
|
||||
auto_tag = False
|
||||
auto_spam = False
|
||||
auto_cal = False
|
||||
if not auto_sum and not auto_reply_draft and not auto_reply_away and not auto_tag and not auto_spam and not auto_cal:
|
||||
return "Nothing to do"
|
||||
|
||||
# Owner of the account being processed. All calendar + mailbox reads/writes
|
||||
@@ -304,11 +585,11 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
|
||||
_c = _sql3.connect(SCHEDULED_DB)
|
||||
_cache_owner_clause, _cache_owner_params = _email_cache_owner_clause(account_owner)
|
||||
_sum_existing = {r[0] for r in _c.execute(
|
||||
_sum_existing = set() if away_only else {r[0] for r in _c.execute(
|
||||
f"SELECT message_id FROM email_summaries WHERE {_cache_owner_clause}",
|
||||
_cache_owner_params,
|
||||
).fetchall()}
|
||||
_reply_existing = {r[0] for r in _c.execute(
|
||||
_reply_existing = set() if away_only else {r[0] for r in _c.execute(
|
||||
f"SELECT message_id FROM email_ai_replies WHERE {_cache_owner_clause}",
|
||||
_cache_owner_params,
|
||||
).fetchall()}
|
||||
@@ -325,7 +606,7 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
).fetchall()}
|
||||
else:
|
||||
_tag_existing = set()
|
||||
_cal_existing = {r[0] for r in _c.execute(
|
||||
_cal_existing = set() if away_only else {r[0] for r in _c.execute(
|
||||
f"SELECT message_id FROM email_calendar_extractions WHERE {_cache_owner_clause}",
|
||||
_cache_owner_params,
|
||||
).fetchall()} if auto_cal else set()
|
||||
@@ -351,12 +632,21 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
if auto_spam and not spam_folder:
|
||||
logger.warning("Auto-spam enabled but no Junk/Spam folder detected — will classify but not move")
|
||||
|
||||
task_candidates = resolve_task_candidates(owner=account_owner)
|
||||
if not task_candidates:
|
||||
return "No model configured"
|
||||
url, model, headers = task_candidates[0]
|
||||
needs_llm = bool(auto_sum or auto_reply_draft or auto_tag or auto_spam or auto_cal)
|
||||
if needs_llm:
|
||||
task_candidates = resolve_task_candidates(owner=account_owner)
|
||||
if not task_candidates:
|
||||
return "No model configured"
|
||||
url, model, headers = task_candidates[0]
|
||||
else:
|
||||
url, model, headers = None, "", None
|
||||
|
||||
writing_style = settings.get("email_writing_style", "")
|
||||
by_account_styles = settings.get("email_writing_styles_by_account") or {}
|
||||
writing_style = ""
|
||||
if account_id and isinstance(by_account_styles, dict):
|
||||
writing_style = str(by_account_styles.get(str(account_id)) or "")
|
||||
if not writing_style:
|
||||
writing_style = settings.get("email_writing_style", "")
|
||||
processed = 0
|
||||
already_cached = 0
|
||||
too_short = 0
|
||||
@@ -366,12 +656,15 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
_events_created = 0
|
||||
_replies_drafted = 0
|
||||
_reply_failed = 0
|
||||
_away_replies_sent = 0
|
||||
_away_replies_skipped = 0
|
||||
_away_replies_failed = 0
|
||||
_detail_lines = []
|
||||
_current_folder = "INBOX"
|
||||
# Calendar extraction is sequential and each row can involve a model
|
||||
# call plus a calendar write. Keep the scheduled calendar-only pass
|
||||
# below the 5-minute action budget instead of timing out mid-run.
|
||||
_default_max_process = 3 if (auto_cal and not auto_sum and not auto_reply and not auto_tag and not auto_spam) else 5
|
||||
_default_max_process = 3 if (auto_cal and not auto_sum and not auto_reply_draft and not auto_reply_away and not auto_tag and not auto_spam) else 5
|
||||
try:
|
||||
_max_process = max(1, int(max_process)) if max_process is not None else _default_max_process
|
||||
except Exception:
|
||||
@@ -402,10 +695,6 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
seed = f"{_folder}|{uid_str}|{msg.get('From','')}|{msg.get('Date','')}|{msg.get('Subject','')}"
|
||||
message_id = f"<synth-{_hl.sha256(seed.encode()).hexdigest()[:16]}@local>"
|
||||
no_msgid += 1
|
||||
need_sum = auto_sum and message_id not in _sum_existing
|
||||
need_reply = auto_reply and message_id not in _reply_existing
|
||||
need_class = (auto_tag or auto_spam) and message_id not in _tag_existing
|
||||
need_cal = bool(settings.get("email_auto_calendar", False)) and message_id not in _cal_existing
|
||||
# Only check urgency on INBOX (received mail), not Sent
|
||||
# Skip messages that are themselves urgency alerts, or that
|
||||
# we sent to ourselves — otherwise the alert loop re-flags
|
||||
@@ -422,17 +711,45 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
except Exception:
|
||||
_from_addr_only = ""
|
||||
_is_self_mail = bool(_self_self_addr) and _from_addr_only.lower() == _self_self_addr
|
||||
need_sum = auto_sum and message_id not in _sum_existing
|
||||
need_reply = auto_reply_draft and message_id not in _reply_existing
|
||||
need_away_reply = bool(
|
||||
auto_reply_away
|
||||
and _folder.upper() == "INBOX"
|
||||
and not _is_self_mail
|
||||
and (away_only or _message_after_away_enabled(settings, msg))
|
||||
and not _away_reply_already_sent(settings, account_owner, account_id, message_id, _from_addr_only)
|
||||
)
|
||||
need_class = (auto_tag or auto_spam) and message_id not in _tag_existing
|
||||
need_cal = bool(settings.get("email_auto_calendar", False)) and message_id not in _cal_existing
|
||||
need_urgent = (auto_urgent and message_id not in _urgent_existing
|
||||
and not _folder.lower().startswith("sent")
|
||||
and "sent" not in _folder.lower()
|
||||
and not _is_alert_echo
|
||||
and not _is_self_mail)
|
||||
if not need_sum and not need_reply and not need_class and not need_cal and not need_urgent:
|
||||
if not need_sum and not need_reply and not need_away_reply and not need_class and not need_cal and not need_urgent:
|
||||
already_cached += 1
|
||||
await _emit_progress(progress_cb, f"Checked {examined}/{len(uid_list)} · {already_cached} already cached")
|
||||
continue
|
||||
subject = _decode_header(msg.get("Subject", ""))
|
||||
sender = _decode_header(msg.get("From", ""))
|
||||
if need_away_reply:
|
||||
try:
|
||||
sent_away, away_detail = _send_away_reply(
|
||||
settings, account_owner, account_id, msg, message_id, sender, subject
|
||||
)
|
||||
if sent_away:
|
||||
_away_replies_sent += 1
|
||||
_uid_text = uid.decode() if isinstance(uid, bytes) else str(uid)
|
||||
_detail_lines.append(f"away reply · {_folder}#{_uid_text} · {subject or '(no subject)'} — {away_detail}")
|
||||
else:
|
||||
_away_replies_skipped += 1
|
||||
logger.info(f"Away reply skipped for uid={uid}: {away_detail}")
|
||||
except Exception as e:
|
||||
_away_replies_failed += 1
|
||||
_uid_text = uid.decode() if isinstance(uid, bytes) else str(uid)
|
||||
_detail_lines.append(f"away reply failed · {_folder}#{_uid_text} · {subject or '(no subject)'}")
|
||||
logger.warning(f"Away reply {uid} failed: {e}")
|
||||
body = _extract_text(msg)
|
||||
# Pull text out of any PDFs / text attachments and append to
|
||||
# the body so summaries / replies can actually reason about
|
||||
@@ -454,7 +771,7 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
elif need_reply:
|
||||
if not body:
|
||||
body = subject
|
||||
elif (not body or len(body) < 100) and not att_text:
|
||||
elif not need_away_reply and (not body or len(body) < 100) and not att_text:
|
||||
too_short += 1
|
||||
continue
|
||||
# Augmented body sent to the LLM: original body + attachment text.
|
||||
@@ -993,7 +1310,8 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
# Build a clear status message
|
||||
ops = []
|
||||
if auto_sum: ops.append("summary")
|
||||
if auto_reply: ops.append("reply")
|
||||
if auto_reply_draft: ops.append("reply")
|
||||
if auto_reply_away: ops.append("away")
|
||||
if auto_tag: ops.append("tag")
|
||||
if auto_spam: ops.append("spam")
|
||||
ops_label = "/".join(ops) or "none"
|
||||
@@ -1002,10 +1320,14 @@ 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 auto_reply:
|
||||
if auto_reply_draft:
|
||||
parts.append(f"drafted {_replies_drafted} repl" + ("y" if _replies_drafted == 1 else "ies"))
|
||||
if _reply_failed:
|
||||
parts.append(f"{_reply_failed} reply failed")
|
||||
if auto_reply_away:
|
||||
parts.append(f"sent {_away_replies_sent} away repl" + ("y" if _away_replies_sent == 1 else "ies"))
|
||||
if _away_replies_failed:
|
||||
parts.append(f"{_away_replies_failed} away failed")
|
||||
if already_cached:
|
||||
parts.append(f"{already_cached} already cached")
|
||||
if too_short:
|
||||
@@ -1032,12 +1354,13 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
|
||||
|
||||
async def _auto_summarize_poller():
|
||||
"""Background loop kept for backward compatibility — calls _auto_summarize_pass every 60s.
|
||||
"""Background loop kept for backward compatibility — calls _auto_summarize_pass periodically.
|
||||
Newer setups should use scheduled tasks instead (summarize_emails, draft_email_replies)."""
|
||||
import asyncio as _asyncio
|
||||
while True:
|
||||
try:
|
||||
await _asyncio.sleep(1800)
|
||||
settings = _load_settings()
|
||||
await _asyncio.sleep(60 if settings.get("email_auto_reply", False) else 1800)
|
||||
await _auto_summarize_pass()
|
||||
except Exception as e:
|
||||
logger.error(f"Auto-summarize poller crash: {e}")
|
||||
@@ -1068,6 +1391,28 @@ def _scheduled_poll_once() -> dict:
|
||||
for r in rows:
|
||||
sid = r[0]
|
||||
try:
|
||||
# Atomically claim this row before doing any work. Two
|
||||
# pollers can race here (the in-process asyncio task and an
|
||||
# externally cron-driven `odysseus-mail poll-scheduled`, or
|
||||
# an admin running the CLI manually alongside the in-process
|
||||
# one despite the ODYSSEUS_INPROCESS_POLLERS=0 guidance) -
|
||||
# both can SELECT the same 'pending' row before either has
|
||||
# updated its status. The UPDATE...WHERE status='pending' is
|
||||
# the atomicity boundary: only the poller whose UPDATE
|
||||
# actually changes a row (rowcount == 1) proceeds to send;
|
||||
# a loser sees rowcount == 0 and skips it instead of sending
|
||||
# a duplicate.
|
||||
claim_conn = sqlite3.connect(SCHEDULED_DB)
|
||||
claim_cur = claim_conn.execute(
|
||||
"UPDATE scheduled_emails SET status='sending' WHERE id=? AND status='pending'",
|
||||
(sid,),
|
||||
)
|
||||
claim_conn.commit()
|
||||
claimed = claim_cur.rowcount == 1
|
||||
claim_conn.close()
|
||||
if not claimed:
|
||||
continue
|
||||
|
||||
attachments = json.loads(r[8] or "[]")
|
||||
row_account_id = r[9] if len(r) > 9 else None
|
||||
odysseus_kind = r[10] if len(r) > 10 else "scheduled"
|
||||
|
||||
+870
-61
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,9 @@
|
||||
"""Gallery routes — browsable library for photos and AI-generated images."""
|
||||
|
||||
import os
|
||||
import base64
|
||||
import hashlib
|
||||
import io
|
||||
import logging
|
||||
import re
|
||||
import uuid
|
||||
@@ -27,6 +29,165 @@ from routes.gallery.gallery_helpers import (
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SAM_STATE: Dict[str, Any] = {}
|
||||
_GROUNDING_STATE: Dict[str, Any] = {}
|
||||
|
||||
|
||||
def _b64_to_pil_image(image_b64: str, *, mode: str = "RGBA"):
|
||||
if not image_b64:
|
||||
raise HTTPException(400, "Missing image")
|
||||
if "," in image_b64 and image_b64.split(",", 1)[0].startswith("data:"):
|
||||
image_b64 = image_b64.split(",", 1)[1]
|
||||
try:
|
||||
from PIL import Image
|
||||
|
||||
raw = base64.b64decode(image_b64)
|
||||
return Image.open(io.BytesIO(raw)).convert(mode)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise HTTPException(400, "Invalid image") from exc
|
||||
|
||||
|
||||
def _pil_image_to_b64(img, *, fmt: str = "PNG") -> str:
|
||||
buf = io.BytesIO()
|
||||
img.save(buf, format=fmt)
|
||||
return base64.b64encode(buf.getvalue()).decode("ascii")
|
||||
|
||||
|
||||
def _load_sam_backend():
|
||||
model_id = os.getenv("ODYSSEUS_SAM_MODEL", "facebook/sam-vit-base")
|
||||
cached = _SAM_STATE.get(model_id)
|
||||
if cached:
|
||||
return cached
|
||||
try:
|
||||
import torch
|
||||
from transformers import SamModel, SamProcessor
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
501,
|
||||
"SAM mask tools are not installed. Install Cookbook Dependencies -> SAM mask tools.",
|
||||
) from exc
|
||||
|
||||
device = "cpu"
|
||||
try:
|
||||
if torch.cuda.is_available():
|
||||
device = "cuda"
|
||||
elif getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
|
||||
device = "mps"
|
||||
except Exception:
|
||||
device = "cpu"
|
||||
|
||||
try:
|
||||
processor = SamProcessor.from_pretrained(model_id)
|
||||
model = SamModel.from_pretrained(model_id)
|
||||
model.to(device)
|
||||
model.eval()
|
||||
except Exception as exc:
|
||||
raise HTTPException(500, f"Failed to load SAM model {model_id}: {exc}") from exc
|
||||
|
||||
cached = {"torch": torch, "processor": processor, "model": model, "device": device, "model_id": model_id}
|
||||
_SAM_STATE[model_id] = cached
|
||||
return cached
|
||||
|
||||
|
||||
def _load_grounding_backend():
|
||||
model_id = os.getenv("ODYSSEUS_GROUNDING_MODEL", "google/owlvit-base-patch32")
|
||||
cached = _GROUNDING_STATE.get(model_id)
|
||||
if cached:
|
||||
return cached
|
||||
try:
|
||||
import torch
|
||||
from transformers import OwlViTForObjectDetection, OwlViTProcessor
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
501,
|
||||
"Object mask tools are not installed. Install Cookbook Dependencies -> SAM mask tools.",
|
||||
) from exc
|
||||
|
||||
device = "cpu"
|
||||
try:
|
||||
if torch.cuda.is_available():
|
||||
device = "cuda"
|
||||
elif getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
|
||||
device = "mps"
|
||||
except Exception:
|
||||
device = "cpu"
|
||||
|
||||
try:
|
||||
processor = OwlViTProcessor.from_pretrained(model_id)
|
||||
model = OwlViTForObjectDetection.from_pretrained(model_id)
|
||||
model.to(device)
|
||||
model.eval()
|
||||
except Exception as exc:
|
||||
raise HTTPException(500, f"Failed to load object mask model {model_id}: {exc}") from exc
|
||||
|
||||
cached = {"torch": torch, "processor": processor, "model": model, "device": device, "model_id": model_id}
|
||||
_GROUNDING_STATE[model_id] = cached
|
||||
return cached
|
||||
|
||||
|
||||
def _ground_text_to_box(image, text: str, *, threshold: float = 0.05):
|
||||
query = (text or "").strip()
|
||||
if not query:
|
||||
raise HTTPException(400, "Missing object text")
|
||||
backend = _load_grounding_backend()
|
||||
torch = backend["torch"]
|
||||
processor = backend["processor"]
|
||||
model = backend["model"]
|
||||
device = backend["device"]
|
||||
|
||||
labels = [query]
|
||||
if not query.lower().startswith(("a ", "an ", "the ")):
|
||||
labels.append(f"a photo of {query}")
|
||||
try:
|
||||
inputs = processor(text=[labels], images=image, return_tensors="pt")
|
||||
model_inputs = {
|
||||
k: (v.to(device) if hasattr(v, "to") else v)
|
||||
for k, v in inputs.items()
|
||||
}
|
||||
with torch.no_grad():
|
||||
outputs = model(**model_inputs)
|
||||
target_sizes = torch.tensor([[image.height, image.width]])
|
||||
if hasattr(processor, "post_process_object_detection"):
|
||||
results = processor.post_process_object_detection(
|
||||
outputs=outputs,
|
||||
target_sizes=target_sizes,
|
||||
threshold=float(threshold),
|
||||
)
|
||||
elif hasattr(processor, "post_process_grounded_object_detection"):
|
||||
results = processor.post_process_grounded_object_detection(
|
||||
outputs=outputs,
|
||||
target_sizes=target_sizes,
|
||||
threshold=float(threshold),
|
||||
text_labels=[labels],
|
||||
)
|
||||
else:
|
||||
raise HTTPException(500, "Installed Transformers does not expose OWL-ViT object detection post-processing")
|
||||
boxes = results[0].get("boxes")
|
||||
scores = results[0].get("scores")
|
||||
labels_idx = results[0].get("labels")
|
||||
text_labels = results[0].get("text_labels") or results[0].get("labels_text")
|
||||
if boxes is None or scores is None or len(boxes) == 0:
|
||||
raise HTTPException(404, f"No visible object matched '{query}'")
|
||||
idx = int(torch.argmax(scores).item())
|
||||
box = [float(v) for v in boxes[idx].detach().cpu().tolist()]
|
||||
label_idx = int(labels_idx[idx].detach().cpu().item()) if labels_idx is not None and len(labels_idx) else 0
|
||||
label = labels[min(label_idx, len(labels) - 1)]
|
||||
if text_labels and len(text_labels) > idx:
|
||||
label = str(text_labels[idx])
|
||||
return {
|
||||
"box": box,
|
||||
"score": float(scores[idx].detach().cpu().item()),
|
||||
"label": label,
|
||||
"model": backend["model_id"],
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.exception("ground_text_to_box failed")
|
||||
raise HTTPException(500, f"Object mask failed: {exc}") from exc
|
||||
|
||||
|
||||
def _current_user_is_admin(request: Request, user: str | None) -> bool:
|
||||
if not user:
|
||||
@@ -1240,20 +1401,89 @@ def setup_gallery_routes() -> APIRouter:
|
||||
except httpx.TimeoutException:
|
||||
raise HTTPException(504, "OpenAI inpaint timed out (120s)")
|
||||
|
||||
# Self-hosted diffusion server path
|
||||
# Self-hosted diffusion server path. Newer Odysseus image
|
||||
# wrappers expose the OpenAI-compatible /v1/images/edits
|
||||
# multipart route even when they are local/self-hosted. Older
|
||||
# diffusion_server.py exposes /v1/images/inpaint as JSON. Try the
|
||||
# OpenAI-compatible local route first, then fall back.
|
||||
try:
|
||||
# Forward chosen_model so the diffusion server can route if it ever
|
||||
# supports multiple models per process. Harmless if ignored.
|
||||
if chosen_model:
|
||||
body["model"] = chosen_model
|
||||
async with httpx.AsyncClient(timeout=120) as client:
|
||||
async with httpx.AsyncClient(timeout=240) as client:
|
||||
try:
|
||||
import base64, io
|
||||
from PIL import Image
|
||||
|
||||
img_bytes = base64.b64decode(body["image"])
|
||||
mask_bytes = base64.b64decode(body["mask"])
|
||||
# Normalize both inputs to PNG bytes. Local MLX and
|
||||
# Diffusers wrappers expect white mask pixels to mean
|
||||
# "edit this region", which matches the editor's mask.
|
||||
source_png = Image.open(io.BytesIO(img_bytes)).convert("RGBA")
|
||||
mask_png = Image.open(io.BytesIO(mask_bytes)).convert("L")
|
||||
src_buf = io.BytesIO()
|
||||
source_png.save(src_buf, format="PNG")
|
||||
mask_buf = io.BytesIO()
|
||||
mask_png.save(mask_buf, format="PNG")
|
||||
files = {
|
||||
"image": ("source.png", src_buf.getvalue(), "image/png"),
|
||||
"mask": ("mask.png", mask_buf.getvalue(), "image/png"),
|
||||
}
|
||||
data = {
|
||||
"model": chosen_model or body.get("model") or "",
|
||||
"prompt": body.get("prompt", ""),
|
||||
"size": f"{int(body.get('width') or source_png.width)}x{int(body.get('height') or source_png.height)}",
|
||||
"n": "1",
|
||||
}
|
||||
r = await client.post(_join_checked_gallery_endpoint(base, "/images/edits"), data=data, files=files)
|
||||
if r.status_code == 200:
|
||||
result = r.json()
|
||||
if isinstance(result, dict) and result.get("data"):
|
||||
item = result["data"][0]
|
||||
if item.get("b64_json"):
|
||||
return {"image": item["b64_json"]}
|
||||
if item.get("url"):
|
||||
raw_b64 = await _fetch_result_image_b64(item["url"])
|
||||
if raw_b64:
|
||||
return {"image": raw_b64}
|
||||
if isinstance(result, dict) and result.get("image"):
|
||||
return {"image": result["image"]}
|
||||
raise HTTPException(502, "Image edit endpoint returned no image")
|
||||
if r.status_code not in (404, 405):
|
||||
logger.warning("inpaint_proxy self-hosted edits: status %s", r.status_code)
|
||||
detail = "Image edit request failed"
|
||||
try:
|
||||
err = r.json()
|
||||
detail = err.get("detail") or err.get("error") or detail
|
||||
except Exception:
|
||||
pass
|
||||
# A plain SD/SDXL checkpoint often exposes
|
||||
# generation only at /images/edits.
|
||||
# That does not mean the endpoint cannot inpaint:
|
||||
# Odysseus diffusion_server.py has a dedicated
|
||||
# /images/inpaint route that can derive/fallback to
|
||||
# inpaint, img2img crop+composite, or txt2img
|
||||
# crop+composite. Fall through to that route instead
|
||||
# of surfacing "does not support image edits".
|
||||
if r.status_code == 400 and "does not support image edits" in str(detail).lower():
|
||||
logger.info("inpaint_proxy self-hosted edits unsupported; falling back to /images/inpaint")
|
||||
else:
|
||||
raise HTTPException(r.status_code, detail)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception:
|
||||
logger.exception("inpaint_proxy: failed to prepare self-hosted edit request")
|
||||
raise HTTPException(400, "Failed to prepare inpaint request")
|
||||
|
||||
r = await client.post(_join_checked_gallery_endpoint(base, "/images/inpaint"), json=body)
|
||||
if r.status_code != 200:
|
||||
logger.error("inpaint_proxy diffusion: status %s", r.status_code)
|
||||
raise HTTPException(r.status_code, "Inpaint request failed")
|
||||
return r.json()
|
||||
except httpx.TimeoutException:
|
||||
raise HTTPException(504, "Inpaint request timed out (120s)")
|
||||
raise HTTPException(504, "Inpaint request timed out (240s)")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception:
|
||||
@@ -1588,6 +1818,135 @@ def setup_gallery_routes() -> APIRouter:
|
||||
return {"error": "AI upscale failed"}
|
||||
|
||||
# ---- POST /api/image/remove-bg ----
|
||||
@router.post("/api/image/mask")
|
||||
async def smart_mask(request: Request):
|
||||
"""Create a neutral segmentation mask from user-provided points or a box.
|
||||
|
||||
This endpoint intentionally does not inspect edit prompts. It only
|
||||
turns explicit visual selection hints into a binary mask that the
|
||||
editor can reuse for wand/layer-mask/inpaint workflows.
|
||||
"""
|
||||
require_privilege(request, "can_generate_images")
|
||||
body = await request.json()
|
||||
image = _b64_to_pil_image(body.get("image") or "", mode="RGB")
|
||||
points = body.get("points") or []
|
||||
box = body.get("box")
|
||||
text = (body.get("text") or body.get("query") or "").strip()
|
||||
grounded = None
|
||||
|
||||
if not points and not box and text:
|
||||
grounded = _ground_text_to_box(image, text)
|
||||
box = grounded["box"]
|
||||
|
||||
if not points and not box:
|
||||
raise HTTPException(400, "Provide at least one point, box, or object text")
|
||||
|
||||
backend = _load_sam_backend()
|
||||
torch = backend["torch"]
|
||||
processor = backend["processor"]
|
||||
model = backend["model"]
|
||||
device = backend["device"]
|
||||
|
||||
kwargs: Dict[str, Any] = {"return_tensors": "pt"}
|
||||
input_points = []
|
||||
if points:
|
||||
input_labels = []
|
||||
for p in points:
|
||||
try:
|
||||
input_points.append([float(p["x"]), float(p["y"])])
|
||||
input_labels.append(int(p.get("label", 1)))
|
||||
except Exception as exc:
|
||||
raise HTTPException(400, "Invalid point format") from exc
|
||||
kwargs["input_points"] = [input_points]
|
||||
kwargs["input_labels"] = [input_labels]
|
||||
if box:
|
||||
if not isinstance(box, list) or len(box) != 4:
|
||||
raise HTTPException(400, "Box must be [x1, y1, x2, y2]")
|
||||
try:
|
||||
kwargs["input_boxes"] = [[[float(v) for v in box]]]
|
||||
except Exception as exc:
|
||||
raise HTTPException(400, "Invalid box format") from exc
|
||||
|
||||
try:
|
||||
inputs = processor(image, **kwargs)
|
||||
model_inputs = {
|
||||
k: (v.to(device) if hasattr(v, "to") else v)
|
||||
for k, v in inputs.items()
|
||||
}
|
||||
with torch.no_grad():
|
||||
outputs = model(**model_inputs)
|
||||
masks = processor.image_processor.post_process_masks(
|
||||
outputs.pred_masks.detach().cpu(),
|
||||
inputs["original_sizes"].detach().cpu(),
|
||||
inputs["reshaped_input_sizes"].detach().cpu(),
|
||||
)
|
||||
mask_tensor = masks[0]
|
||||
while getattr(mask_tensor, "ndim", 0) > 3:
|
||||
mask_tensor = mask_tensor[0]
|
||||
if getattr(mask_tensor, "ndim", 0) == 3:
|
||||
scores = outputs.iou_scores.detach().cpu()[0]
|
||||
while getattr(scores, "ndim", 0) > 1:
|
||||
scores = scores[0]
|
||||
# SAM commonly returns multiple candidates for a click. The
|
||||
# highest-IoU candidate can be the entire image, which is
|
||||
# useless as an editor selection. Prefer a candidate that
|
||||
# contains the clicked point while keeping area reasonable.
|
||||
point_xy = None
|
||||
if input_points:
|
||||
try:
|
||||
point_xy = (
|
||||
int(round(float(input_points[0][0]))),
|
||||
int(round(float(input_points[0][1]))),
|
||||
)
|
||||
except Exception:
|
||||
point_xy = None
|
||||
best_idx = 0
|
||||
best_rank = None
|
||||
total_px = max(1, int(mask_tensor.shape[-1]) * int(mask_tensor.shape[-2]))
|
||||
for i in range(int(mask_tensor.shape[0])):
|
||||
candidate = mask_tensor[i]
|
||||
area_ratio = float(candidate.sum().item()) / float(total_px)
|
||||
if area_ratio >= 0.985:
|
||||
continue
|
||||
contains_click = True
|
||||
if point_xy:
|
||||
px = max(0, min(int(candidate.shape[-1]) - 1, point_xy[0]))
|
||||
py = max(0, min(int(candidate.shape[-2]) - 1, point_xy[1]))
|
||||
contains_click = bool(candidate[py, px].item())
|
||||
if not contains_click:
|
||||
continue
|
||||
score = float(scores[min(i, len(scores) - 1)].item()) if len(scores) else 0.0
|
||||
# Strongly penalize broad masks; a click-selection should
|
||||
# usually be local unless the user gives a box.
|
||||
rank = score - (area_ratio * 0.35)
|
||||
if best_rank is None or rank > best_rank:
|
||||
best_rank = rank
|
||||
best_idx = i
|
||||
if best_rank is None and len(scores):
|
||||
best_idx = int(torch.argmax(scores).item())
|
||||
mask_tensor = mask_tensor[min(best_idx, mask_tensor.shape[0] - 1)]
|
||||
mask_array = (mask_tensor.numpy() > 0).astype("uint8") * 255
|
||||
from PIL import Image
|
||||
|
||||
mask_img = Image.fromarray(mask_array, mode="L")
|
||||
if mask_img.size != image.size:
|
||||
mask_img = mask_img.resize(image.size, Image.NEAREST)
|
||||
bbox = mask_img.getbbox()
|
||||
result = {
|
||||
"mask": _pil_image_to_b64(mask_img),
|
||||
"bbox": list(bbox) if bbox else None,
|
||||
"model": backend["model_id"],
|
||||
"device": device,
|
||||
}
|
||||
if grounded:
|
||||
result["grounding"] = grounded
|
||||
return result
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.exception("smart_mask failed")
|
||||
raise HTTPException(500, f"SAM mask failed: {exc}") from exc
|
||||
|
||||
@router.post("/api/image/remove-bg")
|
||||
async def remove_background(request: Request):
|
||||
"""Remove background from an image. If the client passes a `hint_mask`
|
||||
|
||||
@@ -10,7 +10,9 @@ from fastapi import APIRouter, Request, HTTPException
|
||||
|
||||
from core.models import ChatMessage
|
||||
from core.database import SessionLocal, ChatMessage as DbChatMessage, Session as DbSession
|
||||
from src.auth_helpers import effective_user
|
||||
from src.topic_analyzer import analyze_topics
|
||||
from src.upload_handler import reserve_message_upload_references
|
||||
from routes.session_routes import (
|
||||
_message_role,
|
||||
_message_text,
|
||||
@@ -98,9 +100,29 @@ def _merge_continue_rows_to_delete(db_messages, db1, db2):
|
||||
return to_delete
|
||||
|
||||
|
||||
def setup_history_routes(session_manager) -> APIRouter:
|
||||
def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
router = APIRouter(tags=["history"])
|
||||
|
||||
def _reserve_message_uploads(
|
||||
request: Request,
|
||||
content: Any,
|
||||
metadata: Any = None,
|
||||
) -> None:
|
||||
try:
|
||||
missing_id = reserve_message_upload_references(
|
||||
upload_handler,
|
||||
effective_user(request),
|
||||
content,
|
||||
metadata,
|
||||
)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise HTTPException(400, "Invalid message attachment metadata") from exc
|
||||
if missing_id:
|
||||
raise HTTPException(
|
||||
409,
|
||||
f"Referenced upload is no longer available: {missing_id}",
|
||||
)
|
||||
|
||||
def _db_history_entry(m: DbChatMessage) -> Dict[str, Any]:
|
||||
entry = {"role": m.role, "content": _history_display_content(m.content)}
|
||||
meta = {}
|
||||
@@ -115,6 +137,44 @@ def setup_history_routes(session_manager) -> APIRouter:
|
||||
entry["metadata"] = meta
|
||||
return entry
|
||||
|
||||
def _db_message_metadata(m: DbChatMessage) -> Dict[str, Any]:
|
||||
meta = {}
|
||||
if m.meta_data:
|
||||
try:
|
||||
meta = json.loads(m.meta_data) or {}
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
meta = {}
|
||||
if m.timestamp and "timestamp" not in meta:
|
||||
meta["timestamp"] = m.timestamp.isoformat() + "Z"
|
||||
return meta
|
||||
|
||||
def _hydrate_session_history_from_db(session_id: str, rows: list[DbChatMessage]) -> None:
|
||||
"""Rebuild in-memory context from raw DB rows after a history load.
|
||||
|
||||
The browser history endpoint can return paged/display-trimmed messages,
|
||||
but the next model call reads ``session.history``. After a restart or a
|
||||
stale in-memory session, selecting an old chat through the paged endpoint
|
||||
used to show the transcript while the model only saw fresh context.
|
||||
"""
|
||||
if not rows:
|
||||
return
|
||||
try:
|
||||
session = session_manager.get_session(session_id)
|
||||
except KeyError:
|
||||
return
|
||||
session.history = [
|
||||
ChatMessage(role=m.role, content=m.content, metadata=_db_message_metadata(m) or None)
|
||||
for m in rows
|
||||
]
|
||||
session.message_count = len(session.history)
|
||||
|
||||
def _session_needs_db_history_hydration(session_id: str, total: int) -> bool:
|
||||
try:
|
||||
session = session_manager.get_session(session_id)
|
||||
except KeyError:
|
||||
return False
|
||||
return len(session.history or []) < int(total or 0)
|
||||
|
||||
@router.get("/api/history/{session_id}")
|
||||
async def get_session_history(
|
||||
request: Request,
|
||||
@@ -146,6 +206,14 @@ def setup_history_routes(session_manager) -> APIRouter:
|
||||
.limit(page_limit)
|
||||
.all()
|
||||
)
|
||||
if _session_needs_db_history_hydration(session_id, total):
|
||||
full_rows = (
|
||||
db.query(DbChatMessage)
|
||||
.filter(DbChatMessage.session_id == session_id)
|
||||
.order_by(DbChatMessage.timestamp)
|
||||
.all()
|
||||
)
|
||||
_hydrate_session_history_from_db(session_id, full_rows)
|
||||
history_dict = [
|
||||
entry for entry in (_db_history_entry(m) for m in rows)
|
||||
if not (entry.get("metadata") or {}).get("hidden")
|
||||
@@ -206,10 +274,7 @@ def setup_history_routes(session_manager) -> APIRouter:
|
||||
if db_history:
|
||||
# Rebuild in-memory history from the full set so hidden
|
||||
# messages (e.g. compaction summaries) are kept for AI context.
|
||||
session.history = [
|
||||
ChatMessage(role=m["role"], content=m["content"], metadata=m.get("metadata"))
|
||||
for m in db_history
|
||||
]
|
||||
_hydrate_session_history_from_db(session_id, db_messages)
|
||||
# Response excludes hidden messages, matching the in-memory path.
|
||||
history_dict = [
|
||||
m for m in db_history
|
||||
@@ -251,7 +316,9 @@ def setup_history_routes(session_manager) -> APIRouter:
|
||||
content = body.get("content", "")
|
||||
if not content:
|
||||
raise HTTPException(400, "content is required")
|
||||
msg = ChatMessage(role=role, content=content, metadata=body.get("metadata"))
|
||||
metadata = body.get("metadata")
|
||||
_reserve_message_uploads(request, content, metadata)
|
||||
msg = ChatMessage(role=role, content=content, metadata=metadata)
|
||||
session_manager.add_message(session_id, msg)
|
||||
return {"status": "ok"}
|
||||
except KeyError:
|
||||
@@ -331,6 +398,8 @@ def setup_history_routes(session_manager) -> APIRouter:
|
||||
if not msg_id or content is None:
|
||||
raise HTTPException(400, "msg_id and content are required")
|
||||
|
||||
_reserve_message_uploads(request, content)
|
||||
|
||||
session = session_manager.get_session(session_id)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
@@ -630,6 +699,55 @@ def setup_history_routes(session_manager) -> APIRouter:
|
||||
except Exception as e:
|
||||
raise HTTPException(500, f"Topic analysis failed: {e}")
|
||||
|
||||
@router.get("/api/session/{session_id}/context")
|
||||
async def get_session_context_usage(request: Request, session_id: str) -> Dict[str, Any]:
|
||||
"""Return an estimated whole-chat context usage for the session's model.
|
||||
|
||||
Streaming footers report the prompt size for the last request. This
|
||||
endpoint estimates the persisted session context so the header can show
|
||||
when the whole chat is approaching compaction.
|
||||
"""
|
||||
_verify_session_owner(request, session_id)
|
||||
try:
|
||||
session = session_manager.get_session(session_id)
|
||||
except KeyError:
|
||||
raise HTTPException(404, "Session not found")
|
||||
|
||||
try:
|
||||
from src.model_context import estimate_tokens, get_context_length
|
||||
|
||||
messages = session.get_context_messages()
|
||||
used = int(estimate_tokens(messages))
|
||||
ctx_len = int(get_context_length(session.endpoint_url, session.model) or 0)
|
||||
pct = round((used / ctx_len) * 100, 1) if ctx_len else 0.0
|
||||
pct = max(0.0, min(100.0, pct))
|
||||
visible_messages = sum(
|
||||
1 for m in session.history
|
||||
if not (getattr(m, "metadata", None) or {}).get("hidden")
|
||||
)
|
||||
compacted_messages = sum(
|
||||
1 for m in session.history
|
||||
if (getattr(m, "metadata", None) or {}).get("compacted")
|
||||
)
|
||||
can_compact = used > 0
|
||||
return {
|
||||
"session_id": session_id,
|
||||
"model": session.model,
|
||||
"endpoint_url": session.endpoint_url,
|
||||
"used_tokens": used,
|
||||
"context_length": ctx_len,
|
||||
"context_percent": pct,
|
||||
"messages": visible_messages,
|
||||
"context_messages": len(messages),
|
||||
"compacted_messages": compacted_messages,
|
||||
"can_compact": can_compact,
|
||||
"should_compact": pct >= 70,
|
||||
"auto_compact_threshold": 85,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Context usage error {session_id}: {e}")
|
||||
raise HTTPException(500, str(e))
|
||||
|
||||
@router.post("/api/session/{session_id}/compact")
|
||||
async def compact_session(request: Request, session_id: str):
|
||||
"""Manually trigger context compaction for a session."""
|
||||
@@ -674,7 +792,7 @@ def setup_history_routes(session_manager) -> APIRouter:
|
||||
compact_model = util_model or session.model
|
||||
compact_headers = util_headers if util_url else session.headers
|
||||
|
||||
from src.context_compactor import SELF_SUMMARY_SYSTEM_PROMPT
|
||||
from src.context_compactor import SELF_SUMMARY_SYSTEM_PROMPT, normalize_compaction_summary
|
||||
compaction_count = sum(1 for m in session.history if isinstance(m, ChatMessage) and "[Conversation summary" in (m.content or ""))
|
||||
sys_prompt = SELF_SUMMARY_SYSTEM_PROMPT.replace("{count}", str(len(older))).replace("{n}", str(compaction_count + 1))
|
||||
summary = await llm_call_async(
|
||||
@@ -686,6 +804,7 @@ def setup_history_routes(session_manager) -> APIRouter:
|
||||
temperature=0.2, max_tokens=1024,
|
||||
headers=compact_headers, timeout=30,
|
||||
)
|
||||
summary = normalize_compaction_summary(summary)
|
||||
|
||||
# Replace session history: summary as system message + recent messages
|
||||
# System message holds the full summary for AI context
|
||||
|
||||
+21
-5
@@ -429,11 +429,27 @@ def setup_hwfit_routes():
|
||||
system["available_ram_gb"] = 0
|
||||
system["total_ram_gb"] = 0
|
||||
system = _apply_manual_hardware(system, manual_mode, manual_gpu_count, manual_vram_gb, manual_ram_gb, manual_backend)
|
||||
# Image models use a single GPU — always use per-GPU VRAM
|
||||
gpu_vrams = [float(g.get("vram_gb") or 0) for g in (system.get("gpus") or []) if isinstance(g, dict)]
|
||||
single_vram = max(gpu_vrams) if gpu_vrams else ((system.get("gpu_vram_gb") or 0) / max(system.get("gpu_count") or 1, 1))
|
||||
system["gpu_vram_gb"] = single_vram
|
||||
system["gpu_count"] = 1 if single_vram > 0 else 0
|
||||
try:
|
||||
requested_gpu_count = int(gpu_count) if gpu_count != "" else None
|
||||
except ValueError:
|
||||
requested_gpu_count = None
|
||||
if requested_gpu_count == 0:
|
||||
# Respect the UI's RAM toggle. Before this route always rewrote the
|
||||
# system to best-single-GPU VRAM, so image rows never changed when
|
||||
# switching RAM/GPU.
|
||||
system["has_gpu"] = False
|
||||
system["gpu_vram_gb"] = 0
|
||||
system["gpu_count"] = 0
|
||||
system["gpu_only"] = False
|
||||
else:
|
||||
# Image diffusion backends generally use one device per pipeline,
|
||||
# so rank GPU mode against the best single GPU rather than total
|
||||
# multi-GPU VRAM.
|
||||
gpu_vrams = [float(g.get("vram_gb") or 0) for g in (system.get("gpus") or []) if isinstance(g, dict)]
|
||||
single_vram = max(gpu_vrams) if gpu_vrams else ((system.get("gpu_vram_gb") or 0) / max(system.get("gpu_count") or 1, 1))
|
||||
system["gpu_vram_gb"] = single_vram
|
||||
system["gpu_count"] = 1 if single_vram > 0 else 0
|
||||
system["gpu_only"] = True if single_vram > 0 else False
|
||||
results = rank_image_models(system, search=search or None, sort=sort)
|
||||
return {"system": system, "models": results}
|
||||
|
||||
|
||||
+215
-17
@@ -472,7 +472,11 @@ def _endpoint_kind(ep: Any) -> str:
|
||||
|
||||
|
||||
def _endpoint_refresh_mode(ep: Any, endpoint_kind: str | None = None) -> str:
|
||||
return _normalize_refresh_mode(getattr(ep, "model_refresh_mode", None), endpoint_kind or _endpoint_kind(ep))
|
||||
return _normalize_endpoint_refresh_mode(
|
||||
getattr(ep, "model_refresh_mode", None),
|
||||
endpoint_kind or _endpoint_kind(ep),
|
||||
getattr(ep, "base_url", ""),
|
||||
)
|
||||
|
||||
|
||||
def _endpoint_refresh_interval(ep: Any, category: str) -> float:
|
||||
@@ -851,6 +855,99 @@ def _ollama_model_names(data: Any) -> List[str]:
|
||||
return out
|
||||
|
||||
|
||||
def _is_google_api_base(base_url: str) -> bool:
|
||||
try:
|
||||
return (urlparse(base_url).hostname or "").lower() == "generativelanguage.googleapis.com"
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _normalize_endpoint_refresh_mode(value: Any, endpoint_kind: str = "auto", base_url: str = "") -> str:
|
||||
if not str(value or "").strip() and _is_google_api_base(base_url):
|
||||
return "manual"
|
||||
return _normalize_refresh_mode(value, endpoint_kind)
|
||||
|
||||
|
||||
def _google_native_root(base_url: str) -> str:
|
||||
"""Return the Gemini native API root for a Google endpoint.
|
||||
|
||||
Chat calls may be configured against Google's OpenAI-compatible
|
||||
`/openai` path, but model catalog reads should use the native Models API
|
||||
so we get Google's current Model resource shape.
|
||||
"""
|
||||
try:
|
||||
parsed = urlparse(base_url)
|
||||
except Exception:
|
||||
return "https://generativelanguage.googleapis.com/v1beta"
|
||||
path = (parsed.path or "").rstrip("/")
|
||||
if path.endswith("/openai"):
|
||||
path = path[: -len("/openai")].rstrip("/")
|
||||
if not path:
|
||||
path = "/v1beta"
|
||||
return urlunparse(parsed._replace(path=path, query="", fragment="")).rstrip("/")
|
||||
|
||||
|
||||
def _google_native_models_url(base_url: str) -> str:
|
||||
return _google_native_root(base_url) + "/models"
|
||||
|
||||
|
||||
def _google_model_id_from_item(item: Any) -> str:
|
||||
if not isinstance(item, dict):
|
||||
return ""
|
||||
value = item.get("baseModelId") or item.get("name") or item.get("model") or ""
|
||||
return str(value or "").strip().removeprefix("models/")
|
||||
|
||||
|
||||
def _google_model_supports_chat(item: Any) -> bool:
|
||||
"""Return whether a native Google Model resource supports chat generation."""
|
||||
if not isinstance(item, dict):
|
||||
return False
|
||||
methods = item.get("supportedGenerationMethods")
|
||||
if not isinstance(methods, list):
|
||||
return False
|
||||
chat_methods = {"generateContent", "generateMessage", "generateText", "generateAnswer"}
|
||||
return any(method in chat_methods for method in methods)
|
||||
|
||||
|
||||
def _probe_google_models(base_url: str, api_key: str = None, timeout: int = 5, page_size: int = 1000) -> List[str]:
|
||||
"""Read Google's native paginated Models API.
|
||||
|
||||
This intentionally returns only provider-reported model IDs. Capability
|
||||
mapping is handled by the model capability reader and must not infer from
|
||||
names here.
|
||||
"""
|
||||
url = _google_native_models_url(base_url)
|
||||
try:
|
||||
page_size = min(max(int(page_size or 1000), 1), 1000)
|
||||
except Exception:
|
||||
page_size = 1000
|
||||
headers = {"Accept": "application/json"}
|
||||
if api_key:
|
||||
headers["x-goog-api-key"] = api_key
|
||||
params: Dict[str, Any] = {"pageSize": page_size}
|
||||
models: List[str] = []
|
||||
seen = set()
|
||||
page_token = ""
|
||||
for _ in range(20):
|
||||
request_params = dict(params)
|
||||
if page_token:
|
||||
request_params["pageToken"] = page_token
|
||||
r = httpx.get(url, headers=headers, params=request_params, timeout=timeout, verify=llm_verify())
|
||||
r.raise_for_status()
|
||||
data = r.json()
|
||||
for item in data.get("models") or []:
|
||||
if not _google_model_supports_chat(item):
|
||||
continue
|
||||
model_id = _google_model_id_from_item(item)
|
||||
if model_id and model_id not in seen:
|
||||
seen.add(model_id)
|
||||
models.append(model_id)
|
||||
page_token = str(data.get("nextPageToken") or "").strip()
|
||||
if not page_token:
|
||||
break
|
||||
return models
|
||||
|
||||
|
||||
def _probe_endpoint(base_url: str, api_key: str = None, timeout: int = 5) -> List[str]:
|
||||
"""Probe a base URL's /models endpoint and return list of model IDs.
|
||||
For Anthropic, queries their /v1/models API, falling back to hardcoded list."""
|
||||
@@ -863,6 +960,17 @@ def _probe_endpoint(base_url: str, api_key: str = None, timeout: int = 5) -> Lis
|
||||
if api_key:
|
||||
return fetch_available_models(api_key, timeout=timeout)
|
||||
return []
|
||||
if _is_google_api_base(base):
|
||||
try:
|
||||
models = _probe_google_models(base, api_key, timeout=timeout)
|
||||
if models:
|
||||
return models
|
||||
except httpx.HTTPStatusError as e:
|
||||
status = e.response.status_code if e.response is not None else "unknown"
|
||||
logger.warning(f"Google native models probe failed: HTTP {status}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Google native models probe failed: {e}")
|
||||
return []
|
||||
if provider == "anthropic":
|
||||
# Try Anthropic's /v1/models endpoint first
|
||||
url = _safe_build_models_url(base)
|
||||
@@ -1202,6 +1310,56 @@ def _visible_models(cached_models, hidden_models, pinned_models=None):
|
||||
return [m for m in merged if m not in hidden]
|
||||
|
||||
|
||||
def _picker_requires_pinning(base_url: str, kind: str) -> bool:
|
||||
return _classify_endpoint(base_url, kind) == "api"
|
||||
|
||||
|
||||
def _has_explicit_pinned_models(ep) -> bool:
|
||||
"""Whether pinned_models was deliberately written for this endpoint.
|
||||
|
||||
API endpoints use pinned_models as an allow-list. An explicit empty JSON
|
||||
list means "show no models"; it must not fall back to the old hidden-list
|
||||
migration behavior.
|
||||
"""
|
||||
raw = getattr(ep, "pinned_models", None)
|
||||
return raw is not None and str(raw).strip() != ""
|
||||
|
||||
|
||||
def _legacy_visible_api_models(ep) -> List[str]:
|
||||
"""Return API models selected under the old hidden-list picker.
|
||||
|
||||
Before API endpoints switched to an explicit allow-list, selected models
|
||||
were represented as cached_models minus hidden_models. Existing OpenRouter
|
||||
rows can therefore have many checked models and an empty pinned_models
|
||||
field. Treat that old state as the initial pinned list so settings and chat
|
||||
agree after upgrade.
|
||||
"""
|
||||
return _visible_models(
|
||||
_cached_model_ids(ep),
|
||||
getattr(ep, "hidden_models", None),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def _picker_models_for_endpoint(ep, base_url: str, kind: str):
|
||||
"""Return model IDs that should appear in the picker for an endpoint.
|
||||
|
||||
API providers expose remote inventory from /v1/models. Treat that cache as
|
||||
inventory, not approval: only manually pinned API models should appear in
|
||||
the picker. Local/self-hosted endpoints keep the older hide-list behavior.
|
||||
"""
|
||||
pinned = _normalize_model_ids(getattr(ep, "pinned_models", None))
|
||||
if _picker_requires_pinning(base_url, kind):
|
||||
if not _has_explicit_pinned_models(ep):
|
||||
pinned = _legacy_visible_api_models(ep) if _hidden_model_ids(ep) else []
|
||||
return pinned, pinned
|
||||
return _visible_models(
|
||||
_cached_model_ids(ep),
|
||||
getattr(ep, "hidden_models", None),
|
||||
pinned,
|
||||
), pinned
|
||||
|
||||
|
||||
def _api_key_fingerprint(api_key: Optional[str]) -> str:
|
||||
"""Stable, non-secret label for distinguishing same-URL credentials."""
|
||||
key = (api_key or "").strip()
|
||||
@@ -1400,24 +1558,18 @@ def setup_model_routes(model_discovery):
|
||||
for ep in endpoints:
|
||||
base = _normalize_base(ep.base_url)
|
||||
provider = _safe_detect_provider(base)
|
||||
# Merge cached + pinned models, then filter out hidden ones
|
||||
ep_model_type = getattr(ep, "model_type", None) or "llm"
|
||||
model_ids = _visible_models(
|
||||
_cached_model_ids(ep),
|
||||
ep.hidden_models,
|
||||
getattr(ep, "pinned_models", None),
|
||||
)
|
||||
# Build correct URL based on provider
|
||||
chat_url = build_chat_url(base)
|
||||
kind = _effective_endpoint_kind(ep, base)
|
||||
category = _classify_endpoint(base, kind)
|
||||
model_ids, pinned = _picker_models_for_endpoint(ep, base, kind)
|
||||
|
||||
if model_ids:
|
||||
curated_key = _match_provider_curated(base, None)
|
||||
curated, extra = _curate_models(model_ids, curated_key)
|
||||
# Pinned models are admin-selected — they always belong in the
|
||||
# primary curated list, not buried in extras.
|
||||
pinned = _normalize_model_ids(getattr(ep, "pinned_models", None))
|
||||
for m in pinned:
|
||||
if m not in curated:
|
||||
curated.append(m)
|
||||
@@ -1779,18 +1931,24 @@ def setup_model_routes(model_discovery):
|
||||
_invalidate_models_cache()
|
||||
rows = db.query(ModelEndpoint).order_by(ModelEndpoint.created_at).all()
|
||||
results = []
|
||||
upgraded_legacy_pins = False
|
||||
for r in rows:
|
||||
all_models = _cached_model_ids(r)
|
||||
hidden = _hidden_model_ids(r)
|
||||
pinned = _normalize_model_ids(getattr(r, "pinned_models", None))
|
||||
visible = _visible_models(all_models, r.hidden_models, pinned)
|
||||
# Keep the list route cache-only. It feeds Settings →
|
||||
# Added Models and must render immediately; explicit
|
||||
# Refresh/Probe endpoints do the network work.
|
||||
status = "online" if (all_models or pinned) else ("empty" if r.is_enabled else "offline")
|
||||
ping = None
|
||||
base = _normalize_base(r.base_url)
|
||||
kind = _effective_endpoint_kind(r, base)
|
||||
visible, pinned = _picker_models_for_endpoint(r, base, kind)
|
||||
if _picker_requires_pinning(base, kind) and pinned and not _has_explicit_pinned_models(r):
|
||||
r.pinned_models = json.dumps(pinned)
|
||||
upgraded_legacy_pins = True
|
||||
model_inventory_count = len(_merge_model_ids(all_models, pinned))
|
||||
picker_requires_pinning = _picker_requires_pinning(base, kind)
|
||||
status = "online" if (all_models or visible or pinned) else ("empty" if r.is_enabled else "offline")
|
||||
results.append({
|
||||
"id": r.id,
|
||||
"name": r.name,
|
||||
@@ -1799,6 +1957,8 @@ def setup_model_routes(model_discovery):
|
||||
"api_key_fingerprint": _api_key_fingerprint(r.api_key),
|
||||
"is_enabled": r.is_enabled,
|
||||
"models": visible,
|
||||
"model_count": model_inventory_count,
|
||||
"picker_requires_pinning": picker_requires_pinning,
|
||||
"pinned_models": pinned,
|
||||
"hidden_count": len(hidden),
|
||||
"online": status != "offline",
|
||||
@@ -1812,6 +1972,9 @@ def setup_model_routes(model_discovery):
|
||||
"model_refresh_interval": getattr(r, "model_refresh_interval", None),
|
||||
"model_refresh_timeout": getattr(r, "model_refresh_timeout", None),
|
||||
})
|
||||
if upgraded_legacy_pins:
|
||||
db.commit()
|
||||
_invalidate_models_cache()
|
||||
return results
|
||||
finally:
|
||||
db.close()
|
||||
@@ -1854,7 +2017,7 @@ def setup_model_routes(model_discovery):
|
||||
name = base_url.replace("http://", "").replace("https://", "").split("/")[0]
|
||||
|
||||
requested_kind = _normalize_endpoint_kind(endpoint_kind)
|
||||
refresh_mode = _normalize_refresh_mode(model_refresh_mode, requested_kind)
|
||||
refresh_mode = _normalize_endpoint_refresh_mode(model_refresh_mode, requested_kind, base_url)
|
||||
refresh_interval = _parse_positive_int(model_refresh_interval, minimum=30, maximum=86400)
|
||||
refresh_timeout = _parse_positive_int(model_refresh_timeout, minimum=1, maximum=60)
|
||||
require_model_list = _truthy(require_models)
|
||||
@@ -1915,6 +2078,10 @@ def setup_model_routes(model_discovery):
|
||||
if refresh_timeout is not None:
|
||||
existing.model_refresh_timeout = refresh_timeout
|
||||
changed = True
|
||||
incoming_model_type = (model_type or "").strip() or "llm"
|
||||
if incoming_model_type and (getattr(existing, "model_type", None) or "llm") != incoming_model_type:
|
||||
existing.model_type = incoming_model_type
|
||||
changed = True
|
||||
if api_key.strip() and not existing.api_key:
|
||||
existing.api_key = api_key.strip()
|
||||
changed = True
|
||||
@@ -2147,9 +2314,10 @@ def setup_model_routes(model_discovery):
|
||||
raise HTTPException(404, "Endpoint not found")
|
||||
hidden = _hidden_model_ids(ep)
|
||||
all_models = _cached_model_ids(ep)
|
||||
base = _normalize_base(ep.base_url)
|
||||
kind = _effective_endpoint_kind(ep, base)
|
||||
picker_requires_pinning = _picker_requires_pinning(base, kind)
|
||||
if refresh:
|
||||
base = _normalize_base(ep.base_url)
|
||||
kind = _effective_endpoint_kind(ep, base)
|
||||
category = _classify_endpoint(base, kind)
|
||||
timeout = _manual_refresh_timeout(ep, category, refresh_timeout)
|
||||
try:
|
||||
@@ -2168,6 +2336,8 @@ def setup_model_routes(model_discovery):
|
||||
response.headers["X-Model-Refresh-Status"] = "failed"
|
||||
response.headers["X-Model-Refresh-Warning"] = "Model refresh failed or returned no models; kept cached models."
|
||||
pinned = _normalize_model_ids(getattr(ep, "pinned_models", None))
|
||||
if picker_requires_pinning and not _has_explicit_pinned_models(ep):
|
||||
pinned = _legacy_visible_api_models(ep)
|
||||
pinned_set = set(pinned)
|
||||
return [
|
||||
{
|
||||
@@ -2175,6 +2345,7 @@ def setup_model_routes(model_discovery):
|
||||
"display": m.split("/")[-1],
|
||||
"is_hidden": m in hidden,
|
||||
"is_pinned": m in pinned_set,
|
||||
"picker_requires_pinning": picker_requires_pinning,
|
||||
}
|
||||
for m in _merge_model_ids(all_models, pinned)
|
||||
]
|
||||
@@ -2203,11 +2374,28 @@ def setup_model_routes(model_discovery):
|
||||
hidden = body.get("hidden")
|
||||
if not isinstance(hidden, list):
|
||||
raise HTTPException(400, "hidden must be a list of model IDs")
|
||||
ep.hidden_models = json.dumps(hidden) if hidden else None
|
||||
base = _normalize_base(ep.base_url)
|
||||
kind = _effective_endpoint_kind(ep, base)
|
||||
if _picker_requires_pinning(base, kind):
|
||||
# Compatibility for older/admin UI paths that still submit
|
||||
# the previous hide-list shape. API pickers are allow-lists:
|
||||
# convert "unchecked models" into an explicit pinned list so
|
||||
# Settings summary, /api/models, and chat agree.
|
||||
selected = _visible_models(_cached_model_ids(ep), hidden, None)
|
||||
ep.pinned_models = json.dumps(selected)
|
||||
ep.hidden_models = None
|
||||
else:
|
||||
ep.hidden_models = json.dumps(hidden) if hidden else None
|
||||
# Accept either "pinned" or "pinned_models" for the manual IDs list.
|
||||
if "pinned_models" in body or "pinned" in body:
|
||||
pinned = _normalize_model_ids(body.get("pinned_models", body.get("pinned")))
|
||||
ep.pinned_models = json.dumps(pinned) if pinned else None
|
||||
base = _normalize_base(ep.base_url)
|
||||
kind = _effective_endpoint_kind(ep, base)
|
||||
if _picker_requires_pinning(base, kind):
|
||||
ep.pinned_models = json.dumps(pinned)
|
||||
ep.hidden_models = None
|
||||
else:
|
||||
ep.pinned_models = json.dumps(pinned) if pinned else None
|
||||
db.commit()
|
||||
_invalidate_models_cache()
|
||||
hidden_count = len(json.loads(ep.hidden_models)) if ep.hidden_models else 0
|
||||
@@ -2360,11 +2548,21 @@ def setup_model_routes(model_discovery):
|
||||
ep.model_type = body["model_type"].strip() or ep.model_type
|
||||
if "pinned_models" in body:
|
||||
_pinned = _normalize_model_ids(body["pinned_models"])
|
||||
ep.pinned_models = json.dumps(_pinned) if _pinned else None
|
||||
_base_for_pins = _normalize_base(ep.base_url)
|
||||
_kind_for_pins = _effective_endpoint_kind(ep, _base_for_pins)
|
||||
if _picker_requires_pinning(_base_for_pins, _kind_for_pins):
|
||||
ep.pinned_models = json.dumps(_pinned)
|
||||
ep.hidden_models = None
|
||||
else:
|
||||
ep.pinned_models = json.dumps(_pinned) if _pinned else None
|
||||
if "endpoint_kind" in body:
|
||||
ep.endpoint_kind = _normalize_endpoint_kind(body.get("endpoint_kind"))
|
||||
if "model_refresh_mode" in body:
|
||||
ep.model_refresh_mode = _normalize_refresh_mode(body.get("model_refresh_mode"), _endpoint_kind(ep))
|
||||
ep.model_refresh_mode = _normalize_endpoint_refresh_mode(
|
||||
body.get("model_refresh_mode"),
|
||||
_endpoint_kind(ep),
|
||||
ep.base_url,
|
||||
)
|
||||
if "model_refresh_interval" in body:
|
||||
interval = _parse_positive_int(body.get("model_refresh_interval"), minimum=30, maximum=86400)
|
||||
ep.model_refresh_interval = interval
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Note route domain package (slice 2f, #4082/#4071).
|
||||
|
||||
Contains note_routes.py, migrated from the flat routes/ directory.
|
||||
Backward-compat shim at routes/note_routes.py re-exports from here.
|
||||
"""
|
||||
@@ -0,0 +1,937 @@
|
||||
# routes/note_routes.py
|
||||
"""Google Keep-style notes / checklists API."""
|
||||
|
||||
import json
|
||||
import uuid
|
||||
import logging
|
||||
from typing import Dict, Any, Optional
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
from core.database import SessionLocal, Note
|
||||
from core.middleware import INTERNAL_TOOL_USER
|
||||
from src.auth_helpers import require_user
|
||||
from src.constants import DATA_DIR
|
||||
from src.upload_handler import reserve_upload_references
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request models
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class NoteCreate(BaseModel):
|
||||
title: str = ""
|
||||
content: Optional[str] = None
|
||||
items: Optional[list] = None
|
||||
note_type: str = "note"
|
||||
color: Optional[str] = None
|
||||
label: Optional[str] = None
|
||||
pinned: bool = False
|
||||
due_date: Optional[str] = None
|
||||
source: str = "user"
|
||||
session_id: Optional[str] = None
|
||||
image_url: Optional[str] = None
|
||||
repeat: Optional[str] = "none"
|
||||
sort_order: Optional[int] = None
|
||||
|
||||
|
||||
class NoteUpdate(BaseModel):
|
||||
title: Optional[str] = None
|
||||
content: Optional[str] = None
|
||||
items: Optional[list] = None
|
||||
note_type: Optional[str] = None
|
||||
color: Optional[str] = None
|
||||
label: Optional[str] = None
|
||||
pinned: Optional[bool] = None
|
||||
archived: Optional[bool] = None
|
||||
due_date: Optional[str] = None
|
||||
image_url: Optional[str] = None
|
||||
repeat: Optional[str] = None
|
||||
sort_order: Optional[int] = None
|
||||
agent_session_id: Optional[str] = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _note_to_dict(note: Note) -> Dict[str, Any]:
|
||||
items = None
|
||||
if note.items:
|
||||
try:
|
||||
items = json.loads(note.items)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
items = None
|
||||
ai_cls = None
|
||||
raw_ai = getattr(note, "ai_classification", None)
|
||||
if raw_ai:
|
||||
try:
|
||||
ai_cls = json.loads(raw_ai)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
ai_cls = None
|
||||
return {
|
||||
"id": note.id,
|
||||
"owner": note.owner,
|
||||
"title": note.title,
|
||||
"content": note.content,
|
||||
"items": items,
|
||||
"note_type": note.note_type,
|
||||
"color": note.color,
|
||||
"label": note.label,
|
||||
"pinned": note.pinned,
|
||||
"archived": note.archived,
|
||||
"due_date": note.due_date,
|
||||
"source": note.source,
|
||||
"session_id": note.session_id,
|
||||
"sort_order": note.sort_order or 0,
|
||||
"image_url": note.image_url,
|
||||
"repeat": note.repeat or "none",
|
||||
"ai_classification": ai_cls,
|
||||
"ai_content_hash": getattr(note, "ai_content_hash", None),
|
||||
"agent_session_id": getattr(note, "agent_session_id", None),
|
||||
"created_at": note.created_at.isoformat() if note.created_at else None,
|
||||
"updated_at": note.updated_at.isoformat() if note.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _reminder_text_from_note(note: Note) -> tuple[str, str]:
|
||||
"""Return the reminder title/body from a stored note row."""
|
||||
title = (note.title or "Note reminder").strip() or "Note reminder"
|
||||
if note.items:
|
||||
try:
|
||||
items = json.loads(note.items)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
items = None
|
||||
if isinstance(items, list):
|
||||
pending: list[str] = []
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("done") or item.get("checked"):
|
||||
continue
|
||||
text = str(item.get("text") or "").strip()
|
||||
if text:
|
||||
pending.append(text)
|
||||
if pending:
|
||||
shown = "\n".join(f"- {text}" for text in pending[:8])
|
||||
extra = f"\n...and {len(pending) - 8} more" if len(pending) > 8 else ""
|
||||
return title, f"Pending ({len(pending)}):\n{shown}{extra}"
|
||||
return title, f"{len(items)} item{'s' if len(items) != 1 else ''}"
|
||||
return title, (note.content or "").strip()[:400]
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Reminder dispatch — module-level so background tasks (built-in actions)
|
||||
# can call it directly without an HTTP roundtrip + auth cookie. The route
|
||||
# version below is a thin wrapper that pulls `owner` from the request.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Scheduler reference — set by setup_note_routes() so dispatch_reminder can
|
||||
# push a parallel in-app notification (frontend polls the scheduler's queue
|
||||
# and fires real browser Notification(...) popups). Optional; works without it.
|
||||
_scheduler_ref = None
|
||||
|
||||
|
||||
async def dispatch_reminder(
|
||||
title: str,
|
||||
note_body: str,
|
||||
note_id: str,
|
||||
owner: str = "",
|
||||
queue_browser: bool = True,
|
||||
settings_override: dict | None = None,
|
||||
) -> dict:
|
||||
"""Fire a reminder via the configured channel (browser/email/ntfy/webhook).
|
||||
|
||||
Args:
|
||||
title: short headline shown to the user
|
||||
note_body: longer body text
|
||||
note_id: stable id (used as tag/dedupe in browser notifications)
|
||||
owner: the user this reminder belongs to — scopes SMTP config to
|
||||
their account so we don't cross-leak credentials
|
||||
|
||||
Returns: {synthesis, email_sent, ntfy_sent}. Browser channel is wired via
|
||||
the in-memory notification queue picked up by the frontend poller, so
|
||||
nothing is "sent" synchronously for it — the channel just routes there.
|
||||
"""
|
||||
from src.settings import load_settings
|
||||
settings = {**load_settings(), **(settings_override or {})}
|
||||
channel = settings.get("reminder_channel", "browser")
|
||||
llm_on = bool(settings.get("reminder_llm_synthesis", False))
|
||||
title = (title or "").strip()
|
||||
note_body = (note_body or "").strip()
|
||||
cache_key = str(note_id) if note_id else ""
|
||||
cache = {}
|
||||
cache_path = None
|
||||
if cache_key:
|
||||
try:
|
||||
import json as _json
|
||||
from datetime import datetime as _dt, timezone as _tz, timedelta as _td
|
||||
from pathlib import Path as _P
|
||||
_slug = "".join(c if (c.isalnum() or c in "-_.@") else "_" for c in (owner or "default"))
|
||||
cache_path = _P(DATA_DIR) / f"note_pings_{_slug}.json"
|
||||
if cache_path.exists():
|
||||
cache = _json.loads(cache_path.read_text(encoding="utf-8"))
|
||||
last = cache.get(cache_key)
|
||||
if last:
|
||||
last_channel = None
|
||||
if isinstance(last, dict):
|
||||
last_channel = last.get("channel")
|
||||
last = last.get("at")
|
||||
last_dt = _dt.fromisoformat(str(last))
|
||||
if last_dt.tzinfo is None:
|
||||
last_dt = last_dt.replace(tzinfo=_tz.utc)
|
||||
# Legacy cache values were plain timestamps and could be
|
||||
# written by the frontend even when the email/ntfy send failed.
|
||||
# Treat those as browser-only dedupe so email reminders can be
|
||||
# retried by the backend scanner after a failed frontend path.
|
||||
should_skip = last_dt >= _dt.now(_tz.utc) - _td(minutes=25)
|
||||
if should_skip and channel in ("email", "ntfy", "webhook"):
|
||||
should_skip = last_channel == channel
|
||||
if should_skip:
|
||||
return {
|
||||
"synthesis": None,
|
||||
"email_sent": False,
|
||||
"ntfy_sent": False,
|
||||
"webhook_sent": False,
|
||||
"browser_sent": True,
|
||||
"skipped": True,
|
||||
}
|
||||
except Exception as _e:
|
||||
logger.debug(f"dispatch_reminder: cache read failed: {_e}")
|
||||
|
||||
synthesis = None
|
||||
_SYNTH_FAILED_TAG = "[utility model unavailable — no summary generated]"
|
||||
if llm_on:
|
||||
try:
|
||||
from src.endpoint_resolver import resolve_endpoint
|
||||
from src.llm_core import llm_call_async
|
||||
from src.reminder_personas import synthesis_system_prompt
|
||||
url, model, headers = resolve_endpoint("utility", owner=owner or None)
|
||||
if not url:
|
||||
url, model, headers = resolve_endpoint("default", owner=owner or None)
|
||||
if url and model:
|
||||
persona_id = (settings.get("reminder_llm_persona") or "").strip()
|
||||
sys_prompt = synthesis_system_prompt(persona_id)
|
||||
raw = await llm_call_async(
|
||||
url=url, model=model,
|
||||
messages=[
|
||||
{"role": "system", "content": sys_prompt},
|
||||
{"role": "user", "content": f"Title: {title}\n\n{note_body}".strip()},
|
||||
],
|
||||
temperature=0.7, max_tokens=200, headers=headers, timeout=30,
|
||||
)
|
||||
from src.text_helpers import strip_think as _strip_think
|
||||
# prose=True strips untagged "The user wants me to…" chain-of-thought.
|
||||
# prompt_echo=True strips Qwen-style "Thinking Process:" / leaked
|
||||
# prompt prefixes. Both are safe here because this is a
|
||||
# one-sentence LLM-only output, not user-pasted content.
|
||||
synthesis = _strip_think(raw or "", prose=True, prompt_echo=True)
|
||||
# Reminder synthesis is supposed to be ONE sentence. Strip-think's
|
||||
# paragraph-based heuristic misses cases where the model puts
|
||||
# reasoning + answer on consecutive lines inside one paragraph
|
||||
# (e.g. "I should write... [\n] You have one task waiting...").
|
||||
# Walk lines, drop reasoning/prompt-echo lines, then keep the
|
||||
# last surviving line — that's the actual warm sentence.
|
||||
if synthesis:
|
||||
import re as _re
|
||||
# Tightened: target ACTUAL self-talk (model narrating what
|
||||
# it'll do) rather than any first-person sentence. The old
|
||||
# pattern killed legit warm sentences like "I'll see you
|
||||
# tomorrow" or "I should be done by then". New rules:
|
||||
# • "I (need|should|have|'ll|will) (write|draft|reply|…)"
|
||||
# only matches when followed by a TASK verb taking an
|
||||
# OBJECT (so first-person + intransitive verb passes).
|
||||
# • Self-instructional patterns the model emits verbatim:
|
||||
# "I should write something that reminds them…",
|
||||
# "I need to draft…", "Let me think…".
|
||||
# • Explicit instructions echoed back from the prompt:
|
||||
# "Keep it under 25 words", "No greetings".
|
||||
_reasoning = _re.compile(
|
||||
r"^\s*(?:"
|
||||
# "I should write/draft/compose…" with a task-object follow
|
||||
r"i (?:need|should|have|'ll|will|am going|am)\s+to\s+"
|
||||
r"(?:write|draft|compose|craft|generate|produce|create|"
|
||||
r"summarize|answer|provide|note|address|remind|output)"
|
||||
r"\s+(?:a |an |the |something|this|that|here|them|him|her|"
|
||||
r"you|user|reply|response|sentence|message|line|warm)|"
|
||||
# The model literally narrating about the user
|
||||
r"the user (?:wants|is asking|asks|needs|wrote|said|requested) (?:me )?(?:to|for|that|about|something)|"
|
||||
# "Let me [think/write/draft/…] (about/for/the …)"
|
||||
r"let me (?:think|write|draft|consider|note|see|check)\b\s+(?:about|for|the|this|that|if|whether)|"
|
||||
# "Looking at the/this/that …"
|
||||
r"looking at (?:the|this|that)\b|"
|
||||
# "Based on the/this/what …"
|
||||
r"based on (?:the|this|what|context|that)\b|"
|
||||
# Prompt-echo of length / style instructions
|
||||
r"keep it under \d+ words\b|"
|
||||
r"(?:no greetings|no preamble|no hashtags|just output the)\b"
|
||||
r").*",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
# Echo of the prompt's "Pending:" / "<N> pending" tail.
|
||||
_echo = _re.compile(
|
||||
r"^\s*(?:pending\s*[:.]|(?:\d+|one|two|three|four|five)\s+pending\b)",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
lines = [ln for ln in synthesis.splitlines() if ln.strip()]
|
||||
cleaned = [ln for ln in lines if not _reasoning.match(ln) and not _echo.match(ln)]
|
||||
if cleaned:
|
||||
# The model's actual answer is normally the LAST surviving
|
||||
# line — reasoning leads, answer trails.
|
||||
synthesis = cleaned[-1].strip()
|
||||
else:
|
||||
synthesis = _SYNTH_FAILED_TAG
|
||||
except Exception as e:
|
||||
logger.warning(f"Reminder LLM synthesis failed: {e}")
|
||||
synthesis = _SYNTH_FAILED_TAG
|
||||
if synthesis:
|
||||
_s = synthesis.strip(); _low = _s.lower()
|
||||
if (not _s or _low.startswith("error:") or _low.startswith("[error")
|
||||
or "operation failed" in _low
|
||||
or ("upstream" in _low and "failed" in _low)) and synthesis != _SYNTH_FAILED_TAG:
|
||||
logger.warning(f"Reminder synthesis looked like an error, replacing: {_s[:120]!r}")
|
||||
synthesis = _SYNTH_FAILED_TAG
|
||||
|
||||
email_sent = False
|
||||
email_error = ""
|
||||
if channel == "email":
|
||||
try:
|
||||
from routes.email_routes import _get_email_config, _smtp_ready
|
||||
from email.mime.text import MIMEText
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from datetime import datetime as _dt
|
||||
# `reminder_email_account_id` lets the user pick WHICH email
|
||||
# account to send reminders from (when they have several
|
||||
# configured in Integrations). Falls back to the default
|
||||
# account when no explicit choice is saved.
|
||||
_acc_id = (settings.get("reminder_email_account_id") or "").strip() or None
|
||||
cfg = _get_email_config(account_id=_acc_id, owner=owner or "")
|
||||
if not _smtp_ready(cfg):
|
||||
try:
|
||||
from core.database import SessionLocal as _SL, EmailAccount as _EA
|
||||
from sqlalchemy import and_, or_
|
||||
db = _SL()
|
||||
try:
|
||||
q = db.query(_EA).filter(_EA.enabled == True) # noqa: E712
|
||||
if owner:
|
||||
unowned = or_(_EA.owner == None, _EA.owner == "") # noqa: E711
|
||||
same_mailbox = or_(_EA.imap_user == owner, _EA.from_address == owner)
|
||||
q = q.filter(or_(_EA.owner == owner, and_(unowned, same_mailbox)))
|
||||
for row in q.order_by(_EA.is_default.desc(), _EA.created_at.asc()).all():
|
||||
trial = _get_email_config(account_id=row.id, owner=owner or "")
|
||||
if _smtp_ready(trial):
|
||||
cfg = trial
|
||||
break
|
||||
finally:
|
||||
db.close()
|
||||
except Exception as _fallback_error:
|
||||
logger.debug(f"Reminder SMTP fallback lookup failed: {_fallback_error}")
|
||||
from_addr = (cfg.get("from_address") or cfg.get("smtp_user") or "").strip()
|
||||
recipient = (settings.get("reminder_email_to") or "").strip() or from_addr
|
||||
# Loud diagnostic so we can see WHY a reminder didn't send (the
|
||||
# previous "silently no-op when cfg has no smtp_host" was invisible).
|
||||
logger.info(
|
||||
"dispatch_reminder[email] note_id=%s owner=%r "
|
||||
"has_smtp_host=%s has_smtp_user=%s has_from=%s has_recipient=%s",
|
||||
note_id, owner,
|
||||
bool(cfg.get("smtp_host")), bool(cfg.get("smtp_user")),
|
||||
bool(from_addr), bool(recipient),
|
||||
)
|
||||
missing = []
|
||||
if not cfg.get("smtp_host"):
|
||||
missing.append("SMTP host")
|
||||
if not cfg.get("smtp_user"):
|
||||
missing.append("SMTP user")
|
||||
if not (cfg.get("smtp_password") or cfg.get("oauth_provider")):
|
||||
missing.append("SMTP credentials")
|
||||
if not from_addr:
|
||||
missing.append("from address")
|
||||
if not recipient:
|
||||
missing.append("recipient")
|
||||
if missing:
|
||||
email_error = "Missing " + ", ".join(missing)
|
||||
logger.warning(
|
||||
"Reminder email not sent for note_id=%s account=%r: %s",
|
||||
note_id, cfg.get("account_name"), email_error,
|
||||
)
|
||||
else:
|
||||
msg = MIMEMultipart("alternative")
|
||||
msg["From"] = from_addr
|
||||
msg["To"] = recipient
|
||||
_t = title or 'Note'
|
||||
_t = _t[len('Reminder:'):].strip() if _t.lower().startswith('reminder:') else _t
|
||||
msg["Subject"] = f"Reminder (Odysseus): {_t}"
|
||||
msg["Date"] = _dt.utcnow().strftime("%a, %d %b %Y %H:%M:%S +0000")
|
||||
msg["X-Odysseus-Origin"] = "odysseus-ui"
|
||||
msg["X-Odysseus-Kind"] = "reminder"
|
||||
msg["X-Odysseus-Ref"] = str(note_id)
|
||||
# Body shape: synthesis (warm sentence) → blank line → bold
|
||||
# title header → note details. The title was previously only
|
||||
# in the subject line, so the email read like a faceless
|
||||
# to-do list with no anchor to which note triggered it.
|
||||
_body_chunks = []
|
||||
if synthesis:
|
||||
_body_chunks.append(synthesis)
|
||||
if _t:
|
||||
_body_chunks.append(_t)
|
||||
if note_body:
|
||||
_body_chunks.append(note_body)
|
||||
plain = "\n\n".join(_body_chunks) if _body_chunks else title
|
||||
msg.attach(MIMEText(plain, "plain", "utf-8"))
|
||||
|
||||
def _smtp_send():
|
||||
from routes.email_helpers import _send_smtp_message
|
||||
_send_smtp_message(cfg, from_addr, [recipient], msg.as_string())
|
||||
|
||||
import asyncio as _aio
|
||||
await _aio.to_thread(_smtp_send)
|
||||
email_sent = True
|
||||
except Exception as e:
|
||||
email_error = str(e) or e.__class__.__name__
|
||||
logger.warning(f"Reminder email send failed: {e}")
|
||||
|
||||
webhook_sent = False
|
||||
webhook_error = ""
|
||||
if channel == "webhook":
|
||||
try:
|
||||
import httpx
|
||||
import json as _wjson
|
||||
from src.integrations import load_integrations
|
||||
# Built-in payload defaults for known presets so users don't have
|
||||
# to configure a template just to use a standard service.
|
||||
_PRESET_TEMPLATE_DEFAULTS = {
|
||||
"discord_webhook": '{"embeds": [{"title": "{{title}}", "description": "{{message}}", "color": 5793266}]}',
|
||||
}
|
||||
intg_id = settings.get("reminder_webhook_integration_id", "").strip()
|
||||
template = settings.get("reminder_webhook_payload_template", "").strip()
|
||||
if not intg_id:
|
||||
webhook_error = "No webhook integration selected"
|
||||
else:
|
||||
intg = next(
|
||||
(i for i in load_integrations()
|
||||
if i.get("id") == intg_id and i.get("base_url")),
|
||||
None,
|
||||
)
|
||||
if not intg:
|
||||
webhook_error = f"Integration {intg_id!r} not found or missing base URL"
|
||||
else:
|
||||
# Fall back to a built-in default for known presets so
|
||||
# users don't have to configure a template for standard
|
||||
# services like Discord.
|
||||
if not template:
|
||||
template = _PRESET_TEMPLATE_DEFAULTS.get(intg.get("preset", ""), "")
|
||||
if not template:
|
||||
webhook_error = "No payload template configured"
|
||||
else:
|
||||
# Render template: JSON-escape the values so the result
|
||||
# is always valid JSON regardless of special characters.
|
||||
# dumps() returns `"value"` — strip outer quotes.
|
||||
msg = (synthesis or note_body or title or "Reminder")[:4000]
|
||||
_t = _wjson.dumps(title or "Reminder")[1:-1]
|
||||
_m = _wjson.dumps(msg)[1:-1]
|
||||
rendered = template.replace("{{title}}", _t).replace("{{message}}", _m)
|
||||
hdrs = {"Content-Type": "application/json"}
|
||||
api_key = intg.get("api_key", "")
|
||||
auth_type = (intg.get("auth_type") or "none").lower()
|
||||
if api_key:
|
||||
if auth_type == "bearer":
|
||||
hdrs["Authorization"] = f"Bearer {api_key}"
|
||||
elif auth_type == "header":
|
||||
hdrs[intg.get("auth_header") or "Authorization"] = api_key
|
||||
url = intg["base_url"].rstrip("/")
|
||||
# SSRF guard — matches the pattern used by webhook_routes,
|
||||
# CalDAV, search, and embeddings. Blocks link-local / metadata
|
||||
# addresses (169.254.x.x) by default; set
|
||||
# REMINDER_WEBHOOK_BLOCK_PRIVATE_IPS=true to also block
|
||||
# RFC-1918 ranges for locked-down deployments.
|
||||
import os as _os
|
||||
from src.url_safety import check_outbound_url as _chk
|
||||
_block = _os.getenv("REMINDER_WEBHOOK_BLOCK_PRIVATE_IPS", "false").lower() == "true"
|
||||
_ok, _reason = _chk(url, block_private=_block)
|
||||
if not _ok:
|
||||
webhook_error = f"Webhook URL rejected: {_reason}"
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
resp = await client.post(url, content=rendered.encode(), headers=hdrs)
|
||||
webhook_sent = resp.is_success
|
||||
if not webhook_sent:
|
||||
webhook_error = f"Webhook returned HTTP {resp.status_code}"
|
||||
except Exception as e:
|
||||
webhook_error = str(e) or e.__class__.__name__
|
||||
logger.warning(f"Reminder webhook send failed: {e}")
|
||||
|
||||
ntfy_sent = False
|
||||
ntfy_error = ""
|
||||
if channel == "ntfy":
|
||||
try:
|
||||
from src.integrations import load_integrations
|
||||
import httpx
|
||||
intg = next(
|
||||
(i for i in load_integrations()
|
||||
if i.get("preset") == "ntfy" and i.get("enabled", True) and i.get("base_url")),
|
||||
None,
|
||||
)
|
||||
if intg:
|
||||
base = intg["base_url"].rstrip("/")
|
||||
topic = settings.get("reminder_ntfy_topic") or "reminders"
|
||||
ntfy_body = synthesis or note_body or title
|
||||
# ntfy Title is an ASCII HTTP header; sanitize Unicode and cap its length.
|
||||
_clean_title = (title or "Reminder").encode("ascii", "replace").decode("ascii")[:200]
|
||||
hdrs = {"Title": _clean_title, "Priority": "high", "Tags": "bell"}
|
||||
api_key = intg.get("api_key", "")
|
||||
if api_key:
|
||||
hdrs["Authorization"] = f"Bearer {api_key}"
|
||||
# SSRF guard — same check (and env knob) as the webhook branch
|
||||
# above: link-local / metadata addresses are always rejected;
|
||||
# REMINDER_WEBHOOK_BLOCK_PRIVATE_IPS=true also blocks RFC-1918
|
||||
# so a ntfy base_url can't be pointed at internal services.
|
||||
import os as _os
|
||||
from src.url_safety import check_outbound_url as _chk
|
||||
_block = _os.getenv("REMINDER_WEBHOOK_BLOCK_PRIVATE_IPS", "false").lower() == "true"
|
||||
_ok, _reason = _chk(f"{base}/{topic}", block_private=_block)
|
||||
if not _ok:
|
||||
ntfy_error = f"ntfy URL rejected: {_reason}"
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
resp = await client.post(f"{base}/{topic}", content=ntfy_body, headers=hdrs)
|
||||
ntfy_sent = resp.is_success
|
||||
if not ntfy_sent:
|
||||
ntfy_error = f"ntfy returned HTTP {resp.status_code}"
|
||||
else:
|
||||
ntfy_error = "No enabled ntfy integration"
|
||||
except Exception as e:
|
||||
ntfy_error = str(e) or e.__class__.__name__
|
||||
logger.warning(f"Reminder ntfy send failed: {e}")
|
||||
|
||||
# In-app browser notification ALWAYS fires (regardless of channel). The
|
||||
# frontend polls `/api/tasks/notifications` and turns any entry with a
|
||||
# `body` into a real `Notification(...)` — same surface as task-success
|
||||
# popups. Lets the user see reminders inside the app even when the
|
||||
# primary channel is email/ntfy and the tab is open.
|
||||
browser_sent = False
|
||||
local_browser_sent = (not queue_browser and channel == "browser")
|
||||
if queue_browser and _scheduler_ref is not None:
|
||||
try:
|
||||
_scheduler_ref.add_notification(
|
||||
task_name=title or "Reminder",
|
||||
status="success",
|
||||
task_id=f"reminder-{note_id}",
|
||||
owner=owner or None,
|
||||
body=(synthesis or note_body or title or "").strip()[:500] or "Reminder",
|
||||
)
|
||||
browser_sent = True
|
||||
except Exception as _e:
|
||||
logger.debug(f"dispatch_reminder: in-app notif push failed: {_e}")
|
||||
|
||||
# Dedupe across paths: write to the same cache file `action_ping_notes`
|
||||
# reads, so the background scanner's REPING_MIN window suppresses a
|
||||
# second send for the same note within 25 min. Without this, a note
|
||||
# whose due_date fires while the user has the app open got TWO emails
|
||||
# (frontend-fired here + background-fired by ping_notes 0–5 min later).
|
||||
if (email_sent or ntfy_sent or webhook_sent or browser_sent or local_browser_sent) and note_id:
|
||||
try:
|
||||
import json as _json
|
||||
from datetime import datetime as _dt, timezone as _tz
|
||||
from pathlib import Path as _P
|
||||
# Per-owner cache so the scanner's prune step on user A's run
|
||||
# doesn't drop user B's just-fired entry (review C4).
|
||||
_STATE = cache_path
|
||||
if _STATE is None:
|
||||
_slug = "".join(c if (c.isalnum() or c in "-_.@") else "_" for c in (owner or "default"))
|
||||
_STATE = _P(DATA_DIR) / f"note_pings_{_slug}.json"
|
||||
_STATE.parent.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
_cache = cache or (_json.loads(_STATE.read_text(encoding="utf-8")) if _STATE.exists() else {})
|
||||
except Exception:
|
||||
_cache = {}
|
||||
sent_channel = "email" if email_sent else "ntfy" if ntfy_sent else "webhook" if webhook_sent else "browser"
|
||||
_cache[cache_key or str(note_id)] = {
|
||||
"at": _dt.now(_tz.utc).isoformat(),
|
||||
"channel": sent_channel,
|
||||
}
|
||||
_STATE.write_text(_json.dumps(_cache), encoding="utf-8")
|
||||
except Exception as _e:
|
||||
logger.debug(f"dispatch_reminder: cache write failed: {_e}")
|
||||
|
||||
return {
|
||||
"channel": channel,
|
||||
"synthesis": synthesis,
|
||||
"email_sent": email_sent,
|
||||
"email_error": email_error,
|
||||
"ntfy_sent": ntfy_sent,
|
||||
"ntfy_error": ntfy_error,
|
||||
"webhook_sent": webhook_sent,
|
||||
"webhook_error": webhook_error,
|
||||
"browser_sent": browser_sent or local_browser_sent,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Router factory
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def setup_note_routes(task_scheduler=None, upload_handler=None):
|
||||
# Expose the scheduler to module-level `dispatch_reminder` so reminders
|
||||
# can also push to the in-app notification queue (the polling system
|
||||
# turns each entry into a real browser Notification + the existing
|
||||
# tasks-tab badge / dot system).
|
||||
global _scheduler_ref
|
||||
_scheduler_ref = task_scheduler
|
||||
|
||||
router = APIRouter(prefix="/api/notes", tags=["notes"])
|
||||
|
||||
def _owner(request: Request) -> Optional[str]:
|
||||
# require_user, not bare get_current_user: a request that reaches
|
||||
# these owner-scoped routes with NO identity (auth-middleware
|
||||
# regression, SSRF from a sibling service) must fail closed (401)
|
||||
# when auth is configured — not be treated as the single-user mode
|
||||
# and handed blanket access to every account's notes. The documented
|
||||
# anonymous modes (AUTH_ENABLED=false, LOCALHOST_BYPASS on loopback,
|
||||
# unconfigured first-run) still resolve to None, the single-user
|
||||
# path. fire_reminder below already gated this way; the CRUD routes
|
||||
# did not.
|
||||
return require_user(request) or None
|
||||
|
||||
def _reserve_note_uploads(owner: Optional[str], *values) -> None:
|
||||
missing_id = reserve_upload_references(upload_handler, owner, *values)
|
||||
if missing_id:
|
||||
raise HTTPException(409, f"Referenced upload is no longer available: {missing_id}")
|
||||
|
||||
def _is_admin_or_single_user(request: Request, user: str | None) -> bool:
|
||||
if user == INTERNAL_TOOL_USER:
|
||||
return True
|
||||
if not user:
|
||||
# require_user() already admitted this request, which only happens
|
||||
# for auth-disabled, loopback-bypass, or unconfigured single-user
|
||||
# modes. There is no separate non-admin account boundary there.
|
||||
return True
|
||||
try:
|
||||
from core.auth import AuthManager
|
||||
auth_mgr = getattr(request.app.state, "auth_manager", None) or AuthManager()
|
||||
if not getattr(auth_mgr, "is_configured", True):
|
||||
return True
|
||||
return bool(auth_mgr.is_admin(user))
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
# --- LIST ---
|
||||
@router.get("")
|
||||
def list_notes(
|
||||
request: Request,
|
||||
archived: Optional[bool] = None,
|
||||
label: Optional[str] = None,
|
||||
):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(Note)
|
||||
if user is not None:
|
||||
q = q.filter(Note.owner == user)
|
||||
if archived is not None:
|
||||
q = q.filter(Note.archived == archived)
|
||||
else:
|
||||
q = q.filter(Note.archived == False)
|
||||
if label:
|
||||
q = q.filter(Note.label == label)
|
||||
# Archived view: most recently archived first. Active view: pin + manual order.
|
||||
if archived is True:
|
||||
notes = q.order_by(Note.updated_at.desc()).all()
|
||||
else:
|
||||
notes = q.order_by(Note.pinned.desc(), Note.sort_order.asc(), Note.updated_at.desc()).all()
|
||||
return {"notes": [_note_to_dict(n) for n in notes]}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- CREATE ---
|
||||
@router.post("")
|
||||
def create_note(request: Request, body: NoteCreate):
|
||||
user = _owner(request)
|
||||
_reserve_note_uploads(
|
||||
user,
|
||||
body.image_url,
|
||||
body.color,
|
||||
body.content,
|
||||
json.dumps(body.items) if body.items is not None else None,
|
||||
)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = Note(
|
||||
id=str(uuid.uuid4()),
|
||||
owner=user,
|
||||
title=body.title,
|
||||
content=body.content,
|
||||
items=json.dumps(body.items) if body.items is not None else None,
|
||||
note_type=body.note_type,
|
||||
color=body.color,
|
||||
label=body.label,
|
||||
pinned=body.pinned,
|
||||
due_date=body.due_date,
|
||||
source=body.source,
|
||||
session_id=body.session_id,
|
||||
image_url=body.image_url,
|
||||
repeat=body.repeat or "none",
|
||||
sort_order=body.sort_order if body.sort_order is not None else 0,
|
||||
)
|
||||
db.add(note)
|
||||
db.commit()
|
||||
db.refresh(note)
|
||||
return _note_to_dict(note)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- GET ONE ---
|
||||
@router.get("/{note_id}")
|
||||
def get_note(request: Request, note_id: str):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
return _note_to_dict(note)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- UPDATE ---
|
||||
@router.put("/{note_id}")
|
||||
def update_note(request: Request, note_id: str, body: NoteUpdate):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
|
||||
_reserve_note_uploads(
|
||||
user,
|
||||
body.image_url,
|
||||
body.color,
|
||||
body.content,
|
||||
json.dumps(body.items) if body.items is not None else None,
|
||||
)
|
||||
if body.title is not None:
|
||||
note.title = body.title
|
||||
if body.content is not None:
|
||||
note.content = body.content
|
||||
if body.items is not None:
|
||||
note.items = json.dumps(body.items)
|
||||
flag_modified(note, "items")
|
||||
if body.note_type is not None:
|
||||
note.note_type = body.note_type
|
||||
if body.color is not None:
|
||||
note.color = body.color
|
||||
if body.label is not None:
|
||||
note.label = body.label
|
||||
if body.pinned is not None:
|
||||
note.pinned = body.pinned
|
||||
if body.archived is not None:
|
||||
note.archived = body.archived
|
||||
if body.due_date is not None:
|
||||
note.due_date = body.due_date
|
||||
if body.image_url is not None:
|
||||
note.image_url = body.image_url
|
||||
if body.repeat is not None:
|
||||
note.repeat = body.repeat
|
||||
if body.sort_order is not None:
|
||||
note.sort_order = body.sort_order
|
||||
if body.agent_session_id is not None:
|
||||
note.agent_session_id = body.agent_session_id
|
||||
|
||||
db.commit()
|
||||
db.refresh(note)
|
||||
return _note_to_dict(note)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- DELETE ---
|
||||
@router.delete("/{note_id}")
|
||||
def delete_note(request: Request, note_id: str):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
db.delete(note)
|
||||
db.commit()
|
||||
return {"ok": True}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- TOGGLE PIN ---
|
||||
@router.post("/{note_id}/pin")
|
||||
def toggle_pin(request: Request, note_id: str):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
note.pinned = not note.pinned
|
||||
db.commit()
|
||||
return {"ok": True, "pinned": note.pinned}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- TOGGLE ARCHIVE ---
|
||||
@router.post("/{note_id}/archive")
|
||||
def toggle_archive(request: Request, note_id: str):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
note.archived = not note.archived
|
||||
db.commit()
|
||||
return {"ok": True, "archived": note.archived}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- TOGGLE CHECKLIST ITEM ---
|
||||
@router.post("/{note_id}/items/{index}/toggle")
|
||||
def toggle_item(request: Request, note_id: str, index: int):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
if not note.items:
|
||||
raise HTTPException(400, "Note has no checklist items")
|
||||
items = json.loads(note.items)
|
||||
if index < 0 or index >= len(items):
|
||||
raise HTTPException(400, f"Item index {index} out of range")
|
||||
items[index]["done"] = not items[index].get("done", False)
|
||||
note.items = json.dumps(items)
|
||||
flag_modified(note, "items")
|
||||
db.commit()
|
||||
return {"ok": True, "items": items}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- FIRE REMINDER ---
|
||||
@router.post("/fire-reminder")
|
||||
async def fire_reminder(request: Request):
|
||||
"""Dispatch a reminder according to user settings.
|
||||
|
||||
Called by the frontend when a reminder fires. Optionally generates an
|
||||
LLM synthesis line and/or sends an email through configured SMTP.
|
||||
Returns {synthesis, email_sent}.
|
||||
"""
|
||||
# Gate against anonymous callers — LLM synthesis can burn tokens.
|
||||
user = require_user(request)
|
||||
body = await request.json()
|
||||
note_id = str(body.get("note_id") or "").strip()
|
||||
if not note_id:
|
||||
raise HTTPException(400, "note_id required")
|
||||
|
||||
caller = _owner(request)
|
||||
is_test = note_id.startswith("test-")
|
||||
is_admin = _is_admin_or_single_user(request, user or caller)
|
||||
_override: dict = {}
|
||||
if is_test:
|
||||
if not is_admin:
|
||||
raise HTTPException(403, "Admin only")
|
||||
title = (body.get("title") or "Test Reminder").strip() or "Test Reminder"
|
||||
note_body = (body.get("body") or "").strip()
|
||||
# Optional overrides let the admin settings test button pass the
|
||||
# current UI values directly so it never races a pending save.
|
||||
if body.get("channel"):
|
||||
_override["reminder_channel"] = body["channel"]
|
||||
if body.get("webhook_integration_id"):
|
||||
_override["reminder_webhook_integration_id"] = body["webhook_integration_id"]
|
||||
if body.get("webhook_payload_template"):
|
||||
_override["reminder_webhook_payload_template"] = body["webhook_payload_template"]
|
||||
# Mirror the in-UI AI Synthesis toggle + persona so the test
|
||||
# actually exercises the synthesis path before/without a Save.
|
||||
if "llm_synthesis" in body:
|
||||
_override["reminder_llm_synthesis"] = bool(body["llm_synthesis"])
|
||||
if "llm_persona" in body:
|
||||
_override["reminder_llm_persona"] = str(body["llm_persona"] or "")
|
||||
else:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
if caller is not None and note.owner != caller:
|
||||
raise HTTPException(404, "Note not found")
|
||||
title, note_body = _reminder_text_from_note(note)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return await dispatch_reminder(
|
||||
title=title, note_body=note_body, note_id=note_id,
|
||||
owner=caller or "",
|
||||
queue_browser=False,
|
||||
settings_override=_override or None,
|
||||
)
|
||||
|
||||
# --- REORDER NOTES ---
|
||||
@router.post("/reorder")
|
||||
async def reorder_notes(request: Request):
|
||||
"""Update sort_order for a list of note IDs in the order provided."""
|
||||
user = _owner(request)
|
||||
body = await request.json()
|
||||
ids = body.get("ids", [])
|
||||
if not isinstance(ids, list):
|
||||
raise HTTPException(400, "ids must be a list")
|
||||
# v2 review HIGH-12: drop the legacy `(owner == user) | (owner ==
|
||||
# None)` OR which let an authenticated user silently reorder
|
||||
# every legacy-null-owner note belonging to other accounts. In
|
||||
# an unconfigured (single-user) auth deploy the OR is still safe
|
||||
# because there's no second user to attack; we keep that branch
|
||||
# explicit and gated on AuthManager.is_configured.
|
||||
try:
|
||||
from core.auth import AuthManager
|
||||
_allow_null = not AuthManager().is_configured
|
||||
except Exception:
|
||||
_allow_null = False
|
||||
db = SessionLocal()
|
||||
try:
|
||||
for i, nid in enumerate(ids):
|
||||
q = db.query(Note).filter(Note.id == nid)
|
||||
if user is not None:
|
||||
if _allow_null:
|
||||
q = q.filter((Note.owner == user) | (Note.owner == None)) # noqa: E711
|
||||
else:
|
||||
q = q.filter(Note.owner == user)
|
||||
note = q.first()
|
||||
if note:
|
||||
note.sort_order = i
|
||||
db.commit()
|
||||
return {"ok": True, "count": len(ids)}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return router
|
||||
+14
-911
@@ -1,915 +1,18 @@
|
||||
# routes/note_routes.py
|
||||
"""Google Keep-style notes / checklists API."""
|
||||
"""Backward-compat shim — canonical location is routes/note/note_routes.py.
|
||||
|
||||
import json
|
||||
import uuid
|
||||
import logging
|
||||
from typing import Dict, Any, Optional
|
||||
This module is replaced in ``sys.modules`` by the canonical module object so
|
||||
that ``import routes.note_routes``, ``from routes.note_routes import X``,
|
||||
``importlib.import_module("routes.note_routes")``, and the
|
||||
``import ... as note_routes`` + ``monkeypatch.setattr(note_routes, "SessionLocal",
|
||||
...)`` pattern used by test_note_reminder_fire_scope.py /
|
||||
test_notes_fail_closed_auth.py all operate on the *same* object the
|
||||
application actually uses. Keeps existing import paths working after
|
||||
slice 2f (#4082/#4071). Source-introspection tests read the canonical file
|
||||
by path.
|
||||
"""
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
import sys as _sys
|
||||
|
||||
from core.database import SessionLocal, Note
|
||||
from core.middleware import INTERNAL_TOOL_USER
|
||||
from src.auth_helpers import require_user
|
||||
from src.constants import DATA_DIR
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
from routes.note import note_routes as _canonical # noqa: F401
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request models
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class NoteCreate(BaseModel):
|
||||
title: str = ""
|
||||
content: Optional[str] = None
|
||||
items: Optional[list] = None
|
||||
note_type: str = "note"
|
||||
color: Optional[str] = None
|
||||
label: Optional[str] = None
|
||||
pinned: bool = False
|
||||
due_date: Optional[str] = None
|
||||
source: str = "user"
|
||||
session_id: Optional[str] = None
|
||||
image_url: Optional[str] = None
|
||||
repeat: Optional[str] = "none"
|
||||
sort_order: Optional[int] = None
|
||||
|
||||
|
||||
class NoteUpdate(BaseModel):
|
||||
title: Optional[str] = None
|
||||
content: Optional[str] = None
|
||||
items: Optional[list] = None
|
||||
note_type: Optional[str] = None
|
||||
color: Optional[str] = None
|
||||
label: Optional[str] = None
|
||||
pinned: Optional[bool] = None
|
||||
archived: Optional[bool] = None
|
||||
due_date: Optional[str] = None
|
||||
image_url: Optional[str] = None
|
||||
repeat: Optional[str] = None
|
||||
sort_order: Optional[int] = None
|
||||
agent_session_id: Optional[str] = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _note_to_dict(note: Note) -> Dict[str, Any]:
|
||||
items = None
|
||||
if note.items:
|
||||
try:
|
||||
items = json.loads(note.items)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
items = None
|
||||
ai_cls = None
|
||||
raw_ai = getattr(note, "ai_classification", None)
|
||||
if raw_ai:
|
||||
try:
|
||||
ai_cls = json.loads(raw_ai)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
ai_cls = None
|
||||
return {
|
||||
"id": note.id,
|
||||
"owner": note.owner,
|
||||
"title": note.title,
|
||||
"content": note.content,
|
||||
"items": items,
|
||||
"note_type": note.note_type,
|
||||
"color": note.color,
|
||||
"label": note.label,
|
||||
"pinned": note.pinned,
|
||||
"archived": note.archived,
|
||||
"due_date": note.due_date,
|
||||
"source": note.source,
|
||||
"session_id": note.session_id,
|
||||
"sort_order": note.sort_order or 0,
|
||||
"image_url": note.image_url,
|
||||
"repeat": note.repeat or "none",
|
||||
"ai_classification": ai_cls,
|
||||
"ai_content_hash": getattr(note, "ai_content_hash", None),
|
||||
"agent_session_id": getattr(note, "agent_session_id", None),
|
||||
"created_at": note.created_at.isoformat() if note.created_at else None,
|
||||
"updated_at": note.updated_at.isoformat() if note.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _reminder_text_from_note(note: Note) -> tuple[str, str]:
|
||||
"""Return the reminder title/body from a stored note row."""
|
||||
title = (note.title or "Note reminder").strip() or "Note reminder"
|
||||
if note.items:
|
||||
try:
|
||||
items = json.loads(note.items)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
items = None
|
||||
if isinstance(items, list):
|
||||
pending: list[str] = []
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("done") or item.get("checked"):
|
||||
continue
|
||||
text = str(item.get("text") or "").strip()
|
||||
if text:
|
||||
pending.append(text)
|
||||
if pending:
|
||||
shown = "\n".join(f"- {text}" for text in pending[:8])
|
||||
extra = f"\n...and {len(pending) - 8} more" if len(pending) > 8 else ""
|
||||
return title, f"Pending ({len(pending)}):\n{shown}{extra}"
|
||||
return title, f"{len(items)} item{'s' if len(items) != 1 else ''}"
|
||||
return title, (note.content or "").strip()[:400]
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Reminder dispatch — module-level so background tasks (built-in actions)
|
||||
# can call it directly without an HTTP roundtrip + auth cookie. The route
|
||||
# version below is a thin wrapper that pulls `owner` from the request.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Scheduler reference — set by setup_note_routes() so dispatch_reminder can
|
||||
# push a parallel in-app notification (frontend polls the scheduler's queue
|
||||
# and fires real browser Notification(...) popups). Optional; works without it.
|
||||
_scheduler_ref = None
|
||||
|
||||
|
||||
async def dispatch_reminder(
|
||||
title: str,
|
||||
note_body: str,
|
||||
note_id: str,
|
||||
owner: str = "",
|
||||
queue_browser: bool = True,
|
||||
settings_override: dict | None = None,
|
||||
) -> dict:
|
||||
"""Fire a reminder via the configured channel (browser/email/ntfy/webhook).
|
||||
|
||||
Args:
|
||||
title: short headline shown to the user
|
||||
note_body: longer body text
|
||||
note_id: stable id (used as tag/dedupe in browser notifications)
|
||||
owner: the user this reminder belongs to — scopes SMTP config to
|
||||
their account so we don't cross-leak credentials
|
||||
|
||||
Returns: {synthesis, email_sent, ntfy_sent}. Browser channel is wired via
|
||||
the in-memory notification queue picked up by the frontend poller, so
|
||||
nothing is "sent" synchronously for it — the channel just routes there.
|
||||
"""
|
||||
from src.settings import load_settings
|
||||
settings = {**load_settings(), **(settings_override or {})}
|
||||
channel = settings.get("reminder_channel", "browser")
|
||||
llm_on = bool(settings.get("reminder_llm_synthesis", False))
|
||||
title = (title or "").strip()
|
||||
note_body = (note_body or "").strip()
|
||||
cache_key = str(note_id) if note_id else ""
|
||||
cache = {}
|
||||
cache_path = None
|
||||
if cache_key:
|
||||
try:
|
||||
import json as _json
|
||||
from datetime import datetime as _dt, timezone as _tz, timedelta as _td
|
||||
from pathlib import Path as _P
|
||||
_slug = "".join(c if (c.isalnum() or c in "-_.@") else "_" for c in (owner or "default"))
|
||||
cache_path = _P(DATA_DIR) / f"note_pings_{_slug}.json"
|
||||
if cache_path.exists():
|
||||
cache = _json.loads(cache_path.read_text(encoding="utf-8"))
|
||||
last = cache.get(cache_key)
|
||||
if last:
|
||||
last_channel = None
|
||||
if isinstance(last, dict):
|
||||
last_channel = last.get("channel")
|
||||
last = last.get("at")
|
||||
last_dt = _dt.fromisoformat(str(last))
|
||||
if last_dt.tzinfo is None:
|
||||
last_dt = last_dt.replace(tzinfo=_tz.utc)
|
||||
# Legacy cache values were plain timestamps and could be
|
||||
# written by the frontend even when the email/ntfy send failed.
|
||||
# Treat those as browser-only dedupe so email reminders can be
|
||||
# retried by the backend scanner after a failed frontend path.
|
||||
should_skip = last_dt >= _dt.now(_tz.utc) - _td(minutes=25)
|
||||
if should_skip and channel in ("email", "ntfy", "webhook"):
|
||||
should_skip = last_channel == channel
|
||||
if should_skip:
|
||||
return {
|
||||
"synthesis": None,
|
||||
"email_sent": False,
|
||||
"ntfy_sent": False,
|
||||
"webhook_sent": False,
|
||||
"browser_sent": True,
|
||||
"skipped": True,
|
||||
}
|
||||
except Exception as _e:
|
||||
logger.debug(f"dispatch_reminder: cache read failed: {_e}")
|
||||
|
||||
synthesis = None
|
||||
_SYNTH_FAILED_TAG = "[utility model unavailable — no summary generated]"
|
||||
if llm_on:
|
||||
try:
|
||||
from src.endpoint_resolver import resolve_endpoint
|
||||
from src.llm_core import llm_call_async
|
||||
from src.reminder_personas import synthesis_system_prompt
|
||||
url, model, headers = resolve_endpoint("utility", owner=owner or None)
|
||||
if not url:
|
||||
url, model, headers = resolve_endpoint("default", owner=owner or None)
|
||||
if url and model:
|
||||
persona_id = (settings.get("reminder_llm_persona") or "").strip()
|
||||
sys_prompt = synthesis_system_prompt(persona_id)
|
||||
raw = await llm_call_async(
|
||||
url=url, model=model,
|
||||
messages=[
|
||||
{"role": "system", "content": sys_prompt},
|
||||
{"role": "user", "content": f"Title: {title}\n\n{note_body}".strip()},
|
||||
],
|
||||
temperature=0.7, max_tokens=200, headers=headers, timeout=30,
|
||||
)
|
||||
from src.text_helpers import strip_think as _strip_think
|
||||
# prose=True strips untagged "The user wants me to…" chain-of-thought.
|
||||
# prompt_echo=True strips Qwen-style "Thinking Process:" / leaked
|
||||
# prompt prefixes. Both are safe here because this is a
|
||||
# one-sentence LLM-only output, not user-pasted content.
|
||||
synthesis = _strip_think(raw or "", prose=True, prompt_echo=True)
|
||||
# Reminder synthesis is supposed to be ONE sentence. Strip-think's
|
||||
# paragraph-based heuristic misses cases where the model puts
|
||||
# reasoning + answer on consecutive lines inside one paragraph
|
||||
# (e.g. "I should write... [\n] You have one task waiting...").
|
||||
# Walk lines, drop reasoning/prompt-echo lines, then keep the
|
||||
# last surviving line — that's the actual warm sentence.
|
||||
if synthesis:
|
||||
import re as _re
|
||||
# Tightened: target ACTUAL self-talk (model narrating what
|
||||
# it'll do) rather than any first-person sentence. The old
|
||||
# pattern killed legit warm sentences like "I'll see you
|
||||
# tomorrow" or "I should be done by then". New rules:
|
||||
# • "I (need|should|have|'ll|will) (write|draft|reply|…)"
|
||||
# only matches when followed by a TASK verb taking an
|
||||
# OBJECT (so first-person + intransitive verb passes).
|
||||
# • Self-instructional patterns the model emits verbatim:
|
||||
# "I should write something that reminds them…",
|
||||
# "I need to draft…", "Let me think…".
|
||||
# • Explicit instructions echoed back from the prompt:
|
||||
# "Keep it under 25 words", "No greetings".
|
||||
_reasoning = _re.compile(
|
||||
r"^\s*(?:"
|
||||
# "I should write/draft/compose…" with a task-object follow
|
||||
r"i (?:need|should|have|'ll|will|am going|am)\s+to\s+"
|
||||
r"(?:write|draft|compose|craft|generate|produce|create|"
|
||||
r"summarize|answer|provide|note|address|remind|output)"
|
||||
r"\s+(?:a |an |the |something|this|that|here|them|him|her|"
|
||||
r"you|user|reply|response|sentence|message|line|warm)|"
|
||||
# The model literally narrating about the user
|
||||
r"the user (?:wants|is asking|asks|needs|wrote|said|requested) (?:me )?(?:to|for|that|about|something)|"
|
||||
# "Let me [think/write/draft/…] (about/for/the …)"
|
||||
r"let me (?:think|write|draft|consider|note|see|check)\b\s+(?:about|for|the|this|that|if|whether)|"
|
||||
# "Looking at the/this/that …"
|
||||
r"looking at (?:the|this|that)\b|"
|
||||
# "Based on the/this/what …"
|
||||
r"based on (?:the|this|what|context|that)\b|"
|
||||
# Prompt-echo of length / style instructions
|
||||
r"keep it under \d+ words\b|"
|
||||
r"(?:no greetings|no preamble|no hashtags|just output the)\b"
|
||||
r").*",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
# Echo of the prompt's "Pending:" / "<N> pending" tail.
|
||||
_echo = _re.compile(
|
||||
r"^\s*(?:pending\s*[:.]|(?:\d+|one|two|three|four|five)\s+pending\b)",
|
||||
_re.IGNORECASE,
|
||||
)
|
||||
lines = [ln for ln in synthesis.splitlines() if ln.strip()]
|
||||
cleaned = [ln for ln in lines if not _reasoning.match(ln) and not _echo.match(ln)]
|
||||
if cleaned:
|
||||
# The model's actual answer is normally the LAST surviving
|
||||
# line — reasoning leads, answer trails.
|
||||
synthesis = cleaned[-1].strip()
|
||||
else:
|
||||
synthesis = _SYNTH_FAILED_TAG
|
||||
except Exception as e:
|
||||
logger.warning(f"Reminder LLM synthesis failed: {e}")
|
||||
synthesis = _SYNTH_FAILED_TAG
|
||||
if synthesis:
|
||||
_s = synthesis.strip(); _low = _s.lower()
|
||||
if (not _s or _low.startswith("error:") or _low.startswith("[error")
|
||||
or "operation failed" in _low
|
||||
or ("upstream" in _low and "failed" in _low)) and synthesis != _SYNTH_FAILED_TAG:
|
||||
logger.warning(f"Reminder synthesis looked like an error, replacing: {_s[:120]!r}")
|
||||
synthesis = _SYNTH_FAILED_TAG
|
||||
|
||||
email_sent = False
|
||||
email_error = ""
|
||||
if channel == "email":
|
||||
try:
|
||||
from routes.email_routes import _get_email_config
|
||||
from email.mime.text import MIMEText
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from datetime import datetime as _dt
|
||||
# `reminder_email_account_id` lets the user pick WHICH email
|
||||
# account to send reminders from (when they have several
|
||||
# configured in Integrations). Falls back to the default
|
||||
# account when no explicit choice is saved.
|
||||
_acc_id = (settings.get("reminder_email_account_id") or "").strip() or None
|
||||
cfg = _get_email_config(account_id=_acc_id, owner=owner or "")
|
||||
if not (cfg.get("smtp_host") and cfg.get("smtp_user") and cfg.get("smtp_password")):
|
||||
try:
|
||||
from core.database import SessionLocal as _SL, EmailAccount as _EA
|
||||
from sqlalchemy import and_, or_
|
||||
db = _SL()
|
||||
try:
|
||||
q = db.query(_EA).filter(_EA.enabled == True) # noqa: E712
|
||||
if owner:
|
||||
unowned = or_(_EA.owner == None, _EA.owner == "") # noqa: E711
|
||||
same_mailbox = or_(_EA.imap_user == owner, _EA.from_address == owner)
|
||||
q = q.filter(or_(_EA.owner == owner, and_(unowned, same_mailbox)))
|
||||
for row in q.order_by(_EA.is_default.desc(), _EA.created_at.asc()).all():
|
||||
trial = _get_email_config(account_id=row.id, owner=owner or "")
|
||||
if trial.get("smtp_host") and trial.get("smtp_user") and trial.get("smtp_password"):
|
||||
cfg = trial
|
||||
break
|
||||
finally:
|
||||
db.close()
|
||||
except Exception as _fallback_error:
|
||||
logger.debug(f"Reminder SMTP fallback lookup failed: {_fallback_error}")
|
||||
from_addr = (cfg.get("from_address") or cfg.get("smtp_user") or "").strip()
|
||||
recipient = (settings.get("reminder_email_to") or "").strip() or from_addr
|
||||
# Loud diagnostic so we can see WHY a reminder didn't send (the
|
||||
# previous "silently no-op when cfg has no smtp_host" was invisible).
|
||||
logger.info(
|
||||
"dispatch_reminder[email] note_id=%s owner=%r "
|
||||
"has_smtp_host=%s has_smtp_user=%s has_from=%s has_recipient=%s",
|
||||
note_id, owner,
|
||||
bool(cfg.get("smtp_host")), bool(cfg.get("smtp_user")),
|
||||
bool(from_addr), bool(recipient),
|
||||
)
|
||||
missing = []
|
||||
if not cfg.get("smtp_host"):
|
||||
missing.append("SMTP host")
|
||||
if not cfg.get("smtp_user"):
|
||||
missing.append("SMTP user")
|
||||
if not cfg.get("smtp_password"):
|
||||
missing.append("SMTP password")
|
||||
if not from_addr:
|
||||
missing.append("from address")
|
||||
if not recipient:
|
||||
missing.append("recipient")
|
||||
if missing:
|
||||
email_error = "Missing " + ", ".join(missing)
|
||||
logger.warning(
|
||||
"Reminder email not sent for note_id=%s account=%r: %s",
|
||||
note_id, cfg.get("account_name"), email_error,
|
||||
)
|
||||
else:
|
||||
msg = MIMEMultipart("alternative")
|
||||
msg["From"] = from_addr
|
||||
msg["To"] = recipient
|
||||
_t = title or 'Note'
|
||||
_t = _t[len('Reminder:'):].strip() if _t.lower().startswith('reminder:') else _t
|
||||
msg["Subject"] = f"Reminder (Odysseus): {_t}"
|
||||
msg["Date"] = _dt.utcnow().strftime("%a, %d %b %Y %H:%M:%S +0000")
|
||||
msg["X-Odysseus-Origin"] = "odysseus-ui"
|
||||
msg["X-Odysseus-Kind"] = "reminder"
|
||||
msg["X-Odysseus-Ref"] = str(note_id)
|
||||
# Body shape: synthesis (warm sentence) → blank line → bold
|
||||
# title header → note details. The title was previously only
|
||||
# in the subject line, so the email read like a faceless
|
||||
# to-do list with no anchor to which note triggered it.
|
||||
_body_chunks = []
|
||||
if synthesis:
|
||||
_body_chunks.append(synthesis)
|
||||
if _t:
|
||||
_body_chunks.append(_t)
|
||||
if note_body:
|
||||
_body_chunks.append(note_body)
|
||||
plain = "\n\n".join(_body_chunks) if _body_chunks else title
|
||||
msg.attach(MIMEText(plain, "plain", "utf-8"))
|
||||
|
||||
def _smtp_send():
|
||||
from routes.email_helpers import _send_smtp_message
|
||||
_send_smtp_message(cfg, from_addr, [recipient], msg.as_string())
|
||||
|
||||
import asyncio as _aio
|
||||
await _aio.to_thread(_smtp_send)
|
||||
email_sent = True
|
||||
except Exception as e:
|
||||
email_error = str(e) or e.__class__.__name__
|
||||
logger.warning(f"Reminder email send failed: {e}")
|
||||
|
||||
webhook_sent = False
|
||||
webhook_error = ""
|
||||
if channel == "webhook":
|
||||
try:
|
||||
import httpx
|
||||
import json as _wjson
|
||||
from src.integrations import load_integrations
|
||||
# Built-in payload defaults for known presets so users don't have
|
||||
# to configure a template just to use a standard service.
|
||||
_PRESET_TEMPLATE_DEFAULTS = {
|
||||
"discord_webhook": '{"embeds": [{"title": "{{title}}", "description": "{{message}}", "color": 5793266}]}',
|
||||
}
|
||||
intg_id = settings.get("reminder_webhook_integration_id", "").strip()
|
||||
template = settings.get("reminder_webhook_payload_template", "").strip()
|
||||
if not intg_id:
|
||||
webhook_error = "No webhook integration selected"
|
||||
else:
|
||||
intg = next(
|
||||
(i for i in load_integrations()
|
||||
if i.get("id") == intg_id and i.get("base_url")),
|
||||
None,
|
||||
)
|
||||
if not intg:
|
||||
webhook_error = f"Integration {intg_id!r} not found or missing base URL"
|
||||
else:
|
||||
# Fall back to a built-in default for known presets so
|
||||
# users don't have to configure a template for standard
|
||||
# services like Discord.
|
||||
if not template:
|
||||
template = _PRESET_TEMPLATE_DEFAULTS.get(intg.get("preset", ""), "")
|
||||
if not template:
|
||||
webhook_error = "No payload template configured"
|
||||
else:
|
||||
# Render template: JSON-escape the values so the result
|
||||
# is always valid JSON regardless of special characters.
|
||||
# dumps() returns `"value"` — strip outer quotes.
|
||||
msg = (synthesis or note_body or title or "Reminder")[:4000]
|
||||
_t = _wjson.dumps(title or "Reminder")[1:-1]
|
||||
_m = _wjson.dumps(msg)[1:-1]
|
||||
rendered = template.replace("{{title}}", _t).replace("{{message}}", _m)
|
||||
hdrs = {"Content-Type": "application/json"}
|
||||
api_key = intg.get("api_key", "")
|
||||
auth_type = (intg.get("auth_type") or "none").lower()
|
||||
if api_key:
|
||||
if auth_type == "bearer":
|
||||
hdrs["Authorization"] = f"Bearer {api_key}"
|
||||
elif auth_type == "header":
|
||||
hdrs[intg.get("auth_header") or "Authorization"] = api_key
|
||||
url = intg["base_url"].rstrip("/")
|
||||
# SSRF guard — matches the pattern used by webhook_routes,
|
||||
# CalDAV, search, and embeddings. Blocks link-local / metadata
|
||||
# addresses (169.254.x.x) by default; set
|
||||
# REMINDER_WEBHOOK_BLOCK_PRIVATE_IPS=true to also block
|
||||
# RFC-1918 ranges for locked-down deployments.
|
||||
import os as _os
|
||||
from src.url_safety import check_outbound_url as _chk
|
||||
_block = _os.getenv("REMINDER_WEBHOOK_BLOCK_PRIVATE_IPS", "false").lower() == "true"
|
||||
_ok, _reason = _chk(url, block_private=_block)
|
||||
if not _ok:
|
||||
webhook_error = f"Webhook URL rejected: {_reason}"
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
resp = await client.post(url, content=rendered.encode(), headers=hdrs)
|
||||
webhook_sent = resp.is_success
|
||||
if not webhook_sent:
|
||||
webhook_error = f"Webhook returned HTTP {resp.status_code}"
|
||||
except Exception as e:
|
||||
webhook_error = str(e) or e.__class__.__name__
|
||||
logger.warning(f"Reminder webhook send failed: {e}")
|
||||
|
||||
ntfy_sent = False
|
||||
ntfy_error = ""
|
||||
if channel == "ntfy":
|
||||
try:
|
||||
from src.integrations import load_integrations
|
||||
import httpx
|
||||
intg = next(
|
||||
(i for i in load_integrations()
|
||||
if i.get("preset") == "ntfy" and i.get("enabled", True) and i.get("base_url")),
|
||||
None,
|
||||
)
|
||||
if intg:
|
||||
base = intg["base_url"].rstrip("/")
|
||||
topic = settings.get("reminder_ntfy_topic") or "reminders"
|
||||
ntfy_body = synthesis or note_body or title
|
||||
hdrs = {"Title": title or "Reminder", "Priority": "high", "Tags": "bell"}
|
||||
api_key = intg.get("api_key", "")
|
||||
if api_key:
|
||||
hdrs["Authorization"] = f"Bearer {api_key}"
|
||||
# SSRF guard — same check (and env knob) as the webhook branch
|
||||
# above: link-local / metadata addresses are always rejected;
|
||||
# REMINDER_WEBHOOK_BLOCK_PRIVATE_IPS=true also blocks RFC-1918
|
||||
# so a ntfy base_url can't be pointed at internal services.
|
||||
import os as _os
|
||||
from src.url_safety import check_outbound_url as _chk
|
||||
_block = _os.getenv("REMINDER_WEBHOOK_BLOCK_PRIVATE_IPS", "false").lower() == "true"
|
||||
_ok, _reason = _chk(f"{base}/{topic}", block_private=_block)
|
||||
if not _ok:
|
||||
ntfy_error = f"ntfy URL rejected: {_reason}"
|
||||
else:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
resp = await client.post(f"{base}/{topic}", content=ntfy_body, headers=hdrs)
|
||||
ntfy_sent = resp.is_success
|
||||
if not ntfy_sent:
|
||||
ntfy_error = f"ntfy returned HTTP {resp.status_code}"
|
||||
else:
|
||||
ntfy_error = "No enabled ntfy integration"
|
||||
except Exception as e:
|
||||
ntfy_error = str(e) or e.__class__.__name__
|
||||
logger.warning(f"Reminder ntfy send failed: {e}")
|
||||
|
||||
# In-app browser notification ALWAYS fires (regardless of channel). The
|
||||
# frontend polls `/api/tasks/notifications` and turns any entry with a
|
||||
# `body` into a real `Notification(...)` — same surface as task-success
|
||||
# popups. Lets the user see reminders inside the app even when the
|
||||
# primary channel is email/ntfy and the tab is open.
|
||||
browser_sent = False
|
||||
local_browser_sent = (not queue_browser and channel == "browser")
|
||||
if queue_browser and _scheduler_ref is not None:
|
||||
try:
|
||||
_scheduler_ref.add_notification(
|
||||
task_name=title or "Reminder",
|
||||
status="success",
|
||||
task_id=f"reminder-{note_id}",
|
||||
owner=owner or None,
|
||||
body=(synthesis or note_body or title or "").strip()[:500] or "Reminder",
|
||||
)
|
||||
browser_sent = True
|
||||
except Exception as _e:
|
||||
logger.debug(f"dispatch_reminder: in-app notif push failed: {_e}")
|
||||
|
||||
# Dedupe across paths: write to the same cache file `action_ping_notes`
|
||||
# reads, so the background scanner's REPING_MIN window suppresses a
|
||||
# second send for the same note within 25 min. Without this, a note
|
||||
# whose due_date fires while the user has the app open got TWO emails
|
||||
# (frontend-fired here + background-fired by ping_notes 0–5 min later).
|
||||
if (email_sent or ntfy_sent or webhook_sent or browser_sent or local_browser_sent) and note_id:
|
||||
try:
|
||||
import json as _json
|
||||
from datetime import datetime as _dt, timezone as _tz
|
||||
from pathlib import Path as _P
|
||||
# Per-owner cache so the scanner's prune step on user A's run
|
||||
# doesn't drop user B's just-fired entry (review C4).
|
||||
_STATE = cache_path
|
||||
if _STATE is None:
|
||||
_slug = "".join(c if (c.isalnum() or c in "-_.@") else "_" for c in (owner or "default"))
|
||||
_STATE = _P(DATA_DIR) / f"note_pings_{_slug}.json"
|
||||
_STATE.parent.mkdir(parents=True, exist_ok=True)
|
||||
try:
|
||||
_cache = cache or (_json.loads(_STATE.read_text(encoding="utf-8")) if _STATE.exists() else {})
|
||||
except Exception:
|
||||
_cache = {}
|
||||
sent_channel = "email" if email_sent else "ntfy" if ntfy_sent else "webhook" if webhook_sent else "browser"
|
||||
_cache[cache_key or str(note_id)] = {
|
||||
"at": _dt.now(_tz.utc).isoformat(),
|
||||
"channel": sent_channel,
|
||||
}
|
||||
_STATE.write_text(_json.dumps(_cache), encoding="utf-8")
|
||||
except Exception as _e:
|
||||
logger.debug(f"dispatch_reminder: cache write failed: {_e}")
|
||||
|
||||
return {
|
||||
"channel": channel,
|
||||
"synthesis": synthesis,
|
||||
"email_sent": email_sent,
|
||||
"email_error": email_error,
|
||||
"ntfy_sent": ntfy_sent,
|
||||
"ntfy_error": ntfy_error,
|
||||
"webhook_sent": webhook_sent,
|
||||
"webhook_error": webhook_error,
|
||||
"browser_sent": browser_sent or local_browser_sent,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Router factory
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def setup_note_routes(task_scheduler=None):
|
||||
# Expose the scheduler to module-level `dispatch_reminder` so reminders
|
||||
# can also push to the in-app notification queue (the polling system
|
||||
# turns each entry into a real browser Notification + the existing
|
||||
# tasks-tab badge / dot system).
|
||||
global _scheduler_ref
|
||||
_scheduler_ref = task_scheduler
|
||||
|
||||
router = APIRouter(prefix="/api/notes", tags=["notes"])
|
||||
|
||||
def _owner(request: Request) -> Optional[str]:
|
||||
# require_user, not bare get_current_user: a request that reaches
|
||||
# these owner-scoped routes with NO identity (auth-middleware
|
||||
# regression, SSRF from a sibling service) must fail closed (401)
|
||||
# when auth is configured — not be treated as the single-user mode
|
||||
# and handed blanket access to every account's notes. The documented
|
||||
# anonymous modes (AUTH_ENABLED=false, LOCALHOST_BYPASS on loopback,
|
||||
# unconfigured first-run) still resolve to None, the single-user
|
||||
# path. fire_reminder below already gated this way; the CRUD routes
|
||||
# did not.
|
||||
return require_user(request) or None
|
||||
|
||||
def _is_admin_or_single_user(request: Request, user: str | None) -> bool:
|
||||
if user == INTERNAL_TOOL_USER:
|
||||
return True
|
||||
if not user:
|
||||
# require_user() already admitted this request, which only happens
|
||||
# for auth-disabled, loopback-bypass, or unconfigured single-user
|
||||
# modes. There is no separate non-admin account boundary there.
|
||||
return True
|
||||
try:
|
||||
from core.auth import AuthManager
|
||||
auth_mgr = getattr(request.app.state, "auth_manager", None) or AuthManager()
|
||||
if not getattr(auth_mgr, "is_configured", True):
|
||||
return True
|
||||
return bool(auth_mgr.is_admin(user))
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
# --- LIST ---
|
||||
@router.get("")
|
||||
def list_notes(
|
||||
request: Request,
|
||||
archived: Optional[bool] = None,
|
||||
label: Optional[str] = None,
|
||||
):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(Note)
|
||||
if user is not None:
|
||||
q = q.filter(Note.owner == user)
|
||||
if archived is not None:
|
||||
q = q.filter(Note.archived == archived)
|
||||
else:
|
||||
q = q.filter(Note.archived == False)
|
||||
if label:
|
||||
q = q.filter(Note.label == label)
|
||||
# Archived view: most recently archived first. Active view: pin + manual order.
|
||||
if archived is True:
|
||||
notes = q.order_by(Note.updated_at.desc()).all()
|
||||
else:
|
||||
notes = q.order_by(Note.pinned.desc(), Note.sort_order.asc(), Note.updated_at.desc()).all()
|
||||
return {"notes": [_note_to_dict(n) for n in notes]}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- CREATE ---
|
||||
@router.post("")
|
||||
def create_note(request: Request, body: NoteCreate):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = Note(
|
||||
id=str(uuid.uuid4()),
|
||||
owner=user,
|
||||
title=body.title,
|
||||
content=body.content,
|
||||
items=json.dumps(body.items) if body.items is not None else None,
|
||||
note_type=body.note_type,
|
||||
color=body.color,
|
||||
label=body.label,
|
||||
pinned=body.pinned,
|
||||
due_date=body.due_date,
|
||||
source=body.source,
|
||||
session_id=body.session_id,
|
||||
image_url=body.image_url,
|
||||
repeat=body.repeat or "none",
|
||||
sort_order=body.sort_order if body.sort_order is not None else 0,
|
||||
)
|
||||
db.add(note)
|
||||
db.commit()
|
||||
db.refresh(note)
|
||||
return _note_to_dict(note)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- GET ONE ---
|
||||
@router.get("/{note_id}")
|
||||
def get_note(request: Request, note_id: str):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
return _note_to_dict(note)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- UPDATE ---
|
||||
@router.put("/{note_id}")
|
||||
def update_note(request: Request, note_id: str, body: NoteUpdate):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
|
||||
if body.title is not None:
|
||||
note.title = body.title
|
||||
if body.content is not None:
|
||||
note.content = body.content
|
||||
if body.items is not None:
|
||||
note.items = json.dumps(body.items)
|
||||
flag_modified(note, "items")
|
||||
if body.note_type is not None:
|
||||
note.note_type = body.note_type
|
||||
if body.color is not None:
|
||||
note.color = body.color
|
||||
if body.label is not None:
|
||||
note.label = body.label
|
||||
if body.pinned is not None:
|
||||
note.pinned = body.pinned
|
||||
if body.archived is not None:
|
||||
note.archived = body.archived
|
||||
if body.due_date is not None:
|
||||
note.due_date = body.due_date
|
||||
if body.image_url is not None:
|
||||
note.image_url = body.image_url
|
||||
if body.repeat is not None:
|
||||
note.repeat = body.repeat
|
||||
if body.sort_order is not None:
|
||||
note.sort_order = body.sort_order
|
||||
if body.agent_session_id is not None:
|
||||
note.agent_session_id = body.agent_session_id
|
||||
|
||||
db.commit()
|
||||
db.refresh(note)
|
||||
return _note_to_dict(note)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- DELETE ---
|
||||
@router.delete("/{note_id}")
|
||||
def delete_note(request: Request, note_id: str):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
db.delete(note)
|
||||
db.commit()
|
||||
return {"ok": True}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- TOGGLE PIN ---
|
||||
@router.post("/{note_id}/pin")
|
||||
def toggle_pin(request: Request, note_id: str):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
note.pinned = not note.pinned
|
||||
db.commit()
|
||||
return {"ok": True, "pinned": note.pinned}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- TOGGLE ARCHIVE ---
|
||||
@router.post("/{note_id}/archive")
|
||||
def toggle_archive(request: Request, note_id: str):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
note.archived = not note.archived
|
||||
db.commit()
|
||||
return {"ok": True, "archived": note.archived}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- TOGGLE CHECKLIST ITEM ---
|
||||
@router.post("/{note_id}/items/{index}/toggle")
|
||||
def toggle_item(request: Request, note_id: str, index: int):
|
||||
user = _owner(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
# SECURITY: strict ownership — previously `note.owner and note.owner != user`
|
||||
# let any user touch a row whose owner field was null/empty.
|
||||
if user is not None and note.owner != user:
|
||||
raise HTTPException(404, "Note not found")
|
||||
if not note.items:
|
||||
raise HTTPException(400, "Note has no checklist items")
|
||||
items = json.loads(note.items)
|
||||
if index < 0 or index >= len(items):
|
||||
raise HTTPException(400, f"Item index {index} out of range")
|
||||
items[index]["done"] = not items[index].get("done", False)
|
||||
note.items = json.dumps(items)
|
||||
flag_modified(note, "items")
|
||||
db.commit()
|
||||
return {"ok": True, "items": items}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# --- FIRE REMINDER ---
|
||||
@router.post("/fire-reminder")
|
||||
async def fire_reminder(request: Request):
|
||||
"""Dispatch a reminder according to user settings.
|
||||
|
||||
Called by the frontend when a reminder fires. Optionally generates an
|
||||
LLM synthesis line and/or sends an email through configured SMTP.
|
||||
Returns {synthesis, email_sent}.
|
||||
"""
|
||||
# Gate against anonymous callers — LLM synthesis can burn tokens.
|
||||
user = require_user(request)
|
||||
body = await request.json()
|
||||
note_id = str(body.get("note_id") or "").strip()
|
||||
if not note_id:
|
||||
raise HTTPException(400, "note_id required")
|
||||
|
||||
caller = _owner(request)
|
||||
is_test = note_id.startswith("test-")
|
||||
is_admin = _is_admin_or_single_user(request, user or caller)
|
||||
_override: dict = {}
|
||||
if is_test:
|
||||
if not is_admin:
|
||||
raise HTTPException(403, "Admin only")
|
||||
title = (body.get("title") or "Test Reminder").strip() or "Test Reminder"
|
||||
note_body = (body.get("body") or "").strip()
|
||||
# Optional overrides let the admin settings test button pass the
|
||||
# current UI values directly so it never races a pending save.
|
||||
if body.get("channel"):
|
||||
_override["reminder_channel"] = body["channel"]
|
||||
if body.get("webhook_integration_id"):
|
||||
_override["reminder_webhook_integration_id"] = body["webhook_integration_id"]
|
||||
if body.get("webhook_payload_template"):
|
||||
_override["reminder_webhook_payload_template"] = body["webhook_payload_template"]
|
||||
# Mirror the in-UI AI Synthesis toggle + persona so the test
|
||||
# actually exercises the synthesis path before/without a Save.
|
||||
if "llm_synthesis" in body:
|
||||
_override["reminder_llm_synthesis"] = bool(body["llm_synthesis"])
|
||||
if "llm_persona" in body:
|
||||
_override["reminder_llm_persona"] = str(body["llm_persona"] or "")
|
||||
else:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
note = db.query(Note).filter(Note.id == note_id).first()
|
||||
if not note:
|
||||
raise HTTPException(404, "Note not found")
|
||||
if caller is not None and note.owner != caller:
|
||||
raise HTTPException(404, "Note not found")
|
||||
title, note_body = _reminder_text_from_note(note)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return await dispatch_reminder(
|
||||
title=title, note_body=note_body, note_id=note_id,
|
||||
owner=caller or "",
|
||||
queue_browser=False,
|
||||
settings_override=_override or None,
|
||||
)
|
||||
|
||||
# --- REORDER NOTES ---
|
||||
@router.post("/reorder")
|
||||
async def reorder_notes(request: Request):
|
||||
"""Update sort_order for a list of note IDs in the order provided."""
|
||||
user = _owner(request)
|
||||
body = await request.json()
|
||||
ids = body.get("ids", [])
|
||||
if not isinstance(ids, list):
|
||||
raise HTTPException(400, "ids must be a list")
|
||||
# v2 review HIGH-12: drop the legacy `(owner == user) | (owner ==
|
||||
# None)` OR which let an authenticated user silently reorder
|
||||
# every legacy-null-owner note belonging to other accounts. In
|
||||
# an unconfigured (single-user) auth deploy the OR is still safe
|
||||
# because there's no second user to attack; we keep that branch
|
||||
# explicit and gated on AuthManager.is_configured.
|
||||
try:
|
||||
from core.auth import AuthManager
|
||||
_allow_null = not AuthManager().is_configured
|
||||
except Exception:
|
||||
_allow_null = False
|
||||
db = SessionLocal()
|
||||
try:
|
||||
for i, nid in enumerate(ids):
|
||||
q = db.query(Note).filter(Note.id == nid)
|
||||
if user is not None:
|
||||
if _allow_null:
|
||||
q = q.filter((Note.owner == user) | (Note.owner == None)) # noqa: E711
|
||||
else:
|
||||
q = q.filter(Note.owner == user)
|
||||
note = q.first()
|
||||
if note:
|
||||
note.sort_order = i
|
||||
db.commit()
|
||||
return {"ok": True, "count": len(ids)}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return router
|
||||
_sys.modules[__name__] = _canonical
|
||||
|
||||
@@ -12,7 +12,9 @@ from core.models import ChatMessage
|
||||
from src.request_models import SessionResponse
|
||||
from core.database import Session as DbSession, SessionLocal, Document, GalleryImage, utcnow_naive
|
||||
from src.auth_helpers import effective_user, _auth_disabled, owner_filter
|
||||
from src.session_image_cleanup import _generated_image_path_for_cleanup, session_image_refs
|
||||
from src.session_actions import is_session_recently_active
|
||||
from src.upload_handler import reserve_message_upload_references
|
||||
|
||||
|
||||
def _sanitize_export_filename(name: str) -> str:
|
||||
@@ -203,7 +205,12 @@ def _pick_endpoint_for_sort(owner=None):
|
||||
return url, model, headers
|
||||
return None, None, None
|
||||
|
||||
def setup_session_routes(session_manager: SessionManager, config: dict, webhook_manager=None):
|
||||
def setup_session_routes(
|
||||
session_manager: SessionManager,
|
||||
config: dict,
|
||||
webhook_manager=None,
|
||||
upload_handler=None,
|
||||
):
|
||||
"""Setup session routes with the provided manager and config"""
|
||||
|
||||
REQUEST_TIMEOUT = config.get("REQUEST_TIMEOUT", 20)
|
||||
@@ -214,6 +221,7 @@ def setup_session_routes(session_manager: SessionManager, config: dict, webhook_
|
||||
@router.get("/sessions")
|
||||
def list_sessions(request: Request):
|
||||
user = effective_user(request)
|
||||
active_incognito_id = str(request.query_params.get("active_incognito_id") or "").strip()
|
||||
# Lazy purge: incognito sessions are ephemeral by design — wipe leftovers
|
||||
# from the DB and session_manager so they vanish on the next page refresh.
|
||||
# BUT: skip sessions that were created within the last 10 minutes.
|
||||
@@ -234,6 +242,8 @@ def setup_session_routes(session_manager: SessionManager, config: dict, webhook_
|
||||
DbSession.created_at < _cutoff,
|
||||
).all()
|
||||
for _g in _ghosts:
|
||||
if active_incognito_id and _g.id == active_incognito_id:
|
||||
continue
|
||||
_purge_db.query(_DbMsg).filter(_DbMsg.session_id == _g.id).delete()
|
||||
_purge_db.delete(_g)
|
||||
if hasattr(session_manager, "delete_session"):
|
||||
@@ -537,6 +547,22 @@ def setup_session_routes(session_manager: SessionManager, config: dict, webhook_
|
||||
body = await request.json()
|
||||
messages = body.get("messages", [])
|
||||
from core.models import ChatMessage
|
||||
owner = effective_user(request)
|
||||
try:
|
||||
for message in messages:
|
||||
missing_id = reserve_message_upload_references(
|
||||
upload_handler,
|
||||
owner,
|
||||
message.get("content"),
|
||||
message.get("metadata"),
|
||||
)
|
||||
if missing_id:
|
||||
raise HTTPException(
|
||||
409,
|
||||
f"Referenced upload is no longer available: {missing_id}",
|
||||
)
|
||||
except (AttributeError, TypeError, ValueError) as exc:
|
||||
raise HTTPException(400, "Invalid message attachment metadata") from exc
|
||||
for m in messages:
|
||||
sess.add_message(ChatMessage(m["role"], m["content"], metadata=m.get("metadata")))
|
||||
session_manager.save_sessions()
|
||||
@@ -619,13 +645,43 @@ def setup_session_routes(session_manager: SessionManager, config: dict, webhook_
|
||||
db = SessionLocal()
|
||||
try:
|
||||
from core.database import ChatMessage as DbChatMessage
|
||||
session_ids = [row[0] for row in db.query(DbSession.id).all()]
|
||||
count = db.query(DbSession).count()
|
||||
image_ids: set[str] = set()
|
||||
filenames: set[str] = set()
|
||||
for sid in session_ids:
|
||||
ids, names = session_image_refs(db, sid)
|
||||
image_ids.update(ids)
|
||||
filenames.update(names)
|
||||
image_query = db.query(GalleryImage).filter(GalleryImage.session_id.in_(session_ids)) if session_ids else db.query(GalleryImage).filter(False)
|
||||
if image_ids or filenames:
|
||||
from sqlalchemy import or_
|
||||
clauses = []
|
||||
if session_ids:
|
||||
clauses.append(GalleryImage.session_id.in_(session_ids))
|
||||
if image_ids:
|
||||
clauses.append(GalleryImage.id.in_(list(image_ids)))
|
||||
if filenames:
|
||||
clauses.append(GalleryImage.filename.in_(list(filenames)))
|
||||
image_query = db.query(GalleryImage).filter(or_(*clauses))
|
||||
images = image_query.all()
|
||||
removed_images = 0
|
||||
for img in images:
|
||||
img.is_active = False
|
||||
if img.filename:
|
||||
path = _generated_image_path_for_cleanup(img.filename)
|
||||
if path and path.exists():
|
||||
try:
|
||||
path.unlink()
|
||||
except Exception as exc:
|
||||
logger.warning("Could not remove generated image %s during all-session delete: %s", img.filename, exc)
|
||||
removed_images += 1
|
||||
db.query(DbChatMessage).delete()
|
||||
db.query(DbSession).delete()
|
||||
db.commit()
|
||||
session_manager.sessions.clear()
|
||||
logger.info(f"Admin deleted all {count} sessions")
|
||||
return {"status": "deleted", "count": count}
|
||||
logger.info(f"Admin deleted all {count} sessions and {removed_images} linked images")
|
||||
return {"status": "deleted", "count": count, "images_deleted": removed_images}
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.error(f"Error deleting all sessions: {e}")
|
||||
|
||||
+197
-4
@@ -157,6 +157,7 @@ def _package_installed_from_probe(name: str, probe: dict) -> bool:
|
||||
binaries = probe.get("binaries") if isinstance(probe.get("binaries"), dict) else {}
|
||||
dists = probe.get("dists") if isinstance(probe.get("dists"), dict) else {}
|
||||
modules = probe.get("modules") if isinstance(probe.get("modules"), dict) else {}
|
||||
files = probe.get("files") if isinstance(probe.get("files"), dict) else {}
|
||||
|
||||
if name == "vllm":
|
||||
return bool(binaries.get("vllm"))
|
||||
@@ -166,11 +167,43 @@ def _package_installed_from_probe(name: str, probe: dict) -> bool:
|
||||
return bool(dists.get("sglang") or modules.get("sglang", {}).get("real_module"))
|
||||
if name == "mlx_lm":
|
||||
return bool(dists.get("mlx-lm") or modules.get("mlx_lm", {}).get("real_module"))
|
||||
if name == "mflux":
|
||||
return bool(
|
||||
dists.get("mflux")
|
||||
or modules.get("mflux", {}).get("real_module")
|
||||
or binaries.get("mflux-generate-qwen")
|
||||
or binaries.get("mflux-generate")
|
||||
)
|
||||
if name == "boogu_image_mlx":
|
||||
return bool(
|
||||
dists.get("boogu-image-mlx")
|
||||
or modules.get("boogu_image_mlx", {}).get("real_module")
|
||||
)
|
||||
if name == "mlx_lama_swift":
|
||||
return bool(
|
||||
(binaries.get("odysseus-mlx-inpaint") or binaries.get("mlx-lama-serve"))
|
||||
and (files.get("mlx.metallib") or files.get("default.metallib"))
|
||||
)
|
||||
if name == "mlx_ddcolor_swift":
|
||||
return bool(
|
||||
(binaries.get("odysseus-mlx-colorize") or binaries.get("mlx-ddcolor-serve"))
|
||||
and (files.get("mlx.metallib") or files.get("default.metallib"))
|
||||
)
|
||||
if name == "diffusers":
|
||||
return bool(
|
||||
(dists.get("diffusers") or modules.get("diffusers", {}).get("real_module"))
|
||||
and (dists.get("torch") or modules.get("torch", {}).get("real_module"))
|
||||
)
|
||||
if name == "krea_diffusers":
|
||||
return bool(
|
||||
(dists.get("diffusers") or modules.get("diffusers", {}).get("real_module"))
|
||||
and (dists.get("torch") or modules.get("torch", {}).get("real_module"))
|
||||
)
|
||||
if name == "sam_mask":
|
||||
return bool(
|
||||
(dists.get("transformers") or modules.get("transformers", {}).get("real_module"))
|
||||
and (dists.get("torch") or modules.get("torch", {}).get("real_module"))
|
||||
)
|
||||
if name == "hf_transfer":
|
||||
return bool(
|
||||
dists.get("hf-transfer")
|
||||
@@ -183,6 +216,7 @@ def _package_status_note(name: str, probe: dict) -> str:
|
||||
binaries = probe.get("binaries") if isinstance(probe.get("binaries"), dict) else {}
|
||||
modules = probe.get("modules") if isinstance(probe.get("modules"), dict) else {}
|
||||
dists = probe.get("dists") if isinstance(probe.get("dists"), dict) else {}
|
||||
files = probe.get("files") if isinstance(probe.get("files"), dict) else {}
|
||||
module = modules.get(name) if isinstance(modules.get(name), dict) else {}
|
||||
locations = module.get("locations") or []
|
||||
if name == "vllm":
|
||||
@@ -212,10 +246,53 @@ def _package_status_note(name: str, probe: dict) -> str:
|
||||
if _package_installed_from_probe(name, probe):
|
||||
return f"diffusers {dists.get('diffusers', 'available')} with torch {dists.get('torch', 'available')}"
|
||||
return "Diffusers serving needs both diffusers and torch."
|
||||
if name == "krea_diffusers":
|
||||
if _package_installed_from_probe(name, probe):
|
||||
return f"Latest Diffusers runtime: diffusers {dists.get('diffusers', 'available')} with torch {dists.get('torch', 'available')}. Use Update/Reinstall to pull latest Diffusers from Git."
|
||||
return "Some newer image models need torch plus latest Diffusers from Git."
|
||||
if name == "sam_mask":
|
||||
if _package_installed_from_probe(name, probe):
|
||||
return f"SAM object masks: transformers {dists.get('transformers', 'available')} with torch {dists.get('torch', 'available')}"
|
||||
return "SAM click/object mask selection needs transformers and torch."
|
||||
if name == "mlx_lm":
|
||||
if _package_installed_from_probe(name, probe):
|
||||
return f"MLX LM {dists.get('mlx-lm', 'available')}"
|
||||
return "MLX serving needs mlx-lm on an Apple Silicon Mac."
|
||||
if name == "mflux":
|
||||
if _package_installed_from_probe(name, probe):
|
||||
parts = []
|
||||
if dists.get("mflux"):
|
||||
parts.append(f"mflux {dists['mflux']}")
|
||||
if binaries.get("mflux-generate-qwen"):
|
||||
parts.append(f"Qwen CLI: {binaries['mflux-generate-qwen']}")
|
||||
if binaries.get("mflux-generate"):
|
||||
parts.append(f"Flux CLI: {binaries['mflux-generate']}")
|
||||
return "; ".join(parts) if parts else "mflux available"
|
||||
return "MLX image serving needs mflux on an Apple Silicon Mac."
|
||||
if name == "boogu_image_mlx":
|
||||
if _package_installed_from_probe(name, probe):
|
||||
return f"Boogu MLX pipeline {dists.get('boogu-image-mlx', 'available')}"
|
||||
return "Boogu image models need boogu-image-mlx on an Apple Silicon Mac."
|
||||
if name == "mlx_lama_swift":
|
||||
if _package_installed_from_probe(name, probe):
|
||||
found = [
|
||||
binaries.get("odysseus-mlx-inpaint"),
|
||||
binaries.get("mlx-lama-serve"),
|
||||
]
|
||||
return f"LaMa/MI-GAN Swift MLX runner: {next((p for p in found if p), 'available')}"
|
||||
if binaries.get("odysseus-mlx-inpaint") or binaries.get("mlx-lama-serve"):
|
||||
return "LaMa/MI-GAN Swift runner is installed, but mlx.metallib is missing next to the runner."
|
||||
return "LaMa/MI-GAN inpainting models need an Odysseus-compatible mlx-lama-swift bridge on an Apple Silicon Mac."
|
||||
if name == "mlx_ddcolor_swift":
|
||||
if _package_installed_from_probe(name, probe):
|
||||
found = [
|
||||
binaries.get("odysseus-mlx-colorize"),
|
||||
binaries.get("mlx-ddcolor-serve"),
|
||||
]
|
||||
return f"DDColor Swift MLX runner: {next((p for p in found if p), 'available')}"
|
||||
if binaries.get("odysseus-mlx-colorize") or binaries.get("mlx-ddcolor-serve"):
|
||||
return "DDColor Swift runner is installed, but mlx.metallib is missing next to the runner."
|
||||
return "DDColor colorization models need an Odysseus-compatible mlx-ddcolor-swift bridge on an Apple Silicon Mac."
|
||||
if name in dists:
|
||||
return f"{name} {dists[name]}"
|
||||
return ""
|
||||
@@ -314,12 +391,22 @@ dist_names={{
|
||||
'llama_cpp':['llama-cpp-python'],
|
||||
'sglang':['sglang'],
|
||||
'mlx_lm':['mlx-lm'],
|
||||
'mlx_vlm':['mlx-vlm'],
|
||||
'mflux':['mflux'],
|
||||
'boogu_image_mlx':['boogu-image-mlx'],
|
||||
'mlx_lama_swift':[],
|
||||
'mlx_ddcolor_swift':[],
|
||||
'diffusers':['diffusers','torch'],
|
||||
'krea_diffusers':['diffusers','torch'],
|
||||
'sam_mask':['transformers','torch'],
|
||||
'hf_transfer':['hf-transfer','hf_transfer'],
|
||||
}}
|
||||
bin_names={{
|
||||
'vllm':['vllm'],
|
||||
'llama_cpp':['llama-server'],
|
||||
'mflux':['mflux-generate-qwen', 'mflux-generate'],
|
||||
'mlx_lama_swift':['odysseus-mlx-inpaint', 'mlx-lama-serve'],
|
||||
'mlx_ddcolor_swift':['odysseus-mlx-colorize', 'mlx-ddcolor-serve'],
|
||||
'tmux':['tmux'],
|
||||
}}
|
||||
|
||||
@@ -372,7 +459,19 @@ def probe(n):
|
||||
mods['torch'] = mod_status('torch')
|
||||
dists = dist_status(dist_names.get(n, [n]))
|
||||
bins = {{b: shutil.which(b) for b in bin_names.get(n, [])}}
|
||||
return {{'modules': mods, 'dists': dists, 'binaries': bins}}
|
||||
files = {{}}
|
||||
if n in ('mlx_lama_swift', 'mlx_ddcolor_swift'):
|
||||
for key in ('mlx.metallib', 'default.metallib'):
|
||||
found = None
|
||||
for b in bins.values():
|
||||
if not b:
|
||||
continue
|
||||
p = os.path.join(os.path.dirname(b), key)
|
||||
if os.path.exists(p):
|
||||
found = p
|
||||
break
|
||||
files[key] = found
|
||||
return {{'modules': mods, 'dists': dists, 'binaries': bins, 'files': files}}
|
||||
|
||||
print(json.dumps({{n: probe(n) for n in names}}))
|
||||
"""
|
||||
@@ -1088,6 +1187,8 @@ def setup_shell_routes() -> APIRouter:
|
||||
ssh_port: str | None = None,
|
||||
venv: str | None = None,
|
||||
backend: str | None = None,
|
||||
platform: str | None = None,
|
||||
model_hint: str | None = None,
|
||||
):
|
||||
"""Check which optional packages are installed.
|
||||
|
||||
@@ -1104,6 +1205,14 @@ def setup_shell_routes() -> APIRouter:
|
||||
import site
|
||||
import sys
|
||||
|
||||
platform_l = (platform or "").strip().lower()
|
||||
model_hint_l = (model_hint or "").strip().lower()
|
||||
has_krea_model = "krea" in model_hint_l
|
||||
has_lama_mlx_model = any(
|
||||
key in model_hint_l
|
||||
for key in ("lama", "mi-gan", "migan", "inpainting-mlx")
|
||||
)
|
||||
has_ddcolor_mlx_model = "ddcolor" in model_hint_l
|
||||
_prepend_user_install_bins_to_path()
|
||||
importlib.invalidate_caches()
|
||||
try:
|
||||
@@ -1158,7 +1267,7 @@ def setup_shell_routes() -> APIRouter:
|
||||
"name": "hf_transfer",
|
||||
"pip": "hf_transfer",
|
||||
"desc": "Fast model downloads from HuggingFace",
|
||||
"category": "LLM",
|
||||
"category": "Tools",
|
||||
"target": "remote",
|
||||
},
|
||||
{
|
||||
@@ -1210,8 +1319,52 @@ def setup_shell_routes() -> APIRouter:
|
||||
# ── Image ── editor + diffusion model serving
|
||||
{
|
||||
"name": "diffusers",
|
||||
"pip": "diffusers[torch]",
|
||||
"desc": "Image generation/editing pipelines (SD, Flux) with PyTorch",
|
||||
"pip": "diffusers[torch] torchvision accelerate scipy python-multipart",
|
||||
"desc": "Image generation/editing pipelines with PyTorch and Diffusers",
|
||||
"category": "Image",
|
||||
"target": "remote",
|
||||
},
|
||||
{
|
||||
"name": "krea_diffusers",
|
||||
"pip": "git+https://github.com/huggingface/diffusers.git torchvision accelerate scipy python-multipart",
|
||||
"desc": "Latest Diffusers from Git for newly released image pipelines",
|
||||
"category": "Image",
|
||||
"target": "remote",
|
||||
},
|
||||
{
|
||||
"name": "mflux",
|
||||
"pip": "mflux",
|
||||
"desc": "MLX image generation runtime for Apple Silicon models like Qwen Image",
|
||||
"category": "Image",
|
||||
"target": "remote",
|
||||
},
|
||||
{
|
||||
"name": "boogu_image_mlx",
|
||||
"pip": "git+https://github.com/xocialize/boogu-image-mlx.git",
|
||||
"desc": "MLX image generation pipeline for Boogu Image models on Apple Silicon",
|
||||
"category": "Image",
|
||||
"target": "remote",
|
||||
},
|
||||
{
|
||||
"name": "mlx_lama_swift",
|
||||
"pip": "",
|
||||
"desc": "Swift MLX runtime for LaMa / MI-GAN inpainting and object removal",
|
||||
"category": "Image",
|
||||
"target": "remote",
|
||||
"install_hint": "Build an Odysseus-compatible mlx-lama-swift bridge on the selected Apple Silicon Mac and put odysseus-mlx-inpaint or mlx-lama-serve on PATH. Upstream currently ships Swift libraries plus smoke executables, not a stable image-edit CLI.",
|
||||
},
|
||||
{
|
||||
"name": "mlx_ddcolor_swift",
|
||||
"pip": "",
|
||||
"desc": "Swift MLX runtime for DDColor automatic image colorization",
|
||||
"category": "Image",
|
||||
"target": "remote",
|
||||
"install_hint": "Build an Odysseus-compatible mlx-ddcolor-swift bridge on the selected Apple Silicon Mac and put odysseus-mlx-colorize or mlx-ddcolor-serve on PATH. Upstream currently ships Swift libraries plus smoke executables, not a stable colorize CLI.",
|
||||
},
|
||||
{
|
||||
"name": "mlx_vlm",
|
||||
"pip": "mlx-vlm",
|
||||
"desc": "MLX-VLM backbone used by HiDream image models on Apple Silicon",
|
||||
"category": "Image",
|
||||
"target": "remote",
|
||||
},
|
||||
@@ -1222,6 +1375,13 @@ def setup_shell_routes() -> APIRouter:
|
||||
"category": "Image",
|
||||
"target": "remote",
|
||||
},
|
||||
{
|
||||
"name": "sam_mask",
|
||||
"pip": "torch torchvision transformers accelerate pillow",
|
||||
"desc": "Neutral click/box/object segmentation masks for the image editor",
|
||||
"category": "Image",
|
||||
"target": "local",
|
||||
},
|
||||
{
|
||||
"name": "rembg",
|
||||
"pip": "rembg[gpu]",
|
||||
@@ -1251,6 +1411,21 @@ def setup_shell_routes() -> APIRouter:
|
||||
for pkg in packages:
|
||||
pkg.setdefault("install_cmd", None)
|
||||
pkg.setdefault("update_cmd", None)
|
||||
if not has_krea_model:
|
||||
packages = [
|
||||
p for p in packages
|
||||
if p.get("name") not in {"krea_diffusers", "transformers"}
|
||||
]
|
||||
if not has_lama_mlx_model:
|
||||
packages = [
|
||||
p for p in packages
|
||||
if p.get("name") != "mlx_lama_swift"
|
||||
]
|
||||
if not has_ddcolor_mlx_model:
|
||||
packages = [
|
||||
p for p in packages
|
||||
if p.get("name") != "mlx_ddcolor_swift"
|
||||
]
|
||||
# Remote check: for remote-target packages, probe the selected server's
|
||||
# venv over SSH so a remote `pip install` actually reflects here.
|
||||
remote_status: dict = {}
|
||||
@@ -1381,8 +1556,22 @@ def setup_shell_routes() -> APIRouter:
|
||||
target_os_id = ""
|
||||
if sys.platform == "darwin":
|
||||
target_os_id = "macos"
|
||||
if not target_os_id and platform_l in {"darwin", "macos", "mac"}:
|
||||
target_os_id = "macos"
|
||||
|
||||
for pkg in packages:
|
||||
if pkg.get("name") in {"mflux", "boogu_image_mlx", "mlx_vlm", "mlx_lama_swift", "mlx_ddcolor_swift"}:
|
||||
is_apple_target = target_os_id == "macos" or (
|
||||
not host and IS_APPLE_SILICON
|
||||
)
|
||||
known_non_apple_target = bool(target_os_id and target_os_id != "macos") or (
|
||||
not host and not IS_APPLE_SILICON
|
||||
)
|
||||
pkg["applicable"] = is_apple_target
|
||||
if known_non_apple_target:
|
||||
pkg["installed"] = None
|
||||
pkg["status_note"] = "Only relevant for Apple Silicon / MLX image serving."
|
||||
continue
|
||||
on_remote = bool(host and pkg.get("target") == "remote")
|
||||
probe = None
|
||||
if on_remote:
|
||||
@@ -1588,6 +1777,10 @@ def setup_shell_routes() -> APIRouter:
|
||||
"sglang[all]",
|
||||
"diffusers",
|
||||
"diffusers[torch]",
|
||||
"git+https://github.com/huggingface/diffusers.git",
|
||||
"mflux",
|
||||
"git+https://github.com/xocialize/boogu-image-mlx.git",
|
||||
"mlx-vlm",
|
||||
"transformers",
|
||||
"TTS",
|
||||
"bark",
|
||||
|
||||
+145
-3
@@ -10,16 +10,140 @@ from fastapi import APIRouter, Request, File, UploadFile, HTTPException, Form
|
||||
from typing import List, Optional
|
||||
import logging
|
||||
from core.middleware import require_admin
|
||||
from core.database import SessionLocal, GalleryImage, Session as DbSession
|
||||
from core.database import (
|
||||
SessionLocal,
|
||||
ChatMessage as DbChatMessage,
|
||||
CalendarCal,
|
||||
CalendarEvent,
|
||||
Document,
|
||||
DocumentVersion,
|
||||
GalleryImage,
|
||||
Note,
|
||||
Session as DbSession,
|
||||
)
|
||||
from src.auth_helpers import effective_user
|
||||
from src.attachment_refs import attachment_refs_from_metadata
|
||||
from src.constants import GENERATED_IMAGES_DIR
|
||||
from src.upload_handler import count_recent_uploads
|
||||
from src.upload_handler import (
|
||||
UploadCleanupSafetyError,
|
||||
count_recent_uploads,
|
||||
extract_upload_ids,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/upload", tags=["upload"])
|
||||
UPLOAD_RESPONSE_HEADERS = {"X-Content-Type-Options": "nosniff"}
|
||||
|
||||
def _upload_ids_from_persisted_text(value: object) -> set[str]:
|
||||
"""Return canonical upload IDs embedded in persisted text.
|
||||
|
||||
This covers attachment reference lines/URIs and the PDF source markers
|
||||
stored by the document editor. False positives are intentionally
|
||||
conservative: retaining an extra upload is safer than deleting referenced
|
||||
bytes.
|
||||
"""
|
||||
return extract_upload_ids(value)
|
||||
|
||||
|
||||
def _upload_ids_from_message_metadata(raw_metadata: object) -> set[str]:
|
||||
"""Extract attachment IDs from a persisted chat metadata JSON value.
|
||||
|
||||
Malformed metadata raises instead of being treated as an empty reference
|
||||
set. The admin cleanup route catches that failure and aborts cleanup.
|
||||
"""
|
||||
if raw_metadata in (None, ""):
|
||||
return set()
|
||||
if isinstance(raw_metadata, str):
|
||||
metadata = json.loads(raw_metadata)
|
||||
else:
|
||||
metadata = raw_metadata
|
||||
if not isinstance(metadata, dict):
|
||||
raise ValueError("chat message metadata must be a JSON object")
|
||||
|
||||
attachments = metadata.get("attachments")
|
||||
if attachments is not None:
|
||||
if not isinstance(attachments, list) or any(
|
||||
not isinstance(item, dict) for item in attachments
|
||||
):
|
||||
raise ValueError("chat message attachments metadata is malformed")
|
||||
|
||||
ids = {
|
||||
str(ref["attachment_id"])
|
||||
for ref in attachment_refs_from_metadata(metadata)
|
||||
if ref.get("attachment_id")
|
||||
}
|
||||
# Preserve canonical IDs even in older metadata shapes not normalized by
|
||||
# attachment_refs_from_metadata().
|
||||
ids.update(_upload_ids_from_persisted_text(json.dumps(metadata)))
|
||||
return ids
|
||||
|
||||
|
||||
def _collect_persisted_upload_references() -> tuple[set[str], set[str]]:
|
||||
"""Collect upload IDs/hashes still referenced by durable application data.
|
||||
|
||||
The caller must treat any exception as an incomplete scan and fail closed.
|
||||
There is no distinct artifact table in the current schema; artifact-like
|
||||
attachment references persisted in chat/document text are covered by the
|
||||
canonical-ID scan.
|
||||
"""
|
||||
referenced_ids: set[str] = set()
|
||||
referenced_hashes: set[str] = set()
|
||||
db = SessionLocal()
|
||||
try:
|
||||
for content, raw_metadata in db.query(
|
||||
DbChatMessage.content,
|
||||
DbChatMessage.meta_data,
|
||||
).yield_per(500):
|
||||
referenced_ids.update(_upload_ids_from_persisted_text(content))
|
||||
referenced_ids.update(_upload_ids_from_message_metadata(raw_metadata))
|
||||
|
||||
for (content,) in db.query(Document.current_content).yield_per(500):
|
||||
referenced_ids.update(_upload_ids_from_persisted_text(content))
|
||||
|
||||
for (content,) in db.query(DocumentVersion.content).yield_per(500):
|
||||
referenced_ids.update(_upload_ids_from_persisted_text(content))
|
||||
|
||||
for filename, file_hash in db.query(
|
||||
GalleryImage.filename,
|
||||
GalleryImage.file_hash,
|
||||
).yield_per(500):
|
||||
referenced_ids.update(_upload_ids_from_persisted_text(filename))
|
||||
if file_hash:
|
||||
referenced_hashes.add(str(file_hash))
|
||||
|
||||
for image_url, color, content, items in db.query(
|
||||
Note.image_url,
|
||||
Note.color,
|
||||
Note.content,
|
||||
Note.items,
|
||||
).yield_per(500):
|
||||
for value in (image_url, color, content, items):
|
||||
referenced_ids.update(_upload_ids_from_persisted_text(value))
|
||||
|
||||
for (color,) in db.query(CalendarCal.color).yield_per(500):
|
||||
referenced_ids.update(_upload_ids_from_persisted_text(color))
|
||||
|
||||
for color, description, location in db.query(
|
||||
CalendarEvent.color,
|
||||
CalendarEvent.description,
|
||||
CalendarEvent.location,
|
||||
).yield_per(500):
|
||||
for value in (color, description, location):
|
||||
referenced_ids.update(_upload_ids_from_persisted_text(value))
|
||||
|
||||
return referenced_ids, referenced_hashes
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _run_reference_safe_cleanup(upload_handler) -> int:
|
||||
referenced_ids, referenced_hashes = _collect_persisted_upload_references()
|
||||
return upload_handler.cleanup_old_uploads(
|
||||
referenced_upload_ids=referenced_ids,
|
||||
referenced_upload_hashes=referenced_hashes,
|
||||
)
|
||||
|
||||
def setup_upload_routes(upload_handler):
|
||||
"""Setup upload routes with the provided handler"""
|
||||
|
||||
@@ -172,7 +296,9 @@ def setup_upload_routes(upload_handler):
|
||||
"mime": meta["mime"],
|
||||
"size": meta["size"],
|
||||
"hash": meta["hash"],
|
||||
"checksum_sha256": meta.get("checksum_sha256") or meta["hash"],
|
||||
"uploaded_at": meta["uploaded_at"],
|
||||
"created_at": meta.get("created_at") or meta["uploaded_at"],
|
||||
"width": meta.get("width"),
|
||||
"height": meta.get("height"),
|
||||
"is_duplicate": meta.get("is_duplicate", False)
|
||||
@@ -195,7 +321,23 @@ def setup_upload_routes(upload_handler):
|
||||
async def manual_cleanup(request: Request):
|
||||
"""Manually trigger cleanup of old uploads."""
|
||||
require_admin(request)
|
||||
cleaned_count = upload_handler.cleanup_old_uploads()
|
||||
try:
|
||||
cleaned_count = await asyncio.to_thread(
|
||||
_run_reference_safe_cleanup,
|
||||
upload_handler,
|
||||
)
|
||||
except UploadCleanupSafetyError:
|
||||
logger.exception("Upload cleanup aborted because index safety checks failed")
|
||||
raise HTTPException(
|
||||
503,
|
||||
"Upload cleanup aborted because upload index integrity could not be verified",
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Upload cleanup skipped because reference discovery failed")
|
||||
raise HTTPException(
|
||||
503,
|
||||
"Upload cleanup skipped because persisted references could not be verified",
|
||||
)
|
||||
return {"status": "success", "files_cleaned": cleaned_count}
|
||||
|
||||
@router.get("/stats")
|
||||
|
||||
Reference in New Issue
Block a user