mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-10-09 16:32:21 +02:00
Compare commits
46
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bec4d1805d | ||
|
|
54d794e8de | ||
|
|
3cd6cdb638 | ||
|
|
2c394704c6 | ||
|
|
f9235ebbf1 | ||
|
|
49e4e55d2c | ||
|
|
b2789d04fb | ||
|
|
a6bc86e331 | ||
|
|
c4369305f0 | ||
|
|
b52296471b | ||
|
|
53869d194d | ||
|
|
17ee856d1c | ||
|
|
e7eddbae13 | ||
|
|
93eb10d4f0 | ||
|
|
3f9633c44f | ||
|
|
858c872832 | ||
|
|
1939a6ad2d | ||
|
|
937c883c41 | ||
|
|
e0615cda47 | ||
|
|
1976fe1b60 | ||
|
|
adfe3ab379 | ||
|
|
d87a913729 | ||
|
|
bea48c749c | ||
|
|
5a016e492c | ||
|
|
93653120d6 | ||
|
|
c2b9666def | ||
|
|
1183fe0ff1 | ||
|
|
663d6879b7 | ||
|
|
3bea7a53ee | ||
|
|
22e0af2a58 | ||
|
|
c00ef8f9c2 | ||
|
|
1fef4929cf | ||
|
|
651bf714de | ||
|
|
d449a9d431 | ||
|
|
dbeed4b63f | ||
|
|
96aca52094 | ||
|
|
8f2f483725 | ||
|
|
42da399b4d | ||
|
|
48cf08328f | ||
|
|
e4fa4ae5dd | ||
|
|
378518f6df | ||
|
|
f06a0a30a8 | ||
|
|
99566d28b5 | ||
|
|
f1e96d102e | ||
|
|
36d4098421 | ||
|
|
5ddef23d94 |
@@ -2,7 +2,7 @@ name: CI
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
branches: [main, dev]
|
||||
pull_request:
|
||||
|
||||
# Least privilege: none of the jobs write to the repo.
|
||||
@@ -103,10 +103,7 @@ jobs:
|
||||
python-tests:
|
||||
name: Python tests (pytest)
|
||||
runs-on: ubuntu-latest
|
||||
# Informational for now: the suite has known flaky / environment-dependent
|
||||
# failures (test isolation + embedding-model assertions). Tracked under the
|
||||
# ROADMAP "fresh install smoke tests" item; make this required once green.
|
||||
continue-on-error: true
|
||||
# Make Python test validation authoritative for the configured scope.
|
||||
steps:
|
||||
- uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0
|
||||
with:
|
||||
|
||||
@@ -630,13 +630,24 @@ app.include_router(auth_router)
|
||||
|
||||
@app.post("/api/activity/heartbeat")
|
||||
async def activity_heartbeat():
|
||||
from src.interactive_gate import mark_browser_activity
|
||||
from src.interactive_gate import (
|
||||
mark_browser_activity,
|
||||
maybe_stop_background_tasks_for_heartbeat,
|
||||
)
|
||||
|
||||
await mark_browser_activity()
|
||||
|
||||
async def _stop_background():
|
||||
try:
|
||||
await task_scheduler.stop_background_tasks_for_foreground(reason="browser heartbeat")
|
||||
await maybe_stop_background_tasks_for_heartbeat(
|
||||
task_scheduler.stop_background_tasks_for_foreground
|
||||
)
|
||||
except Exception:
|
||||
logging.getLogger("app.foreground_gate").debug("heartbeat task stop failed", exc_info=True)
|
||||
logging.getLogger("app.foreground_gate").debug(
|
||||
"heartbeat task stop failed",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
asyncio.create_task(_stop_background())
|
||||
return {"ok": True}
|
||||
|
||||
@@ -805,7 +816,7 @@ app.include_router(setup_font_routes())
|
||||
# MCP (Model Context Protocol)
|
||||
from src.mcp_manager import McpManager
|
||||
from src.agent_tools import set_mcp_manager
|
||||
from routes.mcp_routes import setup_mcp_routes
|
||||
from routes.mcp.mcp_routes import setup_mcp_routes
|
||||
|
||||
mcp_manager = McpManager()
|
||||
set_mcp_manager(mcp_manager)
|
||||
|
||||
+8
-4
@@ -15,17 +15,21 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
|
||||
|
||||
def atomic_write_json(path: str, data: Any, *, indent: Optional[int] = None) -> None:
|
||||
"""Atomically persist `data` as JSON at `path`.
|
||||
|
||||
The temp file uses the live PID as a suffix so two processes saving the
|
||||
same file (e.g. unit tests) don't collide on the rename target.
|
||||
The temp file uses a random suffix so two concurrent writers saving the
|
||||
same file don't collide on the rename target. A PID suffix does not do
|
||||
this: the PID is constant for the life of a process, so two writers on
|
||||
the same path within one process (or one single-process container, where
|
||||
the PID never changes at all) still race for the same temp file.
|
||||
"""
|
||||
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||
tmp = f"{path}.tmp.{os.getpid()}"
|
||||
tmp = f"{path}.tmp.{uuid.uuid4().hex}"
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, indent=indent)
|
||||
f.flush()
|
||||
@@ -37,7 +41,7 @@ def atomic_write_text(path: str, text: str) -> None:
|
||||
if not isinstance(text, str):
|
||||
raise TypeError("atomic_write_text expects a string")
|
||||
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||
tmp = f"{path}.tmp.{os.getpid()}"
|
||||
tmp = f"{path}.tmp.{uuid.uuid4().hex}"
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
f.write(text)
|
||||
f.flush()
|
||||
|
||||
+237
-62
@@ -5,7 +5,7 @@ from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from urllib.parse import unquote, urlparse
|
||||
from sqlalchemy import event, create_engine, Column, String, Text, Boolean, DateTime, Integer, ForeignKey, JSON, Index, func, text
|
||||
from sqlalchemy import DDL, event, create_engine, Column, String, Text, Boolean, DateTime, Integer, ForeignKey, JSON, Index, func, inspect, text
|
||||
from sqlalchemy.engine import Engine, make_url
|
||||
from sqlalchemy.types import TypeDecorator
|
||||
from sqlalchemy.ext.declarative import declarative_base, declared_attr
|
||||
@@ -430,6 +430,93 @@ class EmailAccount(TimestampMixin, Base):
|
||||
)
|
||||
|
||||
|
||||
class EmailAccountOwnerLock(Base):
|
||||
"""Durable per-owner mutex for email-account default mutations.
|
||||
|
||||
Row-locking databases serialize mutations by locking this row before they
|
||||
inspect or stage EmailAccount changes. SQLite uses ``BEGIN IMMEDIATE``
|
||||
instead, because it ignores ``SELECT ... FOR UPDATE``; keeping the table in
|
||||
the shared metadata still makes the non-SQLite path available without a
|
||||
separate migration. The empty key represents the normalized legacy /
|
||||
unconfigured scope shared by ``owner IS NULL`` and ``owner = ''`` rows.
|
||||
"""
|
||||
__tablename__ = "email_account_owner_locks"
|
||||
|
||||
owner_key = Column(String, primary_key=True)
|
||||
|
||||
|
||||
_EMAIL_ACCOUNT_DEFAULT_INDEX = "ux_email_accounts_one_default_per_owner"
|
||||
_EMAIL_ACCOUNT_DEFAULT_INDEX_DDL = {
|
||||
"sqlite": (
|
||||
f"CREATE UNIQUE INDEX IF NOT EXISTS {_EMAIL_ACCOUNT_DEFAULT_INDEX} "
|
||||
"ON email_accounts (COALESCE(owner, '')) WHERE is_default = 1"
|
||||
),
|
||||
"postgresql": (
|
||||
f"CREATE UNIQUE INDEX IF NOT EXISTS {_EMAIL_ACCOUNT_DEFAULT_INDEX} "
|
||||
"ON email_accounts ((COALESCE(owner, ''))) WHERE is_default IS TRUE"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
# SQLAlchemy cannot express one portable partial, functional index across the
|
||||
# two supported database families. Register dialect-specific DDL so fresh
|
||||
# databases get the invariant as part of create_all(); the startup migration
|
||||
# below installs the same index on existing databases after normalizing legacy
|
||||
# duplicate rows.
|
||||
for _dialect_name, _index_ddl in _EMAIL_ACCOUNT_DEFAULT_INDEX_DDL.items():
|
||||
event.listen(
|
||||
EmailAccount.__table__,
|
||||
"after_create",
|
||||
DDL(_index_ddl).execute_if(dialect=_dialect_name),
|
||||
)
|
||||
|
||||
|
||||
def lock_email_account_owner_mutations(db, *owners: str) -> None:
|
||||
"""Lock normalized email-account owner scopes in canonical order.
|
||||
|
||||
``NULL`` and the empty string are one legacy/single-user owner partition,
|
||||
matching the unique default-account index. SQLite has only a database
|
||||
writer reservation, while row-locking databases use durable mutex rows.
|
||||
Sorting all requested owner keys keeps multi-owner operations such as user
|
||||
rename from deadlocking with another mutation that requests the same keys
|
||||
in the opposite order.
|
||||
"""
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
owner_keys = sorted({owner or "" for owner in owners} or {""})
|
||||
if db.get_bind().dialect.name == "sqlite":
|
||||
db.execute(text("BEGIN IMMEDIATE"))
|
||||
return
|
||||
|
||||
for owner_key in owner_keys:
|
||||
lock_row = db.get(
|
||||
EmailAccountOwnerLock,
|
||||
owner_key,
|
||||
with_for_update=True,
|
||||
)
|
||||
if lock_row is not None:
|
||||
continue
|
||||
|
||||
inserted = False
|
||||
try:
|
||||
with db.begin_nested():
|
||||
db.add(EmailAccountOwnerLock(owner_key=owner_key))
|
||||
db.flush()
|
||||
inserted = True
|
||||
except IntegrityError:
|
||||
# A competing transaction created the mutex row first. Once its
|
||||
# insert commits, lock that durable row before touching accounts.
|
||||
pass
|
||||
|
||||
if not inserted:
|
||||
(
|
||||
db.query(EmailAccountOwnerLock)
|
||||
.filter(EmailAccountOwnerLock.owner_key == owner_key)
|
||||
.with_for_update()
|
||||
.one()
|
||||
)
|
||||
|
||||
|
||||
class ModelEndpoint(TimestampMixin, Base):
|
||||
"""Admin-configured model endpoints. Models are auto-discovered via /v1/models."""
|
||||
__tablename__ = "model_endpoints"
|
||||
@@ -1404,8 +1491,25 @@ def _migrate_assign_legacy_owner():
|
||||
with open(prefs_path, "r", encoding="utf-8") as f:
|
||||
prefs = _json.load(f)
|
||||
if "_users" not in prefs and prefs:
|
||||
# Flat format → nest under admin user
|
||||
new_prefs = {"_users": {admin_user: prefs}}
|
||||
# Flat format → nest ordinary preferences under the admin
|
||||
# user. Foreground fallback is an explicit per-owner opt-in,
|
||||
# so auth-disabled consent must remain inert at the flat root
|
||||
# rather than becoming consent for the first named owner.
|
||||
foreground_keys = {
|
||||
"foreground_fallback_enabled",
|
||||
"foreground_model_fallbacks",
|
||||
}
|
||||
named_prefs = {
|
||||
key: value
|
||||
for key, value in prefs.items()
|
||||
if key not in foreground_keys
|
||||
}
|
||||
new_prefs = {
|
||||
key: prefs[key]
|
||||
for key in foreground_keys
|
||||
if key in prefs
|
||||
}
|
||||
new_prefs["_users"] = {admin_user: named_prefs}
|
||||
with open(prefs_path, "w", encoding="utf-8") as f:
|
||||
_json.dump(new_prefs, f, indent=2)
|
||||
logger.info(f"Migrated user_prefs.json to per-user format under '{admin_user}'")
|
||||
@@ -1812,72 +1916,142 @@ class Integration(TimestampMixin, Base):
|
||||
|
||||
|
||||
|
||||
def _migrate_seed_email_account():
|
||||
"""If email_accounts is empty and settings.json has legacy flat imap_host/smtp_host
|
||||
keys, create a single default account from them so nothing breaks for users who
|
||||
upgraded. Safe to run repeatedly — it short-circuits once any row exists."""
|
||||
def _migrate_email_account_default_invariant():
|
||||
"""Normalize legacy duplicates and install durable at-most-one enforcement.
|
||||
|
||||
Older databases only had a non-unique ``(owner, is_default)`` lookup index.
|
||||
Keep the oldest default deterministically in each normalized owner scope,
|
||||
then add the same partial functional unique index used for fresh schemas.
|
||||
"""
|
||||
dialect_name = engine.dialect.name
|
||||
index_ddl = _EMAIL_ACCOUNT_DEFAULT_INDEX_DDL.get(dialect_name)
|
||||
if index_ddl is None:
|
||||
logger.warning(
|
||||
"Email-account default uniqueness is not available for database "
|
||||
"dialect %s; mutations remain serialized but are not protected by "
|
||||
"a database constraint",
|
||||
dialect_name,
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
with engine.connect() as conn:
|
||||
tables = [r[0] for r in conn.execute(text(
|
||||
"SELECT name FROM sqlite_master WHERE type='table' AND name='email_accounts'"
|
||||
))]
|
||||
if "email_accounts" not in tables:
|
||||
return
|
||||
existing = conn.execute(text("SELECT COUNT(*) FROM email_accounts")).scalar() or 0
|
||||
if existing > 0:
|
||||
with engine.begin() as conn:
|
||||
if not inspect(conn).has_table(EmailAccount.__tablename__):
|
||||
return
|
||||
default_rows = conn.execute(text("""
|
||||
SELECT id, owner
|
||||
FROM email_accounts
|
||||
WHERE is_default IS TRUE
|
||||
ORDER BY
|
||||
COALESCE(owner, ''),
|
||||
CASE WHEN created_at IS NULL THEN 1 ELSE 0 END,
|
||||
created_at,
|
||||
id
|
||||
""")).mappings()
|
||||
seen_owner_keys = set()
|
||||
duplicate_ids = []
|
||||
for row in default_rows:
|
||||
owner_key = row["owner"] or ""
|
||||
if owner_key in seen_owner_keys:
|
||||
duplicate_ids.append(row["id"])
|
||||
else:
|
||||
seen_owner_keys.add(owner_key)
|
||||
|
||||
import json as _json
|
||||
import uuid as _uuid
|
||||
from pathlib import Path
|
||||
settings_file = Path(SETTINGS_FILE)
|
||||
if not settings_file.exists():
|
||||
return
|
||||
try:
|
||||
s = _json.loads(settings_file.read_text(encoding="utf-8"))
|
||||
except Exception:
|
||||
return
|
||||
for account_id in duplicate_ids:
|
||||
conn.execute(
|
||||
text("UPDATE email_accounts SET is_default = :value WHERE id = :id"),
|
||||
{"value": False, "id": account_id},
|
||||
)
|
||||
conn.execute(text(index_ddl))
|
||||
|
||||
imap_host = (s.get("imap_host") or "").strip()
|
||||
smtp_host = (s.get("smtp_host") or "").strip()
|
||||
if not imap_host and not smtp_host:
|
||||
return # nothing to migrate
|
||||
if duplicate_ids:
|
||||
logger.warning(
|
||||
"Normalized %d duplicate default email account(s) before "
|
||||
"installing %s",
|
||||
len(duplicate_ids),
|
||||
_EMAIL_ACCOUNT_DEFAULT_INDEX,
|
||||
)
|
||||
except Exception:
|
||||
# Starting without the constraint would silently retain the race this
|
||||
# migration is intended to close. Fail startup so an operator sees and
|
||||
# can repair an incompatible schema instead of accepting unsafe writes.
|
||||
logger.exception("Failed to enforce the email-account default invariant")
|
||||
raise
|
||||
|
||||
|
||||
def _migrate_seed_email_account():
|
||||
"""Atomically seed one legacy default account when no account exists.
|
||||
|
||||
Reading settings is intentionally done before taking the owner mutex. The
|
||||
decisive emptiness check and insert share one locked transaction, so two
|
||||
application workers starting together cannot both seed a default row.
|
||||
"""
|
||||
import json as _json
|
||||
import uuid as _uuid
|
||||
|
||||
settings_file = Path(SETTINGS_FILE)
|
||||
if not settings_file.exists():
|
||||
return
|
||||
try:
|
||||
s = _json.loads(settings_file.read_text(encoding="utf-8"))
|
||||
except Exception:
|
||||
return
|
||||
|
||||
imap_host = (s.get("imap_host") or "").strip()
|
||||
smtp_host = (s.get("smtp_host") or "").strip()
|
||||
if not imap_host and not smtp_host:
|
||||
return
|
||||
|
||||
db = None
|
||||
try:
|
||||
if not inspect(engine).has_table(EmailAccount.__tablename__):
|
||||
return
|
||||
db = SessionLocal()
|
||||
lock_email_account_owner_mutations(db, "")
|
||||
existing = db.execute(text("SELECT COUNT(*) FROM email_accounts")).scalar() or 0
|
||||
if existing > 0:
|
||||
return
|
||||
|
||||
now = utcnow_naive()
|
||||
with engine.begin() as conn:
|
||||
conn.execute(text("""
|
||||
INSERT INTO email_accounts
|
||||
(id, owner, name, is_default, enabled,
|
||||
imap_host, imap_port, imap_user, imap_password, imap_starttls,
|
||||
smtp_host, smtp_port, smtp_user, smtp_password,
|
||||
from_address, created_at, updated_at)
|
||||
VALUES
|
||||
(:id, :owner, :name, :is_default, :enabled,
|
||||
:imap_host, :imap_port, :imap_user, :imap_password, :imap_starttls,
|
||||
:smtp_host, :smtp_port, :smtp_user, :smtp_password,
|
||||
:from_address, :created_at, :updated_at)
|
||||
"""), {
|
||||
"id": _uuid.uuid4().hex,
|
||||
"owner": None,
|
||||
"name": "Default",
|
||||
"is_default": True,
|
||||
"enabled": True,
|
||||
"imap_host": imap_host,
|
||||
"imap_port": int(s.get("imap_port") or 993),
|
||||
"imap_user": s.get("imap_user") or "",
|
||||
"imap_password": s.get("imap_password") or "",
|
||||
"imap_starttls": bool(s.get("imap_starttls", True)),
|
||||
"smtp_host": smtp_host,
|
||||
"smtp_port": int(s.get("smtp_port") or 465),
|
||||
"smtp_user": s.get("smtp_user") or "",
|
||||
"smtp_password": s.get("smtp_password") or "",
|
||||
"from_address": s.get("email_from") or "",
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
})
|
||||
logging.getLogger(__name__).info("Seeded email_accounts 'Default' from settings.json")
|
||||
db.execute(text("""
|
||||
INSERT INTO email_accounts
|
||||
(id, owner, name, is_default, enabled,
|
||||
imap_host, imap_port, imap_user, imap_password, imap_starttls,
|
||||
smtp_host, smtp_port, smtp_user, smtp_password,
|
||||
from_address, created_at, updated_at)
|
||||
VALUES
|
||||
(:id, :owner, :name, :is_default, :enabled,
|
||||
:imap_host, :imap_port, :imap_user, :imap_password, :imap_starttls,
|
||||
:smtp_host, :smtp_port, :smtp_user, :smtp_password,
|
||||
:from_address, :created_at, :updated_at)
|
||||
"""), {
|
||||
"id": _uuid.uuid4().hex,
|
||||
"owner": None,
|
||||
"name": "Default",
|
||||
"is_default": True,
|
||||
"enabled": True,
|
||||
"imap_host": imap_host,
|
||||
"imap_port": int(s.get("imap_port") or 993),
|
||||
"imap_user": s.get("imap_user") or "",
|
||||
"imap_password": s.get("imap_password") or "",
|
||||
"imap_starttls": bool(s.get("imap_starttls", True)),
|
||||
"smtp_host": smtp_host,
|
||||
"smtp_port": int(s.get("smtp_port") or 465),
|
||||
"smtp_user": s.get("smtp_user") or "",
|
||||
"smtp_password": s.get("smtp_password") or "",
|
||||
"from_address": s.get("email_from") or "",
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
})
|
||||
db.commit()
|
||||
logger.info("Seeded email_accounts 'Default' from settings.json")
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"seed email account migration: {e}")
|
||||
if db is not None:
|
||||
db.rollback()
|
||||
logger.warning("seed email account migration: %s", e)
|
||||
finally:
|
||||
if db is not None:
|
||||
db.close()
|
||||
|
||||
|
||||
# WARNING: Foreign-key enforcement is enabled globally for all SQLite connections.
|
||||
@@ -1960,6 +2134,7 @@ def init_db():
|
||||
_migrate_add_crew_member_id()
|
||||
_migrate_add_assistant_columns()
|
||||
_migrate_add_email_smtp_security()
|
||||
_migrate_email_account_default_invariant()
|
||||
_migrate_seed_email_account()
|
||||
_migrate_add_calendar_metadata()
|
||||
_migrate_add_calendar_is_utc()
|
||||
|
||||
+41
-12
@@ -194,7 +194,12 @@ class SessionManager:
|
||||
is_important=getattr(db_session, 'is_important', False) or False,
|
||||
)
|
||||
|
||||
session.message_count = getattr(db_session, 'message_count', len(history))
|
||||
# The rows just loaded are the whole transcript, so they — not the
|
||||
# denormalized sessions.message_count column — are the truth for this
|
||||
# cached object. get_session's hydration gate compares against this
|
||||
# number; seeding it from a drifted column would ask for a reload that
|
||||
# can never close the gap.
|
||||
session.message_count = len(history)
|
||||
return session
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -398,30 +403,50 @@ class SessionManager:
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def get_session(self, session_id: str) -> Session:
|
||||
"""Get a session by ID, loading from DB if needed.
|
||||
"""Get a session by ID, loading complete DB history when needed.
|
||||
|
||||
Sessions seeded by `load_sessions` start with empty history. The
|
||||
first read here hydrates them with the message rows.
|
||||
Sessions seeded by ``load_sessions`` start with empty history, and a
|
||||
cached session can also become partially stale. Refresh metadata first,
|
||||
then hydrate whenever the cached transcript is short of the stored rows.
|
||||
Model-send routes enter through this method before building context,
|
||||
while paginated display history reads SQLite directly.
|
||||
|
||||
The gate compares against ``sync_session_metadata``'s reconciled count
|
||||
(the real ``chat_messages`` total), never the denormalized column, so a
|
||||
hydrate always closes the gap and the next read is a cache hit.
|
||||
"""
|
||||
if session_id not in self.sessions:
|
||||
self._load_session_from_db(session_id)
|
||||
else:
|
||||
cached = self.sessions[session_id]
|
||||
# Lazy hydrate: metadata-only entries get their messages on first read.
|
||||
if not cached.history and getattr(cached, "message_count", 0) > 0:
|
||||
self._load_session_from_db(session_id)
|
||||
|
||||
# Keep model/endpoint metadata fresh. Endpoint deletion can clear the
|
||||
# DB row while a session object is still cached in RAM.
|
||||
# DB row while a session object is still cached in RAM. Refreshing first
|
||||
# also exposes the authoritative message count before completeness is
|
||||
# checked.
|
||||
self.sync_session_metadata(session_id)
|
||||
|
||||
cached = self.sessions[session_id]
|
||||
cached_count = len(cached.history or [])
|
||||
stored_count = int(getattr(cached, "message_count", 0) or 0)
|
||||
if cached_count < stored_count:
|
||||
self._load_session_from_db(session_id)
|
||||
|
||||
# Update last_accessed
|
||||
self._touch_session(session_id)
|
||||
|
||||
return self.sessions[session_id]
|
||||
|
||||
def sync_session_metadata(self, session_id: str) -> bool:
|
||||
"""Refresh non-message session fields from the DB into the cached object."""
|
||||
"""Refresh non-message session fields from the DB into the cached object.
|
||||
|
||||
``message_count`` is reconciled against the real ``chat_messages`` rows
|
||||
rather than copied from the denormalized ``sessions.message_count``
|
||||
column. That column drifts in normal operation — ``_persist_message``
|
||||
swallows a failed insert but ``add_message`` has already appended in
|
||||
memory, so the next successful persist writes rows+1, and a persist for
|
||||
an uncached session writes 0. Hydration keys off this number: a
|
||||
drifted-high column would reload the whole transcript on every warm
|
||||
read, and a drifted-low one would leave the model a truncated one.
|
||||
"""
|
||||
session = self.sessions.get(session_id)
|
||||
if session is None:
|
||||
return False
|
||||
@@ -444,7 +469,11 @@ class SessionManager:
|
||||
session.archived = db_session.archived
|
||||
session.owner = getattr(db_session, "owner", None)
|
||||
session.is_important = getattr(db_session, "is_important", False) or False
|
||||
session.message_count = getattr(db_session, "message_count", session.message_count) or 0
|
||||
session.message_count = (
|
||||
db.query(DbChatMessage)
|
||||
.filter(DbChatMessage.session_id == session_id)
|
||||
.count()
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Error syncing session metadata {session_id}: {e}")
|
||||
|
||||
@@ -129,12 +129,17 @@ services:
|
||||
fi
|
||||
sed "s|__SEARXNG_SECRET__|$$secret|g" /tmp/searxng-settings.yml.template > /etc/searxng/settings.yml
|
||||
fi
|
||||
# Advisory: a settings file the migration cannot parse or rewrite must
|
||||
# not be what stops searxng from booting. It explains itself on stderr
|
||||
# and we carry on, letting searxng report anything genuinely wrong.
|
||||
/usr/local/searxng/.venv/bin/python /tmp/migrate-searxng-settings.py /etc/searxng/settings.yml || true
|
||||
exec /usr/local/searxng/entrypoint.sh
|
||||
ports:
|
||||
- "127.0.0.1:8080:8080"
|
||||
volumes:
|
||||
- searxng-data:/etc/searxng
|
||||
- ./config/searxng/settings.yml:/tmp/searxng-settings.yml.template:ro,z
|
||||
- ./scripts/migrate_searxng_settings.py:/tmp/migrate-searxng-settings.py:ro,z
|
||||
environment:
|
||||
- SEARXNG_BASE_URL=http://localhost:8080/
|
||||
- SEARXNG_SECRET=${SEARXNG_SECRET:-}
|
||||
|
||||
@@ -132,12 +132,17 @@ services:
|
||||
fi
|
||||
sed "s|__SEARXNG_SECRET__|$$secret|g" /tmp/searxng-settings.yml.template > /etc/searxng/settings.yml
|
||||
fi
|
||||
# Advisory: a settings file the migration cannot parse or rewrite must
|
||||
# not be what stops searxng from booting. It explains itself on stderr
|
||||
# and we carry on, letting searxng report anything genuinely wrong.
|
||||
/usr/local/searxng/.venv/bin/python /tmp/migrate-searxng-settings.py /etc/searxng/settings.yml || true
|
||||
exec /usr/local/searxng/entrypoint.sh
|
||||
ports:
|
||||
- "127.0.0.1:8080:8080"
|
||||
volumes:
|
||||
- searxng-data:/etc/searxng
|
||||
- ./config/searxng/settings.yml:/tmp/searxng-settings.yml.template:ro,z
|
||||
- ./scripts/migrate_searxng_settings.py:/tmp/migrate-searxng-settings.py:ro,z
|
||||
environment:
|
||||
- SEARXNG_BASE_URL=http://localhost:8080/
|
||||
- SEARXNG_SECRET=${SEARXNG_SECRET:-}
|
||||
|
||||
@@ -110,12 +110,17 @@ services:
|
||||
fi
|
||||
sed "s|__SEARXNG_SECRET__|$$secret|g" /tmp/searxng-settings.yml.template > /etc/searxng/settings.yml
|
||||
fi
|
||||
# Advisory: a settings file the migration cannot parse or rewrite must
|
||||
# not be what stops searxng from booting. It explains itself on stderr
|
||||
# and we carry on, letting searxng report anything genuinely wrong.
|
||||
/usr/local/searxng/.venv/bin/python /tmp/migrate-searxng-settings.py /etc/searxng/settings.yml || true
|
||||
exec /usr/local/searxng/entrypoint.sh
|
||||
ports:
|
||||
- "127.0.0.1:8080:8080"
|
||||
volumes:
|
||||
- searxng-data:/etc/searxng
|
||||
- ./config/searxng/settings.yml:/tmp/searxng-settings.yml.template:ro,z
|
||||
- ./scripts/migrate_searxng_settings.py:/tmp/migrate-searxng-settings.py:ro,z
|
||||
environment:
|
||||
- SEARXNG_BASE_URL=http://localhost:8080/
|
||||
- SEARXNG_SECRET=${SEARXNG_SECRET:-}
|
||||
|
||||
+174
@@ -309,6 +309,32 @@ container. Cookbook **Serve** is a separate workflow for serving downloaded
|
||||
models through Odysseus/llama.cpp, so Windows users with an existing Ollama
|
||||
install usually only need to add the endpoint in Settings.
|
||||
|
||||
**Tool calls not firing on a manually-added Ollama `/v1` endpoint.** By
|
||||
design, a local Ollama `/v1` endpoint defaults to the conservative
|
||||
text-based (fenced-block) tool-calling path rather than native structured
|
||||
tool calls, since some locally-served models mishandle native schemas (see
|
||||
#1567). This is correct for most local setups, but if you know your specific
|
||||
model reliably supports native tool calling (check `ollama show <model>` for
|
||||
`tools` under Capabilities), you can opt that endpoint in explicitly. There
|
||||
is currently no UI control for this on manually-added endpoints (see #5192);
|
||||
the flag can still be set directly against the existing API, from a browser
|
||||
console on an authenticated admin session:
|
||||
|
||||
```js
|
||||
fetch('/api/model-endpoints/<endpoint-id>', {
|
||||
method: 'PATCH',
|
||||
credentials: 'same-origin',
|
||||
headers: {'Content-Type': 'application/json'},
|
||||
body: JSON.stringify({supports_tools: true})
|
||||
}).then(r => r.json()).then(console.log)
|
||||
```
|
||||
|
||||
Find `<endpoint-id>` by inspecting the `/api/model-endpoints` response (or
|
||||
your browser's network tab while Settings loads the endpoint list). Send
|
||||
`supports_tools: false` to disable native structured tool calls and force the
|
||||
conservative fenced/text path, or `supports_tools: null` to return the endpoint
|
||||
to the Auto heuristic.
|
||||
|
||||
**Useful checks.**
|
||||
|
||||
```bash
|
||||
@@ -471,6 +497,154 @@ Odysseus serves plain HTTP on its app port. Docker Compose binds Odysseus and th
|
||||
Cloudflare Access, Tailscale, Caddy, nginx, and Traefik can all fit this pattern; none are required by Odysseus. If your access layer reaches Odysseus on the same host, proxy to `http://127.0.0.1:7000` and keep `AUTH_ENABLED=true`, `LOCALHOST_BYPASS=false`, and `SECURE_COOKIES=true`.
|
||||
`ALLOWED_ORIGINS` lists exact permitted origins for cross-origin browser/API clients; ordinary same-origin reverse-proxy access usually does not need a special CORS entry.
|
||||
|
||||
#### Faster over the network: HTTP/2
|
||||
|
||||
The frontend is raw ES modules with no bundler, so a page load is a few hundred
|
||||
small same-origin requests. Over HTTP/1.1 browsers typically allow only a small
|
||||
number of concurrent connections per host (commonly around six), so many of
|
||||
those requests are serialized across multiple round trips. On localhost that
|
||||
costs almost nothing. Over a LAN, VPN, or remote link it can become a major
|
||||
part of load time, especially as latency increases.
|
||||
|
||||
HTTP/2 multiplexes them onto one connection and the serialisation disappears.
|
||||
Odysseus needs no changes for this — uvicorn keeps speaking HTTP/1.1 on
|
||||
loopback and the proxy speaks HTTP/2 to the browser. Mainstream browsers
|
||||
negotiate HTTP/2 for normal web pages over TLS; they do not use the cleartext
|
||||
h2c mode here, so browser-facing HTTP/2 requires a certificate. The
|
||||
`--ssl-certfile` route in *HTTPS + LAN/Tailscale exposure* above gives you
|
||||
HTTPS but not HTTP/2 — uvicorn does not speak it.
|
||||
|
||||
**1. Install Caddy.** See the [install docs](https://caddyserver.com/docs/install)
|
||||
for your platform; on macOS, `brew install caddy`.
|
||||
|
||||
**2. Write a `Caddyfile`.** Pick the block that matches how you reach the
|
||||
machine. Replace `7000` if Odysseus listens elsewhere — the macOS start script
|
||||
uses `7860`.
|
||||
|
||||
Public domain, Caddy obtains and renews the certificate itself:
|
||||
|
||||
```
|
||||
odysseus.example.com {
|
||||
reverse_proxy 127.0.0.1:7000
|
||||
}
|
||||
```
|
||||
|
||||
Tailscale, no public DNS needed — `tailscale cert` issues a browser-trusted
|
||||
certificate for a tailnet name and writes `<domain>.crt` and `<domain>.key`:
|
||||
|
||||
```bash
|
||||
tailscale cert myhost.tailnet-name.ts.net
|
||||
```
|
||||
|
||||
```
|
||||
myhost.tailnet-name.ts.net {
|
||||
tls /path/to/myhost.tailnet-name.ts.net.crt /path/to/myhost.tailnet-name.ts.net.key
|
||||
reverse_proxy 127.0.0.1:7000
|
||||
}
|
||||
```
|
||||
|
||||
LAN with your own certificate — same shape, your own files:
|
||||
|
||||
```
|
||||
odysseus.lan {
|
||||
tls /path/to/cert.pem /path/to/key.pem
|
||||
reverse_proxy 127.0.0.1:7000
|
||||
}
|
||||
```
|
||||
|
||||
Give `tls` absolute paths: a service starts in a working directory you did not
|
||||
choose. If port 443 is already taken, append a port to the site address
|
||||
(`odysseus.example.com:8443`) and use it in the URL. That alone does not free
|
||||
port 80 — Caddy still binds it for the HTTP-to-HTTPS redirect, and fails to
|
||||
start with `listen tcp :80: bind: address already in use` if something else
|
||||
holds it. Turn the redirect off with a global block at the top of the file:
|
||||
|
||||
```
|
||||
{
|
||||
auto_https disable_redirects
|
||||
}
|
||||
```
|
||||
|
||||
**3. Run it in the foreground first:**
|
||||
|
||||
```bash
|
||||
caddy run --config ./Caddyfile
|
||||
```
|
||||
|
||||
Once that works, run it as a service:
|
||||
|
||||
```bash
|
||||
brew services start caddy # macOS — reads $(brew --prefix)/etc/Caddyfile, not ./Caddyfile
|
||||
sudo systemctl enable --now caddy # Linux, if your package installed the unit
|
||||
```
|
||||
|
||||
Odysseus's own service is unchanged; the proxy runs alongside it. Under Docker,
|
||||
run the proxy as another container, or on the host pointing at the published
|
||||
port.
|
||||
|
||||
**4. Point Odysseus at the new origin** in `.env`, then restart it:
|
||||
|
||||
```bash
|
||||
SECURE_COOKIES=true
|
||||
# only if you use remote MCP servers with OAuth:
|
||||
OAUTH_REDIRECT_BASE_URL=https://odysseus.example.com
|
||||
```
|
||||
|
||||
Gmail OAuth needs nothing here when the proxy runs on the same host: the
|
||||
redirect URI is built from the incoming request, and uvicorn rewrites the
|
||||
scheme from `X-Forwarded-Proto` for proxies it trusts — by default only
|
||||
`127.0.0.1`. A proxy in a separate container or on another machine is not
|
||||
trusted, so pin the URI there:
|
||||
|
||||
```bash
|
||||
GOOGLE_OAUTH_REDIRECT_URI=https://odysseus.example.com/api/email/oauth/google/callback
|
||||
```
|
||||
|
||||
(uvicorn's own `FORWARDED_ALLOW_IPS` widens that trust, but it has to be in the
|
||||
environment uvicorn starts with — `.env` is read by the app afterwards, too
|
||||
late for it to take effect.)
|
||||
|
||||
**5. Confirm HTTP/2 is really on:**
|
||||
|
||||
```bash
|
||||
curl -s -o /dev/null -w '%{http_version}\n' https://odysseus.example.com/
|
||||
# 2
|
||||
```
|
||||
|
||||
The status code is not the thing to check here — a logged-out request redirects
|
||||
to the login page, so `curl -I` shows `HTTP/2 302`, and the `HTTP/2` prefix is
|
||||
the part that matters. The browser reports the same in the Network panel's
|
||||
Protocol column (`h2`); in Chrome and Firefox that column is hidden until you
|
||||
enable it by right-clicking the column headers.
|
||||
|
||||
Three things bite when moving an existing install behind TLS:
|
||||
|
||||
- Set `SECURE_COOKIES=true` **at the same time** you stop serving plain HTTP,
|
||||
not before. The flag is applied to every login regardless of the scheme the
|
||||
request arrived on, so while an HTTP entrypoint is still reachable the
|
||||
browser will reject the `Secure` cookie there and login will appear to loop.
|
||||
- `OAUTH_REDIRECT_BASE_URL` defaults to `http://localhost:7000`. Unlike the
|
||||
Gmail redirect URI it cannot be derived from a request — it is registered
|
||||
with each MCP authorization server up front — so set it to the external
|
||||
origin if you use remote MCP servers over OAuth.
|
||||
- Odysseus sends `Strict-Transport-Security` once it sees `X-Forwarded-Proto:
|
||||
https`. HSTS applies to the whole hostname and ignores the port, so any other
|
||||
plain-HTTP service on that same hostname becomes unreachable in browsers that
|
||||
have visited Odysseus. Give Odysseus its own hostname, or strip the header at
|
||||
the proxy (`header_down -Strict-Transport-Security` in Caddy).
|
||||
|
||||
Server-sent events are not buffered by this configuration, so chat streaming
|
||||
arrives token by token; add `flush_interval -1` inside the `reverse_proxy`
|
||||
block if you want that pinned explicitly. nginx needs `proxy_buffering off;`
|
||||
for the same reason.
|
||||
|
||||
Changing the external origin also affects state scoped to it. Service workers
|
||||
and their caches are origin-scoped, so moving to a different origin starts with
|
||||
a cold load. Cookies follow their own domain/path/security rules rather than
|
||||
being port-scoped: changing the hostname normally requires a new login, while
|
||||
changing only the scheme or port does not by itself guarantee that existing
|
||||
cookies disappear.
|
||||
|
||||
Common internal-only ports from the default docs/compose setup:
|
||||
|
||||
| Port | Service |
|
||||
|
||||
@@ -1802,7 +1802,6 @@ async def _ai_draft_reply_to_email(uid, folder="INBOX", reply_all=False, account
|
||||
from src.endpoint_resolver import (
|
||||
resolve_endpoint,
|
||||
resolve_utility_fallback_candidates,
|
||||
resolve_chat_fallback_candidates,
|
||||
)
|
||||
from src.llm_core import llm_call_async_with_fallback
|
||||
except Exception as exc:
|
||||
@@ -1843,13 +1842,6 @@ async def _ai_draft_reply_to_email(uid, folder="INBOX", reply_all=False, account
|
||||
utility_fallbacks = resolve_utility_fallback_candidates() or []
|
||||
for cand in utility_fallbacks:
|
||||
_add(*cand)
|
||||
try:
|
||||
chat_fallbacks = resolve_chat_fallback_candidates(owner=None) or []
|
||||
except TypeError:
|
||||
chat_fallbacks = resolve_chat_fallback_candidates() or []
|
||||
for cand in chat_fallbacks:
|
||||
_add(*cand)
|
||||
|
||||
if not candidates:
|
||||
return {"error": "No LLM endpoint configured for AI reply"}
|
||||
|
||||
|
||||
+59
-3
@@ -22,6 +22,8 @@ from src.settings import (
|
||||
load_features as _load_features,
|
||||
save_features as _save_features,
|
||||
DEFAULT_SETTINGS,
|
||||
RETIRED_SETTING_KEYS,
|
||||
without_retired_settings,
|
||||
)
|
||||
from src.integrations import (
|
||||
load_integrations,
|
||||
@@ -345,9 +347,61 @@ def setup_auth_routes(auth_manager: AuthManager) -> APIRouter:
|
||||
# docs, email accounts, tasks, etc.
|
||||
try:
|
||||
from sqlalchemy import func
|
||||
from core.database import Base, SessionLocal
|
||||
from core.database import (
|
||||
Base,
|
||||
EmailAccount,
|
||||
SessionLocal,
|
||||
lock_email_account_owner_mutations,
|
||||
)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
# Email-account defaults are protected by per-owner mutex rows.
|
||||
# A rename crosses two owner partitions, so lock both in the
|
||||
# shared helper's canonical order before inspecting either.
|
||||
lock_email_account_owner_mutations(
|
||||
db, old_username, new_username
|
||||
)
|
||||
|
||||
source_default_ids = [
|
||||
row[0]
|
||||
for row in (
|
||||
db.query(EmailAccount.id)
|
||||
.filter(
|
||||
func.lower(EmailAccount.owner) == old_username,
|
||||
EmailAccount.is_default == True, # noqa: E712
|
||||
)
|
||||
.order_by(EmailAccount.created_at.asc(), EmailAccount.id.asc())
|
||||
.all()
|
||||
)
|
||||
]
|
||||
destination_default_ids = [
|
||||
row[0]
|
||||
for row in (
|
||||
db.query(EmailAccount.id)
|
||||
.filter(
|
||||
func.lower(EmailAccount.owner) == new_username,
|
||||
EmailAccount.is_default == True, # noqa: E712
|
||||
)
|
||||
.order_by(EmailAccount.created_at.asc(), EmailAccount.id.asc())
|
||||
.all()
|
||||
)
|
||||
]
|
||||
if destination_default_ids:
|
||||
clear_default_ids = (
|
||||
destination_default_ids[1:] + source_default_ids
|
||||
)
|
||||
else:
|
||||
clear_default_ids = source_default_ids[1:]
|
||||
if clear_default_ids:
|
||||
(
|
||||
db.query(EmailAccount)
|
||||
.filter(EmailAccount.id.in_(clear_default_ids))
|
||||
.update(
|
||||
{EmailAccount.is_default: False},
|
||||
synchronize_session=False,
|
||||
)
|
||||
)
|
||||
|
||||
for mapper in Base.registry.mappers:
|
||||
model = mapper.class_
|
||||
if not hasattr(model, "owner"):
|
||||
@@ -637,7 +691,7 @@ def setup_auth_routes(auth_manager: AuthManager) -> APIRouter:
|
||||
a scrubbed copy with secret keys blanked. The frontend uses this
|
||||
for keybinds + TTS prefs, so it stays callable without admin."""
|
||||
user = _get_current_user(request)
|
||||
settings = _load_settings()
|
||||
settings = without_retired_settings(_load_settings())
|
||||
if user and auth_manager.is_admin(user):
|
||||
return settings
|
||||
return scrub_settings(settings)
|
||||
@@ -657,6 +711,8 @@ def setup_auth_routes(auth_manager: AuthManager) -> APIRouter:
|
||||
"agent_max_tool_calls": (0, 1000), # 0 = unlimited
|
||||
}
|
||||
for key in DEFAULT_SETTINGS:
|
||||
if key in RETIRED_SETTING_KEYS:
|
||||
continue
|
||||
if key not in body:
|
||||
continue
|
||||
val = body[key]
|
||||
@@ -669,7 +725,7 @@ def setup_auth_routes(auth_manager: AuthManager) -> APIRouter:
|
||||
val = max(lo, min(val, hi))
|
||||
current[key] = val
|
||||
_save_settings(current)
|
||||
return current
|
||||
return without_retired_settings(current)
|
||||
|
||||
# ---- Integrations CRUD ----
|
||||
|
||||
|
||||
+115
-7
@@ -10,6 +10,7 @@ from typing import Optional, List
|
||||
from fastapi import APIRouter, HTTPException, Request, UploadFile, File
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import or_, and_
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from dateutil.rrule import rrulestr
|
||||
|
||||
from core.database import SessionLocal, CalendarCal, CalendarDeletedEvent, CalendarEvent
|
||||
@@ -221,22 +222,125 @@ class EventUpdate(BaseModel):
|
||||
|
||||
# ── Helpers ──
|
||||
|
||||
_DEFAULT_CALENDAR_NAMESPACE = uuid.UUID("4840613a-9847-4a3b-bd75-19e6bc5fc3ce")
|
||||
|
||||
|
||||
def _default_calendar_id(owner: str, collision_index: int = 0) -> str:
|
||||
"""Return one stable primary-key candidate for an owner's lazy default.
|
||||
|
||||
Slot zero preserves the original owner-derived identifier. Later slots
|
||||
let a username be reused after its prior calendar was migrated to another
|
||||
owner during a rename, without making concurrent first use choose random
|
||||
and therefore divergent identifiers.
|
||||
"""
|
||||
if collision_index == 0:
|
||||
candidate_name = owner
|
||||
else:
|
||||
candidate_name = json.dumps(
|
||||
[owner, collision_index],
|
||||
ensure_ascii=False,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return str(uuid.uuid5(_DEFAULT_CALENDAR_NAMESPACE, candidate_name))
|
||||
|
||||
|
||||
def _begin_sqlite_default_write(db) -> None:
|
||||
"""Serialize an absent-default check with other SQLite writers.
|
||||
|
||||
SQLite's default deferred transactions allow two workers to both read an
|
||||
empty calendar set before either writes. ``BEGIN IMMEDIATE`` acquires the
|
||||
writer reservation before the second, authoritative lookup. We issue it
|
||||
only when the driver has not already opened a write transaction; a caller
|
||||
with a pending write already owns the required reservation.
|
||||
"""
|
||||
connection = db.connection()
|
||||
dbapi_connection = connection.connection
|
||||
driver_connection = getattr(
|
||||
dbapi_connection,
|
||||
"driver_connection",
|
||||
dbapi_connection,
|
||||
)
|
||||
if not getattr(driver_connection, "in_transaction", False):
|
||||
connection.exec_driver_sql("BEGIN IMMEDIATE")
|
||||
|
||||
|
||||
def _ensure_default_calendar(db, owner: str = None) -> CalendarCal:
|
||||
"""Create default calendar if none exist for this owner."""
|
||||
"""Return the owner's calendar, staging a default in the caller's transaction.
|
||||
|
||||
A stable owner-derived primary key makes concurrent first-use inserts
|
||||
converge on one row on every SQL backend. SQLite additionally serializes
|
||||
the absent-row check because its deferred transactions otherwise permit
|
||||
both workers to read the gap before either writes. Other backends recover
|
||||
a lost insert race inside a savepoint so the caller's event transaction
|
||||
remains usable and atomic.
|
||||
"""
|
||||
owner = owner or FALLBACK_OWNER
|
||||
cal = db.query(CalendarCal).filter(CalendarCal.owner == owner).first()
|
||||
if not cal:
|
||||
if cal:
|
||||
return cal
|
||||
|
||||
dialect = db.get_bind().dialect.name
|
||||
if dialect == "sqlite":
|
||||
_begin_sqlite_default_write(db)
|
||||
# Another worker may have committed while BEGIN IMMEDIATE waited.
|
||||
cal = db.query(CalendarCal).filter(CalendarCal.owner == owner).first()
|
||||
if cal:
|
||||
return cal
|
||||
|
||||
collision_index = 0
|
||||
while True:
|
||||
default_id = _default_calendar_id(owner, collision_index)
|
||||
|
||||
if dialect == "sqlite":
|
||||
# BEGIN IMMEDIATE above makes this occupancy check authoritative:
|
||||
# another SQLite writer cannot rename, delete, or claim this slot
|
||||
# until the caller commits or rolls back.
|
||||
occupant = db.query(CalendarCal).filter(
|
||||
CalendarCal.id == default_id,
|
||||
).first()
|
||||
if occupant is not None:
|
||||
if occupant.owner == owner:
|
||||
return occupant
|
||||
collision_index += 1
|
||||
continue
|
||||
|
||||
cal = CalendarCal(
|
||||
id=str(uuid.uuid4()),
|
||||
id=default_id,
|
||||
owner=owner,
|
||||
name="Personal",
|
||||
color="#5b8abf",
|
||||
source="local",
|
||||
)
|
||||
db.add(cal)
|
||||
db.commit()
|
||||
db.refresh(cal)
|
||||
return cal
|
||||
|
||||
if dialect == "sqlite":
|
||||
db.add(cal)
|
||||
db.flush()
|
||||
return cal
|
||||
|
||||
try:
|
||||
# A uniqueness failure rolls back only this savepoint, not an event
|
||||
# or reminder already staged by the caller's outer transaction.
|
||||
with db.begin_nested():
|
||||
db.add(cal)
|
||||
db.flush()
|
||||
return cal
|
||||
except IntegrityError:
|
||||
# Use a locking/current read so repeatable-read backends can observe
|
||||
# the row that won after our transaction's original empty snapshot.
|
||||
occupant = db.query(CalendarCal).filter(
|
||||
CalendarCal.id == default_id,
|
||||
).with_for_update().first()
|
||||
if occupant is None:
|
||||
# Do not misclassify an unrelated integrity failure as an ID
|
||||
# collision and loop forever. A concurrently deleted winner is
|
||||
# safe for the caller to retry as a fresh transaction.
|
||||
raise
|
||||
if occupant.owner == owner:
|
||||
return occupant
|
||||
# A renamed calendar owns this deterministic slot. Advance to the
|
||||
# next stable slot; concurrent callers for this owner will still
|
||||
# converge there.
|
||||
collision_index += 1
|
||||
|
||||
|
||||
# Per-request user time context. chat_routes sets this from browser timezone
|
||||
@@ -1015,6 +1119,9 @@ def setup_calendar_routes(upload_handler=None) -> APIRouter:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
_ensure_default_calendar(db, owner)
|
||||
# Listing calendars intentionally lazily creates a durable default.
|
||||
# Other callers commit it with the event they are creating.
|
||||
db.commit()
|
||||
cals = db.query(CalendarCal).filter(CalendarCal.owner == owner).all()
|
||||
return {"calendars": [
|
||||
{"name": c.name, "href": c.id, "color": c.color, "source": c.source}
|
||||
@@ -1023,6 +1130,7 @@ def setup_calendar_routes(upload_handler=None) -> APIRouter:
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.error("Failed to list calendars: %s", e)
|
||||
raise HTTPException(500, "Failed to list calendars")
|
||||
finally:
|
||||
|
||||
+47
-100
@@ -15,7 +15,7 @@ from core.database import Session as DBSession, ModelEndpoint
|
||||
from src.llm_core import normalize_model_id
|
||||
from src.endpoint_resolver import normalize_base
|
||||
from src.context_compactor import maybe_compact, trim_for_context
|
||||
from src.model_context import estimate_tokens
|
||||
from src.model_context import estimate_tokens, get_context_length
|
||||
from src.auth_helpers import effective_user
|
||||
from src.prompt_security import untrusted_context_message
|
||||
from src.attachment_refs import attachment_ref
|
||||
@@ -152,10 +152,38 @@ class ChatContext:
|
||||
# Uploads attached to this user turn, resolved and owner-checked for the
|
||||
# agent's private context. This is not emitted to the browser.
|
||||
uploaded_files: list = field(default_factory=list)
|
||||
# Route-neutral prompt before any model-window compaction/trimming. This is
|
||||
# retained only when explicit foreground fallbacks are enabled so each
|
||||
# concrete candidate can apply its own context budget independently.
|
||||
route_messages: list = field(default_factory=list)
|
||||
|
||||
|
||||
# ── Helpers ────────────────────────────────────────────────────────────── #
|
||||
|
||||
def _allowed_models_from_privileges(privs: dict) -> Optional[frozenset[str]]:
|
||||
if privs.get("block_all_models"):
|
||||
return frozenset()
|
||||
allowed_raw = privs.get("allowed_models")
|
||||
allowed = allowed_raw if isinstance(allowed_raw, list) else []
|
||||
restricted = bool(privs.get("allowed_models_restricted")) or bool(allowed)
|
||||
return frozenset(model for model in allowed if isinstance(model, str)) if restricted else None
|
||||
|
||||
|
||||
def _allowed_models_for_request(request) -> Optional[frozenset[str]]:
|
||||
"""Return the caller's model allowlist, or ``None`` when unrestricted."""
|
||||
|
||||
try:
|
||||
user = effective_user(request)
|
||||
except Exception:
|
||||
user = None
|
||||
if not user:
|
||||
return None
|
||||
auth_manager = getattr(getattr(request.app, "state", None), "auth_manager", None)
|
||||
if not auth_manager:
|
||||
return None
|
||||
privs = auth_manager.get_privileges(user) or {}
|
||||
return _allowed_models_from_privileges(privs)
|
||||
|
||||
def _enforce_chat_privileges(request, sess) -> None:
|
||||
"""Apply the per-user privilege gates (allowed_models + max_messages_per_day)
|
||||
that both /api/chat and /api/chat_stream must enforce BEFORE any LLM work.
|
||||
@@ -185,10 +213,8 @@ def _enforce_chat_privileges(request, sess) -> None:
|
||||
if privs.get("block_all_models"):
|
||||
raise HTTPException(403, f"Your account is not allowed to use model '{sess.model}'.")
|
||||
|
||||
allowed_raw = privs.get("allowed_models")
|
||||
allowed = allowed_raw if isinstance(allowed_raw, list) else []
|
||||
restricted = bool(privs.get("allowed_models_restricted")) or bool(allowed)
|
||||
if restricted and sess.model and sess.model not in allowed:
|
||||
allowed_models = _allowed_models_from_privileges(privs)
|
||||
if allowed_models is not None and sess.model and sess.model not in allowed_models:
|
||||
raise HTTPException(403, f"Your account is not allowed to use model '{sess.model}'.")
|
||||
|
||||
cap = int(privs.get("max_messages_per_day") or 0)
|
||||
@@ -287,96 +313,6 @@ async def auto_name_session(session_manager, sess):
|
||||
logger.error(f"Auto-name failed for {sess.id}: {e}\n{traceback.format_exc()}")
|
||||
|
||||
|
||||
def try_fallback_endpoint(sess, session_id: str) -> dict | None:
|
||||
"""Find an alternative working endpoint when the current one fails.
|
||||
|
||||
Returns {"model": ..., "endpoint_url": ..., "endpoint_name": ...} or None.
|
||||
"""
|
||||
import requests as _req
|
||||
from src.endpoint_resolver import (
|
||||
build_chat_url,
|
||||
build_headers,
|
||||
build_models_url,
|
||||
normalize_base,
|
||||
resolve_endpoint_runtime,
|
||||
)
|
||||
from src.chatgpt_subscription import is_chatgpt_subscription_base
|
||||
|
||||
current_url = sess.endpoint_url or ""
|
||||
owner = getattr(sess, "owner", None)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(ModelEndpoint).filter(
|
||||
ModelEndpoint.is_enabled == True
|
||||
)
|
||||
if owner:
|
||||
from src.auth_helpers import owner_filter
|
||||
q = owner_filter(q, ModelEndpoint, owner)
|
||||
endpoints = q.all()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
for ep in endpoints:
|
||||
base = normalize_base(ep.base_url)
|
||||
# Skip current endpoint
|
||||
if current_url and base in current_url:
|
||||
continue
|
||||
try:
|
||||
base, api_key = resolve_endpoint_runtime(ep, owner=owner)
|
||||
except Exception:
|
||||
continue
|
||||
ping_url = build_models_url(base)
|
||||
headers = build_headers(api_key, base)
|
||||
try:
|
||||
if ping_url:
|
||||
r = _req.get(ping_url, headers=headers, timeout=5)
|
||||
r.raise_for_status()
|
||||
data = r.json()
|
||||
models = [m.get("id") for m in (data.get("data") or []) if m.get("id")]
|
||||
if not models:
|
||||
models = [
|
||||
m.get("name") or m.get("model")
|
||||
for m in (data.get("models") or [])
|
||||
if m.get("name") or m.get("model")
|
||||
]
|
||||
else:
|
||||
models = json.loads(ep.cached_models or "[]")
|
||||
if not models:
|
||||
continue
|
||||
# Found a working endpoint — update session
|
||||
new_model = models[0]
|
||||
chat_url = build_chat_url(base)
|
||||
new_headers = build_headers(api_key, base)
|
||||
persisted_headers = {} if is_chatgpt_subscription_base(base) else new_headers
|
||||
|
||||
sess.model = new_model
|
||||
sess.endpoint_url = chat_url
|
||||
sess.headers = new_headers
|
||||
|
||||
# Persist
|
||||
_db = SessionLocal()
|
||||
try:
|
||||
_db.query(DBSession).filter(DBSession.id == session_id).update({
|
||||
"model": new_model,
|
||||
"endpoint_url": chat_url,
|
||||
"headers": persisted_headers,
|
||||
})
|
||||
_db.commit()
|
||||
finally:
|
||||
_db.close()
|
||||
|
||||
logger.info(f"Fallback: switched session {session_id} from {current_url} to {ep.name} ({new_model})")
|
||||
return {
|
||||
"model": new_model,
|
||||
"endpoint_url": chat_url,
|
||||
"endpoint_name": ep.name,
|
||||
}
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def extract_preset(chat_handler, preset_id) -> PresetInfo:
|
||||
"""Extract preset parameters via chat_handler."""
|
||||
temperature, max_tokens, system_prompt, char_name = (
|
||||
@@ -687,6 +623,7 @@ async def build_chat_context(
|
||||
use_enhanced_message: bool = False,
|
||||
agent_mode: bool = False,
|
||||
allow_tool_preprocessing: bool = True,
|
||||
defer_context_shaping: bool = False,
|
||||
) -> ChatContext:
|
||||
"""Build the full context (preface + messages) for an LLM call.
|
||||
|
||||
@@ -830,13 +767,22 @@ async def build_chat_context(
|
||||
except Exception:
|
||||
logger.debug("Failed to add current date/time context", exc_info=True)
|
||||
|
||||
# Auto-compact
|
||||
messages, context_length, was_compacted = await maybe_compact(
|
||||
sess, sess.endpoint_url, sess.model, messages, sess.headers, owner=user,
|
||||
)
|
||||
route_messages = list(messages)
|
||||
# Explicit fallback routing must shape from the same route-neutral prompt
|
||||
# for every candidate. Running selected-model compaction here would mutate
|
||||
# session history before we know which route can answer and would make a
|
||||
# later larger-context candidate unable to recover discarded history.
|
||||
if defer_context_shaping:
|
||||
context_length = get_context_length(sess.endpoint_url, sess.model)
|
||||
was_compacted = False
|
||||
else:
|
||||
messages, context_length, was_compacted = await maybe_compact(
|
||||
sess, sess.endpoint_url, sess.model, messages, sess.headers, owner=user,
|
||||
)
|
||||
_before_trim_messages = len(messages)
|
||||
_before_trim_tokens = estimate_tokens(messages)
|
||||
messages = trim_for_context(messages, context_length)
|
||||
if not defer_context_shaping:
|
||||
messages = trim_for_context(messages, context_length)
|
||||
_after_trim_messages = len(messages)
|
||||
_after_trim_tokens = estimate_tokens(messages)
|
||||
_context_trimmed = _after_trim_messages < _before_trim_messages or _after_trim_tokens < _before_trim_tokens
|
||||
@@ -860,6 +806,7 @@ async def build_chat_context(
|
||||
context_tokens_after_trim=_after_trim_tokens,
|
||||
auto_opened_docs=auto_opened_docs,
|
||||
uploaded_files=uploaded_files,
|
||||
route_messages=route_messages,
|
||||
)
|
||||
|
||||
|
||||
|
||||
+569
-45
@@ -15,12 +15,28 @@ from pydantic import ValidationError
|
||||
|
||||
from core.models import ChatMessage
|
||||
from src.request_models import ChatRequest
|
||||
from src.llm_core import llm_call_async, stream_llm, stream_llm_with_fallback
|
||||
from src.llm_core import (
|
||||
_normalize_http_status,
|
||||
llm_call_async,
|
||||
llm_call_async_with_route_fallback,
|
||||
stream_llm,
|
||||
stream_llm_with_fallback,
|
||||
)
|
||||
from src.agent_loop import stream_agent_loop
|
||||
from src import agent_runs
|
||||
from src.model_context import estimate_tokens
|
||||
from src.context_compactor import (
|
||||
apply_compaction_state,
|
||||
maybe_compact,
|
||||
trim_for_context,
|
||||
)
|
||||
from src.chat_helpers import coerce_message_and_session
|
||||
from src.endpoint_resolver import normalize_base as _normalize_base, build_chat_url
|
||||
from src.foreground_model_routing import (
|
||||
build_foreground_model_candidates,
|
||||
build_foreground_route_descriptors,
|
||||
resolve_foreground_model_policy,
|
||||
)
|
||||
from src.session_search import search_session_messages
|
||||
from src.prompt_security import untrusted_context_message
|
||||
from core.exceptions import SessionNotFoundError
|
||||
@@ -38,7 +54,9 @@ from routes.chat_helpers import (
|
||||
build_chat_context,
|
||||
save_assistant_response,
|
||||
run_post_response_tasks,
|
||||
accumulate_token_usage,
|
||||
clean_thinking_for_save,
|
||||
_allowed_models_for_request,
|
||||
_enforce_chat_privileges,
|
||||
)
|
||||
from src.action_intents import ToolIntent, classify_tool_intent as _classify_tool_intent
|
||||
@@ -56,6 +74,74 @@ logger = logging.getLogger(__name__)
|
||||
_active_streams: Dict[str, dict] = {}
|
||||
|
||||
|
||||
def _stream_failure_status(chunk: str) -> Optional[int]:
|
||||
"""Extract a provider status without retaining provider-supplied detail."""
|
||||
|
||||
try:
|
||||
for line in str(chunk or "").splitlines():
|
||||
if not line.startswith("data: "):
|
||||
continue
|
||||
status = json.loads(line[6:]).get("status")
|
||||
return _normalize_http_status(status)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _chat_candidate_request_factory(
|
||||
messages,
|
||||
fallback_context_length: int = 0,
|
||||
*,
|
||||
session=None,
|
||||
owner: Optional[str] = None,
|
||||
):
|
||||
"""Shape one route-neutral Chat prompt for each candidate window."""
|
||||
|
||||
state = {
|
||||
"requests": {},
|
||||
"context_lengths": {},
|
||||
"trim_stats": {},
|
||||
"compactions": {},
|
||||
"was_compacted": {},
|
||||
}
|
||||
|
||||
async def factory(index, candidate_url, candidate_model, candidate_headers):
|
||||
compaction_state = {}
|
||||
candidate_messages, context_length, was_compacted = await maybe_compact(
|
||||
session,
|
||||
candidate_url,
|
||||
candidate_model,
|
||||
list(messages),
|
||||
candidate_headers,
|
||||
owner=owner,
|
||||
persist=False,
|
||||
compaction_state=compaction_state,
|
||||
)
|
||||
if not context_length:
|
||||
context_length = fallback_context_length
|
||||
request_messages = trim_for_context(candidate_messages, context_length)
|
||||
state["requests"][index] = request_messages
|
||||
state["context_lengths"][index] = context_length
|
||||
state["compactions"][index] = compaction_state
|
||||
state["was_compacted"][index] = was_compacted
|
||||
state["trim_stats"][index] = {
|
||||
"messages_before": len(messages),
|
||||
"messages_after": len(request_messages),
|
||||
"tokens_before": estimate_tokens(messages),
|
||||
"tokens_after": estimate_tokens(request_messages),
|
||||
}
|
||||
return {"messages": request_messages}
|
||||
|
||||
return factory, state
|
||||
|
||||
|
||||
def _candidate_index(candidates, actual_candidate) -> int:
|
||||
for index, candidate in enumerate(candidates):
|
||||
if candidate == actual_candidate:
|
||||
return index
|
||||
return 0
|
||||
|
||||
|
||||
def _stream_set(session_id: str, **fields) -> None:
|
||||
"""Update fields on the active-stream entry for `session_id`, or
|
||||
no-op if the entry has already been popped. Using .get() avoids a
|
||||
@@ -589,8 +675,8 @@ def setup_chat_routes(
|
||||
# ------------------------------------------------------------------ #
|
||||
# POST /api/chat (non-streaming)
|
||||
# ------------------------------------------------------------------ #
|
||||
@router.post("/api/chat", response_model=Dict[str, str])
|
||||
async def chat_endpoint(request: Request, chat_request: ChatRequest) -> Dict[str, str]:
|
||||
@router.post("/api/chat", response_model=Dict[str, Any])
|
||||
async def chat_endpoint(request: Request, chat_request: ChatRequest) -> Dict[str, Any]:
|
||||
_set_user_time_from_request(request)
|
||||
|
||||
message = chat_request.message
|
||||
@@ -622,6 +708,8 @@ def setup_chat_routes(
|
||||
400,
|
||||
"No model selected for this chat. Open the model picker and choose one before sending.",
|
||||
)
|
||||
if not (getattr(sess, "endpoint_url", "") or "").strip():
|
||||
raise HTTPException(400, "Selected model endpoint is not configured")
|
||||
|
||||
# Same allowed_models + daily-cap gate as chat_stream (mirror so the
|
||||
# non-streaming path can't be used to bypass).
|
||||
@@ -637,6 +725,11 @@ def setup_chat_routes(
|
||||
if memory_response:
|
||||
return {"response": memory_response}
|
||||
|
||||
foreground_policy = resolve_foreground_model_policy(
|
||||
owner=owner,
|
||||
allowed_models=_allowed_models_for_request(request),
|
||||
)
|
||||
|
||||
# Build shared context (preset, preprocess, preface, compact)
|
||||
ctx = await build_chat_context(
|
||||
sess, request, chat_handler, chat_processor,
|
||||
@@ -648,6 +741,7 @@ def setup_chat_routes(
|
||||
time_filter=time_filter,
|
||||
webhook_manager=webhook_manager,
|
||||
allow_tool_preprocessing=allow_tool_preprocessing,
|
||||
defer_context_shaping=foreground_policy.enabled,
|
||||
)
|
||||
|
||||
# Research injection
|
||||
@@ -661,24 +755,88 @@ def setup_chat_routes(
|
||||
research_ctx = await research_handler.call_research_service(
|
||||
message, _r_ep, _r_model, llm_headers=_r_headers
|
||||
)
|
||||
ctx.messages.insert(
|
||||
len(ctx.preface),
|
||||
untrusted_context_message("research context", research_ctx),
|
||||
)
|
||||
research_message = untrusted_context_message("research context", research_ctx)
|
||||
ctx.messages.insert(len(ctx.preface), research_message)
|
||||
if foreground_policy.enabled:
|
||||
getattr(ctx, "route_messages", ctx.messages).insert(
|
||||
len(ctx.preface),
|
||||
research_message,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Research failed: {e}")
|
||||
|
||||
reply = await llm_call_async(
|
||||
foreground_candidates = build_foreground_model_candidates(
|
||||
sess.endpoint_url,
|
||||
sess.model,
|
||||
ctx.messages,
|
||||
headers=sess.headers,
|
||||
sess.headers,
|
||||
owner=owner,
|
||||
policy=foreground_policy,
|
||||
)
|
||||
route_descriptors = build_foreground_route_descriptors(
|
||||
sess.endpoint_url,
|
||||
sess.model,
|
||||
sess.headers,
|
||||
owner=owner,
|
||||
policy=foreground_policy,
|
||||
selected_endpoint_id=chat_request.selected_endpoint_id,
|
||||
)
|
||||
candidate_request_factory = None
|
||||
selected_context_length = getattr(ctx, "context_length", 0)
|
||||
candidate_request_state = {
|
||||
"context_lengths": {0: selected_context_length},
|
||||
"requests": {0: ctx.messages},
|
||||
"trim_stats": {},
|
||||
}
|
||||
request_messages = ctx.messages
|
||||
if foreground_policy.enabled:
|
||||
request_messages = getattr(ctx, "route_messages", ctx.messages)
|
||||
candidate_request_factory, candidate_request_state = _chat_candidate_request_factory(
|
||||
request_messages,
|
||||
selected_context_length,
|
||||
session=sess,
|
||||
owner=owner,
|
||||
)
|
||||
requested_model = sess.model
|
||||
reply, actual_candidate, actual_model = await llm_call_async_with_route_fallback(
|
||||
foreground_candidates,
|
||||
request_messages,
|
||||
fallback_statuses=foreground_policy.eligible_statuses,
|
||||
candidate_request_factory=candidate_request_factory,
|
||||
temperature=ctx.preset.temperature,
|
||||
max_tokens=ctx.preset.max_tokens,
|
||||
prompt_type=preset_id,
|
||||
session_id=session,
|
||||
)
|
||||
_clean_reply, _clean_md = clean_thinking_for_save(reply, {"model": sess.model})
|
||||
actual_index = _candidate_index(foreground_candidates, actual_candidate)
|
||||
apply_compaction_state(
|
||||
sess,
|
||||
candidate_request_state.get("compactions", {}).get(actual_index),
|
||||
)
|
||||
requested_route = route_descriptors[0]
|
||||
actual_route = route_descriptors[actual_index]
|
||||
actual_trim = candidate_request_state.get("trim_stats", {}).get(actual_index, {})
|
||||
_clean_reply, _clean_md = clean_thinking_for_save(
|
||||
reply,
|
||||
{
|
||||
"model": actual_model,
|
||||
"requested_model": requested_model,
|
||||
"endpoint_id": actual_route.get("endpoint_id"),
|
||||
"endpoint_label": actual_route.get("endpoint_label"),
|
||||
"requested_endpoint_id": requested_route.get("endpoint_id"),
|
||||
"requested_endpoint_label": requested_route.get("endpoint_label"),
|
||||
"context_length": candidate_request_state["context_lengths"].get(
|
||||
actual_index,
|
||||
selected_context_length,
|
||||
),
|
||||
"context_trimmed": bool(
|
||||
actual_trim
|
||||
and (
|
||||
actual_trim.get("messages_after") < actual_trim.get("messages_before")
|
||||
or actual_trim.get("tokens_after") < actual_trim.get("tokens_before")
|
||||
)
|
||||
),
|
||||
},
|
||||
)
|
||||
sess.add_message(ChatMessage("assistant", _clean_reply, metadata=_clean_md))
|
||||
|
||||
from core.database import update_session_last_accessed
|
||||
@@ -694,7 +852,15 @@ def setup_chat_routes(
|
||||
allow_background_extraction=not tool_policy.block_all_tool_calls,
|
||||
)
|
||||
|
||||
return {"response": reply}
|
||||
return {
|
||||
"response": reply,
|
||||
"requested_model": requested_model,
|
||||
"model": actual_model,
|
||||
"requested_endpoint_id": requested_route.get("endpoint_id"),
|
||||
"requested_endpoint_label": requested_route.get("endpoint_label"),
|
||||
"endpoint_id": actual_route.get("endpoint_id"),
|
||||
"endpoint_label": actual_route.get("endpoint_label"),
|
||||
}
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# POST /api/chat_stream
|
||||
@@ -723,6 +889,11 @@ def setup_chat_routes(
|
||||
use_research = form_data.get("use_research")
|
||||
time_filter = form_data.get("time_filter")
|
||||
preset_id = form_data.get("preset_id")
|
||||
selected_endpoint_id = str(
|
||||
form_data.get("selected_endpoint_id")
|
||||
or (body or {}).get("selected_endpoint_id")
|
||||
or ""
|
||||
).strip()
|
||||
# Issue #3229: API callers send JSON, not FormData. Read from the
|
||||
# JSON body as fallback so callers who send {"allow_bash": true}
|
||||
# actually get bash enabled.
|
||||
@@ -895,6 +1066,8 @@ def setup_chat_routes(
|
||||
400,
|
||||
"No model selected for this chat. Open the model picker and choose one before sending.",
|
||||
)
|
||||
if not (getattr(sess, "endpoint_url", "") or "").strip():
|
||||
raise HTTPException(400, "Selected model endpoint is not configured")
|
||||
if (
|
||||
chat_mode == "chat"
|
||||
and isinstance(message, str)
|
||||
@@ -970,6 +1143,10 @@ def setup_chat_routes(
|
||||
last_user_message=message,
|
||||
)
|
||||
allow_tool_preprocessing = not pre_context_tool_policy.block_all_tool_calls
|
||||
foreground_policy = resolve_foreground_model_policy(
|
||||
owner=owner,
|
||||
allowed_models=_allowed_models_for_request(request),
|
||||
)
|
||||
|
||||
# Build shared context (stream path uses enhanced_message for context preface)
|
||||
ctx = await build_chat_context(
|
||||
@@ -992,6 +1169,7 @@ def setup_chat_routes(
|
||||
# index would be useless / unwanted noise.
|
||||
agent_mode=(chat_mode == "agent"),
|
||||
allow_tool_preprocessing=allow_tool_preprocessing,
|
||||
defer_context_shaping=foreground_policy.enabled,
|
||||
)
|
||||
|
||||
_research_flags = {"do": do_research} # Mutable container for generator scope
|
||||
@@ -1291,6 +1469,8 @@ def setup_chat_routes(
|
||||
"what aspects matter most, are they comparing to something, what's their context "
|
||||
"(moving, traveling, curiosity). Be conversational. Keep it short."
|
||||
})
|
||||
if foreground_policy.enabled:
|
||||
getattr(ctx, "route_messages", ctx.messages).insert(0, dict(ctx.messages[0]))
|
||||
_skip_research = True
|
||||
else:
|
||||
_skip_research = False
|
||||
@@ -1387,7 +1567,12 @@ def setup_chat_routes(
|
||||
_active_streams.pop(session, None)
|
||||
return
|
||||
|
||||
messages = _ensure_current_request_is_latest_user(ctx.messages, message)
|
||||
context_source = (
|
||||
getattr(ctx, "route_messages", ctx.messages)
|
||||
if foreground_policy.enabled
|
||||
else ctx.messages
|
||||
)
|
||||
messages = _ensure_current_request_is_latest_user(context_source, message)
|
||||
|
||||
# Auto-compact notification
|
||||
if ctx.was_compacted:
|
||||
@@ -1399,25 +1584,56 @@ def setup_chat_routes(
|
||||
thinking_response = ""
|
||||
last_metrics = None
|
||||
|
||||
# Configured fallback chain for the default chat model. Tried in
|
||||
# order if the session's primary model fails before producing
|
||||
# output. Resolved once per request.
|
||||
try:
|
||||
from src.endpoint_resolver import resolve_chat_fallback_candidates
|
||||
_fallback_candidates = resolve_chat_fallback_candidates(owner=_user)
|
||||
except Exception:
|
||||
_fallback_candidates = []
|
||||
# Foreground Chat and Agent requests share one explicit owner-aware
|
||||
# policy. Strict mode is the default; legacy values are unrelated.
|
||||
_foreground_policy = foreground_policy
|
||||
_foreground_candidates = build_foreground_model_candidates(
|
||||
sess.endpoint_url,
|
||||
sess.model,
|
||||
sess.headers,
|
||||
owner=_user,
|
||||
policy=_foreground_policy,
|
||||
)
|
||||
_foreground_route_descriptors = build_foreground_route_descriptors(
|
||||
sess.endpoint_url,
|
||||
sess.model,
|
||||
sess.headers,
|
||||
owner=_user,
|
||||
policy=_foreground_policy,
|
||||
selected_endpoint_id=selected_endpoint_id,
|
||||
)
|
||||
_chat_request_factory = None
|
||||
_selected_context_length = getattr(ctx, "context_length", 0)
|
||||
_chat_request_state = {
|
||||
"context_lengths": {0: _selected_context_length},
|
||||
"requests": {0: messages},
|
||||
"trim_stats": {},
|
||||
}
|
||||
if _foreground_policy.enabled:
|
||||
_chat_request_factory, _chat_request_state = _chat_candidate_request_factory(
|
||||
messages,
|
||||
_selected_context_length,
|
||||
session=sess,
|
||||
owner=_user,
|
||||
)
|
||||
|
||||
# Send model name early so the frontend can show it during streaming
|
||||
_model_suffix = "Research" if effective_do_research else None
|
||||
_model_info = {"type": "model_info", "model": sess.model}
|
||||
_selected_route = _foreground_route_descriptors[0]
|
||||
_model_info = {
|
||||
"type": "model_info",
|
||||
"model": sess.model,
|
||||
"endpoint_id": _selected_route.get("endpoint_id"),
|
||||
"endpoint_label": _selected_route.get("endpoint_label"),
|
||||
}
|
||||
if _model_suffix:
|
||||
_model_info["suffix"] = _model_suffix
|
||||
if ctx.preset.character_name:
|
||||
_model_info["character_name"] = ctx.preset.character_name
|
||||
yield f'data: {json.dumps(_model_info)}\n\n'
|
||||
|
||||
if image_generation_session:
|
||||
_terminal_saved = False
|
||||
if _is_image_generation_session(sess, owner=_user):
|
||||
from src.settings import get_setting
|
||||
if tool_policy.blocks("generate_image"):
|
||||
_blocked_msg = tool_policy.reason_for("generate_image")
|
||||
@@ -1520,11 +1736,20 @@ def setup_chat_routes(
|
||||
_answered_by = None # set if the selected model failed and a fallback answered
|
||||
_requested_model = sess.model
|
||||
_actual_model = None
|
||||
_requested_route = _foreground_route_descriptors[0]
|
||||
_actual_route = _requested_route
|
||||
_actual_candidate_index = 0
|
||||
_chat_terminal_saved = False
|
||||
def _commit_chat_compaction(candidate_index: int) -> bool:
|
||||
return apply_compaction_state(
|
||||
sess,
|
||||
_chat_request_state.get("compactions", {}).get(candidate_index),
|
||||
)
|
||||
|
||||
# ── Chat mode: call stream_llm directly, NO tools, NO document access ──
|
||||
try:
|
||||
_chat_candidates = [(sess.endpoint_url, sess.model, sess.headers)] + _fallback_candidates
|
||||
async for chunk in stream_llm_with_fallback(
|
||||
_chat_candidates,
|
||||
_foreground_candidates,
|
||||
messages,
|
||||
temperature=ctx.preset.temperature,
|
||||
# Respect the preset; 0/unset = let the server decide (no
|
||||
@@ -1536,11 +1761,21 @@ def setup_chat_routes(
|
||||
prompt_type=preset_id,
|
||||
tools=None,
|
||||
session_id=session,
|
||||
fallback_statuses=_foreground_policy.eligible_statuses,
|
||||
fallback_on_empty=_foreground_policy.fallback_on_empty,
|
||||
candidate_request_factory=_chat_request_factory,
|
||||
candidate_route_descriptors=_foreground_route_descriptors,
|
||||
):
|
||||
if chunk.startswith("data: ") and not chunk.startswith("data: [DONE]"):
|
||||
try:
|
||||
data = json.loads(chunk[6:])
|
||||
if "delta" in data:
|
||||
if _commit_chat_compaction(_actual_candidate_index):
|
||||
_compacted_length = _chat_request_state["context_lengths"].get(
|
||||
_actual_candidate_index,
|
||||
_selected_context_length,
|
||||
)
|
||||
yield f'data: {json.dumps({"type": "compacted", "context_length": _compacted_length})}\n\n'
|
||||
# Reasoning tokens arrive flagged thinking:true.
|
||||
# Forward them so the client can show a thinking
|
||||
# indicator, but don't fold them into the saved
|
||||
@@ -1556,29 +1791,82 @@ def setup_chat_routes(
|
||||
# Forward the notice and remember the real model.
|
||||
_answered_by = data.get("answered_by") or _answered_by
|
||||
_actual_model = _actual_model or _answered_by
|
||||
_actual_candidate_index = data.get("candidate_index", 0)
|
||||
if not isinstance(_actual_candidate_index, int):
|
||||
_actual_candidate_index = 0
|
||||
if 0 <= _actual_candidate_index < len(_foreground_route_descriptors):
|
||||
_actual_route = _foreground_route_descriptors[_actual_candidate_index]
|
||||
if _commit_chat_compaction(_actual_candidate_index):
|
||||
_compacted_length = _chat_request_state["context_lengths"].get(
|
||||
_actual_candidate_index,
|
||||
_selected_context_length,
|
||||
)
|
||||
yield f'data: {json.dumps({"type": "compacted", "context_length": _compacted_length})}\n\n'
|
||||
data["selected_model"] = data.get("selected_model") or _requested_model
|
||||
yield chunk
|
||||
yield f'data: {json.dumps(data)}\n\n'
|
||||
elif data.get("type") == "model_actual":
|
||||
if _commit_chat_compaction(_actual_candidate_index):
|
||||
_compacted_length = _chat_request_state["context_lengths"].get(
|
||||
_actual_candidate_index,
|
||||
_selected_context_length,
|
||||
)
|
||||
yield f'data: {json.dumps({"type": "compacted", "context_length": _compacted_length})}\n\n'
|
||||
_actual_model = data.get("model") or _actual_model
|
||||
data["requested_model"] = _requested_model
|
||||
data["requested_endpoint_id"] = _requested_route.get("endpoint_id")
|
||||
data["requested_endpoint_label"] = _requested_route.get("endpoint_label")
|
||||
data["endpoint_id"] = _actual_route.get("endpoint_id")
|
||||
data["endpoint_label"] = _actual_route.get("endpoint_label")
|
||||
yield f'data: {json.dumps(data)}\n\n'
|
||||
elif data.get("type") == "usage":
|
||||
if _commit_chat_compaction(_actual_candidate_index):
|
||||
_compacted_length = _chat_request_state["context_lengths"].get(
|
||||
_actual_candidate_index,
|
||||
_selected_context_length,
|
||||
)
|
||||
yield f'data: {json.dumps({"type": "compacted", "context_length": _compacted_length})}\n\n'
|
||||
last_metrics = data.get("data", {})
|
||||
_reported_model = last_metrics.get("model")
|
||||
last_metrics["requested_model"] = _requested_model
|
||||
last_metrics["model"] = _reported_model or _actual_model or _answered_by or _requested_model
|
||||
if ctx.context_trimmed:
|
||||
last_metrics["requested_endpoint_id"] = _requested_route.get("endpoint_id")
|
||||
last_metrics["requested_endpoint_label"] = _requested_route.get("endpoint_label")
|
||||
last_metrics["endpoint_id"] = _actual_route.get("endpoint_id")
|
||||
last_metrics["endpoint_label"] = _actual_route.get("endpoint_label")
|
||||
if isinstance(
|
||||
_actual_route.get("endpoint_cost_tracked"),
|
||||
bool,
|
||||
):
|
||||
last_metrics["endpoint_cost_tracked"] = _actual_route.get(
|
||||
"endpoint_cost_tracked"
|
||||
)
|
||||
_actual_context_length = _chat_request_state["context_lengths"].get(
|
||||
_actual_candidate_index,
|
||||
_selected_context_length,
|
||||
)
|
||||
_route_trim = _chat_request_state.get("trim_stats", {}).get(
|
||||
_actual_candidate_index,
|
||||
{},
|
||||
)
|
||||
if _route_trim and (
|
||||
_route_trim.get("messages_after") < _route_trim.get("messages_before")
|
||||
or _route_trim.get("tokens_after") < _route_trim.get("tokens_before")
|
||||
):
|
||||
last_metrics["context_trimmed"] = True
|
||||
last_metrics["context_messages_before_trim"] = _route_trim.get("messages_before")
|
||||
last_metrics["context_messages_after_trim"] = _route_trim.get("messages_after")
|
||||
last_metrics["context_tokens_before_trim"] = _route_trim.get("tokens_before")
|
||||
last_metrics["context_tokens_after_trim"] = _route_trim.get("tokens_after")
|
||||
elif ctx.context_trimmed:
|
||||
last_metrics["context_trimmed"] = True
|
||||
last_metrics["context_messages_before_trim"] = ctx.context_messages_before_trim
|
||||
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
|
||||
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)
|
||||
if _actual_context_length and last_metrics.get("input_tokens"):
|
||||
pct = min(round((last_metrics["input_tokens"] / _actual_context_length) * 100, 1), 100.0)
|
||||
last_metrics["context_percent"] = pct
|
||||
last_metrics["context_length"] = ctx.context_length
|
||||
last_metrics["context_length"] = _actual_context_length
|
||||
# The frontend reads `tokens_per_second`; the raw usage event
|
||||
# carries the backend's true gen speed as `gen_tps` (llama.cpp
|
||||
# timings). Map it through so this direct-chat path shows real
|
||||
@@ -1593,17 +1881,121 @@ def setup_chat_routes(
|
||||
yield chunk
|
||||
elif chunk.startswith("event: error"):
|
||||
logger.warning(f"Stream error for {sess.model} on {sess.endpoint_url}: {chunk!r}")
|
||||
if (
|
||||
not _chat_terminal_saved
|
||||
and (full_response.strip() or thinking_response.strip())
|
||||
):
|
||||
_failure_status = _stream_failure_status(chunk)
|
||||
_failure_message = (
|
||||
f"Model request failed (HTTP {_failure_status})"
|
||||
if _failure_status is not None
|
||||
else "Model request failed"
|
||||
)
|
||||
_terminal_content = full_response.strip()
|
||||
_failure_note = f"[Response stopped: {_failure_message}]"
|
||||
_terminal_content = (
|
||||
f"{_terminal_content}\n\n{_failure_note}"
|
||||
if _terminal_content
|
||||
else _failure_note
|
||||
)
|
||||
_had_terminal_usage = bool(last_metrics)
|
||||
_terminal_metrics = dict(last_metrics or {})
|
||||
if not _had_terminal_usage:
|
||||
_actual_request_messages = _chat_request_state["requests"].get(
|
||||
_actual_candidate_index,
|
||||
messages,
|
||||
)
|
||||
_actual_context_length = _chat_request_state["context_lengths"].get(
|
||||
_actual_candidate_index,
|
||||
_selected_context_length,
|
||||
)
|
||||
_estimated_input = estimate_tokens(_actual_request_messages)
|
||||
_estimated_output = max(
|
||||
len(full_response + thinking_response) // 4,
|
||||
0,
|
||||
)
|
||||
_terminal_metrics.update({
|
||||
"input_tokens": _estimated_input,
|
||||
"output_tokens": _estimated_output,
|
||||
"total_tokens": _estimated_input + _estimated_output,
|
||||
"usage_source": "estimated",
|
||||
"response_time": round(time.time() - _chat_start, 2),
|
||||
"context_length": _actual_context_length,
|
||||
"context_percent": (
|
||||
min(
|
||||
round(
|
||||
(_estimated_input / _actual_context_length) * 100,
|
||||
1,
|
||||
),
|
||||
100.0,
|
||||
)
|
||||
if _actual_context_length
|
||||
else 0
|
||||
),
|
||||
})
|
||||
_terminal_metrics.update({
|
||||
"failed": True,
|
||||
"failure": {
|
||||
"status": _failure_status,
|
||||
"message": _failure_message,
|
||||
},
|
||||
"model": _actual_model or _answered_by or _requested_model,
|
||||
"requested_model": _requested_model,
|
||||
"endpoint_id": _actual_route.get("endpoint_id"),
|
||||
"endpoint_label": _actual_route.get("endpoint_label"),
|
||||
"requested_endpoint_id": _requested_route.get("endpoint_id"),
|
||||
"requested_endpoint_label": _requested_route.get("endpoint_label"),
|
||||
})
|
||||
if isinstance(
|
||||
_actual_route.get("endpoint_cost_tracked"),
|
||||
bool,
|
||||
):
|
||||
_terminal_metrics["endpoint_cost_tracked"] = _actual_route.get(
|
||||
"endpoint_cost_tracked"
|
||||
)
|
||||
if thinking_response.strip():
|
||||
_terminal_metrics["thinking"] = thinking_response.strip()
|
||||
_commit_chat_compaction(_actual_candidate_index)
|
||||
_saved_id = save_assistant_response(
|
||||
sess,
|
||||
session_manager,
|
||||
session,
|
||||
_terminal_content,
|
||||
_terminal_metrics,
|
||||
character_name=ctx.preset.character_name,
|
||||
incognito=incognito,
|
||||
)
|
||||
accumulate_token_usage(session, _terminal_metrics)
|
||||
_chat_terminal_saved = True
|
||||
_stream_set(session, status="error")
|
||||
if _saved_id:
|
||||
yield f'data: {json.dumps({"type": "message_saved", "id": _saved_id})}\n\n'
|
||||
yield f'data: {json.dumps({"type": "chat_terminal", "data": _terminal_metrics})}\n\n'
|
||||
yield chunk
|
||||
elif chunk.startswith("event: "):
|
||||
yield chunk
|
||||
elif chunk == "data: [DONE]\n\n":
|
||||
if _chat_terminal_saved:
|
||||
# Some providers append DONE after a terminal
|
||||
# error. The failed partial is already saved;
|
||||
# never re-save/post-process it as a success or
|
||||
# advertise successful completion to the client.
|
||||
continue
|
||||
# Generate fallback metrics if LLM didn't send usage
|
||||
if not last_metrics and full_response:
|
||||
_elapsed = time.time() - _chat_start
|
||||
_est_in = estimate_tokens(messages)
|
||||
_est_out = len(full_response) // 4
|
||||
_tps = round(_est_out / _elapsed, 2) if _elapsed > 0 else 0
|
||||
_ctx_pct = min(round((_est_in / ctx.context_length) * 100, 1), 100.0) if ctx.context_length else 0
|
||||
_actual_context_length = _chat_request_state["context_lengths"].get(
|
||||
_actual_candidate_index,
|
||||
_selected_context_length,
|
||||
)
|
||||
_actual_request_messages = _chat_request_state["requests"].get(
|
||||
_actual_candidate_index,
|
||||
messages,
|
||||
)
|
||||
_est_in = estimate_tokens(_actual_request_messages)
|
||||
_ctx_pct = min(round((_est_in / _actual_context_length) * 100, 1), 100.0) if _actual_context_length else 0
|
||||
last_metrics = {
|
||||
"response_time": round(_elapsed, 2),
|
||||
"input_tokens": _est_in,
|
||||
@@ -1611,13 +2003,25 @@ def setup_chat_routes(
|
||||
"tokens_per_second": _tps,
|
||||
"request_context_tokens": _est_in,
|
||||
"context_percent": _ctx_pct,
|
||||
"context_length": ctx.context_length,
|
||||
"context_length": _actual_context_length,
|
||||
"model": _actual_model or _answered_by or _requested_model,
|
||||
"requested_model": _requested_model,
|
||||
"requested_endpoint_id": _requested_route.get("endpoint_id"),
|
||||
"requested_endpoint_label": _requested_route.get("endpoint_label"),
|
||||
"endpoint_id": _actual_route.get("endpoint_id"),
|
||||
"endpoint_label": _actual_route.get("endpoint_label"),
|
||||
"usage_source": "estimated",
|
||||
}
|
||||
if isinstance(
|
||||
_actual_route.get("endpoint_cost_tracked"),
|
||||
bool,
|
||||
):
|
||||
last_metrics["endpoint_cost_tracked"] = _actual_route.get(
|
||||
"endpoint_cost_tracked"
|
||||
)
|
||||
yield f'data: {json.dumps({"type": "metrics", "data": last_metrics})}\n\n'
|
||||
if full_response:
|
||||
_commit_chat_compaction(_actual_candidate_index)
|
||||
_metrics_to_save = dict(last_metrics or {})
|
||||
if thinking_response.strip() and not _metrics_to_save.get("thinking"):
|
||||
_metrics_to_save["thinking"] = thinking_response.strip()
|
||||
@@ -1652,6 +2056,10 @@ def setup_chat_routes(
|
||||
"stopped": True,
|
||||
"model": _actual_model or _answered_by or _requested_model,
|
||||
"requested_model": _requested_model,
|
||||
"endpoint_id": _actual_route.get("endpoint_id"),
|
||||
"endpoint_label": _actual_route.get("endpoint_label"),
|
||||
"requested_endpoint_id": _requested_route.get("endpoint_id"),
|
||||
"requested_endpoint_label": _requested_route.get("endpoint_label"),
|
||||
},
|
||||
)
|
||||
sess.add_message(ChatMessage("assistant", _stopped_content, metadata=_stopped_md))
|
||||
@@ -1666,6 +2074,12 @@ def setup_chat_routes(
|
||||
_answered_by = None # set if the selected model failed and a fallback answered
|
||||
_requested_model = sess.model
|
||||
_actual_model = None
|
||||
_agent_requested_route = _foreground_route_descriptors[0]
|
||||
_agent_actual_endpoint_id = _agent_requested_route.get("endpoint_id")
|
||||
_agent_actual_endpoint_label = _agent_requested_route.get("endpoint_label")
|
||||
_agent_round_models = {1: _requested_model}
|
||||
_agent_round_endpoint_ids = {1: _agent_actual_endpoint_id}
|
||||
_agent_round_endpoint_labels = {1: _agent_actual_endpoint_label}
|
||||
try:
|
||||
from src.settings import get_setting
|
||||
from src.agent_tools import MAX_AGENT_ROUNDS as _DEFAULT_ROUNDS
|
||||
@@ -1703,19 +2117,24 @@ def setup_chat_routes(
|
||||
prompt_type=preset_id,
|
||||
max_tool_calls=_tool_budget,
|
||||
max_rounds=_max_rounds,
|
||||
context_length=ctx.context_length,
|
||||
context_length=_selected_context_length,
|
||||
active_document=active_doc,
|
||||
active_email=active_email_ctx,
|
||||
session_id=session,
|
||||
history_session=sess,
|
||||
disabled_tools=disabled_tools if disabled_tools else None,
|
||||
tool_policy=tool_policy,
|
||||
owner=_user,
|
||||
fallbacks=_fallback_candidates,
|
||||
fallbacks=_foreground_candidates[1:],
|
||||
route_descriptors=_foreground_route_descriptors,
|
||||
fallback_statuses=_foreground_policy.eligible_statuses,
|
||||
fallback_on_empty=_foreground_policy.fallback_on_empty,
|
||||
plan_mode=plan_mode,
|
||||
approved_plan=approved_plan or None,
|
||||
workspace=workspace or None,
|
||||
forced_tools=_forced_tools,
|
||||
uploaded_files=ctx.uploaded_files,
|
||||
defer_context_shaping=_foreground_policy.enabled,
|
||||
):
|
||||
if chunk.startswith("data: ") and not chunk.startswith("data: [DONE]"):
|
||||
try:
|
||||
@@ -1744,7 +2163,20 @@ def setup_chat_routes(
|
||||
"plan_update",
|
||||
):
|
||||
if data.get("type") == "agent_step":
|
||||
_agent_rounds = max(_agent_rounds, data.get("round", 1))
|
||||
_event_round = data.get("round", 1)
|
||||
_agent_rounds = max(_agent_rounds, _event_round)
|
||||
_agent_round_models.setdefault(
|
||||
_event_round,
|
||||
_actual_model or _answered_by or _requested_model,
|
||||
)
|
||||
_agent_round_endpoint_ids.setdefault(
|
||||
_event_round,
|
||||
_agent_actual_endpoint_id,
|
||||
)
|
||||
_agent_round_endpoint_labels.setdefault(
|
||||
_event_round,
|
||||
_agent_actual_endpoint_label,
|
||||
)
|
||||
elif data.get("type") == "tool_start":
|
||||
_agent_tool_calls += 1
|
||||
yield chunk
|
||||
@@ -1754,13 +2186,70 @@ def setup_chat_routes(
|
||||
# model so metrics reflect it, not the masked
|
||||
# selected model.
|
||||
_answered_by = data.get("answered_by") or _answered_by
|
||||
_actual_model = _actual_model or _answered_by
|
||||
_actual_model = _answered_by or _actual_model
|
||||
if "answered_by_endpoint_id" in data:
|
||||
_agent_actual_endpoint_id = data.get("answered_by_endpoint_id")
|
||||
if data.get("answered_by_endpoint_label"):
|
||||
_agent_actual_endpoint_label = data.get("answered_by_endpoint_label")
|
||||
_event_round = data.get("round") or max(_agent_rounds, 1)
|
||||
_agent_round_models[_event_round] = _answered_by or _requested_model
|
||||
_agent_round_endpoint_ids[_event_round] = _agent_actual_endpoint_id
|
||||
_agent_round_endpoint_labels[_event_round] = _agent_actual_endpoint_label
|
||||
data["selected_model"] = data.get("selected_model") or _requested_model
|
||||
yield chunk
|
||||
elif data.get("type") == "model_actual":
|
||||
_actual_model = data.get("model") or _actual_model
|
||||
if "endpoint_id" in data:
|
||||
_agent_actual_endpoint_id = data.get("endpoint_id")
|
||||
if data.get("endpoint_label"):
|
||||
_agent_actual_endpoint_label = data.get("endpoint_label")
|
||||
_event_round = data.get("round") or max(_agent_rounds, 1)
|
||||
_agent_round_models[_event_round] = _actual_model or _requested_model
|
||||
_agent_round_endpoint_ids[_event_round] = _agent_actual_endpoint_id
|
||||
_agent_round_endpoint_labels[_event_round] = _agent_actual_endpoint_label
|
||||
data["requested_model"] = _requested_model
|
||||
yield f'data: {json.dumps(data)}\n\n'
|
||||
elif data.get("type") == "agent_terminal":
|
||||
terminal_metadata = dict(data.get("data") or {})
|
||||
last_metrics = terminal_metadata
|
||||
failure = terminal_metadata.get("failure") or {}
|
||||
failure_status = _normalize_http_status(
|
||||
failure.get("status")
|
||||
)
|
||||
failure_message = (
|
||||
f"Model request failed (HTTP {failure_status})"
|
||||
if failure_status is not None
|
||||
else "Model request failed"
|
||||
)
|
||||
terminal_metadata["failure"] = {
|
||||
"status": failure_status,
|
||||
"message": failure_message,
|
||||
}
|
||||
terminal_content = full_response.strip()
|
||||
failure_note = f"[Agent stopped: {failure_message}]"
|
||||
if terminal_content:
|
||||
terminal_content = f"{terminal_content}\n\n{failure_note}"
|
||||
else:
|
||||
terminal_content = failure_note
|
||||
if not _terminal_saved:
|
||||
_saved_id = save_assistant_response(
|
||||
sess,
|
||||
session_manager,
|
||||
session,
|
||||
terminal_content,
|
||||
terminal_metadata,
|
||||
character_name=ctx.preset.character_name,
|
||||
web_sources=web_sources,
|
||||
rag_sources=ctx.rag_sources,
|
||||
used_memories=ctx.used_memories,
|
||||
incognito=incognito,
|
||||
)
|
||||
_terminal_saved = True
|
||||
accumulate_token_usage(session, terminal_metadata)
|
||||
_stream_set(session, status="error")
|
||||
if _saved_id:
|
||||
yield f'data: {json.dumps({"type": "message_saved", "id": _saved_id})}\n\n'
|
||||
yield chunk
|
||||
elif data.get("type") == "metrics":
|
||||
last_metrics = data.get("data", {})
|
||||
_reported_model = last_metrics.get("model")
|
||||
@@ -1772,7 +2261,16 @@ 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
|
||||
yield f'data: {json.dumps({"type": "metrics", "data": last_metrics})}\n\n'
|
||||
_metrics_event = {"type": "metrics", "data": last_metrics}
|
||||
# Inline teacher escalation marks its
|
||||
# recursively emitted events at the SSE
|
||||
# envelope. Preserve that non-secret marker
|
||||
# when normalizing metrics so the browser's
|
||||
# replay-stable ledger keeps primary and
|
||||
# teacher segments distinct.
|
||||
if data.get("teacher") is True:
|
||||
_metrics_event["teacher"] = True
|
||||
yield f'data: {json.dumps(_metrics_event)}\n\n'
|
||||
except json.JSONDecodeError:
|
||||
yield chunk
|
||||
elif chunk.startswith("event: "):
|
||||
@@ -1824,6 +2322,22 @@ def setup_chat_routes(
|
||||
"stopped": True,
|
||||
"model": _actual_model or _answered_by or _requested_model,
|
||||
"requested_model": _requested_model,
|
||||
"endpoint_id": _agent_actual_endpoint_id,
|
||||
"endpoint_label": _agent_actual_endpoint_label,
|
||||
"requested_endpoint_id": _agent_requested_route.get("endpoint_id"),
|
||||
"requested_endpoint_label": _agent_requested_route.get("endpoint_label"),
|
||||
"round_models": [
|
||||
_agent_round_models.get(i, _actual_model or _requested_model)
|
||||
for i in range(1, max(_agent_round_models, default=1) + 1)
|
||||
],
|
||||
"round_endpoint_ids": [
|
||||
_agent_round_endpoint_ids.get(i)
|
||||
for i in range(1, max(_agent_round_models, default=1) + 1)
|
||||
],
|
||||
"round_endpoint_labels": [
|
||||
_agent_round_endpoint_labels.get(i)
|
||||
for i in range(1, max(_agent_round_models, default=1) + 1)
|
||||
],
|
||||
},
|
||||
)
|
||||
sess.add_message(ChatMessage("assistant", _stopped_content2, metadata=_stopped_md2))
|
||||
@@ -1866,8 +2380,12 @@ def setup_chat_routes(
|
||||
if compare_mode:
|
||||
return StreamingResponse(_safe_stream(), media_type="text/event-stream")
|
||||
|
||||
agent_runs.start(session, _safe_stream())
|
||||
return StreamingResponse(agent_runs.subscribe(session), media_type="text/event-stream")
|
||||
_detached_run = agent_runs.start(session, _safe_stream())
|
||||
return StreamingResponse(
|
||||
agent_runs.subscribe(session, _detached_run),
|
||||
media_type="text/event-stream",
|
||||
headers={"X-Odysseus-Run-Id": _detached_run.run_id},
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# GET /api/chat/resume — reconnect to a detached run that's still going
|
||||
@@ -1876,9 +2394,14 @@ def setup_chat_routes(
|
||||
@router.get("/api/chat/resume/{session_id}")
|
||||
async def chat_resume(request: Request, session_id: str) -> StreamingResponse:
|
||||
_verify_session_owner(request, session_id)
|
||||
if not agent_runs.is_active(session_id):
|
||||
_active_run = agent_runs.get_active_run(session_id)
|
||||
if _active_run is None:
|
||||
raise HTTPException(404, "No active run for this session")
|
||||
return StreamingResponse(agent_runs.subscribe(session_id), media_type="text/event-stream")
|
||||
return StreamingResponse(
|
||||
agent_runs.subscribe(session_id, _active_run),
|
||||
media_type="text/event-stream",
|
||||
headers={"X-Odysseus-Run-Id": _active_run.run_id},
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# POST /api/chat/stop — cancel a detached run (Stop button). Closing the SSE
|
||||
@@ -1887,7 +2410,8 @@ def setup_chat_routes(
|
||||
@router.post("/api/chat/stop/{session_id}")
|
||||
async def chat_stop(request: Request, session_id: str) -> Dict[str, Any]:
|
||||
_verify_session_owner(request, session_id)
|
||||
stopped = agent_runs.stop(session_id)
|
||||
_expected_run_id = request.headers.get("X-Odysseus-Run-Id")
|
||||
stopped = agent_runs.stop(session_id, _expected_run_id)
|
||||
return {"stopped": stopped}
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
@@ -73,6 +73,30 @@ _HF_TOKEN_STATUS_SNIPPET = (
|
||||
)
|
||||
|
||||
|
||||
def _windows_local_pid_record_line(pid_path: Path, ready_path: Path) -> str:
|
||||
"""Build the Git Bash prelude that records a Win32-stoppable PID.
|
||||
|
||||
Python publishes the detached outer process's Win32 PID first, then touches
|
||||
``ready_path``. The inner Git Bash runner waits for that publication before
|
||||
replacing the fallback with its own Win32 PID from /proc/<msys-pid>/winpid.
|
||||
|
||||
Missing, malformed, or late mappings leave the valid outer PID untouched.
|
||||
"""
|
||||
pp = shlex.quote(pid_path.as_posix())
|
||||
rp = shlex.quote(ready_path.as_posix())
|
||||
return (
|
||||
"i=0; "
|
||||
f"while [ ! -e {rp} ] && [ \"$i\" -lt 500 ]; do "
|
||||
"i=$((i+1)); sleep 0.01; done; "
|
||||
f"if [ -e {rp} ]; then "
|
||||
"winpid=\"$(cat /proc/$$/winpid 2>/dev/null || true)\"; "
|
||||
"case \"$winpid\" in ''|*[!0-9]*) ;; "
|
||||
f"*) printf '%s\\n' \"$winpid\" > {pp} ;; esac; "
|
||||
"fi; "
|
||||
f"rm -f {rp}"
|
||||
)
|
||||
|
||||
|
||||
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"
|
||||
@@ -978,15 +1002,18 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
directly (simple commands only). Returns the launched job record."""
|
||||
log_path = TMUX_LOG_DIR / f"{session_id}.log"
|
||||
pid_path = TMUX_LOG_DIR / f"{session_id}.pid"
|
||||
pid_ready_path: Path | None = None
|
||||
bash = find_bash()
|
||||
if bash:
|
||||
# Run the existing bash wrapper verbatim through Git Bash, redirecting
|
||||
# all output to the log the poller reads. Paths handed to bash use
|
||||
# POSIX form + shell-quoting so drive paths / spaces survive.
|
||||
inner = TMUX_LOG_DIR / f"{session_id}_run.sh"
|
||||
pp = shlex.quote(pid_path.as_posix())
|
||||
pid_ready_path = TMUX_LOG_DIR / f"{session_id}.pid.ready"
|
||||
pid_ready_path.unlink(missing_ok=True)
|
||||
inner.write_text(
|
||||
f"printf '%s\\n' \"$$\" > {pp}\n" + "\n".join(bash_lines) + "\n",
|
||||
_windows_local_pid_record_line(pid_path, pid_ready_path) + "\n"
|
||||
+ "\n".join(bash_lines) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
lp = shlex.quote(log_path.as_posix())
|
||||
@@ -1020,7 +1047,18 @@ def setup_cookbook_routes() -> APIRouter:
|
||||
env=env,
|
||||
**detached_popen_kwargs(),
|
||||
)
|
||||
# Publish a valid Win32 ancestor first. The Git Bash runner may then
|
||||
# replace it with its own Win32 pid, but never before this fallback exists.
|
||||
pid_path.write_text(str(proc.pid), encoding="utf-8")
|
||||
if pid_ready_path is not None:
|
||||
try:
|
||||
pid_ready_path.touch()
|
||||
except OSError as e:
|
||||
logger.warning(
|
||||
"Could not publish Windows local PID handoff for %s: %s",
|
||||
session_id,
|
||||
e,
|
||||
)
|
||||
return {"pid": proc.pid, "log_path": str(log_path)}
|
||||
|
||||
@router.post("/api/model/download")
|
||||
|
||||
@@ -247,6 +247,7 @@ import re as _re_reply
|
||||
_REPLY_OPEN_RE = _re_reply.compile(r"<<<\s*(?:REPLY|SUMMARY|OUTPUT)\s*>>+", _re_reply.I)
|
||||
_REPLY_CLOSE_RE = _re_reply.compile(r"<<<\s*END\s*>>+", _re_reply.I)
|
||||
_REPLY_ROLE_MARKER_RE = _re_reply.compile(r"</?\|(?:assistant|assistan|user|system|tool)\|>?|</\|end\|>?", _re_reply.I)
|
||||
_SUMMARY_BULLET_RE = _re_reply.compile(r"^(?:[-*\u2022]\s+|\d+[.)]\s+)")
|
||||
|
||||
|
||||
def _extract_reply(text: str) -> str:
|
||||
@@ -277,6 +278,125 @@ def _extract_reply(text: str) -> str:
|
||||
return _strip_think(t).strip()
|
||||
|
||||
|
||||
def _build_email_summary_messages(sender: str, subject: str, body_for_llm: str) -> list[dict[str, str]]:
|
||||
return [
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
"You are an email summarizer. Format: 1-3 short bullet points "
|
||||
"(use '- '). Cover: main point, action items, deadlines. If the "
|
||||
"email has attachments (marked '--- ATTACHMENTS ---'), USE THEIR "
|
||||
"CONTENTS - pull invoice totals, deadlines, key clauses, concrete "
|
||||
"numbers/dates from PDFs/docs into the bullets. Be terse.\n\n"
|
||||
"OUTPUT FORMAT: Put ONLY the bullet points between these exact "
|
||||
"markers, each on its own line:\n"
|
||||
"<<<SUMMARY>>>\n"
|
||||
"- ...\n"
|
||||
"<<<END>>>\n"
|
||||
"Any reasoning must come BEFORE <<<SUMMARY>>> (ideally inside "
|
||||
"<think>...</think>). Only the text between the markers is kept."
|
||||
),
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"From: {sender}\nSubject: {subject}\n\n{body_for_llm[:12000]}"
|
||||
"\n\n---\n\nSummarize the email. Output the bullets between "
|
||||
"<<<SUMMARY>>> and <<<END>>>."
|
||||
),
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
async def _generate_email_summary(
|
||||
url: str,
|
||||
model: str,
|
||||
sender: str,
|
||||
subject: str,
|
||||
body_for_llm: str,
|
||||
*,
|
||||
headers: dict | None = None,
|
||||
max_tokens: int = 8192,
|
||||
timeout: int = 180,
|
||||
) -> str:
|
||||
"""Generate an interactive email summary through the shared LLM adapter."""
|
||||
from src.llm_core import llm_call_async
|
||||
|
||||
raw = await llm_call_async(
|
||||
url=url,
|
||||
model=model,
|
||||
messages=_build_email_summary_messages(sender, subject, body_for_llm),
|
||||
temperature=0.3,
|
||||
max_tokens=max_tokens,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
workload="foreground",
|
||||
)
|
||||
return _normalize_email_summary(raw)
|
||||
|
||||
|
||||
async def _generate_scheduled_email_summary(
|
||||
url: str,
|
||||
model: str,
|
||||
sender: str,
|
||||
subject: str,
|
||||
body_for_llm: str,
|
||||
*,
|
||||
headers: dict | None = None,
|
||||
owner: str | None = None,
|
||||
max_tokens: int = 8192,
|
||||
timeout: int = 180,
|
||||
) -> str:
|
||||
"""Generate a scheduled summary through the background task candidate chain."""
|
||||
from src.task_endpoint import task_llm_call_async
|
||||
|
||||
raw = await task_llm_call_async(
|
||||
messages=_build_email_summary_messages(sender, subject, body_for_llm),
|
||||
fallback_url=url,
|
||||
fallback_model=model,
|
||||
fallback_headers=headers,
|
||||
owner=owner,
|
||||
temperature=0.3,
|
||||
max_tokens=max_tokens,
|
||||
timeout=timeout,
|
||||
)
|
||||
return _normalize_email_summary(raw)
|
||||
|
||||
|
||||
def _normalize_email_summary(raw) -> str:
|
||||
"""Extract a stable cache/UI summary from provider output."""
|
||||
raw_text = raw or ""
|
||||
if _REPLY_OPEN_RE.search(raw_text):
|
||||
summary = _extract_reply(raw_text)
|
||||
if summary:
|
||||
return summary
|
||||
|
||||
cleaned = _strip_think(raw_text).strip()
|
||||
bullets = [
|
||||
line.strip()
|
||||
for line in cleaned.splitlines()
|
||||
if _SUMMARY_BULLET_RE.match(line.strip())
|
||||
]
|
||||
if bullets:
|
||||
return "\n".join(bullets)
|
||||
return cleaned.strip()
|
||||
|
||||
|
||||
EMAIL_SUMMARY_ERROR_CODE = "email_summary_unavailable"
|
||||
EMAIL_SUMMARY_ERROR_MESSAGE = "Failed to summarize"
|
||||
|
||||
|
||||
def _email_summary_failure_log_detail(exc: BaseException) -> str:
|
||||
"""Return useful provider-failure metadata without echoing exception text."""
|
||||
detail = f"type={type(exc).__name__}"
|
||||
status = getattr(exc, "status_code", None)
|
||||
if status is None:
|
||||
status = getattr(getattr(exc, "response", None), "status_code", None)
|
||||
if isinstance(status, int):
|
||||
detail += f" status={status}"
|
||||
return detail
|
||||
|
||||
|
||||
def _apply_email_style_mechanics(text: str) -> str:
|
||||
"""Enforce deterministic writing-style mechanics that models often miss."""
|
||||
if not text:
|
||||
|
||||
+23
-9
@@ -40,6 +40,7 @@ from routes.email_helpers import (
|
||||
_pre_retrieve_context,
|
||||
_attach_compose_uploads, _cleanup_compose_uploads, _q,
|
||||
SCHEDULED_DB, _EMAIL_REPLY_SYS_PROMPT_BASE, _email_cache_owner_clause,
|
||||
_generate_scheduled_email_summary, _email_summary_failure_log_detail,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -653,6 +654,7 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
no_msgid = 0
|
||||
examined = 0
|
||||
_summaries_created = 0
|
||||
_summary_failed = 0
|
||||
_events_created = 0
|
||||
_replies_drafted = 0
|
||||
_reply_failed = 0
|
||||
@@ -785,16 +787,17 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
|
||||
if need_sum:
|
||||
try:
|
||||
summary = await task_llm_call_async(
|
||||
messages=[
|
||||
{"role": "system", "content": "You are an email summarizer. Format: 1-3 short bullet points (use '- '). Cover: main point, action items, deadlines. If the email has attachments (marked '--- ATTACHMENTS ---'), USE THEIR CONTENTS — pull out invoice totals, deadlines, key clauses, any concrete numbers/dates in PDFs/docs, and reflect them in the bullets. Be terse.\n\nOUTPUT FORMAT: Put ONLY the bullet points between these exact markers, each on its own line:\n<<<SUMMARY>>>\n- ...\n<<<END>>>\nAny reasoning or planning must come BEFORE <<<SUMMARY>>> (ideally inside <think>...</think>). Only the text between the markers is kept."},
|
||||
{"role": "user", "content": f"From: {sender}\nSubject: {subject}\n\n{body_for_llm[:12000]}\n\n---\n\nSummarize the email. Output the bullets between <<<SUMMARY>>> and <<<END>>>."},
|
||||
],
|
||||
fallback_url=url, fallback_model=model, fallback_headers=headers,
|
||||
summary = await _generate_scheduled_email_summary(
|
||||
url=url,
|
||||
model=model,
|
||||
sender=sender,
|
||||
subject=subject,
|
||||
body_for_llm=body_for_llm,
|
||||
headers=req_headers,
|
||||
owner=account_owner or None,
|
||||
temperature=0.3, max_tokens=16384, timeout=240,
|
||||
max_tokens=16384,
|
||||
timeout=240,
|
||||
)
|
||||
summary = _extract_reply((summary or "").strip())
|
||||
if summary:
|
||||
_c = _sql3.connect(SCHEDULED_DB)
|
||||
_c.execute("""
|
||||
@@ -808,10 +811,19 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
_summaries_created += 1
|
||||
_uid_text = uid.decode() if isinstance(uid, bytes) else str(uid)
|
||||
_detail_lines.append(f"summary · {_folder}#{_uid_text} · {subject or '(no subject)'} — {sender or '(unknown sender)'}")
|
||||
else:
|
||||
_summary_failed += 1
|
||||
_uid_text = uid.decode() if isinstance(uid, bytes) else str(uid)
|
||||
_detail_lines.append(f"summary empty · {_folder}#{_uid_text} · {subject or '(no subject)'} — {sender or '(unknown sender)'}")
|
||||
except Exception as e:
|
||||
_summary_failed += 1
|
||||
_uid_text = uid.decode() if isinstance(uid, bytes) else str(uid)
|
||||
_detail_lines.append(f"summary failed · {_folder}#{_uid_text} · {subject or '(no subject)'} — {sender or '(unknown sender)'}")
|
||||
logger.warning(f"Auto-summary {uid} failed: {e}")
|
||||
logger.warning(
|
||||
"Auto-summary uid=%s failed %s",
|
||||
_uid_text,
|
||||
_email_summary_failure_log_detail(e),
|
||||
)
|
||||
|
||||
if need_reply:
|
||||
await _emit_progress(progress_cb, f"Drafting reply {processed + 1}/{_max_process} · checked {examined}/{len(uid_list)}")
|
||||
@@ -1320,6 +1332,8 @@ async def _auto_summarize_pass_single(days_back: int = 1, account_id: str | None
|
||||
parts.append(f"processed {processed} new")
|
||||
if auto_sum:
|
||||
parts.append(f"summarized {_summaries_created}")
|
||||
if _summary_failed:
|
||||
parts.append(f"{_summary_failed} summary failed")
|
||||
if auto_reply_draft:
|
||||
parts.append(f"drafted {_replies_drafted} repl" + ("y" if _replies_drafted == 1 else "ies"))
|
||||
if _reply_failed:
|
||||
|
||||
+241
-121
@@ -45,6 +45,7 @@ from src.upload_limits import read_upload_limited, EMAIL_COMPOSE_UPLOAD_MAX_BYTE
|
||||
|
||||
from routes.email_helpers import (
|
||||
_strip_think, _extract_reply, _apply_email_style_mechanics, require_owner, require_user, _assert_owns_account,
|
||||
_account_visible_to_owner,
|
||||
_q, _attach_compose_uploads, _cleanup_compose_uploads,
|
||||
_load_settings, _save_settings, _get_email_config,
|
||||
_send_smtp_message, _smtp_security_mode,
|
||||
@@ -57,7 +58,8 @@ from routes.email_helpers import (
|
||||
_extract_attachment_to_disk, _extract_html, _extract_text,
|
||||
_fetch_sender_thread_context, _pre_retrieve_context,
|
||||
_EMAIL_REPLY_SYS_PROMPT_BASE, _POOL_HOOKS,
|
||||
_friendly_email_auth_error,
|
||||
_friendly_email_auth_error, _email_summary_failure_log_detail,
|
||||
_generate_email_summary, EMAIL_SUMMARY_ERROR_CODE, EMAIL_SUMMARY_ERROR_MESSAGE,
|
||||
SendEmailRequest, ExtractStyleRequest,
|
||||
ATTACHMENTS_DIR, COMPOSE_UPLOADS_DIR, SCHEDULED_DB,
|
||||
attachment_extract_dir, _email_cache_owner_clause, email_translation_body_hash,
|
||||
@@ -194,6 +196,64 @@ def _coerce_port(value, default):
|
||||
return None, f"Invalid port {value!r}; must be a whole number"
|
||||
|
||||
|
||||
def _lock_email_account_owner_mutation(db, *owners: str) -> None:
|
||||
"""Delegate account/default serialization to the shared DB primitive."""
|
||||
from core.database import lock_email_account_owner_mutations
|
||||
|
||||
lock_email_account_owner_mutations(db, *owners)
|
||||
|
||||
|
||||
def _email_account_owner_scope(query, owner: str):
|
||||
"""Restrict a query to one normalized EmailAccount owner partition."""
|
||||
from core.database import EmailAccount
|
||||
from sqlalchemy import or_
|
||||
|
||||
if owner:
|
||||
return query.filter(EmailAccount.owner == owner)
|
||||
return query.filter(or_(EmailAccount.owner == None, EmailAccount.owner == "")) # noqa: E711
|
||||
|
||||
|
||||
def _discover_email_account_mutation_scope(account_id: str, owner: str) -> str:
|
||||
"""Read the initial lock key and fail closed before a mutation session."""
|
||||
from core.database import EmailAccount, SessionLocal
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
row = db.get(EmailAccount, account_id)
|
||||
if row is None or (owner and not _account_visible_to_owner(row, owner)):
|
||||
raise HTTPException(404, "Account not found")
|
||||
return row.owner or ""
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error("Account-owner mutation check failed: %s", exc)
|
||||
raise HTTPException(503, "Account check failed")
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _lock_and_reload_email_account(db, account_id: str, owner: str, scope: str):
|
||||
"""Lock, reload, and revalidate an account, retrying if its owner moved."""
|
||||
from core.database import EmailAccount
|
||||
|
||||
owner_scopes = {scope or ""}
|
||||
while True:
|
||||
_lock_email_account_owner_mutation(db, *owner_scopes)
|
||||
row = db.get(EmailAccount, account_id, populate_existing=True)
|
||||
if row is None or (owner and not _account_visible_to_owner(row, owner)):
|
||||
raise HTTPException(404, "Account not found")
|
||||
|
||||
current_scope = row.owner or ""
|
||||
if current_scope in owner_scopes or db.get_bind().dialect.name == "sqlite":
|
||||
return row
|
||||
|
||||
# The account changed owner after discovery but before lock acquisition.
|
||||
# Release the partial lock set and reacquire all observed scopes in the
|
||||
# shared helper's canonical order, then validate from the database again.
|
||||
db.rollback()
|
||||
owner_scopes.add(current_scope)
|
||||
|
||||
|
||||
def _email_tag_owner_aliases(account_id: str | None, owner: str = "") -> list[str]:
|
||||
aliases = [owner or ""]
|
||||
try:
|
||||
@@ -2860,13 +2920,22 @@ def setup_email_routes():
|
||||
return indexed_response
|
||||
return {"emails": [], "total": 0, "error": "Mail operation failed"}
|
||||
|
||||
def _read_email_sync(uid, folder, account_id, owner, mark_seen=True, full=False):
|
||||
def _read_email_sync(uid, folder, account_id, owner, mark_seen=False, full=False):
|
||||
"""Sync IMAP read — wrapped in to_thread by the async handler.
|
||||
|
||||
The normal reader path fetches the headers plus a bounded body prefix.
|
||||
That avoids downloading multi-megabyte attachments just to open a
|
||||
message. Full-message fetch remains available for flows that need
|
||||
attachment metadata immediately, such as forwarding.
|
||||
|
||||
`mark_seen` defaults to False because it mutates provider state: it
|
||||
selects the mailbox read-write and issues a STORE. Only a foreground
|
||||
open should ask for it, and it has to ask explicitly.
|
||||
|
||||
A failed \\Seen transition is reported as `mark_seen_failed` on an
|
||||
otherwise normal response, never as an error. The body has already been
|
||||
fetched at that point, so refusing to return it would turn a cosmetic
|
||||
flag failure into an unreadable message.
|
||||
"""
|
||||
import time as _t
|
||||
_t0 = _t.monotonic()
|
||||
@@ -2874,9 +2943,28 @@ def setup_email_routes():
|
||||
preview_bytes = 384 * 1024
|
||||
_t_select = 0.0
|
||||
_t_fetch = 0.0
|
||||
mark_seen_failed = False
|
||||
try:
|
||||
with _imap(account_id, owner=owner) as conn:
|
||||
conn.select(_q(folder), readonly=True)
|
||||
# A foreground open owns both the body fetch and the \Seen
|
||||
# transition. Keep them on one read-write IMAP selection so the
|
||||
# route never schedules a second connection that can race the
|
||||
# response. Prefetch/read-only callers retain BODY.PEEK and a
|
||||
# read-only mailbox selection.
|
||||
try:
|
||||
conn.select(_q(folder), readonly=not mark_seen)
|
||||
except Exception as select_exc:
|
||||
if not mark_seen:
|
||||
raise
|
||||
# Read-only mailboxes (shared archives, some provider
|
||||
# folders) reject a read-write SELECT. Serve the message
|
||||
# read-only and report the flag failure.
|
||||
logger.warning(
|
||||
f"read-write SELECT rejected for {folder!r}; "
|
||||
f"serving read-only without \\Seen: {select_exc}"
|
||||
)
|
||||
conn.select(_q(folder), readonly=True)
|
||||
mark_seen_failed = True
|
||||
_t_select = _t.monotonic() - _t0
|
||||
fetch_query = "(BODY.PEEK[])" if full else f"(BODY.PEEK[HEADER] BODY.PEEK[TEXT]<0.{preview_bytes}>)"
|
||||
status, msg_data = _imap_uid_fetch(conn, uid, fetch_query)
|
||||
@@ -2902,22 +2990,44 @@ def setup_email_routes():
|
||||
header_part = msg_data[0][1] or b""
|
||||
raw = header_part + b"\r\n" + text_part
|
||||
|
||||
msg = email_mod.message_from_bytes(raw)
|
||||
# Parse the fetched payload before mutating provider state. If
|
||||
# the message is malformed enough that the reader cannot build
|
||||
# a response, the caller gets an error while the message stays
|
||||
# unread instead of receiving a false optimistic rollback.
|
||||
msg = email_mod.message_from_bytes(raw)
|
||||
|
||||
subject = _decode_header(msg.get("Subject", "(no subject)"))
|
||||
sender = _decode_header(msg.get("From", "unknown"))
|
||||
to = _decode_header(msg.get("To", ""))
|
||||
cc = _decode_header(msg.get("Cc", ""))
|
||||
date_str = msg.get("Date", "")
|
||||
message_id = msg.get("Message-ID", "")
|
||||
in_reply_to = msg.get("In-Reply-To", "")
|
||||
references = msg.get("References", "")
|
||||
body = _extract_text(msg)
|
||||
body_html = _extract_html(msg)
|
||||
subject = _decode_header(msg.get("Subject", "(no subject)"))
|
||||
sender = _decode_header(msg.get("From", "unknown"))
|
||||
to = _decode_header(msg.get("To", ""))
|
||||
cc = _decode_header(msg.get("Cc", ""))
|
||||
date_str = msg.get("Date", "")
|
||||
message_id = msg.get("Message-ID", "")
|
||||
in_reply_to = msg.get("In-Reply-To", "")
|
||||
references = msg.get("References", "")
|
||||
body = _extract_text(msg)
|
||||
body_html = _extract_html(msg)
|
||||
|
||||
sender_name, sender_addr = email.utils.parseaddr(sender)
|
||||
parsed_date = email.utils.parsedate_to_datetime(date_str) if date_str else None
|
||||
attachments = _list_attachments_from_msg(msg) if full else (_email_attachment_meta_cache_get(owner, account_id, folder, uid) or [])
|
||||
|
||||
if mark_seen and not mark_seen_failed:
|
||||
seen_status, _ = conn.uid("STORE", _uid_bytes(uid), "+FLAGS", "(\\Seen)")
|
||||
if seen_status != "OK":
|
||||
# Report, don't raise. The parsed body below is still a
|
||||
# valid response; only the flag claim is untrue.
|
||||
logger.warning(
|
||||
f"IMAP STORE \\Seen failed for UID {uid} in {folder!r}: {seen_status}"
|
||||
)
|
||||
mark_seen_failed = True
|
||||
|
||||
# Only record the local flag transition when the provider actually
|
||||
# accepted it, so the index and list cache cannot drift ahead of
|
||||
# the mailbox.
|
||||
if mark_seen and not mark_seen_failed:
|
||||
_email_index_update_flags(owner, account_id, folder, uid, "\\Seen", True)
|
||||
_update_list_cache_seen(account_id, folder, uid, True)
|
||||
|
||||
sender_name, sender_addr = email.utils.parseaddr(sender)
|
||||
parsed_date = email.utils.parsedate_to_datetime(date_str) if date_str else None
|
||||
attachments = _list_attachments_from_msg(msg) if full else (_email_attachment_meta_cache_get(owner, account_id, folder, uid) or [])
|
||||
related_attachments = []
|
||||
if full and not _has_visible_attachments(msg):
|
||||
related_attachments = _related_thread_attachments_sync(
|
||||
@@ -3038,20 +3148,29 @@ def setup_email_routes():
|
||||
"boundaries": cached_boundaries,
|
||||
"thread_turns": cached_turns,
|
||||
"sender_signature": cached_sender_sig,
|
||||
# Per-request, not part of the message: the route strips this
|
||||
# before caching so a one-off flag failure is never replayed to
|
||||
# later readers.
|
||||
"mark_seen_failed": mark_seen_failed,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to read email {uid}: {e}")
|
||||
return {"error": "Mail operation failed"}
|
||||
|
||||
def _mark_email_seen_sync(uid, folder, account_id, owner):
|
||||
"""Synchronously mark a cached email seen and report success."""
|
||||
try:
|
||||
with _imap(account_id, owner=owner) as conn:
|
||||
conn.select(_q(folder))
|
||||
conn.uid("STORE", _uid_bytes(uid), "+FLAGS", "\\Seen")
|
||||
conn.select(_q(folder), readonly=False)
|
||||
status, _ = conn.uid("STORE", _uid_bytes(uid), "+FLAGS", "(\\Seen)")
|
||||
if status != "OK":
|
||||
return False
|
||||
_email_index_update_flags(owner, account_id, folder, uid, "\\Seen", True)
|
||||
_update_list_cache_seen(account_id, folder, uid, True)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.debug(f"mark-seen after cached read failed uid={uid}: {e}")
|
||||
logger.warning(f"mark-seen after cached read failed uid={uid}: {e}")
|
||||
return False
|
||||
|
||||
@router.get("/read/{uid}")
|
||||
async def read_email_by_uid(
|
||||
@@ -3077,32 +3196,32 @@ def setup_email_routes():
|
||||
if cached.get("attachment_version") != EMAIL_READ_ATTACHMENT_VERSION:
|
||||
cached = None
|
||||
if cached is not None:
|
||||
if mark_seen:
|
||||
try:
|
||||
_asyncio.create_task(_asyncio.to_thread(_mark_email_seen_sync, uid, folder, account_id, owner))
|
||||
except RuntimeError:
|
||||
pass
|
||||
# A cache hit already holds a complete, valid message. Await the
|
||||
# STORE so the response reports the real flag state, but never let
|
||||
# a failed STORE withhold a body we are holding in memory.
|
||||
if mark_seen and not await _asyncio.to_thread(
|
||||
_mark_email_seen_sync, uid, folder, account_id, owner
|
||||
):
|
||||
return {**cached, "mark_seen_failed": True}
|
||||
return cached
|
||||
if not full:
|
||||
persisted = _email_preview_cache_get(owner, account_id, folder, uid)
|
||||
if persisted and persisted.get("attachment_version") == EMAIL_READ_ATTACHMENT_VERSION:
|
||||
_read_cache_put(ck, persisted)
|
||||
if mark_seen:
|
||||
try:
|
||||
_asyncio.create_task(_asyncio.to_thread(_mark_email_seen_sync, uid, folder, account_id, owner))
|
||||
except RuntimeError:
|
||||
pass
|
||||
if mark_seen and not await _asyncio.to_thread(
|
||||
_mark_email_seen_sync, uid, folder, account_id, owner
|
||||
):
|
||||
return {**persisted, "mark_seen_failed": True}
|
||||
return persisted
|
||||
result = await _asyncio.to_thread(_read_email_sync, uid, folder, account_id, owner, mark_seen, full)
|
||||
if result and not result.get("error"):
|
||||
_read_cache_put(ck, result)
|
||||
# `mark_seen_failed` describes this request, not the message, so it
|
||||
# must not enter either cache — a later reader would otherwise be
|
||||
# told a STORE failed that it never issued.
|
||||
cacheable = {k: v for k, v in result.items() if k != "mark_seen_failed"}
|
||||
_read_cache_put(ck, cacheable)
|
||||
if not full:
|
||||
_email_preview_cache_put(owner, account_id, folder, uid, result)
|
||||
if mark_seen:
|
||||
try:
|
||||
_asyncio.create_task(_asyncio.to_thread(_mark_email_seen_sync, uid, folder, account_id, owner))
|
||||
except RuntimeError:
|
||||
pass
|
||||
_email_preview_cache_put(owner, account_id, folder, uid, cacheable)
|
||||
return result
|
||||
|
||||
def _schedule_recent_email_warm(emails: list, folder: str, account_id: str | None, owner: str):
|
||||
@@ -4766,8 +4885,6 @@ def setup_email_routes():
|
||||
"""Generate a quick AI summary of an email body."""
|
||||
try:
|
||||
from src.endpoint_resolver import resolve_endpoint
|
||||
from src.llm_core import _uses_max_completion_tokens, _restricts_temperature
|
||||
import requests as _req
|
||||
|
||||
body = data.get("body", "")
|
||||
subject = data.get("subject", "")
|
||||
@@ -4778,7 +4895,11 @@ def setup_email_routes():
|
||||
if account_id:
|
||||
_assert_owns_account(account_id, owner)
|
||||
if not body:
|
||||
return {"success": False, "error": "No body provided"}
|
||||
return {
|
||||
"success": False,
|
||||
"error": "No body provided",
|
||||
"error_code": "email_summary_missing_body",
|
||||
}
|
||||
|
||||
# If we know which UID this is, fetch the raw message and pull
|
||||
# attachment text so the summary can reference invoice totals,
|
||||
@@ -4807,53 +4928,43 @@ def setup_email_routes():
|
||||
if not url:
|
||||
url, model, headers = resolve_endpoint("default", owner=owner)
|
||||
if not url or not model:
|
||||
return {"success": False, "error": "No LLM endpoint configured"}
|
||||
return {
|
||||
"success": False,
|
||||
"error": "No model configured for email summaries",
|
||||
"error_code": "email_summary_not_configured",
|
||||
}
|
||||
|
||||
req_headers = {"Content-Type": "application/json"}
|
||||
if headers:
|
||||
req_headers.update(headers)
|
||||
tok_key = "max_completion_tokens" if _uses_max_completion_tokens(model) else "max_tokens"
|
||||
payload = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are an email summarizer. Format: 1-3 short bullet points (use '- '). Cover: main point, action items, deadlines. If the email has attachments (marked '--- ATTACHMENTS ---'), USE THEIR CONTENTS — pull invoice totals, deadlines, key clauses, concrete numbers/dates from PDFs/docs into the bullets. Be terse.\n\nOUTPUT FORMAT: Put ONLY the bullet points between these exact markers, each on its own line:\n<<<SUMMARY>>>\n- ...\n<<<END>>>\nAny reasoning must come BEFORE <<<SUMMARY>>> (ideally inside <think>...</think>). Only the text between the markers is kept."},
|
||||
{"role": "user", "content": f"From: {sender}\nSubject: {subject}\n\n{body_for_llm[:12000]}\n\n---\n\nSummarize the email. Output the bullets between <<<SUMMARY>>> and <<<END>>>."},
|
||||
],
|
||||
tok_key: 8192,
|
||||
"temperature": 0.3,
|
||||
"stream": False,
|
||||
}
|
||||
# Reasoning models (o1/o3/o4/gpt-5) reject an explicit temperature.
|
||||
if _restricts_temperature(model):
|
||||
payload.pop("temperature", None)
|
||||
resp = await asyncio.to_thread(
|
||||
_req.post, url, json=payload, headers=req_headers, timeout=180
|
||||
)
|
||||
if not resp.ok:
|
||||
return {"success": False, "error": f"LLM HTTP {resp.status_code}"}
|
||||
rdata = resp.json()
|
||||
msg = (rdata.get("choices") or [{}])[0].get("message", {})
|
||||
content = (msg.get("content") or "").strip()
|
||||
content = _extract_reply(content)
|
||||
try:
|
||||
content = await _generate_email_summary(
|
||||
url=url,
|
||||
model=model,
|
||||
sender=sender,
|
||||
subject=subject,
|
||||
body_for_llm=body_for_llm,
|
||||
headers=req_headers,
|
||||
max_tokens=8192,
|
||||
timeout=180,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Email summary LLM call failed %s",
|
||||
_email_summary_failure_log_detail(e),
|
||||
)
|
||||
return {
|
||||
"success": False,
|
||||
"error": EMAIL_SUMMARY_ERROR_MESSAGE,
|
||||
"error_code": EMAIL_SUMMARY_ERROR_CODE,
|
||||
}
|
||||
|
||||
if not content:
|
||||
# Model put everything in reasoning_content — extract bullet points
|
||||
rc = (msg.get("reasoning_content") or "").strip()
|
||||
# Find bullet-point style output (lines starting with -, •, *, or numbered)
|
||||
bullet_lines = []
|
||||
for line in rc.split("\n"):
|
||||
stripped = line.strip()
|
||||
if re.match(r"^[-•*]\s+|^\d+[.)]\s+", stripped):
|
||||
bullet_lines.append(stripped)
|
||||
if bullet_lines:
|
||||
content = "\n".join(bullet_lines)
|
||||
else:
|
||||
# Last resort: take the last paragraph
|
||||
paragraphs = [p.strip() for p in rc.split("\n\n") if p.strip()]
|
||||
content = paragraphs[-1] if paragraphs else rc[:500]
|
||||
|
||||
if not content:
|
||||
return {"success": False, "error": "Empty response from model"}
|
||||
return {
|
||||
"success": False,
|
||||
"error": "The model returned an empty summary",
|
||||
"error_code": "email_summary_empty",
|
||||
}
|
||||
|
||||
# Cache the summary if we have a message_id
|
||||
mid = data.get("message_id", "")
|
||||
@@ -4876,8 +4987,15 @@ def setup_email_routes():
|
||||
|
||||
return {"success": True, "summary": content, "model_used": model}
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to summarize: {e}")
|
||||
return {"success": False, "error": "Mail operation failed"}
|
||||
logger.error(
|
||||
"Email summary route failed %s",
|
||||
_email_summary_failure_log_detail(e),
|
||||
)
|
||||
return {
|
||||
"success": False,
|
||||
"error": EMAIL_SUMMARY_ERROR_MESSAGE,
|
||||
"error_code": EMAIL_SUMMARY_ERROR_CODE,
|
||||
}
|
||||
|
||||
@router.post("/translate")
|
||||
async def translate_email(data: dict, owner: str = Depends(require_owner)):
|
||||
@@ -4886,7 +5004,6 @@ def setup_email_routes():
|
||||
from src.endpoint_resolver import (
|
||||
resolve_endpoint,
|
||||
resolve_utility_fallback_candidates,
|
||||
resolve_chat_fallback_candidates,
|
||||
)
|
||||
from src.llm_core import llm_call_async_with_fallback
|
||||
|
||||
@@ -4948,8 +5065,6 @@ def setup_email_routes():
|
||||
pass
|
||||
for cand in resolve_utility_fallback_candidates(owner=owner) or []:
|
||||
_add(*cand)
|
||||
for cand in resolve_chat_fallback_candidates(owner=owner) or []:
|
||||
_add(*cand)
|
||||
if not candidates:
|
||||
return {"success": False, "error": "No LLM endpoint configured"}
|
||||
|
||||
@@ -5209,13 +5324,11 @@ def setup_email_routes():
|
||||
# Build a candidate chain so a stale session-stored API key
|
||||
# (the most common cause of "authentication failed" here)
|
||||
# doesn't kill AI Reply outright — fall through to the
|
||||
# user's Utility / Default endpoints AND their configured
|
||||
# fallback chains. Dedupe by url+model so we don't retry
|
||||
# the same broken endpoint.
|
||||
# user's Utility / Default endpoints and active Utility fallback
|
||||
# chain. Dedupe by url+model so we don't retry the same endpoint.
|
||||
from src.llm_core import llm_call_async_with_fallback
|
||||
from src.endpoint_resolver import (
|
||||
resolve_utility_fallback_candidates,
|
||||
resolve_chat_fallback_candidates,
|
||||
)
|
||||
_seen = set()
|
||||
_candidates = []
|
||||
@@ -5240,11 +5353,9 @@ def setup_email_routes():
|
||||
_add(_d_url, _d_model, _d_headers)
|
||||
except Exception:
|
||||
pass
|
||||
# Configured fallback chains last.
|
||||
# Active Utility fallbacks last.
|
||||
for cand in resolve_utility_fallback_candidates(owner=owner) or []:
|
||||
_add(*cand)
|
||||
for cand in resolve_chat_fallback_candidates(owner=owner) or []:
|
||||
_add(*cand)
|
||||
_messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_msg},
|
||||
@@ -5428,9 +5539,9 @@ def setup_email_routes():
|
||||
import uuid as _uuid
|
||||
db = SessionLocal()
|
||||
try:
|
||||
_lock_email_account_owner_mutation(db, owner)
|
||||
q = db.query(EmailAccount).filter(EmailAccount.is_default == True) # noqa: E712
|
||||
if owner:
|
||||
q = q.filter(EmailAccount.owner == owner)
|
||||
q = _email_account_owner_scope(q, owner)
|
||||
row = q.first()
|
||||
if row is None:
|
||||
row = EmailAccount(id=_uuid.uuid4().hex, owner=owner, name="Default", is_default=True, enabled=True)
|
||||
@@ -5456,8 +5567,7 @@ def setup_email_routes():
|
||||
if data.get("smtp_password"):
|
||||
row.smtp_password = _enc(data["smtp_password"])
|
||||
clear_q = db.query(EmailAccount).filter(EmailAccount.id != row.id)
|
||||
if owner:
|
||||
clear_q = clear_q.filter(EmailAccount.owner == owner)
|
||||
clear_q = _email_account_owner_scope(clear_q, owner)
|
||||
clear_q.update({EmailAccount.is_default: False})
|
||||
db.commit()
|
||||
finally:
|
||||
@@ -5552,6 +5662,7 @@ def setup_email_routes():
|
||||
return {"ok": False, "error": port_err}
|
||||
db = SessionLocal()
|
||||
try:
|
||||
_lock_email_account_owner_mutation(db, owner)
|
||||
row = EmailAccount(
|
||||
id=_uuid.uuid4().hex,
|
||||
name=name,
|
||||
@@ -5578,9 +5689,7 @@ def setup_email_routes():
|
||||
# the one-default invariant — but scope it to THIS user's accounts,
|
||||
# otherwise creating a default would clear every other user's
|
||||
# default flag too.
|
||||
scope_q = db.query(EmailAccount)
|
||||
if owner:
|
||||
scope_q = scope_q.filter(EmailAccount.owner == owner)
|
||||
scope_q = _email_account_owner_scope(db.query(EmailAccount), owner)
|
||||
existing_count = scope_q.count()
|
||||
if row.is_default or existing_count == 0:
|
||||
scope_q.update({EmailAccount.is_default: False})
|
||||
@@ -5631,28 +5740,39 @@ def setup_email_routes():
|
||||
|
||||
@router.delete("/accounts/{account_id}")
|
||||
async def delete_email_account(account_id: str, owner: str = Depends(require_user)):
|
||||
_assert_owns_account(account_id, owner)
|
||||
initial_scope = _discover_email_account_mutation_scope(account_id, owner)
|
||||
from core.database import SessionLocal, EmailAccount
|
||||
db = SessionLocal()
|
||||
try:
|
||||
row = db.get(EmailAccount, account_id)
|
||||
if not row:
|
||||
return {"ok": False, "error": "Account not found"}
|
||||
row = _lock_and_reload_email_account(
|
||||
db, account_id, owner, initial_scope
|
||||
)
|
||||
row_scope = row.owner or ""
|
||||
was_default = bool(row.is_default)
|
||||
db.delete(row)
|
||||
db.commit()
|
||||
# Flush the removal before staging a replacement default. The
|
||||
# partial unique index is checked statement-by-statement, and the
|
||||
# ORM is otherwise free to UPDATE the promoted row before DELETE.
|
||||
db.flush()
|
||||
# If the deleted row was default, promote the next-oldest enabled
|
||||
# row owned by THIS user. Without the owner filter we'd promote
|
||||
# another user's account and the deleter would silently inherit
|
||||
# it as their default.
|
||||
if was_default:
|
||||
promote_q = db.query(EmailAccount).filter(EmailAccount.enabled == True) # noqa: E712
|
||||
if owner:
|
||||
promote_q = promote_q.filter(EmailAccount.owner == owner)
|
||||
promote = promote_q.order_by(EmailAccount.created_at.asc()).first()
|
||||
promote_q = db.query(EmailAccount).filter(
|
||||
EmailAccount.id != account_id,
|
||||
EmailAccount.enabled == True, # noqa: E712
|
||||
)
|
||||
promote_q = _email_account_owner_scope(promote_q, row_scope)
|
||||
promote = promote_q.order_by(
|
||||
EmailAccount.created_at.asc(), EmailAccount.id.asc()
|
||||
).first()
|
||||
if promote:
|
||||
promote.is_default = True
|
||||
db.commit()
|
||||
# Deletion and any replacement promotion are one durable state
|
||||
# transition, so another worker can never observe or race the old
|
||||
# split-commit gap.
|
||||
db.commit()
|
||||
return {"ok": True}
|
||||
finally:
|
||||
db.close()
|
||||
@@ -5865,18 +5985,18 @@ def setup_email_routes():
|
||||
|
||||
@router.post("/accounts/{account_id}/set-default")
|
||||
async def set_default_account(account_id: str, owner: str = Depends(require_user)):
|
||||
_assert_owns_account(account_id, owner)
|
||||
initial_scope = _discover_email_account_mutation_scope(account_id, owner)
|
||||
from core.database import SessionLocal, EmailAccount
|
||||
db = SessionLocal()
|
||||
try:
|
||||
row = db.get(EmailAccount, account_id)
|
||||
if not row:
|
||||
return {"ok": False, "error": "Account not found"}
|
||||
# SECURITY: scope the "clear other defaults" sweep to this user's
|
||||
# accounts so we don't unset another user's default flag.
|
||||
clear_q = db.query(EmailAccount)
|
||||
if owner:
|
||||
clear_q = clear_q.filter(EmailAccount.owner == owner)
|
||||
row = _lock_and_reload_email_account(
|
||||
db, account_id, owner, initial_scope
|
||||
)
|
||||
# Scope the sweep to the target row's normalized owner partition;
|
||||
# this also handles visible legacy NULL/empty-owner accounts.
|
||||
clear_q = _email_account_owner_scope(
|
||||
db.query(EmailAccount), row.owner or ""
|
||||
)
|
||||
clear_q.update({EmailAccount.is_default: False})
|
||||
row.is_default = True
|
||||
db.commit()
|
||||
@@ -5895,7 +6015,7 @@ def setup_email_routes():
|
||||
raise HTTPException(400, "GOOGLE_OAUTH_CLIENT_ID not set — add it to .env")
|
||||
redirect_uri = (
|
||||
os.environ.get("GOOGLE_OAUTH_REDIRECT_URI")
|
||||
or f"http://{request.headers.get('host', 'localhost:7000')}/api/email/oauth/google/callback"
|
||||
or f"{request.url.scheme}://{request.headers.get('host', 'localhost:7000')}/api/email/oauth/google/callback"
|
||||
)
|
||||
state = make_oauth_state(account_id, owner)
|
||||
params = urllib.parse.urlencode({
|
||||
@@ -5932,7 +6052,7 @@ def setup_email_routes():
|
||||
client_secret = os.environ.get("GOOGLE_OAUTH_CLIENT_SECRET", "")
|
||||
redirect_uri = (
|
||||
os.environ.get("GOOGLE_OAUTH_REDIRECT_URI")
|
||||
or f"http://{request.headers.get('host', 'localhost:7000')}/api/email/oauth/google/callback"
|
||||
or f"{request.url.scheme}://{request.headers.get('host', 'localhost:7000')}/api/email/oauth/google/callback"
|
||||
)
|
||||
import httpx as _httpx
|
||||
try:
|
||||
|
||||
@@ -127,6 +127,25 @@ def _load_grounding_backend():
|
||||
return cached
|
||||
|
||||
|
||||
def _model_input_to_device(value, device: str, torch):
|
||||
if not hasattr(value, "to"):
|
||||
return value
|
||||
if (
|
||||
device == "mps"
|
||||
and hasattr(torch, "float64")
|
||||
and getattr(value, "dtype", None) == torch.float64
|
||||
):
|
||||
return value.to(device=device, dtype=torch.float32)
|
||||
return value.to(device)
|
||||
|
||||
|
||||
def _model_inputs_to_device(inputs, device: str, torch) -> Dict[str, Any]:
|
||||
return {
|
||||
key: _model_input_to_device(value, device, torch)
|
||||
for key, value in inputs.items()
|
||||
}
|
||||
|
||||
|
||||
def _ground_text_to_box(image, text: str, *, threshold: float = 0.05):
|
||||
query = (text or "").strip()
|
||||
if not query:
|
||||
@@ -142,10 +161,7 @@ def _ground_text_to_box(image, text: str, *, threshold: float = 0.05):
|
||||
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()
|
||||
}
|
||||
model_inputs = _model_inputs_to_device(inputs, device, torch)
|
||||
with torch.no_grad():
|
||||
outputs = model(**model_inputs)
|
||||
target_sizes = torch.tensor([[image.height, image.width]])
|
||||
@@ -1869,10 +1885,7 @@ def setup_gallery_routes() -> APIRouter:
|
||||
|
||||
try:
|
||||
inputs = processor(image, **kwargs)
|
||||
model_inputs = {
|
||||
k: (v.to(device) if hasattr(v, "to") else v)
|
||||
for k, v in inputs.items()
|
||||
}
|
||||
model_inputs = _model_inputs_to_device(inputs, device, torch)
|
||||
with torch.no_grad():
|
||||
outputs = model(**model_inputs)
|
||||
masks = processor.image_processor.post_process_masks(
|
||||
|
||||
@@ -137,44 +137,6 @@ def setup_history_routes(session_manager, upload_handler=None) -> 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,
|
||||
@@ -198,6 +160,8 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
)
|
||||
page_offset = int(offset) if offset is not None else max(total - page_limit, 0)
|
||||
page_offset = max(0, min(page_offset, total))
|
||||
# Keep display pagination page-scoped. ``get_session`` is the
|
||||
# full model-context hydration seam and must not be entered here.
|
||||
rows = (
|
||||
db.query(DbChatMessage)
|
||||
.filter(DbChatMessage.session_id == session_id)
|
||||
@@ -206,14 +170,6 @@ def setup_history_routes(session_manager, upload_handler=None) -> 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")
|
||||
@@ -258,7 +214,10 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
entry["metadata"] = msg["metadata"]
|
||||
history_dict.append(entry)
|
||||
|
||||
# Fallback: load from DB if in-memory is empty
|
||||
# Fallback: load from DB if in-memory renders empty. Display only —
|
||||
# get_session above is the hydration seam, so nothing here writes back
|
||||
# into session.history — rebuilding it from raw rows would overwrite
|
||||
# parsed multimodal content and the _db_id edit/delete keys it just set.
|
||||
if not history_dict:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
@@ -268,17 +227,10 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
.order_by(DbChatMessage.timestamp)
|
||||
.all()
|
||||
)
|
||||
db_history = []
|
||||
for m in db_messages:
|
||||
db_history.append(_db_history_entry(m))
|
||||
if db_history:
|
||||
# Rebuild in-memory history from the full set so hidden
|
||||
# messages (e.g. compaction summaries) are kept for AI context.
|
||||
_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
|
||||
if not (m.get("metadata") or {}).get("hidden")
|
||||
entry for entry in (_db_history_entry(m) for m in db_messages)
|
||||
if not (entry.get("metadata") or {}).get("hidden")
|
||||
]
|
||||
except Exception as e:
|
||||
logger.error(f"DB fallback failed for {session_id}: {e}")
|
||||
@@ -645,8 +597,14 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
body = await request.json()
|
||||
keep_count = body.get("keep_count", 0)
|
||||
|
||||
# Get the source session
|
||||
source = session_manager.sessions.get(session_id)
|
||||
# Get the source session. keep_count indexes into source.history,
|
||||
# so this must go through get_session — reading the cache directly
|
||||
# forks an empty transcript out of a metadata-only session after a
|
||||
# restart (display pagination no longer hydrates it).
|
||||
try:
|
||||
source = session_manager.get_session(session_id)
|
||||
except KeyError:
|
||||
raise HTTPException(404, "Session not found")
|
||||
if not source:
|
||||
raise HTTPException(404, "Session not found")
|
||||
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""MCP route domain package (slice 2o, #4082/#4071).
|
||||
|
||||
Contains mcp_routes.py, migrated from the flat routes/ directory.
|
||||
Backward-compat shim at routes/mcp_routes.py re-exports from here.
|
||||
"""
|
||||
@@ -0,0 +1,697 @@
|
||||
# routes/mcp_routes.py
|
||||
"""MCP (Model Context Protocol) server management routes."""
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
import urllib.parse
|
||||
import html
|
||||
from pathlib import Path
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from fastapi.responses import RedirectResponse, HTMLResponse
|
||||
import logging
|
||||
import httpx
|
||||
|
||||
from core.database import McpServer, SessionLocal
|
||||
from core.middleware import require_admin
|
||||
from src.constants import DATA_DIR, MCP_OAUTH_DIR
|
||||
from src.mcp_manager import McpManager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/api/mcp", tags=["mcp"])
|
||||
|
||||
|
||||
def _mcp_oauth_base_dir() -> Path:
|
||||
"""Directory that may contain OAuth files managed by Odysseus."""
|
||||
return Path(MCP_OAUTH_DIR).resolve(strict=False)
|
||||
|
||||
|
||||
def _resolve_mcp_oauth_path(raw_path, field_name: str) -> str:
|
||||
"""Resolve an MCP OAuth path and keep it under DATA_DIR/mcp_oauth."""
|
||||
raw = str(raw_path or "").strip()
|
||||
if not raw:
|
||||
return ""
|
||||
|
||||
base = _mcp_oauth_base_dir()
|
||||
path = Path(os.path.expanduser(raw))
|
||||
if not path.is_absolute():
|
||||
path = base / path
|
||||
resolved = path.resolve(strict=False)
|
||||
|
||||
try:
|
||||
resolved.relative_to(base)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
400,
|
||||
f"Invalid OAuth {field_name}: path must stay under {base}",
|
||||
) from exc
|
||||
return str(resolved)
|
||||
|
||||
|
||||
def _sanitize_mcp_oauth_config(oauth_cfg):
|
||||
"""Return an OAuth config copy with file paths confined to mcp_oauth."""
|
||||
if not oauth_cfg:
|
||||
return oauth_cfg
|
||||
if not isinstance(oauth_cfg, dict):
|
||||
return {}
|
||||
sanitized = dict(oauth_cfg)
|
||||
for field_name in ("keys_file", "token_file"):
|
||||
if sanitized.get(field_name):
|
||||
sanitized[field_name] = _resolve_mcp_oauth_path(
|
||||
sanitized[field_name],
|
||||
field_name,
|
||||
)
|
||||
return sanitized
|
||||
|
||||
|
||||
def _mcp_oauth_token_missing(oauth_cfg, *, strict: bool = True) -> bool:
|
||||
"""Check token existence without letting legacy bad paths break listing."""
|
||||
if not isinstance(oauth_cfg, dict):
|
||||
return False
|
||||
try:
|
||||
token_file = _resolve_mcp_oauth_path(oauth_cfg.get("token_file", ""), "token_file")
|
||||
except HTTPException:
|
||||
if strict:
|
||||
raise
|
||||
logger.warning("Ignoring MCP OAuth config with unsafe token_file")
|
||||
return True
|
||||
return bool(token_file and not os.path.exists(token_file))
|
||||
|
||||
|
||||
def _apply_mcp_oauth_env(env: dict, oauth_cfg) -> None:
|
||||
"""Pass sanitized Gmail package paths to MCP servers that honor them."""
|
||||
if not oauth_cfg or not isinstance(env, dict):
|
||||
return
|
||||
keys_file = oauth_cfg.get("keys_file")
|
||||
token_file = oauth_cfg.get("token_file")
|
||||
if keys_file:
|
||||
env["GMAIL_OAUTH_PATH"] = keys_file
|
||||
if token_file:
|
||||
env["GMAIL_CREDENTIALS_PATH"] = token_file
|
||||
|
||||
|
||||
def _load_disabled_map():
|
||||
"""Load per-server disabled tool sets from DB."""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
disabled_map = {}
|
||||
for srv in db.query(McpServer).all():
|
||||
if srv.disabled_tools:
|
||||
try:
|
||||
names = json.loads(srv.disabled_tools)
|
||||
if names:
|
||||
disabled_map[srv.id] = set(names)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
return disabled_map
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _mcp_oauth_redirect_uri() -> str:
|
||||
"""Shared callback URL for legacy Google and generic MCP OAuth flows."""
|
||||
from src.mcp_oauth import REDIRECT_URI
|
||||
return REDIRECT_URI
|
||||
|
||||
|
||||
def setup_mcp_routes(mcp_manager: McpManager):
|
||||
"""Setup MCP routes with the provided manager."""
|
||||
|
||||
@router.get("/servers")
|
||||
def list_servers(request: Request):
|
||||
"""List all configured MCP servers with connection status."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
servers = db.query(McpServer).all()
|
||||
result = []
|
||||
for srv in servers:
|
||||
status = mcp_manager.get_server_status(srv.id)
|
||||
oauth_cfg = json.loads(srv.oauth_config) if srv.oauth_config else None
|
||||
needs_oauth = False
|
||||
if oauth_cfg:
|
||||
needs_oauth = _mcp_oauth_token_missing(oauth_cfg, strict=False)
|
||||
disabled_list = json.loads(srv.disabled_tools) if srv.disabled_tools else []
|
||||
total_tools = status.get("tool_count", 0)
|
||||
result.append({
|
||||
"id": srv.id,
|
||||
"name": srv.name,
|
||||
"transport": srv.transport,
|
||||
"command": srv.command,
|
||||
"args": json.loads(srv.args) if srv.args else [],
|
||||
"env": json.loads(srv.env) if srv.env else {},
|
||||
"url": srv.url,
|
||||
"is_enabled": srv.is_enabled,
|
||||
"status": status.get("status", "disconnected"),
|
||||
"tool_count": total_tools,
|
||||
"disabled_tool_count": len(disabled_list),
|
||||
"enabled_tool_count": max(0, total_tools - len(disabled_list)),
|
||||
"error": status.get("error"),
|
||||
"auth_url": status.get("auth_url"),
|
||||
"has_oauth": oauth_cfg is not None,
|
||||
"needs_oauth": needs_oauth,
|
||||
})
|
||||
return result
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.post("/servers")
|
||||
async def add_server(
|
||||
request: Request,
|
||||
name: str = Form(...),
|
||||
transport: str = Form("stdio"),
|
||||
command: str = Form(None),
|
||||
args: str = Form("[]"),
|
||||
env: str = Form("{}"),
|
||||
url: str = Form(None),
|
||||
oauth_file: str = Form(None),
|
||||
oauth_config: str = Form(None),
|
||||
):
|
||||
"""Add a new MCP server config and attempt connection. Admin-only:
|
||||
registering a stdio server is equivalent to executing arbitrary
|
||||
binaries on the host."""
|
||||
require_admin(request)
|
||||
server_id = str(uuid.uuid4())[:8]
|
||||
|
||||
# Validate
|
||||
if transport == "stdio" and not command:
|
||||
raise HTTPException(400, "command is required for stdio transport")
|
||||
if transport == "sse" and not url:
|
||||
raise HTTPException(400, "url is required for SSE transport")
|
||||
if transport == "http" and not url:
|
||||
raise HTTPException(400, "url is required for HTTP transport")
|
||||
|
||||
# Parse JSON fields
|
||||
try:
|
||||
parsed_args = json.loads(args) if args else []
|
||||
except json.JSONDecodeError:
|
||||
parsed_args = []
|
||||
try:
|
||||
parsed_env = json.loads(env) if env else {}
|
||||
except json.JSONDecodeError:
|
||||
parsed_env = {}
|
||||
if not isinstance(parsed_env, dict):
|
||||
parsed_env = {}
|
||||
|
||||
# Parse OAuth config
|
||||
parsed_oauth_config = None
|
||||
if oauth_config:
|
||||
try:
|
||||
parsed_oauth_config = _sanitize_mcp_oauth_config(json.loads(oauth_config))
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
_apply_mcp_oauth_env(parsed_env, parsed_oauth_config)
|
||||
|
||||
# Write OAuth credentials file if provided (for Google MCP servers)
|
||||
logger.info(f"MCP add_server: oauth_file={oauth_file!r}")
|
||||
if oauth_file:
|
||||
try:
|
||||
oauth_data = json.loads(oauth_file)
|
||||
oauth_dir = _resolve_mcp_oauth_path(oauth_data.get("dir", ""), "dir")
|
||||
oauth_filename = oauth_data.get("filename", "")
|
||||
client_id = oauth_data.get("client_id", "")
|
||||
client_secret = oauth_data.get("client_secret", "")
|
||||
if oauth_dir and oauth_filename and client_id and client_secret:
|
||||
filepath = _resolve_mcp_oauth_path(
|
||||
Path(oauth_dir) / str(oauth_filename),
|
||||
"filename",
|
||||
)
|
||||
os.makedirs(os.path.dirname(filepath), exist_ok=True)
|
||||
creds = {
|
||||
"installed": {
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"redirect_uris": ["http://localhost"],
|
||||
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
|
||||
"token_uri": "https://accounts.google.com/o/oauth2/token",
|
||||
}
|
||||
}
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
json.dump(creds, f, indent=2)
|
||||
logger.info(f"Wrote OAuth credentials to {filepath}")
|
||||
parsed_env.pop("GOOGLE_CLIENT_ID", None)
|
||||
parsed_env.pop("GOOGLE_CLIENT_SECRET", None)
|
||||
except (json.JSONDecodeError, OSError) as e:
|
||||
logger.warning(f"Failed to write OAuth file: {e}")
|
||||
|
||||
# Save to DB
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = McpServer(
|
||||
id=server_id,
|
||||
name=name,
|
||||
transport=transport,
|
||||
command=command,
|
||||
args=json.dumps(parsed_args),
|
||||
env=json.dumps(parsed_env),
|
||||
url=url,
|
||||
is_enabled=True,
|
||||
oauth_config=json.dumps(parsed_oauth_config) if parsed_oauth_config else None,
|
||||
)
|
||||
db.add(srv)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# Check if OAuth token already exists — skip connection attempt if not
|
||||
needs_oauth = False
|
||||
if parsed_oauth_config:
|
||||
needs_oauth = _mcp_oauth_token_missing(parsed_oauth_config)
|
||||
|
||||
connected = False
|
||||
if not needs_oauth:
|
||||
connected = await mcp_manager.connect_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport=transport,
|
||||
command=command,
|
||||
args=parsed_args,
|
||||
env=parsed_env,
|
||||
url=url,
|
||||
)
|
||||
|
||||
status = mcp_manager.get_server_status(server_id)
|
||||
needs_auth = status.get("status") == "needs_auth"
|
||||
return {
|
||||
"id": server_id,
|
||||
"name": name,
|
||||
"connected": connected,
|
||||
"status": "needs_oauth" if needs_oauth else status.get("status", "disconnected"),
|
||||
"tool_count": status.get("tool_count", 0),
|
||||
"error": "OAuth authorization required" if needs_oauth else status.get("error"),
|
||||
"needs_oauth": needs_oauth,
|
||||
"needs_auth": needs_auth,
|
||||
"auth_url": status.get("auth_url"),
|
||||
}
|
||||
|
||||
@router.post("/servers/{server_id}/reconnect")
|
||||
async def reconnect_server(server_id: str, request: Request):
|
||||
"""Reconnect to an MCP server."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
|
||||
await mcp_manager.disconnect_server(server_id)
|
||||
|
||||
args = json.loads(srv.args) if srv.args else []
|
||||
env = json.loads(srv.env) if srv.env else {}
|
||||
connected = await mcp_manager.connect_server(
|
||||
server_id=server_id,
|
||||
name=srv.name,
|
||||
transport=srv.transport,
|
||||
command=srv.command,
|
||||
args=args,
|
||||
env=env,
|
||||
url=srv.url,
|
||||
)
|
||||
|
||||
status = mcp_manager.get_server_status(server_id)
|
||||
return {
|
||||
"connected": connected,
|
||||
"status": status.get("status", "disconnected"),
|
||||
"tool_count": status.get("tool_count", 0),
|
||||
"error": status.get("error"),
|
||||
"auth_url": status.get("auth_url"),
|
||||
"needs_auth": status.get("status") == "needs_auth",
|
||||
}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.patch("/servers/{server_id}")
|
||||
async def toggle_server(server_id: str, request: Request, is_enabled: str = Form(...)):
|
||||
"""Enable or disable an MCP server."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
|
||||
enabled = str(is_enabled).lower() == "true"
|
||||
srv.is_enabled = enabled
|
||||
db.commit()
|
||||
|
||||
if enabled:
|
||||
args = json.loads(srv.args) if srv.args else []
|
||||
env = json.loads(srv.env) if srv.env else {}
|
||||
await mcp_manager.connect_server(
|
||||
server_id=server_id,
|
||||
name=srv.name,
|
||||
transport=srv.transport,
|
||||
command=srv.command,
|
||||
args=args,
|
||||
env=env,
|
||||
url=srv.url,
|
||||
)
|
||||
else:
|
||||
await mcp_manager.disconnect_server(server_id)
|
||||
|
||||
return {"id": server_id, "is_enabled": enabled}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.delete("/servers/{server_id}")
|
||||
async def delete_server(server_id: str, request: Request):
|
||||
"""Remove an MCP server."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
|
||||
await mcp_manager.disconnect_server(server_id)
|
||||
|
||||
db.delete(srv)
|
||||
db.commit()
|
||||
return {"status": "deleted"}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.get("/tools")
|
||||
def list_tools(request: Request):
|
||||
"""List all discovered MCP tools across all connected servers."""
|
||||
require_admin(request)
|
||||
disabled_map = _load_disabled_map()
|
||||
return mcp_manager.get_all_tools(disabled_map)
|
||||
|
||||
@router.get("/servers/{server_id}/tools")
|
||||
def list_server_tools(server_id: str, request: Request):
|
||||
"""List all tools for a specific MCP server with enabled/disabled state."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
disabled_list = json.loads(srv.disabled_tools) if srv.disabled_tools else []
|
||||
disabled_set = set(disabled_list)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
all_tools = mcp_manager.get_all_tools()
|
||||
server_tools = [t for t in all_tools if t["server_id"] == server_id]
|
||||
for t in server_tools:
|
||||
t["is_disabled"] = t["name"] in disabled_set
|
||||
return server_tools
|
||||
|
||||
@router.patch("/servers/{server_id}/tools")
|
||||
async def update_disabled_tools(server_id: str, request: Request):
|
||||
"""Bulk update disabled tools list for a server.
|
||||
|
||||
Expects JSON body: {"disabled": ["tool_name_1", "tool_name_2"]}
|
||||
"""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
|
||||
body = await request.json()
|
||||
disabled = body.get("disabled", [])
|
||||
if not isinstance(disabled, list):
|
||||
raise HTTPException(400, "disabled must be a list of tool names")
|
||||
|
||||
srv.disabled_tools = json.dumps(disabled) if disabled else None
|
||||
db.commit()
|
||||
|
||||
return {"id": server_id, "disabled_count": len(disabled)}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# ── OAuth flow for Google MCP servers ──────────────────────────
|
||||
|
||||
@router.get("/oauth/authorize/{server_id}")
|
||||
def oauth_authorize(server_id: str, request: Request):
|
||||
"""Show OAuth authorization page with Google sign-in link."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
if not srv.oauth_config:
|
||||
raise HTTPException(400, "Server has no OAuth config")
|
||||
|
||||
oauth_cfg = _sanitize_mcp_oauth_config(json.loads(srv.oauth_config))
|
||||
keys_file = oauth_cfg.get("keys_file", "")
|
||||
if not keys_file or not os.path.exists(keys_file):
|
||||
raise HTTPException(400, "OAuth keys file not found")
|
||||
|
||||
with open(keys_file, encoding="utf-8") as f:
|
||||
keys_data = json.load(f)
|
||||
keys = keys_data.get("installed") or keys_data.get("web")
|
||||
if not keys:
|
||||
raise HTTPException(400, "Invalid OAuth keys file format")
|
||||
|
||||
client_id = keys["client_id"]
|
||||
scopes = oauth_cfg.get("scopes", [])
|
||||
|
||||
# For Desktop App creds, default to localhost — the user will
|
||||
# paste the resulting URL back if they're on a different device.
|
||||
redirect_uri = _mcp_oauth_redirect_uri()
|
||||
|
||||
params = {
|
||||
"client_id": client_id,
|
||||
"redirect_uri": redirect_uri,
|
||||
"response_type": "code",
|
||||
"scope": " ".join(scopes),
|
||||
"access_type": "offline",
|
||||
"prompt": "consent",
|
||||
"state": server_id,
|
||||
}
|
||||
auth_url = "https://accounts.google.com/o/oauth2/v2/auth?" + urllib.parse.urlencode(params)
|
||||
|
||||
# Determine if user is accessing from the same machine
|
||||
host = request.headers.get("host", "")
|
||||
is_local = host.startswith("localhost") or host.startswith("127.0.0.1")
|
||||
|
||||
if is_local:
|
||||
# Same machine — just redirect, callback will work directly
|
||||
return RedirectResponse(auth_url)
|
||||
else:
|
||||
# Remote device — show paste-back page
|
||||
return HTMLResponse(_oauth_authorize_page(auth_url, server_id, host, redirect_uri))
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.get("/oauth/callback")
|
||||
async def oauth_callback(code: str, state: str, request: Request):
|
||||
"""Handle OAuth callback. Generic MCP OAuth flows resolve via the
|
||||
pending-state registry; Google flows fall through to the legacy path."""
|
||||
require_admin(request)
|
||||
from src.mcp_oauth import resolve_pending
|
||||
if resolve_pending(state, code):
|
||||
return HTMLResponse(_oauth_result_page(
|
||||
"Authorization Successful",
|
||||
"The MCP server is connecting. You can close this window and return to Odysseus.",
|
||||
success=True,
|
||||
))
|
||||
# Legacy Google path: state is the server_id
|
||||
return await _exchange_and_connect(state, code, request)
|
||||
|
||||
@router.post("/oauth/exchange/{server_id}")
|
||||
async def oauth_exchange(server_id: str, request: Request, callback_url: str = Form(...)):
|
||||
"""Manual code exchange — user pastes the callback URL from their browser."""
|
||||
require_admin(request)
|
||||
try:
|
||||
parsed = urllib.parse.urlparse(callback_url)
|
||||
params = urllib.parse.parse_qs(parsed.query)
|
||||
code = params.get("code", [None])[0]
|
||||
if not code:
|
||||
return HTMLResponse(_oauth_result_page("Error", "No authorization code found in the URL. Make sure you copied the full URL from your browser."), status_code=400)
|
||||
except Exception:
|
||||
return HTMLResponse(_oauth_result_page("Error", "Invalid URL format."), status_code=400)
|
||||
|
||||
# Generic MCP OAuth: if the pasted URL carries a state we are waiting on,
|
||||
# resolve it directly (the background connect finishes the handshake).
|
||||
state = params.get("state", [None])[0]
|
||||
from src.mcp_oauth import resolve_pending
|
||||
if state and resolve_pending(state, code):
|
||||
return HTMLResponse(_oauth_result_page(
|
||||
"Authorization Successful",
|
||||
"The MCP server is connecting. You can close this window and return to Odysseus.",
|
||||
success=True,
|
||||
))
|
||||
|
||||
return await _exchange_and_connect(server_id, code, request)
|
||||
|
||||
async def _exchange_and_connect(server_id: str, code: str, request: Request):
|
||||
"""Exchange auth code for tokens and connect the MCP server."""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
return HTMLResponse(_oauth_result_page("Error", "Server not found."), status_code=404)
|
||||
if not srv.oauth_config:
|
||||
return HTMLResponse(_oauth_result_page("Error", "No OAuth config."), status_code=400)
|
||||
|
||||
oauth_cfg = _sanitize_mcp_oauth_config(json.loads(srv.oauth_config))
|
||||
keys_file = oauth_cfg.get("keys_file", "")
|
||||
token_file = oauth_cfg.get("token_file", "")
|
||||
if not keys_file or not token_file:
|
||||
raise HTTPException(400, "OAuth keys/token file not configured")
|
||||
|
||||
with open(keys_file, encoding="utf-8") as f:
|
||||
keys_data = json.load(f)
|
||||
keys = keys_data.get("installed") or keys_data.get("web")
|
||||
client_id = keys["client_id"]
|
||||
client_secret = keys["client_secret"]
|
||||
|
||||
redirect_uri = _mcp_oauth_redirect_uri()
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
resp = await client.post(
|
||||
"https://oauth2.googleapis.com/token",
|
||||
data={
|
||||
"code": code,
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"redirect_uri": redirect_uri,
|
||||
"grant_type": "authorization_code",
|
||||
},
|
||||
)
|
||||
|
||||
if resp.status_code != 200:
|
||||
err = resp.text
|
||||
logger.error(f"OAuth token exchange failed: {err}")
|
||||
return HTMLResponse(_oauth_result_page("Authorization Failed", f"Google returned an error: {err}"), status_code=400)
|
||||
|
||||
tokens = resp.json()
|
||||
logger.info(f"OAuth tokens received for server {server_id}")
|
||||
|
||||
# Save tokens to the file the MCP package expects
|
||||
os.makedirs(os.path.dirname(token_file), exist_ok=True)
|
||||
with open(token_file, "w", encoding="utf-8") as f:
|
||||
json.dump(tokens, f, indent=2)
|
||||
logger.info(f"Saved OAuth tokens to {token_file}")
|
||||
|
||||
# Attempt to connect the MCP server now
|
||||
args = json.loads(srv.args) if srv.args else []
|
||||
env = json.loads(srv.env) if srv.env else {}
|
||||
connected = await mcp_manager.connect_server(
|
||||
server_id=server_id,
|
||||
name=srv.name,
|
||||
transport=srv.transport,
|
||||
command=srv.command,
|
||||
args=args,
|
||||
env=env,
|
||||
url=srv.url,
|
||||
)
|
||||
|
||||
if connected:
|
||||
status = mcp_manager.get_server_status(server_id)
|
||||
tool_count = status.get("tool_count", 0)
|
||||
return HTMLResponse(_oauth_result_page(
|
||||
"Authorization Successful",
|
||||
f"{srv.name} connected with {tool_count} tools. You can close this window.",
|
||||
success=True,
|
||||
))
|
||||
else:
|
||||
status = mcp_manager.get_server_status(server_id)
|
||||
return HTMLResponse(_oauth_result_page(
|
||||
"Authorized but Connection Failed",
|
||||
f"Tokens saved, but the server failed to connect: {status.get('error', 'unknown error')}. Try reconnecting from Settings.",
|
||||
))
|
||||
except HTTPException as e:
|
||||
logger.warning(f"OAuth callback rejected: {e.detail}")
|
||||
return HTMLResponse(_oauth_result_page("Error", str(e.detail)), status_code=e.status_code)
|
||||
except Exception as e:
|
||||
logger.exception(f"OAuth callback error: {e}")
|
||||
return HTMLResponse(_oauth_result_page("Error", str(e)), status_code=500)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return router
|
||||
|
||||
|
||||
def _oauth_authorize_page(
|
||||
auth_url: str,
|
||||
server_id: str,
|
||||
host: str,
|
||||
redirect_uri: str = "http://localhost:7000/api/mcp/oauth/callback",
|
||||
) -> str:
|
||||
"""Page with Google sign-in link and URL paste-back form for remote access."""
|
||||
# Escape values interpolated into the page: `host` comes from the request
|
||||
# Host header and `server_id` from the OAuth state — neither is trusted.
|
||||
auth_url = html.escape(auth_url, quote=True)
|
||||
server_id = html.escape(server_id, quote=True)
|
||||
host = html.escape(host, quote=True)
|
||||
redirect_uri = html.escape(redirect_uri, quote=True)
|
||||
return f"""<!DOCTYPE html>
|
||||
<html><head>
|
||||
<meta charset="UTF-8"><title>Authorize — Odysseus</title>
|
||||
<style>
|
||||
body {{ font-family: 'Fira Code', monospace; background: #0f0f0f; color: #e0e0e0;
|
||||
display: flex; justify-content: center; align-items: center; min-height: 100vh; }}
|
||||
.card {{ background: #1a1a1a; border: 1px solid #333; border-radius: 12px;
|
||||
padding: 2rem; max-width: 480px; text-align: center; }}
|
||||
h2 {{ color: #e06c75; margin-bottom: 0.5rem; font-size: 1.1rem; }}
|
||||
p {{ color: #aaa; font-size: 0.82rem; line-height: 1.6; margin: 0.8rem 0; }}
|
||||
.step {{ text-align: left; color: #ccc; font-size: 0.82rem; line-height: 1.7; margin: 1rem 0; }}
|
||||
.step b {{ color: #e06c75; }}
|
||||
a.auth-link {{
|
||||
display: inline-block; margin: 1rem 0; padding: 0.6rem 1.5rem;
|
||||
background: #e06c75; color: #fff; text-decoration: none; border-radius: 6px;
|
||||
font-weight: 600; font-size: 0.9rem;
|
||||
}}
|
||||
a.auth-link:hover {{ background: #c55; }}
|
||||
input[type=text] {{
|
||||
width: 100%; padding: 0.5rem; margin: 0.5rem 0;
|
||||
background: #0f0f0f; border: 1px solid #333; border-radius: 6px;
|
||||
color: #e0e0e0; font-family: 'Fira Code', monospace; font-size: 0.8rem;
|
||||
}}
|
||||
input:focus {{ outline: none; border-color: #e06c75; }}
|
||||
button {{
|
||||
padding: 0.5rem 1.5rem; border: none; border-radius: 6px;
|
||||
background: #e06c75; color: #fff; font-weight: 600; cursor: pointer;
|
||||
font-family: 'Fira Code', monospace; font-size: 0.85rem; margin-top: 0.3rem;
|
||||
}}
|
||||
button:hover {{ background: #c55; }}
|
||||
.divider {{ border-top: 1px solid #333; margin: 1.2rem 0; }}
|
||||
</style></head>
|
||||
<body><div class="card">
|
||||
<h2>Authorize Google Account</h2>
|
||||
<div class="step">
|
||||
<b>1.</b> Click the button below to sign in with Google<br>
|
||||
<b>2.</b> After approving, your browser will show an error page — that's normal<br>
|
||||
<b>3.</b> Copy the full URL from your browser's address bar<br>
|
||||
<b>4.</b> Paste it below and click Connect
|
||||
</div>
|
||||
<a class="auth-link" href="{auth_url}" target="_blank" rel="noopener">Sign in with Google</a>
|
||||
<div class="divider"></div>
|
||||
<form method="POST" action="http://{host}/api/mcp/oauth/exchange/{server_id}">
|
||||
<p>Paste the URL from your browser after signing in:</p>
|
||||
<input type="text" name="callback_url" placeholder="{redirect_uri}?code=..." required>
|
||||
<br><button type="submit">Connect</button>
|
||||
</form>
|
||||
</div></body></html>"""
|
||||
|
||||
|
||||
def _oauth_result_page(title: str, message: str, success: bool = False) -> str:
|
||||
"""Generate a simple HTML page for the OAuth result."""
|
||||
safe_title = html.escape(title)
|
||||
safe_message = html.escape(message)
|
||||
color = "#00661a" if success else "#e06c75"
|
||||
icon = "✓" if success else "✗"
|
||||
return f"""<!DOCTYPE html>
|
||||
<html><head>
|
||||
<meta charset="UTF-8"><title>{safe_title}</title>
|
||||
<style>
|
||||
body {{ font-family: 'Fira Code', monospace; background: #0f0f0f; color: #e0e0e0;
|
||||
display: flex; justify-content: center; align-items: center; min-height: 100vh; }}
|
||||
.card {{ background: #1a1a1a; border: 1px solid #333; border-radius: 12px;
|
||||
padding: 2rem; max-width: 420px; text-align: center; }}
|
||||
.icon {{ font-size: 3rem; color: {color}; margin-bottom: 1rem; }}
|
||||
h2 {{ color: {color}; margin-bottom: 0.5rem; font-size: 1.1rem; }}
|
||||
p {{ color: #aaa; font-size: 0.85rem; line-height: 1.5; }}
|
||||
</style></head>
|
||||
<body><div class="card">
|
||||
<div class="icon">{icon}</div>
|
||||
<h2>{safe_title}</h2>
|
||||
<p>{safe_message}</p>
|
||||
</div></body></html>"""
|
||||
+14
-693
@@ -1,697 +1,18 @@
|
||||
# routes/mcp_routes.py
|
||||
"""MCP (Model Context Protocol) server management routes."""
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
import urllib.parse
|
||||
import html
|
||||
from pathlib import Path
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from fastapi.responses import RedirectResponse, HTMLResponse
|
||||
import logging
|
||||
import httpx
|
||||
"""Backward-compat shim — canonical location is routes/mcp/mcp_routes.py.
|
||||
|
||||
from core.database import McpServer, SessionLocal
|
||||
from core.middleware import require_admin
|
||||
from src.constants import DATA_DIR, MCP_OAUTH_DIR
|
||||
from src.mcp_manager import McpManager
|
||||
This module is replaced in ``sys.modules`` by the canonical module object so
|
||||
that ``import routes.mcp_routes``, ``from routes.mcp_routes import X``,
|
||||
``importlib.import_module("routes.mcp_routes")``, the
|
||||
``sys.modules.pop("routes.mcp_routes")`` + re-import pattern in
|
||||
test_security_regressions.py, and the ``monkeypatch.setattr(mcp_routes,
|
||||
"MCP_OAUTH_DIR", ...)`` pattern all operate on the *same* object. This also
|
||||
makes ``mcp_routes.__file__`` resolve to the canonical file (which the
|
||||
source-introspection at line 839 reads). Keeps existing import paths working
|
||||
after slice 2o (#4082/#4071).
|
||||
"""
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
import sys as _sys
|
||||
|
||||
router = APIRouter(prefix="/api/mcp", tags=["mcp"])
|
||||
from routes.mcp import mcp_routes as _canonical # noqa: F401
|
||||
|
||||
|
||||
def _mcp_oauth_base_dir() -> Path:
|
||||
"""Directory that may contain OAuth files managed by Odysseus."""
|
||||
return Path(MCP_OAUTH_DIR).resolve(strict=False)
|
||||
|
||||
|
||||
def _resolve_mcp_oauth_path(raw_path, field_name: str) -> str:
|
||||
"""Resolve an MCP OAuth path and keep it under DATA_DIR/mcp_oauth."""
|
||||
raw = str(raw_path or "").strip()
|
||||
if not raw:
|
||||
return ""
|
||||
|
||||
base = _mcp_oauth_base_dir()
|
||||
path = Path(os.path.expanduser(raw))
|
||||
if not path.is_absolute():
|
||||
path = base / path
|
||||
resolved = path.resolve(strict=False)
|
||||
|
||||
try:
|
||||
resolved.relative_to(base)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
400,
|
||||
f"Invalid OAuth {field_name}: path must stay under {base}",
|
||||
) from exc
|
||||
return str(resolved)
|
||||
|
||||
|
||||
def _sanitize_mcp_oauth_config(oauth_cfg):
|
||||
"""Return an OAuth config copy with file paths confined to mcp_oauth."""
|
||||
if not oauth_cfg:
|
||||
return oauth_cfg
|
||||
if not isinstance(oauth_cfg, dict):
|
||||
return {}
|
||||
sanitized = dict(oauth_cfg)
|
||||
for field_name in ("keys_file", "token_file"):
|
||||
if sanitized.get(field_name):
|
||||
sanitized[field_name] = _resolve_mcp_oauth_path(
|
||||
sanitized[field_name],
|
||||
field_name,
|
||||
)
|
||||
return sanitized
|
||||
|
||||
|
||||
def _mcp_oauth_token_missing(oauth_cfg, *, strict: bool = True) -> bool:
|
||||
"""Check token existence without letting legacy bad paths break listing."""
|
||||
if not isinstance(oauth_cfg, dict):
|
||||
return False
|
||||
try:
|
||||
token_file = _resolve_mcp_oauth_path(oauth_cfg.get("token_file", ""), "token_file")
|
||||
except HTTPException:
|
||||
if strict:
|
||||
raise
|
||||
logger.warning("Ignoring MCP OAuth config with unsafe token_file")
|
||||
return True
|
||||
return bool(token_file and not os.path.exists(token_file))
|
||||
|
||||
|
||||
def _apply_mcp_oauth_env(env: dict, oauth_cfg) -> None:
|
||||
"""Pass sanitized Gmail package paths to MCP servers that honor them."""
|
||||
if not oauth_cfg or not isinstance(env, dict):
|
||||
return
|
||||
keys_file = oauth_cfg.get("keys_file")
|
||||
token_file = oauth_cfg.get("token_file")
|
||||
if keys_file:
|
||||
env["GMAIL_OAUTH_PATH"] = keys_file
|
||||
if token_file:
|
||||
env["GMAIL_CREDENTIALS_PATH"] = token_file
|
||||
|
||||
|
||||
def _load_disabled_map():
|
||||
"""Load per-server disabled tool sets from DB."""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
disabled_map = {}
|
||||
for srv in db.query(McpServer).all():
|
||||
if srv.disabled_tools:
|
||||
try:
|
||||
names = json.loads(srv.disabled_tools)
|
||||
if names:
|
||||
disabled_map[srv.id] = set(names)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
return disabled_map
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _mcp_oauth_redirect_uri() -> str:
|
||||
"""Shared callback URL for legacy Google and generic MCP OAuth flows."""
|
||||
from src.mcp_oauth import REDIRECT_URI
|
||||
return REDIRECT_URI
|
||||
|
||||
|
||||
def setup_mcp_routes(mcp_manager: McpManager):
|
||||
"""Setup MCP routes with the provided manager."""
|
||||
|
||||
@router.get("/servers")
|
||||
def list_servers(request: Request):
|
||||
"""List all configured MCP servers with connection status."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
servers = db.query(McpServer).all()
|
||||
result = []
|
||||
for srv in servers:
|
||||
status = mcp_manager.get_server_status(srv.id)
|
||||
oauth_cfg = json.loads(srv.oauth_config) if srv.oauth_config else None
|
||||
needs_oauth = False
|
||||
if oauth_cfg:
|
||||
needs_oauth = _mcp_oauth_token_missing(oauth_cfg, strict=False)
|
||||
disabled_list = json.loads(srv.disabled_tools) if srv.disabled_tools else []
|
||||
total_tools = status.get("tool_count", 0)
|
||||
result.append({
|
||||
"id": srv.id,
|
||||
"name": srv.name,
|
||||
"transport": srv.transport,
|
||||
"command": srv.command,
|
||||
"args": json.loads(srv.args) if srv.args else [],
|
||||
"env": json.loads(srv.env) if srv.env else {},
|
||||
"url": srv.url,
|
||||
"is_enabled": srv.is_enabled,
|
||||
"status": status.get("status", "disconnected"),
|
||||
"tool_count": total_tools,
|
||||
"disabled_tool_count": len(disabled_list),
|
||||
"enabled_tool_count": max(0, total_tools - len(disabled_list)),
|
||||
"error": status.get("error"),
|
||||
"auth_url": status.get("auth_url"),
|
||||
"has_oauth": oauth_cfg is not None,
|
||||
"needs_oauth": needs_oauth,
|
||||
})
|
||||
return result
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.post("/servers")
|
||||
async def add_server(
|
||||
request: Request,
|
||||
name: str = Form(...),
|
||||
transport: str = Form("stdio"),
|
||||
command: str = Form(None),
|
||||
args: str = Form("[]"),
|
||||
env: str = Form("{}"),
|
||||
url: str = Form(None),
|
||||
oauth_file: str = Form(None),
|
||||
oauth_config: str = Form(None),
|
||||
):
|
||||
"""Add a new MCP server config and attempt connection. Admin-only:
|
||||
registering a stdio server is equivalent to executing arbitrary
|
||||
binaries on the host."""
|
||||
require_admin(request)
|
||||
server_id = str(uuid.uuid4())[:8]
|
||||
|
||||
# Validate
|
||||
if transport == "stdio" and not command:
|
||||
raise HTTPException(400, "command is required for stdio transport")
|
||||
if transport == "sse" and not url:
|
||||
raise HTTPException(400, "url is required for SSE transport")
|
||||
if transport == "http" and not url:
|
||||
raise HTTPException(400, "url is required for HTTP transport")
|
||||
|
||||
# Parse JSON fields
|
||||
try:
|
||||
parsed_args = json.loads(args) if args else []
|
||||
except json.JSONDecodeError:
|
||||
parsed_args = []
|
||||
try:
|
||||
parsed_env = json.loads(env) if env else {}
|
||||
except json.JSONDecodeError:
|
||||
parsed_env = {}
|
||||
if not isinstance(parsed_env, dict):
|
||||
parsed_env = {}
|
||||
|
||||
# Parse OAuth config
|
||||
parsed_oauth_config = None
|
||||
if oauth_config:
|
||||
try:
|
||||
parsed_oauth_config = _sanitize_mcp_oauth_config(json.loads(oauth_config))
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
_apply_mcp_oauth_env(parsed_env, parsed_oauth_config)
|
||||
|
||||
# Write OAuth credentials file if provided (for Google MCP servers)
|
||||
logger.info(f"MCP add_server: oauth_file={oauth_file!r}")
|
||||
if oauth_file:
|
||||
try:
|
||||
oauth_data = json.loads(oauth_file)
|
||||
oauth_dir = _resolve_mcp_oauth_path(oauth_data.get("dir", ""), "dir")
|
||||
oauth_filename = oauth_data.get("filename", "")
|
||||
client_id = oauth_data.get("client_id", "")
|
||||
client_secret = oauth_data.get("client_secret", "")
|
||||
if oauth_dir and oauth_filename and client_id and client_secret:
|
||||
filepath = _resolve_mcp_oauth_path(
|
||||
Path(oauth_dir) / str(oauth_filename),
|
||||
"filename",
|
||||
)
|
||||
os.makedirs(os.path.dirname(filepath), exist_ok=True)
|
||||
creds = {
|
||||
"installed": {
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"redirect_uris": ["http://localhost"],
|
||||
"auth_uri": "https://accounts.google.com/o/oauth2/auth",
|
||||
"token_uri": "https://accounts.google.com/o/oauth2/token",
|
||||
}
|
||||
}
|
||||
with open(filepath, "w", encoding="utf-8") as f:
|
||||
json.dump(creds, f, indent=2)
|
||||
logger.info(f"Wrote OAuth credentials to {filepath}")
|
||||
parsed_env.pop("GOOGLE_CLIENT_ID", None)
|
||||
parsed_env.pop("GOOGLE_CLIENT_SECRET", None)
|
||||
except (json.JSONDecodeError, OSError) as e:
|
||||
logger.warning(f"Failed to write OAuth file: {e}")
|
||||
|
||||
# Save to DB
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = McpServer(
|
||||
id=server_id,
|
||||
name=name,
|
||||
transport=transport,
|
||||
command=command,
|
||||
args=json.dumps(parsed_args),
|
||||
env=json.dumps(parsed_env),
|
||||
url=url,
|
||||
is_enabled=True,
|
||||
oauth_config=json.dumps(parsed_oauth_config) if parsed_oauth_config else None,
|
||||
)
|
||||
db.add(srv)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# Check if OAuth token already exists — skip connection attempt if not
|
||||
needs_oauth = False
|
||||
if parsed_oauth_config:
|
||||
needs_oauth = _mcp_oauth_token_missing(parsed_oauth_config)
|
||||
|
||||
connected = False
|
||||
if not needs_oauth:
|
||||
connected = await mcp_manager.connect_server(
|
||||
server_id=server_id,
|
||||
name=name,
|
||||
transport=transport,
|
||||
command=command,
|
||||
args=parsed_args,
|
||||
env=parsed_env,
|
||||
url=url,
|
||||
)
|
||||
|
||||
status = mcp_manager.get_server_status(server_id)
|
||||
needs_auth = status.get("status") == "needs_auth"
|
||||
return {
|
||||
"id": server_id,
|
||||
"name": name,
|
||||
"connected": connected,
|
||||
"status": "needs_oauth" if needs_oauth else status.get("status", "disconnected"),
|
||||
"tool_count": status.get("tool_count", 0),
|
||||
"error": "OAuth authorization required" if needs_oauth else status.get("error"),
|
||||
"needs_oauth": needs_oauth,
|
||||
"needs_auth": needs_auth,
|
||||
"auth_url": status.get("auth_url"),
|
||||
}
|
||||
|
||||
@router.post("/servers/{server_id}/reconnect")
|
||||
async def reconnect_server(server_id: str, request: Request):
|
||||
"""Reconnect to an MCP server."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
|
||||
await mcp_manager.disconnect_server(server_id)
|
||||
|
||||
args = json.loads(srv.args) if srv.args else []
|
||||
env = json.loads(srv.env) if srv.env else {}
|
||||
connected = await mcp_manager.connect_server(
|
||||
server_id=server_id,
|
||||
name=srv.name,
|
||||
transport=srv.transport,
|
||||
command=srv.command,
|
||||
args=args,
|
||||
env=env,
|
||||
url=srv.url,
|
||||
)
|
||||
|
||||
status = mcp_manager.get_server_status(server_id)
|
||||
return {
|
||||
"connected": connected,
|
||||
"status": status.get("status", "disconnected"),
|
||||
"tool_count": status.get("tool_count", 0),
|
||||
"error": status.get("error"),
|
||||
"auth_url": status.get("auth_url"),
|
||||
"needs_auth": status.get("status") == "needs_auth",
|
||||
}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.patch("/servers/{server_id}")
|
||||
async def toggle_server(server_id: str, request: Request, is_enabled: str = Form(...)):
|
||||
"""Enable or disable an MCP server."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
|
||||
enabled = str(is_enabled).lower() == "true"
|
||||
srv.is_enabled = enabled
|
||||
db.commit()
|
||||
|
||||
if enabled:
|
||||
args = json.loads(srv.args) if srv.args else []
|
||||
env = json.loads(srv.env) if srv.env else {}
|
||||
await mcp_manager.connect_server(
|
||||
server_id=server_id,
|
||||
name=srv.name,
|
||||
transport=srv.transport,
|
||||
command=srv.command,
|
||||
args=args,
|
||||
env=env,
|
||||
url=srv.url,
|
||||
)
|
||||
else:
|
||||
await mcp_manager.disconnect_server(server_id)
|
||||
|
||||
return {"id": server_id, "is_enabled": enabled}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.delete("/servers/{server_id}")
|
||||
async def delete_server(server_id: str, request: Request):
|
||||
"""Remove an MCP server."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
|
||||
await mcp_manager.disconnect_server(server_id)
|
||||
|
||||
db.delete(srv)
|
||||
db.commit()
|
||||
return {"status": "deleted"}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.get("/tools")
|
||||
def list_tools(request: Request):
|
||||
"""List all discovered MCP tools across all connected servers."""
|
||||
require_admin(request)
|
||||
disabled_map = _load_disabled_map()
|
||||
return mcp_manager.get_all_tools(disabled_map)
|
||||
|
||||
@router.get("/servers/{server_id}/tools")
|
||||
def list_server_tools(server_id: str, request: Request):
|
||||
"""List all tools for a specific MCP server with enabled/disabled state."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
disabled_list = json.loads(srv.disabled_tools) if srv.disabled_tools else []
|
||||
disabled_set = set(disabled_list)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
all_tools = mcp_manager.get_all_tools()
|
||||
server_tools = [t for t in all_tools if t["server_id"] == server_id]
|
||||
for t in server_tools:
|
||||
t["is_disabled"] = t["name"] in disabled_set
|
||||
return server_tools
|
||||
|
||||
@router.patch("/servers/{server_id}/tools")
|
||||
async def update_disabled_tools(server_id: str, request: Request):
|
||||
"""Bulk update disabled tools list for a server.
|
||||
|
||||
Expects JSON body: {"disabled": ["tool_name_1", "tool_name_2"]}
|
||||
"""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
|
||||
body = await request.json()
|
||||
disabled = body.get("disabled", [])
|
||||
if not isinstance(disabled, list):
|
||||
raise HTTPException(400, "disabled must be a list of tool names")
|
||||
|
||||
srv.disabled_tools = json.dumps(disabled) if disabled else None
|
||||
db.commit()
|
||||
|
||||
return {"id": server_id, "disabled_count": len(disabled)}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# ── OAuth flow for Google MCP servers ──────────────────────────
|
||||
|
||||
@router.get("/oauth/authorize/{server_id}")
|
||||
def oauth_authorize(server_id: str, request: Request):
|
||||
"""Show OAuth authorization page with Google sign-in link."""
|
||||
require_admin(request)
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
raise HTTPException(404, "Server not found")
|
||||
if not srv.oauth_config:
|
||||
raise HTTPException(400, "Server has no OAuth config")
|
||||
|
||||
oauth_cfg = _sanitize_mcp_oauth_config(json.loads(srv.oauth_config))
|
||||
keys_file = oauth_cfg.get("keys_file", "")
|
||||
if not keys_file or not os.path.exists(keys_file):
|
||||
raise HTTPException(400, "OAuth keys file not found")
|
||||
|
||||
with open(keys_file, encoding="utf-8") as f:
|
||||
keys_data = json.load(f)
|
||||
keys = keys_data.get("installed") or keys_data.get("web")
|
||||
if not keys:
|
||||
raise HTTPException(400, "Invalid OAuth keys file format")
|
||||
|
||||
client_id = keys["client_id"]
|
||||
scopes = oauth_cfg.get("scopes", [])
|
||||
|
||||
# For Desktop App creds, default to localhost — the user will
|
||||
# paste the resulting URL back if they're on a different device.
|
||||
redirect_uri = _mcp_oauth_redirect_uri()
|
||||
|
||||
params = {
|
||||
"client_id": client_id,
|
||||
"redirect_uri": redirect_uri,
|
||||
"response_type": "code",
|
||||
"scope": " ".join(scopes),
|
||||
"access_type": "offline",
|
||||
"prompt": "consent",
|
||||
"state": server_id,
|
||||
}
|
||||
auth_url = "https://accounts.google.com/o/oauth2/v2/auth?" + urllib.parse.urlencode(params)
|
||||
|
||||
# Determine if user is accessing from the same machine
|
||||
host = request.headers.get("host", "")
|
||||
is_local = host.startswith("localhost") or host.startswith("127.0.0.1")
|
||||
|
||||
if is_local:
|
||||
# Same machine — just redirect, callback will work directly
|
||||
return RedirectResponse(auth_url)
|
||||
else:
|
||||
# Remote device — show paste-back page
|
||||
return HTMLResponse(_oauth_authorize_page(auth_url, server_id, host, redirect_uri))
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.get("/oauth/callback")
|
||||
async def oauth_callback(code: str, state: str, request: Request):
|
||||
"""Handle OAuth callback. Generic MCP OAuth flows resolve via the
|
||||
pending-state registry; Google flows fall through to the legacy path."""
|
||||
require_admin(request)
|
||||
from src.mcp_oauth import resolve_pending
|
||||
if resolve_pending(state, code):
|
||||
return HTMLResponse(_oauth_result_page(
|
||||
"Authorization Successful",
|
||||
"The MCP server is connecting. You can close this window and return to Odysseus.",
|
||||
success=True,
|
||||
))
|
||||
# Legacy Google path: state is the server_id
|
||||
return await _exchange_and_connect(state, code, request)
|
||||
|
||||
@router.post("/oauth/exchange/{server_id}")
|
||||
async def oauth_exchange(server_id: str, request: Request, callback_url: str = Form(...)):
|
||||
"""Manual code exchange — user pastes the callback URL from their browser."""
|
||||
require_admin(request)
|
||||
try:
|
||||
parsed = urllib.parse.urlparse(callback_url)
|
||||
params = urllib.parse.parse_qs(parsed.query)
|
||||
code = params.get("code", [None])[0]
|
||||
if not code:
|
||||
return HTMLResponse(_oauth_result_page("Error", "No authorization code found in the URL. Make sure you copied the full URL from your browser."), status_code=400)
|
||||
except Exception:
|
||||
return HTMLResponse(_oauth_result_page("Error", "Invalid URL format."), status_code=400)
|
||||
|
||||
# Generic MCP OAuth: if the pasted URL carries a state we are waiting on,
|
||||
# resolve it directly (the background connect finishes the handshake).
|
||||
state = params.get("state", [None])[0]
|
||||
from src.mcp_oauth import resolve_pending
|
||||
if state and resolve_pending(state, code):
|
||||
return HTMLResponse(_oauth_result_page(
|
||||
"Authorization Successful",
|
||||
"The MCP server is connecting. You can close this window and return to Odysseus.",
|
||||
success=True,
|
||||
))
|
||||
|
||||
return await _exchange_and_connect(server_id, code, request)
|
||||
|
||||
async def _exchange_and_connect(server_id: str, code: str, request: Request):
|
||||
"""Exchange auth code for tokens and connect the MCP server."""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
srv = db.query(McpServer).filter(McpServer.id == server_id).first()
|
||||
if not srv:
|
||||
return HTMLResponse(_oauth_result_page("Error", "Server not found."), status_code=404)
|
||||
if not srv.oauth_config:
|
||||
return HTMLResponse(_oauth_result_page("Error", "No OAuth config."), status_code=400)
|
||||
|
||||
oauth_cfg = _sanitize_mcp_oauth_config(json.loads(srv.oauth_config))
|
||||
keys_file = oauth_cfg.get("keys_file", "")
|
||||
token_file = oauth_cfg.get("token_file", "")
|
||||
if not keys_file or not token_file:
|
||||
raise HTTPException(400, "OAuth keys/token file not configured")
|
||||
|
||||
with open(keys_file, encoding="utf-8") as f:
|
||||
keys_data = json.load(f)
|
||||
keys = keys_data.get("installed") or keys_data.get("web")
|
||||
client_id = keys["client_id"]
|
||||
client_secret = keys["client_secret"]
|
||||
|
||||
redirect_uri = _mcp_oauth_redirect_uri()
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
resp = await client.post(
|
||||
"https://oauth2.googleapis.com/token",
|
||||
data={
|
||||
"code": code,
|
||||
"client_id": client_id,
|
||||
"client_secret": client_secret,
|
||||
"redirect_uri": redirect_uri,
|
||||
"grant_type": "authorization_code",
|
||||
},
|
||||
)
|
||||
|
||||
if resp.status_code != 200:
|
||||
err = resp.text
|
||||
logger.error(f"OAuth token exchange failed: {err}")
|
||||
return HTMLResponse(_oauth_result_page("Authorization Failed", f"Google returned an error: {err}"), status_code=400)
|
||||
|
||||
tokens = resp.json()
|
||||
logger.info(f"OAuth tokens received for server {server_id}")
|
||||
|
||||
# Save tokens to the file the MCP package expects
|
||||
os.makedirs(os.path.dirname(token_file), exist_ok=True)
|
||||
with open(token_file, "w", encoding="utf-8") as f:
|
||||
json.dump(tokens, f, indent=2)
|
||||
logger.info(f"Saved OAuth tokens to {token_file}")
|
||||
|
||||
# Attempt to connect the MCP server now
|
||||
args = json.loads(srv.args) if srv.args else []
|
||||
env = json.loads(srv.env) if srv.env else {}
|
||||
connected = await mcp_manager.connect_server(
|
||||
server_id=server_id,
|
||||
name=srv.name,
|
||||
transport=srv.transport,
|
||||
command=srv.command,
|
||||
args=args,
|
||||
env=env,
|
||||
url=srv.url,
|
||||
)
|
||||
|
||||
if connected:
|
||||
status = mcp_manager.get_server_status(server_id)
|
||||
tool_count = status.get("tool_count", 0)
|
||||
return HTMLResponse(_oauth_result_page(
|
||||
"Authorization Successful",
|
||||
f"{srv.name} connected with {tool_count} tools. You can close this window.",
|
||||
success=True,
|
||||
))
|
||||
else:
|
||||
status = mcp_manager.get_server_status(server_id)
|
||||
return HTMLResponse(_oauth_result_page(
|
||||
"Authorized but Connection Failed",
|
||||
f"Tokens saved, but the server failed to connect: {status.get('error', 'unknown error')}. Try reconnecting from Settings.",
|
||||
))
|
||||
except HTTPException as e:
|
||||
logger.warning(f"OAuth callback rejected: {e.detail}")
|
||||
return HTMLResponse(_oauth_result_page("Error", str(e.detail)), status_code=e.status_code)
|
||||
except Exception as e:
|
||||
logger.exception(f"OAuth callback error: {e}")
|
||||
return HTMLResponse(_oauth_result_page("Error", str(e)), status_code=500)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
return router
|
||||
|
||||
|
||||
def _oauth_authorize_page(
|
||||
auth_url: str,
|
||||
server_id: str,
|
||||
host: str,
|
||||
redirect_uri: str = "http://localhost:7000/api/mcp/oauth/callback",
|
||||
) -> str:
|
||||
"""Page with Google sign-in link and URL paste-back form for remote access."""
|
||||
# Escape values interpolated into the page: `host` comes from the request
|
||||
# Host header and `server_id` from the OAuth state — neither is trusted.
|
||||
auth_url = html.escape(auth_url, quote=True)
|
||||
server_id = html.escape(server_id, quote=True)
|
||||
host = html.escape(host, quote=True)
|
||||
redirect_uri = html.escape(redirect_uri, quote=True)
|
||||
return f"""<!DOCTYPE html>
|
||||
<html><head>
|
||||
<meta charset="UTF-8"><title>Authorize — Odysseus</title>
|
||||
<style>
|
||||
body {{ font-family: 'Fira Code', monospace; background: #0f0f0f; color: #e0e0e0;
|
||||
display: flex; justify-content: center; align-items: center; min-height: 100vh; }}
|
||||
.card {{ background: #1a1a1a; border: 1px solid #333; border-radius: 12px;
|
||||
padding: 2rem; max-width: 480px; text-align: center; }}
|
||||
h2 {{ color: #e06c75; margin-bottom: 0.5rem; font-size: 1.1rem; }}
|
||||
p {{ color: #aaa; font-size: 0.82rem; line-height: 1.6; margin: 0.8rem 0; }}
|
||||
.step {{ text-align: left; color: #ccc; font-size: 0.82rem; line-height: 1.7; margin: 1rem 0; }}
|
||||
.step b {{ color: #e06c75; }}
|
||||
a.auth-link {{
|
||||
display: inline-block; margin: 1rem 0; padding: 0.6rem 1.5rem;
|
||||
background: #e06c75; color: #fff; text-decoration: none; border-radius: 6px;
|
||||
font-weight: 600; font-size: 0.9rem;
|
||||
}}
|
||||
a.auth-link:hover {{ background: #c55; }}
|
||||
input[type=text] {{
|
||||
width: 100%; padding: 0.5rem; margin: 0.5rem 0;
|
||||
background: #0f0f0f; border: 1px solid #333; border-radius: 6px;
|
||||
color: #e0e0e0; font-family: 'Fira Code', monospace; font-size: 0.8rem;
|
||||
}}
|
||||
input:focus {{ outline: none; border-color: #e06c75; }}
|
||||
button {{
|
||||
padding: 0.5rem 1.5rem; border: none; border-radius: 6px;
|
||||
background: #e06c75; color: #fff; font-weight: 600; cursor: pointer;
|
||||
font-family: 'Fira Code', monospace; font-size: 0.85rem; margin-top: 0.3rem;
|
||||
}}
|
||||
button:hover {{ background: #c55; }}
|
||||
.divider {{ border-top: 1px solid #333; margin: 1.2rem 0; }}
|
||||
</style></head>
|
||||
<body><div class="card">
|
||||
<h2>Authorize Google Account</h2>
|
||||
<div class="step">
|
||||
<b>1.</b> Click the button below to sign in with Google<br>
|
||||
<b>2.</b> After approving, your browser will show an error page — that's normal<br>
|
||||
<b>3.</b> Copy the full URL from your browser's address bar<br>
|
||||
<b>4.</b> Paste it below and click Connect
|
||||
</div>
|
||||
<a class="auth-link" href="{auth_url}" target="_blank" rel="noopener">Sign in with Google</a>
|
||||
<div class="divider"></div>
|
||||
<form method="POST" action="http://{host}/api/mcp/oauth/exchange/{server_id}">
|
||||
<p>Paste the URL from your browser after signing in:</p>
|
||||
<input type="text" name="callback_url" placeholder="{redirect_uri}?code=..." required>
|
||||
<br><button type="submit">Connect</button>
|
||||
</form>
|
||||
</div></body></html>"""
|
||||
|
||||
|
||||
def _oauth_result_page(title: str, message: str, success: bool = False) -> str:
|
||||
"""Generate a simple HTML page for the OAuth result."""
|
||||
safe_title = html.escape(title)
|
||||
safe_message = html.escape(message)
|
||||
color = "#00661a" if success else "#e06c75"
|
||||
icon = "✓" if success else "✗"
|
||||
return f"""<!DOCTYPE html>
|
||||
<html><head>
|
||||
<meta charset="UTF-8"><title>{safe_title}</title>
|
||||
<style>
|
||||
body {{ font-family: 'Fira Code', monospace; background: #0f0f0f; color: #e0e0e0;
|
||||
display: flex; justify-content: center; align-items: center; min-height: 100vh; }}
|
||||
.card {{ background: #1a1a1a; border: 1px solid #333; border-radius: 12px;
|
||||
padding: 2rem; max-width: 420px; text-align: center; }}
|
||||
.icon {{ font-size: 3rem; color: {color}; margin-bottom: 1rem; }}
|
||||
h2 {{ color: {color}; margin-bottom: 0.5rem; font-size: 1.1rem; }}
|
||||
p {{ color: #aaa; font-size: 0.85rem; line-height: 1.5; }}
|
||||
</style></head>
|
||||
<body><div class="card">
|
||||
<div class="icon">{icon}</div>
|
||||
<h2>{safe_title}</h2>
|
||||
<p>{safe_message}</p>
|
||||
</div></body></html>"""
|
||||
_sys.modules[__name__] = _canonical
|
||||
|
||||
+9
-33
@@ -46,10 +46,12 @@ _ENDPOINT_SETTING_FIELDS = {
|
||||
}
|
||||
|
||||
_ENDPOINT_FALLBACK_FIELDS = {
|
||||
"default_model_fallbacks": "Default Model Fallbacks",
|
||||
"foreground_model_fallbacks": "Foreground Model Fallbacks",
|
||||
"utility_model_fallbacks": "Utility Model Fallbacks",
|
||||
"vision_model_fallbacks": "Vision Model Fallbacks",
|
||||
}
|
||||
# `default_model_fallbacks` is intentionally absent. The legacy data remains
|
||||
# stored as-is even when an endpoint is removed, but no longer affects routing.
|
||||
|
||||
|
||||
def _speech_settings_using_endpoint(settings: dict, ep_id: str) -> list:
|
||||
@@ -179,7 +181,12 @@ def _clear_user_pref_endpoint_refs(all_prefs: dict, ep_id: str) -> int:
|
||||
if not isinstance(all_prefs, dict):
|
||||
return 0
|
||||
users = all_prefs.get("_users")
|
||||
pref_sets = users.values() if isinstance(users, dict) else [all_prefs]
|
||||
# A mixed store can contain auth-disabled foreground policy at the root
|
||||
# alongside named-owner preferences. Both are active namespaces; legacy
|
||||
# `default_model_fallbacks` remains untouched by the field allowlist.
|
||||
pref_sets = [all_prefs]
|
||||
if isinstance(users, dict):
|
||||
pref_sets.extend(users.values())
|
||||
cleared_users = 0
|
||||
for prefs in pref_sets:
|
||||
if isinstance(prefs, dict) and _clear_endpoint_settings_for_endpoint(prefs, ep_id):
|
||||
@@ -2437,7 +2444,6 @@ def setup_model_routes(model_discovery):
|
||||
_user_prefs = _load_for_user(_user) or {}
|
||||
ep_id = (_user_prefs.get("default_endpoint_id") or "").strip()
|
||||
model = (_user_prefs.get("default_model") or "").strip()
|
||||
_fallbacks = _user_prefs.get("default_model_fallbacks") or []
|
||||
# If user has no personal default, fall back to global default
|
||||
# But only based on the "share_defaults_with_users" flag
|
||||
# (only if share_defaults_with_users is enabled)
|
||||
@@ -2446,12 +2452,9 @@ def setup_model_routes(model_discovery):
|
||||
ep_id = settings.get("default_endpoint_id", "")
|
||||
if not model:
|
||||
model = settings.get("default_model", "")
|
||||
if not _fallbacks:
|
||||
_fallbacks = settings.get("default_model_fallbacks") or []
|
||||
else:
|
||||
ep_id = settings.get("default_endpoint_id", "")
|
||||
model = settings.get("default_model", "")
|
||||
_fallbacks = settings.get("default_model_fallbacks") or []
|
||||
db = SessionLocal()
|
||||
try:
|
||||
ep = None
|
||||
@@ -2466,33 +2469,6 @@ def setup_model_routes(model_discovery):
|
||||
if _user and not _is_admin:
|
||||
ep_q = owner_filter(ep_q, ModelEndpoint, _user)
|
||||
ep = ep_q.first()
|
||||
# Configured fallback chain — when the chosen default endpoint is
|
||||
# gone/disabled, honor the user's configured `default_model_fallbacks`
|
||||
# in order BEFORE arbitrarily grabbing the first enabled endpoint.
|
||||
# (Previously this jumped straight to "first enabled", which is why
|
||||
# deleting/changing the main endpoint silently reassigned the default
|
||||
# chat to some unrelated endpoint instead of the fallback.)
|
||||
if not ep:
|
||||
for entry in _fallbacks:
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
fid = (entry.get("endpoint_id") or "").strip()
|
||||
if not fid:
|
||||
continue
|
||||
cand_q = db.query(ModelEndpoint).filter(
|
||||
ModelEndpoint.id == fid, ModelEndpoint.is_enabled == True
|
||||
)
|
||||
if _user and not _is_admin:
|
||||
cand_q = owner_filter(cand_q, ModelEndpoint, _user)
|
||||
cand = cand_q.first()
|
||||
if cand:
|
||||
ep = cand
|
||||
# Use the fallback entry's model. Reset even when empty
|
||||
# so we don't carry the prior endpoint's stale model onto
|
||||
# this fallback — the cached-models lookup below then
|
||||
# fills it from the fallback endpoint.
|
||||
model = (entry.get("model") or "").strip()
|
||||
break
|
||||
# Last resort: first enabled endpoint owned by THIS user. Do not
|
||||
# include null-owner/shared endpoints here: a brand-new user with
|
||||
# no explicit default should not auto-open a pending chat using an
|
||||
|
||||
+163
-92
@@ -1,11 +1,13 @@
|
||||
# routes/personal_routes.py
|
||||
"""Routes for personal documents management."""
|
||||
import asyncio
|
||||
import os
|
||||
import logging
|
||||
import shutil
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Tuple
|
||||
from fastapi import APIRouter, HTTPException, Query, Request, UploadFile, File, Depends
|
||||
from fastapi.concurrency import run_in_threadpool
|
||||
from src.request_models import DirectoryRequest
|
||||
from core.constants import BASE_DIR, PERSONAL_DIR, PERSONAL_UPLOADS_DIR
|
||||
from src.rag_singleton import get_rag_manager
|
||||
@@ -18,7 +20,6 @@ UPLOADS_DIR = PERSONAL_UPLOADS_DIR
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _personal_upload_dir_for_owner(owner: str | None, *, create: bool = True) -> str:
|
||||
"""Return the per-owner upload directory used for direct RAG uploads."""
|
||||
owner_segment = secure_filename((owner or "local").strip())[:80] or "local"
|
||||
@@ -141,6 +142,22 @@ def setup_personal_routes(personal_docs_manager, rag_manager, rag_available):
|
||||
"""
|
||||
router = APIRouter(prefix="/api/personal")
|
||||
|
||||
# Serializes directory index jobs across requests. Indexing runs in the
|
||||
# threadpool (#5558), so concurrent requests would otherwise run in parallel
|
||||
# and race PersonalDocsManager's unsynchronized list mutations and file
|
||||
# writes; before the threadpool move they serialized on the blocked event
|
||||
# loop, so one-at-a-time is behavior parity.
|
||||
#
|
||||
# An asyncio.Lock acquired in the async handler BEFORE offloading: a waiting
|
||||
# request parks on the event loop instead of pinning a threadpool worker (an
|
||||
# earlier threading.Lock taken INSIDE the worker meant queued jobs held pool
|
||||
# tokens while blocked, starving every other run_in_threadpool caller).
|
||||
# add/remove/reload all take this lock, so their mutations never interleave.
|
||||
# Per-router (not module-global) so each app binds it to its own event loop.
|
||||
# Scope is the single process: multi-worker deployments would need a shared
|
||||
# lock (out of scope for #5558).
|
||||
_index_job_lock = asyncio.Lock()
|
||||
|
||||
def _rag():
|
||||
"""Get the current RAG manager, retrying init if needed."""
|
||||
return get_rag_manager()
|
||||
@@ -172,8 +189,12 @@ def setup_personal_routes(personal_docs_manager, rag_manager, rag_available):
|
||||
return {"files": files, "directories": directories}
|
||||
|
||||
@router.post("/reload")
|
||||
def api_personal_reload(owner: str = Depends(require_user), _admin: None = Depends(require_admin)):
|
||||
personal_docs_manager.refresh_index()
|
||||
async def api_personal_reload(owner: str = Depends(require_user), _admin: None = Depends(require_admin)):
|
||||
# refresh_index() re-extracts text across every tracked directory —
|
||||
# blocking work. Take the shared job lock (so it cannot race an add /
|
||||
# remove) and run it off the event loop.
|
||||
async with _index_job_lock:
|
||||
await run_in_threadpool(personal_docs_manager.refresh_index)
|
||||
return {"ok": True, "count": len(personal_docs_manager.index)}
|
||||
|
||||
@router.post("/add_directory")
|
||||
@@ -207,12 +228,26 @@ def setup_personal_routes(personal_docs_manager, rag_manager, rag_available):
|
||||
# Use the RAGManager to index the directory
|
||||
rag = _rag()
|
||||
if rag:
|
||||
result = rag.index_personal_documents(directory, owner=owner)
|
||||
|
||||
def _index_directory():
|
||||
result = rag.index_personal_documents(directory, owner=owner)
|
||||
if result["success"]:
|
||||
# Also update the personal_docs_manager to track this
|
||||
# directory. Kept inside the offloaded call: it triggers
|
||||
# refresh_index(), which re-extracts text across tracked
|
||||
# directories.
|
||||
personal_docs_manager.add_directory(directory, index=False)
|
||||
return result
|
||||
|
||||
# Indexing walks, embeds, and stores the whole tree — minutes
|
||||
# on a real directory. The handler is async, so calling it
|
||||
# inline runs it on the event loop and every other request
|
||||
# queues behind it until it finishes (#5558). Serialize on the
|
||||
# async job lock BEFORE offloading so a queued request parks on
|
||||
# the loop instead of pinning a threadpool worker.
|
||||
async with _index_job_lock:
|
||||
result = await run_in_threadpool(_index_directory)
|
||||
|
||||
if result["success"]:
|
||||
# Also update the personal_docs_manager to track this directory
|
||||
personal_docs_manager.add_directory(directory, index=False)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"Successfully indexed {result['indexed_count']} chunks from {directory}",
|
||||
@@ -251,17 +286,25 @@ def setup_personal_routes(personal_docs_manager, rag_manager, rag_available):
|
||||
|
||||
logger.info(f"Removing directory from RAG: {directory}")
|
||||
|
||||
# Always remove from personal_docs_manager tracking
|
||||
if hasattr(personal_docs_manager, 'remove_directory'):
|
||||
personal_docs_manager.remove_directory(directory)
|
||||
|
||||
# Remove from RAG vector store (best-effort)
|
||||
rag = _rag()
|
||||
if rag:
|
||||
try:
|
||||
rag.remove_directory(directory)
|
||||
except Exception as e:
|
||||
logger.warning(f"RAG removal failed for directory {directory}: {e}")
|
||||
|
||||
def _remove_directory():
|
||||
# Always remove from personal_docs_manager tracking. This
|
||||
# mutates the same unsynchronized list/index an add job touches
|
||||
# and re-extracts text (refresh_index), so it is blocking work.
|
||||
if hasattr(personal_docs_manager, 'remove_directory'):
|
||||
personal_docs_manager.remove_directory(directory)
|
||||
# Remove from RAG vector store (best-effort).
|
||||
if rag:
|
||||
try:
|
||||
rag.remove_directory(directory)
|
||||
except Exception as e:
|
||||
logger.warning(f"RAG removal failed for directory {directory}: {e}")
|
||||
|
||||
# Same job lock as add/reload so remove cannot interleave with an
|
||||
# in-flight add; offloaded off the event loop.
|
||||
async with _index_job_lock:
|
||||
await run_in_threadpool(_remove_directory)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
@@ -289,54 +332,73 @@ def setup_personal_routes(personal_docs_manager, rag_manager, rag_available):
|
||||
total_failed = 0
|
||||
uploaded_files = []
|
||||
|
||||
for upload in files:
|
||||
try:
|
||||
file_path, stored_name, safe_name = _unique_personal_upload_path(upload_dir, upload.filename)
|
||||
content_bytes = await upload.read(PERSONAL_UPLOAD_MAX_BYTES + 1)
|
||||
if len(content_bytes) > PERSONAL_UPLOAD_MAX_BYTES:
|
||||
logger.warning(f"Rejected oversized personal upload: {upload.filename!r}")
|
||||
total_failed += 1
|
||||
continue
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(content_bytes)
|
||||
|
||||
ext = os.path.splitext(safe_name)[1].lower()
|
||||
if ext == ".pdf":
|
||||
from src.personal_docs import extract_pdf_text
|
||||
text = extract_pdf_text(file_path)
|
||||
else:
|
||||
text = content_bytes.decode("utf-8", errors="replace")
|
||||
|
||||
if not text or not text.strip():
|
||||
total_failed += 1
|
||||
continue
|
||||
|
||||
# Chunk and index
|
||||
chunks = rag._split_into_chunks(text, chunk_size=500)
|
||||
for i, chunk in enumerate(chunks):
|
||||
metadata = {
|
||||
"source": file_path,
|
||||
"filename": safe_name,
|
||||
"stored_filename": stored_name,
|
||||
"directory": upload_dir,
|
||||
"type": ext,
|
||||
"chunk_id": i,
|
||||
}
|
||||
if user:
|
||||
metadata["owner"] = user
|
||||
if rag.add_document(chunk, metadata):
|
||||
total_indexed += 1
|
||||
else:
|
||||
# Chunking, embedding and the tracking update are blocking work over the
|
||||
# same vector/tracking state add_directory mutates (#5634). Take the
|
||||
# shared job lock BEFORE offloading so a queued request parks on the loop
|
||||
# instead of pinning a threadpool worker, matching add_directory.
|
||||
# Read and process one capped payload at a time so a multi-file request
|
||||
# cannot retain len(files) * PERSONAL_UPLOAD_MAX_BYTES in memory.
|
||||
async with _index_job_lock:
|
||||
for upload in files:
|
||||
try:
|
||||
file_path, stored_name, safe_name = _unique_personal_upload_path(
|
||||
upload_dir, upload.filename
|
||||
)
|
||||
content_bytes = await upload.read(PERSONAL_UPLOAD_MAX_BYTES + 1)
|
||||
if len(content_bytes) > PERSONAL_UPLOAD_MAX_BYTES:
|
||||
logger.warning(f"Rejected oversized personal upload: {upload.filename!r}")
|
||||
total_failed += 1
|
||||
continue
|
||||
|
||||
uploaded_files.append(safe_name)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to upload/index {upload.filename}: {e}")
|
||||
total_failed += 1
|
||||
def _index_upload():
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(content_bytes)
|
||||
|
||||
# Track uploads directory
|
||||
if uploaded_files and hasattr(personal_docs_manager, "add_directory"):
|
||||
personal_docs_manager.add_directory(upload_dir, index=False)
|
||||
ext = os.path.splitext(safe_name)[1].lower()
|
||||
if ext == ".pdf":
|
||||
from src.personal_docs import extract_pdf_text
|
||||
text = extract_pdf_text(file_path)
|
||||
else:
|
||||
text = content_bytes.decode("utf-8", errors="replace")
|
||||
|
||||
if not text or not text.strip():
|
||||
return 0, 1, None
|
||||
|
||||
indexed = 0
|
||||
failed = 0
|
||||
chunks = rag._split_into_chunks(text, chunk_size=500)
|
||||
for i, chunk in enumerate(chunks):
|
||||
metadata = {
|
||||
"source": file_path,
|
||||
"filename": safe_name,
|
||||
"stored_filename": stored_name,
|
||||
"directory": upload_dir,
|
||||
"type": ext,
|
||||
"chunk_id": i,
|
||||
}
|
||||
if user:
|
||||
metadata["owner"] = user
|
||||
if rag.add_document(chunk, metadata):
|
||||
indexed += 1
|
||||
else:
|
||||
failed += 1
|
||||
return indexed, failed, safe_name
|
||||
|
||||
indexed, failed, uploaded_name = await run_in_threadpool(_index_upload)
|
||||
total_indexed += indexed
|
||||
total_failed += failed
|
||||
if uploaded_name:
|
||||
uploaded_files.append(uploaded_name)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to upload/index {upload.filename}: {e}")
|
||||
total_failed += 1
|
||||
|
||||
# Same transition, same lock: the tracking update must not land
|
||||
# while another job is mid-write over the same state.
|
||||
if uploaded_files and hasattr(personal_docs_manager, "add_directory"):
|
||||
await run_in_threadpool(
|
||||
personal_docs_manager.add_directory, upload_dir, index=False
|
||||
)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
@@ -349,38 +411,47 @@ def setup_personal_routes(personal_docs_manager, rag_manager, rag_available):
|
||||
async def delete_file_from_rag(filepath: str = Query(...), owner: str = Depends(require_user), _admin: None = Depends(require_admin)):
|
||||
"""Delete a specific file from RAG index and optionally from disk."""
|
||||
try:
|
||||
# Remove chunks from RAG vector store (best-effort)
|
||||
removed = 0
|
||||
rag = _rag()
|
||||
if rag:
|
||||
try:
|
||||
removed = rag.delete_by_source(filepath)
|
||||
except Exception as e:
|
||||
logger.warning(f"RAG removal failed for {filepath}: {e}")
|
||||
def _delete_file():
|
||||
# Remove chunks from RAG vector store (best-effort)
|
||||
removed = 0
|
||||
rag = _rag()
|
||||
if rag:
|
||||
try:
|
||||
removed = rag.delete_by_source(filepath)
|
||||
except Exception as e:
|
||||
logger.warning(f"RAG removal failed for {filepath}: {e}")
|
||||
|
||||
# Delete file from disk if it's in the caller's own uploads dir.
|
||||
# Scope to the per-owner subdir, not the shared uploads root, so one
|
||||
# admin can't delete another user's personal files by path.
|
||||
deleted_from_disk = False
|
||||
try:
|
||||
abs_target = os.path.realpath(filepath)
|
||||
base_abs = os.path.realpath(_personal_upload_dir_for_owner(owner, create=False))
|
||||
in_uploads = (
|
||||
abs_target == base_abs
|
||||
or os.path.commonpath([abs_target, base_abs]) == base_abs
|
||||
)
|
||||
except ValueError:
|
||||
# commonpath raises on mixed drives / non-comparable paths
|
||||
in_uploads = False
|
||||
if in_uploads and abs_target != base_abs:
|
||||
# Delete file from disk if it's in the caller's own uploads dir.
|
||||
# Scope to the per-owner subdir, not the shared uploads root, so one
|
||||
# admin can't delete another user's personal files by path.
|
||||
deleted_from_disk = False
|
||||
try:
|
||||
os.remove(abs_target)
|
||||
deleted_from_disk = True
|
||||
except FileNotFoundError:
|
||||
pass # already gone — race with another request or cleanup
|
||||
abs_target = os.path.realpath(filepath)
|
||||
base_abs = os.path.realpath(_personal_upload_dir_for_owner(owner, create=False))
|
||||
in_uploads = (
|
||||
abs_target == base_abs
|
||||
or os.path.commonpath([abs_target, base_abs]) == base_abs
|
||||
)
|
||||
except ValueError:
|
||||
# commonpath raises on mixed drives / non-comparable paths
|
||||
in_uploads = False
|
||||
if in_uploads and abs_target != base_abs:
|
||||
try:
|
||||
os.remove(abs_target)
|
||||
deleted_from_disk = True
|
||||
except FileNotFoundError:
|
||||
pass # already gone — race with another request or cleanup
|
||||
|
||||
# Exclude the file from the listing (persists across restarts)
|
||||
personal_docs_manager.exclude_file(filepath)
|
||||
# Exclude the file from the listing (persists across restarts)
|
||||
personal_docs_manager.exclude_file(filepath)
|
||||
return removed, deleted_from_disk
|
||||
|
||||
# Vector removal, the disk unlink and the exclusion write are one
|
||||
# transition over the same state add_directory mutates (#5634), and
|
||||
# all three block. Take the shared job lock BEFORE offloading, as
|
||||
# add_directory does.
|
||||
async with _index_job_lock:
|
||||
removed, deleted_from_disk = await run_in_threadpool(_delete_file)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
|
||||
+53
-19
@@ -1,12 +1,16 @@
|
||||
"""User preferences API — per-user key/value store backed by a JSON file."""
|
||||
import json
|
||||
import os
|
||||
from typing import Optional
|
||||
from fastapi import APIRouter, Request
|
||||
from core.atomic_io import atomic_write_json
|
||||
from src.auth_helpers import get_current_user
|
||||
from src.constants import USER_PREFS_FILE
|
||||
|
||||
PREFS_FILE = USER_PREFS_FILE
|
||||
_FOREGROUND_POLICY_KEYS = (
|
||||
"foreground_fallback_enabled",
|
||||
"foreground_model_fallbacks",
|
||||
)
|
||||
|
||||
|
||||
def _load():
|
||||
@@ -20,26 +24,33 @@ def _load():
|
||||
|
||||
|
||||
def _save(prefs):
|
||||
os.makedirs(os.path.dirname(PREFS_FILE) or ".", exist_ok=True)
|
||||
tmp = f"{PREFS_FILE}.tmp.{os.getpid()}"
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
json.dump(prefs, f, indent=2)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp, PREFS_FILE)
|
||||
atomic_write_json(PREFS_FILE, prefs, indent=2)
|
||||
|
||||
|
||||
def _load_for_user(user: Optional[str] = None) -> dict:
|
||||
"""Load preferences for a specific user."""
|
||||
all_prefs = _load()
|
||||
if "_users" in all_prefs:
|
||||
users = all_prefs.get("_users")
|
||||
if isinstance(users, dict):
|
||||
if user is None:
|
||||
# Auth disabled — return first user's prefs for backward compat
|
||||
users = all_prefs["_users"]
|
||||
return dict(next(iter(users.values()), {}))
|
||||
return dict(all_prefs["_users"].get(user, {}))
|
||||
# Legacy flat format — return as-is
|
||||
return dict(all_prefs)
|
||||
prefs = dict(next(iter(users.values()), {}))
|
||||
# Foreground fallback consent is never borrowed from a named
|
||||
# owner. Auth-disabled operation has a separate flat/root opt-in
|
||||
# that remains inert when authentication is enabled again.
|
||||
for key in _FOREGROUND_POLICY_KEYS:
|
||||
prefs.pop(key, None)
|
||||
if key in all_prefs:
|
||||
prefs[key] = all_prefs[key]
|
||||
return prefs
|
||||
prefs = users.get(user, {})
|
||||
return dict(prefs) if isinstance(prefs, dict) else {}
|
||||
# A legacy flat store belongs only to auth-disabled single-user mode.
|
||||
# Copying it into the first named user's new `_users` record during an
|
||||
# auth transition would silently transfer another user's preferences and,
|
||||
# critically, foreground fallback consent. Named owners therefore start
|
||||
# with an empty record and must write their own preferences explicitly.
|
||||
return dict(all_prefs) if user is None else {}
|
||||
|
||||
|
||||
def _save_for_user(user: Optional[str], prefs: dict):
|
||||
@@ -51,17 +62,40 @@ def _save_for_user(user: Optional[str], prefs: dict):
|
||||
# `prefs` flat would overwrite the whole `_users` map and destroy every
|
||||
# other user's preferences. Instead write back into the same (first)
|
||||
# slot _load_for_user(None) reads from, preserving the others.
|
||||
if "_users" in all_prefs:
|
||||
users = all_prefs["_users"]
|
||||
users = all_prefs.get("_users")
|
||||
if isinstance(users, dict):
|
||||
first_key = next(iter(users), None)
|
||||
if first_key is not None:
|
||||
users[first_key] = prefs
|
||||
existing_named = users.get(first_key)
|
||||
existing_named = (
|
||||
dict(existing_named)
|
||||
if isinstance(existing_named, dict)
|
||||
else {}
|
||||
)
|
||||
named_foreground = {
|
||||
key: existing_named[key]
|
||||
for key in _FOREGROUND_POLICY_KEYS
|
||||
if key in existing_named
|
||||
}
|
||||
users[first_key] = {
|
||||
key: value
|
||||
for key, value in prefs.items()
|
||||
if key not in _FOREGROUND_POLICY_KEYS
|
||||
}
|
||||
users[first_key].update(named_foreground)
|
||||
for key in _FOREGROUND_POLICY_KEYS:
|
||||
if key in prefs:
|
||||
all_prefs[key] = prefs[key]
|
||||
_save(all_prefs)
|
||||
return
|
||||
_save(prefs)
|
||||
return
|
||||
if "_users" not in all_prefs:
|
||||
all_prefs = {"_users": {}}
|
||||
if not isinstance(all_prefs.get("_users"), dict):
|
||||
# Preserve the flat single-user object as inert legacy data while
|
||||
# creating the first named-owner namespace. In particular, historical
|
||||
# fallback values must not be deleted or copied into the new owner.
|
||||
all_prefs = dict(all_prefs)
|
||||
all_prefs["_users"] = {}
|
||||
all_prefs["_users"][user] = prefs
|
||||
_save(all_prefs)
|
||||
|
||||
|
||||
@@ -801,15 +801,6 @@ def setup_session_routes(
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
@router.get("/history/{sid}")
|
||||
def get_history(request: Request, sid: str):
|
||||
_verify_session_owner(request, sid)
|
||||
try:
|
||||
session = session_manager.get_session(sid)
|
||||
except KeyError:
|
||||
raise HTTPException(404, f"Session {sid} not found")
|
||||
return {"history": [msg.to_dict() for msg in session.history]}
|
||||
|
||||
@router.get("/session/{sid}/export")
|
||||
def export_session(request: Request, sid: str, fmt: str = "md", filename: str = ""):
|
||||
"""Export conversation history as a downloadable file.
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Create/remove the switchable, non-default 'Demo' EmailAccount in Odysseus.
|
||||
"""Create/remove the switchable 'Demo' EmailAccount in Odysseus.
|
||||
|
||||
Mirrors the existing local-Dovecot account (localhost:31143, STARTTLS) but points
|
||||
at the throwaway demo@odysseus.local mailbox. Password is stored Fernet-encrypted
|
||||
@@ -20,7 +20,14 @@ from pathlib import Path
|
||||
ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
sys.path.insert(0, str(ROOT))
|
||||
|
||||
from core.database import SessionLocal, EmailAccount, Base, engine # noqa: E402
|
||||
from core.database import ( # noqa: E402
|
||||
Base,
|
||||
EmailAccount,
|
||||
SessionLocal,
|
||||
engine,
|
||||
lock_email_account_owner_mutations,
|
||||
)
|
||||
from sqlalchemy import or_ # noqa: E402
|
||||
from src.secret_storage import encrypt # noqa: E402
|
||||
|
||||
NAME = "Demo"
|
||||
@@ -31,18 +38,98 @@ IMAP_PASSWORD = "demodemo"
|
||||
OWNER = ""
|
||||
|
||||
|
||||
def setup() -> int:
|
||||
Base.metadata.create_all(bind=engine)
|
||||
def _owner_scope(query, owner: str):
|
||||
if owner:
|
||||
return query.filter(EmailAccount.owner == owner)
|
||||
return query.filter(or_(EmailAccount.owner == None, EmailAccount.owner == "")) # noqa: E711
|
||||
|
||||
|
||||
def _discover_demo_scopes() -> set[str]:
|
||||
db = SessionLocal()
|
||||
try:
|
||||
acct = db.query(EmailAccount).filter(
|
||||
EmailAccount.name == NAME, EmailAccount.imap_user == IMAP_USER
|
||||
).first()
|
||||
return {
|
||||
row.owner or ""
|
||||
for row in db.query(EmailAccount).filter(
|
||||
EmailAccount.name == NAME,
|
||||
EmailAccount.imap_user == IMAP_USER,
|
||||
).all()
|
||||
}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _lock_and_load_demo_rows(db, scopes: set[str]):
|
||||
"""Reload Demo rows under every observed owner lock."""
|
||||
scopes = set(scopes) or {OWNER}
|
||||
while True:
|
||||
lock_email_account_owner_mutations(db, *scopes)
|
||||
rows = (
|
||||
db.query(EmailAccount)
|
||||
.filter(
|
||||
EmailAccount.name == NAME,
|
||||
EmailAccount.imap_user == IMAP_USER,
|
||||
)
|
||||
.order_by(EmailAccount.created_at.asc(), EmailAccount.id.asc())
|
||||
.all()
|
||||
)
|
||||
current_scopes = {row.owner or "" for row in rows}
|
||||
if current_scopes.issubset(scopes) or db.get_bind().dialect.name == "sqlite":
|
||||
return rows
|
||||
db.rollback()
|
||||
scopes.update(current_scopes)
|
||||
|
||||
|
||||
def _promote_oldest_enabled(db, owner: str, excluded_ids: list[str]) -> None:
|
||||
remaining = _owner_scope(
|
||||
db.query(EmailAccount).filter(
|
||||
EmailAccount.enabled == True, # noqa: E712
|
||||
~EmailAccount.id.in_(excluded_ids),
|
||||
),
|
||||
owner,
|
||||
)
|
||||
if remaining.filter(EmailAccount.is_default == True).first() is not None: # noqa: E712
|
||||
return
|
||||
promote = remaining.order_by(
|
||||
EmailAccount.created_at.asc(), EmailAccount.id.asc()
|
||||
).first()
|
||||
if promote is not None:
|
||||
promote.is_default = True
|
||||
|
||||
|
||||
def setup() -> int:
|
||||
Base.metadata.create_all(bind=engine)
|
||||
scopes = _discover_demo_scopes() | {OWNER}
|
||||
db = SessionLocal()
|
||||
try:
|
||||
rows = _lock_and_load_demo_rows(db, scopes)
|
||||
acct = rows[0] if rows else None
|
||||
if acct is None:
|
||||
acct = EmailAccount(id=uuid.uuid4().hex, name=NAME)
|
||||
db.add(acct)
|
||||
old_scope = acct.owner or ""
|
||||
was_default = bool(acct.is_default)
|
||||
if old_scope != OWNER:
|
||||
# Move a non-default row first so the unique index cannot see two
|
||||
# defaults transiently while SQLAlchemy flushes the owner move and
|
||||
# old-scope promotion in separate UPDATE statements.
|
||||
acct.is_default = False
|
||||
acct.owner = OWNER
|
||||
db.flush()
|
||||
if was_default:
|
||||
_promote_oldest_enabled(db, old_scope, [acct.id])
|
||||
|
||||
target_default = _owner_scope(
|
||||
db.query(EmailAccount).filter(
|
||||
EmailAccount.id != acct.id,
|
||||
EmailAccount.is_default == True, # noqa: E712
|
||||
),
|
||||
OWNER,
|
||||
).first()
|
||||
acct.owner = OWNER
|
||||
acct.is_default = False # never default — user switches to it
|
||||
# Keep Demo non-default when a real default exists. If it is the only
|
||||
# enabled account, it must be default to preserve normal create
|
||||
# semantics and avoid leaving the owner partition without one.
|
||||
acct.is_default = target_default is None
|
||||
acct.enabled = True
|
||||
acct.imap_host = "localhost"
|
||||
acct.imap_port = 31143
|
||||
@@ -57,20 +144,27 @@ def setup() -> int:
|
||||
acct.smtp_password = encrypt(IMAP_PASSWORD)
|
||||
acct.from_address = IMAP_USER
|
||||
db.commit()
|
||||
print(f"'{NAME}' account ready (id={acct.id}, non-default, switchable).")
|
||||
state = "default" if acct.is_default else "non-default"
|
||||
print(f"'{NAME}' account ready (id={acct.id}, {state}, switchable).")
|
||||
return 0
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def teardown() -> int:
|
||||
scopes = _discover_demo_scopes()
|
||||
db = SessionLocal()
|
||||
try:
|
||||
rows = db.query(EmailAccount).filter(
|
||||
EmailAccount.name == NAME, EmailAccount.imap_user == IMAP_USER
|
||||
).all()
|
||||
rows = _lock_and_load_demo_rows(db, scopes)
|
||||
deleted_ids = [row.id for row in rows]
|
||||
default_scopes = {row.owner or "" for row in rows if row.is_default}
|
||||
for r in rows:
|
||||
db.delete(r)
|
||||
# Ensure the old default DELETE reaches the database before a
|
||||
# replacement UPDATE; the unique index is enforced per statement.
|
||||
db.flush()
|
||||
for owner in default_scopes:
|
||||
_promote_oldest_enabled(db, owner, deleted_ids)
|
||||
db.commit()
|
||||
print(f"removed {len(rows)} '{NAME}' account row(s).")
|
||||
return 0
|
||||
|
||||
@@ -0,0 +1,166 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Make retained SearXNG settings inherit defaults without replacing them."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import stat
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
from yaml.nodes import MappingNode
|
||||
from yaml.tokens import BlockMappingStartToken, FlowMappingStartToken
|
||||
|
||||
|
||||
_UTF8_BOM = b"\xef\xbb\xbf"
|
||||
|
||||
|
||||
def _parse_root_mapping(text: str) -> tuple[MappingNode | None, dict]:
|
||||
"""Parse settings with the same safe YAML semantics SearXNG uses."""
|
||||
try:
|
||||
loaded = yaml.safe_load(text)
|
||||
node = yaml.compose(text, Loader=yaml.SafeLoader)
|
||||
except yaml.YAMLError:
|
||||
raise ValueError("settings file is not valid single-document YAML") from None
|
||||
|
||||
if loaded is None and node is None:
|
||||
return None, {}
|
||||
if not isinstance(loaded, dict) or not isinstance(node, MappingNode):
|
||||
raise ValueError("settings root is not a mapping")
|
||||
return node, loaded
|
||||
|
||||
|
||||
def _flow_mapping_start(text: str) -> int:
|
||||
"""Return the root flow mapping's opening-brace character offset."""
|
||||
try:
|
||||
for token in yaml.scan(text, Loader=yaml.SafeLoader):
|
||||
if isinstance(token, FlowMappingStartToken):
|
||||
return token.start_mark.index
|
||||
except yaml.YAMLError:
|
||||
pass
|
||||
raise ValueError("flow-style settings mapping has no opening brace")
|
||||
|
||||
|
||||
def _newline_for(contents: bytes) -> bytes:
|
||||
first_lf = contents.find(b"\n")
|
||||
if first_lf > 0 and contents[first_lf - 1 : first_lf + 1] == b"\r\n":
|
||||
return b"\r\n"
|
||||
return b"\n"
|
||||
|
||||
|
||||
def _block_mapping_position(text: str, root: MappingNode | None) -> tuple[int, int]:
|
||||
"""Return a safe character offset and indent for a root block mapping key."""
|
||||
if root is None:
|
||||
return len(text), 0
|
||||
|
||||
try:
|
||||
for token in yaml.scan(text, Loader=yaml.SafeLoader):
|
||||
if not isinstance(token, BlockMappingStartToken):
|
||||
continue
|
||||
line_start = token.start_mark.index - token.start_mark.column
|
||||
if not text[line_start : token.start_mark.index].strip():
|
||||
return line_start, token.start_mark.column
|
||||
return root.end_mark.index, token.start_mark.column
|
||||
except yaml.YAMLError:
|
||||
pass
|
||||
return root.end_mark.index, root.start_mark.column
|
||||
|
||||
|
||||
def _add_block_default_inheritance(
|
||||
contents: bytes, text: str, root: MappingNode | None
|
||||
) -> bytes:
|
||||
newline = _newline_for(contents)
|
||||
character_offset, indent_width = _block_mapping_position(text, root)
|
||||
bom_length = len(_UTF8_BOM) if contents.startswith(_UTF8_BOM) else 0
|
||||
offset = bom_length + len(text[:character_offset].encode("utf-8"))
|
||||
separator = b""
|
||||
if offset not in (0, bom_length) and not contents[:offset].endswith((b"\n", b"\r")):
|
||||
separator = newline
|
||||
addition = (
|
||||
separator
|
||||
+ b" " * indent_width
|
||||
+ b"use_default_settings: true"
|
||||
+ newline
|
||||
)
|
||||
return contents[:offset] + addition + contents[offset:]
|
||||
|
||||
|
||||
def migrate_settings(path: Path) -> bool:
|
||||
"""Add the missing inheritance key atomically; return whether the file changed."""
|
||||
source_stat = path.lstat()
|
||||
if not stat.S_ISREG(source_stat.st_mode):
|
||||
raise ValueError(f"settings path is not a regular file: {path}")
|
||||
|
||||
contents = path.read_bytes()
|
||||
if not contents:
|
||||
return False
|
||||
|
||||
text = contents.decode("utf-8-sig")
|
||||
root, loaded = _parse_root_mapping(text)
|
||||
if "use_default_settings" in loaded:
|
||||
return False
|
||||
|
||||
if root is not None and root.flow_style:
|
||||
start = _flow_mapping_start(text)
|
||||
bom_length = len(_UTF8_BOM) if contents.startswith(_UTF8_BOM) else 0
|
||||
offset = bom_length + len(text[: start + 1].encode("utf-8"))
|
||||
separator = b", " if root.value else b""
|
||||
updated = (
|
||||
contents[:offset]
|
||||
+ b"use_default_settings: true"
|
||||
+ separator
|
||||
+ contents[offset:]
|
||||
)
|
||||
else:
|
||||
updated = _add_block_default_inheritance(contents, text, root)
|
||||
fd, temporary_name = tempfile.mkstemp(
|
||||
prefix=f".{path.name}.odysseus-", dir=path.parent
|
||||
)
|
||||
temporary = Path(temporary_name)
|
||||
try:
|
||||
# chmod before chown: the Compose cap set is `cap_drop: ALL` plus
|
||||
# CHOWN/SETGID/SETUID/DAC_OVERRIDE, with no FOWNER. Once the temporary
|
||||
# file belongs to searxng:searxng — which every retained settings file
|
||||
# does, because searxng's entrypoint chowns /etc/searxng — root can no
|
||||
# longer chmod it and the migration dies with EPERM.
|
||||
os.fchmod(fd, stat.S_IMODE(source_stat.st_mode))
|
||||
os.fchown(fd, source_stat.st_uid, source_stat.st_gid)
|
||||
with os.fdopen(fd, "wb") as handle:
|
||||
fd = -1
|
||||
handle.write(updated)
|
||||
handle.flush()
|
||||
os.fsync(handle.fileno())
|
||||
os.replace(temporary, path)
|
||||
directory_fd = os.open(path.parent, os.O_RDONLY | os.O_DIRECTORY)
|
||||
try:
|
||||
os.fsync(directory_fd)
|
||||
finally:
|
||||
os.close(directory_fd)
|
||||
finally:
|
||||
if fd >= 0:
|
||||
os.close(fd)
|
||||
temporary.unlink(missing_ok=True)
|
||||
return True
|
||||
|
||||
|
||||
def main(argv: list[str]) -> int:
|
||||
if len(argv) > 2:
|
||||
print(f"usage: {Path(argv[0]).name} [settings.yml]", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
path = Path(argv[1]) if len(argv) == 2 else Path("/etc/searxng/settings.yml")
|
||||
try:
|
||||
changed = migrate_settings(path)
|
||||
except (OSError, UnicodeError, ValueError) as exc:
|
||||
print(f"SearXNG settings migration failed: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
if changed:
|
||||
print("Added use_default_settings inheritance to retained SearXNG settings")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main(sys.argv))
|
||||
@@ -100,6 +100,18 @@ def _parse_scalar(raw: str) -> Any:
|
||||
if raw.lower() in ("null", "none", "~"):
|
||||
return None
|
||||
if (raw[0] == raw[-1]) and raw[0] in ("'", '"'):
|
||||
if raw[0] == '"':
|
||||
# _emit_scalar writes double-quoted scalars with json.dumps, so
|
||||
# decode the escapes instead of only stripping the quotes. Without
|
||||
# this, `\"` / `\\` / `\uXXXX` stayed verbatim in the value and the
|
||||
# next save escaped their backslashes again, doubling them on every
|
||||
# load/save cycle (issue #5210).
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except ValueError:
|
||||
# Hand-written file using escapes JSON rejects (e.g. a bare
|
||||
# Windows path). Keep the previous literal reading.
|
||||
pass
|
||||
return raw[1:-1]
|
||||
# Try number
|
||||
try:
|
||||
@@ -171,6 +183,26 @@ def parse_frontmatter(text: str) -> tuple[Dict[str, Any], str]:
|
||||
return fm, body
|
||||
|
||||
|
||||
# Characters that force a quoted scalar. The punctuation would otherwise change
|
||||
# how the value reads back; the second row is every character str.splitlines()
|
||||
# treats as a line break, and parse_frontmatter() reads one scalar per line, so
|
||||
# emitting one of those bare would split the value across lines.
|
||||
_FM_MUST_QUOTE = (
|
||||
":", "#", "[", "]", "{", "}", ",", "&", "*", "!", "|", ">", "'", '"', "%", "@",
|
||||
"\n", "\r", "\v", "\f", "\x1c", "\x1d", "\x1e", "\x85", "\u2028", "\u2029",
|
||||
)
|
||||
|
||||
# json.dumps escapes every C0 control character, but with ensure_ascii=False it
|
||||
# passes NEL / LINE SEPARATOR / PARAGRAPH SEPARATOR through literally, and
|
||||
# str.splitlines() still breaks on all three. Re-escape exactly those, which
|
||||
# json.loads decodes again on the way in, so the pair stays symmetric.
|
||||
_FM_POST_DUMPS_ESCAPES = (
|
||||
("\x85", "\\u0085"),
|
||||
("\u2028", "\\u2028"),
|
||||
("\u2029", "\\u2029"),
|
||||
)
|
||||
|
||||
|
||||
def _emit_scalar(v: Any) -> str:
|
||||
if v is None:
|
||||
return "null"
|
||||
@@ -181,8 +213,15 @@ def _emit_scalar(v: Any) -> str:
|
||||
if isinstance(v, list):
|
||||
return "[" + ", ".join(_emit_scalar(x) for x in v) + "]"
|
||||
s = str(v)
|
||||
if any(c in s for c in (":", "#", "\n", "[", "]", "{", "}", ",", "&", "*", "!", "|", ">", "'", '"', "%", "@")):
|
||||
return json.dumps(s)
|
||||
if any(c in s for c in _FM_MUST_QUOTE):
|
||||
# ensure_ascii=False keeps non-ASCII text as itself. SKILL.md is UTF-8 at
|
||||
# both ends (skills.py reads it, atomic_write_text writes it), so the
|
||||
# \uXXXX form bought nothing and leaked into the parsed value (#5210).
|
||||
out = json.dumps(s, ensure_ascii=False)
|
||||
for ch, esc in _FM_POST_DUMPS_ESCAPES:
|
||||
if ch in out:
|
||||
out = out.replace(ch, esc)
|
||||
return out
|
||||
return s
|
||||
|
||||
|
||||
|
||||
@@ -1,16 +1,18 @@
|
||||
"""Import SKILL.md bundles from public GitHub (or skills.sh → GitHub) URLs."""
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
from typing import Dict, Iterable, List, Optional, Tuple, cast
|
||||
from urllib.parse import quote, urljoin, urlparse
|
||||
|
||||
import httpcore
|
||||
import httpx
|
||||
|
||||
from src.url_safety import check_outbound_url
|
||||
from src.url_safety import _default_resolver, check_outbound_url
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -25,6 +27,7 @@ TEXT_NAMES = {"skill.md", "license", "license.md", "readme.md"}
|
||||
_GITHUB_HOSTS = frozenset({
|
||||
"github.com", "www.github.com", "api.github.com", "raw.githubusercontent.com",
|
||||
})
|
||||
_SKILLS_SH_HOSTS = frozenset({"skills.sh", "www.skills.sh"})
|
||||
|
||||
|
||||
def _github_host(url: str) -> str:
|
||||
@@ -72,18 +75,158 @@ def _is_text_file(name: str) -> bool:
|
||||
_MAX_FETCH_REDIRECTS = 5
|
||||
|
||||
|
||||
def _check_fetch_url(url: str) -> None:
|
||||
"""SSRF guard for skill-import fetches (defense-in-depth).
|
||||
def _validated_ips(raw_ips: List[str]) -> List[ipaddress._BaseAddress]:
|
||||
"""Parse and de-duplicate one resolver snapshot in resolver order."""
|
||||
ips: List[ipaddress._BaseAddress] = []
|
||||
seen = set()
|
||||
for raw in raw_ips:
|
||||
if not isinstance(raw, str):
|
||||
continue
|
||||
try:
|
||||
ip = ipaddress.ip_address(raw.split("%", 1)[0])
|
||||
except ValueError:
|
||||
continue
|
||||
if ip in seen:
|
||||
continue
|
||||
seen.add(ip)
|
||||
ips.append(ip)
|
||||
return ips
|
||||
|
||||
Skill bundles only ever come from public GitHub, never an internal
|
||||
address, so block private/loopback/link-local targets on every hop —
|
||||
matching the hardened web-fetch path in
|
||||
``services/search/content.py:_get_public_url`` rather than the lenient
|
||||
default used for admin-configured model endpoints.
|
||||
"""
|
||||
ok, reason = check_outbound_url(url, block_private=True)
|
||||
|
||||
def _resolve_and_check_url(url: str) -> List[ipaddress._BaseAddress]:
|
||||
"""Return the exact address snapshot approved for one fetch hop."""
|
||||
resolved_ips: List[str] = []
|
||||
|
||||
def _recording_resolver(host: str) -> List[str]:
|
||||
answers = list(_default_resolver(host))
|
||||
resolved_ips[:] = answers
|
||||
return answers
|
||||
|
||||
ok, reason = check_outbound_url(
|
||||
url,
|
||||
block_private=True,
|
||||
resolver=_recording_resolver,
|
||||
)
|
||||
if not ok:
|
||||
raise SkillImportError(reason)
|
||||
raise SkillImportError(f"outbound URL blocked: {reason}")
|
||||
|
||||
pinned_ips = _validated_ips(resolved_ips)
|
||||
if not pinned_ips:
|
||||
raise SkillImportError("outbound URL blocked: host did not resolve to a usable address")
|
||||
return pinned_ips
|
||||
|
||||
|
||||
# Backward compatibility alias for tests importing _check_fetch_url directly
|
||||
_check_fetch_url = _resolve_and_check_url
|
||||
|
||||
|
||||
class _PinnedBackend(httpcore.NetworkBackend):
|
||||
"""Connect only to addresses from one validated DNS snapshot."""
|
||||
|
||||
def __init__(self, ips: List[ipaddress._BaseAddress]):
|
||||
self._ips = [str(ip) for ip in ips]
|
||||
self._real = httpcore.SyncBackend()
|
||||
|
||||
def connect_tcp(
|
||||
self,
|
||||
host: str,
|
||||
port: int,
|
||||
timeout: float | None = None,
|
||||
local_address: str | None = None,
|
||||
socket_options=None,
|
||||
):
|
||||
deadline = None if timeout is None else time.monotonic() + timeout
|
||||
last_exc: Optional[Exception] = None
|
||||
for ip in self._ips:
|
||||
remaining = None if deadline is None else max(0.0, deadline - time.monotonic())
|
||||
try:
|
||||
return self._real.connect_tcp(
|
||||
ip,
|
||||
port,
|
||||
remaining,
|
||||
local_address,
|
||||
socket_options,
|
||||
)
|
||||
except (httpcore.ConnectError, httpcore.ConnectTimeout) as exc:
|
||||
last_exc = exc
|
||||
if deadline is not None and time.monotonic() >= deadline:
|
||||
break
|
||||
if last_exc is not None:
|
||||
raise last_exc
|
||||
raise httpcore.ConnectError("no validated address available")
|
||||
|
||||
def connect_unix_socket(self, path, timeout=None, socket_options=None):
|
||||
return self._real.connect_unix_socket(path, timeout, socket_options)
|
||||
|
||||
def sleep(self, seconds: float) -> None:
|
||||
return self._real.sleep(seconds)
|
||||
|
||||
|
||||
_HTTPCORE_TO_HTTPX_EXC = {
|
||||
httpcore.ConnectError: httpx.ConnectError,
|
||||
httpcore.ConnectTimeout: httpx.ConnectTimeout,
|
||||
httpcore.LocalProtocolError: httpx.LocalProtocolError,
|
||||
httpcore.NetworkError: httpx.NetworkError,
|
||||
httpcore.PoolTimeout: httpx.PoolTimeout,
|
||||
httpcore.ProtocolError: httpx.ProtocolError,
|
||||
httpcore.ProxyError: httpx.ProxyError,
|
||||
httpcore.ReadError: httpx.ReadError,
|
||||
httpcore.ReadTimeout: httpx.ReadTimeout,
|
||||
httpcore.RemoteProtocolError: httpx.RemoteProtocolError,
|
||||
httpcore.TimeoutException: httpx.TimeoutException,
|
||||
httpcore.UnsupportedProtocol: httpx.UnsupportedProtocol,
|
||||
httpcore.WriteError: httpx.WriteError,
|
||||
httpcore.WriteTimeout: httpx.WriteTimeout,
|
||||
}
|
||||
|
||||
|
||||
class _PinnedTransport(httpx.BaseTransport):
|
||||
"""Pin socket connects while preserving URL authority, Host, and TLS SNI."""
|
||||
|
||||
def __init__(self, ips: List[ipaddress._BaseAddress]):
|
||||
self._pinned_ips = list(ips)
|
||||
self._pool = httpcore.ConnectionPool(
|
||||
ssl_context=httpx.create_ssl_context(),
|
||||
http1=True,
|
||||
http2=False,
|
||||
network_backend=_PinnedBackend(ips),
|
||||
)
|
||||
|
||||
def handle_request(self, request: httpx.Request) -> httpx.Response:
|
||||
core_request = httpcore.Request(
|
||||
method=request.method,
|
||||
url=httpcore.URL(
|
||||
scheme=request.url.raw_scheme,
|
||||
host=request.url.raw_host,
|
||||
port=request.url.port,
|
||||
target=request.url.raw_path,
|
||||
),
|
||||
headers=request.headers.raw,
|
||||
content=request.stream,
|
||||
extensions=request.extensions,
|
||||
)
|
||||
core_response = None
|
||||
try:
|
||||
core_response = self._pool.handle_request(core_request)
|
||||
content = b"".join(cast(Iterable[bytes], core_response.stream))
|
||||
except Exception as exc:
|
||||
mapped = _HTTPCORE_TO_HTTPX_EXC.get(type(exc))
|
||||
if mapped is not None:
|
||||
raise mapped(str(exc)) from exc
|
||||
raise
|
||||
finally:
|
||||
if core_response is not None:
|
||||
core_response.close()
|
||||
|
||||
return httpx.Response(
|
||||
status_code=core_response.status,
|
||||
headers=core_response.headers,
|
||||
content=content,
|
||||
extensions=core_response.extensions,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
self._pool.close()
|
||||
|
||||
|
||||
def _get_checked(
|
||||
@@ -100,49 +243,76 @@ def _get_checked(
|
||||
hand lets us re-validate every hop, closing that blind-SSRF gap.
|
||||
"""
|
||||
current = url
|
||||
with httpx.Client(follow_redirects=False, timeout=timeout) as client:
|
||||
for _ in range(_MAX_FETCH_REDIRECTS + 1):
|
||||
_check_fetch_url(current)
|
||||
for _ in range(_MAX_FETCH_REDIRECTS + 1):
|
||||
pinned_ips = _resolve_and_check_url(current)
|
||||
with httpx.Client(
|
||||
transport=_PinnedTransport(pinned_ips),
|
||||
follow_redirects=False,
|
||||
timeout=timeout,
|
||||
) as client:
|
||||
r = client.get(current, headers=headers)
|
||||
if r.status_code in (301, 302, 303, 307, 308):
|
||||
location = r.headers.get("location")
|
||||
if not location:
|
||||
return r
|
||||
current = urljoin(str(r.url), location)
|
||||
continue
|
||||
return r
|
||||
|
||||
if r.status_code in (301, 302, 303, 307, 308):
|
||||
location = r.headers.get("location")
|
||||
if not location:
|
||||
return r
|
||||
current = urljoin(str(r.url), location)
|
||||
continue
|
||||
return r
|
||||
raise SkillImportError("too many redirects while fetching skill bundle")
|
||||
|
||||
|
||||
def parse_skill_source(url: str) -> ResolvedSource:
|
||||
"""Normalize skills.sh / GitHub web URLs into owner/repo/ref/path."""
|
||||
raw = (url or "").strip()
|
||||
if not raw:
|
||||
url = (url or "").strip()
|
||||
if not url:
|
||||
raise SkillImportError("URL is required")
|
||||
|
||||
# skills.sh often links to GitHub; try to unwrap ?url= or redirect target later.
|
||||
if "skills.sh" in raw and "github.com" not in raw:
|
||||
r = _get_checked(raw, timeout=20.0)
|
||||
# ``urlparse`` only reports an unambiguous scheme when the URL carries the
|
||||
# ``scheme://`` form. Opaque schemes (``mailto:``, ``javascript:``) and a
|
||||
# schemeless ``host:port`` both parse a "scheme" that is not one, so they
|
||||
# fall through to the host check below and are rejected on the host instead.
|
||||
scheme = urlparse(url).scheme.lower()
|
||||
if scheme not in ("http", "https"):
|
||||
if scheme and url.lower().startswith(f"{scheme}://"):
|
||||
raise SkillImportError(f"unsupported URL scheme: {scheme}")
|
||||
# Schemeless "github.com/owner/repo" — accept only a supported host.
|
||||
rough_host = (urlparse("//" + url).hostname or "").lower()
|
||||
if rough_host not in _GITHUB_HOSTS and rough_host not in _SKILLS_SH_HOSTS:
|
||||
raise SkillImportError("Only GitHub or skills.sh URLs are supported")
|
||||
url = "https://" + url
|
||||
|
||||
parsed = urlparse(url)
|
||||
hostname = (parsed.hostname or "").lower()
|
||||
if hostname not in _GITHUB_HOSTS and hostname not in _SKILLS_SH_HOSTS:
|
||||
raise SkillImportError("Only GitHub or skills.sh URLs are supported")
|
||||
|
||||
# A skills.sh link is only usable if it redirects to an exact supported
|
||||
# GitHub host. Scraping the page body for a github.com link cannot work:
|
||||
# skill pages only ever link the repository root, never the skill's
|
||||
# subdirectory, so the scrape resolves every skill in a repo to the same
|
||||
# (wrong) bundle. Fail with an actionable message instead.
|
||||
if hostname in _SKILLS_SH_HOSTS:
|
||||
r = _get_checked(url, timeout=20.0)
|
||||
if r.status_code >= 400:
|
||||
raise _github_response_error(r)
|
||||
final = str(r.url)
|
||||
_assert_github_url(final, context="redirect target")
|
||||
# Page may embed a github link; prefer final URL if redirected.
|
||||
if "github.com" in final:
|
||||
raw = final
|
||||
else:
|
||||
m = re.search(r"https?://github\.com/[^\s\"')]+", r.text or "")
|
||||
if m:
|
||||
raw = m.group(0).rstrip(".,)")
|
||||
if _github_host(final) not in _GITHUB_HOSTS:
|
||||
raise SkillImportError(
|
||||
"skills.sh did not redirect to GitHub — open the skill's "
|
||||
"repository on GitHub, navigate to the exact skill folder or "
|
||||
"SKILL.md file, and paste that URL; the repository-root link "
|
||||
"alone is not sufficient"
|
||||
)
|
||||
url = final
|
||||
|
||||
parsed = urlparse(raw)
|
||||
host = _github_host(raw)
|
||||
if host not in _GITHUB_HOSTS:
|
||||
raise SkillImportError(
|
||||
"Only GitHub URLs are supported (https://github.com/... or raw.githubusercontent.com/...)"
|
||||
)
|
||||
# Update parsed and hostname to reflect the new GitHub URL
|
||||
parsed = urlparse(url)
|
||||
hostname = (parsed.hostname or "").lower()
|
||||
|
||||
if host == "raw.githubusercontent.com":
|
||||
_assert_github_url(url)
|
||||
|
||||
if hostname == "raw.githubusercontent.com":
|
||||
# /owner/repo/ref/path/to/file
|
||||
bits = [p for p in parsed.path.split("/") if p]
|
||||
if len(bits) < 4:
|
||||
|
||||
+1090
-291
File diff suppressed because it is too large
Load Diff
+84
-26
@@ -17,13 +17,14 @@ close / navigation / refresh). It does NOT survive a server restart.
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from typing import AsyncGenerator, Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _Run:
|
||||
__slots__ = ("buffer", "subscribers", "status", "task", "evict_task")
|
||||
__slots__ = ("buffer", "subscribers", "status", "task", "evict_task", "run_id")
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.buffer: list = [] # ordered SSE event strings (replay log)
|
||||
@@ -31,6 +32,9 @@ class _Run:
|
||||
self.status: str = "running" # running | done | error | stopped
|
||||
self.task: Optional[asyncio.Task] = None
|
||||
self.evict_task: Optional[asyncio.Task] = None
|
||||
# Stable across every subscription/replay of this exact detached run.
|
||||
# The browser uses it to make local cost accounting replay-idempotent.
|
||||
self.run_id: str = uuid.uuid4().hex
|
||||
|
||||
|
||||
_RUNS: Dict[str, _Run] = {}
|
||||
@@ -53,13 +57,24 @@ def _publish(run: _Run, ev: str) -> None:
|
||||
pass
|
||||
|
||||
|
||||
def _schedule_evict(session_id: str) -> None:
|
||||
def _wake_run_subscribers(run: _Run) -> None:
|
||||
"""Close subscribers even when the drain task never reached its body."""
|
||||
for q in list(run.subscribers):
|
||||
try:
|
||||
q.put_nowait((None, None))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _schedule_evict(session_id: str, expected_run: Optional[_Run] = None) -> None:
|
||||
"""(Re)arm a grace-period eviction for a terminal run with no subscribers.
|
||||
Identity-checked so a run that gets replaced/reused is never evicted by a
|
||||
stale timer."""
|
||||
run = _RUNS.get(session_id)
|
||||
if run is None:
|
||||
return
|
||||
if expected_run is not None and run is not expected_run:
|
||||
return
|
||||
if run.evict_task and not run.evict_task.done():
|
||||
run.evict_task.cancel()
|
||||
|
||||
@@ -85,25 +100,38 @@ def get_status(session_id: str) -> Optional[str]:
|
||||
return r.status if r else None
|
||||
|
||||
|
||||
async def _drain(session_id: str, agen: AsyncGenerator[str, None],
|
||||
def get_run_id(session_id: str) -> Optional[str]:
|
||||
"""Return the opaque identity of the current detached run, if present."""
|
||||
r = _RUNS.get(session_id)
|
||||
return r.run_id if r else None
|
||||
|
||||
|
||||
def get_active_run(session_id: str) -> Optional[_Run]:
|
||||
"""Return the exact active run currently registered for a session."""
|
||||
r = _RUNS.get(session_id)
|
||||
return r if r and r.status == "running" else None
|
||||
|
||||
|
||||
async def _drain(session_id: str, run: _Run, agen: AsyncGenerator[str, None],
|
||||
prev_task: Optional[asyncio.Task] = None) -> None:
|
||||
"""Pull every event from the wrapped generator into the run buffer, fanning
|
||||
each out to live subscribers. Runs to completion regardless of subscribers."""
|
||||
run = _RUNS.get(session_id)
|
||||
if run is None:
|
||||
return
|
||||
subscribers_woken = False
|
||||
|
||||
def _wake_subscribers() -> None:
|
||||
nonlocal subscribers_woken
|
||||
if subscribers_woken:
|
||||
return
|
||||
subscribers_woken = True
|
||||
_wake_run_subscribers(run)
|
||||
|
||||
# If this run replaced an in-flight one (rapid double-send), wait for that
|
||||
# one to fully finish first. Its CancelledError handler calls aclose(), which
|
||||
# persists its partial response — letting it complete before we start writing
|
||||
# keeps the two runs' session saves sequential instead of interleaved.
|
||||
if prev_task is not None and not prev_task.done():
|
||||
try:
|
||||
await asyncio.wait({prev_task})
|
||||
except asyncio.CancelledError:
|
||||
raise # our own cancellation — propagate
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
if prev_task is not None and not prev_task.done():
|
||||
await asyncio.wait({prev_task})
|
||||
async for ev in agen:
|
||||
_publish(run, ev)
|
||||
if run.status == "running":
|
||||
@@ -116,6 +144,16 @@ async def _drain(session_id: str, agen: AsyncGenerator[str, None],
|
||||
await agen.aclose()
|
||||
except Exception:
|
||||
pass
|
||||
# A rapid third replacement can cancel this task while it is still
|
||||
# waiting for its predecessor. Close this run's subscribers promptly,
|
||||
# but keep the task alive until the predecessor finishes so the next
|
||||
# run still observes the transitive session-save ordering barrier.
|
||||
_wake_subscribers()
|
||||
if prev_task is not None and not prev_task.done():
|
||||
try:
|
||||
await asyncio.shield(prev_task)
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.error("[agent-run] %s failed: %s", session_id, e, exc_info=True)
|
||||
run.status = "error"
|
||||
@@ -127,15 +165,11 @@ async def _drain(session_id: str, agen: AsyncGenerator[str, None],
|
||||
_publish(run, "data: [DONE]\n\n")
|
||||
finally:
|
||||
# Wake every subscriber with the end sentinel so their SSE closes.
|
||||
for q in list(run.subscribers):
|
||||
try:
|
||||
q.put_nowait((None, None))
|
||||
except Exception:
|
||||
pass
|
||||
_wake_subscribers()
|
||||
# Run is terminal — arm the grace timer so it (and its buffer) is
|
||||
# eventually freed even if nobody ever reconnects. subscribe() cancels
|
||||
# this on connect and re-arms on disconnect.
|
||||
_schedule_evict(session_id)
|
||||
_schedule_evict(session_id, run)
|
||||
|
||||
|
||||
def start(session_id: str, agen: AsyncGenerator[str, None]) -> _Run:
|
||||
@@ -145,20 +179,37 @@ def start(session_id: str, agen: AsyncGenerator[str, None]) -> _Run:
|
||||
prev_task: Optional[asyncio.Task] = None
|
||||
if prev:
|
||||
if prev.task and not prev.task.done():
|
||||
# A task cancelled before its first instruction never enters
|
||||
# _drain(), so its except/finally blocks cannot update status or
|
||||
# wake a response already bound to this exact run. Terminalize it
|
||||
# synchronously before cancelling; _drain's cleanup is idempotent
|
||||
# when the task had already started.
|
||||
if prev.status == "running":
|
||||
prev.status = "stopped"
|
||||
_wake_run_subscribers(prev)
|
||||
prev.task.cancel()
|
||||
prev_task = prev.task # new run awaits this before it starts writing
|
||||
if prev.evict_task and not prev.evict_task.done():
|
||||
prev.evict_task.cancel()
|
||||
run = _Run()
|
||||
_RUNS[session_id] = run
|
||||
run.task = asyncio.create_task(_drain(session_id, agen, prev_task))
|
||||
run.task = asyncio.create_task(_drain(session_id, run, agen, prev_task))
|
||||
return run
|
||||
|
||||
|
||||
async def subscribe(session_id: str) -> AsyncGenerator[str, None]:
|
||||
async def subscribe(
|
||||
session_id: str,
|
||||
expected_run: Optional[_Run] = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""Replay the run's buffer from the start, then stream live until it ends.
|
||||
Safe to call repeatedly (reconnect) and from multiple clients at once."""
|
||||
run = _RUNS.get(session_id)
|
||||
Safe to call repeatedly (reconnect) and from multiple clients at once.
|
||||
|
||||
``expected_run`` binds a lazy StreamingResponse body to the same run whose
|
||||
identity was put in its response headers. Without that binding, a rapid
|
||||
replacement between response construction and body iteration could replay
|
||||
the replacement run under the prior run's identity.
|
||||
"""
|
||||
run = expected_run or _RUNS.get(session_id)
|
||||
if run is None:
|
||||
return
|
||||
q: asyncio.Queue = asyncio.Queue()
|
||||
@@ -201,12 +252,19 @@ async def subscribe(session_id: str) -> AsyncGenerator[str, None]:
|
||||
# Last subscriber gone on a finished run — (re)arm eviction so the
|
||||
# buffer doesn't linger indefinitely.
|
||||
if not run.subscribers and run.status != "running":
|
||||
_schedule_evict(session_id)
|
||||
_schedule_evict(session_id, run)
|
||||
|
||||
|
||||
def stop(session_id: str) -> bool:
|
||||
"""Cancel an in-flight run (the wrapped generator saves its partial)."""
|
||||
def stop(session_id: str, expected_run_id: Optional[str] = None) -> bool:
|
||||
"""Cancel the matching in-flight run (which saves its partial output).
|
||||
|
||||
A stale browser may issue Stop after another tab has replaced the session's
|
||||
run. Once the caller knows its opaque run identity, fail closed rather than
|
||||
cancelling that newer run.
|
||||
"""
|
||||
run = _RUNS.get(session_id)
|
||||
if not expected_run_id or run is None or run.run_id != expected_run_id:
|
||||
return False
|
||||
if run and run.task and not run.task.done():
|
||||
run.task.cancel()
|
||||
return True
|
||||
|
||||
@@ -510,7 +510,12 @@ async def do_manage_settings(content: str, owner: Optional[str] = None) -> Dict:
|
||||
# set/get/list/delete operate on the REAL app settings (the same store
|
||||
# the Settings panel writes), so changing a model / voice / search
|
||||
# engine / reminder channel from chat actually takes effect.
|
||||
from src.settings import load_settings, save_settings, DEFAULT_SETTINGS
|
||||
from src.settings import (
|
||||
DEFAULT_SETTINGS,
|
||||
RETIRED_SETTING_KEYS,
|
||||
load_settings,
|
||||
save_settings,
|
||||
)
|
||||
|
||||
# Secrets/credentials the agent must NOT write: kept read-only (masked)
|
||||
# so API keys never flow through chat. User sets these in the panel.
|
||||
@@ -562,6 +567,9 @@ async def do_manage_settings(content: str, owner: Optional[str] = None) -> Dict:
|
||||
return k2
|
||||
return _ALIASES_SET.get(k2, (k or "").strip())
|
||||
|
||||
def _is_managed_key(key):
|
||||
return key in DEFAULT_SETTINGS and key not in RETIRED_SETTING_KEYS
|
||||
|
||||
_ENUMS = {
|
||||
"image_quality": ["low", "medium", "high"],
|
||||
"reminder_channel": ["browser", "email", "ntfy", "webhook"],
|
||||
@@ -624,14 +632,18 @@ async def do_manage_settings(content: str, owner: Optional[str] = None) -> Dict:
|
||||
|
||||
if action == "list":
|
||||
s = load_settings()
|
||||
shown = {k: _mask(k, v) for k, v in s.items() if k in DEFAULT_SETTINGS and not isinstance(v, dict)}
|
||||
shown = {
|
||||
k: _mask(k, v)
|
||||
for k, v in s.items()
|
||||
if _is_managed_key(k) and not isinstance(v, dict)
|
||||
}
|
||||
return {"response": f"{len(shown)} settings (use get/set with a key)", "settings": shown, "exit_code": 0}
|
||||
|
||||
elif action == "get":
|
||||
key = _resolve(args.get("key", ""))
|
||||
if not key:
|
||||
return {"error": "key is required", "exit_code": 1}
|
||||
if key not in DEFAULT_SETTINGS:
|
||||
if not _is_managed_key(key):
|
||||
return {"error": f"Unknown setting '{args.get('key')}'. Use action='list' to see them.", "exit_code": 1}
|
||||
val = load_settings().get(key, DEFAULT_SETTINGS.get(key))
|
||||
return {"response": f"{key} = {_mask(key, val)}", "value": _mask(key, val), "exit_code": 0}
|
||||
@@ -642,11 +654,11 @@ async def do_manage_settings(content: str, owner: Optional[str] = None) -> Dict:
|
||||
if not raw:
|
||||
return {"error": "key is required", "exit_code": 1}
|
||||
key = _resolve(raw)
|
||||
if key not in DEFAULT_SETTINGS:
|
||||
if not _is_managed_key(key):
|
||||
return {"error": f"Unknown setting '{raw}'. Use action='list' to see available settings.", "exit_code": 1}
|
||||
if _is_secret(key):
|
||||
return {"response": f"'{key}' is a credential/secret. For security I can't set it from chat. Open Settings and set it there.", "exit_code": 0}
|
||||
# Structured settings (dicts/lists like keybinds, default_model_fallbacks)
|
||||
# Structured settings (dicts/lists like keybinds or vision fallbacks)
|
||||
# have no safe scalar coercion; _coerce would pass a bare string
|
||||
# straight through and clobber the structure. Refuse them here; they're
|
||||
# edited in their dedicated panels. (reset/delete still restore the
|
||||
@@ -675,7 +687,7 @@ async def do_manage_settings(content: str, owner: Optional[str] = None) -> Dict:
|
||||
|
||||
elif action == "delete" or action == "reset":
|
||||
key = _resolve(args.get("key", ""))
|
||||
if key not in DEFAULT_SETTINGS:
|
||||
if not _is_managed_key(key):
|
||||
return {"error": f"Unknown setting '{args.get('key')}'.", "exit_code": 1}
|
||||
if _is_secret(key):
|
||||
return {"response": f"'{key}' is a credential. Reset it in the panel.", "exit_code": 0}
|
||||
|
||||
@@ -6,6 +6,7 @@ import sys
|
||||
import time
|
||||
import collections
|
||||
from typing import Optional, Callable, Awaitable, Tuple, Dict
|
||||
from core.platform_compat import IS_WINDOWS, find_bash
|
||||
from src.constants import MAX_OUTPUT_CHARS
|
||||
|
||||
DEFAULT_BASH_TIMEOUT = 60 * 60 # 1 hour
|
||||
@@ -16,6 +17,27 @@ PROGRESS_TAIL_LINES = 12
|
||||
TMUX_CAPTURE_LINES = 2000
|
||||
|
||||
|
||||
async def _create_bash_subprocess(command: str, **kwargs):
|
||||
"""Start the agent shell with Bash semantics on every supported OS.
|
||||
|
||||
``asyncio.create_subprocess_shell`` delegates to ``cmd.exe`` on native
|
||||
Windows. That contradicts the Bash tool contract and makes POSIX commands
|
||||
such as ``pwd``, ``ls -la``, and ``cat`` unreliable even when the launcher
|
||||
has found Git Bash. Pass the selected workspace as a structural ``cwd``
|
||||
argument; Git Bash inherits that native Windows directory and exposes it
|
||||
using its normal ``/c/...`` representation.
|
||||
"""
|
||||
if IS_WINDOWS:
|
||||
bash = find_bash()
|
||||
if not bash:
|
||||
raise RuntimeError(
|
||||
"Git Bash is required for the Bash tool on Windows; "
|
||||
"install Git for Windows and restart Odysseus"
|
||||
)
|
||||
return await asyncio.create_subprocess_exec(bash, "-c", command, **kwargs)
|
||||
return await asyncio.create_subprocess_shell(command, **kwargs)
|
||||
|
||||
|
||||
def _tmux_session_name(session_id: Optional[str]) -> str:
|
||||
raw = re.sub(r"[^A-Za-z0-9_.-]+", "-", str(session_id or "default")).strip("-")
|
||||
return f"ody-agent-{raw[:80] or 'default'}"
|
||||
@@ -280,7 +302,10 @@ class BashTool:
|
||||
progress_cb = ctx.get("progress_cb")
|
||||
_subproc_env = ctx.get("subproc_env")
|
||||
session_id = ctx.get("session_id")
|
||||
if session_id and shutil.which("tmux"):
|
||||
# tmux is a POSIX persistence path. A stray MSYS/Cygwin tmux.exe on
|
||||
# native Windows must not bypass the Git Bash launcher below: the tmux
|
||||
# setup hard-codes /bin/bash and cannot safely consume a native cwd.
|
||||
if session_id and not IS_WINDOWS and shutil.which("tmux"):
|
||||
stdout, stderr, rc, timed_out = await _run_tmux_bash(
|
||||
content,
|
||||
session_id=str(session_id),
|
||||
@@ -307,13 +332,16 @@ class BashTool:
|
||||
"tmux_session": _tmux_session_name(str(session_id)),
|
||||
}
|
||||
|
||||
proc = await asyncio.create_subprocess_shell(
|
||||
content,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
env=_subproc_env,
|
||||
cwd=agent_cwd(),
|
||||
)
|
||||
try:
|
||||
proc = await _create_bash_subprocess(
|
||||
content,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
env=_subproc_env,
|
||||
cwd=agent_cwd(),
|
||||
)
|
||||
except RuntimeError as e:
|
||||
return {"error": f"bash: {e}", "exit_code": 1}
|
||||
stdout, stderr, rc, timed_out = await _run_subprocess_streaming(
|
||||
proc,
|
||||
timeout=DEFAULT_BASH_TIMEOUT,
|
||||
|
||||
+682
-105
@@ -20,6 +20,395 @@ from src.interactive_gate import wait_for_interactive_quiet
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _read_email_urgency_state(state_path):
|
||||
"""Read one atomic urgency checkpoint, tolerating the legacy shape."""
|
||||
from pathlib import Path
|
||||
|
||||
state_path = Path(state_path)
|
||||
try:
|
||||
state = (
|
||||
json.loads(state_path.read_text(encoding="utf-8"))
|
||||
if state_path.exists()
|
||||
else {}
|
||||
)
|
||||
except Exception:
|
||||
return {}
|
||||
return state if isinstance(state, dict) else {}
|
||||
|
||||
|
||||
def _email_urgency_account_generations(state):
|
||||
"""Return normalized per-account checkpoint/complete generations.
|
||||
|
||||
Checkpoint generations fence every accepted state mutation. Complete
|
||||
generations advance only for a non-stale complete scan. Missing metadata
|
||||
is the legacy generation zero.
|
||||
"""
|
||||
raw = state.get("account_generations", {}) if isinstance(state, dict) else {}
|
||||
if not isinstance(raw, dict):
|
||||
return {}
|
||||
|
||||
generations = {}
|
||||
for account_id, value in raw.items():
|
||||
if isinstance(value, dict):
|
||||
checkpoint = value.get("checkpoint", 0)
|
||||
complete = value.get("complete", 0)
|
||||
else:
|
||||
# Tolerate an intermediate scalar representation as one completed
|
||||
# checkpoint generation instead of discarding its fence.
|
||||
checkpoint = value
|
||||
complete = value
|
||||
try:
|
||||
checkpoint = max(0, int(checkpoint))
|
||||
except (TypeError, ValueError):
|
||||
checkpoint = 0
|
||||
try:
|
||||
complete = max(0, int(complete))
|
||||
except (TypeError, ValueError):
|
||||
complete = 0
|
||||
generations[str(account_id)] = {
|
||||
"checkpoint": checkpoint,
|
||||
"complete": complete,
|
||||
}
|
||||
return generations
|
||||
|
||||
|
||||
def _email_urgency_string_set(value):
|
||||
if not isinstance(value, (list, tuple, set, frozenset)):
|
||||
return set()
|
||||
return {str(item) for item in value if isinstance(item, (str, int))}
|
||||
|
||||
|
||||
def _acquire_email_urgency_state_lock(
|
||||
state_path,
|
||||
lock_db_path,
|
||||
cancel_event,
|
||||
timeout_seconds=120,
|
||||
):
|
||||
"""Acquire the cross-process urgency lock without blocking the app loop."""
|
||||
import sqlite3
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
state_path = Path(state_path)
|
||||
state_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
deadline = time.monotonic() + timeout_seconds
|
||||
|
||||
while not cancel_event.is_set():
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
raise sqlite3.OperationalError("timed out waiting for urgency state lock")
|
||||
conn = sqlite3.connect(
|
||||
str(lock_db_path),
|
||||
timeout=min(0.25, max(0.01, remaining)),
|
||||
check_same_thread=False,
|
||||
)
|
||||
try:
|
||||
conn.execute("BEGIN IMMEDIATE")
|
||||
except sqlite3.OperationalError as exc:
|
||||
conn.close()
|
||||
if "locked" not in str(exc).lower():
|
||||
raise
|
||||
cancel_event.wait(min(0.05, max(0.0, remaining)))
|
||||
continue
|
||||
except BaseException:
|
||||
conn.close()
|
||||
raise
|
||||
|
||||
if cancel_event.is_set():
|
||||
conn.rollback()
|
||||
conn.close()
|
||||
return None, None
|
||||
return conn, _read_email_urgency_state(state_path)
|
||||
|
||||
return None, None
|
||||
|
||||
|
||||
def _close_email_urgency_state_lock(conn):
|
||||
if conn is None:
|
||||
return
|
||||
try:
|
||||
try:
|
||||
conn.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def _commit_email_urgency_state(conn, state_path, next_state):
|
||||
"""Atomically publish JSON before releasing the SQLite write lock."""
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
state_path = Path(state_path)
|
||||
temp_path = state_path.with_name(
|
||||
f".{state_path.name}.{uuid.uuid4().hex}.tmp"
|
||||
)
|
||||
try:
|
||||
temp_path.write_text(json.dumps(next_state), encoding="utf-8")
|
||||
temp_path.replace(state_path)
|
||||
conn.commit()
|
||||
except BaseException:
|
||||
conn.rollback()
|
||||
raise
|
||||
finally:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
conn.close()
|
||||
|
||||
|
||||
async def _run_email_urgency_state_transaction(
|
||||
state_path,
|
||||
lock_db_path,
|
||||
operation,
|
||||
):
|
||||
"""Serialize one urgency decision while keeping async work on this loop.
|
||||
|
||||
Only lock acquisition waits in a worker thread. ``operation`` is awaited
|
||||
on the caller's long-lived event loop, where shared async clients, locks,
|
||||
and the browser-notification queue belong. Cancellation rolls back the
|
||||
SQLite transaction and never publishes a checkpoint.
|
||||
"""
|
||||
import asyncio
|
||||
import threading
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
cancel_event = threading.Event()
|
||||
acquire_future = loop.run_in_executor(
|
||||
None,
|
||||
_acquire_email_urgency_state_lock,
|
||||
state_path,
|
||||
lock_db_path,
|
||||
cancel_event,
|
||||
)
|
||||
try:
|
||||
conn, prior = await asyncio.shield(acquire_future)
|
||||
except asyncio.CancelledError as cancelled:
|
||||
cancel_event.set()
|
||||
# The acquisition worker owns any connection until it returns. Wait
|
||||
# for its short busy-poll to observe cancellation, then close a lock it
|
||||
# may have won concurrently with the cancellation request.
|
||||
while True:
|
||||
try:
|
||||
conn, _prior = await asyncio.shield(acquire_future)
|
||||
break
|
||||
except asyncio.CancelledError:
|
||||
continue
|
||||
except Exception:
|
||||
conn = None
|
||||
break
|
||||
_close_email_urgency_state_lock(conn)
|
||||
raise cancelled
|
||||
|
||||
if conn is None:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
try:
|
||||
result, next_state = await operation(prior)
|
||||
# Keep this small atomic publish synchronous. There is no await between
|
||||
# the successful operation and commit, so cancellation cannot be
|
||||
# observed and then followed by a checkpoint.
|
||||
try:
|
||||
_commit_email_urgency_state(conn, state_path, next_state)
|
||||
finally:
|
||||
conn = None
|
||||
return result
|
||||
except BaseException:
|
||||
_close_email_urgency_state_lock(conn)
|
||||
raise
|
||||
|
||||
|
||||
def _email_urgency_account_key(message_key):
|
||||
return str(message_key).split(":", 1)[0]
|
||||
|
||||
|
||||
def _email_urgency_payload_account_ids(state):
|
||||
"""Return account IDs that still own user-visible urgency payload."""
|
||||
if not isinstance(state, dict):
|
||||
return set()
|
||||
|
||||
per_uid = state.get("per_uid", {})
|
||||
per_uid_keys = per_uid if isinstance(per_uid, dict) else {}
|
||||
return {
|
||||
_email_urgency_account_key(key) for key in per_uid_keys
|
||||
} | {
|
||||
_email_urgency_account_key(key)
|
||||
for key in _email_urgency_string_set(state.get("notified_uids", []))
|
||||
}
|
||||
|
||||
|
||||
def _email_urgency_known_account_ids(state):
|
||||
"""Return payload owners plus generation-only active/retired markers."""
|
||||
return _email_urgency_payload_account_ids(state) | set(
|
||||
_email_urgency_account_generations(state)
|
||||
)
|
||||
|
||||
|
||||
def _email_urgency_stale_accounts(
|
||||
prior,
|
||||
base_account_generations,
|
||||
account_ids,
|
||||
):
|
||||
prior_generations = _email_urgency_account_generations(prior)
|
||||
base_generations = _email_urgency_account_generations(
|
||||
{"account_generations": base_account_generations}
|
||||
)
|
||||
return {
|
||||
str(account_id)
|
||||
for account_id in account_ids
|
||||
if prior_generations.get(str(account_id), {}).get("checkpoint", 0)
|
||||
!= base_generations.get(str(account_id), {}).get("checkpoint", 0)
|
||||
}
|
||||
|
||||
|
||||
def _merge_email_urgency_state(
|
||||
prior,
|
||||
*,
|
||||
owner,
|
||||
per_uid_scores,
|
||||
notified_uids,
|
||||
all_unread_keys,
|
||||
fully_scanned_account_ids,
|
||||
base_account_generations,
|
||||
timestamp,
|
||||
retired_account_ids=(),
|
||||
base_payload_account_ids=(),
|
||||
known_account_ids=(),
|
||||
):
|
||||
"""Merge a scan without letting an older snapshot erase newer facts."""
|
||||
prior_per_uid = prior.get("per_uid", {})
|
||||
if not isinstance(prior_per_uid, dict):
|
||||
prior_per_uid = {}
|
||||
complete = {str(account_id) for account_id in fully_scanned_account_ids}
|
||||
prior_generations = _email_urgency_account_generations(prior)
|
||||
retire_requested = {str(account_id) for account_id in retired_account_ids}
|
||||
observed_accounts = {
|
||||
_email_urgency_account_key(key) for key in per_uid_scores
|
||||
} | complete | retire_requested
|
||||
stale_accounts = _email_urgency_stale_accounts(
|
||||
prior,
|
||||
base_account_generations,
|
||||
observed_accounts,
|
||||
)
|
||||
prior_payload_accounts = _email_urgency_payload_account_ids(prior)
|
||||
base_payload_accounts = {
|
||||
str(account_id) for account_id in base_payload_account_ids
|
||||
}
|
||||
# A selected account can be absent from the base snapshot. If another
|
||||
# worker creates its first payload before this transaction wins the lock,
|
||||
# membership itself is a fence even when both snapshots normalize to the
|
||||
# legacy generation zero.
|
||||
retired_accounts = {
|
||||
account_id
|
||||
for account_id in retire_requested - stale_accounts
|
||||
if not (
|
||||
account_id in prior_payload_accounts
|
||||
and account_id not in base_payload_accounts
|
||||
)
|
||||
}
|
||||
fresh_complete = complete - stale_accounts - retired_accounts
|
||||
changed_accounts = set(fresh_complete)
|
||||
|
||||
merged_per_uid = {
|
||||
key: value
|
||||
for key, value in prior_per_uid.items()
|
||||
if _email_urgency_account_key(key) not in retired_accounts
|
||||
}
|
||||
for key in list(merged_per_uid):
|
||||
account_id = _email_urgency_account_key(key)
|
||||
if account_id in fresh_complete:
|
||||
merged_per_uid.pop(key, None)
|
||||
changed_accounts.add(account_id)
|
||||
# Partial scans may add or refresh facts, but absence from a partial scan
|
||||
# is not evidence that another checkpoint or UI row is stale. When another
|
||||
# worker committed after this scan captured its base generation, discard
|
||||
# this account's whole stale snapshot. A key absent from the newer state
|
||||
# may have been removed/read, so even a stale-only key is not safely
|
||||
# additive without another fresh scan.
|
||||
for key, value in per_uid_scores.items():
|
||||
account_id = _email_urgency_account_key(key)
|
||||
if account_id in stale_accounts or account_id in retired_accounts:
|
||||
continue
|
||||
if merged_per_uid.get(key) != value:
|
||||
changed_accounts.add(account_id)
|
||||
merged_per_uid[key] = value
|
||||
|
||||
prior_notified = _email_urgency_string_set(prior.get("notified_uids", []))
|
||||
merged_notified = {
|
||||
key
|
||||
for key in prior_notified
|
||||
if _email_urgency_account_key(key) not in retired_accounts
|
||||
}
|
||||
for key in _email_urgency_string_set(notified_uids) - prior_notified:
|
||||
account_id = _email_urgency_account_key(key)
|
||||
if account_id in stale_accounts or account_id in retired_accounts:
|
||||
continue
|
||||
merged_notified.add(key)
|
||||
changed_accounts.add(account_id)
|
||||
for key in list(merged_notified):
|
||||
if (
|
||||
_email_urgency_account_key(key) in fresh_complete
|
||||
and key not in all_unread_keys
|
||||
):
|
||||
merged_notified.discard(key)
|
||||
changed_accounts.add(_email_urgency_account_key(key))
|
||||
|
||||
next_generations = {
|
||||
account_id: dict(value)
|
||||
for account_id, value in prior_generations.items()
|
||||
}
|
||||
for account_id in changed_accounts:
|
||||
generation = next_generations.setdefault(
|
||||
account_id,
|
||||
{"checkpoint": 0, "complete": 0},
|
||||
)
|
||||
generation["checkpoint"] += 1
|
||||
if account_id in fresh_complete:
|
||||
generation["complete"] += 1
|
||||
for account_id in {str(value) for value in known_account_ids}:
|
||||
next_generations.setdefault(
|
||||
account_id,
|
||||
{"checkpoint": 0, "complete": 0},
|
||||
)
|
||||
for account_id in retired_accounts:
|
||||
# Every authoritative absence advances its generation, even when the
|
||||
# prior state is already a payload-empty tombstone. A re-enabled scan
|
||||
# may have captured that previous tombstone immediately before the
|
||||
# account was disabled/deleted again; monotonic advancement is what
|
||||
# makes that in-flight scan stale.
|
||||
generation = next_generations.setdefault(
|
||||
account_id,
|
||||
{"checkpoint": 0, "complete": 0},
|
||||
)
|
||||
generation["checkpoint"] += 1
|
||||
|
||||
total_unread = 0
|
||||
total_urgent = 0
|
||||
max_score = 0
|
||||
for value in merged_per_uid.values():
|
||||
if not isinstance(value, dict):
|
||||
continue
|
||||
try:
|
||||
score = max(0, min(3, int(value.get("score", 0))))
|
||||
except (TypeError, ValueError):
|
||||
score = 0
|
||||
max_score = max(max_score, score)
|
||||
if value.get("unread"):
|
||||
total_unread += 1
|
||||
if score >= 2:
|
||||
total_urgent += 1
|
||||
|
||||
return {
|
||||
"ts": timestamp,
|
||||
"owner": owner or "",
|
||||
"total_unread": total_unread,
|
||||
"total_urgent": total_urgent,
|
||||
"max_score": max_score,
|
||||
"per_uid": merged_per_uid,
|
||||
"notified_uids": sorted(merged_notified),
|
||||
"account_generations": next_generations,
|
||||
}
|
||||
|
||||
|
||||
class TaskNoop(BaseException):
|
||||
"""Raised by an action when it determined there's nothing to do.
|
||||
|
||||
@@ -1878,6 +2267,7 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
# filename for single-user installs (matches prior behaviour).
|
||||
_owner_slug = "".join(c if (c.isalnum() or c in "-_.@") else "_" for c in (owner or "default"))
|
||||
STATE_PATH = _P(DATA_DIR) / f"email_urgency_state_{_owner_slug}.json"
|
||||
STATE_LOCK_DB = STATE_PATH.with_suffix(".lock.sqlite3")
|
||||
CACHE_DIR = _P(EMAIL_URGENCY_CACHE_DIR)
|
||||
CACHE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
STATE_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
@@ -1892,35 +2282,144 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
"shopping", "social", "work", "personal", "legal", "support", "promo",
|
||||
}
|
||||
|
||||
# ── 1. Resolve LLM candidates (utility primary + utility fallbacks; fall
|
||||
# through to default chat as a last resort).
|
||||
# Resolve with the task owner as before, but defer the availability
|
||||
# gate until after authoritative account cleanup. State retirement must
|
||||
# still run when no model is configured.
|
||||
from src.task_endpoint import resolve_task_candidates
|
||||
candidates = resolve_task_candidates(owner=owner)
|
||||
if not candidates:
|
||||
return "No LLM endpoint available", False
|
||||
|
||||
target_account_id = _email_task_account_id(kwargs)
|
||||
|
||||
# ── 2. Enumerate enabled accounts. Match this task's owner AND fall
|
||||
# ── 1. Enumerate enabled accounts. Match this task's owner AND fall
|
||||
# back to the legacy "unowned account whose imap_user / from_address
|
||||
# == this owner" pattern — same rule `_get_email_config` uses, so a
|
||||
# pre-multi-user account row still gets picked up for the seeded task.
|
||||
db = _SL()
|
||||
try:
|
||||
from sqlalchemy import and_ as _and, or_ as _or
|
||||
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)))
|
||||
if target_account_id:
|
||||
q = q.filter(_EA.id == target_account_id)
|
||||
accounts = q.all()
|
||||
finally:
|
||||
db.close()
|
||||
def _enumerate_enabled_accounts():
|
||||
db = _SL()
|
||||
try:
|
||||
from sqlalchemy import and_ as _and, or_ as _or
|
||||
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))
|
||||
)
|
||||
if target_account_id:
|
||||
q = q.filter(_EA.id == target_account_id)
|
||||
return q.all()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
initial_accounts = _enumerate_enabled_accounts()
|
||||
initial_account_ids = {
|
||||
str(account.id) for account in initial_accounts
|
||||
}
|
||||
|
||||
# Register every account before IMAP work, including its first-ever
|
||||
# scan. A concurrent zero-account cleanup can then advance this marker
|
||||
# and fence delivery even before the scan has produced payload.
|
||||
registered_state = None
|
||||
if initial_account_ids:
|
||||
async def _register_accounts(prior):
|
||||
next_state = _merge_email_urgency_state(
|
||||
prior,
|
||||
owner=owner,
|
||||
per_uid_scores={},
|
||||
notified_uids=prior.get("notified_uids", []),
|
||||
all_unread_keys=set(),
|
||||
fully_scanned_account_ids=set(),
|
||||
base_account_generations=(
|
||||
_email_urgency_account_generations(prior)
|
||||
),
|
||||
timestamp=_time.time(),
|
||||
known_account_ids=initial_account_ids,
|
||||
)
|
||||
# Return the exact state committed by registration. This is
|
||||
# the scan's generation token: adopting a later checkpoint
|
||||
# after account cleanup would let the stale scan appear fresh.
|
||||
return next_state, next_state
|
||||
|
||||
registered_state = await _run_email_urgency_state_transaction(
|
||||
STATE_PATH,
|
||||
STATE_LOCK_DB,
|
||||
_register_accounts,
|
||||
)
|
||||
|
||||
# Revalidate after registration. If deletion/disable and its cleanup
|
||||
# completed before the marker was published, this second enumeration
|
||||
# observes the absence and this action retires its own marker instead
|
||||
# of starting IMAP. Accounts newly appearing between the two reads are
|
||||
# left for the next pass rather than scanned without prior registration.
|
||||
verified_accounts = _enumerate_enabled_accounts()
|
||||
enabled_account_ids = {
|
||||
str(account.id) for account in verified_accounts
|
||||
}
|
||||
accounts = [
|
||||
account
|
||||
for account in verified_accounts
|
||||
if str(account.id) in initial_account_ids
|
||||
]
|
||||
|
||||
# Capture the checkpoint basis before cleanup or IMAP. A full
|
||||
# owner-wide enumeration authoritatively retires all known state IDs
|
||||
# absent from the current enabled/visible set. A scoped task may retire
|
||||
# only its selected missing/disabled account. Existing accounts remain
|
||||
# present even if their later network scan fails, so transient IMAP
|
||||
# failure never erases their last known state.
|
||||
base_state = (
|
||||
registered_state
|
||||
if registered_state is not None
|
||||
else _read_email_urgency_state(STATE_PATH)
|
||||
)
|
||||
base_account_generations = _email_urgency_account_generations(
|
||||
base_state
|
||||
)
|
||||
base_payload_account_ids = _email_urgency_payload_account_ids(base_state)
|
||||
known_state_account_ids = _email_urgency_known_account_ids(base_state)
|
||||
if target_account_id:
|
||||
retired_account_ids = (
|
||||
{str(target_account_id)}
|
||||
if str(target_account_id) not in enabled_account_ids
|
||||
else set()
|
||||
)
|
||||
else:
|
||||
retired_account_ids = (
|
||||
known_state_account_ids - enabled_account_ids
|
||||
)
|
||||
|
||||
if retired_account_ids:
|
||||
async def _retire_accounts(prior):
|
||||
next_state = _merge_email_urgency_state(
|
||||
prior,
|
||||
owner=owner,
|
||||
per_uid_scores={},
|
||||
notified_uids=prior.get("notified_uids", []),
|
||||
all_unread_keys=set(),
|
||||
fully_scanned_account_ids=set(),
|
||||
base_account_generations=base_account_generations,
|
||||
timestamp=_time.time(),
|
||||
retired_account_ids=retired_account_ids,
|
||||
base_payload_account_ids=base_payload_account_ids,
|
||||
)
|
||||
return None, next_state
|
||||
|
||||
await _run_email_urgency_state_transaction(
|
||||
STATE_PATH,
|
||||
STATE_LOCK_DB,
|
||||
_retire_accounts,
|
||||
)
|
||||
if not accounts:
|
||||
raise TaskNoop("no email accounts configured")
|
||||
|
||||
# ── 2. Account retirement above is state maintenance and does not
|
||||
# depend on model availability. Scanning still requires the utility
|
||||
# primary/fallback candidates resolved for this task owner.
|
||||
if not candidates:
|
||||
return "No LLM endpoint available", False
|
||||
|
||||
urgency_prompt = settings.get("urgent_email_prompt", "")
|
||||
per_uid_scores = {} # key = "<acc_id>:<uid>" → {"score": 0-3, "reason": "..."}
|
||||
all_unread_keys = set()
|
||||
@@ -1929,6 +2428,7 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
failed_classifications = []
|
||||
tag_write_details = []
|
||||
scanned = 0
|
||||
fully_scanned_account_ids = set()
|
||||
|
||||
def _heuristic_email_verdict(item: dict) -> dict:
|
||||
blob = (
|
||||
@@ -2024,16 +2524,27 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
def _scan_one(account=acc, cache_uids=cache.get("uids", {})):
|
||||
"""Sync IMAP work runs in a thread."""
|
||||
results = []
|
||||
scan_complete = True
|
||||
conn = _imap_connect(account.id)
|
||||
try:
|
||||
conn.select("INBOX", readonly=True)
|
||||
select_status, _select_data = conn.select("INBOX", readonly=True)
|
||||
if select_status != "OK":
|
||||
return results, False
|
||||
# Tag recent inbox mail, not only unread mail. Urgency
|
||||
# reminders below still only notify for unread messages.
|
||||
since_str = AGE_CUTOFF.strftime("%d-%b-%Y")
|
||||
status, data = conn.uid("SEARCH", None, f'(SINCE {since_str})')
|
||||
if status != "OK" or not data or not data[0]:
|
||||
return results
|
||||
uids = data[0].split()[-30:]
|
||||
if status != "OK":
|
||||
return results, False
|
||||
if not data or not data[0]:
|
||||
return results, True
|
||||
matching_uids = data[0].split()
|
||||
if len(matching_uids) > 30:
|
||||
# The scale guard deliberately processes only the most
|
||||
# recent 30. That is a partial account snapshot, so it
|
||||
# cannot justify pruning older checkpoint facts.
|
||||
scan_complete = False
|
||||
uids = matching_uids[-30:]
|
||||
for uid_b in uids:
|
||||
uid = uid_b.decode() if isinstance(uid_b, bytes) else str(uid_b)
|
||||
key = f"{account.id}:{uid}"
|
||||
@@ -2041,12 +2552,41 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
cached_ok = isinstance(cached, dict) and cached.get("triage_version") == TRIAGE_VERSION
|
||||
results.append({"key": key, "uid": uid, "cached": cached if cached_ok else None})
|
||||
if cached_ok:
|
||||
# Already classified — skip the fetch.
|
||||
# Cached verdicts still need a lightweight FLAGS
|
||||
# refresh. Without it a cached unread message looks
|
||||
# read and its successful notification checkpoint
|
||||
# is pruned on the next pass.
|
||||
try:
|
||||
st, flag_data = conn.uid("FETCH", uid_b, "(UID FLAGS)")
|
||||
if st != "OK" or not flag_data:
|
||||
scan_complete = False
|
||||
results.pop()
|
||||
continue
|
||||
flag_parts = []
|
||||
for part in flag_data:
|
||||
if isinstance(part, (bytes, bytearray)):
|
||||
flag_parts.append(bytes(part))
|
||||
elif (
|
||||
isinstance(part, tuple)
|
||||
and part
|
||||
and isinstance(part[0], (bytes, bytearray))
|
||||
):
|
||||
flag_parts.append(bytes(part[0]))
|
||||
flags_blob = b" ".join(flag_parts)
|
||||
results[-1]["unread"] = b"\\Seen" not in flags_blob
|
||||
except Exception as _fe:
|
||||
scan_complete = False
|
||||
results.pop()
|
||||
logger.debug(
|
||||
f"urgency: flag fetch for uid {uid} failed: {_fe}"
|
||||
)
|
||||
continue
|
||||
# Pull headers + first ~800 chars of plaintext body.
|
||||
try:
|
||||
st, msg_data = conn.uid("FETCH", uid_b, "(UID FLAGS RFC822.HEADER BODY.PEEK[TEXT]<0.800>)")
|
||||
if st != "OK" or not msg_data:
|
||||
scan_complete = False
|
||||
results.pop()
|
||||
continue
|
||||
flags_blob = b" ".join(
|
||||
part[0] for part in msg_data
|
||||
@@ -2060,6 +2600,8 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
if isinstance(part, tuple) and part[1]:
|
||||
raw += part[1] + b"\n\n"
|
||||
if not raw:
|
||||
scan_complete = False
|
||||
results.pop()
|
||||
continue
|
||||
msg = _email_mod.message_from_bytes(raw)
|
||||
# Skip Odysseus-generated reminders so the scanner
|
||||
@@ -2115,17 +2657,21 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
"unread": is_unread,
|
||||
})
|
||||
except Exception as _fe:
|
||||
scan_complete = False
|
||||
results.pop()
|
||||
logger.debug(f"urgency: header fetch for uid {uid} failed: {_fe}")
|
||||
finally:
|
||||
try: conn.logout()
|
||||
except Exception: pass
|
||||
return results
|
||||
return results, scan_complete
|
||||
|
||||
try:
|
||||
items = await _aio.to_thread(_scan_one)
|
||||
items, scan_complete = await _aio.to_thread(_scan_one)
|
||||
except Exception as e:
|
||||
logger.warning(f"urgency: IMAP scan failed for account {acc.id}: {e}")
|
||||
continue
|
||||
if scan_complete:
|
||||
fully_scanned_account_ids.add(str(acc.id))
|
||||
|
||||
for item in items:
|
||||
scanned += 1
|
||||
@@ -2262,13 +2808,13 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
logger.debug(f"urgency: LLM classify failed for {key}: {e}")
|
||||
continue
|
||||
|
||||
# ── Prune cache entries for UIDs that are no longer in the recent
|
||||
# scan window. Read messages remain cached because tags are useful
|
||||
# on read mail too; unread state is refreshed per scan above.
|
||||
seen_uids = {it["uid"] for it in items}
|
||||
cache_uids = cache.get("uids", {})
|
||||
for stale in [u for u in cache_uids if u not in seen_uids]:
|
||||
cache_uids.pop(stale, None)
|
||||
if scan_complete:
|
||||
# Only a complete account scan proves a cached UID left the
|
||||
# recent window. Partial/failing scans preserve prior facts.
|
||||
seen_uids = {it["uid"] for it in items}
|
||||
cache_uids = cache.get("uids", {})
|
||||
for stale in [u for u in cache_uids if u not in seen_uids]:
|
||||
cache_uids.pop(stale, None)
|
||||
|
||||
try:
|
||||
cache_file.write_text(_json.dumps(cache), encoding="utf-8")
|
||||
@@ -2372,40 +2918,34 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
|
||||
# ── 4. Aggregate state. urgent = score ≥ 2.
|
||||
urgent_keys = [k for k, v in per_uid_scores.items() if v.get("score", 0) >= 2 and v.get("unread")]
|
||||
max_score = max((v.get("score", 0) for v in per_uid_scores.values()), default=0)
|
||||
total_urgent = len(urgent_keys)
|
||||
|
||||
# Load prior state to know which urgent UIDs we've already notified.
|
||||
try:
|
||||
prior = _json.loads(STATE_PATH.read_text(encoding="utf-8")) if STATE_PATH.exists() else {}
|
||||
except Exception:
|
||||
prior = {}
|
||||
notified_uids = set(prior.get("notified_uids", []))
|
||||
|
||||
# ── 5. Fire reminder ONLY when a previously-unnotified UID scores urgent.
|
||||
new_urgent = [k for k in urgent_keys if k not in notified_uids]
|
||||
# ── 5. Fire a reminder only when a previously-unnotified UID scores
|
||||
# urgent. The read, decision, delivery, and checkpoint are serialized
|
||||
# below so two scheduler workers cannot both act on the same stale
|
||||
# state or overwrite each other's successful checkpoint.
|
||||
newly_notified = set()
|
||||
notify_failed = set()
|
||||
if new_urgent:
|
||||
title = "Urgent email" if total_urgent == 1 else f"{total_urgent} urgent emails"
|
||||
# Build a real listing — subject · sender · reason for each urgent
|
||||
# one — so the reminder email tells you which messages to act on,
|
||||
# not just "4 needing reply". Optional deep-link when the user has
|
||||
# `app_public_url` configured in Settings (so the email row links
|
||||
# straight into the Odysseus Email tab).
|
||||
# Sort: highest-scored UIDs first; cap at 10 to keep the email tidy.
|
||||
|
||||
def _urgency_reminder_payload(reminder_keys):
|
||||
total = len(reminder_keys)
|
||||
title = "Urgent email" if total == 1 else f"{total} urgent emails"
|
||||
sorted_urgent = sorted(
|
||||
((k, per_uid_scores[k]) for k in urgent_keys),
|
||||
key=lambda kv: kv[1].get("score", 0), reverse=True,
|
||||
((key, per_uid_scores[key]) for key in reminder_keys),
|
||||
key=lambda item: item[1].get("score", 0),
|
||||
reverse=True,
|
||||
)[:10]
|
||||
_pub = (settings.get("app_public_url") or "").strip().rstrip("/")
|
||||
from urllib.parse import quote as _quote
|
||||
lines = [f"{total_urgent} email" + ("" if total_urgent == 1 else "s") + " need an urgent reply:", ""]
|
||||
for i, (k, v) in enumerate(sorted_urgent, 1):
|
||||
subj = (v.get("subject") or "(no subject)")[:160]
|
||||
frm = v.get("from") or ""
|
||||
why = v.get("reason") or ""
|
||||
uid_for_link = str(k).split(":", 1)[-1]
|
||||
lines = [
|
||||
f"{total} email" + ("" if total == 1 else "s")
|
||||
+ " need an urgent reply:",
|
||||
"",
|
||||
]
|
||||
for i, (key, value) in enumerate(sorted_urgent, 1):
|
||||
subj = (value.get("subject") or "(no subject)")[:160]
|
||||
frm = value.get("from") or ""
|
||||
why = value.get("reason") or ""
|
||||
uid_for_link = str(key).split(":", 1)[-1]
|
||||
hash_link = f"#email={_quote('INBOX', safe='')}:{uid_for_link}"
|
||||
open_link = f"{_pub}/{hash_link}" if _pub else hash_link
|
||||
line = f"{i}. {subj}"
|
||||
@@ -2415,57 +2955,94 @@ async def action_check_email_urgency(owner: str, **kwargs) -> Tuple[str, bool]:
|
||||
line += f" · {why}"
|
||||
lines.append(line)
|
||||
lines.append(f" Open email: {open_link}")
|
||||
if total_urgent > len(sorted_urgent):
|
||||
if total > len(sorted_urgent):
|
||||
lines.append("")
|
||||
lines.append(f"…and {total_urgent - len(sorted_urgent)} more.")
|
||||
body = "\n".join(lines)
|
||||
try:
|
||||
# Call dispatch_reminder DIRECTLY (no HTTP/auth roundtrip — the
|
||||
# endpoint version 401's the background scheduler because it
|
||||
# has no session cookie).
|
||||
from routes.note_routes import dispatch_reminder
|
||||
dispatch_result = await dispatch_reminder(
|
||||
title=title, note_body=body, note_id="urgent-email",
|
||||
owner=owner or "",
|
||||
)
|
||||
channel = (settings.get("reminder_channel") or "browser").strip().lower()
|
||||
delivered = bool(dispatch_result.get("browser_sent"))
|
||||
if channel == "email":
|
||||
delivered = bool(dispatch_result.get("email_sent"))
|
||||
elif channel == "ntfy":
|
||||
delivered = bool(dispatch_result.get("ntfy_sent"))
|
||||
elif channel == "webhook":
|
||||
delivered = bool(dispatch_result.get("webhook_sent"))
|
||||
if delivered:
|
||||
newly_notified.update(new_urgent)
|
||||
else:
|
||||
lines.append(f"…and {total - len(sorted_urgent)} more.")
|
||||
return title, "\n".join(lines)
|
||||
|
||||
async def _dispatch_urgency_reminder(reminder_keys):
|
||||
# Call dispatch_reminder directly: a scheduler has no browser
|
||||
# session cookie with which to call the HTTP endpoint.
|
||||
from routes.note_routes import dispatch_reminder
|
||||
title, body = _urgency_reminder_payload(reminder_keys)
|
||||
return await dispatch_reminder(
|
||||
title=title,
|
||||
note_body=body,
|
||||
note_id="urgent-email",
|
||||
owner=owner or "",
|
||||
)
|
||||
|
||||
async def _dispatch_and_checkpoint(prior):
|
||||
notified_uids = _email_urgency_string_set(
|
||||
prior.get("notified_uids", [])
|
||||
)
|
||||
observed_accounts = {
|
||||
_email_urgency_account_key(key) for key in per_uid_scores
|
||||
} | fully_scanned_account_ids
|
||||
stale_accounts = _email_urgency_stale_accounts(
|
||||
prior,
|
||||
base_account_generations,
|
||||
observed_accounts,
|
||||
)
|
||||
# Generation fencing must happen before delivery, not only during
|
||||
# merge. A stale-only unread UID may have been removed, read, or
|
||||
# downgraded by the newer completed scan.
|
||||
deliverable_urgent = [
|
||||
key
|
||||
for key in urgent_keys
|
||||
if _email_urgency_account_key(key) not in stale_accounts
|
||||
]
|
||||
new_urgent = [
|
||||
key
|
||||
for key in deliverable_urgent
|
||||
if key not in notified_uids
|
||||
]
|
||||
if new_urgent:
|
||||
try:
|
||||
dispatch_result = await _dispatch_urgency_reminder(
|
||||
deliverable_urgent
|
||||
)
|
||||
channel = (settings.get("reminder_channel") or "browser").strip().lower()
|
||||
delivered = bool(dispatch_result.get("browser_sent"))
|
||||
if channel == "email":
|
||||
delivered = bool(dispatch_result.get("email_sent"))
|
||||
elif channel == "ntfy":
|
||||
delivered = bool(dispatch_result.get("ntfy_sent"))
|
||||
elif channel == "webhook":
|
||||
delivered = bool(dispatch_result.get("webhook_sent"))
|
||||
if delivered:
|
||||
newly_notified.update(new_urgent)
|
||||
notified_uids.update(new_urgent)
|
||||
else:
|
||||
notify_failed.update(new_urgent)
|
||||
logger.warning(
|
||||
"urgency: reminder dispatch returned no successful "
|
||||
f"delivery path: {dispatch_result}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"urgency: reminder dispatch failed: {e}")
|
||||
notify_failed.update(new_urgent)
|
||||
logger.warning(f"urgency: reminder dispatch returned no successful delivery path: {dispatch_result}")
|
||||
except Exception as e:
|
||||
logger.warning(f"urgency: reminder dispatch failed: {e}")
|
||||
notify_failed.update(new_urgent)
|
||||
# Mark only successfully delivered UIDs as notified so a transient
|
||||
# SMTP/ntfy/browser failure retries instead of lying forever.
|
||||
notified_uids.update(newly_notified)
|
||||
|
||||
# Prune notified_uids that aren't unread anymore (so a future re-urgent
|
||||
# message with the same UID — rare but possible after archive→unarchive
|
||||
# — can re-notify). Keep only UIDs still in `all_unread_keys`.
|
||||
notified_uids = {u for u in notified_uids if u in all_unread_keys}
|
||||
next_state = _merge_email_urgency_state(
|
||||
prior,
|
||||
owner=owner,
|
||||
per_uid_scores=per_uid_scores,
|
||||
notified_uids=notified_uids,
|
||||
all_unread_keys=all_unread_keys,
|
||||
fully_scanned_account_ids=fully_scanned_account_ids,
|
||||
base_account_generations=base_account_generations,
|
||||
timestamp=_time.time(),
|
||||
)
|
||||
return notified_uids, next_state
|
||||
|
||||
state = {
|
||||
"ts": _time.time(),
|
||||
"owner": owner or "",
|
||||
"total_unread": len(all_unread_keys),
|
||||
"total_urgent": total_urgent,
|
||||
"max_score": max_score,
|
||||
"per_uid": per_uid_scores,
|
||||
"notified_uids": sorted(notified_uids),
|
||||
}
|
||||
try:
|
||||
STATE_PATH.write_text(_json.dumps(state), encoding="utf-8")
|
||||
await _run_email_urgency_state_transaction(
|
||||
STATE_PATH,
|
||||
STATE_LOCK_DB,
|
||||
_dispatch_and_checkpoint,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"urgency: state write failed: {e}")
|
||||
logger.warning(f"urgency: state transaction failed: {e}")
|
||||
|
||||
# ── 6. Activity-log summary — counts line on top, then per-tier
|
||||
# bulleted breakdown so the user can see WHICH emails ranked where
|
||||
|
||||
@@ -282,7 +282,9 @@ def trim_for_context(messages: List[Dict], context_length: int, reserve_tokens:
|
||||
if essential_system:
|
||||
sys_text = essential_system[0].get("content", "")
|
||||
if len(sys_text) > 2000:
|
||||
essential_system[0] = {"role": "system", "content": sys_text[:2000] + "\n[System prompt truncated for context limits]"}
|
||||
truncated_system = dict(essential_system[0])
|
||||
truncated_system["content"] = sys_text[:2000] + "\n[System prompt truncated for context limits]"
|
||||
essential_system[0] = truncated_system
|
||||
trimmed = essential_system + convo_msgs
|
||||
if estimate_tokens(trimmed) <= budget:
|
||||
return _sanitize_tool_messages(essential_system + protected_msgs + convo_msgs)
|
||||
@@ -325,6 +327,9 @@ async def maybe_compact(
|
||||
messages: List[Dict],
|
||||
headers: Optional[Dict] = None,
|
||||
owner: Optional[str] = None,
|
||||
*,
|
||||
persist: bool = True,
|
||||
compaction_state: Optional[Dict[str, Any]] = None,
|
||||
) -> tuple:
|
||||
"""Check context usage and compact if above threshold.
|
||||
|
||||
@@ -416,7 +421,17 @@ async def maybe_compact(
|
||||
# offset — session.history INCLUDES the system messages, but
|
||||
# split_point is indexed against convo_msgs which does NOT. Without
|
||||
# this, the slice drops the leading system message(s).
|
||||
_update_session_history(session, split_point, summary, system_msg_count=len(system_msgs))
|
||||
if compaction_state is not None:
|
||||
compaction_state.update({
|
||||
"split_point": split_point,
|
||||
"summary": summary,
|
||||
"system_msg_count": len(system_msgs),
|
||||
"applied": False,
|
||||
})
|
||||
if persist:
|
||||
_update_session_history(session, split_point, summary, system_msg_count=len(system_msgs))
|
||||
if compaction_state is not None:
|
||||
compaction_state["applied"] = True
|
||||
|
||||
new_used = estimate_tokens(compacted)
|
||||
logger.info(
|
||||
@@ -427,6 +442,51 @@ async def maybe_compact(
|
||||
return compacted, context_length, True
|
||||
|
||||
|
||||
def apply_compaction_state(session, compaction_state: Optional[Dict[str, Any]]) -> bool:
|
||||
"""Persist a route-specific compaction after that route commits output.
|
||||
|
||||
Candidate prompts may be compacted speculatively while an explicit
|
||||
foreground fallback chain is being tried. Persisting at construction time
|
||||
would let an unavailable route rewrite history before another route answers,
|
||||
so callers hold this small plan and apply only the winning route's plan.
|
||||
"""
|
||||
|
||||
state = compaction_state if isinstance(compaction_state, dict) else None
|
||||
if not state or state.get("applied"):
|
||||
return False
|
||||
summary = state.get("summary")
|
||||
split_point = state.get("split_point")
|
||||
system_msg_count = state.get("system_msg_count", 0)
|
||||
if not isinstance(summary, str) or not isinstance(split_point, int):
|
||||
return False
|
||||
_update_session_history(
|
||||
session,
|
||||
split_point,
|
||||
summary,
|
||||
system_msg_count=system_msg_count if isinstance(system_msg_count, int) else 0,
|
||||
)
|
||||
state["applied"] = True
|
||||
return True
|
||||
|
||||
|
||||
def apply_compaction_state_for_session(
|
||||
session_id: Optional[str],
|
||||
compaction_state: Optional[Dict[str, Any]],
|
||||
) -> bool:
|
||||
"""Resolve an in-memory session and apply a deferred compaction plan."""
|
||||
|
||||
if not session_id:
|
||||
return False
|
||||
try:
|
||||
from core.models import get_session_manager_instance
|
||||
|
||||
manager = get_session_manager_instance()
|
||||
session = manager.get_session(session_id) if manager else None
|
||||
except Exception:
|
||||
session = None
|
||||
return apply_compaction_state(session, compaction_state) if session else False
|
||||
|
||||
|
||||
def _update_session_history(session, split_point: int, summary: str,
|
||||
system_msg_count: int = 0):
|
||||
"""Update the in-memory session history after compaction.
|
||||
|
||||
+215
-33
@@ -5,6 +5,7 @@ Consolidates the 4+ copies of normalize_base / resolve_endpoint logic into one p
|
||||
"""
|
||||
|
||||
import json
|
||||
import ipaddress
|
||||
import logging
|
||||
import socket
|
||||
import subprocess
|
||||
@@ -27,6 +28,43 @@ _NON_CHAT_MODEL = (
|
||||
)
|
||||
|
||||
|
||||
def endpoint_cost_tracked(url: str, endpoint_kind: Optional[str] = None) -> bool:
|
||||
"""Return whether token cost should be tracked for a concrete route.
|
||||
|
||||
This is intentionally a non-secret route classification. It mirrors the
|
||||
frontend's local/subscription exclusions without exposing endpoint URLs to
|
||||
message metadata.
|
||||
"""
|
||||
|
||||
try:
|
||||
parsed = urlparse(url or "")
|
||||
host = (parsed.hostname or "").lower().rstrip(".")
|
||||
path = (parsed.path or "").rstrip("/")
|
||||
except Exception:
|
||||
return False
|
||||
if not host:
|
||||
return False
|
||||
if host == "chatgpt.com" and (
|
||||
path == "/backend-api/codex" or path.startswith("/backend-api/codex/")
|
||||
):
|
||||
return False
|
||||
kind = str(endpoint_kind or "auto").strip().lower()
|
||||
if kind == "local":
|
||||
return False
|
||||
if kind in {"api", "proxy"}:
|
||||
return True
|
||||
if host in {"localhost", "0.0.0.0", "host.docker.internal"} or host.endswith(".local"):
|
||||
return False
|
||||
try:
|
||||
ip = ipaddress.ip_address(host)
|
||||
return ip.is_global
|
||||
except ValueError:
|
||||
pass
|
||||
if "." not in host:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _first_chat_model(models) -> Optional[str]:
|
||||
"""First model that isn't an embedding/tts/etc.; falls back to models[0]."""
|
||||
for m in (models or []):
|
||||
@@ -396,10 +434,14 @@ def resolve_endpoint(
|
||||
db.close()
|
||||
|
||||
|
||||
def resolve_endpoint_by_id(
|
||||
ep_id: str, model: Optional[str] = None, owner: Optional[str] = None
|
||||
) -> Optional[Tuple[str, str, Dict]]:
|
||||
"""Resolve a specific endpoint id (+ optional model) to (chat_url, model, headers).
|
||||
def _resolve_endpoint_by_id_with_descriptor(
|
||||
ep_id: str,
|
||||
model: Optional[str] = None,
|
||||
owner: Optional[str] = None,
|
||||
*,
|
||||
require_exact_model: bool = False,
|
||||
) -> Optional[Tuple[Tuple[str, str, Dict], dict]]:
|
||||
"""Resolve a concrete endpoint/model plus its non-secret descriptor.
|
||||
|
||||
Returns None if the endpoint doesn't exist or is disabled. Used to turn
|
||||
a configured fallback entry ({endpoint_id, model}) into a dispatch target.
|
||||
@@ -426,15 +468,34 @@ def resolve_endpoint_by_id(
|
||||
chat_url = build_chat_url(base)
|
||||
headers = build_headers(api_key, base)
|
||||
m = (model or "").strip()
|
||||
# Drop a model the user disabled on the endpoint, then pick the first
|
||||
# enabled chat model rather than a hidden one.
|
||||
if m and m in _endpoint_hidden_models(ep):
|
||||
m = ""
|
||||
if not m:
|
||||
m = _first_chat_model(_endpoint_enabled_models(ep)) or ""
|
||||
enabled_models = _endpoint_enabled_models(ep)
|
||||
if require_exact_model:
|
||||
# Explicit foreground fallback entries are concrete choices. A
|
||||
# hidden or known-missing model must disable the entry instead of
|
||||
# silently substituting another model from the endpoint.
|
||||
if not m or m in _endpoint_hidden_models(ep):
|
||||
return None
|
||||
if enabled_models and m not in enabled_models:
|
||||
return None
|
||||
else:
|
||||
# Legacy Utility/Vision chains retain their model-repair behavior.
|
||||
if m and m in _endpoint_hidden_models(ep):
|
||||
m = ""
|
||||
if not m:
|
||||
m = _first_chat_model(enabled_models) or ""
|
||||
if not m:
|
||||
return None
|
||||
return chat_url, m, headers
|
||||
return (
|
||||
(chat_url, m, headers),
|
||||
{
|
||||
"endpoint_id": ep.id,
|
||||
"endpoint_label": getattr(ep, "name", None) or ep.id,
|
||||
"endpoint_cost_tracked": endpoint_cost_tracked(
|
||||
chat_url,
|
||||
getattr(ep, "endpoint_kind", None),
|
||||
),
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Could not resolve endpoint {ep_id}: {e}")
|
||||
return None
|
||||
@@ -442,29 +503,105 @@ def resolve_endpoint_by_id(
|
||||
db.close()
|
||||
|
||||
|
||||
def resolve_chat_fallback_candidates(owner: Optional[str] = None) -> list:
|
||||
"""Build the configured default-chat fallback chain as a list of
|
||||
(chat_url, model, headers) tuples, skipping any that can't resolve.
|
||||
def resolve_endpoint_by_id(
|
||||
ep_id: str,
|
||||
model: Optional[str] = None,
|
||||
owner: Optional[str] = None,
|
||||
*,
|
||||
require_exact_model: bool = False,
|
||||
) -> Optional[Tuple[str, str, Dict]]:
|
||||
"""Resolve a specific endpoint id (+ optional model) to its runtime route."""
|
||||
|
||||
The primary model is NOT included — callers prepend their session's
|
||||
current (url, model, headers) so per-session model overrides are honored.
|
||||
resolved = _resolve_endpoint_by_id_with_descriptor(
|
||||
ep_id,
|
||||
model,
|
||||
owner=owner,
|
||||
require_exact_model=require_exact_model,
|
||||
)
|
||||
return resolved[0] if resolved else None
|
||||
|
||||
|
||||
def resolve_route_descriptor(
|
||||
endpoint_url: str,
|
||||
model: str,
|
||||
headers: Optional[Dict] = None,
|
||||
owner: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""Return the visible endpoint identity for an already-resolved route.
|
||||
|
||||
Headers are compared only inside the process so two endpoints using the
|
||||
same provider URL/model but different credentials remain distinguishable.
|
||||
No credential material is returned or logged.
|
||||
"""
|
||||
return _resolve_fallback_candidates("default_model_fallbacks", owner=owner)
|
||||
|
||||
if not endpoint_url or not model:
|
||||
return {
|
||||
"endpoint_id": None,
|
||||
"endpoint_label": "Selected route",
|
||||
"endpoint_cost_tracked": endpoint_cost_tracked(endpoint_url),
|
||||
}
|
||||
db = SessionLocal()
|
||||
try:
|
||||
q = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True)
|
||||
if owner:
|
||||
from src.auth_helpers import owner_filter
|
||||
q = owner_filter(q, ModelEndpoint, owner)
|
||||
expected = (endpoint_url.rstrip("/"), model, headers or {})
|
||||
for ep in q.all():
|
||||
resolved = _resolve_endpoint_by_id_with_descriptor(
|
||||
ep.id,
|
||||
model,
|
||||
owner=owner,
|
||||
require_exact_model=True,
|
||||
)
|
||||
if not resolved:
|
||||
continue
|
||||
candidate, descriptor = resolved
|
||||
actual = (candidate[0].rstrip("/"), candidate[1], candidate[2] or {})
|
||||
if actual == expected:
|
||||
return descriptor
|
||||
except Exception as e:
|
||||
logger.debug("Could not identify selected endpoint route: %s", e)
|
||||
finally:
|
||||
db.close()
|
||||
return {
|
||||
"endpoint_id": None,
|
||||
"endpoint_label": "Selected route",
|
||||
"endpoint_cost_tracked": endpoint_cost_tracked(endpoint_url),
|
||||
}
|
||||
|
||||
|
||||
def resolve_route_descriptor_by_id(
|
||||
endpoint_id: str,
|
||||
endpoint_url: str,
|
||||
model: str,
|
||||
headers: Optional[Dict] = None,
|
||||
owner: Optional[str] = None,
|
||||
) -> Optional[dict]:
|
||||
"""Resolve a selected route's identity without relying on row order.
|
||||
|
||||
The explicit endpoint id is still verified against the resolved runtime
|
||||
route. This prevents stale or mismatched request metadata from being used
|
||||
for attribution while disambiguating endpoints whose routes are otherwise
|
||||
identical.
|
||||
"""
|
||||
|
||||
resolved = _resolve_endpoint_by_id_with_descriptor(
|
||||
endpoint_id,
|
||||
model,
|
||||
owner=owner,
|
||||
require_exact_model=True,
|
||||
)
|
||||
if not resolved:
|
||||
return None
|
||||
candidate, descriptor = resolved
|
||||
expected = ((endpoint_url or "").rstrip("/"), model, headers or {})
|
||||
actual = (candidate[0].rstrip("/"), candidate[1], candidate[2] or {})
|
||||
return descriptor if actual == expected else None
|
||||
|
||||
|
||||
def resolve_utility_fallback_candidates(owner: Optional[str] = None) -> list:
|
||||
"""Configured fallback chain for the Utility model (`utility_model_fallbacks`)."""
|
||||
try:
|
||||
from src.settings import get_user_setting, load_settings
|
||||
settings = load_settings()
|
||||
utility_ep = (get_user_setting("utility_endpoint_id", owner or "", settings.get("utility_endpoint_id", "")) or "").strip()
|
||||
if not utility_ep:
|
||||
utility_chain = get_user_setting("utility_model_fallbacks", owner or "", settings.get("utility_model_fallbacks") or []) or []
|
||||
if utility_chain:
|
||||
return _resolve_fallback_candidates("utility_model_fallbacks", owner=owner)
|
||||
return _resolve_fallback_candidates("default_model_fallbacks", owner=owner)
|
||||
except Exception:
|
||||
pass
|
||||
return _resolve_fallback_candidates("utility_model_fallbacks", owner=owner)
|
||||
|
||||
|
||||
@@ -474,17 +611,62 @@ def resolve_vision_fallback_candidates(owner: Optional[str] = None) -> list:
|
||||
|
||||
|
||||
def _resolve_fallback_candidates(setting_key: str, owner: Optional[str] = None) -> list:
|
||||
out = []
|
||||
try:
|
||||
from src.settings import get_user_setting, load_settings
|
||||
settings = load_settings()
|
||||
chain = get_user_setting(setting_key, owner or "", settings.get(setting_key) or []) or []
|
||||
except Exception:
|
||||
return out
|
||||
for entry in chain:
|
||||
return []
|
||||
return resolve_fallback_entries(chain, owner=owner)
|
||||
|
||||
|
||||
def resolve_fallback_entries(
|
||||
entries,
|
||||
owner: Optional[str] = None,
|
||||
*,
|
||||
require_exact_model: bool = False,
|
||||
) -> list:
|
||||
"""Resolve ordered endpoint/model entries within the caller's owner scope."""
|
||||
|
||||
out = []
|
||||
for entry in entries or []:
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
resolved = resolve_endpoint_by_id(entry.get("endpoint_id", ""), entry.get("model", ""), owner=owner)
|
||||
if resolved:
|
||||
resolved = resolve_endpoint_by_id(
|
||||
entry.get("endpoint_id", ""),
|
||||
entry.get("model", ""),
|
||||
owner=owner,
|
||||
require_exact_model=require_exact_model,
|
||||
)
|
||||
if resolved and resolved not in out:
|
||||
out.append(resolved)
|
||||
return out
|
||||
|
||||
|
||||
def resolve_fallback_entries_with_descriptors(
|
||||
entries,
|
||||
owner: Optional[str] = None,
|
||||
*,
|
||||
require_exact_model: bool = False,
|
||||
) -> list:
|
||||
"""Resolve ordered entries while retaining safe endpoint provenance."""
|
||||
|
||||
out = []
|
||||
seen = []
|
||||
for entry in entries or []:
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
resolved = _resolve_endpoint_by_id_with_descriptor(
|
||||
entry.get("endpoint_id", ""),
|
||||
entry.get("model", ""),
|
||||
owner=owner,
|
||||
require_exact_model=require_exact_model,
|
||||
)
|
||||
if not resolved:
|
||||
continue
|
||||
candidate, descriptor = resolved
|
||||
if any(candidate == prior for prior in seen):
|
||||
continue
|
||||
seen.append(candidate)
|
||||
out.append((candidate, descriptor))
|
||||
return out
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
"""Explicit foreground Chat and Agent model-routing policy."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Collection, Dict, FrozenSet, Optional, Tuple
|
||||
|
||||
from src.endpoint_resolver import (
|
||||
endpoint_cost_tracked,
|
||||
resolve_fallback_entries,
|
||||
resolve_fallback_entries_with_descriptors,
|
||||
resolve_route_descriptor,
|
||||
resolve_route_descriptor_by_id,
|
||||
)
|
||||
|
||||
_DEFAULT_FALLBACK_ENTRY_RESOLVER = resolve_fallback_entries
|
||||
|
||||
|
||||
FOREGROUND_FALLBACK_ENABLED_KEY = "foreground_fallback_enabled"
|
||||
FOREGROUND_FALLBACK_LIST_KEY = "foreground_model_fallbacks"
|
||||
FOREGROUND_AVAILABILITY_STATUSES: FrozenSet[int] = frozenset({
|
||||
408, 425, 429, 500, 502, 503, 504, 507, 508, 529,
|
||||
})
|
||||
MAX_FOREGROUND_FALLBACKS = 10
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ForegroundModelPolicy:
|
||||
"""Resolved per-user foreground fallback policy."""
|
||||
|
||||
enabled: bool = False
|
||||
fallback_candidates: Tuple[tuple, ...] = ()
|
||||
fallback_descriptors: Tuple[dict, ...] = ()
|
||||
eligible_statuses: FrozenSet[int] = FOREGROUND_AVAILABILITY_STATUSES
|
||||
fallback_on_empty: bool = False
|
||||
|
||||
|
||||
def _load_policy_preferences(owner: Optional[str]) -> dict:
|
||||
"""Load only preferences that explicitly belong to ``owner``.
|
||||
|
||||
The generic preferences loader intentionally treats a legacy flat store as
|
||||
the single-user preferences object. That compatibility must not cross an
|
||||
authentication transition: once a named owner is present, foreground
|
||||
fallback consent exists only in an actual ``_users[owner]`` dictionary.
|
||||
"""
|
||||
|
||||
from routes import prefs_routes
|
||||
|
||||
if owner is None:
|
||||
prefs = prefs_routes._load_for_user(None)
|
||||
return dict(prefs) if isinstance(prefs, dict) else {}
|
||||
|
||||
raw = prefs_routes._load()
|
||||
users = raw.get("_users") if isinstance(raw, dict) else None
|
||||
if not isinstance(users, dict):
|
||||
return {}
|
||||
prefs = users.get(owner)
|
||||
return dict(prefs) if isinstance(prefs, dict) else {}
|
||||
|
||||
|
||||
def resolve_foreground_model_policy(
|
||||
owner: Optional[str] = None,
|
||||
allowed_models: Optional[Collection[str]] = None,
|
||||
) -> ForegroundModelPolicy:
|
||||
"""Resolve an explicit owner-scoped policy, failing closed to strict mode.
|
||||
|
||||
The policy is stored in user preferences even when authentication is
|
||||
disabled. Historical ``default_model_fallbacks`` values are deliberately
|
||||
unrelated and are never read or migrated.
|
||||
"""
|
||||
|
||||
try:
|
||||
prefs = _load_policy_preferences(owner)
|
||||
except Exception:
|
||||
return ForegroundModelPolicy()
|
||||
|
||||
if prefs.get(FOREGROUND_FALLBACK_ENABLED_KEY) is not True:
|
||||
return ForegroundModelPolicy()
|
||||
|
||||
entries = prefs.get(FOREGROUND_FALLBACK_LIST_KEY)
|
||||
if not isinstance(entries, list) or not entries:
|
||||
return ForegroundModelPolicy()
|
||||
if allowed_models is not None:
|
||||
allowed = frozenset(allowed_models)
|
||||
entries = [
|
||||
entry for entry in entries
|
||||
if (
|
||||
isinstance(entry, dict)
|
||||
and isinstance(entry.get("model"), str)
|
||||
and entry.get("model") in allowed
|
||||
)
|
||||
]
|
||||
if not entries:
|
||||
return ForegroundModelPolicy()
|
||||
entries = entries[:MAX_FOREGROUND_FALLBACKS]
|
||||
|
||||
if resolve_fallback_entries is not _DEFAULT_FALLBACK_ENTRY_RESOLVER:
|
||||
# Preserve the long-standing resolver seam used by downstream tests and
|
||||
# integrations. Production uses the descriptor-aware resolver below.
|
||||
compatibility_candidates = resolve_fallback_entries(
|
||||
entries,
|
||||
owner=owner,
|
||||
require_exact_model=True,
|
||||
)
|
||||
# Known limitation of this test-only seam: alignment matches on model
|
||||
# alone, so when two entries share a model and the resolver skips the
|
||||
# first, the surviving candidate inherits the skipped entry's
|
||||
# endpoint_id. Production uses the descriptor-aware branch below,
|
||||
# which is unaffected.
|
||||
resolved_routes = []
|
||||
remaining_entries = list(entries)
|
||||
for candidate in compatibility_candidates:
|
||||
matching_index = next(
|
||||
(
|
||||
index for index, entry in enumerate(remaining_entries)
|
||||
if isinstance(entry, dict)
|
||||
and entry.get("model") == candidate[1]
|
||||
),
|
||||
None,
|
||||
)
|
||||
matching_entry = (
|
||||
remaining_entries.pop(matching_index)
|
||||
if matching_index is not None
|
||||
else {}
|
||||
)
|
||||
descriptor = {
|
||||
"endpoint_id": matching_entry.get("endpoint_id"),
|
||||
"endpoint_label": matching_entry.get("endpoint_id") or "Fallback route",
|
||||
"endpoint_cost_tracked": endpoint_cost_tracked(candidate[0]),
|
||||
}
|
||||
resolved_routes.append((candidate, descriptor))
|
||||
else:
|
||||
resolved_routes = resolve_fallback_entries_with_descriptors(
|
||||
entries,
|
||||
owner=owner,
|
||||
require_exact_model=True,
|
||||
)
|
||||
candidates = [candidate for candidate, _descriptor in resolved_routes]
|
||||
if not candidates:
|
||||
return ForegroundModelPolicy()
|
||||
|
||||
return ForegroundModelPolicy(
|
||||
enabled=True,
|
||||
fallback_candidates=tuple(candidates),
|
||||
fallback_descriptors=tuple(
|
||||
dict(descriptor) for _candidate, descriptor in resolved_routes
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def resolve_foreground_fallback_candidates(owner: Optional[str] = None) -> list:
|
||||
"""Return only candidates explicitly enabled by the current user."""
|
||||
|
||||
return list(resolve_foreground_model_policy(owner).fallback_candidates)
|
||||
|
||||
|
||||
def build_foreground_model_candidates(
|
||||
endpoint_url: str,
|
||||
model: str,
|
||||
headers: Optional[Dict[str, Any]] = None,
|
||||
owner: Optional[str] = None,
|
||||
policy: Optional[ForegroundModelPolicy] = None,
|
||||
) -> list:
|
||||
"""Build the ordered candidate list for a foreground request."""
|
||||
|
||||
policy = policy or resolve_foreground_model_policy(owner)
|
||||
primary = (endpoint_url, model, headers or {})
|
||||
candidates = [primary]
|
||||
for candidate in policy.fallback_candidates:
|
||||
if candidate not in candidates:
|
||||
candidates.append(candidate)
|
||||
return candidates
|
||||
|
||||
|
||||
def build_foreground_route_descriptors(
|
||||
endpoint_url: str,
|
||||
model: str,
|
||||
headers: Optional[Dict[str, Any]] = None,
|
||||
owner: Optional[str] = None,
|
||||
policy: Optional[ForegroundModelPolicy] = None,
|
||||
selected_endpoint_id: Optional[str] = None,
|
||||
) -> list:
|
||||
"""Build safe route metadata parallel to foreground candidates."""
|
||||
|
||||
policy = policy or resolve_foreground_model_policy(owner)
|
||||
selected = None
|
||||
if selected_endpoint_id:
|
||||
selected = resolve_route_descriptor_by_id(
|
||||
selected_endpoint_id,
|
||||
endpoint_url,
|
||||
model,
|
||||
headers or {},
|
||||
owner=owner,
|
||||
)
|
||||
if selected is None:
|
||||
selected = resolve_route_descriptor(endpoint_url, model, headers or {}, owner=owner)
|
||||
primary = (endpoint_url, model, headers or {})
|
||||
candidates = [primary]
|
||||
descriptors = [selected]
|
||||
for candidate, descriptor in zip(
|
||||
policy.fallback_candidates,
|
||||
policy.fallback_descriptors,
|
||||
):
|
||||
if candidate in candidates:
|
||||
continue
|
||||
candidates.append(candidate)
|
||||
descriptors.append(dict(descriptor))
|
||||
return descriptors
|
||||
@@ -63,8 +63,11 @@ _PASSIVE_EXACT_PATHS = {
|
||||
"/api/activity/heartbeat",
|
||||
"/api/client-perf",
|
||||
"/api/tasks/notifications",
|
||||
"/api/tasks/runs/recent",
|
||||
"/api/research/active",
|
||||
"/api/email/urgency-state",
|
||||
# UI idle poll sibling of urgency-state; must not pre-empt background tasks.
|
||||
"/api/email/unread-state",
|
||||
}
|
||||
|
||||
_PASSIVE_PREFIXES = (
|
||||
@@ -74,6 +77,19 @@ _PASSIVE_PREFIXES = (
|
||||
)
|
||||
|
||||
|
||||
async def maybe_stop_background_tasks_for_heartbeat(stop_background) -> bool:
|
||||
"""Stop background work for browser activity only when the gate is enabled.
|
||||
|
||||
``stop_background`` is injected by the application boundary so this policy
|
||||
remains independently testable without importing the full FastAPI app.
|
||||
"""
|
||||
if not _enabled():
|
||||
return False
|
||||
|
||||
await stop_background(reason="browser heartbeat")
|
||||
return True
|
||||
|
||||
|
||||
def should_track_interactive_request(path: str, method: str = "GET") -> bool:
|
||||
if not _enabled():
|
||||
return False
|
||||
|
||||
+960
-137
File diff suppressed because it is too large
Load Diff
@@ -12,6 +12,7 @@ class ChatRequest(BaseModel):
|
||||
use_research: Optional[bool] = Field(default=False, description="Enable deep research")
|
||||
time_filter: Optional[str] = Field(default=None, description="Time filter for search")
|
||||
preset_id: Optional[str] = Field(default=None, description="Preset identifier")
|
||||
selected_endpoint_id: Optional[str] = Field(default=None, description="Selected model endpoint ID")
|
||||
|
||||
@field_validator('message')
|
||||
@classmethod
|
||||
|
||||
+25
-8
@@ -14,6 +14,13 @@ from src.constants import SETTINGS_FILE, FEATURES_FILE
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Keys retained in the raw settings store for compatibility and rollback, but
|
||||
# deliberately unavailable through generic settings APIs or agent tools. They
|
||||
# must stay in ``DEFAULT_SETTINGS`` so old files continue to load without data
|
||||
# loss; callers that present or mutate settings should use this set as a
|
||||
# tombstone boundary.
|
||||
RETIRED_SETTING_KEYS = frozenset({"default_model_fallbacks"})
|
||||
|
||||
# Tiny TTL cache for settings/features. get_setting() is called on hot paths
|
||||
# (every chat, every preprocess); without this it re-parses the JSON each call.
|
||||
# Picks up edits within _CACHE_TTL seconds, which is fine for human-edited config.
|
||||
@@ -138,14 +145,13 @@ DEFAULT_SETTINGS = {
|
||||
# Email replies use email_writing_style instead because greetings,
|
||||
# signatures, and mailbox identity rules are medium-specific.
|
||||
"document_writing_style": "",
|
||||
# Ordered fallback chain for the default chat model. Each entry is
|
||||
# {"endpoint_id": "...", "model": "..."}. If the primary model fails
|
||||
# before producing output (endpoint offline / errors), the chat
|
||||
# dispatch retries the next entry in order.
|
||||
# Legacy ordered fallback chain for the default chat model. Values remain
|
||||
# stored for compatibility and rollback reference, but model routing no
|
||||
# longer reads this key.
|
||||
"default_model_fallbacks": [],
|
||||
# When True, non-admin users inherit global default model/endpoint/fallbacks
|
||||
# when they have no personal defaults. When False, users only use their
|
||||
# personal defaults (no global fallback). Default is False.
|
||||
# When True, non-admin users inherit the global default model/endpoint when
|
||||
# they have no personal defaults. When False, users only use their personal
|
||||
# defaults. Default is False.
|
||||
"share_defaults_with_users": False,
|
||||
"utility_endpoint_id": "",
|
||||
"utility_model": "",
|
||||
@@ -198,6 +204,17 @@ DEFAULT_SETTINGS = {
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def without_retired_settings(settings: dict) -> dict:
|
||||
"""Return a shallow copy suitable for generic settings interfaces."""
|
||||
if not isinstance(settings, dict):
|
||||
return {}
|
||||
return {
|
||||
key: value
|
||||
for key, value in settings.items()
|
||||
if key not in RETIRED_SETTING_KEYS
|
||||
}
|
||||
|
||||
DEFAULT_FEATURES = {
|
||||
"web_search": True,
|
||||
"web_fetch": True,
|
||||
@@ -270,7 +287,7 @@ _PER_USER_KEYS = {
|
||||
# Default chat endpoint / model — without per-user resolution every new
|
||||
# account inherited whatever the most-recent admin picked, which then
|
||||
# got injected into the chat composer on first open.
|
||||
"default_endpoint_id", "default_model", "default_model_fallbacks",
|
||||
"default_endpoint_id", "default_model",
|
||||
"utility_endpoint_id", "utility_model", "utility_model_fallbacks",
|
||||
"research_endpoint_id", "research_model",
|
||||
}
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""Shared resolver for background-task AI endpoints."""
|
||||
|
||||
from src.endpoint_resolver import (
|
||||
resolve_chat_fallback_candidates,
|
||||
resolve_endpoint,
|
||||
resolve_utility_fallback_candidates,
|
||||
)
|
||||
@@ -32,7 +31,6 @@ def resolve_task_candidates(
|
||||
2. Utility endpoint/model
|
||||
3. Default endpoint/model
|
||||
4. Utility fallback chain
|
||||
5. Default fallback chain
|
||||
"""
|
||||
candidates = []
|
||||
|
||||
@@ -49,9 +47,6 @@ def resolve_task_candidates(
|
||||
_append(*resolve_endpoint("default", owner=owner))
|
||||
for url, model, headers in resolve_utility_fallback_candidates(owner=owner):
|
||||
_append(url, model, headers)
|
||||
for url, model, headers in resolve_chat_fallback_candidates(owner=owner):
|
||||
_append(url, model, headers)
|
||||
|
||||
return candidates
|
||||
|
||||
|
||||
|
||||
@@ -233,7 +233,8 @@ async def _call_teacher(teacher_model_spec: str, prompt: str,
|
||||
owner: Optional[str] = None) -> Optional[str]:
|
||||
"""Call the configured teacher endpoint with the escalation prompt."""
|
||||
from src.llm_core import llm_call_async
|
||||
from src.ai_interaction import _resolve_model, _TEACHER_SYSTEM_PROMPT
|
||||
from src.ai_interaction import _resolve_model
|
||||
from src.agent_tools.model_interaction_tools import _TEACHER_SYSTEM_PROMPT
|
||||
try:
|
||||
url, model, headers = await asyncio.to_thread(_resolve_model, teacher_model_spec, owner=owner)
|
||||
except Exception as e:
|
||||
|
||||
+69
-4
@@ -187,9 +187,13 @@ _FUNCTION_MODEL_NAME_RE = re.compile(
|
||||
_FUNCTION_MODEL_PARAMS_OPEN_RE = re.compile(r"<parameters>\s*", re.IGNORECASE)
|
||||
_FUNCTION_MODEL_PARAMS_CLOSE_RE = re.compile(r"</parameters>", re.IGNORECASE)
|
||||
_QWEN_ROLE_MARKER_RE = re.compile(r"</?\|(?:assistant|assistan|user|system|tool)\|>?|</\|end\|>?", re.IGNORECASE)
|
||||
# At least one pipe is required around `end`. Both pipes used to be optional
|
||||
# (`\|?end\|?`), which also matched a bare `end` on its own line and deleted it
|
||||
# from ordinary prose and from Ruby/Lua/shell snippets that close blocks with
|
||||
# one; see #5547. `|end`, `end|`, `|end|` and `/|end|` still strip as before.
|
||||
_QWEN_BARE_MARKER_RE = re.compile(
|
||||
r"(?:^|[\t\r\n ])(?:\|?end\|?|/?\|end\|)(?=[\t\r\n ]|$)|"
|
||||
r"(?:^|[\t\r\n ])assistan(?:t)?(?=[\t\r\n ]|$)",
|
||||
r"(?:^|[\t\r\n ])(?:/?\|end\||\|end|end\|)(?=[\t\r\n ]|$)|"
|
||||
r"(?:^|[\r\n])[ \t]*assistan(?:t)?[ \t]*(?=[\r\n]|$)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
@@ -925,6 +929,46 @@ def _parse_xml_direct_tool(name, body) -> Optional[ToolBlock]:
|
||||
return function_call_to_tool_block(mapped, json.dumps(params))
|
||||
|
||||
|
||||
def _looks_like_json_body(body: str) -> bool:
|
||||
"""True when a <tool_call> wrapper body is JSON, not XML markup."""
|
||||
return body.lstrip()[:1] in ("{", "[")
|
||||
|
||||
|
||||
def _parse_json_tool_call_body(body: str) -> Optional[ToolBlock]:
|
||||
"""Parse a Qwen/Hermes text-mode wrapper body: bare JSON inside <tool_call>.
|
||||
|
||||
<tool_call>
|
||||
{"name": "bash", "arguments": {"command": "mkdir -p agent-test"}}
|
||||
</tool_call>
|
||||
|
||||
Strict by design (issue #5187 / tracker #5333): the body must decode to an
|
||||
object with a string "name", and "arguments" — when present — must itself
|
||||
be an object. Anything else returns None rather than being coerced, so a
|
||||
malformed call is dropped instead of dispatching with mangled arguments.
|
||||
raw_decode tolerates trailing chatter after the JSON object; the trailing
|
||||
text is never scanned for tool markup. Conversion goes through
|
||||
function_call_to_tool_block so aliases and per-tool argument formatting
|
||||
stay identical to the XML invoke path.
|
||||
"""
|
||||
stripped = body.strip()
|
||||
if not stripped.startswith("{"):
|
||||
return None
|
||||
try:
|
||||
parsed, _end = json.JSONDecoder().raw_decode(stripped)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
if not isinstance(parsed, dict):
|
||||
return None
|
||||
name = parsed.get("name")
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
return None
|
||||
if "arguments" in parsed and not isinstance(parsed["arguments"], dict):
|
||||
return None
|
||||
args = parsed.get("arguments", {})
|
||||
from src.tool_schemas import function_call_to_tool_block
|
||||
return function_call_to_tool_block(name.strip().lower(), json.dumps(args))
|
||||
|
||||
|
||||
def _iter_stepfun_tool_calls(text: str):
|
||||
"""Yield StepFun native tool-call token bodies without regex backtracking."""
|
||||
pos = 0
|
||||
@@ -1326,10 +1370,21 @@ def parse_tool_blocks(text: str, skip_fenced: bool = False) -> List[ToolBlock]:
|
||||
if blocks:
|
||||
return blocks
|
||||
# Try wrapped: <tool_call><invoke ...>...</invoke></tool_call>
|
||||
# A wrapper body that is JSON (Qwen/Hermes text mode, issue #5187) is
|
||||
# parsed as JSON or dropped — never scanned by the XML iterators, so
|
||||
# XML-like text inside JSON argument values stays data instead of
|
||||
# selecting a different tool.
|
||||
json_body_seen = False
|
||||
for _ms, inner_start, inner_end, _me in _iter_delimited(
|
||||
text, _XML_TOOL_CALL_OPEN_RE, _XML_TOOL_CALL_CLOSE_RE
|
||||
):
|
||||
body = text[inner_start:inner_end]
|
||||
if _looks_like_json_body(body):
|
||||
json_body_seen = True
|
||||
block = _parse_json_tool_call_body(body)
|
||||
if block:
|
||||
blocks.append(block)
|
||||
continue
|
||||
for inv_name, inv_body in _iter_xml_invoke(body):
|
||||
block = _parse_xml_invoke(inv_name, inv_body)
|
||||
if block:
|
||||
@@ -1344,6 +1399,13 @@ def parse_tool_blocks(text: str, skip_fenced: bool = False) -> List[ToolBlock]:
|
||||
if not blocks:
|
||||
for m in _XML_OPEN_TOOL_CALL_RE.finditer(text):
|
||||
body = m.group(1)
|
||||
if _looks_like_json_body(body):
|
||||
# Same fail-closed rule as above for an unclosed wrapper.
|
||||
json_body_seen = True
|
||||
block = _parse_json_tool_call_body(body)
|
||||
if block:
|
||||
blocks.append(block)
|
||||
break
|
||||
for inv_name, inv_body in _iter_xml_invoke(body):
|
||||
block = _parse_xml_invoke(inv_name, inv_body)
|
||||
if block:
|
||||
@@ -1354,8 +1416,11 @@ def parse_tool_blocks(text: str, skip_fenced: bool = False) -> List[ToolBlock]:
|
||||
block = _parse_xml_direct_tool(d_name, d_body)
|
||||
if block:
|
||||
blocks.append(block)
|
||||
# Try bare <invoke> without wrapper
|
||||
if not blocks:
|
||||
# Try bare <invoke> without wrapper. Skipped when a JSON wrapper body
|
||||
# was seen but produced no block: this rescan covers the full text,
|
||||
# wrapper bodies included, and <invoke> markup inside a (possibly
|
||||
# malformed) JSON payload must stay data rather than dispatch.
|
||||
if not blocks and not json_body_seen:
|
||||
for inv_name, inv_body in _iter_xml_invoke(text):
|
||||
block = _parse_xml_invoke(inv_name, inv_body)
|
||||
if block:
|
||||
|
||||
@@ -196,6 +196,9 @@ async def do_manage_calendar(content: str, owner: Optional[str] = None) -> Dict:
|
||||
try:
|
||||
if action == "list_calendars":
|
||||
_ensure_default_calendar(db, owner)
|
||||
# This read path intentionally persists the lazily-created default;
|
||||
# event creation commits it in the event's transaction instead.
|
||||
db.commit()
|
||||
cals = _calendar_query().all()
|
||||
result = [{"name": c.name, "href": c.id} for c in cals]
|
||||
if result:
|
||||
|
||||
+105
-38
@@ -35,6 +35,16 @@ import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
UploadIndexFileSignature = tuple[
|
||||
str,
|
||||
Optional[int],
|
||||
Optional[int],
|
||||
Optional[int],
|
||||
Optional[int],
|
||||
Optional[int],
|
||||
]
|
||||
UploadIndexSignature = tuple[UploadIndexFileSignature, ...]
|
||||
|
||||
|
||||
class UploadCleanupSafetyError(RuntimeError):
|
||||
"""Raised when cleanup cannot prove that destructive work is safe."""
|
||||
@@ -242,7 +252,7 @@ class UploadHandler:
|
||||
|
||||
# In-memory index cache to avoid O(N) disk I/O on every request
|
||||
self._index_cache: Optional[Dict[str, Any]] = None
|
||||
self._index_mtime: float = 0.0
|
||||
self._index_signature: Optional[UploadIndexSignature] = None
|
||||
|
||||
def inside_base_dir(self, path: str) -> bool:
|
||||
"""Check if path is inside base directory"""
|
||||
@@ -727,62 +737,119 @@ class UploadHandler:
|
||||
# Update cache if this is the main index
|
||||
if path.endswith("uploads.json"):
|
||||
self._index_cache = data
|
||||
self._index_signature = self._upload_index_signature(
|
||||
(path, path + ".bak")
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _upload_index_signature(
|
||||
paths: tuple[str, ...],
|
||||
) -> Optional[UploadIndexSignature]:
|
||||
"""Return file identities strong enough to validate the index cache.
|
||||
|
||||
Modification time alone is insufficient: a torn write can change a
|
||||
file without receiving a strictly newer timestamp on some filesystems.
|
||||
Size, inode, and nanosecond change times make those mutations visible
|
||||
while preserving the cache fast path for unchanged files.
|
||||
"""
|
||||
signature: list[UploadIndexFileSignature] = []
|
||||
for candidate in paths:
|
||||
try:
|
||||
self._index_mtime = os.path.getmtime(path)
|
||||
stat_result = os.stat(candidate)
|
||||
except FileNotFoundError:
|
||||
signature.append((candidate, None, None, None, None, None))
|
||||
continue
|
||||
except OSError:
|
||||
self._index_mtime = time.time()
|
||||
return None
|
||||
signature.append(
|
||||
(
|
||||
candidate,
|
||||
stat_result.st_dev,
|
||||
stat_result.st_ino,
|
||||
stat_result.st_size,
|
||||
stat_result.st_mtime_ns,
|
||||
stat_result.st_ctime_ns,
|
||||
)
|
||||
)
|
||||
return tuple(signature)
|
||||
|
||||
def _load_upload_index(self, *, fail_on_error: bool = False) -> Dict[str, Any]:
|
||||
"""Load the upload index from disk/cache. Uses mtime-based validation
|
||||
to avoid redundant parsing on hot paths. When ``fail_on_error`` is
|
||||
true, a missing, malformed, or unreadable live index raises so
|
||||
destructive callers cannot mistake corruption for an empty store.
|
||||
"""Load the upload index from disk/cache. Uses file-identity validation
|
||||
to avoid redundant parsing on hot paths without missing same-timestamp
|
||||
mutations. When ``fail_on_error`` is true, a missing, malformed, or
|
||||
unreadable live index raises so destructive callers cannot mistake
|
||||
corruption for an empty store.
|
||||
"""
|
||||
uploads_db_path = os.path.join(self.upload_dir, "uploads.json")
|
||||
candidates = (uploads_db_path, uploads_db_path + ".bak")
|
||||
if fail_on_error:
|
||||
# A backup is intentionally the previous snapshot. It is useful for
|
||||
# non-destructive reads, but cannot authorize deletion when the live
|
||||
# index is missing or corrupt.
|
||||
if not os.path.exists(uploads_db_path):
|
||||
raise ValueError("live uploads database is missing")
|
||||
existing_candidates = [uploads_db_path]
|
||||
else:
|
||||
existing_candidates = [path for path in candidates if os.path.exists(path)]
|
||||
if not existing_candidates:
|
||||
self._index_cache = {}
|
||||
self._index_mtime = 0.0
|
||||
return {}
|
||||
for _attempt in range(3):
|
||||
signature = self._upload_index_signature(candidates)
|
||||
if fail_on_error:
|
||||
# A backup is intentionally the previous snapshot. It is useful for
|
||||
# non-destructive reads, but cannot authorize deletion when the live
|
||||
# index is missing or corrupt.
|
||||
if not os.path.exists(uploads_db_path):
|
||||
raise ValueError("live uploads database is missing")
|
||||
existing_candidates = [uploads_db_path]
|
||||
else:
|
||||
existing_candidates = [
|
||||
path for path in candidates if os.path.exists(path)
|
||||
]
|
||||
if not existing_candidates:
|
||||
self._index_cache = {}
|
||||
self._index_signature = signature
|
||||
return {}
|
||||
|
||||
# Check cache validity
|
||||
try:
|
||||
mtime = max(os.path.getmtime(path) for path in existing_candidates)
|
||||
# Check cache validity
|
||||
if (
|
||||
not fail_on_error
|
||||
and signature is not None
|
||||
and self._index_cache is not None
|
||||
and mtime <= self._index_mtime
|
||||
and signature == self._index_signature
|
||||
):
|
||||
return self._index_cache
|
||||
except OSError:
|
||||
mtime = 0.0
|
||||
|
||||
# Try the live file first, fall back to the .bak sibling if the
|
||||
# live file is truncated/corrupted.
|
||||
for candidate in existing_candidates:
|
||||
try:
|
||||
with open(candidate, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
if isinstance(data, dict):
|
||||
self._index_cache = data
|
||||
self._index_mtime = mtime
|
||||
return data
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to read uploads database ({candidate}): {e}")
|
||||
# Try the live file first, fall back to the .bak sibling if the
|
||||
# live file is truncated/corrupted. A candidate parsed from an old
|
||||
# inode is accepted only when the whole index signature stays
|
||||
# stable through the read; otherwise retry so the cache cannot pair
|
||||
# stale data with a fresh replacement signature.
|
||||
index_changed_during_read = False
|
||||
for candidate in existing_candidates:
|
||||
try:
|
||||
with open(candidate, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
verified_signature = self._upload_index_signature(candidates)
|
||||
if (
|
||||
signature is not None
|
||||
and verified_signature is not None
|
||||
and verified_signature != signature
|
||||
):
|
||||
index_changed_during_read = True
|
||||
break
|
||||
if isinstance(data, dict):
|
||||
self._index_cache = data
|
||||
self._index_signature = verified_signature
|
||||
return data
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to read uploads database ({candidate}): {e}")
|
||||
verified_signature = self._upload_index_signature(candidates)
|
||||
if (
|
||||
signature is not None
|
||||
and verified_signature is not None
|
||||
and verified_signature != signature
|
||||
):
|
||||
index_changed_during_read = True
|
||||
break
|
||||
continue
|
||||
if index_changed_during_read:
|
||||
continue
|
||||
break
|
||||
|
||||
if fail_on_error:
|
||||
raise ValueError("live uploads database is unreadable")
|
||||
self._index_cache = {}
|
||||
self._index_signature = self._upload_index_signature(candidates)
|
||||
return {}
|
||||
|
||||
def get_upload_info(self, upload_id: str) -> Optional[Dict[str, Any]]:
|
||||
|
||||
+50
-147
@@ -10,18 +10,25 @@ import modelsModule from './js/models.js?v=20260715startupcalm2';
|
||||
import ragModule from './js/rag.js';
|
||||
import presetsModule from './js/presets.js';
|
||||
import searchModule from './js/search.js';
|
||||
import chatModule from './js/chat.js?v=20260722ctxheader4';
|
||||
import chatModule from './js/chat.js?v=20260801fix1';
|
||||
import compareModule from './js/compare/index.js?v=20260723compareicon2';
|
||||
import documentModule from './js/document.js?v=20260722emailfastindex1';
|
||||
import searchChatModule from './js/search-chat.js';
|
||||
import { makeWindowDraggable } from './js/windowDrag.js';
|
||||
import {
|
||||
revealApplicationShellAfterPaint,
|
||||
runDeferredRouteOpener,
|
||||
deferRouteOpener,
|
||||
settleSessionHydration
|
||||
} from './js/startupShell.js';
|
||||
import markdownModule from './js/markdown.js';
|
||||
import chatRenderer from './js/chatRenderer.js?v=20260722emailfastindex1';
|
||||
import sessionModule from './js/sessions.js?v=20260722ctxheader4';
|
||||
import sessionModule from './js/sessions.js';
|
||||
import memoryModule from './js/memory.js?v=20260722memoryloading1';
|
||||
import voiceRecorderModule from './js/voiceRecorder.js';
|
||||
import censorModule from './js/censor.js';
|
||||
import galleryModule from './js/gallery.js';
|
||||
import { UI_VIS_DEFAULT_OFF, resolveVisibility } from './js/ui_visibility.js';
|
||||
import tasksModule from './js/tasks.js?v=20260723tasksbulkfeedback1';
|
||||
import calendarModule from './js/calendar.js';
|
||||
import notesModule from './js/notes.js';
|
||||
@@ -1217,12 +1224,13 @@ function initializeEventListeners() {
|
||||
'/library': () => sessionModule && sessionModule.openLibrary && sessionModule.openLibrary(),
|
||||
};
|
||||
const _opener = _routeOpen[urlPath];
|
||||
// Defer the opener — at this point in init, the modules whose handlers
|
||||
// we trigger (#rail-new-session click handler, the email-section header
|
||||
// click handler in emailInbox, sessionModule's loaded session list) are
|
||||
// still being wired up further down in this same function. Stash the
|
||||
// opener so it runs from sessionModule.loadSessions().finally() below.
|
||||
if (_opener) window._odysseusRouteOpener = _opener;
|
||||
// Defer the opener — at this point in init, the modules whose handlers we
|
||||
// trigger (#rail-new-session click handler, the email-section header click
|
||||
// handler in emailInbox, sessionModule) are still being wired up further
|
||||
// down in this same function. startupShell decides when it can run: as soon
|
||||
// as wiring completes, or — for the routes that read the session list —
|
||||
// once /api/sessions has settled.
|
||||
deferRouteOpener(urlPath, _opener);
|
||||
|
||||
// Archive browser tool button
|
||||
const toolLibraryBtn = el('tool-library-btn');
|
||||
@@ -1689,12 +1697,20 @@ function initializeEventListeners() {
|
||||
|
||||
const newMemoryInput = el('new-memory-input');
|
||||
if (newMemoryInput) {
|
||||
newMemoryInput.addEventListener('keypress', (e) => {
|
||||
if (e.key === 'Enter') {
|
||||
// keydown, not the deprecated keypress: keypress is not guaranteed to
|
||||
// fire for Enter everywhere, which left the Add Memory form with no
|
||||
// working submit path (#5828).
|
||||
newMemoryInput.addEventListener('keydown', (e) => {
|
||||
if (e.key === 'Enter' && !e.isComposing) {
|
||||
e.preventDefault();
|
||||
memoryModule.addNewMemory();
|
||||
}
|
||||
});
|
||||
}
|
||||
const newMemoryAddBtn = el('new-memory-add-btn');
|
||||
if (newMemoryAddBtn) {
|
||||
newMemoryAddBtn.addEventListener('click', () => memoryModule.addNewMemory());
|
||||
}
|
||||
|
||||
// Voice recording is handled by the dual-purpose send/mic button (see below)
|
||||
|
||||
@@ -2710,46 +2726,6 @@ function initializeEventListeners() {
|
||||
// ── UI Visibility (Customize UI modal) ──
|
||||
const UI_VIS_KEY = 'odysseus-ui-visibility';
|
||||
|
||||
// Selector map: key → CSS selector(s) for targets
|
||||
const UI_VIS_MAP = {
|
||||
'sidebar-brand': '.sidebar-brand-title',
|
||||
'sidebar-new-chat': '#sidebar-new-chat-btn',
|
||||
'sidebar-search': '#sidebar-search-btn',
|
||||
'sessions-section': '#sessions-section',
|
||||
'email-section': '#email-section',
|
||||
'tools-section': '#tools-section',
|
||||
// Per-tool visibility — fine-grained control over which entries show
|
||||
// inside the Tools section in the sidebar.
|
||||
'tool-calendar': '#tool-calendar-btn',
|
||||
'tool-compare': '#tool-compare-btn',
|
||||
'tool-cookbook': '#tool-cookbook-btn',
|
||||
'tool-research': '#tool-research-btn',
|
||||
'tool-gallery': '#tool-gallery-btn',
|
||||
'tool-library': '#tool-library-btn',
|
||||
'tool-memory': '#tool-memory-btn',
|
||||
'tool-notes': '#tool-notes-btn',
|
||||
'tool-tasks': '#tool-tasks-btn',
|
||||
'tool-theme': '#tool-theme-btn',
|
||||
'user-bar': '#user-bar-profile',
|
||||
'sidebar-settings-btn':'#user-bar-settings',
|
||||
'chat-meta': '.chat-meta-overlay',
|
||||
'welcome-text': '.welcome-name, .welcome-sub, #welcome-tip',
|
||||
'incognito-btn': '.incognito-btn',
|
||||
'web-toggle-btn': '#web-toggle-btn',
|
||||
'doc-toggle-btn': '#overflow-doc-btn',
|
||||
'rag-toggle-btn': '#overflow-rag-btn',
|
||||
'bash-toggle-btn': '#bash-toggle-btn',
|
||||
'overflow-plus-btn': '.overflow-wrapper',
|
||||
'mode-toggle': '.mode-toggle',
|
||||
'preset-mini-btn': '#overflow-preset-btn',
|
||||
'attach-btn': '#overflow-attach-btn',
|
||||
'research-btn': '#overflow-research-btn',
|
||||
'rail-new-chat': '#rail-new-session',
|
||||
};
|
||||
|
||||
// Keys hidden by default on first run (no localStorage yet)
|
||||
const UI_VIS_DEFAULT_OFF = new Set(['rag-toggle-btn', 'text-emojis', 'chat-fullwidth']);
|
||||
|
||||
// Keys that need admin to toggle off (reserved for future use)
|
||||
const UI_VIS_ADMIN_ONLY = new Set([]);
|
||||
|
||||
@@ -2762,14 +2738,14 @@ function initializeEventListeners() {
|
||||
}
|
||||
|
||||
function applyUIVis(state) {
|
||||
Object.entries(UI_VIS_MAP).forEach(([key, selector]) => {
|
||||
// section-drag-reorder uses a body class instead of inline styles
|
||||
if (key === 'section-drag-reorder') return;
|
||||
const visible = key in state ? state[key] !== false : !UI_VIS_DEFAULT_OFF.has(key);
|
||||
// resolveVisibility computes selector→visible (pure; ui_visibility.js),
|
||||
// including the tools-section parent rule that hides every tool rail
|
||||
// launcher when Tools is off. Apply the result to the DOM here.
|
||||
for (const [selector, visible] of Object.entries(resolveVisibility(state))) {
|
||||
document.querySelectorAll(selector).forEach(el => {
|
||||
el.style.display = visible ? '' : 'none';
|
||||
});
|
||||
});
|
||||
}
|
||||
// Drag reorder: use body class so dynamically created handles are covered
|
||||
const dragEnabled = state['section-drag-reorder'] === true;
|
||||
document.body.classList.toggle('rearrange-mode', dragEnabled);
|
||||
@@ -3908,85 +3884,10 @@ function startOdysseusApp() {
|
||||
const messageInput = el('message');
|
||||
const modelPickerWrap = document.getElementById('model-picker-wrap');
|
||||
|
||||
function _readComposerPromptHistory() {
|
||||
const chatBox = document.getElementById('chat-history');
|
||||
if (!chatBox) return [];
|
||||
return Array.from(chatBox.querySelectorAll('.msg-user'))
|
||||
.reverse()
|
||||
.map(msg => {
|
||||
const body = msg.querySelector('.body');
|
||||
return msg.dataset?.raw || (body ? body.textContent : '') || '';
|
||||
})
|
||||
.filter(Boolean);
|
||||
}
|
||||
|
||||
if (messageInput && !messageInput._odysseusPromptRecallCapture) {
|
||||
messageInput._odysseusPromptRecallCapture = true;
|
||||
let recallHistory = [];
|
||||
let recallIndex = -1;
|
||||
let lastRecalled = '';
|
||||
const norm = (v) => String(v || '').replace(/\r\n/g, '\n').trimEnd();
|
||||
messageInput.addEventListener('input', () => {
|
||||
if (norm(messageInput.value) === norm(lastRecalled)) return;
|
||||
recallHistory = [];
|
||||
recallIndex = -1;
|
||||
lastRecalled = '';
|
||||
try { delete messageInput.dataset.odysseusRecallIndex; } catch {}
|
||||
}, true);
|
||||
messageInput.addEventListener('keydown', (e) => {
|
||||
if (e.key !== 'ArrowUp' && e.key !== 'ArrowDown') return;
|
||||
if (e.shiftKey || e.altKey || e.ctrlKey || e.metaKey || e.isComposing) return;
|
||||
if (window._ghostAutocomplete?.isActive?.()) return;
|
||||
const fresh = _readComposerPromptHistory();
|
||||
const history = fresh.length ? fresh : recallHistory;
|
||||
if (!history.length) return;
|
||||
const current = norm(messageInput.value);
|
||||
let currentIndex = current ? history.findIndex(item => norm(item) === current) : -1;
|
||||
if (current && currentIndex < 0 && current === norm(lastRecalled)) currentIndex = recallIndex;
|
||||
if (current && currentIndex < 0) {
|
||||
const markedIndex = Number(messageInput.dataset.odysseusRecallIndex);
|
||||
if (Number.isInteger(markedIndex) && markedIndex >= 0 && markedIndex < history.length) {
|
||||
currentIndex = markedIndex;
|
||||
}
|
||||
}
|
||||
e.preventDefault();
|
||||
e.stopPropagation();
|
||||
e.stopImmediatePropagation();
|
||||
if (e.key === 'ArrowDown') {
|
||||
if (currentIndex < 0) return;
|
||||
const nextIndex = currentIndex - 1;
|
||||
if (nextIndex < 0) {
|
||||
recallHistory = history;
|
||||
recallIndex = -1;
|
||||
lastRecalled = '';
|
||||
try { delete messageInput.dataset.odysseusRecallIndex; } catch {}
|
||||
messageInput.value = '';
|
||||
try { messageInput.selectionStart = messageInput.selectionEnd = 0; } catch {}
|
||||
try { uiModule.autoResize(messageInput); } catch {}
|
||||
return;
|
||||
}
|
||||
const recalled = history[nextIndex];
|
||||
recallHistory = history;
|
||||
recallIndex = nextIndex;
|
||||
lastRecalled = recalled;
|
||||
try { messageInput.dataset.odysseusRecallIndex = String(nextIndex); } catch {}
|
||||
messageInput.value = recalled;
|
||||
try { messageInput.selectionStart = messageInput.selectionEnd = recalled.length; } catch {}
|
||||
try { uiModule.autoResize(messageInput); } catch {}
|
||||
return;
|
||||
}
|
||||
const nextIndex = currentIndex >= 0 ? Math.min(currentIndex + 1, history.length - 1) : 0;
|
||||
const recalled = history[nextIndex];
|
||||
if (!recalled) return;
|
||||
recallHistory = history;
|
||||
recallIndex = nextIndex;
|
||||
lastRecalled = recalled;
|
||||
try { messageInput.dataset.odysseusRecallIndex = String(nextIndex); } catch {}
|
||||
messageInput.value = recalled;
|
||||
try { messageInput.selectionStart = messageInput.selectionEnd = recalled.length; } catch {}
|
||||
try { uiModule.autoResize(messageInput); } catch {}
|
||||
}, true);
|
||||
}
|
||||
// ArrowUp/ArrowDown prompt recall on #message lives in
|
||||
// static/js/composerArrowUpRecall.js (wired from chat.js). Do not re-add a
|
||||
// copy here: two capture-phase listeners on the same textarea meant the one
|
||||
// without the draft guard won and ate unsent multi-line prompts (#5862).
|
||||
|
||||
const _sendIcon = '<svg width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2.5" stroke-linecap="round" stroke-linejoin="round"><path d="M12 19V5M5 12l7-7 7 7"/></svg>';
|
||||
const _micIcon = '<svg width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><path d="M12 1a3 3 0 0 0-3 3v8a3 3 0 0 0 6 0V4a3 3 0 0 0-3-3z"/><path d="M19 10v2a7 7 0 0 1-14 0v-2"/><line x1="12" y1="19" x2="12" y2="23"/><line x1="8" y1="23" x2="16" y2="23"/></svg>';
|
||||
@@ -4382,6 +4283,10 @@ function startOdysseusApp() {
|
||||
// Load initial data
|
||||
presetsModule.loadPresets(uiModule.showError);
|
||||
|
||||
// Core wiring is complete for this turn — reveal the shell independently of
|
||||
// the session-list request.
|
||||
revealApplicationShellAfterPaint();
|
||||
|
||||
if (sessionModule) {
|
||||
sessionModule.initDependencies({
|
||||
API_BASE: API_BASE,
|
||||
@@ -4393,21 +4298,19 @@ function startOdysseusApp() {
|
||||
scrollHistory: uiModule.scrollHistoryInstant
|
||||
});
|
||||
|
||||
// Load sessions first (critical path) — remove loader when done
|
||||
sessionModule.loadSessions()
|
||||
.catch(e => console.warn('loadSessions error:', e))
|
||||
.finally(() => {
|
||||
const loader = document.getElementById('app-loader');
|
||||
if (loader) { loader.style.opacity = '0'; setTimeout(() => loader.remove(), 300); }
|
||||
// Fire any URL route opener now that sessions + module wiring are
|
||||
// ready. Deferred from up top of init for exactly this reason.
|
||||
if (window._odysseusRouteOpener) {
|
||||
try { window._odysseusRouteOpener(); } catch (_) {}
|
||||
window._odysseusRouteOpener = null;
|
||||
}
|
||||
});
|
||||
// sessionModule is now wired, so every route opener has the modules it
|
||||
// drives. The ones that read no session data open here rather than
|
||||
// queueing behind /api/sessions.
|
||||
runDeferredRouteOpener();
|
||||
|
||||
// The shell is already usable at this point; session hydration is
|
||||
// sidebar-local and settles on its own schedule.
|
||||
settleSessionHydration(() => sessionModule.loadSessions());
|
||||
} else {
|
||||
console.error('Session module not loaded!');
|
||||
// Nothing will hydrate. Settle immediately so the sidebar exposes the
|
||||
// failure; session-dependent routes must remain unopened without data.
|
||||
settleSessionHydration(null);
|
||||
}
|
||||
|
||||
const runNonCriticalStartup = (fn, delay = 4000) => {
|
||||
|
||||
+32
-17
@@ -248,11 +248,20 @@
|
||||
}, { once: true });
|
||||
})();
|
||||
</script>
|
||||
<link rel="stylesheet" href="/static/style.css?v=20260723tasksbulkfeedback1">
|
||||
<link rel="modulepreload" href="/static/app.js?v=20260723tasksbulkfeedback1">
|
||||
<link rel="modulepreload" href="/static/js/chat.js?v=20260722ctxheader4">
|
||||
<!-- Preload the two faces first paint actually uses: Fira Code 400 and 600,
|
||||
the app font and the weight the sidebar and header text render at. They
|
||||
are declared in style.css, so without a hint they are only discovered
|
||||
after the stylesheet parses and then queue behind the module graph.
|
||||
crossorigin is required even though these are same-origin: fonts are
|
||||
always fetched in CORS mode, and a preload whose mode does not match the
|
||||
real request is discarded and the font fetched a second time. -->
|
||||
<link rel="preload" as="font" type="font/woff2" crossorigin href="/static/fonts/FiraCode-Regular.woff2">
|
||||
<link rel="preload" as="font" type="font/woff2" crossorigin href="/static/fonts/FiraCode-SemiBold.woff2">
|
||||
<link rel="stylesheet" href="/static/style.css?v=20260808startupshell1">
|
||||
<link rel="modulepreload" href="/static/app.js?v=20260808startupshell1">
|
||||
<link rel="modulepreload" href="/static/js/chat.js?v=20260801fix1">
|
||||
<link rel="modulepreload" href="/static/js/ui.js">
|
||||
<link rel="modulepreload" href="/static/js/sessions.js?v=20260722ctxheader4">
|
||||
<link rel="modulepreload" href="/static/js/sessions.js">
|
||||
<link rel="modulepreload" href="/static/js/markdown.js">
|
||||
</head>
|
||||
<body>
|
||||
@@ -286,7 +295,13 @@
|
||||
if(!document.getElementById('app-loader')){clearInterval(iv);return}
|
||||
render();
|
||||
},150);
|
||||
setTimeout(function(){var l=document.getElementById('app-loader');if(l){l.style.opacity='0';setTimeout(function(){l.remove()},300)}},5000);
|
||||
// startupShell.js hides the loader as soon as the shell is wired; it calls
|
||||
// back here to stop the wave because this interval is owned by this script.
|
||||
window.__odysseusLoaderWaveStop=function(){clearInterval(iv)};
|
||||
// Last-resort fallback for a boot that never reaches app.js at all. Must
|
||||
// still REMOVE the node: sessions.js reads its presence as "startup in
|
||||
// progress" and stops clearing the composer while it is around.
|
||||
setTimeout(function(){var l=document.getElementById('app-loader');if(l){clearInterval(iv);l.style.opacity='0';setTimeout(function(){l.remove()},300)}},5000);
|
||||
})();
|
||||
</script>
|
||||
<!-- Memory Management Modal -->
|
||||
@@ -365,6 +380,7 @@
|
||||
<span class="skill-rich-ph"><span class="k">Add a memory</span> — e.g. 'I prefer concise replies' <svg class="k" width="12" height="12" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" style="vertical-align:-2px;margin-left:4px;" aria-hidden="true"><polyline points="9 10 4 15 9 20"/><path d="M20 4v7a4 4 0 0 1-4 4H4"/></svg></span>
|
||||
</div>
|
||||
<select id="new-memory-category" class="memory-edit-cat-select" aria-label="Memory category"></select>
|
||||
<button type="button" id="new-memory-add-btn" class="theme-io-btn" title="Save this memory" style="flex:none;height:28px;font-size:12px;"><svg width="13" height="13" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" style="vertical-align:-2px;margin-right:4px;" aria-hidden="true"><line x1="12" y1="5" x2="12" y2="19"/><line x1="5" y1="12" x2="19" y2="12"/></svg>Add</button>
|
||||
</div>
|
||||
</div>
|
||||
<div class="admin-card">
|
||||
@@ -812,7 +828,13 @@
|
||||
<button class="session-bulk-btn" id="session-bulk-cancel" title="Cancel"><svg width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2.5" stroke-linecap="round"><line x1="18" y1="6" x2="6" y2="18"/><line x1="6" y1="6" x2="18" y2="18"/></svg></button>
|
||||
</div>
|
||||
</div>
|
||||
<div id="session-list" role="listbox"></div>
|
||||
<div id="session-list" role="listbox">
|
||||
<!-- Sidebar-local bootstrap state. renderSessionList() replaces the
|
||||
whole list on first hydration, so this row is transient. -->
|
||||
<div id="session-list-loading" class="list-item session-list-bootstrap" role="option" aria-disabled="true" aria-live="polite" aria-atomic="true">
|
||||
<span class="grow muted" data-session-list-status>Loading chats…</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<!-- Hidden dropdown for session actions -->
|
||||
<div id="session-actions-dropdown" class="dropdown hidden">
|
||||
@@ -1005,7 +1027,7 @@
|
||||
var tips = mobile ? phone : desktop;
|
||||
var el = document.getElementById('welcome-tip');
|
||||
if (el) {
|
||||
el.textContent = 'Pick a model if you want, or just type.';
|
||||
el.textContent = tips[Math.floor(Math.random() * tips.length)];
|
||||
}
|
||||
fetch('/api/version').then(function(r){return r.json()}).then(function(d){
|
||||
if (d.version) window._appVersion = d.version;
|
||||
@@ -1482,13 +1504,6 @@
|
||||
<span class="adm-model-logo" id="set-defaultModelSelect-logo" style="display:inline-flex;align-items:center;justify-content:center;width:18px;height:18px;flex-shrink:0;opacity:0.9;color:var(--fg);"></span>
|
||||
<select id="set-defaultModelSelect" class="settings-select"></select>
|
||||
</div>
|
||||
<div class="settings-row" style="align-items:flex-start;">
|
||||
<label class="settings-label" style="margin-top:6px;">Fallbacks</label>
|
||||
<div style="flex:1;display:flex;flex-direction:column;gap:6px;">
|
||||
<div id="set-defaultFallbacks" class="settings-fallbacks"></div>
|
||||
<button type="button" class="settings-fallback-add" id="set-defaultAddFallback" title="Add a model to try if the one above fails">+ Add fallback</button>
|
||||
</div>
|
||||
</div>
|
||||
<div id="set-defaultChatMsg" style="font-size:11px;color:color-mix(in srgb, var(--fg) 45%, transparent);"></div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -2504,7 +2519,7 @@
|
||||
<script type="module" src="/static/js/ui.js"></script>
|
||||
<script type="module" src="/static/js/markdown.js"></script>
|
||||
<script type="module" src="/static/js/dragSort.js"></script>
|
||||
<script type="module" src="/static/js/sessions.js?v=20260722ctxheader4"></script>
|
||||
<script type="module" src="/static/js/sessions.js"></script>
|
||||
<script type="module" src="/static/js/memory.js?v=20260722memoryloading1"></script>
|
||||
<script type="module" src="/static/js/skills.js"></script>
|
||||
<script type="module" src="/static/js/tourHints.js"></script>
|
||||
@@ -2522,7 +2537,7 @@
|
||||
<script type="module" src="/static/js/chatRenderer.js?v=20260722emailfastindex1"></script>
|
||||
<script type="module" src="/static/js/codeRunner.js"></script>
|
||||
<script type="module" src="/static/js/chatStream.js?v=20260722emailfastindex1"></script>
|
||||
<script type="module" src="/static/js/chat.js?v=20260722ctxheader4"></script>
|
||||
<script type="module" src="/static/js/chat.js?v=20260801fix1"></script>
|
||||
<script type="module" src="/static/js/cookbook.js"></script>
|
||||
<script src="/static/js/cookbookSchedule.js"></script>
|
||||
<script type="module" src="/static/js/search-chat.js"></script>
|
||||
@@ -2530,7 +2545,7 @@
|
||||
<script type="module" src="/static/js/censor.js"></script>
|
||||
<script type="module" src="/static/js/settings.js?v=20260723compareicon1"></script>
|
||||
<script type="module" src="/static/js/assistant.js"></script>
|
||||
<script type="module" src="/static/app.js?v=20260723tasksbulkfeedback1"></script> <!-- app.js must be LAST -->
|
||||
<script type="module" src="/static/app.js?v=20260808startupshell1"></script> <!-- app.js must be LAST -->
|
||||
<script type="module" src="/static/js/init.js?v=20260715freshroot3"></script>
|
||||
<script type="module" src="/static/js/a11y.js"></script>
|
||||
<script nonce="{{CSP_NONCE}}">if('serviceWorker' in navigator){navigator.serviceWorker.register('/static/sw.js').catch(()=>{});}</script>
|
||||
|
||||
@@ -61,6 +61,7 @@ The largest and most central subsystem. Chat submission → backend SSE → prog
|
||||
| **`chatRenderer.js`** | Message DOM construction: `addMessage`, role labels, model route labels, color coding, footers, metrics, code blocks, sources boxes (`web`/`research`/`RAG`), findings box, images, report links, ask-user cards, welcome screen, and transcript utilities. |
|
||||
| **`streamingRenderer.js`** | Incremental streaming renderer used by `chat.js`. Freezes finalized DOM blocks and only re-renders the growing tail to avoid flicker and O(N²) re-parsing. |
|
||||
| **`streamingSegmenter.js`** | Splits a token stream into display units (text vs code fences) for `streamingRenderer.js`. |
|
||||
| **`liveThinkingThrottle.js`** | Trailing-edge coalescer for the live thinking block in `chat.js`: one DOM commit per 100 ms carrying the latest reasoning text, with `flush`/`cancel` for terminal and session-switch paths. |
|
||||
| **`slashCommands.js`** | Slash-command registry (`/help`, `/setup`, etc.), parsing, and dispatch handlers. Exported functions are consumed by `chat.js` and `slashAutocomplete.js`. |
|
||||
| **`slashAutocomplete.js`** | Composer autocomplete popup for `/` commands. |
|
||||
| **`composerArrowUpRecall.js`** | Recall last user message with `↑` on an empty composer. |
|
||||
|
||||
+1032
-349
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,104 @@
|
||||
/** Select and update the response holder for a route-provenance event. */
|
||||
export function applyModelRouteEventState(event, holder, roundHolder, defaultModel = '') {
|
||||
const target = event && event.round && roundHolder ? roundHolder : holder;
|
||||
if (!target) return null;
|
||||
|
||||
target._requestedModel = (
|
||||
event.requested_model
|
||||
|| event.selected_model
|
||||
|| target._requestedModel
|
||||
|| defaultModel
|
||||
);
|
||||
target._actualModel = (
|
||||
event.model
|
||||
|| event.answered_by
|
||||
|| target._actualModel
|
||||
|| target._requestedModel
|
||||
);
|
||||
const hasEndpointRoute = Boolean(
|
||||
event.requested_endpoint_id
|
||||
|| event.selected_endpoint_id
|
||||
|| event.endpoint_id
|
||||
|| event.answered_by_endpoint_id
|
||||
|| event.requested_endpoint_label
|
||||
|| event.selected_endpoint_label
|
||||
|| event.endpoint_label
|
||||
|| event.answered_by_endpoint_label
|
||||
|| target._requestedEndpointLabel
|
||||
);
|
||||
if (hasEndpointRoute) {
|
||||
target._requestedEndpointId = (
|
||||
event.requested_endpoint_id
|
||||
|| event.selected_endpoint_id
|
||||
|| target._requestedEndpointId
|
||||
|| null
|
||||
);
|
||||
target._requestedEndpointLabel = (
|
||||
event.requested_endpoint_label
|
||||
|| event.selected_endpoint_label
|
||||
|| target._requestedEndpointLabel
|
||||
|| 'Selected route'
|
||||
);
|
||||
target._actualEndpointId = (
|
||||
event.endpoint_id
|
||||
|| event.answered_by_endpoint_id
|
||||
|| target._actualEndpointId
|
||||
|| target._requestedEndpointId
|
||||
|| null
|
||||
);
|
||||
target._actualEndpointLabel = (
|
||||
event.endpoint_label
|
||||
|| event.answered_by_endpoint_label
|
||||
|| target._actualEndpointLabel
|
||||
|| target._requestedEndpointLabel
|
||||
);
|
||||
}
|
||||
return target;
|
||||
}
|
||||
|
||||
/** Copy the active route into the bubble created for the next Agent round. */
|
||||
export function inheritModelRouteState(holder, roundHolder, target, defaultModel = '') {
|
||||
if (!target) return null;
|
||||
const source = roundHolder || holder;
|
||||
target._requestedModel = source?._requestedModel || defaultModel;
|
||||
target._actualModel = source?._actualModel || target._requestedModel;
|
||||
if (source?._requestedEndpointLabel || source?._actualEndpointLabel) {
|
||||
target._requestedEndpointId = source?._requestedEndpointId || null;
|
||||
target._requestedEndpointLabel = source?._requestedEndpointLabel || 'Selected route';
|
||||
target._actualEndpointId = source?._actualEndpointId || target._requestedEndpointId;
|
||||
target._actualEndpointLabel = source?._actualEndpointLabel || target._requestedEndpointLabel;
|
||||
}
|
||||
return target;
|
||||
}
|
||||
|
||||
/** Apply final/metrics provenance to the active round, not the first bubble. */
|
||||
export function applyModelMetricsState(metrics, holder, roundHolder, defaultModel = '') {
|
||||
const target = roundHolder || holder;
|
||||
if (!target || !metrics) return target || null;
|
||||
const roundModels = Array.isArray(metrics.round_models) ? metrics.round_models : [];
|
||||
const roundModel = roundHolder && roundModels.length
|
||||
? roundModels[roundModels.length - 1]
|
||||
: null;
|
||||
target._requestedModel = metrics.requested_model || target._requestedModel || defaultModel;
|
||||
target._actualModel = roundModel || metrics.model || target._actualModel || target._requestedModel;
|
||||
const roundEndpointIds = Array.isArray(metrics.round_endpoint_ids) ? metrics.round_endpoint_ids : [];
|
||||
const roundEndpointLabels = Array.isArray(metrics.round_endpoint_labels) ? metrics.round_endpoint_labels : [];
|
||||
if (
|
||||
metrics.requested_endpoint_label
|
||||
|| metrics.endpoint_label
|
||||
|| roundEndpointLabels.length
|
||||
|| target._requestedEndpointLabel
|
||||
) {
|
||||
target._requestedEndpointId = metrics.requested_endpoint_id || target._requestedEndpointId || null;
|
||||
target._requestedEndpointLabel = metrics.requested_endpoint_label || target._requestedEndpointLabel || 'Selected route';
|
||||
const hasRoundEndpointId = Boolean(roundHolder && roundEndpointIds.length);
|
||||
const hasRoundEndpointLabel = Boolean(roundHolder && roundEndpointLabels.length);
|
||||
target._actualEndpointId = hasRoundEndpointId
|
||||
? roundEndpointIds[roundEndpointIds.length - 1]
|
||||
: (metrics.endpoint_id || target._actualEndpointId || target._requestedEndpointId);
|
||||
target._actualEndpointLabel = hasRoundEndpointLabel
|
||||
? roundEndpointLabels[roundEndpointLabels.length - 1]
|
||||
: (metrics.endpoint_label || target._actualEndpointLabel || target._requestedEndpointLabel);
|
||||
}
|
||||
return target;
|
||||
}
|
||||
+258
-47
@@ -478,7 +478,10 @@ const DSML_STRAY_RE = /<\s*\/?\s*[||]+\s*DSML\s*[||]+[^>]*>/gi;
|
||||
const DSML_INVOKE_RE = /<\s*[||]+\s*DSML\s*[||]+\s*invoke\b[^>]*>[\s\S]*?(?:<\s*\/\s*[||]+\s*DSML\s*[||]+\s*invoke\s*>|$)/gi;
|
||||
const RAW_OPENAI_TOOL_JSON_RE = /(?:\[\s*)?\{\s*"function"\s*:\s*\{[\s\S]*?\}\s*,\s*"id"\s*:\s*"[^"]*"\s*,\s*"type"\s*:\s*"function"\s*\}\s*\]?/gi;
|
||||
const QWEN_ROLE_MARKER_RE = /<\/?\|(?:assistant|assistan|user|system|tool)\|>?|<\/\|end\|>?/gi;
|
||||
const QWEN_BARE_MARKER_RE = /(?:^|[\t\r\n ])(?:\|?end\|?|\/?\|end\|)(?=[\t\r\n ]|$)|(?:^|[\t\r\n ])assistan(?:t)?(?=[\t\r\n ]|$)/gi;
|
||||
// Keep in sync with _QWEN_BARE_MARKER_RE in src/tool_parsing.py. At least one
|
||||
// pipe is required around `end`: with both optional (`\|?end\|?`) this also ate
|
||||
// a bare `end` on its own line, breaking Ruby/Lua/shell snippets (#5547).
|
||||
const QWEN_BARE_MARKER_RE = /(?:^|[\t\r\n ])(?:\/?\|end\||\|end|end\|)(?=[\t\r\n ]|$)|(?:^|[\r\n])[ \t]*assistan(?:t)?[ \t]*(?=[\r\n]|$)/gi;
|
||||
// Self-narration about tool results (model echoing stdout/exit_code)
|
||||
const TOOL_NARRATION_RE = /(?:The (?:result|output) shows?:?\s*)?-?\s*(?:stdout|stderr|exit_code):\s*.+/gi;
|
||||
|
||||
@@ -612,10 +615,36 @@ export function sameModelName(left, right) {
|
||||
|| shortModel(a).toLowerCase() === shortModel(b).toLowerCase();
|
||||
}
|
||||
|
||||
export function modelRouteLabel(requestedModel, actualModel) {
|
||||
function shortEndpointLabel(label) {
|
||||
const value = modelValue(label);
|
||||
if (!value) return '';
|
||||
return value.length > 18 ? value.slice(0, 17) + '…' : value;
|
||||
}
|
||||
|
||||
export function modelRouteLabel(
|
||||
requestedModel,
|
||||
actualModel,
|
||||
requestedEndpointLabel = '',
|
||||
actualEndpointLabel = '',
|
||||
requestedEndpointId = '',
|
||||
actualEndpointId = '',
|
||||
) {
|
||||
const requested = modelValue(requestedModel);
|
||||
const actual = modelValue(actualModel) || requested;
|
||||
if (!requested || sameModelName(requested, actual)) return shortModel(actual || requested);
|
||||
const requestedRoute = modelValue(requestedEndpointId || requestedEndpointLabel);
|
||||
const actualRoute = modelValue(actualEndpointId || actualEndpointLabel);
|
||||
const routeChanged = Boolean(
|
||||
actualRoute
|
||||
&& requestedRoute
|
||||
&& actualRoute !== requestedRoute
|
||||
);
|
||||
if (!requested || sameModelName(requested, actual)) {
|
||||
const model = shortModel(actual || requested);
|
||||
if (!routeChanged) return model;
|
||||
const from = shortEndpointLabel(requestedEndpointLabel || 'Selected route');
|
||||
const to = shortEndpointLabel(actualEndpointLabel || actualEndpointId);
|
||||
return model + ' (' + from + ' -> ' + to + ')';
|
||||
}
|
||||
return shortModel(requested) + ' -> ' + shortModel(actual);
|
||||
}
|
||||
|
||||
@@ -626,10 +655,24 @@ export function replyModelPair(modelName, metadata) {
|
||||
if (actualFromMeta || requestedFromMeta) {
|
||||
const actual = actualFromMeta || requestedFromMeta || modelValue(modelName);
|
||||
const requested = requestedFromMeta || actual;
|
||||
return { requestedModel: requested, actualModel: actual };
|
||||
return {
|
||||
requestedModel: requested,
|
||||
actualModel: actual,
|
||||
requestedEndpointId: meta.requested_endpoint_id || null,
|
||||
requestedEndpointLabel: meta.requested_endpoint_label || 'Selected route',
|
||||
actualEndpointId: meta.endpoint_id || null,
|
||||
actualEndpointLabel: meta.endpoint_label || meta.requested_endpoint_label || 'Selected route',
|
||||
};
|
||||
}
|
||||
const fallback = modelValue(modelName);
|
||||
return { requestedModel: fallback, actualModel: fallback };
|
||||
return {
|
||||
requestedModel: fallback,
|
||||
actualModel: fallback,
|
||||
requestedEndpointId: null,
|
||||
requestedEndpointLabel: 'Selected route',
|
||||
actualEndpointId: null,
|
||||
actualEndpointLabel: 'Selected route',
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -821,12 +864,50 @@ export function isCostTrackedEndpoint(url) {
|
||||
}
|
||||
|
||||
/** Cost for the current turn, returning null for non-billable endpoints. */
|
||||
function _billableCost(model, inputTokens, outputTokens) {
|
||||
const url = _currentEndpointUrl();
|
||||
if (!isCostTrackedEndpoint(url)) return null;
|
||||
function _billableCost(model, inputTokens, outputTokens, endpointCostTracked, selectedEndpointUrl) {
|
||||
// Foreground fallback can answer on a different endpoint than the session's
|
||||
// selected route. Prefer the backend's non-secret actual-route
|
||||
// classification; retain the selected-endpoint check for older history.
|
||||
if (endpointCostTracked === false) return null;
|
||||
const selectedUrl = selectedEndpointUrl === undefined
|
||||
? _currentEndpointUrl()
|
||||
: selectedEndpointUrl;
|
||||
if (endpointCostTracked !== true && !isCostTrackedEndpoint(selectedUrl)) {
|
||||
return null;
|
||||
}
|
||||
return getModelCost(model, inputTokens, outputTokens);
|
||||
}
|
||||
|
||||
/** Sum cost using the route/model that produced each Agent round. */
|
||||
function _metricsBillableCost(metrics, model, inputTokens, outputTokens, selectedEndpointUrl) {
|
||||
const buckets = Array.isArray(metrics.usage_buckets) ? metrics.usage_buckets : [];
|
||||
if (!buckets.length) {
|
||||
return _billableCost(
|
||||
model,
|
||||
inputTokens,
|
||||
outputTokens,
|
||||
metrics.endpoint_cost_tracked,
|
||||
selectedEndpointUrl,
|
||||
);
|
||||
}
|
||||
let total = 0;
|
||||
let hasPricedUsage = false;
|
||||
for (const bucket of buckets) {
|
||||
if (!bucket || typeof bucket !== 'object') continue;
|
||||
const bucketCost = _billableCost(
|
||||
bucket.model || model,
|
||||
Number(bucket.input_tokens) || 0,
|
||||
Number(bucket.output_tokens) || 0,
|
||||
bucket.endpoint_cost_tracked,
|
||||
selectedEndpointUrl,
|
||||
);
|
||||
if (bucketCost === null) continue;
|
||||
total += bucketCost;
|
||||
hasPricedUsage = true;
|
||||
}
|
||||
return hasPricedUsage ? total : null;
|
||||
}
|
||||
|
||||
export function getImageCost(model, quality, size) {
|
||||
if (!model) return null;
|
||||
const m = model.toLowerCase();
|
||||
@@ -841,6 +922,9 @@ export function getImageCost(model, quality, size) {
|
||||
|
||||
/* ── Session cost helpers ─────────────────────────────────────────── */
|
||||
const _COST_KEY = 'ody-session-cost';
|
||||
const _COST_RUNS_KEY = 'ody-session-cost-runs';
|
||||
const _MAX_COST_RUNS_PER_SESSION = 256;
|
||||
const _COST_LEDGER_LOCK = 'odysseus-session-cost-ledger';
|
||||
|
||||
/** Return the accumulated cost for the current (or given) session. */
|
||||
export function getSessionCost(sessionId) {
|
||||
@@ -848,7 +932,14 @@ export function getSessionCost(sessionId) {
|
||||
if (!sid) return 0;
|
||||
try {
|
||||
const costs = JSON.parse(localStorage.getItem(_COST_KEY) || '{}');
|
||||
return costs[sid] || 0;
|
||||
const runCosts = JSON.parse(localStorage.getItem(_COST_RUNS_KEY) || '{}');
|
||||
const recordedRuns = runCosts[sid] && typeof runCosts[sid] === 'object'
|
||||
? Object.values(runCosts[sid])
|
||||
: [];
|
||||
return (costs[sid] || 0) + recordedRuns.reduce(
|
||||
(total, value) => total + (Number(value) || 0),
|
||||
0,
|
||||
);
|
||||
} catch (_e) { return 0; }
|
||||
}
|
||||
|
||||
@@ -860,6 +951,9 @@ export function resetSessionCost(sessionId) {
|
||||
const costs = JSON.parse(localStorage.getItem(_COST_KEY) || '{}');
|
||||
delete costs[sid];
|
||||
localStorage.setItem(_COST_KEY, JSON.stringify(costs));
|
||||
const runCosts = JSON.parse(localStorage.getItem(_COST_RUNS_KEY) || '{}');
|
||||
delete runCosts[sid];
|
||||
localStorage.setItem(_COST_RUNS_KEY, JSON.stringify(runCosts));
|
||||
} catch (_e) { /* ignore */ }
|
||||
updateSessionCostUI();
|
||||
}
|
||||
@@ -868,21 +962,8 @@ export function resetSessionCost(sessionId) {
|
||||
export function updateSessionCostUI() {
|
||||
const el = document.getElementById('session-cost-display');
|
||||
if (!el) return;
|
||||
// Non-billable endpoint? Hide the badge and clear stale cost that a previous
|
||||
// cloud-rate calculation may have left in localStorage for this session.
|
||||
const _url = _currentEndpointUrl();
|
||||
if (!isCostTrackedEndpoint(_url)) {
|
||||
const sid = window.sessionModule && window.sessionModule.getCurrentSessionId();
|
||||
if (sid && getSessionCost(sid) > 0) {
|
||||
try {
|
||||
const costs = JSON.parse(localStorage.getItem(_COST_KEY) || '{}');
|
||||
delete costs[sid];
|
||||
localStorage.setItem(_COST_KEY, JSON.stringify(costs));
|
||||
} catch (_e) { /* ignore */ }
|
||||
}
|
||||
el.style.display = 'none';
|
||||
return;
|
||||
}
|
||||
// The ledger records billable work already performed in this session. A
|
||||
// selected local endpoint does not erase cost from a paid fallback route.
|
||||
const cost = getSessionCost();
|
||||
if (cost > 0) {
|
||||
el.textContent = '$' + (cost < 0.01 ? cost.toFixed(4) : cost < 1 ? cost.toFixed(3) : cost.toFixed(2));
|
||||
@@ -892,6 +973,94 @@ export function updateSessionCostUI() {
|
||||
}
|
||||
}
|
||||
|
||||
/** Record one metrics payload in a session ledger at most once. */
|
||||
export function recordSessionMetricsCost(metrics, sessionId, selectedEndpointUrl) {
|
||||
if (!metrics || typeof metrics !== 'object') return null;
|
||||
const cost = _metricsBillableCost(
|
||||
metrics,
|
||||
metrics.model || 'Unknown',
|
||||
metrics.input_tokens || 0,
|
||||
metrics.output_tokens || 0,
|
||||
selectedEndpointUrl,
|
||||
);
|
||||
if (metrics._fromHistory) return cost;
|
||||
const sid = sessionId || (
|
||||
window.sessionModule && window.sessionModule.getCurrentSessionId()
|
||||
);
|
||||
if (!sid || cost === null) return cost;
|
||||
const runId = typeof metrics._costRecordId === 'string'
|
||||
? metrics._costRecordId.trim()
|
||||
: '';
|
||||
if ((metrics._costRecorded || metrics._costRecordPending) && !runId) return cost;
|
||||
// Recorded is only set once the write actually runs; pending covers the
|
||||
// window while the write waits on the cross-tab lock, so a replay in that
|
||||
// window cannot double-add and a tab closed mid-queue never claims recorded.
|
||||
metrics._costRecordPending = true;
|
||||
const writeCost = () => {
|
||||
if (runId) {
|
||||
try {
|
||||
const runCosts = JSON.parse(localStorage.getItem(_COST_RUNS_KEY) || '{}');
|
||||
const sessionRuns = runCosts[sid] && typeof runCosts[sid] === 'object'
|
||||
? runCosts[sid]
|
||||
: {};
|
||||
// Assigning by detached-run identity is replay-idempotent even when a
|
||||
// refresh produces a fresh metrics object. The Web Lock around this
|
||||
// read/modify/write also keeps distinct runs from two tabs from
|
||||
// overwriting one another's stale snapshot.
|
||||
sessionRuns[runId] = cost;
|
||||
const entries = Object.entries(sessionRuns);
|
||||
if (entries.length > _MAX_COST_RUNS_PER_SESSION) {
|
||||
const overflow = entries.slice(0, entries.length - _MAX_COST_RUNS_PER_SESSION);
|
||||
const costs = JSON.parse(localStorage.getItem(_COST_KEY) || '{}');
|
||||
costs[sid] = (costs[sid] || 0) + overflow.reduce(
|
||||
(total, entry) => total + (Number(entry[1]) || 0),
|
||||
0,
|
||||
);
|
||||
overflow.forEach(([oldRunId]) => delete sessionRuns[oldRunId]);
|
||||
localStorage.setItem(_COST_KEY, JSON.stringify(costs));
|
||||
}
|
||||
runCosts[sid] = sessionRuns;
|
||||
localStorage.setItem(_COST_RUNS_KEY, JSON.stringify(runCosts));
|
||||
} catch (_e) { /* ignore */ }
|
||||
} else {
|
||||
try {
|
||||
const costs = JSON.parse(localStorage.getItem(_COST_KEY) || '{}');
|
||||
costs[sid] = (costs[sid] || 0) + cost;
|
||||
localStorage.setItem(_COST_KEY, JSON.stringify(costs));
|
||||
} catch (_e) { /* ignore */ }
|
||||
}
|
||||
metrics._costRecorded = true;
|
||||
metrics._costRecordPending = false;
|
||||
const currentSid = window.sessionModule && window.sessionModule.getCurrentSessionId();
|
||||
if (currentSid === sid) updateSessionCostUI();
|
||||
};
|
||||
|
||||
let writeStarted = false;
|
||||
const guardedWrite = () => {
|
||||
writeStarted = true;
|
||||
writeCost();
|
||||
};
|
||||
try {
|
||||
if (
|
||||
typeof navigator !== 'undefined'
|
||||
&& navigator.locks
|
||||
&& typeof navigator.locks.request === 'function'
|
||||
) {
|
||||
const pendingWrite = navigator.locks.request(_COST_LEDGER_LOCK, guardedWrite);
|
||||
if (pendingWrite && typeof pendingWrite.catch === 'function') {
|
||||
pendingWrite.catch(() => {
|
||||
if (!writeStarted) guardedWrite();
|
||||
});
|
||||
}
|
||||
} else {
|
||||
guardedWrite();
|
||||
}
|
||||
} catch (_e) {
|
||||
if (!writeStarted) guardedWrite();
|
||||
}
|
||||
return cost;
|
||||
}
|
||||
|
||||
/** Create a timestamp span for role labels.
|
||||
* Pass an ISO string / Date / epoch-ms to render the message's own time
|
||||
* (used when replaying history). Falls back to "now" when no value is given. */
|
||||
@@ -1871,23 +2040,19 @@ export function displayMetrics(messageElement, metrics) {
|
||||
const isReal = metrics.usage_source === 'real';
|
||||
const ctxPct = metrics.context_percent;
|
||||
const model = metrics.model || 'Unknown';
|
||||
const cost = _billableCost(model, inputTokens, outputTokens);
|
||||
const cost = _metricsBillableCost(
|
||||
metrics,
|
||||
model,
|
||||
inputTokens,
|
||||
outputTokens,
|
||||
);
|
||||
|
||||
// Nothing useful to show — bail out (only if ALL metrics are missing)
|
||||
if (!responseTime && !inputTokens && !outputTokens && tps == null && !ctxPct) return;
|
||||
|
||||
// Accumulate session cost (only on fresh metrics, not history reload)
|
||||
if (!metrics._fromHistory) {
|
||||
const _sid = window.sessionModule && window.sessionModule.getCurrentSessionId();
|
||||
if (_sid && cost !== null) {
|
||||
try {
|
||||
const _costs = JSON.parse(localStorage.getItem(_COST_KEY) || '{}');
|
||||
_costs[_sid] = (_costs[_sid] || 0) + cost;
|
||||
localStorage.setItem(_COST_KEY, JSON.stringify(_costs));
|
||||
} catch (_e) { /* ignore */ }
|
||||
updateSessionCostUI();
|
||||
}
|
||||
}
|
||||
// Rendering can occur when metrics arrive and again after [DONE]. The
|
||||
// ledger mutation is idempotent for that shared payload.
|
||||
recordSessionMetricsCost(metrics);
|
||||
|
||||
// Keep token counts in the Message Stats popup; the footer should stay slim.
|
||||
const costStr0 = cost !== null ? `$${cost < 0.01 ? cost.toFixed(4) : cost.toFixed(3)}` : null;
|
||||
@@ -2304,9 +2469,19 @@ export function addMessage(role, content, modelName, metadata) {
|
||||
const textRaw = Array.isArray(content) ? markdownModule.renderContent(content) : content;
|
||||
|
||||
// --- Agent multi-bubble reconstruction from saved metadata ---
|
||||
if (role === 'assistant' && metadata && metadata.tool_events && metadata.tool_events.length > 0) {
|
||||
if (
|
||||
role === 'assistant'
|
||||
&& metadata
|
||||
&& (
|
||||
(Array.isArray(metadata.tool_events) && metadata.tool_events.length > 0)
|
||||
|| (Array.isArray(metadata.round_texts) && metadata.round_texts.length > 1)
|
||||
)
|
||||
) {
|
||||
const roundTexts = metadata.round_texts || [];
|
||||
const toolEvents = metadata.tool_events;
|
||||
const roundModels = metadata.round_models || [];
|
||||
const roundEndpointIds = metadata.round_endpoint_ids || [];
|
||||
const roundEndpointLabels = metadata.round_endpoint_labels || [];
|
||||
const toolEvents = metadata.tool_events || [];
|
||||
let pendingAskUser = null;
|
||||
let lastWrap = null;
|
||||
let firstMsgAi = null;
|
||||
@@ -2319,7 +2494,8 @@ export function addMessage(role, content, modelName, metadata) {
|
||||
toolsByRound[r].push(ev);
|
||||
}
|
||||
|
||||
const maxRound = Math.max(...Object.keys(toolsByRound).map(Number), roundTexts.length);
|
||||
const toolRounds = Object.keys(toolsByRound).map(Number);
|
||||
const maxRound = Math.max(toolRounds.length ? Math.max(...toolRounds) : 0, roundTexts.length);
|
||||
|
||||
for (let r = 0; r < maxRound; r++) {
|
||||
const roundNum = r + 1;
|
||||
@@ -2331,10 +2507,31 @@ export function addMessage(role, content, modelName, metadata) {
|
||||
const roleEl = document.createElement('div');
|
||||
roleEl.className = 'role';
|
||||
const pair = replyModelPair(modelName, metadata);
|
||||
const contModel = pair.actualModel || pair.requestedModel;
|
||||
roleEl.textContent = modelRouteLabel(pair.requestedModel, contModel);
|
||||
if (pair.requestedModel && contModel && !sameModelName(pair.requestedModel, contModel)) {
|
||||
roleEl.title = pair.requestedModel + ' -> ' + contModel;
|
||||
const contModel = roundModels[r] || pair.actualModel || pair.requestedModel;
|
||||
const contEndpointId = r < roundEndpointIds.length
|
||||
? roundEndpointIds[r]
|
||||
: pair.actualEndpointId;
|
||||
const contEndpointLabel = r < roundEndpointLabels.length
|
||||
? roundEndpointLabels[r]
|
||||
: pair.actualEndpointLabel;
|
||||
roleEl.textContent = modelRouteLabel(
|
||||
pair.requestedModel,
|
||||
contModel,
|
||||
pair.requestedEndpointLabel,
|
||||
contEndpointLabel,
|
||||
pair.requestedEndpointId,
|
||||
contEndpointId,
|
||||
);
|
||||
if (
|
||||
pair.requestedModel
|
||||
&& contModel
|
||||
&& (
|
||||
!sameModelName(pair.requestedModel, contModel)
|
||||
|| (pair.requestedEndpointId && contEndpointId && pair.requestedEndpointId !== contEndpointId)
|
||||
)
|
||||
) {
|
||||
roleEl.title = pair.requestedModel + ' -> ' + contModel
|
||||
+ ' (' + pair.requestedEndpointLabel + ' -> ' + contEndpointLabel + ')';
|
||||
}
|
||||
applyModelColor(roleEl, contModel);
|
||||
if (r === 0) roleEl.appendChild(roleTimestamp(metadata?.timestamp));
|
||||
@@ -2489,7 +2686,14 @@ export function addMessage(role, content, modelName, metadata) {
|
||||
const isCompacted = metadata?.compacted;
|
||||
const replyModels = replyModelPair(modelName, metadata);
|
||||
const resolvedModel = replyModels.actualModel || replyModels.requestedModel;
|
||||
var _roleText = role === 'user' ? 'You' : (isSlash || isCompacted) ? 'Odysseus' : modelRouteLabel(replyModels.requestedModel, resolvedModel);
|
||||
var _roleText = role === 'user' ? 'You' : (isSlash || isCompacted) ? 'Odysseus' : modelRouteLabel(
|
||||
replyModels.requestedModel,
|
||||
resolvedModel,
|
||||
replyModels.requestedEndpointLabel,
|
||||
replyModels.actualEndpointLabel,
|
||||
replyModels.requestedEndpointId,
|
||||
replyModels.actualEndpointId,
|
||||
);
|
||||
if (role === 'assistant' && (metadata?.research || metadata?.research_clarification)) {
|
||||
_roleText += ' (Research)';
|
||||
}
|
||||
@@ -2500,8 +2704,14 @@ export function addMessage(role, content, modelName, metadata) {
|
||||
}
|
||||
r.textContent = _roleText;
|
||||
if (role !== 'user') {
|
||||
if (!isSlash && !isCompacted && replyModels.requestedModel && resolvedModel && !sameModelName(replyModels.requestedModel, resolvedModel)) {
|
||||
r.title = replyModels.requestedModel + ' -> ' + resolvedModel;
|
||||
const endpointChanged = Boolean(
|
||||
replyModels.requestedEndpointId
|
||||
&& replyModels.actualEndpointId
|
||||
&& replyModels.requestedEndpointId !== replyModels.actualEndpointId
|
||||
);
|
||||
if (!isSlash && !isCompacted && replyModels.requestedModel && resolvedModel && (!sameModelName(replyModels.requestedModel, resolvedModel) || endpointChanged)) {
|
||||
r.title = replyModels.requestedModel + ' -> ' + resolvedModel
|
||||
+ ' (' + replyModels.requestedEndpointLabel + ' -> ' + replyModels.actualEndpointLabel + ')';
|
||||
}
|
||||
if (!isSlash && !isCompacted) applyModelColor(r, resolvedModel);
|
||||
r.appendChild(roleTimestamp(metadata?.timestamp));
|
||||
@@ -2785,6 +2995,7 @@ const chatRenderer = {
|
||||
getSessionCost,
|
||||
resetSessionCost,
|
||||
updateSessionCostUI,
|
||||
recordSessionMetricsCost,
|
||||
roleTimestamp,
|
||||
stripToolBlocks,
|
||||
copyMessageText,
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
/** Build a terminal stream error while preserving provider-supplied text. */
|
||||
export function createTerminalStreamError(payload = {}) {
|
||||
const rawError = payload.error;
|
||||
const message = (
|
||||
payload.text
|
||||
|| (typeof rawError === 'string' ? rawError : rawError?.message)
|
||||
|| `Error ${payload.status || 'unknown'}`
|
||||
);
|
||||
const error = new Error(message);
|
||||
error.name = 'TerminalStreamError';
|
||||
error.terminalStreamError = true;
|
||||
error.status = payload.status;
|
||||
return error;
|
||||
}
|
||||
|
||||
/** Only connection-class stream failures are safe to resubmit automatically. */
|
||||
export function isRecoverableStreamError(error) {
|
||||
if (!error || error.terminalStreamError || error.name === 'TerminalStreamError') return false;
|
||||
if (error.name === 'TypeError') return true;
|
||||
const message = (error.message || '').toLowerCase();
|
||||
if (/\btool\b|unsupported|json|parse|\b4\d\d\b|\b5\d\d\b/.test(message)) return false;
|
||||
return /network|fetch|connection|reset|closed|aborted|stream|tim(?:e|ed)\s?out|econn|eof/.test(message);
|
||||
}
|
||||
@@ -143,9 +143,9 @@ export function wireArrowUpRecall(composer, getUserMessages, options = {}) {
|
||||
return;
|
||||
}
|
||||
|
||||
// ArrowUp owns prompt history in the chat composer. If the current text
|
||||
// is not already a recalled prompt, start from newest instead of letting
|
||||
// the browser move the caret inside the textarea.
|
||||
// ArrowUp walks older prompts. An unmatched draft already returned above,
|
||||
// so reaching here means the composer is empty or holds a recalled prompt
|
||||
// — the caret-navigation case is never hijacked.
|
||||
const nextIndex = currentIndex >= 0 ? Math.min(currentIndex + 1, history.length - 1) : 0;
|
||||
const recalled = history[nextIndex];
|
||||
if (!recalled) {
|
||||
|
||||
+49
-10
@@ -149,6 +149,7 @@ let _loading = false;
|
||||
let _expanded = false;
|
||||
let _docModule = null;
|
||||
let _listSpinner = null;
|
||||
let _openEmailRequestSeq = 0;
|
||||
let _senderFilter = null; // email address (lowercased) to filter by, or null
|
||||
let _senderFilterLabel = null; // display label for the active filter chip
|
||||
let _showEmailTags = localStorage.getItem('odysseus.email.showTags') !== '0';
|
||||
@@ -187,7 +188,7 @@ export function init(documentModule) {
|
||||
} catch (_) {}
|
||||
if (opts.compose) { _composeNew(); return; }
|
||||
if (opts.email) {
|
||||
await _openEmail(opts.email, null, opts.emailData, opts.mode || 'reply', opts.noteHint || '');
|
||||
await _openEmail(opts.email, null, opts.emailData, opts.mode || 'reply', opts.noteHint || '', '', opts.mailboxContext || null);
|
||||
}
|
||||
},
|
||||
});
|
||||
@@ -751,7 +752,21 @@ function _createEmailItem(em) {
|
||||
return item;
|
||||
}
|
||||
|
||||
async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', noteHint = '', prefilledBody = '') {
|
||||
async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', noteHint = '', prefilledBody = '', mailboxContext = null) {
|
||||
const openRequestSeq = ++_openEmailRequestSeq;
|
||||
const folderAtStart = mailboxContext?.messageFolder || _currentFolder;
|
||||
const accountAtStart = mailboxContext?.accountId ?? (window.__odysseusActiveEmailAccount || '');
|
||||
const accountQueryAtStart = accountAtStart ? `&account_id=${encodeURIComponent(accountAtStart)}` : '';
|
||||
const mailboxContextIsCurrent = typeof mailboxContext?.isCurrent === 'function'
|
||||
? mailboxContext.isCurrent
|
||||
: () => (
|
||||
folderAtStart === _currentFolder &&
|
||||
accountAtStart === (window.__odysseusActiveEmailAccount || '')
|
||||
);
|
||||
const isCurrentOpen = () => (
|
||||
openRequestSeq === _openEmailRequestSeq &&
|
||||
mailboxContextIsCurrent()
|
||||
);
|
||||
const aiReplyMode = mode === 'ai-reply-fast' ? 'fast' : '';
|
||||
const wantsAiReply = mode === 'ai-reply' || !!aiReplyMode;
|
||||
// Body pre-fill from the agent's open_email_reply tool call takes the
|
||||
@@ -780,9 +795,10 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
let data = preloadedData;
|
||||
if (!data) {
|
||||
const fullQS = mode === 'forward' ? '&full=1' : '';
|
||||
const res = await fetch(`${API_BASE}/api/email/read/${em.uid}?folder=${encodeURIComponent(_currentFolder)}${_acct()}${fullQS}`);
|
||||
const res = await fetch(`${API_BASE}/api/email/read/${em.uid}?folder=${encodeURIComponent(folderAtStart)}${accountQueryAtStart}&mark_seen=true${fullQS}`);
|
||||
data = await res.json();
|
||||
}
|
||||
if (!isCurrentOpen()) return;
|
||||
if (data.error) {
|
||||
console.error('Failed to read email:', data.error);
|
||||
return;
|
||||
@@ -808,7 +824,7 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
message_id: _fallback(data.message_id, em.message_id),
|
||||
};
|
||||
if (wantsAiReply) {
|
||||
const activeReplyAccount = data.account_id || em.account_id || window.__odysseusActiveEmailAccount || '';
|
||||
const activeReplyAccount = data.account_id || em.account_id || accountAtStart;
|
||||
if (data.cached_ai_reply && !noteHint && !activeReplyAccount) {
|
||||
aiSuggestedBody = _cleanAiReplyText(data.cached_ai_reply);
|
||||
} else {
|
||||
@@ -834,7 +850,7 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
session_id: currentSessionId,
|
||||
message_id: data.message_id || '',
|
||||
uid: String(em.uid || ''),
|
||||
folder: _currentFolder,
|
||||
folder: folderAtStart,
|
||||
account_id: activeReplyAccount,
|
||||
fast: true,
|
||||
user_hint: (noteHint || '').trim() || undefined,
|
||||
@@ -842,6 +858,7 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
});
|
||||
const result = await res.json();
|
||||
if (draftToastTimer) clearTimeout(draftToastTimer);
|
||||
if (!isCurrentOpen()) return;
|
||||
if (result.success && result.reply) {
|
||||
aiSuggestedBody = _cleanAiReplyText(result.reply);
|
||||
} else {
|
||||
@@ -855,6 +872,7 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
}
|
||||
} catch (e) {
|
||||
if (draftToastTimer) clearTimeout(draftToastTimer);
|
||||
if (!isCurrentOpen()) return;
|
||||
console.error('AI reply generation failed:', e);
|
||||
import('./ui.js').then(m => m.showError && m.showError('AI reply failed: ' + (e.message || e))).catch(() => {});
|
||||
return;
|
||||
@@ -862,8 +880,12 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
}
|
||||
}
|
||||
|
||||
em.is_read = true;
|
||||
if (itemEl) itemEl.classList.remove('email-unread');
|
||||
if (!isCurrentOpen()) return;
|
||||
// Only claim the message is read when the provider accepted the \Seen
|
||||
// transition. A failed STORE still opens the message; it just stays unread.
|
||||
const markedSeen = !data.mark_seen_failed;
|
||||
em.is_read = markedSeen;
|
||||
if (itemEl) itemEl.classList.toggle('email-unread', !markedSeen);
|
||||
|
||||
// Addresses to exclude from Reply All. Prefer the full set of configured
|
||||
// accounts (so a multi-account user's other mailboxes are excluded too),
|
||||
@@ -911,7 +933,7 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
if (mode !== 'forward' && data.message_id) content += `\nIn-Reply-To: ${data.message_id}`;
|
||||
if (mode !== 'forward' && data.message_id) content += `\nReferences: ${data.references ? data.references + ' ' + data.message_id : data.message_id}`;
|
||||
content += `\nX-Source-UID: ${em.uid}`;
|
||||
content += `\nX-Source-Folder: ${_currentFolder}`;
|
||||
content += `\nX-Source-Folder: ${folderAtStart}`;
|
||||
if (data.attachments && data.attachments.length > 0) {
|
||||
const attStr = data.attachments.map(a => `${a.index}:${a.filename}:${a.size}`).join('|');
|
||||
content += `\nX-Attachments: ${attStr}`;
|
||||
@@ -980,21 +1002,27 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
// and block Send on long threads.
|
||||
const reuseExisting = mode !== 'forward' && !!aiSuggestedBody;
|
||||
const existingDocId = (reuseExisting && _docModule.findEmailDocId)
|
||||
? _docModule.findEmailDocId(em.uid, _currentFolder)
|
||||
? _docModule.findEmailDocId(em.uid, folderAtStart)
|
||||
: null;
|
||||
if (existingDocId) {
|
||||
if (!_docModule.isPanelOpen()) _docModule.openPanel();
|
||||
await new Promise(r => requestAnimationFrame(() => requestAnimationFrame(r)));
|
||||
if (!isCurrentOpen()) return;
|
||||
await _docModule.loadDocument(existingDocId);
|
||||
if (!isCurrentOpen()) return;
|
||||
if (typeof _docModule.ensureEmailDraftEnvelope === 'function') {
|
||||
await _docModule.ensureEmailDraftEnvelope(existingDocId, content);
|
||||
if (!isCurrentOpen()) return;
|
||||
}
|
||||
if (aiSuggestedBody && typeof _docModule.replaceEmailReplyBody === 'function') {
|
||||
await _docModule.replaceEmailReplyBody(existingDocId, aiSuggestedBody, { force: false });
|
||||
if (!isCurrentOpen()) return;
|
||||
}
|
||||
_bringEmailReplyDraftToFrontOnMobile();
|
||||
} else {
|
||||
if (!isCurrentOpen()) return;
|
||||
let activeSid = await _createEmailChat(data, { forceNew: true });
|
||||
if (!isCurrentOpen()) return;
|
||||
if (!activeSid) {
|
||||
console.error('reply: could not obtain a session_id');
|
||||
import('./ui.js').then(m => m.showError && m.showError('Could not start a reply chat.')).catch(() => {});
|
||||
@@ -1012,13 +1040,20 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
}),
|
||||
});
|
||||
let docRes = await createReplyDoc(activeSid);
|
||||
if (!isCurrentOpen()) return;
|
||||
if (docRes.status === 404) {
|
||||
console.warn('[reply-debug] draft session rejected; retrying in a fresh email chat', activeSid);
|
||||
if (!isCurrentOpen()) return;
|
||||
activeSid = await _createEmailChat(data, { forceNew: true });
|
||||
if (activeSid) docRes = await createReplyDoc(activeSid);
|
||||
if (!isCurrentOpen()) return;
|
||||
if (activeSid) {
|
||||
docRes = await createReplyDoc(activeSid);
|
||||
if (!isCurrentOpen()) return;
|
||||
}
|
||||
}
|
||||
if (!docRes.ok) {
|
||||
const errText = await docRes.text();
|
||||
if (!isCurrentOpen()) return;
|
||||
console.error('[reply-debug] POST /api/document failed', docRes.status, errText);
|
||||
// uiModule isn't statically imported here — use the dynamic
|
||||
// import pattern the rest of this file uses. (Previously this
|
||||
@@ -1028,10 +1063,12 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
return;
|
||||
}
|
||||
const doc = await docRes.json();
|
||||
if (!isCurrentOpen()) return;
|
||||
if (doc.id) {
|
||||
const wasOpen = _docModule.isPanelOpen();
|
||||
if (!wasOpen) _docModule.openPanel();
|
||||
await new Promise(r => requestAnimationFrame(() => requestAnimationFrame(r)));
|
||||
if (!isCurrentOpen()) return;
|
||||
// Use the doc dict from the POST directly — avoids a 404 race
|
||||
// when the GET fires before the new row is visible to the read
|
||||
// connection (or when caching is interfering). loadDocument's
|
||||
@@ -1040,12 +1077,14 @@ async function _openEmail(em, itemEl, preloadedData = null, mode = 'reply', note
|
||||
_docModule.injectFreshDoc(doc);
|
||||
} else {
|
||||
await _docModule.loadDocument(doc.id);
|
||||
if (!isCurrentOpen()) return;
|
||||
}
|
||||
_bringEmailReplyDraftToFrontOnMobile();
|
||||
}
|
||||
}
|
||||
}
|
||||
} catch (e) {
|
||||
if (!isCurrentOpen()) return;
|
||||
console.error('Failed to open email:', e);
|
||||
// Surface the failure so a silent throw in the reply flow doesn't
|
||||
// look like "nothing happened". Dynamic import — uiModule isn't a
|
||||
|
||||
+423
-157
@@ -13,7 +13,7 @@ import { makeWindowDraggable } from './windowDrag.js';
|
||||
import {
|
||||
_esc, _escLinkify, _extractName, _parseTurnMeta,
|
||||
_formatBubbleDate, _formatRecipients, _senderColor, _initials,
|
||||
_sanitizeHtml,
|
||||
_sanitizeHtml, _renderEmailSummaryError,
|
||||
_TALON_WROTE, _TALON_FROM, _TALON_SENT, _TALON_SUBJ, _TALON_TO,
|
||||
_TALON_ORIG_RE, _SIG_BLOAT_MIN_CHARS,
|
||||
} from './emailLibrary/utils.js';
|
||||
@@ -30,6 +30,10 @@ import { bindMenuDismiss, dismissOrRemove } from './escMenuStack.js';
|
||||
const API_BASE = window.location.origin;
|
||||
let _emailUnreadChipClickWired = false;
|
||||
let _libLoadSeq = 0;
|
||||
let _emailMailboxGeneration = 0;
|
||||
let _emailCardOpenSeq = 0;
|
||||
let _emailReadMutationSeq = 0;
|
||||
const _emailReadMutations = new Map();
|
||||
let _libFolderSeq = 0;
|
||||
let _libSearchSeq = 0;
|
||||
let _libSearchHadResults = false;
|
||||
@@ -837,14 +841,41 @@ document.addEventListener('keydown', (e) => {
|
||||
e.stopImmediatePropagation?.();
|
||||
}, true);
|
||||
|
||||
function _syncEmailReadState(uid, isRead = true) {
|
||||
function _emailReadContextKey(context) {
|
||||
return [context.accountId, context.folder, context.uid].map(value => String(value || '')).join('\u0000');
|
||||
}
|
||||
|
||||
function _emailReadContextIsCurrent(context) {
|
||||
if (!context) return true;
|
||||
return (
|
||||
String(state._libAccountId || '') === context.accountId &&
|
||||
String(state._libFolder || 'INBOX') === context.libraryFolder &&
|
||||
_emailMailboxGeneration === context.mailboxGeneration
|
||||
);
|
||||
}
|
||||
|
||||
function _emailMatchesReadContext(email, context) {
|
||||
if (String(email?.uid || '') !== context.uid) return false;
|
||||
const accountId = String(email?.account_id || context.accountId);
|
||||
const folder = String(email?.folder || context.folder);
|
||||
return accountId === context.accountId && folder === context.folder;
|
||||
}
|
||||
|
||||
function _syncEmailReadState(uid, isRead = true, context = null) {
|
||||
if (uid == null) return;
|
||||
const uidStr = String(uid);
|
||||
const read = !!isRead;
|
||||
const match = (state._libEmails || []).find(x => String(x.uid) === uidStr);
|
||||
if (context && (!_emailReadContextIsCurrent(context) || uidStr !== context.uid)) return;
|
||||
const match = (state._libEmails || []).find(x => (
|
||||
context ? _emailMatchesReadContext(x, context) : String(x.uid) === uidStr
|
||||
));
|
||||
if (match) match.is_read = read;
|
||||
|
||||
document.querySelectorAll('.doclib-card[data-uid="' + CSS.escape(uidStr) + '"]').forEach(card => {
|
||||
if (context && (
|
||||
String(card.dataset.emailAccount || '') !== context.accountId ||
|
||||
String(card.dataset.emailFolder || '') !== context.folder
|
||||
)) return;
|
||||
card.classList.toggle('email-card-unread', !read);
|
||||
const titleRow = card.querySelector('.email-card-titlerow');
|
||||
if (read) {
|
||||
@@ -1762,11 +1793,18 @@ function _rememberedEmailAccountId() {
|
||||
// results and __scheduled__ are deliberately not cached.
|
||||
const _libListCache = new Map();
|
||||
const _LIB_CACHE_MAX = 24;
|
||||
const _LIB_INITIAL_PAGE_SIZE = 100;
|
||||
const _LIB_SESSION_CACHE_PREFIX = 'odysseus.email.list.';
|
||||
const _LIB_SESSION_CACHE_TTL_MS = 10 * 60 * 1000;
|
||||
const _LIB_LAST_ACCOUNT_KEY = 'odysseus.email.lastAccountId';
|
||||
let _libPrewarmTimer = null;
|
||||
const _LIB_PREWARM_COOLDOWN_MS = 5 * 60 * 1000;
|
||||
let _libPrewarmDelayTimer = null;
|
||||
let _libPrewarmIdleHandle = null;
|
||||
let _libPrewarmPromise = null;
|
||||
let _libPrewarmResolve = null;
|
||||
let _libPrewarmAbortController = null;
|
||||
let _libPrewarmDetachPriorityListeners = null;
|
||||
let _libPrewarmGeneration = 0;
|
||||
let _libLastPrewarmAt = 0;
|
||||
let _libUnreadPrewarmKey = '';
|
||||
let _libUnreadPrewarmAt = 0;
|
||||
@@ -1908,6 +1946,7 @@ function _resetEmailListForFreshLoad({ useCache = true } = {}) {
|
||||
_exitEmailReaderModeForList();
|
||||
_resetBulkSelectionForContextChange();
|
||||
state._libOffset = 0;
|
||||
_emailMailboxGeneration += 1;
|
||||
_libLoadSeq += 1;
|
||||
const ck = _libCacheKey();
|
||||
const cached = useCache ? _libCacheGet(ck) : null;
|
||||
@@ -2076,162 +2115,319 @@ function _isChatInteractionBusy() {
|
||||
}
|
||||
}
|
||||
|
||||
function _loadEmailsWhenChatIdle({ delay = 50, retries = 180, options = {} } = {}) {
|
||||
const run = () => {
|
||||
if (!state._libOpen || !document.getElementById('email-lib-modal')) return;
|
||||
if (_isChatInteractionBusy() && retries > 0) {
|
||||
setTimeout(() => _loadEmailsWhenChatIdle({ delay: 1000, retries: retries - 1, options }), 1000);
|
||||
function _canRunEmailPrewarm() {
|
||||
if (state._libOpen || state._libLoading || _libSearchInFlight) return false;
|
||||
if (document.visibilityState && document.visibilityState !== 'visible') return false;
|
||||
return !_isChatInteractionBusy();
|
||||
}
|
||||
|
||||
function _isEmailPrewarmTemporarilyBlocked() {
|
||||
if (state._libOpen || state._libLoading || _libSearchInFlight) return false;
|
||||
if (document.visibilityState && document.visibilityState !== 'visible') return false;
|
||||
return _isChatInteractionBusy();
|
||||
}
|
||||
|
||||
function _isEmailPrewarmCurrent(generation, signal) {
|
||||
return generation === _libPrewarmGeneration
|
||||
&& !signal?.aborted
|
||||
&& _canRunEmailPrewarm();
|
||||
}
|
||||
|
||||
function _settleEmailPrewarm(generation, value = false) {
|
||||
if (generation !== _libPrewarmGeneration) return;
|
||||
const resolve = _libPrewarmResolve;
|
||||
const detachPriorityListeners = _libPrewarmDetachPriorityListeners;
|
||||
_libPrewarmDelayTimer = null;
|
||||
_libPrewarmIdleHandle = null;
|
||||
_libPrewarmPromise = null;
|
||||
_libPrewarmResolve = null;
|
||||
_libPrewarmAbortController = null;
|
||||
_libPrewarmDetachPriorityListeners = null;
|
||||
detachPriorityListeners?.();
|
||||
resolve?.(value);
|
||||
}
|
||||
|
||||
function _cancelEmailPrewarm() {
|
||||
const resolve = _libPrewarmResolve;
|
||||
const detachPriorityListeners = _libPrewarmDetachPriorityListeners;
|
||||
_libPrewarmGeneration += 1;
|
||||
if (_libPrewarmDelayTimer !== null) {
|
||||
clearTimeout(_libPrewarmDelayTimer);
|
||||
}
|
||||
if (_libPrewarmIdleHandle !== null && typeof window.cancelIdleCallback === 'function') {
|
||||
try { window.cancelIdleCallback(_libPrewarmIdleHandle); } catch (_) {}
|
||||
}
|
||||
try { _libPrewarmAbortController?.abort(); } catch (_) {}
|
||||
_libPrewarmDelayTimer = null;
|
||||
_libPrewarmIdleHandle = null;
|
||||
_libPrewarmPromise = null;
|
||||
_libPrewarmResolve = null;
|
||||
_libPrewarmAbortController = null;
|
||||
_libPrewarmDetachPriorityListeners = null;
|
||||
detachPriorityListeners?.();
|
||||
resolve?.(false);
|
||||
}
|
||||
|
||||
function _scheduleEmailPrewarm(task, { delay = 0 } = {}) {
|
||||
if (_libPrewarmPromise) return _libPrewarmPromise;
|
||||
// Do not disguise a timer as idle work. Browsers without the genuine idle
|
||||
// callback simply skip this optional optimization and load on demand.
|
||||
if (typeof window.requestIdleCallback !== 'function') return Promise.resolve(false);
|
||||
|
||||
const generation = ++_libPrewarmGeneration;
|
||||
_libPrewarmPromise = new Promise(resolve => { _libPrewarmResolve = resolve; });
|
||||
const promise = _libPrewarmPromise;
|
||||
let attemptPending = false;
|
||||
let retryRequested = false;
|
||||
|
||||
function clearScheduledAttempt() {
|
||||
if (_libPrewarmDelayTimer !== null) clearTimeout(_libPrewarmDelayTimer);
|
||||
if (_libPrewarmIdleHandle !== null && typeof window.cancelIdleCallback === 'function') {
|
||||
try { window.cancelIdleCallback(_libPrewarmIdleHandle); } catch (_) {}
|
||||
}
|
||||
_libPrewarmDelayTimer = null;
|
||||
_libPrewarmIdleHandle = null;
|
||||
}
|
||||
|
||||
function scheduleIdleRetry(delay = 500) {
|
||||
if (generation !== _libPrewarmGeneration) return;
|
||||
retryRequested = true;
|
||||
if (attemptPending || _libPrewarmDelayTimer !== null || _libPrewarmIdleHandle !== null) return;
|
||||
if (document.visibilityState && document.visibilityState !== 'visible') return;
|
||||
_libPrewarmDelayTimer = setTimeout(requestIdle, Math.max(50, Number(delay) || 500));
|
||||
}
|
||||
|
||||
function handlePriorityChange() {
|
||||
if (generation !== _libPrewarmGeneration) return;
|
||||
if (_canRunEmailPrewarm()) {
|
||||
scheduleIdleRetry(50);
|
||||
return;
|
||||
}
|
||||
_loadEmails(options);
|
||||
|
||||
const priorityBlocked = _isChatInteractionBusy()
|
||||
|| (document.visibilityState && document.visibilityState !== 'visible');
|
||||
if (!priorityBlocked) return;
|
||||
|
||||
retryRequested = true;
|
||||
clearScheduledAttempt();
|
||||
const controller = _libPrewarmAbortController;
|
||||
_libPrewarmAbortController = null;
|
||||
try { controller?.abort(); } catch (_) {}
|
||||
// A hidden page waits for visibilitychange. Chat priority also retains the
|
||||
// timer fallback for busy-until windows whose final transition has no event.
|
||||
if (!document.visibilityState || document.visibilityState === 'visible') {
|
||||
scheduleIdleRetry();
|
||||
}
|
||||
}
|
||||
|
||||
window.addEventListener('odysseus:chat-busy-change', handlePriorityChange);
|
||||
document.addEventListener('visibilitychange', handlePriorityChange);
|
||||
_libPrewarmDetachPriorityListeners = () => {
|
||||
window.removeEventListener('odysseus:chat-busy-change', handlePriorityChange);
|
||||
document.removeEventListener('visibilitychange', handlePriorityChange);
|
||||
};
|
||||
setTimeout(run, Math.max(0, Number(delay) || 0));
|
||||
|
||||
function requestIdle() {
|
||||
if (generation !== _libPrewarmGeneration) return;
|
||||
_libPrewarmDelayTimer = null;
|
||||
try {
|
||||
_libPrewarmIdleHandle = window.requestIdleCallback((deadline) => {
|
||||
if (generation !== _libPrewarmGeneration) return;
|
||||
_libPrewarmIdleHandle = null;
|
||||
const hasIdleBudget = Boolean(
|
||||
deadline
|
||||
&& !deadline.didTimeout
|
||||
&& typeof deadline.timeRemaining === 'function'
|
||||
&& deadline.timeRemaining() > 0
|
||||
);
|
||||
if (!_canRunEmailPrewarm()) {
|
||||
if (_isEmailPrewarmTemporarilyBlocked()) {
|
||||
scheduleIdleRetry();
|
||||
} else {
|
||||
_settleEmailPrewarm(generation, false);
|
||||
}
|
||||
return;
|
||||
}
|
||||
if (!hasIdleBudget) {
|
||||
scheduleIdleRetry();
|
||||
return;
|
||||
}
|
||||
if (generation !== _libPrewarmGeneration) {
|
||||
_settleEmailPrewarm(generation, false);
|
||||
return;
|
||||
}
|
||||
const controller = new AbortController();
|
||||
_libPrewarmAbortController = controller;
|
||||
attemptPending = true;
|
||||
retryRequested = false;
|
||||
Promise.resolve()
|
||||
.then(() => task({ signal: controller.signal, generation }))
|
||||
.then(value => {
|
||||
if (controller !== _libPrewarmAbortController || controller.signal.aborted) return;
|
||||
_settleEmailPrewarm(generation, Boolean(value));
|
||||
})
|
||||
.catch(() => {
|
||||
if (controller !== _libPrewarmAbortController || controller.signal.aborted) return;
|
||||
_settleEmailPrewarm(generation, false);
|
||||
})
|
||||
.finally(() => {
|
||||
attemptPending = false;
|
||||
if (generation !== _libPrewarmGeneration) return;
|
||||
if (retryRequested) scheduleIdleRetry();
|
||||
});
|
||||
});
|
||||
} catch (_) {
|
||||
_settleEmailPrewarm(generation, false);
|
||||
}
|
||||
}
|
||||
|
||||
const wait = Math.max(0, Number(delay) || 0);
|
||||
if (wait > 0) _libPrewarmDelayTimer = setTimeout(requestIdle, wait);
|
||||
else requestIdle();
|
||||
return promise;
|
||||
}
|
||||
|
||||
export function prewarmEmailLibrary({ delay = 2500 } = {}) {
|
||||
if (_libPrewarmTimer || _libPrewarmPromise) return;
|
||||
if (_libPrewarmPromise) return _libPrewarmPromise;
|
||||
const elapsed = Date.now() - _libLastPrewarmAt;
|
||||
if (elapsed >= 0 && elapsed < 5 * 60 * 1000) return;
|
||||
_libPrewarmTimer = setTimeout(() => {
|
||||
_libPrewarmTimer = null;
|
||||
_libPrewarmPromise = _prewarmEmailViews()
|
||||
.catch(() => {})
|
||||
.finally(() => { _libPrewarmPromise = null; });
|
||||
}, Math.max(0, Number(delay) || 0));
|
||||
if (elapsed >= 0 && elapsed < _LIB_PREWARM_COOLDOWN_MS) return Promise.resolve(false);
|
||||
return _scheduleEmailPrewarm(_prewarmEmailViews, { delay });
|
||||
}
|
||||
|
||||
async function _ensureEmailAccountsForPrewarm() {
|
||||
function _chooseEmailPrewarmAccountId(accounts) {
|
||||
const enabled = Array.isArray(accounts) ? accounts.filter(a => a && a.enabled !== false) : [];
|
||||
const remembered = _rememberedEmailAccountId();
|
||||
const current = String(state._libAccountId || '').trim();
|
||||
const chosen = enabled.find(a => String(a.id || '') === remembered)
|
||||
|| enabled.find(a => String(a.id || '') === current)
|
||||
|| enabled.find(a => a.is_default)
|
||||
|| enabled[0]
|
||||
|| null;
|
||||
return String(chosen?.id || '').trim();
|
||||
}
|
||||
|
||||
async function _ensureEmailAccountsForPrewarm({ signal, generation } = {}) {
|
||||
if (!_isEmailPrewarmCurrent(generation, signal)) return null;
|
||||
const accountsFresh = _libAccountsLoadedAt && (Date.now() - _libAccountsLoadedAt) < _LIB_ACCOUNTS_TTL_MS;
|
||||
if (Array.isArray(state._libAccounts) && state._libAccounts.length && accountsFresh) {
|
||||
if (!state._libAccountId) {
|
||||
const def = state._libAccounts.find(a => a.is_default) || state._libAccounts[0];
|
||||
state._libAccountId = def?.id || null;
|
||||
_publishActiveAccount();
|
||||
}
|
||||
return;
|
||||
}
|
||||
try {
|
||||
const accountsRes = await fetch(`${API_BASE}/api/email/accounts`, { credentials: 'same-origin' });
|
||||
if (!accountsRes.ok) return;
|
||||
const accountsData = await accountsRes.json().catch(() => ({}));
|
||||
if (Array.isArray(accountsData.accounts)) {
|
||||
state._libAccounts = accountsData.accounts;
|
||||
_libAccountsLoadedAt = Date.now();
|
||||
if (!state._libAccountId && state._libAccounts.length) {
|
||||
const def = state._libAccounts.find(a => a.is_default) || state._libAccounts[0];
|
||||
state._libAccountId = def?.id || null;
|
||||
_publishActiveAccount();
|
||||
if (!(Array.isArray(state._libAccounts) && state._libAccounts.length && accountsFresh)) {
|
||||
try {
|
||||
const accountsRes = await fetch(`${API_BASE}/api/email/accounts`, {
|
||||
credentials: 'same-origin',
|
||||
signal,
|
||||
});
|
||||
if (!_isEmailPrewarmCurrent(generation, signal)) return null;
|
||||
if (accountsRes.ok) {
|
||||
const accountsData = await accountsRes.json().catch(() => ({}));
|
||||
if (!_isEmailPrewarmCurrent(generation, signal)) return null;
|
||||
if (Array.isArray(accountsData.accounts)) {
|
||||
state._libAccounts = accountsData.accounts;
|
||||
_libAccountsLoadedAt = Date.now();
|
||||
}
|
||||
}
|
||||
} catch (err) {
|
||||
if (err?.name === 'AbortError') return null;
|
||||
}
|
||||
} catch (_) {}
|
||||
}
|
||||
|
||||
const accountId = _chooseEmailPrewarmAccountId(state._libAccounts);
|
||||
if (!_isEmailPrewarmCurrent(generation, signal)) return null;
|
||||
if (!accountId) return null;
|
||||
if (accountId && state._libAccountId !== accountId) {
|
||||
state._libAccountId = accountId;
|
||||
_publishActiveAccount();
|
||||
}
|
||||
return accountId;
|
||||
}
|
||||
|
||||
export async function prewarmUnreadEmails({ limit = 8, maxUid = 0 } = {}) {
|
||||
if (state._libOpen) return;
|
||||
await _ensureEmailAccountsForPrewarm();
|
||||
if (state._libOpen) return;
|
||||
const accountId = state._libAccountId || '';
|
||||
export function prewarmUnreadEmails({ limit = 8, maxUid = 0 } = {}) {
|
||||
return _scheduleEmailPrewarm(
|
||||
context => _prewarmUnreadEmailsNow({ limit, maxUid }, context),
|
||||
{ delay: 0 }
|
||||
);
|
||||
}
|
||||
|
||||
async function _prewarmUnreadEmailsNow({ limit = 8, maxUid = 0 } = {}, { signal, generation } = {}) {
|
||||
if (!_isEmailPrewarmCurrent(generation, signal)) return false;
|
||||
const accountId = await _ensureEmailAccountsForPrewarm({ signal, generation });
|
||||
if (accountId === null || !_isEmailPrewarmCurrent(generation, signal)) return false;
|
||||
const n = Math.max(1, Math.min(20, Number(limit) || 8));
|
||||
const key = `${accountId}|${maxUid || 0}|${n}`;
|
||||
if (_libUnreadPrewarmKey === key && (Date.now() - _libUnreadPrewarmAt) < 60 * 1000) return;
|
||||
_libUnreadPrewarmKey = key;
|
||||
_libUnreadPrewarmAt = Date.now();
|
||||
if (_libUnreadPrewarmKey === key && (Date.now() - _libUnreadPrewarmAt) < 60 * 1000) return true;
|
||||
try {
|
||||
const folder = 'INBOX';
|
||||
const res = await fetch(emailApiUrl('/api/email/list', {
|
||||
folder,
|
||||
limit: n,
|
||||
offset: 0,
|
||||
filter: 'unread',
|
||||
account_id: accountId || undefined,
|
||||
}), { credentials: 'same-origin' });
|
||||
if (state._libOpen) return;
|
||||
if (!res.ok) return;
|
||||
const res = await fetch(emailApiUrl('/api/email/list', {
|
||||
folder,
|
||||
limit: n,
|
||||
offset: 0,
|
||||
filter: 'unread',
|
||||
account_id: accountId || undefined,
|
||||
}), {
|
||||
credentials: 'same-origin',
|
||||
signal,
|
||||
});
|
||||
if (!_isEmailPrewarmCurrent(generation, signal) || !res.ok) return false;
|
||||
const data = await res.json().catch(() => null);
|
||||
if (!data || data.error || !Array.isArray(data.emails) || !data.emails.length) return;
|
||||
if (!_isEmailPrewarmCurrent(generation, signal)) return false;
|
||||
if (!data || data.error || !Array.isArray(data.emails) || !data.emails.length) return false;
|
||||
const sync = data.sync || {};
|
||||
_libCachePut(_libCacheKeyFor(accountId, folder, 'unread', false), {
|
||||
emails: data.emails,
|
||||
total: data.total || data.emails.length,
|
||||
sync,
|
||||
});
|
||||
} catch (_) {}
|
||||
_libUnreadPrewarmKey = key;
|
||||
_libUnreadPrewarmAt = Date.now();
|
||||
return true;
|
||||
} catch (_) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
function _sleep(ms) {
|
||||
return new Promise(resolve => setTimeout(resolve, ms));
|
||||
}
|
||||
|
||||
async function _prewarmEmailViews() {
|
||||
if (state._libOpen) return;
|
||||
_libLastPrewarmAt = Date.now();
|
||||
async function _prewarmEmailViews({ signal, generation } = {}) {
|
||||
if (!_isEmailPrewarmCurrent(generation, signal)) return false;
|
||||
_setEmailSyncStatus({ warming: true });
|
||||
const folder = 'INBOX';
|
||||
const filter = 'all';
|
||||
|
||||
// The accounts request is cheap and warms the account strip for first open.
|
||||
// Then folder/list requests warm both the client cache and the backend
|
||||
// IMAP/read caches. Failure stays silent: no configured mail should not nag.
|
||||
try {
|
||||
const accountsRes = await fetch(`${API_BASE}/api/email/accounts`, { credentials: 'same-origin' });
|
||||
if (accountsRes.ok) {
|
||||
const accountsData = await accountsRes.json().catch(() => ({}));
|
||||
if (Array.isArray(accountsData.accounts)) {
|
||||
state._libAccounts = accountsData.accounts;
|
||||
_libAccountsLoadedAt = Date.now();
|
||||
}
|
||||
const accountId = await _ensureEmailAccountsForPrewarm({ signal, generation });
|
||||
if (accountId === null || !_isEmailPrewarmCurrent(generation, signal)) return false;
|
||||
const ck = _libCacheKeyFor(accountId, folder, filter, false);
|
||||
if (_libCacheGet(ck)) {
|
||||
_libLastPrewarmAt = Date.now();
|
||||
return true;
|
||||
}
|
||||
} catch (_) {}
|
||||
|
||||
const accounts = Array.isArray(state._libAccounts) ? state._libAccounts.filter(a => a && a.enabled !== false) : [];
|
||||
const preferred = state._libAccountId
|
||||
|| (accounts.find(a => a.is_default)?.id)
|
||||
|| (accounts[0]?.id)
|
||||
|| '';
|
||||
if (!state._libAccountId && preferred) {
|
||||
state._libAccountId = preferred;
|
||||
_publishActiveAccount();
|
||||
}
|
||||
const orderedAccountIds = [
|
||||
preferred,
|
||||
...accounts.map(a => a.id).filter(id => id && id !== preferred),
|
||||
].filter((id, idx, arr) => arr.indexOf(id) === idx);
|
||||
if (!orderedAccountIds.length) orderedAccountIds.push('');
|
||||
|
||||
try {
|
||||
for (const accountId of orderedAccountIds.slice(0, 4)) {
|
||||
if (state._libOpen) return;
|
||||
const ck = _libCacheKeyFor(accountId, folder, filter, false);
|
||||
if (_libCacheGet(ck)) continue;
|
||||
await fetch(emailApiUrl('/api/email/folders', { account_id: accountId || undefined }), { credentials: 'same-origin' }).catch(() => null);
|
||||
await fetch(emailApiUrl('/api/email/unread-state', { folder, account_id: accountId || undefined }), { credentials: 'same-origin' }).catch(() => null);
|
||||
const res = await fetch(emailApiUrl('/api/email/list', {
|
||||
folder,
|
||||
limit: 100,
|
||||
offset: 0,
|
||||
filter,
|
||||
account_id: accountId || undefined,
|
||||
}), {
|
||||
credentials: 'same-origin',
|
||||
});
|
||||
if (res.ok) {
|
||||
const data = await res.json().catch(() => null);
|
||||
if (data && !data.error) {
|
||||
const sync = data.sync || {};
|
||||
_libCachePut(ck, {
|
||||
emails: data.emails || [],
|
||||
total: data.total || 0,
|
||||
sync,
|
||||
});
|
||||
_setEmailSyncStatus({
|
||||
updatedAt: sync.updated_at || new Date().toISOString(),
|
||||
source: sync.source || '',
|
||||
warming: true,
|
||||
});
|
||||
}
|
||||
}
|
||||
await _sleep(900);
|
||||
}
|
||||
// One optional first-page request only. Folder metadata, unread state, and
|
||||
// other accounts remain demand-driven so startup cannot fan out into IMAP.
|
||||
const res = await fetch(emailApiUrl('/api/email/list', {
|
||||
folder,
|
||||
limit: _LIB_INITIAL_PAGE_SIZE,
|
||||
offset: 0,
|
||||
filter,
|
||||
account_id: accountId || undefined,
|
||||
}), {
|
||||
credentials: 'same-origin',
|
||||
signal,
|
||||
});
|
||||
if (!_isEmailPrewarmCurrent(generation, signal) || !res.ok) return false;
|
||||
const data = await res.json().catch(() => null);
|
||||
if (!_isEmailPrewarmCurrent(generation, signal)) return false;
|
||||
if (!data || data.error || !Array.isArray(data.emails)) return false;
|
||||
const sync = data.sync || {};
|
||||
_libCachePut(ck, {
|
||||
emails: data.emails,
|
||||
total: data.total || 0,
|
||||
sync,
|
||||
});
|
||||
_libLastPrewarmAt = Date.now();
|
||||
_setEmailSyncStatus({
|
||||
updatedAt: sync.updated_at || new Date().toISOString(),
|
||||
source: sync.source || '',
|
||||
warming: true,
|
||||
});
|
||||
return true;
|
||||
} catch (_) {
|
||||
return false;
|
||||
} finally {
|
||||
_setEmailSyncStatus({ warming: false });
|
||||
}
|
||||
@@ -2286,16 +2482,34 @@ function _publishActiveAccount() {
|
||||
|
||||
export function initEmailLibrary(config) {
|
||||
state._docModule = config.documentModule;
|
||||
state._onEmailClick = config.onEmailClick;
|
||||
const onEmailClick = config.onEmailClick;
|
||||
state._onEmailClick = typeof onEmailClick === 'function' ? (options = {}) => {
|
||||
const accountId = String(state._libAccountId || '');
|
||||
const libraryFolder = String(state._libFolder || 'INBOX');
|
||||
const messageFolder = String(options.email?.folder || libraryFolder);
|
||||
const mailboxGeneration = _emailMailboxGeneration;
|
||||
const mailboxContext = Object.freeze({
|
||||
accountId,
|
||||
libraryFolder,
|
||||
messageFolder,
|
||||
mailboxGeneration,
|
||||
isCurrent: () => (
|
||||
String(state._libAccountId || '') === accountId &&
|
||||
String(state._libFolder || 'INBOX') === libraryFolder &&
|
||||
_emailMailboxGeneration === mailboxGeneration
|
||||
),
|
||||
});
|
||||
return onEmailClick({ ...options, mailboxContext });
|
||||
} : null;
|
||||
}
|
||||
|
||||
export function isOpen() { return state._libOpen; }
|
||||
|
||||
export function openEmailLibrary(opts = {}) {
|
||||
if (_libPrewarmTimer) {
|
||||
clearTimeout(_libPrewarmTimer);
|
||||
_libPrewarmTimer = null;
|
||||
}
|
||||
// Foreground email always wins: cancel a delayed/idle callback and abort the
|
||||
// one optional request if it has already started. Generation checks make a
|
||||
// non-abortable response harmless if it races this transition.
|
||||
_cancelEmailPrewarm();
|
||||
// Force-clean any stale state from previous attempts
|
||||
const existing = document.getElementById('email-lib-modal');
|
||||
if (existing) existing.remove();
|
||||
@@ -2303,6 +2517,7 @@ export function openEmailLibrary(opts = {}) {
|
||||
document.removeEventListener('keydown', state._libEscHandler, true);
|
||||
state._libEscHandler = null;
|
||||
}
|
||||
_emailMailboxGeneration += 1;
|
||||
state._libOpen = true;
|
||||
// On mobile the sidebar overlays content — close it so the email view isn't
|
||||
// opened behind it (same pattern as session-switch/delete).
|
||||
@@ -2926,7 +3141,7 @@ export function openEmailLibrary(opts = {}) {
|
||||
}
|
||||
const fastAccountAtOpen = state._libAccountId || '';
|
||||
if (fastAccountAtOpen) {
|
||||
_loadEmailsWhenChatIdle({ delay: 0 });
|
||||
_loadEmails({ useCache: true });
|
||||
}
|
||||
// If we already know the previous/default account, paint that inbox first
|
||||
// from the durable index and validate accounts in parallel. Cold refreshes
|
||||
@@ -2936,7 +3151,7 @@ export function openEmailLibrary(opts = {}) {
|
||||
_loadFolders();
|
||||
_loadEmailReminderBellVisibility();
|
||||
if (!fastAccountAtOpen || fastAccountAtOpen !== (state._libAccountId || '')) {
|
||||
_loadEmailsWhenChatIdle();
|
||||
_loadEmails({ useCache: true });
|
||||
}
|
||||
})();
|
||||
}
|
||||
@@ -3121,6 +3336,7 @@ export async function openEmailLibrarySettings() {
|
||||
}
|
||||
|
||||
export function closeEmailLibrary() {
|
||||
_cancelEmailPrewarm();
|
||||
const modal = document.getElementById('email-lib-modal');
|
||||
if (modal) modal.remove();
|
||||
if (_libSyncTicker) {
|
||||
@@ -4554,7 +4770,7 @@ async function _loadEmails({ force = false, useCache = true } = {}) {
|
||||
const ctrl = new AbortController();
|
||||
const timer = setTimeout(() => ctrl.abort(), 450);
|
||||
try {
|
||||
const fastRes = await fetch(`${API_BASE}/api/email/list?folder=${encodeURIComponent(folderAtStart)}${accountQS}&limit=100&offset=${offsetAtStart}&filter=${filterAtStart}${attQS}&cached_only=1`, {
|
||||
const fastRes = await fetch(`${API_BASE}/api/email/list?folder=${encodeURIComponent(folderAtStart)}${accountQS}&limit=${_LIB_INITIAL_PAGE_SIZE}&offset=${offsetAtStart}&filter=${filterAtStart}${attQS}&cached_only=1`, {
|
||||
signal: ctrl.signal,
|
||||
});
|
||||
const fastData = await fastRes.json().catch(() => null);
|
||||
@@ -4581,7 +4797,7 @@ async function _loadEmails({ force = false, useCache = true } = {}) {
|
||||
// opens omit it so rapid close/reopen returns instantly; the
|
||||
// Refresh button passes `force: true` to add it back.
|
||||
const buster = force ? `&_=${Date.now()}` : '';
|
||||
const res = await fetch(`${API_BASE}/api/email/list?folder=${encodeURIComponent(folderAtStart)}${accountQS}&limit=100&offset=${offsetAtStart}&filter=${filterAtStart}${attQS}${buster}`);
|
||||
const res = await fetch(`${API_BASE}/api/email/list?folder=${encodeURIComponent(folderAtStart)}${accountQS}&limit=${_LIB_INITIAL_PAGE_SIZE}&offset=${offsetAtStart}&filter=${filterAtStart}${attQS}${buster}`);
|
||||
const data = await res.json();
|
||||
if (seq !== _libLoadSeq || accountAtStart !== (state._libAccountId || '')) return;
|
||||
if (data.error) throw new Error(data.error);
|
||||
@@ -4836,6 +5052,8 @@ function _createCard(em) {
|
||||
else if (!em.is_read) cls += ' email-card-unread';
|
||||
card.className = cls;
|
||||
card.dataset.uid = String(em.uid);
|
||||
card.dataset.emailAccount = String(em.account_id || state._libAccountId || '');
|
||||
card.dataset.emailFolder = String(em.folder || state._libFolder || 'INBOX');
|
||||
if (state._selectMode && state._selectedUids.has(em.uid)) card.classList.add('selected');
|
||||
|
||||
// Checkbox in select mode
|
||||
@@ -5162,6 +5380,25 @@ async function _toggleCardPreview(card, em) {
|
||||
// currently-selected folder for normal inbox cards.
|
||||
const folderAtStart = (em && em.folder) || libraryFolderAtStart;
|
||||
const uidAtStart = String(em?.uid || card?.dataset?.uid || '');
|
||||
const wasReadAtStart = !!em?.is_read;
|
||||
const openGeneration = ++_emailCardOpenSeq;
|
||||
const readContext = Object.freeze({
|
||||
accountId: String(accountAtStart),
|
||||
libraryFolder: String(libraryFolderAtStart),
|
||||
folder: String(folderAtStart),
|
||||
uid: uidAtStart,
|
||||
mailboxGeneration: _emailMailboxGeneration,
|
||||
});
|
||||
const readContextKey = _emailReadContextKey(readContext);
|
||||
const isCurrentOpen = () => (
|
||||
openGeneration === _emailCardOpenSeq &&
|
||||
_emailReadContextIsCurrent(readContext) &&
|
||||
accountAtStart === (state._libAccountId || '') &&
|
||||
libraryFolderAtStart === (state._libFolder || 'INBOX') &&
|
||||
uidAtStart === String(card?.dataset?.uid || '') &&
|
||||
card.isConnected &&
|
||||
card.classList.contains('email-card-expanded')
|
||||
);
|
||||
const grid = card.closest('.doclib-grid');
|
||||
const gridRect = grid?.getBoundingClientRect?.();
|
||||
const modal = document.getElementById('email-lib-modal');
|
||||
@@ -5186,6 +5423,30 @@ async function _toggleCardPreview(card, em) {
|
||||
return;
|
||||
}
|
||||
|
||||
// Every authoritative open supersedes any older optimistic mutation for the
|
||||
// same immutable mailbox identity. Carry the original unread state forward
|
||||
// so a close/reopen followed by failure still rolls back exactly once, while
|
||||
// a late failure from the superseded request cannot undo a newer success.
|
||||
const previousMutation = _emailReadMutations.get(readContextKey);
|
||||
const readMutation = {
|
||||
generation: ++_emailReadMutationSeq,
|
||||
rollbackUnread: !wasReadAtStart || !!previousMutation?.rollbackUnread,
|
||||
};
|
||||
_emailReadMutations.set(readContextKey, readMutation);
|
||||
const restoreUnreadState = () => {
|
||||
if (_emailReadMutations.get(readContextKey)?.generation !== readMutation.generation) return;
|
||||
_emailReadMutations.delete(readContextKey);
|
||||
if (readMutation.rollbackUnread) _syncEmailReadState(uidAtStart, false, readContext);
|
||||
};
|
||||
const commitReadState = () => {
|
||||
// A successful STORE/mark_seen is authoritative for this immutable
|
||||
// mailbox identity even when a newer open is still pending. Retire that
|
||||
// newer rollback token too, otherwise its later failure could restore an
|
||||
// unread state that no longer exists at the provider.
|
||||
_emailReadMutations.delete(readContextKey);
|
||||
_syncEmailReadState(uidAtStart, true, readContext);
|
||||
};
|
||||
|
||||
// Collapse any other expanded card
|
||||
if (grid) {
|
||||
grid.querySelectorAll('.email-card-expanded').forEach(c => {
|
||||
@@ -5207,10 +5468,10 @@ async function _toggleCardPreview(card, em) {
|
||||
requestAnimationFrame(() => {
|
||||
try { card.scrollIntoView({ behavior: 'smooth', block: 'start' }); } catch (_) {}
|
||||
});
|
||||
if (!em.is_read) {
|
||||
_syncEmailReadState(em.uid, true);
|
||||
fetch(`${API_BASE}/api/email/mark-read/${em.uid}?folder=${encodeURIComponent(folderAtStart)}${_acct()}`, { method: 'POST' })
|
||||
.catch(err => console.error('Failed to mark email read:', err));
|
||||
if (!wasReadAtStart) {
|
||||
// Keep the current optimistic visual update, but let the read request below
|
||||
// own the provider-side \Seen transition. A failure restores unread state.
|
||||
_syncEmailReadState(uidAtStart, true, readContext);
|
||||
}
|
||||
// Class hook on the modal so the header-hide / padding rules work on
|
||||
// browsers without :has() support (Firefox mobile) — the :has() versions
|
||||
@@ -5239,25 +5500,28 @@ async function _toggleCardPreview(card, em) {
|
||||
} catch (_) {}
|
||||
};
|
||||
|
||||
let authoritativeReadSucceeded = false;
|
||||
try {
|
||||
const res = await fetch(`${API_BASE}/api/email/read/${em.uid}?folder=${encodeURIComponent(folderAtStart)}${_acct()}`);
|
||||
const accountQueryAtStart = accountAtStart ? `&account_id=${encodeURIComponent(accountAtStart)}` : '';
|
||||
const res = await fetch(`${API_BASE}/api/email/read/${encodeURIComponent(uidAtStart)}?folder=${encodeURIComponent(folderAtStart)}${accountQueryAtStart}&mark_seen=true`);
|
||||
if (!res.ok) throw new Error(`HTTP ${res.status}`);
|
||||
const data = await res.json();
|
||||
if (
|
||||
accountAtStart !== (state._libAccountId || '') ||
|
||||
libraryFolderAtStart !== (state._libFolder || 'INBOX') ||
|
||||
uidAtStart !== String(card?.dataset?.uid || '') ||
|
||||
!card.isConnected ||
|
||||
!card.classList.contains('email-card-expanded')
|
||||
) {
|
||||
return;
|
||||
}
|
||||
if (data.error) {
|
||||
showFailedReader(`Failed to load email: ${data.error}`);
|
||||
restoreUnreadState();
|
||||
if (isCurrentOpen()) showFailedReader(`Failed to load email: ${data.error}`);
|
||||
return;
|
||||
}
|
||||
// Mark as read locally
|
||||
_syncEmailReadState(em.uid, true);
|
||||
if (data.mark_seen_failed) {
|
||||
// The body is authoritative even when the provider refused the \Seen
|
||||
// transition. Render the message and roll the unread marker back so the
|
||||
// list keeps telling the truth, rather than refusing to open a message
|
||||
// we successfully read.
|
||||
restoreUnreadState();
|
||||
} else {
|
||||
authoritativeReadSucceeded = true;
|
||||
commitReadState();
|
||||
}
|
||||
if (!isCurrentOpen()) return;
|
||||
_stampReaderContext(reader, { ...em, ...data }, state._libFolder, state._libAccountId);
|
||||
|
||||
// Build the attachments wrap using the shared helper so the signature-
|
||||
@@ -5439,7 +5703,10 @@ async function _toggleCardPreview(card, em) {
|
||||
// Always stop bubbling so the card's click doesn't fire while reading.
|
||||
reader.addEventListener('click', (ev) => { ev.stopPropagation(); });
|
||||
} catch (e) {
|
||||
showFailedReader(e?.message ? `Failed to load email: ${e.message}` : 'Failed to load email');
|
||||
if (!authoritativeReadSucceeded) restoreUnreadState();
|
||||
if (isCurrentOpen()) {
|
||||
showFailedReader(e?.message ? `Failed to load email: ${e.message}` : 'Failed to load email');
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7259,12 +7526,11 @@ async function _generateSummary(reader, data, btn) {
|
||||
if (label) label.textContent = 'Summary';
|
||||
}
|
||||
} else {
|
||||
content.innerHTML = `<span style="color:var(--red)">${_esc(result.error || 'Failed to summarize')}</span>`;
|
||||
panel.remove();
|
||||
_renderEmailSummaryError(content, result);
|
||||
}
|
||||
} catch (e) {
|
||||
sp.destroy();
|
||||
panel.remove();
|
||||
_renderEmailSummaryError(content, null);
|
||||
if (uiModule) uiModule.showError?.('Failed to summarize');
|
||||
} finally {
|
||||
if (btn) btn.disabled = false;
|
||||
|
||||
@@ -30,6 +30,25 @@ export function _esc(text) {
|
||||
return div.innerHTML;
|
||||
}
|
||||
|
||||
const _EMAIL_SUMMARY_ERROR_MESSAGES = Object.freeze({
|
||||
email_summary_missing_body: 'No email body to summarize',
|
||||
email_summary_not_configured: 'No model configured for email summaries',
|
||||
email_summary_empty: 'The model returned an empty summary',
|
||||
email_summary_unavailable: 'Failed to summarize',
|
||||
});
|
||||
|
||||
export function _emailSummaryErrorMessage(result) {
|
||||
const code = String(result?.error_code || '');
|
||||
return _EMAIL_SUMMARY_ERROR_MESSAGES[code] || 'Failed to summarize';
|
||||
}
|
||||
|
||||
export function _renderEmailSummaryError(container, result) {
|
||||
const message = container.ownerDocument.createElement('span');
|
||||
message.style.color = 'var(--red)';
|
||||
message.textContent = _emailSummaryErrorMessage(result);
|
||||
container.replaceChildren(message);
|
||||
}
|
||||
|
||||
function _attrEsc(text) {
|
||||
return String(text ?? '')
|
||||
.replace(/"/g, '"')
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
// liveThinkingThrottle.js
|
||||
//
|
||||
// Pure trailing-edge coalescer for the live "thinking" block in chat.js.
|
||||
//
|
||||
// A reasoning stream delivers deltas far faster than a human can read them, and
|
||||
// the only thing that matters on screen is the LATEST cumulative text. Committing
|
||||
// every delta to the DOM makes the work grow with the length of the stream. This
|
||||
// throttle collapses a burst of updates into one commit per `delay` ms, always
|
||||
// carrying the most recent value.
|
||||
//
|
||||
// Timers are injected so the behaviour is testable without a browser or a clock:
|
||||
//
|
||||
// const throttle = createLiveThinkingThrottle(commit, { prepare, schedule, cancel });
|
||||
//
|
||||
// Lifecycle contract, which the terminal paths in chat.js depend on:
|
||||
//
|
||||
// update(value) queue `value`; schedule a commit if one is not already pending
|
||||
// flush() commit any pending value NOW and drop the timer; returns whether
|
||||
// a commit happened, so a clean flush cannot duplicate a commit
|
||||
// cancel() drop the timer AND the pending value — nothing lands later
|
||||
//
|
||||
// `cancel()` is what stops a finished (or backgrounded) stream from mutating a
|
||||
// view the user has since navigated away to.
|
||||
|
||||
export function stripLiveThinkingTags(text) {
|
||||
return String(text ?? '').replace(
|
||||
/<\/?(?:think(?:ing)?|thought)(?:\s+[^>]*)?>/gi,
|
||||
'',
|
||||
);
|
||||
}
|
||||
|
||||
const THINKING_BOUNDARY_RE = /<\/?(?:(?:mm:)?think(?:ing)?|thought)(?:\s+[^>]*)?>|<\|channel>(?:thought|response)|<channel\|>/gi;
|
||||
const REPLY_PREFIX_SOURCE = "(?:Hey|Hi |Hi!|Hello|Sure|Yes|No |No,|Yo|OK|Here|Absolutely|Of course|Great|Alright|Thanks|Welcome|Good |I'm happy|I'd be)";
|
||||
const REPLY_LINE_RE = new RegExp('(?:^|\\n)\\s*' + REPLY_PREFIX_SOURCE, 'gi');
|
||||
const REPLY_INLINE_RE = new RegExp('[.!?]\\s*' + REPLY_PREFIX_SOURCE, 'gi');
|
||||
const REASONING_PREFIX_CANDIDATES = [
|
||||
'thinking:', 'thinking process:', 'the user ', 'user wants', 'we need ',
|
||||
'i need ', 'i should ', 'i will ', "i'll ", 'i am going ', 'let me think',
|
||||
'let me look', 'let me see', 'let me check', 'let me read', 'let me review',
|
||||
'let me analyze', 'let me parse', 'let me figure', 'let me draft', 'let me write',
|
||||
'they are ', 'the question ', 'i can ',
|
||||
];
|
||||
|
||||
const DISPLAY_FILTER_BOUNDARY_RE = /\[\/?TOOL_CALL\]|```(?:create_document|documen(?:t)?)(?:\s|$)|```[\w-]+[ \t]*[\[{]|<(?:[\w]+:)?(?:tool_call|function_call)>|<invoke\b|<\s*\/?\s*[||]+\s*DSML\s*[||]+|(?:\[\s*)?\{\s*"function"\s*:|<\/?\|(?:assistant|assistan|user|system|tool|end)\|?>|(?:^|[\r\n])\s*(?:stdout|stderr|exit_code):/i;
|
||||
|
||||
function hasFreshMatch(text, regex, cursor, minStart = 0) {
|
||||
regex.lastIndex = 0;
|
||||
for (const match of text.matchAll(regex)) {
|
||||
const end = match.index + match[0].length;
|
||||
if (end > cursor && match.index >= minStart) return true;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
// Incrementally decides when chat.js needs its compatibility-heavy cumulative
|
||||
// thinking analysis. The gate inspects only a short overlap plus the new text;
|
||||
// ordinary answer/reasoning deltas therefore stay O(delta) while split tags,
|
||||
// namespaced tags, non-tag reply boundaries, and false-close grace deadlines
|
||||
// still request the canonical full analysis.
|
||||
export function createThinkingAnalysisGate({
|
||||
startsWithReasoningPrefix = () => false,
|
||||
now = () => Date.now(),
|
||||
overlap = 512,
|
||||
} = {}) {
|
||||
let cursor = 0;
|
||||
let prefixSettled = false;
|
||||
let prefixProbe = '';
|
||||
|
||||
return {
|
||||
shouldAnalyze(text, {
|
||||
isThinking = false,
|
||||
nonTagThinking = false,
|
||||
recheckAt = 0,
|
||||
} = {}) {
|
||||
const fullText = String(text ?? '');
|
||||
if (fullText.length < cursor) {
|
||||
cursor = 0;
|
||||
prefixSettled = false;
|
||||
prefixProbe = '';
|
||||
}
|
||||
const previousCursor = cursor;
|
||||
if (!prefixSettled && prefixProbe.length < overlap) {
|
||||
// Build the initial probe from deltas so arbitrary leading whitespace
|
||||
// cannot strand the gate in its undecided state. The retained state is
|
||||
// bounded even if a provider emits a very large whitespace prefix.
|
||||
prefixProbe = (prefixProbe + fullText.slice(previousCursor))
|
||||
.trimStart()
|
||||
.slice(0, overlap);
|
||||
}
|
||||
const scanStart = Math.max(0, previousCursor - overlap);
|
||||
const freshText = fullText.slice(scanStart);
|
||||
const relativeCursor = previousCursor - scanStart;
|
||||
const hasBoundary = hasFreshMatch(freshText, THINKING_BOUNDARY_RE, relativeCursor);
|
||||
const hasReplyBoundary = nonTagThinking && (
|
||||
hasFreshMatch(freshText, REPLY_LINE_RE, relativeCursor)
|
||||
|| hasFreshMatch(freshText, REPLY_INLINE_RE, relativeCursor, Math.max(0, 20 - scanStart))
|
||||
);
|
||||
cursor = fullText.length;
|
||||
|
||||
if (hasBoundary || hasReplyBoundary) return true;
|
||||
if (isThinking) return recheckAt > 0 && now() >= recheckAt;
|
||||
if (prefixSettled) return false;
|
||||
|
||||
if (!prefixProbe) return false;
|
||||
if (startsWithReasoningPrefix(prefixProbe)) {
|
||||
prefixSettled = true;
|
||||
return true;
|
||||
}
|
||||
const lowerProbe = prefixProbe.toLowerCase();
|
||||
if (REASONING_PREFIX_CANDIDATES.some((candidate) => candidate.startsWith(lowerProbe))) {
|
||||
return false;
|
||||
}
|
||||
prefixSettled = true;
|
||||
return false;
|
||||
},
|
||||
reset() {
|
||||
cursor = 0;
|
||||
prefixSettled = false;
|
||||
prefixProbe = '';
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
// Keep the common prose path append-only. At the first structured/tool
|
||||
// boundary, filter only the preceding visible prefix and hide the structured
|
||||
// tail until the authoritative terminal render.
|
||||
export function createIncrementalDisplayProjector(filter, { overlap = 512 } = {}) {
|
||||
let projected = '';
|
||||
let boundaryTail = '';
|
||||
let rawLength = 0;
|
||||
let structuredTailHidden = false;
|
||||
|
||||
return {
|
||||
append(delta, fullText) {
|
||||
const chunk = String(delta ?? '');
|
||||
const raw = String(fullText ?? '');
|
||||
if (raw.length < rawLength) this.reset();
|
||||
const boundaryProbe = boundaryTail + chunk;
|
||||
const boundaryMatch = !structuredTailHidden
|
||||
? DISPLAY_FILTER_BOUNDARY_RE.exec(boundaryProbe)
|
||||
: null;
|
||||
if (boundaryMatch) {
|
||||
// Filter the visible prefix, not the incomplete marker itself: several
|
||||
// compatibility regexes intentionally match only completed blocks.
|
||||
const boundaryStart = Math.max(0, raw.length - boundaryProbe.length + boundaryMatch.index);
|
||||
structuredTailHidden = true;
|
||||
projected = String(filter(raw.slice(0, boundaryStart)) ?? '');
|
||||
} else if (!structuredTailHidden) {
|
||||
projected += chunk;
|
||||
}
|
||||
boundaryTail = (boundaryTail + chunk).slice(-overlap);
|
||||
rawLength = raw.length;
|
||||
return projected;
|
||||
},
|
||||
current() {
|
||||
return projected;
|
||||
},
|
||||
reset() {
|
||||
projected = '';
|
||||
boundaryTail = '';
|
||||
rawLength = 0;
|
||||
structuredTailHidden = false;
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
export function createLiveThinkingThrottle(commit, {
|
||||
delay = 100,
|
||||
prepare = (value) => String(value ?? ''),
|
||||
schedule = (callback, ms) => setTimeout(callback, ms),
|
||||
cancel = (timer) => clearTimeout(timer),
|
||||
} = {}) {
|
||||
let timer = null;
|
||||
let latest = null;
|
||||
let dirty = false;
|
||||
|
||||
const commitLatest = () => {
|
||||
timer = null;
|
||||
if (!dirty) return false;
|
||||
dirty = false;
|
||||
commit(prepare(latest));
|
||||
return true;
|
||||
};
|
||||
|
||||
return {
|
||||
update(value) {
|
||||
latest = value;
|
||||
dirty = true;
|
||||
if (timer === null) timer = schedule(commitLatest, delay);
|
||||
},
|
||||
flush() {
|
||||
if (timer !== null) {
|
||||
cancel(timer);
|
||||
timer = null;
|
||||
}
|
||||
return commitLatest();
|
||||
},
|
||||
cancel() {
|
||||
if (timer !== null) cancel(timer);
|
||||
timer = null;
|
||||
dirty = false;
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
export default createLiveThinkingThrottle;
|
||||
+25
-1
@@ -1683,8 +1683,21 @@ export async function loadSessions() {
|
||||
url += `?active_incognito_id=${encodeURIComponent(currentSessionId)}`;
|
||||
}
|
||||
const res = await fetch(url);
|
||||
if (!res.ok) {
|
||||
let detail = '';
|
||||
try {
|
||||
const payload = await res.json();
|
||||
detail = payload?.detail || payload?.error || '';
|
||||
} catch (_) {}
|
||||
const error = new Error(detail || `Session request failed (HTTP ${res.status})`);
|
||||
error.status = res.status;
|
||||
throw error;
|
||||
}
|
||||
fetched = await res.json();
|
||||
}
|
||||
if (!Array.isArray(fetched)) {
|
||||
throw new Error('Session request returned an invalid response');
|
||||
}
|
||||
sessions = _normalizeSessionsList(fetched);
|
||||
renderSessionList();
|
||||
|
||||
@@ -1807,9 +1820,15 @@ export async function loadSessions() {
|
||||
_autoCreateInProgress = false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
} catch (error) {
|
||||
console.error('Error in loadSessions:', error);
|
||||
uiModule.showError('Failed to load sessions: ' + error.message);
|
||||
// app.js's global fetch wrapper owns expired-auth navigation. Avoid
|
||||
// flashing a redundant session error while that 401 redirect is pending.
|
||||
if (error?.status !== 401) {
|
||||
uiModule.showError('Failed to load sessions: ' + error.message);
|
||||
}
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1847,6 +1866,10 @@ export async function selectSession(id, { keepSidebar = false, showLoading = tru
|
||||
const _isTransientChat = !!_meta && (_meta.folder === 'Assistant' || _meta.folder === 'Tasks');
|
||||
if (!_isTransientChat) {
|
||||
Storage.set('lastSessionId', id);
|
||||
// Update URL hash without triggering hashchange handler
|
||||
if (window.location.hash !== '#' + id) {
|
||||
history.replaceState(null, '', '#' + id);
|
||||
}
|
||||
}
|
||||
// Restore character preset for persistent chats
|
||||
try {
|
||||
@@ -2313,6 +2336,7 @@ export async function materializePendingSession() {
|
||||
currentSessionId = payload.id;
|
||||
if (!isIncognito) {
|
||||
Storage.set('lastSessionId', payload.id);
|
||||
history.replaceState(null, '', '#' + payload.id);
|
||||
}
|
||||
|
||||
// Reload the sidebar in the background. Awaiting this used to block the first
|
||||
|
||||
+1
-81
@@ -445,14 +445,7 @@ async function initDefaultChat() {
|
||||
var epSel = el('set-defaultEpSelect');
|
||||
var modelSel = el('set-defaultModelSelect');
|
||||
var msg = el('set-defaultChatMsg');
|
||||
var fbContainer = el('set-defaultFallbacks');
|
||||
var addFbBtn = el('set-defaultAddFallback');
|
||||
var _endpoints = [];
|
||||
var _fallbacks = []; // [{endpoint_id, model}] — tried in order if primary fails
|
||||
|
||||
function enabledEndpoints() {
|
||||
return _endpoints.filter(function(e) { return e.is_enabled; });
|
||||
}
|
||||
|
||||
// Fill any <select> with the models for a given endpoint id.
|
||||
function fillModels(selectEl, epId, selected) {
|
||||
@@ -469,64 +462,6 @@ async function initDefaultChat() {
|
||||
function refreshEndpointOptions(selectedEndpoint, selectedModel) {
|
||||
_fillEndpointSelect(epSel, _endpoints, selectedEndpoint !== undefined ? selectedEndpoint : epSel.value, false);
|
||||
refreshModels(selectedModel !== undefined ? selectedModel : modelSel.value);
|
||||
renderFallbacks();
|
||||
}
|
||||
|
||||
// Render the fallback chain. Each row is endpoint + model + remove.
|
||||
function renderFallbacks() {
|
||||
fbContainer.innerHTML = '';
|
||||
_fallbacks.forEach(function(fb, idx) {
|
||||
var row = document.createElement('div');
|
||||
row.className = 'settings-fallback-row';
|
||||
|
||||
var num = document.createElement('span');
|
||||
num.className = 'settings-fallback-num';
|
||||
num.textContent = (idx + 1) + '.';
|
||||
|
||||
var epS = document.createElement('select');
|
||||
epS.className = 'settings-select';
|
||||
enabledEndpoints().forEach(function(ep) {
|
||||
var o = document.createElement('option');
|
||||
o.value = ep.id;
|
||||
o.textContent = ep.name + (ep.online ? '' : ' (offline)');
|
||||
epS.appendChild(o);
|
||||
});
|
||||
var first = enabledEndpoints()[0];
|
||||
epS.value = fb.endpoint_id || (first ? first.id : '');
|
||||
|
||||
var mS = document.createElement('select');
|
||||
mS.className = 'settings-select';
|
||||
fillModels(mS, epS.value, fb.model);
|
||||
|
||||
// Keep the model in sync with the values actually shown.
|
||||
fb.endpoint_id = epS.value;
|
||||
fb.model = mS.value;
|
||||
|
||||
epS.addEventListener('change', function() {
|
||||
fb.endpoint_id = epS.value;
|
||||
fillModels(mS, epS.value, '');
|
||||
fb.model = mS.value;
|
||||
saveDefault();
|
||||
});
|
||||
mS.addEventListener('change', function() { fb.model = mS.value; saveDefault(); });
|
||||
|
||||
var rm = document.createElement('button');
|
||||
rm.type = 'button';
|
||||
rm.className = 'settings-fallback-remove';
|
||||
rm.title = 'Remove fallback';
|
||||
rm.innerHTML = '<svg width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><polyline points="3 6 5 6 21 6"/><path d="M19 6l-1 14a2 2 0 0 1-2 2H8a2 2 0 0 1-2-2L5 6"/><path d="M10 11v6"/><path d="M14 11v6"/><path d="M9 6V4a1 1 0 0 1 1-1h4a1 1 0 0 1 1 1v2"/></svg>';
|
||||
rm.addEventListener('click', function() {
|
||||
_fallbacks.splice(idx, 1);
|
||||
renderFallbacks();
|
||||
saveDefault();
|
||||
});
|
||||
|
||||
row.appendChild(num);
|
||||
row.appendChild(epS);
|
||||
row.appendChild(mS);
|
||||
row.appendChild(rm);
|
||||
fbContainer.appendChild(row);
|
||||
});
|
||||
}
|
||||
|
||||
try {
|
||||
@@ -534,12 +469,6 @@ async function initDefaultChat() {
|
||||
var settings = await res.json();
|
||||
if (settings.default_endpoint_id) epSel.value = settings.default_endpoint_id;
|
||||
refreshModels(settings.default_model || '');
|
||||
_fallbacks = Array.isArray(settings.default_model_fallbacks)
|
||||
? settings.default_model_fallbacks.map(function(f) {
|
||||
return { endpoint_id: (f && f.endpoint_id) || '', model: (f && f.model) || '' };
|
||||
})
|
||||
: [];
|
||||
renderFallbacks();
|
||||
} catch (e) { console.warn('Failed to load default chat settings', e); }
|
||||
|
||||
epSel.addEventListener('change', function() { refreshModels(''); saveDefault(); });
|
||||
@@ -547,13 +476,11 @@ async function initDefaultChat() {
|
||||
|
||||
async function saveDefault() {
|
||||
try {
|
||||
var clean = _fallbacks.filter(function(f) { return f.endpoint_id && f.model; });
|
||||
await fetch('/api/auth/settings', { method: 'POST', credentials: 'same-origin',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({
|
||||
default_endpoint_id: epSel.value,
|
||||
default_model: modelSel.value,
|
||||
default_model_fallbacks: clean
|
||||
default_model: modelSel.value
|
||||
})
|
||||
});
|
||||
msg.textContent = 'Saved'; msg.style.color = 'var(--fg)';
|
||||
@@ -561,13 +488,6 @@ async function initDefaultChat() {
|
||||
} catch (e) { msg.textContent = 'Failed to save'; msg.style.color = 'var(--red)'; }
|
||||
}
|
||||
|
||||
if (addFbBtn) addFbBtn.addEventListener('click', function() {
|
||||
var first = enabledEndpoints()[0];
|
||||
_fallbacks.push({ endpoint_id: first ? first.id : '', model: '' });
|
||||
renderFallbacks();
|
||||
saveDefault();
|
||||
});
|
||||
|
||||
_registerAiEndpointRefresh(function(endpoints) {
|
||||
_endpoints = endpoints;
|
||||
refreshEndpointOptions(epSel.value, modelSel.value);
|
||||
|
||||
+3
-5
@@ -83,11 +83,9 @@ export async function loadSkills(cascade = false) {
|
||||
// Play the domino-in entrance on this load (set when the tab is opened,
|
||||
// not for the silent re-loads after an edit/delete).
|
||||
if (cascade) _cascadeNext = true;
|
||||
if (cascade && loaded && !_loadPromise && _playSkillsCascade()) {
|
||||
_cascadeNext = false;
|
||||
updateCount();
|
||||
return;
|
||||
}
|
||||
// Always re-fetch when the tab is explicitly opened — the cascade
|
||||
// animation is handled inside renderSkillsList() via _cascadeNext.
|
||||
// Skipping the fetch here caused stale data on panel close/reopen (#5870).
|
||||
if (_loadPromise) return _loadPromise;
|
||||
_loadPromise = (async () => {
|
||||
try {
|
||||
|
||||
@@ -2027,12 +2027,12 @@ async function _cmdUsage(args, ctx) {
|
||||
const messageCount = Number(session?.message_count || 0);
|
||||
const totalTokens = Number(session?.total_tokens || 0);
|
||||
const costTracked = chatRenderer.isCostTrackedEndpoint ? chatRenderer.isCostTrackedEndpoint(endpointUrl) : true;
|
||||
const cost = costTracked && chatRenderer.getSessionCost ? Number(chatRenderer.getSessionCost(sid) || 0) : 0;
|
||||
const costLine = costTracked
|
||||
? (cost > 0
|
||||
? `Estimated local cost: $${cost < 0.01 ? cost.toFixed(4) : cost.toFixed(3)}`
|
||||
: 'Estimated local cost: unavailable or zero')
|
||||
: 'Estimated local cost: not tracked for this endpoint';
|
||||
const cost = chatRenderer.getSessionCost ? Number(chatRenderer.getSessionCost(sid) || 0) : 0;
|
||||
const costLine = cost > 0
|
||||
? `Estimated local cost: $${cost < 0.01 ? cost.toFixed(4) : cost.toFixed(3)}`
|
||||
: costTracked
|
||||
? 'Estimated local cost: unavailable or zero'
|
||||
: 'Estimated local cost: no billable usage recorded';
|
||||
|
||||
slashReply(`<pre>${[
|
||||
`Session: ${ctx.esc(session?.name || 'Current chat')}`,
|
||||
|
||||
+86
-13
@@ -4,6 +4,13 @@
|
||||
* ASCII Spinner Module for AI thinking/processing status
|
||||
*/
|
||||
|
||||
// How long a canvas spinner may keep animating before its element has ever
|
||||
// been inserted into the document. start() runs synchronously, before the
|
||||
// caller appends the element, so frame 1 is always disconnected. Callers do
|
||||
// append in the same task, so anything past this window means the element is
|
||||
// never coming and the frames are drawing for nobody.
|
||||
const UNATTACHED_GRACE_MS = 2000;
|
||||
|
||||
class Spinner {
|
||||
constructor(message = "AI is processing", style = "right", animation = "spinner") {
|
||||
// Different animation frames
|
||||
@@ -21,6 +28,9 @@ class Spinner {
|
||||
this.intervalId = null;
|
||||
this.rafId = null;
|
||||
this.element = null;
|
||||
this._wpWasConnected = false;
|
||||
this._wpUnattachedSince = null;
|
||||
this._visHandler = null;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -74,6 +84,7 @@ class Spinner {
|
||||
}
|
||||
|
||||
_drawSineWave() {
|
||||
if (!this.isRunning) return;
|
||||
const ctx = this._ctx;
|
||||
const W = this._canvas.width;
|
||||
const H = this._canvas.height;
|
||||
@@ -120,9 +131,7 @@ class Spinner {
|
||||
ctx.fillStyle = 'rgba(156, 222, 242, 0.9)';
|
||||
ctx.fill();
|
||||
|
||||
if (this.isRunning) {
|
||||
this.rafId = requestAnimationFrame(() => this._drawSineWave());
|
||||
}
|
||||
if (this.isRunning) this._requestFrame();
|
||||
}
|
||||
|
||||
_createWhirlpoolElement() {
|
||||
@@ -158,6 +167,7 @@ class Spinner {
|
||||
}
|
||||
|
||||
_drawWhirlpool() {
|
||||
if (!this.isRunning) return;
|
||||
const ctx = this._wpCtx;
|
||||
const W = this._wpCanvas.width;
|
||||
const H = this._wpCanvas.height;
|
||||
@@ -229,18 +239,77 @@ class Spinner {
|
||||
ctx.fill();
|
||||
ctx.globalAlpha = 1;
|
||||
|
||||
if (!this.isRunning) return;
|
||||
// Leak-safe self-terminate: stop once our element WAS in the DOM and then
|
||||
// got removed (e.g. a loading row replaced by results). But keep spinning
|
||||
// before it's first appended — start() runs synchronously, before the
|
||||
// caller inserts the element, so it isn't connected on frame 1.
|
||||
// Leak-safe self-terminate. "Nobody can see this spinner" has two shapes
|
||||
// and we have to catch both:
|
||||
// 1. the element WAS in the DOM and then got removed (a loading row
|
||||
// replaced by results);
|
||||
// 2. the element was NEVER inserted, and the grace window for inserting
|
||||
// it has expired. The caller started a spinner and then took an early
|
||||
// return (aborted request, panel that resolved from cache), so no
|
||||
// frame we draw will ever be observed.
|
||||
// Case 2 is why this needs a deadline at all: while the element has never
|
||||
// been connected, `!this._wpWasConnected` stays true forever, so without
|
||||
// the grace check the loop re-arms until the tab closes.
|
||||
const connected = !!(this.element && this.element.isConnected);
|
||||
if (connected) this._wpWasConnected = true;
|
||||
if (connected || !this._wpWasConnected) {
|
||||
this.rafId = requestAnimationFrame(() => this._drawWhirlpool());
|
||||
} else {
|
||||
this.isRunning = false;
|
||||
if (connected) {
|
||||
this._wpWasConnected = true;
|
||||
this._wpUnattachedSince = null;
|
||||
} else if (!this._wpWasConnected) {
|
||||
if (this._wpUnattachedSince === null) this._wpUnattachedSince = performance.now();
|
||||
if (performance.now() - this._wpUnattachedSince > UNATTACHED_GRACE_MS) {
|
||||
this.stop();
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
if (connected || !this._wpWasConnected) {
|
||||
this._requestFrame();
|
||||
} else {
|
||||
this.stop();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Arm the next animation frame. Clearing rafId as the callback enters keeps
|
||||
* it a truthful "a frame is pending" flag, which is what stop() and the
|
||||
* visibility handler cancel against.
|
||||
*/
|
||||
_requestFrame() {
|
||||
this.rafId = requestAnimationFrame(() => {
|
||||
this.rafId = null;
|
||||
if (this.animation === 'sinewave') this._drawSineWave();
|
||||
else this._drawWhirlpool();
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* Stop drawing while the tab is hidden. Browsers throttle background rAF but
|
||||
* do not reliably stop the canvas work, and a spinner nobody is looking at
|
||||
* should cost nothing. The listener is owned by start()/stop() so it is never
|
||||
* left behind on a dead spinner.
|
||||
*/
|
||||
_armVisibilityPause() {
|
||||
if (this._visHandler) return;
|
||||
this._visHandler = () => {
|
||||
if (document.hidden) {
|
||||
if (this.rafId) {
|
||||
cancelAnimationFrame(this.rafId);
|
||||
this.rafId = null;
|
||||
}
|
||||
} else if (this.isRunning && !this.rafId) {
|
||||
// Reset the wave clock so the hidden interval doesn't arrive as one
|
||||
// huge dt and skip the animation forward.
|
||||
this._wavePrev = performance.now();
|
||||
this._requestFrame();
|
||||
}
|
||||
};
|
||||
document.addEventListener('visibilitychange', this._visHandler);
|
||||
}
|
||||
|
||||
_disarmVisibilityPause() {
|
||||
if (!this._visHandler) return;
|
||||
document.removeEventListener('visibilitychange', this._visHandler);
|
||||
this._visHandler = null;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -272,12 +341,15 @@ class Spinner {
|
||||
|
||||
if (this.animation === 'sinewave') {
|
||||
this._wavePrev = performance.now();
|
||||
this._armVisibilityPause();
|
||||
this._drawSineWave();
|
||||
return;
|
||||
}
|
||||
|
||||
if (this.animation === 'whirlpool') {
|
||||
this._wpStartedAt = performance.now();
|
||||
this._wpUnattachedSince = null;
|
||||
this._armVisibilityPause();
|
||||
this._drawWhirlpool();
|
||||
return;
|
||||
}
|
||||
@@ -302,6 +374,7 @@ class Spinner {
|
||||
cancelAnimationFrame(this.rafId);
|
||||
this.rafId = null;
|
||||
}
|
||||
this._disarmVisibilityPause();
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -0,0 +1,153 @@
|
||||
// Odysseus UI — startup shell sequencing
|
||||
// ES6 module — no application dependencies, DOM only.
|
||||
//
|
||||
// Revealing the application shell, retiring the boot loader, settling the
|
||||
// sidebar's own loading state, and firing a deferred URL route are separate
|
||||
// startup concerns that used to sit inline in app.js behind a single promise.
|
||||
// They live here so each step has one owner and so the whole contract can be
|
||||
// exercised directly (tests/test_startup_shell_js.py) without booting the app.
|
||||
|
||||
const LOADER_ID = 'app-loader';
|
||||
const SESSION_BOOTSTRAP_ROW_ID = 'session-list-loading';
|
||||
|
||||
// Route openers that read the hydrated session list. Everything else only
|
||||
// needs module wiring and must not wait on /api/sessions. `/email` spawns a
|
||||
// fresh chat, and that path falls back to the most recent session's model
|
||||
// (_createDirectChatFromPreferredModel in app.js) when there is no default
|
||||
// chat configured, so it genuinely needs the list.
|
||||
const ROUTES_NEEDING_SESSIONS = new Set(['/email']);
|
||||
|
||||
let _routeOpener = null;
|
||||
let _routeOpenerNeedsSessions = false;
|
||||
|
||||
function _loader() {
|
||||
return document.getElementById(LOADER_ID);
|
||||
}
|
||||
|
||||
/** Run `fn` after the next paint has committed (two animation frames). */
|
||||
export function afterNextPaint(fn) {
|
||||
requestAnimationFrame(() => requestAnimationFrame(fn));
|
||||
}
|
||||
|
||||
// The loader node stays in the DOM while sessions hydrate — sidebar-layout.js
|
||||
// and sessions.js both read its presence as a "still starting up" sentinel —
|
||||
// but it must stop covering, announcing, and animating over a usable shell.
|
||||
function _makeLoaderInert(loader) {
|
||||
if (!loader || loader.dataset.shellRevealed === 'true') return;
|
||||
loader.dataset.shellRevealed = 'true';
|
||||
loader.setAttribute('aria-hidden', 'true');
|
||||
loader.style.pointerEvents = 'none';
|
||||
loader.style.opacity = '0';
|
||||
// index.html's inline bootstrap animates the wave on a 150ms interval.
|
||||
// Nothing of it is visible any more, so stop rendering into it.
|
||||
try { window.__odysseusLoaderWaveStop?.(); } catch (_) {}
|
||||
}
|
||||
|
||||
/**
|
||||
* Hand the shell to the user once core wiring is done. Deferred by one paint
|
||||
* so the first frame lands with the app already laid out.
|
||||
*/
|
||||
export function revealApplicationShellAfterPaint() {
|
||||
const loader = _loader();
|
||||
if (!loader || loader.dataset.shellRevealScheduled === 'true') return;
|
||||
loader.dataset.shellRevealScheduled = 'true';
|
||||
afterNextPaint(() => _makeLoaderInert(_loader()));
|
||||
}
|
||||
|
||||
/** Retire the loader node for good. Safe to call after a reveal. */
|
||||
export function removeApplicationLoader() {
|
||||
const loader = _loader();
|
||||
if (!loader) return;
|
||||
_makeLoaderInert(loader);
|
||||
setTimeout(() => loader.remove(), 300);
|
||||
}
|
||||
|
||||
/**
|
||||
* Turn the sidebar's bootstrap row into a failure row. The write is delayed
|
||||
* until the session renderer's frame has committed so a late success cannot
|
||||
* leave stale failure text behind.
|
||||
*/
|
||||
export function markSessionListUnavailableIfStillBootstrapping() {
|
||||
afterNextPaint(() => {
|
||||
const row = document.getElementById(SESSION_BOOTSTRAP_ROW_ID);
|
||||
if (!row) return;
|
||||
const status = row.querySelector('[data-session-list-status]') || row;
|
||||
status.textContent = 'Chats unavailable';
|
||||
});
|
||||
}
|
||||
|
||||
/** True when `path`'s route opener reads the hydrated session list. */
|
||||
export function routeNeedsSessionData(path) {
|
||||
return ROUTES_NEEDING_SESSIONS.has(path);
|
||||
}
|
||||
|
||||
/**
|
||||
* Stash a URL route opener for later. At the point app.js resolves the route,
|
||||
* the modules its handlers drive (the rail new-chat handler, the email
|
||||
* section header handler, sessionModule) are still being wired further down
|
||||
* the same init pass, so the opener cannot run inline.
|
||||
*/
|
||||
export function deferRouteOpener(path, opener) {
|
||||
if (!opener) return;
|
||||
_routeOpener = opener;
|
||||
_routeOpenerNeedsSessions = routeNeedsSessionData(path);
|
||||
}
|
||||
|
||||
/**
|
||||
* Fire the deferred route opener if its data is ready. Called once when
|
||||
* wiring completes and again after authoritative session hydration; a route
|
||||
* that needs no session data takes the first call, one that does takes the
|
||||
* second.
|
||||
*
|
||||
* @returns {boolean} whether an opener ran.
|
||||
*/
|
||||
export function runDeferredRouteOpener({ sessionsSettled = false } = {}) {
|
||||
if (!_routeOpener) return false;
|
||||
if (_routeOpenerNeedsSessions && !sessionsSettled) return false;
|
||||
const opener = _routeOpener;
|
||||
_routeOpener = null;
|
||||
_routeOpenerNeedsSessions = false;
|
||||
try { opener(); } catch (e) { console.warn('route opener failed:', e); }
|
||||
return true;
|
||||
}
|
||||
|
||||
/**
|
||||
* Drive session hydration and everything that hangs off it settling: the
|
||||
* sidebar's failure row, the loader node, and any session-dependent route.
|
||||
*
|
||||
* @param {(() => Promise<boolean>)|null} loadSessions — resolves true only
|
||||
* after the session list was authoritatively loaded and applied. Null means
|
||||
* the session module failed to load.
|
||||
*/
|
||||
export function settleSessionHydration(loadSessions) {
|
||||
const settle = (succeeded) => {
|
||||
if (!succeeded) {
|
||||
markSessionListUnavailableIfStillBootstrapping();
|
||||
// A later unrelated caller must not be able to release a stale startup
|
||||
// opener against unknown session state.
|
||||
_routeOpener = null;
|
||||
_routeOpenerNeedsSessions = false;
|
||||
}
|
||||
removeApplicationLoader();
|
||||
if (succeeded) runDeferredRouteOpener({ sessionsSettled: true });
|
||||
return succeeded;
|
||||
};
|
||||
if (!loadSessions) {
|
||||
return Promise.resolve(settle(false));
|
||||
}
|
||||
// Kick the request off synchronously — a microtask hop here would delay the
|
||||
// fetch this whole change exists to get off the critical path.
|
||||
let pending;
|
||||
try {
|
||||
pending = loadSessions();
|
||||
} catch (e) {
|
||||
console.warn('loadSessions error:', e);
|
||||
return Promise.resolve(settle(false));
|
||||
}
|
||||
return Promise.resolve(pending)
|
||||
.then(result => settle(result === true))
|
||||
.catch(e => {
|
||||
console.warn('loadSessions error:', e);
|
||||
return settle(false);
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
// static/js/ui_visibility.js
|
||||
//
|
||||
// Per-item visibility for the sidebar and collapsed icon rail. Drives the
|
||||
// Settings → Appearance ("Customize UI") checkboxes, persisted in localStorage
|
||||
// under `odysseus-ui-visibility` (loaded/saved by app.js).
|
||||
//
|
||||
// Each key maps to the CSS selector(s) it controls. Tool/section selectors pair
|
||||
// the full-sidebar element with its #rail-* launcher so a tab hidden in the
|
||||
// full view also hides when the sidebar is minimized to the icon rail (id
|
||||
// mapping mirrors _railToolMap in app.js; #tool-library-btn ↔ #rail-archive).
|
||||
|
||||
// Selector map: UI customization key → CSS selector(s) for its target(s).
|
||||
export const UI_VIS_MAP = {
|
||||
'sidebar-brand': '.sidebar-brand-title',
|
||||
'sidebar-new-chat': '#sidebar-new-chat-btn',
|
||||
'sidebar-search': '#sidebar-search-btn',
|
||||
'sessions-section': '#sessions-section',
|
||||
'email-section': '#email-section, #rail-email',
|
||||
'tools-section': '#tools-section',
|
||||
// Per-tool entries pair the sidebar button with its rail launcher.
|
||||
'tool-calendar': '#tool-calendar-btn, #rail-calendar',
|
||||
'tool-compare': '#tool-compare-btn, #rail-compare',
|
||||
'tool-cookbook': '#tool-cookbook-btn, #rail-cookbook',
|
||||
'tool-research': '#tool-research-btn, #rail-research',
|
||||
'tool-gallery': '#tool-gallery-btn, #rail-gallery',
|
||||
'tool-library': '#tool-library-btn, #rail-archive',
|
||||
'tool-memory': '#tool-memory-btn, #rail-memory',
|
||||
'tool-notes': '#tool-notes-btn, #rail-notes',
|
||||
'tool-tasks': '#tool-tasks-btn, #rail-tasks',
|
||||
'tool-theme': '#tool-theme-btn, #rail-theme',
|
||||
'user-bar': '#user-bar-profile',
|
||||
'sidebar-settings-btn':'#user-bar-settings',
|
||||
'chat-meta': '.chat-meta-overlay',
|
||||
'welcome-text': '.welcome-name, .welcome-sub, #welcome-tip',
|
||||
'incognito-btn': '.incognito-btn',
|
||||
'web-toggle-btn': '#web-toggle-btn',
|
||||
'doc-toggle-btn': '#overflow-doc-btn',
|
||||
'rag-toggle-btn': '#overflow-rag-btn',
|
||||
'bash-toggle-btn': '#bash-toggle-btn',
|
||||
'overflow-plus-btn': '.overflow-wrapper',
|
||||
'mode-toggle': '.mode-toggle',
|
||||
'preset-mini-btn': '#overflow-preset-btn',
|
||||
'attach-btn': '#overflow-attach-btn',
|
||||
'research-btn': '#overflow-research-btn',
|
||||
'rail-new-chat': '#rail-new-session',
|
||||
};
|
||||
|
||||
// Keys hidden by default on first run (no localStorage yet).
|
||||
export const UI_VIS_DEFAULT_OFF = new Set(['rag-toggle-btn', 'text-emojis', 'chat-fullwidth']);
|
||||
|
||||
/**
|
||||
* Resolve every UI_VIS_MAP selector to visible (true) or hidden (false) for the
|
||||
* given saved state. Pure: no DOM, no localStorage — app.js applies the result.
|
||||
*
|
||||
* A key is visible when its stored value is not `false`, defaulting to on
|
||||
* unless it is in UI_VIS_DEFAULT_OFF. Per-tool entries also require the Tools
|
||||
* section to be on: hiding Tools hides every tool, mirroring the full sidebar
|
||||
* where the #tools-section container already hides them (the rail has no
|
||||
* container, so this rule keeps it in sync).
|
||||
*
|
||||
* @param {Record<string, boolean>} state
|
||||
* @returns {Record<string, boolean>} selector → visible
|
||||
*/
|
||||
export const resolveVisibility = (state = {}) => {
|
||||
const toolsOn = state['tools-section'] !== false;
|
||||
const out = {};
|
||||
for (const [key, selector] of Object.entries(UI_VIS_MAP)) {
|
||||
let visible = key in state ? state[key] !== false : !UI_VIS_DEFAULT_OFF.has(key);
|
||||
if (!toolsOn && key.startsWith('tool-')) visible = false;
|
||||
out[selector] = visible;
|
||||
}
|
||||
return out;
|
||||
};
|
||||
@@ -38015,6 +38015,12 @@ button.cal-add-btn.cal-add-btn-text.cal-add-btn-sm:hover .cal-add-label {
|
||||
outline-offset: 2px;
|
||||
border-radius: 5px;
|
||||
}
|
||||
/* Bootstrap row shown while the session list hydrates, and on load failure.
|
||||
Reads as a normal list row but is not selectable. */
|
||||
.session-list-bootstrap {
|
||||
cursor: default;
|
||||
pointer-events: none;
|
||||
}
|
||||
#email-lib-grid .date-section-header {
|
||||
padding: 10px 5px 3px;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,277 @@
|
||||
// Tests for the live-thinking throttle that bounds DOM work during long
|
||||
// reasoning streams (see static/js/liveThinkingThrottle.js).
|
||||
//
|
||||
// The throttle's contract is what the terminal paths in chat.js lean on:
|
||||
// a burst of deltas becomes ONE commit carrying the latest text; flush()
|
||||
// lands trailing text synchronously and cannot double-commit; cancel()
|
||||
// guarantees nothing lands after a stream is finished or backgrounded.
|
||||
//
|
||||
// Timers are injected, so this runs with no DOM and no real clock.
|
||||
import assert from 'node:assert/strict';
|
||||
import test from 'node:test';
|
||||
|
||||
import {
|
||||
createIncrementalDisplayProjector,
|
||||
createLiveThinkingThrottle,
|
||||
createThinkingAnalysisGate,
|
||||
stripLiveThinkingTags,
|
||||
} from '../static/js/liveThinkingThrottle.js';
|
||||
|
||||
function fakeTimers() {
|
||||
let nextId = 1;
|
||||
const callbacks = new Map();
|
||||
const delays = [];
|
||||
return {
|
||||
schedule(callback, delay) {
|
||||
const id = nextId++;
|
||||
callbacks.set(id, callback);
|
||||
delays.push(delay);
|
||||
return id;
|
||||
},
|
||||
cancel(id) {
|
||||
callbacks.delete(id);
|
||||
},
|
||||
run(id) {
|
||||
const callback = callbacks.get(id);
|
||||
assert.ok(callback, `missing timer ${id}`);
|
||||
callbacks.delete(id);
|
||||
callback();
|
||||
},
|
||||
pendingIds() {
|
||||
return [...callbacks.keys()];
|
||||
},
|
||||
delays,
|
||||
};
|
||||
}
|
||||
|
||||
test('coalesces a burst and commits only the latest text after 100 ms', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
const throttle = createLiveThinkingThrottle((value) => commits.push(value), timers);
|
||||
|
||||
throttle.update('a');
|
||||
throttle.update('ab');
|
||||
throttle.update('abc');
|
||||
|
||||
assert.deepEqual(commits, []);
|
||||
assert.deepEqual(timers.delays, [100], 'a burst must schedule exactly one commit');
|
||||
const [timer] = timers.pendingIds();
|
||||
timers.run(timer);
|
||||
assert.deepEqual(commits, ['abc']);
|
||||
});
|
||||
|
||||
test('commit count stays flat as the stream grows', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
const throttle = createLiveThinkingThrottle((value) => commits.push(value), timers);
|
||||
|
||||
// 500 deltas arriving inside one window is the regression this guards:
|
||||
// the old code committed once per delta, so work grew with stream length.
|
||||
let text = '';
|
||||
for (let i = 0; i < 500; i++) {
|
||||
text += 'token ';
|
||||
throttle.update(text);
|
||||
}
|
||||
assert.deepEqual(commits, []);
|
||||
assert.equal(timers.pendingIds().length, 1);
|
||||
timers.run(timers.pendingIds()[0]);
|
||||
assert.equal(commits.length, 1);
|
||||
assert.equal(commits[0], text);
|
||||
});
|
||||
|
||||
test('prepares a 200K cumulative stream only at scheduled commit cadence', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
let prepareCalls = 0;
|
||||
let scannedCharacters = 0;
|
||||
const throttle = createLiveThinkingThrottle((value) => commits.push(value), {
|
||||
...timers,
|
||||
prepare(value) {
|
||||
prepareCalls += 1;
|
||||
scannedCharacters += value.length;
|
||||
return stripLiveThinkingTags(value);
|
||||
},
|
||||
});
|
||||
|
||||
const delta = 'reasoning '.repeat(10); // 100 characters
|
||||
let cumulative = '';
|
||||
for (let i = 0; i < 2000; i++) {
|
||||
cumulative += delta;
|
||||
throttle.update(cumulative);
|
||||
}
|
||||
|
||||
assert.equal(cumulative.length, 200_000);
|
||||
assert.equal(prepareCalls, 0, 'cumulative extraction must not run per delta');
|
||||
assert.equal(timers.pendingIds().length, 1);
|
||||
timers.run(timers.pendingIds()[0]);
|
||||
assert.equal(prepareCalls, 1);
|
||||
assert.equal(scannedCharacters, 200_000);
|
||||
assert.deepEqual(commits, [cumulative]);
|
||||
});
|
||||
|
||||
test('ordinary answers and reasoning deltas do not request cumulative analysis', () => {
|
||||
const startsReasoning = (text) => /^\s*thinking(?:\s+process)?\s*:/i.test(text);
|
||||
const ordinaryGate = createThinkingAnalysisGate({ startsWithReasoningPrefix: startsReasoning });
|
||||
let ordinary = '';
|
||||
let ordinaryAnalyses = 0;
|
||||
for (let i = 0; i < 2000; i++) {
|
||||
ordinary += i === 0 ? 'Here is the answer. ' : 'answer '.repeat(10);
|
||||
if (ordinaryGate.shouldAnalyze(ordinary)) ordinaryAnalyses += 1;
|
||||
}
|
||||
assert.equal(ordinaryAnalyses, 0);
|
||||
|
||||
const thinkingGate = createThinkingAnalysisGate({ startsWithReasoningPrefix: startsReasoning });
|
||||
let thinking = 'Thin';
|
||||
assert.equal(thinkingGate.shouldAnalyze(thinking), false);
|
||||
thinking += 'king: inspect the problem';
|
||||
assert.equal(thinkingGate.shouldAnalyze(thinking), true);
|
||||
for (let i = 0; i < 2000; i++) {
|
||||
thinking += ' reasoning'.repeat(10);
|
||||
assert.equal(thinkingGate.shouldAnalyze(thinking, { isThinking: true, nonTagThinking: true }), false);
|
||||
}
|
||||
thinking += '\n\nHere is the answer';
|
||||
assert.equal(thinkingGate.shouldAnalyze(thinking, { isThinking: true, nonTagThinking: true }), true);
|
||||
|
||||
const whitespaceGate = createThinkingAnalysisGate({ startsWithReasoningPrefix: startsReasoning });
|
||||
let whitespaceThinking = ' '.repeat(250);
|
||||
assert.equal(whitespaceGate.shouldAnalyze(whitespaceThinking), false);
|
||||
whitespaceThinking += 'Thinking: bounded probe';
|
||||
assert.equal(whitespaceGate.shouldAnalyze(whitespaceThinking), true);
|
||||
});
|
||||
|
||||
test('split namespaced closes and false-close deadlines request analysis', () => {
|
||||
let clock = 100;
|
||||
const gate = createThinkingAnalysisGate({ now: () => clock });
|
||||
let text = '<mm:think>x</mm:';
|
||||
assert.equal(gate.shouldAnalyze(text, { isThinking: true }), true, 'fresh opening tag is analyzed');
|
||||
text += 'think>answer';
|
||||
assert.equal(gate.shouldAnalyze(text, { isThinking: true }), true, 'split namespaced close is analyzed');
|
||||
|
||||
text += ' still waiting';
|
||||
assert.equal(gate.shouldAnalyze(text, { isThinking: true, recheckAt: 500 }), false);
|
||||
clock = 500;
|
||||
text += ' next delta';
|
||||
assert.equal(gate.shouldAnalyze(text, { isThinking: true, recheckAt: 500 }), true);
|
||||
|
||||
const attributedGate = createThinkingAnalysisGate();
|
||||
let attributed = `<think data-provider="${'x'.repeat(400)}"`;
|
||||
assert.equal(attributedGate.shouldAnalyze(attributed), false);
|
||||
attributed += '>reasoning';
|
||||
assert.equal(attributedGate.shouldAnalyze(attributed), true, 'bounded carry preserves split tag attributes');
|
||||
});
|
||||
|
||||
test('display projection is append-only and filters a structured tail once', () => {
|
||||
let filterCalls = 0;
|
||||
let filteredCharacters = 0;
|
||||
const projector = createIncrementalDisplayProjector((text) => {
|
||||
filterCalls += 1;
|
||||
filteredCharacters += text.length;
|
||||
return text.replace(/\[TOOL_CALL\][\s\S]*$/i, '');
|
||||
});
|
||||
|
||||
let text = '';
|
||||
for (let i = 0; i < 2000; i++) {
|
||||
const delta = i === 0 ? 'Here is the answer. ' : 'ordinary text ';
|
||||
text += delta;
|
||||
assert.equal(projector.append(delta, text), text);
|
||||
}
|
||||
assert.equal(filterCalls, 0, 'ordinary deltas never run the cumulative filter');
|
||||
|
||||
text += '[TOOL_';
|
||||
projector.append('[TOOL_', text);
|
||||
text += 'CALL]{"name":"read"}';
|
||||
const beforeToolPayload = projector.append('CALL]{"name":"read"}', text);
|
||||
for (let i = 0; i < 2000; i++) {
|
||||
const delta = 'payload ';
|
||||
text += delta;
|
||||
assert.equal(projector.append(delta, text), beforeToolPayload);
|
||||
}
|
||||
assert.equal(filterCalls, 1, 'structured payload filtering happens only at its boundary');
|
||||
assert.ok(filteredCharacters < text.length, 'filter work is bounded by the first structured boundary');
|
||||
});
|
||||
|
||||
test('literal escaped tags survive and malformed live tags retain trailing text', () => {
|
||||
assert.equal(
|
||||
stripLiveThinkingTags('<think>literal</think>'),
|
||||
'<think>literal</think>',
|
||||
);
|
||||
assert.equal(
|
||||
stripLiveThinkingTags('<think>first</think> middle <thinking mode="deep">trailing'),
|
||||
'first middle trailing',
|
||||
);
|
||||
assert.equal(stripLiveThinkingTags('answer with 2 < 3 and 5 > 4'), 'answer with 2 < 3 and 5 > 4');
|
||||
});
|
||||
|
||||
test('terminal flush prepares and commits the complete trailing cumulative text', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
const throttle = createLiveThinkingThrottle((value) => commits.push(value), {
|
||||
...timers,
|
||||
prepare: stripLiveThinkingTags,
|
||||
});
|
||||
|
||||
throttle.update('<think>reasoning without a closing tag');
|
||||
assert.equal(throttle.flush(), true);
|
||||
assert.deepEqual(commits, ['reasoning without a closing tag']);
|
||||
assert.deepEqual(timers.pendingIds(), []);
|
||||
});
|
||||
|
||||
test('independent throttles cannot commit cancelled text into another session', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
const first = createLiveThinkingThrottle((value) => commits.push(['first', value]), timers);
|
||||
const second = createLiveThinkingThrottle((value) => commits.push(['second', value]), timers);
|
||||
|
||||
first.update('stale first-session text');
|
||||
second.update('current second-session text');
|
||||
first.cancel();
|
||||
assert.equal(second.flush(), true);
|
||||
|
||||
assert.deepEqual(timers.pendingIds(), []);
|
||||
assert.deepEqual(commits, [['second', 'current second-session text']]);
|
||||
});
|
||||
|
||||
test('flush synchronously preserves trailing text and cancels the pending callback', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
const throttle = createLiveThinkingThrottle((value) => commits.push(value), timers);
|
||||
|
||||
throttle.update('trailing text');
|
||||
assert.equal(throttle.flush(), true);
|
||||
assert.deepEqual(commits, ['trailing text']);
|
||||
assert.deepEqual(timers.pendingIds(), []);
|
||||
assert.equal(throttle.flush(), false, 'clean flush must not duplicate the commit');
|
||||
});
|
||||
|
||||
test('cancel discards pending work without a late DOM commit', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
const throttle = createLiveThinkingThrottle((value) => commits.push(value), timers);
|
||||
|
||||
throttle.update('stale session text');
|
||||
throttle.cancel();
|
||||
assert.deepEqual(timers.pendingIds(), []);
|
||||
assert.deepEqual(commits, []);
|
||||
});
|
||||
|
||||
test('a cancelled throttle accepts new work again', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
const throttle = createLiveThinkingThrottle((value) => commits.push(value), timers);
|
||||
|
||||
throttle.update('discarded');
|
||||
throttle.cancel();
|
||||
throttle.update('fresh');
|
||||
assert.equal(throttle.flush(), true);
|
||||
assert.deepEqual(commits, ['fresh']);
|
||||
});
|
||||
|
||||
test('coerces nullish updates instead of committing undefined', () => {
|
||||
const timers = fakeTimers();
|
||||
const commits = [];
|
||||
const throttle = createLiveThinkingThrottle((value) => commits.push(value), timers);
|
||||
|
||||
throttle.update(null);
|
||||
throttle.flush();
|
||||
assert.deepEqual(commits, ['']);
|
||||
});
|
||||
@@ -0,0 +1,394 @@
|
||||
"""Regression guard for #5558 — POST /api/personal/add_directory must not run
|
||||
the indexing job on the event loop.
|
||||
|
||||
The handler is ``async def`` but called ``rag.index_personal_documents``
|
||||
(os.walk + file reads + per-chunk embedding + Chroma inserts) inline, so
|
||||
FastAPI ran the whole job on the event loop and every other request queued
|
||||
behind it: indexing a real directory froze the UI and API for 25+ minutes.
|
||||
``personal_docs_manager.add_directory`` sits in the same blocking section — it
|
||||
triggers ``refresh_index()``, which re-extracts text across tracked dirs.
|
||||
|
||||
These tests build the real router with fake managers and compare the thread
|
||||
the indexing work runs on against the event loop's thread.
|
||||
"""
|
||||
import asyncio
|
||||
import os
|
||||
import threading
|
||||
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///:memory:")
|
||||
|
||||
import httpx
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
|
||||
def _serialization_probe():
|
||||
"""Shared counter proving two critical sections never overlap."""
|
||||
state = {"active": 0, "max_active": 0}
|
||||
lock = threading.Lock()
|
||||
|
||||
def enter():
|
||||
with lock:
|
||||
state["active"] += 1
|
||||
state["max_active"] = max(state["max_active"], state["active"])
|
||||
|
||||
def leave():
|
||||
with lock:
|
||||
state["active"] -= 1
|
||||
|
||||
return state, enter, leave
|
||||
|
||||
|
||||
# Concurrency tests are `async def` (pyproject asyncio_mode="auto") and drive the
|
||||
# ASGI app through httpx.ASGITransport + AsyncClient + asyncio.gather, NOT starlette
|
||||
# TestClient + ThreadPoolExecutor: the job lock is an asyncio.Lock acquired in the
|
||||
# async handler, and TestClient's portal-thread dispatch deadlocks against it (same
|
||||
# reason test_notes_fail_closed_auth.py uses ASGITransport). asyncio.gather runs both
|
||||
# requests on the test's own loop.
|
||||
def _async_client(app):
|
||||
return httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://t")
|
||||
|
||||
import routes.personal_routes as personal_routes
|
||||
from core.middleware import require_admin
|
||||
from src.auth_helpers import require_user
|
||||
|
||||
|
||||
class _FakeRag:
|
||||
def __init__(self, record):
|
||||
self._record = record
|
||||
|
||||
def index_personal_documents(self, directory, owner=None):
|
||||
self._record["index_thread"] = threading.get_ident()
|
||||
return {"success": True, "indexed_count": 3, "failed_count": 0}
|
||||
|
||||
def _split_into_chunks(self, text, chunk_size=500):
|
||||
return [text]
|
||||
|
||||
def add_document(self, chunk, metadata):
|
||||
self._record["add_document_thread"] = threading.get_ident()
|
||||
return True
|
||||
|
||||
def delete_by_source(self, filepath):
|
||||
self._record["delete_thread"] = threading.get_ident()
|
||||
return 1
|
||||
|
||||
|
||||
class _FakeDocsManager:
|
||||
def __init__(self, record):
|
||||
self._record = record
|
||||
self.index = []
|
||||
|
||||
def add_directory(self, directory, *, index=True, owner=None):
|
||||
self._record["bookkeeping_thread"] = threading.get_ident()
|
||||
self._record["bookkeeping_index_flag"] = index
|
||||
|
||||
def exclude_file(self, filepath):
|
||||
self._record["exclude_thread"] = threading.get_ident()
|
||||
|
||||
|
||||
def _build_app(tmp_path, monkeypatch, record):
|
||||
monkeypatch.setattr(personal_routes, "PERSONAL_DIR", str(tmp_path))
|
||||
monkeypatch.setattr(personal_routes, "get_rag_manager", lambda: _FakeRag(record))
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(
|
||||
personal_routes.setup_personal_routes(_FakeDocsManager(record), None, True)
|
||||
)
|
||||
app.dependency_overrides[require_user] = lambda: "tester"
|
||||
app.dependency_overrides[require_admin] = lambda: None
|
||||
|
||||
@app.get("/loop-thread")
|
||||
async def loop_thread_probe():
|
||||
return {"thread": threading.get_ident()}
|
||||
|
||||
return app
|
||||
|
||||
|
||||
def test_indexing_runs_off_the_event_loop(tmp_path, monkeypatch):
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
target = tmp_path / "docs"
|
||||
target.mkdir()
|
||||
|
||||
# Context-manager client: one portal/event loop serves both requests, so
|
||||
# the probe and the POST are guaranteed to see the same loop thread.
|
||||
with TestClient(app) as client:
|
||||
loop_thread = client.get("/loop-thread").json()["thread"]
|
||||
resp = client.post(
|
||||
"/api/personal/add_directory", json={"directory": str(target)}
|
||||
)
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert record["index_thread"] != loop_thread, (
|
||||
"index_personal_documents ran on the event loop thread — every other "
|
||||
"request queues behind the indexing job (#5558)"
|
||||
)
|
||||
assert record["bookkeeping_thread"] != loop_thread, (
|
||||
"personal_docs_manager.add_directory (refresh_index) ran on the event "
|
||||
"loop thread"
|
||||
)
|
||||
|
||||
|
||||
def test_response_and_bookkeeping_unchanged(tmp_path, monkeypatch):
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
target = tmp_path / "docs"
|
||||
target.mkdir()
|
||||
|
||||
client = TestClient(app)
|
||||
resp = client.post("/api/personal/add_directory", json={"directory": str(target)})
|
||||
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["success"] is True
|
||||
assert body["indexed_count"] == 3
|
||||
assert body["failed_count"] == 0
|
||||
assert body["directory"] == os.path.realpath(str(target))
|
||||
assert record["bookkeeping_index_flag"] is False
|
||||
|
||||
|
||||
async def test_concurrent_add_directory_requests_serialize_indexing(tmp_path, monkeypatch):
|
||||
"""Off-loop execution must not mean parallel index jobs: concurrent
|
||||
requests would race PersonalDocsManager's unsynchronized list mutations
|
||||
and file writes (save_directories/_save_excluded are plain open('w'))."""
|
||||
import time
|
||||
|
||||
state, enter, leave = _serialization_probe()
|
||||
|
||||
def _slow_index(self, directory, owner=None):
|
||||
enter(); time.sleep(0.2); leave()
|
||||
return {"success": True, "indexed_count": 1, "failed_count": 0}
|
||||
|
||||
monkeypatch.setattr(_FakeRag, "index_personal_documents", _slow_index)
|
||||
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
for name in ("docs_a", "docs_b"):
|
||||
(tmp_path / name).mkdir()
|
||||
|
||||
async with _async_client(app) as ac:
|
||||
results = await asyncio.gather(
|
||||
ac.post("/api/personal/add_directory", json={"directory": str(tmp_path / "docs_a")}),
|
||||
ac.post("/api/personal/add_directory", json={"directory": str(tmp_path / "docs_b")}),
|
||||
)
|
||||
|
||||
assert all(r.status_code == 200 for r in results)
|
||||
assert state["max_active"] == 1, (
|
||||
f"{state['max_active']} index jobs ran in parallel — concurrent "
|
||||
"add_directory requests must serialize"
|
||||
)
|
||||
|
||||
|
||||
def test_failed_indexing_still_returns_500(tmp_path, monkeypatch):
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
target = tmp_path / "docs"
|
||||
target.mkdir()
|
||||
|
||||
def _fail(directory, owner=None):
|
||||
return {"success": False, "message": "boom"}
|
||||
|
||||
monkeypatch.setattr(_FakeRag, "index_personal_documents", staticmethod(_fail))
|
||||
|
||||
client = TestClient(app)
|
||||
resp = client.post("/api/personal/add_directory", json={"directory": str(target)})
|
||||
assert resp.status_code == 500
|
||||
assert "boom" in resp.json()["detail"]
|
||||
|
||||
|
||||
async def test_add_and_remove_serialize(tmp_path, monkeypatch):
|
||||
"""#5634: remove must hold the SAME job lock as add. Otherwise a remove
|
||||
running while an add job is in flight races PersonalDocsManager's
|
||||
unsynchronized list/index mutations — the inconsistent state the PR's
|
||||
'add/remove are serialized' guarantee claims to prevent."""
|
||||
import time
|
||||
|
||||
state, enter, leave = _serialization_probe()
|
||||
|
||||
def _slow_index(self, directory, owner=None):
|
||||
enter(); time.sleep(0.25); leave()
|
||||
return {"success": True, "indexed_count": 1, "failed_count": 0}
|
||||
|
||||
def _slow_remove(self, directory):
|
||||
enter(); time.sleep(0.25); leave()
|
||||
|
||||
monkeypatch.setattr(_FakeRag, "index_personal_documents", _slow_index)
|
||||
monkeypatch.setattr(_FakeDocsManager, "remove_directory", _slow_remove, raising=False)
|
||||
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
(tmp_path / "docs_a").mkdir()
|
||||
(tmp_path / "docs_b").mkdir()
|
||||
|
||||
async with _async_client(app) as ac:
|
||||
results = await asyncio.gather(
|
||||
ac.post("/api/personal/add_directory", json={"directory": str(tmp_path / "docs_a")}),
|
||||
ac.delete("/api/personal/remove_directory", params={"directory": str(tmp_path / "docs_b")}),
|
||||
)
|
||||
|
||||
assert all(r.status_code == 200 for r in results)
|
||||
assert state["max_active"] == 1, (
|
||||
f"{state['max_active']} add/remove critical sections overlapped — "
|
||||
"remove must hold the same index job lock as add"
|
||||
)
|
||||
|
||||
|
||||
async def test_add_and_upload_serialize(tmp_path, monkeypatch):
|
||||
"""#5634 follow-up: POST /upload writes chunks into the vector store and then
|
||||
calls personal_docs_manager.add_directory — the same vector/tracking state
|
||||
add_directory mutates. It must hold the SAME job lock, or an upload landing
|
||||
mid-add interleaves two writers over unsynchronized state."""
|
||||
import time
|
||||
|
||||
state, enter, leave = _serialization_probe()
|
||||
|
||||
def _slow_index(self, directory, owner=None):
|
||||
enter(); time.sleep(0.25); leave()
|
||||
return {"success": True, "indexed_count": 1, "failed_count": 0}
|
||||
|
||||
def _slow_add_document(self, chunk, metadata):
|
||||
self._record["add_document_thread"] = threading.get_ident()
|
||||
enter(); time.sleep(0.25); leave()
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(_FakeRag, "index_personal_documents", _slow_index)
|
||||
monkeypatch.setattr(_FakeRag, "add_document", _slow_add_document)
|
||||
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
monkeypatch.setattr(personal_routes, "UPLOADS_DIR", str(tmp_path / "uploads"))
|
||||
monkeypatch.setattr(personal_routes, "require_privilege", lambda request, key: "tester")
|
||||
(tmp_path / "docs_a").mkdir()
|
||||
|
||||
async with _async_client(app) as ac:
|
||||
results = await asyncio.gather(
|
||||
ac.post("/api/personal/add_directory", json={"directory": str(tmp_path / "docs_a")}),
|
||||
ac.post("/api/personal/upload", files={"files": ("a.txt", b"hello world", "text/plain")}),
|
||||
)
|
||||
|
||||
assert all(r.status_code == 200 for r in results)
|
||||
# The test coroutine runs on the event loop, so this IS the loop thread.
|
||||
assert record["add_document_thread"] != threading.get_ident(), (
|
||||
"rag.add_document ran on the event loop thread — chunk writes block "
|
||||
"every other request for the duration of the upload"
|
||||
)
|
||||
assert state["max_active"] == 1, (
|
||||
f"{state['max_active']} add/upload critical sections overlapped — "
|
||||
"upload must hold the same index job lock as add"
|
||||
)
|
||||
|
||||
|
||||
async def test_upload_processes_each_payload_before_reading_the_next(tmp_path, monkeypatch):
|
||||
"""A multi-file upload must retain at most one capped payload at a time."""
|
||||
from starlette.datastructures import UploadFile as StarletteUploadFile
|
||||
|
||||
reads = []
|
||||
original_read = StarletteUploadFile.read
|
||||
|
||||
async def _recording_read(upload, size=-1):
|
||||
reads.append(upload.filename)
|
||||
return await original_read(upload, size)
|
||||
|
||||
def _record_first_index(self, chunk, metadata):
|
||||
self._record.setdefault("reads_at_first_index", len(reads))
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(StarletteUploadFile, "read", _recording_read)
|
||||
monkeypatch.setattr(_FakeRag, "add_document", _record_first_index)
|
||||
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
monkeypatch.setattr(personal_routes, "UPLOADS_DIR", str(tmp_path / "uploads"))
|
||||
monkeypatch.setattr(personal_routes, "require_privilege", lambda request, key: "tester")
|
||||
|
||||
files = [
|
||||
("files", ("a.txt", b"alpha", "text/plain")),
|
||||
("files", ("b.txt", b"bravo", "text/plain")),
|
||||
("files", ("c.txt", b"charlie", "text/plain")),
|
||||
]
|
||||
async with _async_client(app) as ac:
|
||||
response = await ac.post("/api/personal/upload", files=files)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["uploaded"] == ["a.txt", "b.txt", "c.txt"]
|
||||
assert reads == ["a.txt", "b.txt", "c.txt"]
|
||||
assert record["reads_at_first_index"] == 1, (
|
||||
"all upload bodies were retained before worker processing began"
|
||||
)
|
||||
|
||||
|
||||
async def test_add_and_delete_file_serialize(tmp_path, monkeypatch):
|
||||
"""#5634 follow-up: DELETE /file removes chunks from the vector store and
|
||||
calls personal_docs_manager.exclude_file. Both mutate state add_directory
|
||||
also touches, so the delete must hold the SAME job lock as add."""
|
||||
import time
|
||||
|
||||
state, enter, leave = _serialization_probe()
|
||||
|
||||
def _slow_index(self, directory, owner=None):
|
||||
enter(); time.sleep(0.25); leave()
|
||||
return {"success": True, "indexed_count": 1, "failed_count": 0}
|
||||
|
||||
def _slow_delete(self, filepath):
|
||||
self._record["delete_thread"] = threading.get_ident()
|
||||
enter(); time.sleep(0.25); leave()
|
||||
return 1
|
||||
|
||||
monkeypatch.setattr(_FakeRag, "index_personal_documents", _slow_index)
|
||||
monkeypatch.setattr(_FakeRag, "delete_by_source", _slow_delete)
|
||||
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
monkeypatch.setattr(personal_routes, "UPLOADS_DIR", str(tmp_path / "uploads"))
|
||||
(tmp_path / "docs_a").mkdir()
|
||||
doomed = tmp_path / "doomed.txt"
|
||||
doomed.write_text("bye")
|
||||
|
||||
async with _async_client(app) as ac:
|
||||
results = await asyncio.gather(
|
||||
ac.post("/api/personal/add_directory", json={"directory": str(tmp_path / "docs_a")}),
|
||||
ac.delete("/api/personal/file", params={"filepath": str(doomed)}),
|
||||
)
|
||||
|
||||
assert all(r.status_code == 200 for r in results)
|
||||
assert record["delete_thread"] != threading.get_ident(), (
|
||||
"rag.delete_by_source ran on the event loop thread"
|
||||
)
|
||||
assert state["max_active"] == 1, (
|
||||
f"{state['max_active']} add/delete critical sections overlapped — "
|
||||
"delete must hold the same index job lock as add"
|
||||
)
|
||||
|
||||
|
||||
async def test_reload_serializes_with_add(tmp_path, monkeypatch):
|
||||
"""#5634: POST /reload rebuilds the index via refresh_index(); it must hold
|
||||
the same job lock so it cannot race an in-flight add job."""
|
||||
import time
|
||||
|
||||
state, enter, leave = _serialization_probe()
|
||||
|
||||
def _slow_index(self, directory, owner=None):
|
||||
enter(); time.sleep(0.25); leave()
|
||||
return {"success": True, "indexed_count": 1, "failed_count": 0}
|
||||
|
||||
def _slow_refresh(self):
|
||||
enter(); time.sleep(0.25); leave()
|
||||
|
||||
monkeypatch.setattr(_FakeRag, "index_personal_documents", _slow_index)
|
||||
monkeypatch.setattr(_FakeDocsManager, "refresh_index", _slow_refresh, raising=False)
|
||||
|
||||
record = {}
|
||||
app = _build_app(tmp_path, monkeypatch, record)
|
||||
(tmp_path / "docs_a").mkdir()
|
||||
|
||||
async with _async_client(app) as ac:
|
||||
results = await asyncio.gather(
|
||||
ac.post("/api/personal/add_directory", json={"directory": str(tmp_path / "docs_a")}),
|
||||
ac.post("/api/personal/reload"),
|
||||
)
|
||||
|
||||
assert all(r.status_code == 200 for r in results)
|
||||
assert state["max_active"] == 1, (
|
||||
f"{state['max_active']} add/reload critical sections overlapped — "
|
||||
"reload must hold the same index job lock as add"
|
||||
)
|
||||
@@ -0,0 +1,128 @@
|
||||
"""Windows execution contract for the agent Bash tool."""
|
||||
|
||||
import pytest
|
||||
|
||||
from src.agent_tools import subprocess_tools
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windows_bash_uses_git_bash_with_structural_cwd(monkeypatch):
|
||||
captured = {}
|
||||
bash = r"C:\Program Files\Git\bin\bash.exe"
|
||||
workspace = r"D:\Workspaces\Project with spaces"
|
||||
process = object()
|
||||
|
||||
monkeypatch.setattr(subprocess_tools, "IS_WINDOWS", True)
|
||||
monkeypatch.setattr(subprocess_tools, "find_bash", lambda: bash)
|
||||
|
||||
async def fake_exec(*argv, **kwargs):
|
||||
captured["argv"] = argv
|
||||
captured["kwargs"] = kwargs
|
||||
return process
|
||||
|
||||
async def fail_shell(*_args, **_kwargs):
|
||||
pytest.fail("native Windows Bash must not execute through cmd.exe")
|
||||
|
||||
monkeypatch.setattr(subprocess_tools.asyncio, "create_subprocess_exec", fake_exec)
|
||||
monkeypatch.setattr(subprocess_tools.asyncio, "create_subprocess_shell", fail_shell)
|
||||
|
||||
result = await subprocess_tools._create_bash_subprocess(
|
||||
"pwd; cat package.json",
|
||||
cwd=workspace,
|
||||
env={"HOME": r"C:\Odysseus\data"},
|
||||
)
|
||||
|
||||
assert result is process
|
||||
assert captured["argv"] == (bash, "-c", "pwd; cat package.json")
|
||||
assert captured["kwargs"]["cwd"] == workspace
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windows_bash_without_git_bash_fails_clearly(monkeypatch):
|
||||
monkeypatch.setattr(subprocess_tools, "IS_WINDOWS", True)
|
||||
monkeypatch.setattr(subprocess_tools, "find_bash", lambda: None)
|
||||
|
||||
async def fail_spawn(*_args, **_kwargs):
|
||||
pytest.fail("no subprocess should start without Git Bash")
|
||||
|
||||
monkeypatch.setattr(subprocess_tools.asyncio, "create_subprocess_exec", fail_spawn)
|
||||
monkeypatch.setattr(subprocess_tools.asyncio, "create_subprocess_shell", fail_spawn)
|
||||
|
||||
with pytest.raises(RuntimeError, match="Git Bash is required"):
|
||||
await subprocess_tools._create_bash_subprocess("pwd", cwd=r"C:\Work")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bash_tool_returns_install_hint_when_git_bash_is_missing(monkeypatch):
|
||||
monkeypatch.setattr(subprocess_tools, "IS_WINDOWS", True)
|
||||
monkeypatch.setattr(subprocess_tools, "find_bash", lambda: None)
|
||||
|
||||
result = await subprocess_tools.BashTool().execute(
|
||||
"pwd",
|
||||
{"subproc_env": {}, "session_id": None},
|
||||
)
|
||||
|
||||
assert result["exit_code"] == 1
|
||||
assert "install Git for Windows" in result["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_windows_bash_does_not_use_a_stray_tmux_executable(monkeypatch):
|
||||
captured = {}
|
||||
workspace = r"D:\Workspaces\Project with spaces"
|
||||
|
||||
monkeypatch.setattr(subprocess_tools, "IS_WINDOWS", True)
|
||||
monkeypatch.setattr(
|
||||
subprocess_tools.shutil,
|
||||
"which",
|
||||
lambda name: r"C:\msys64\usr\bin\tmux.exe",
|
||||
)
|
||||
monkeypatch.setattr("src.tool_execution.agent_cwd", lambda: workspace)
|
||||
|
||||
async def fail_tmux(*_args, **_kwargs):
|
||||
pytest.fail("native Windows must not enter the POSIX tmux path")
|
||||
|
||||
async def fake_create(command, **kwargs):
|
||||
captured["command"] = command
|
||||
captured["kwargs"] = kwargs
|
||||
return object()
|
||||
|
||||
async def fake_stream(_process, **_kwargs):
|
||||
return "ok", "", 0, False
|
||||
|
||||
monkeypatch.setattr(subprocess_tools, "_run_tmux_bash", fail_tmux)
|
||||
monkeypatch.setattr(subprocess_tools, "_create_bash_subprocess", fake_create)
|
||||
monkeypatch.setattr(subprocess_tools, "_run_subprocess_streaming", fake_stream)
|
||||
|
||||
result = await subprocess_tools.BashTool().execute(
|
||||
"pwd",
|
||||
{"subproc_env": {}, "session_id": "chat-1"},
|
||||
)
|
||||
|
||||
assert result == {"output": "ok", "exit_code": 0}
|
||||
assert captured["command"] == "pwd"
|
||||
assert captured["kwargs"]["cwd"] == workspace
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_posix_bash_keeps_existing_shell_path(monkeypatch):
|
||||
captured = {}
|
||||
process = object()
|
||||
|
||||
monkeypatch.setattr(subprocess_tools, "IS_WINDOWS", False)
|
||||
|
||||
async def fake_shell(command, **kwargs):
|
||||
captured["command"] = command
|
||||
captured["kwargs"] = kwargs
|
||||
return process
|
||||
|
||||
async def fail_exec(*_args, **_kwargs):
|
||||
pytest.fail("POSIX behavior must continue through create_subprocess_shell")
|
||||
|
||||
monkeypatch.setattr(subprocess_tools.asyncio, "create_subprocess_shell", fake_shell)
|
||||
monkeypatch.setattr(subprocess_tools.asyncio, "create_subprocess_exec", fail_exec)
|
||||
|
||||
result = await subprocess_tools._create_bash_subprocess("pwd", cwd="/tmp/work")
|
||||
|
||||
assert result is process
|
||||
assert captured == {"command": "pwd", "kwargs": {"cwd": "/tmp/work"}}
|
||||
@@ -314,17 +314,21 @@ class TestComputeFinalMetrics:
|
||||
def test_tool_events_included(self):
|
||||
events = [{"tool": "bash", "duration": 1.0}]
|
||||
texts = ["round 1 text"]
|
||||
models = ["round-1-model"]
|
||||
m = _compute_final_metrics(**self._base_args(
|
||||
tool_events=events,
|
||||
round_texts=texts,
|
||||
round_models=models,
|
||||
))
|
||||
assert m["tool_events"] == events
|
||||
assert m["round_texts"] == texts
|
||||
assert m["round_models"] == models
|
||||
|
||||
def test_no_tool_events_excluded(self):
|
||||
m = _compute_final_metrics(**self._base_args(tool_events=[], round_texts=[]))
|
||||
assert "tool_events" not in m
|
||||
assert "round_texts" not in m
|
||||
assert "round_models" not in m
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
"""Saved Agent rounds must render and bill with actual per-round provenance."""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
_SOURCE = (
|
||||
Path(__file__).resolve().parents[1] / "static" / "js" / "chatRenderer.js"
|
||||
).read_text(encoding="utf-8")
|
||||
_CHAT_SOURCE = (
|
||||
Path(__file__).resolve().parents[1] / "static" / "js" / "chat.js"
|
||||
).read_text(encoding="utf-8")
|
||||
_SLASH_SOURCE = (
|
||||
Path(__file__).resolve().parents[1] / "static" / "js" / "slashCommands.js"
|
||||
).read_text(encoding="utf-8")
|
||||
_HAS_NODE = shutil.which("node") is not None
|
||||
|
||||
|
||||
def _function_source(name):
|
||||
match = re.search(
|
||||
rf"^(?:export )?function {name}\(.*?^\}}",
|
||||
_SOURCE,
|
||||
re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
assert match, f"{name} not found"
|
||||
return match.group(0).replace("export function", "function", 1)
|
||||
|
||||
|
||||
def _run_node(source):
|
||||
proc = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=source,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
)
|
||||
assert proc.returncode == 0, proc.stderr
|
||||
return json.loads(proc.stdout.strip())
|
||||
|
||||
|
||||
def test_saved_agent_rounds_prefer_round_model_provenance():
|
||||
assert "const roundModels = metadata.round_models || [];" in _SOURCE
|
||||
assert "const contModel = roundModels[r] || pair.actualModel || pair.requestedModel;" in _SOURCE
|
||||
assert "Array.isArray(metadata.round_texts) && metadata.round_texts.length > 1" in _SOURCE
|
||||
assert "const roundEndpointIds = metadata.round_endpoint_ids || [];" in _SOURCE
|
||||
assert "const roundEndpointLabels = metadata.round_endpoint_labels || [];" in _SOURCE
|
||||
assert "r < roundEndpointIds.length" in _SOURCE
|
||||
assert "r < roundEndpointLabels.length" in _SOURCE
|
||||
assert "roundEndpointIds[r] || pair.actualEndpointId" not in _SOURCE
|
||||
|
||||
|
||||
def test_metrics_cost_uses_actual_fallback_endpoint_classification():
|
||||
assert "metrics.endpoint_cost_tracked" in _SOURCE
|
||||
assert "endpointCostTracked === false" in _SOURCE
|
||||
assert "endpointCostTracked !== true && !isCostTrackedEndpoint(selectedUrl)" in _SOURCE
|
||||
assert "Array.isArray(metrics.usage_buckets)" in _SOURCE
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
def test_agent_usage_buckets_sum_only_billable_answering_routes():
|
||||
source = "\n".join([
|
||||
"let currentUrl = '';",
|
||||
"function _currentEndpointUrl() { return currentUrl; }",
|
||||
"function isCostTrackedEndpoint(url) { return url === 'paid'; }",
|
||||
"function getModelCost(_model, inputTokens, outputTokens) { return (inputTokens + outputTokens) / 1000; }",
|
||||
_function_source("_billableCost"),
|
||||
_function_source("_metricsBillableCost"),
|
||||
"const paidSelected = {usage_buckets: [",
|
||||
" {model: 'selected', input_tokens: 100, output_tokens: 10, endpoint_cost_tracked: true},",
|
||||
" {model: 'local-fallback', input_tokens: 200, output_tokens: 20, endpoint_cost_tracked: false},",
|
||||
"]};",
|
||||
"const localSelected = {usage_buckets: [",
|
||||
" {model: 'selected', input_tokens: 100, output_tokens: 10, endpoint_cost_tracked: false},",
|
||||
" {model: 'paid-fallback', input_tokens: 200, output_tokens: 20, endpoint_cost_tracked: true},",
|
||||
"]};",
|
||||
"currentUrl = 'local';",
|
||||
"const paidToLocal = _metricsBillableCost(paidSelected, 'final', 300, 30);",
|
||||
"currentUrl = 'paid';",
|
||||
"const localToPaid = _metricsBillableCost(localSelected, 'final', 300, 30);",
|
||||
"console.log(JSON.stringify({paidToLocal, localToPaid}));",
|
||||
])
|
||||
|
||||
assert _run_node(source) == {"paidToLocal": 0.11, "localToPaid": 0.22}
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
def test_force_answer_synthesis_segment_is_included_in_fallback_cost():
|
||||
source = "\n".join([
|
||||
"function _currentEndpointUrl() { return 'local-selected'; }",
|
||||
"function isCostTrackedEndpoint() { return false; }",
|
||||
"function getModelCost(_model, inputTokens, outputTokens) { return (inputTokens + outputTokens) / 1000; }",
|
||||
_function_source("_billableCost"),
|
||||
_function_source("_metricsBillableCost"),
|
||||
"const metrics = {usage_buckets: [",
|
||||
" {round: 6, model: 'paid-fallback', input_tokens: 100, output_tokens: 0, endpoint_cost_tracked: true},",
|
||||
" {round: 6, model: 'paid-fallback', input_tokens: 80, output_tokens: 20, endpoint_cost_tracked: true},",
|
||||
"]};",
|
||||
"console.log(JSON.stringify({cost: _metricsBillableCost(metrics, 'paid-fallback', 180, 20)}));",
|
||||
])
|
||||
|
||||
assert _run_node(source) == {"cost": 0.2}
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
def test_repeated_live_metrics_render_records_session_cost_once():
|
||||
source = "\n".join([
|
||||
"const _COST_KEY = 'ody-session-cost';",
|
||||
"const state = {};",
|
||||
"const localStorage = {",
|
||||
" getItem(key) { return state[key] || null; },",
|
||||
" setItem(key, value) { state[key] = value; },",
|
||||
"};",
|
||||
"const window = {sessionModule: {getCurrentSessionId() { return 'session'; }}};",
|
||||
"function updateSessionCostUI() {}",
|
||||
"function _currentEndpointUrl() { return 'local'; }",
|
||||
"function isCostTrackedEndpoint(url) { return url === 'paid'; }",
|
||||
"function getModelCost(_model, inputTokens, outputTokens) { return (inputTokens + outputTokens) / 1000; }",
|
||||
_function_source("_billableCost"),
|
||||
_function_source("_metricsBillableCost"),
|
||||
_function_source("recordSessionMetricsCost"),
|
||||
"const metrics = {model: 'paid-model', input_tokens: 100, output_tokens: 10, endpoint_cost_tracked: true};",
|
||||
"recordSessionMetricsCost(metrics);",
|
||||
"recordSessionMetricsCost(metrics);",
|
||||
"console.log(JSON.stringify({cost: JSON.parse(state[_COST_KEY]).session, recorded: metrics._costRecorded}));",
|
||||
])
|
||||
|
||||
assert _run_node(source) == {"cost": 0.11, "recorded": True}
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
def test_replayed_metrics_use_run_identity_for_durable_cost_deduplication():
|
||||
source = "\n".join([
|
||||
"const _COST_KEY = 'ody-session-cost';",
|
||||
"const _COST_RUNS_KEY = 'ody-session-cost-runs';",
|
||||
"const _MAX_COST_RUNS_PER_SESSION = 256;",
|
||||
"const state = {};",
|
||||
"const localStorage = {",
|
||||
" getItem(key) { return state[key] || null; },",
|
||||
" setItem(key, value) { state[key] = value; },",
|
||||
"};",
|
||||
"const window = {sessionModule: {getCurrentSessionId() { return 'session'; }}};",
|
||||
"function updateSessionCostUI() {}",
|
||||
"function _currentEndpointUrl() { return 'local'; }",
|
||||
"function isCostTrackedEndpoint(url) { return url === 'paid'; }",
|
||||
"function getModelCost(_model, inputTokens, outputTokens) { return (inputTokens + outputTokens) / 1000; }",
|
||||
_function_source("_billableCost"),
|
||||
_function_source("_metricsBillableCost"),
|
||||
_function_source("recordSessionMetricsCost"),
|
||||
_function_source("getSessionCost"),
|
||||
"const firstObject = {model: 'paid-model', input_tokens: 100, output_tokens: 10, endpoint_cost_tracked: true, _costRecordId: 'run-1'};",
|
||||
"const replayedObject = {...firstObject};",
|
||||
"recordSessionMetricsCost(firstObject);",
|
||||
"recordSessionMetricsCost(replayedObject);",
|
||||
"console.log(JSON.stringify({cost: getSessionCost('session'), runs: JSON.parse(state[_COST_RUNS_KEY]).session}));",
|
||||
])
|
||||
|
||||
assert _run_node(source) == {"cost": 0.11, "runs": {"run-1": 0.11}}
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
def test_run_cost_ledger_sums_segments_and_updates_repeated_segment_metrics():
|
||||
source = "\n".join([
|
||||
"const _COST_KEY = 'ody-session-cost';",
|
||||
"const _COST_RUNS_KEY = 'ody-session-cost-runs';",
|
||||
"const _MAX_COST_RUNS_PER_SESSION = 256;",
|
||||
"const state = {};",
|
||||
"const localStorage = {",
|
||||
" getItem(key) { return state[key] || null; },",
|
||||
" setItem(key, value) { state[key] = value; },",
|
||||
"};",
|
||||
"const window = {sessionModule: {getCurrentSessionId() { return 'session'; }}};",
|
||||
"function updateSessionCostUI() {}",
|
||||
"function _currentEndpointUrl() { return 'paid'; }",
|
||||
"function isCostTrackedEndpoint() { return true; }",
|
||||
"function getModelCost(_model, inputTokens, outputTokens) { return (inputTokens + outputTokens) / 1000; }",
|
||||
_function_source("_billableCost"),
|
||||
_function_source("_metricsBillableCost"),
|
||||
_function_source("recordSessionMetricsCost"),
|
||||
_function_source("getSessionCost"),
|
||||
"recordSessionMetricsCost({model: 'student', input_tokens: 100, output_tokens: 10, _costRecordId: 'run:primary'});",
|
||||
"recordSessionMetricsCost({model: 'student', input_tokens: 120, output_tokens: 20, _costRecordId: 'run:primary'});",
|
||||
"recordSessionMetricsCost({model: 'teacher', input_tokens: 200, output_tokens: 30, _costRecordId: 'run:teacher'});",
|
||||
"console.log(JSON.stringify({cost: getSessionCost('session'), runs: JSON.parse(state[_COST_RUNS_KEY]).session}));",
|
||||
])
|
||||
|
||||
assert _run_node(source) == {
|
||||
"cost": 0.37,
|
||||
"runs": {"run:primary": 0.14, "run:teacher": 0.23},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
def test_local_selected_endpoint_does_not_erase_paid_fallback_ledger():
|
||||
source = "\n".join([
|
||||
"const _COST_KEY = 'ody-session-cost';",
|
||||
"const _COST_RUNS_KEY = 'ody-session-cost-runs';",
|
||||
"const state = {'ody-session-cost': JSON.stringify({session: 0.125})};",
|
||||
"const localStorage = {",
|
||||
" getItem(key) { return state[key] || null; },",
|
||||
" setItem(key, value) { state[key] = value; },",
|
||||
"};",
|
||||
"const badge = {style: {}, textContent: ''};",
|
||||
"const document = {getElementById() { return badge; }};",
|
||||
"const window = {sessionModule: {getCurrentSessionId() { return 'session'; }, getCurrentEndpointUrl() { return 'local'; }}};",
|
||||
_function_source("getSessionCost"),
|
||||
_function_source("updateSessionCostUI"),
|
||||
"updateSessionCostUI();",
|
||||
"console.log(JSON.stringify({stored: JSON.parse(state[_COST_KEY]).session, display: badge.style.display, text: badge.textContent}));",
|
||||
])
|
||||
|
||||
assert _run_node(source) == {
|
||||
"stored": 0.125,
|
||||
"display": "",
|
||||
"text": "$0.125",
|
||||
}
|
||||
|
||||
|
||||
def test_live_and_resumed_terminal_events_apply_usage_metrics_before_reload():
|
||||
assert "metrics = json.data || metrics;" in _CHAT_SOURCE
|
||||
assert "displayMetrics(terminalMetricsTarget, metrics);" in _CHAT_SOURCE
|
||||
assert "metricsData = json.data || metricsData;" in _CHAT_SOURCE
|
||||
assert "displayMetrics(holder, metricsData);" in _CHAT_SOURCE
|
||||
assert "json.type === 'agent_terminal' || json.type === 'chat_terminal'" in _CHAT_SOURCE
|
||||
assert "chatRenderer.recordSessionMetricsCost(metrics, streamSessionId);" in _CHAT_SOURCE
|
||||
assert "chatRenderer.recordSessionMetricsCost(metricsData, sessionId);" in _CHAT_SOURCE
|
||||
assert "metricsData._costRecordId = _metricsCostRecordId(resumeRunId, json);" in _CHAT_SOURCE
|
||||
assert "bgTerminal.status = 'completed';" in _CHAT_SOURCE
|
||||
|
||||
|
||||
def test_usage_command_does_not_hide_existing_fallback_cost_for_local_selection():
|
||||
assert "const cost = chatRenderer.getSessionCost" in _SLASH_SOURCE
|
||||
assert "const cost = costTracked && chatRenderer.getSessionCost" not in _SLASH_SOURCE
|
||||
+44
-6
@@ -1,17 +1,19 @@
|
||||
"""Tests for ``core.atomic_io`` durability and crash-safety behavior.
|
||||
|
||||
``core.atomic_io`` provides ``atomic_write_json`` and ``atomic_write_text``.
|
||||
Both write to a sibling ``.tmp.<pid>`` file, ``fsync`` it, then ``os.replace``
|
||||
into place so a crash mid-write leaves the previous good copy untouched rather
|
||||
than a truncated/empty file.
|
||||
Both write to a sibling ``.tmp.<random>`` file, ``fsync`` it, then
|
||||
``os.replace`` into place so a crash mid-write leaves the previous good copy
|
||||
untouched rather than a truncated/empty file.
|
||||
|
||||
These tests cover the happy path (round-trip, indent, parent-dir creation,
|
||||
full overwrite, no leftover tmp) and the two failure paths the implementation
|
||||
guarantees: the target file is preserved when serialization fails before the
|
||||
replace, and when ``os.replace`` itself fails.
|
||||
full overwrite, no leftover tmp), the two failure paths the implementation
|
||||
guarantees (the target file is preserved when serialization fails before the
|
||||
replace, and when ``os.replace`` itself fails), and that two concurrent
|
||||
writers to the same path don't collide on the same temp file.
|
||||
"""
|
||||
import importlib.util
|
||||
import json
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
@@ -84,6 +86,42 @@ def test_atomic_write_json_leaves_no_tmp_file(tmp_path):
|
||||
assert _tmp_siblings(tmp_path, "data.json") == []
|
||||
|
||||
|
||||
def test_atomic_write_json_concurrent_writers_do_not_collide(tmp_path):
|
||||
# Both writers run in this same process, so a PID-based tmp suffix is
|
||||
# identical for both: whichever writer finishes first unlinks the tmp
|
||||
# file (via os.replace) out from under the other, which then raises
|
||||
# FileNotFoundError on its own os.replace instead of landing its write.
|
||||
target = tmp_path / "settings.json"
|
||||
orig_dump = json.dump
|
||||
barrier = threading.Barrier(2)
|
||||
errors = []
|
||||
|
||||
def slow_dump(obj, fp, **kwargs):
|
||||
orig_dump(obj, fp, **kwargs)
|
||||
fp.flush()
|
||||
barrier.wait()
|
||||
|
||||
def write(payload):
|
||||
try:
|
||||
atomic_write_json(str(target), payload)
|
||||
except Exception as exc: # noqa: BLE001 - captured for the assertion below
|
||||
errors.append(exc)
|
||||
|
||||
json.dump = slow_dump
|
||||
try:
|
||||
t1 = threading.Thread(target=write, args=({"writer": "A"},))
|
||||
t2 = threading.Thread(target=write, args=({"writer": "B"},))
|
||||
t1.start()
|
||||
t2.start()
|
||||
t1.join()
|
||||
t2.join()
|
||||
finally:
|
||||
json.dump = orig_dump
|
||||
|
||||
assert errors == []
|
||||
assert json.loads(target.read_text(encoding="utf-8"))["writer"] in ("A", "B")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# atomic_write_json — failure path: target preserved on serialization error.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -0,0 +1,538 @@
|
||||
"""Default calendar creation belongs to the caller's transaction.
|
||||
|
||||
Before this regression, ``_ensure_default_calendar`` committed independently.
|
||||
If event persistence then failed, the event rolled back but a new ``Personal``
|
||||
calendar remained (``calendar_count=1``, ``event_count=0``).
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import create_engine, event
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
from sqlalchemy.pool import NullPool
|
||||
|
||||
from tests.helpers.import_state import clear_fake_database_modules
|
||||
|
||||
clear_fake_database_modules()
|
||||
|
||||
import core.database as cdb # noqa: E402
|
||||
import routes.calendar_routes as calendar_routes # noqa: E402
|
||||
from core.database import CalendarCal, CalendarEvent # noqa: E402
|
||||
from routes.calendar_routes import EventCreate # noqa: E402
|
||||
from routes.calendar_routes import ( # noqa: E402
|
||||
_default_calendar_id,
|
||||
_ensure_default_calendar,
|
||||
)
|
||||
|
||||
|
||||
class _RejectEventCommit(Session):
|
||||
"""Reproduce an event commit failure after default-calendar creation."""
|
||||
|
||||
def commit(self):
|
||||
if any(isinstance(row, CalendarEvent) for row in self.new):
|
||||
raise RuntimeError("commit guard rejected event commit")
|
||||
return super().commit()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def session_factory(tmp_path, monkeypatch):
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'calendar.db'}",
|
||||
connect_args={"check_same_thread": False},
|
||||
poolclass=NullPool,
|
||||
)
|
||||
cdb.Base.metadata.create_all(engine)
|
||||
factory = sessionmaker(
|
||||
bind=engine,
|
||||
autoflush=False,
|
||||
autocommit=False,
|
||||
class_=_RejectEventCommit,
|
||||
)
|
||||
monkeypatch.setattr(cdb, "SessionLocal", factory)
|
||||
monkeypatch.setattr(calendar_routes, "SessionLocal", factory)
|
||||
try:
|
||||
yield factory
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def _request():
|
||||
return SimpleNamespace(state=SimpleNamespace(current_user="alice"))
|
||||
|
||||
|
||||
def _endpoint(method, suffix):
|
||||
router = calendar_routes.setup_calendar_routes()
|
||||
for route in router.routes:
|
||||
if route.path.endswith(suffix) and method in route.methods:
|
||||
return route.endpoint
|
||||
raise RuntimeError(f"{method} *{suffix} not found")
|
||||
|
||||
|
||||
def _counts(factory):
|
||||
db = factory()
|
||||
try:
|
||||
return db.query(CalendarCal).count(), db.query(CalendarEvent).count()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
async def test_route_event_failure_rolls_back_new_default_calendar(session_factory):
|
||||
create_event = _endpoint("POST", "/events")
|
||||
|
||||
with pytest.raises(HTTPException) as caught:
|
||||
await create_event(
|
||||
_request(),
|
||||
EventCreate(summary="Planning", dtstart="2126-07-20T09:00:00Z"),
|
||||
)
|
||||
|
||||
assert caught.value.status_code == 500
|
||||
assert _counts(session_factory) == (0, 0)
|
||||
|
||||
|
||||
async def test_route_event_validation_failure_rolls_back_new_default_calendar(
|
||||
session_factory,
|
||||
):
|
||||
create_event = _endpoint("POST", "/events")
|
||||
|
||||
with pytest.raises(HTTPException) as caught:
|
||||
await create_event(
|
||||
_request(),
|
||||
EventCreate(summary="Planning", dtstart="not-a-datetime"),
|
||||
)
|
||||
|
||||
assert caught.value.status_code == 500
|
||||
assert _counts(session_factory) == (0, 0)
|
||||
|
||||
|
||||
async def test_tool_event_failure_rolls_back_new_default_calendar(session_factory):
|
||||
from src.tools.calendar import do_manage_calendar
|
||||
|
||||
result = await do_manage_calendar(
|
||||
json.dumps({
|
||||
"action": "create_event",
|
||||
"summary": "Planning",
|
||||
"dtstart": "2126-07-20T09:00:00Z",
|
||||
}),
|
||||
owner="alice",
|
||||
)
|
||||
|
||||
assert result["exit_code"] == 1
|
||||
assert "commit guard rejected event commit" in result["error"]
|
||||
assert _counts(session_factory) == (0, 0)
|
||||
|
||||
|
||||
async def test_tool_event_validation_failure_rolls_back_new_default_calendar(
|
||||
session_factory,
|
||||
):
|
||||
from src.tools.calendar import do_manage_calendar
|
||||
|
||||
result = await do_manage_calendar(
|
||||
json.dumps({
|
||||
"action": "create_event",
|
||||
"summary": "Planning",
|
||||
"dtstart": "not-a-datetime",
|
||||
}),
|
||||
owner="alice",
|
||||
)
|
||||
|
||||
assert result["exit_code"] == 1
|
||||
assert "Could not parse dtstart" in result["error"]
|
||||
assert _counts(session_factory) == (0, 0)
|
||||
|
||||
|
||||
async def test_route_list_calendars_persists_lazy_default(session_factory):
|
||||
list_calendars = _endpoint("GET", "/calendars")
|
||||
|
||||
result = await list_calendars(_request())
|
||||
|
||||
assert [calendar["name"] for calendar in result["calendars"]] == ["Personal"]
|
||||
assert _counts(session_factory) == (1, 0)
|
||||
|
||||
|
||||
async def test_tool_list_calendars_persists_lazy_default(session_factory):
|
||||
from src.tools.calendar import do_manage_calendar
|
||||
|
||||
result = await do_manage_calendar(
|
||||
json.dumps({"action": "list_calendars"}),
|
||||
owner="alice",
|
||||
)
|
||||
|
||||
assert result["exit_code"] == 0
|
||||
assert [calendar["name"] for calendar in result["calendars"]] == ["Personal"]
|
||||
assert _counts(session_factory) == (1, 0)
|
||||
|
||||
|
||||
def test_repeated_rename_and_reuse_uses_stable_fallback_ids(tmp_path):
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'renamed-calendar.db'}",
|
||||
connect_args={"check_same_thread": False},
|
||||
poolclass=NullPool,
|
||||
)
|
||||
cdb.Base.metadata.create_all(engine)
|
||||
factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
|
||||
db = factory()
|
||||
try:
|
||||
first = _ensure_default_calendar(db, "alice")
|
||||
assert first.id == _default_calendar_id("alice")
|
||||
db.commit()
|
||||
|
||||
# The supported user-rename migration changes owner columns while
|
||||
# deliberately preserving durable row identifiers.
|
||||
first.owner = "bob"
|
||||
db.commit()
|
||||
|
||||
second = _ensure_default_calendar(db, "alice")
|
||||
assert second.id == _default_calendar_id("alice", 1)
|
||||
db.commit()
|
||||
|
||||
# Repeating the same lifecycle must advance deterministically instead
|
||||
# of failing or choosing a random identifier.
|
||||
second.owner = "carol"
|
||||
db.commit()
|
||||
|
||||
third = _ensure_default_calendar(db, "alice")
|
||||
assert third.id == _default_calendar_id("alice", 2)
|
||||
db.commit()
|
||||
|
||||
rows = db.query(CalendarCal).order_by(CalendarCal.owner).all()
|
||||
assert [(row.owner, row.id) for row in rows] == [
|
||||
("alice", _default_calendar_id("alice", 2)),
|
||||
("bob", _default_calendar_id("alice")),
|
||||
("carol", _default_calendar_id("alice", 1)),
|
||||
]
|
||||
finally:
|
||||
db.close()
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def _assert_concurrent_first_use(tmp_path, occupied_owner=None):
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'concurrent-calendar.db'}",
|
||||
connect_args={"check_same_thread": False, "timeout": 10},
|
||||
poolclass=NullPool,
|
||||
)
|
||||
cdb.Base.metadata.create_all(engine)
|
||||
factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
|
||||
expected_collision_index = 0
|
||||
if occupied_owner is not None:
|
||||
seed = factory()
|
||||
try:
|
||||
seed.add(CalendarCal(
|
||||
id=_default_calendar_id("alice"),
|
||||
owner=occupied_owner,
|
||||
name="Personal",
|
||||
source="local",
|
||||
))
|
||||
seed.commit()
|
||||
expected_collision_index = 1
|
||||
finally:
|
||||
seed.close()
|
||||
first_staged = threading.Event()
|
||||
second_selected = threading.Event()
|
||||
errors = []
|
||||
|
||||
@event.listens_for(engine, "after_cursor_execute")
|
||||
def observe_second_gap(conn, cursor, statement, parameters, context, executemany):
|
||||
if (
|
||||
threading.current_thread().name == "calendar-worker-second"
|
||||
and statement.lstrip().upper().startswith("SELECT")
|
||||
and "FROM calendars" in statement
|
||||
):
|
||||
second_selected.set()
|
||||
|
||||
def create_default(worker, hold=False):
|
||||
db = factory()
|
||||
try:
|
||||
if not hold:
|
||||
assert first_staged.wait(5)
|
||||
cal = _ensure_default_calendar(db, "alice")
|
||||
start = datetime(2126, 7, 20, 9 if hold else 10)
|
||||
db.add(CalendarEvent(
|
||||
uid=worker,
|
||||
calendar_id=cal.id,
|
||||
summary=f"Event {worker}",
|
||||
dtstart=start,
|
||||
dtend=start + timedelta(hours=1),
|
||||
))
|
||||
if hold:
|
||||
first_staged.set()
|
||||
# The second session has observed the uncommitted gap before
|
||||
# this transaction releases its writer reservation.
|
||||
assert second_selected.wait(5)
|
||||
db.commit()
|
||||
assert cal.id == _default_calendar_id("alice", expected_collision_index)
|
||||
except BaseException as exc: # pragma: no cover - asserted below
|
||||
errors.append((worker, exc))
|
||||
db.rollback()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
first = threading.Thread(
|
||||
target=create_default,
|
||||
args=("first", True),
|
||||
name="calendar-worker-first",
|
||||
)
|
||||
second = threading.Thread(
|
||||
target=create_default,
|
||||
args=("second",),
|
||||
name="calendar-worker-second",
|
||||
)
|
||||
first.start()
|
||||
second.start()
|
||||
first.join(10)
|
||||
second.join(10)
|
||||
|
||||
try:
|
||||
assert not first.is_alive() and not second.is_alive()
|
||||
assert errors == []
|
||||
db = factory()
|
||||
try:
|
||||
rows = db.query(CalendarCal).filter(CalendarCal.owner == "alice").all()
|
||||
assert [(row.id, row.name) for row in rows] == [
|
||||
(_default_calendar_id("alice", expected_collision_index), "Personal")
|
||||
]
|
||||
assert db.query(CalendarEvent).count() == 2
|
||||
if occupied_owner is not None:
|
||||
occupied = db.query(CalendarCal).filter(
|
||||
CalendarCal.id == _default_calendar_id("alice"),
|
||||
).one()
|
||||
assert occupied.owner == occupied_owner
|
||||
finally:
|
||||
db.close()
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_concurrent_first_use_creates_one_sqlite_default(tmp_path):
|
||||
_assert_concurrent_first_use(tmp_path)
|
||||
|
||||
|
||||
def test_concurrent_first_use_after_rename_creates_one_fallback_default(tmp_path):
|
||||
_assert_concurrent_first_use(tmp_path, occupied_owner="bob")
|
||||
|
||||
|
||||
def test_sqlite_default_stays_in_callers_transaction(session_factory):
|
||||
db = session_factory()
|
||||
try:
|
||||
cal = _ensure_default_calendar(db, "rollback-owner")
|
||||
assert cal.id == _default_calendar_id("rollback-owner")
|
||||
db.rollback()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
verify = session_factory()
|
||||
try:
|
||||
assert (
|
||||
verify.query(CalendarCal)
|
||||
.filter(CalendarCal.owner == "rollback-owner")
|
||||
.count()
|
||||
== 0
|
||||
)
|
||||
finally:
|
||||
verify.close()
|
||||
|
||||
|
||||
def test_sqlite_fallback_default_stays_in_callers_transaction(session_factory):
|
||||
seed = session_factory()
|
||||
try:
|
||||
seed.add(CalendarCal(
|
||||
id=_default_calendar_id("alice"),
|
||||
owner="bob",
|
||||
name="Personal",
|
||||
source="local",
|
||||
))
|
||||
seed.commit()
|
||||
finally:
|
||||
seed.close()
|
||||
|
||||
db = session_factory()
|
||||
try:
|
||||
cal = _ensure_default_calendar(db, "alice")
|
||||
assert cal.id == _default_calendar_id("alice", 1)
|
||||
db.rollback()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
verify = session_factory()
|
||||
try:
|
||||
assert verify.query(CalendarCal).filter(CalendarCal.owner == "alice").count() == 0
|
||||
assert verify.query(CalendarCal).filter(CalendarCal.owner == "bob").count() == 1
|
||||
finally:
|
||||
verify.close()
|
||||
|
||||
|
||||
class _FakeDialect:
|
||||
name = "postgresql"
|
||||
|
||||
|
||||
class _FakeBind:
|
||||
dialect = _FakeDialect()
|
||||
|
||||
|
||||
class _FakeQuery:
|
||||
def __init__(self, session):
|
||||
self.session = session
|
||||
|
||||
def filter(self, *conditions):
|
||||
return self
|
||||
|
||||
def with_for_update(self):
|
||||
self.session.locking_read = True
|
||||
return self
|
||||
|
||||
def first(self):
|
||||
self.session.query_count += 1
|
||||
if self.session.query_count == 1:
|
||||
return None
|
||||
return self.session.winner
|
||||
|
||||
|
||||
class _GenericRaceSession:
|
||||
"""Minimal non-SQLite session that loses the deterministic-ID race."""
|
||||
|
||||
def __init__(self):
|
||||
self.query_count = 0
|
||||
self.nested_entries = 0
|
||||
self.locking_read = False
|
||||
self.candidate = None
|
||||
self.winner = CalendarCal(
|
||||
id=_default_calendar_id("alice"),
|
||||
owner="alice",
|
||||
name="Personal",
|
||||
source="local",
|
||||
)
|
||||
|
||||
def get_bind(self):
|
||||
return _FakeBind()
|
||||
|
||||
def query(self, model):
|
||||
assert model is CalendarCal
|
||||
return _FakeQuery(self)
|
||||
|
||||
@contextmanager
|
||||
def begin_nested(self):
|
||||
self.nested_entries += 1
|
||||
yield
|
||||
|
||||
def add(self, row):
|
||||
self.candidate = row
|
||||
|
||||
def flush(self):
|
||||
raise IntegrityError("insert", {}, RuntimeError("duplicate primary key"))
|
||||
|
||||
|
||||
def test_generic_backend_lost_race_recovers_inside_savepoint():
|
||||
db = _GenericRaceSession()
|
||||
|
||||
winner = _ensure_default_calendar(db, "alice")
|
||||
|
||||
assert winner is db.winner
|
||||
assert db.nested_entries == 1
|
||||
assert db.locking_read is True
|
||||
assert db.candidate.id == db.winner.id
|
||||
|
||||
|
||||
def test_generic_backend_unattributed_integrity_error_is_not_retried():
|
||||
db = _GenericRaceSession()
|
||||
db.winner = None
|
||||
|
||||
with pytest.raises(IntegrityError):
|
||||
_ensure_default_calendar(db, "alice")
|
||||
|
||||
assert db.nested_entries == 1
|
||||
|
||||
|
||||
class _GenericRenamedSlotSession(_GenericRaceSession):
|
||||
"""A different owner occupies slot zero; slot one remains available."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.candidates = []
|
||||
self.winner = CalendarCal(
|
||||
id=_default_calendar_id("alice"),
|
||||
owner="bob",
|
||||
name="Personal",
|
||||
source="local",
|
||||
)
|
||||
|
||||
def add(self, row):
|
||||
self.candidate = row
|
||||
self.candidates.append(row)
|
||||
|
||||
def flush(self):
|
||||
if len(self.candidates) == 1:
|
||||
raise IntegrityError("insert", {}, RuntimeError("duplicate primary key"))
|
||||
|
||||
|
||||
def test_generic_backend_renamed_slot_advances_inside_savepoint():
|
||||
db = _GenericRenamedSlotSession()
|
||||
|
||||
fallback = _ensure_default_calendar(db, "alice")
|
||||
|
||||
assert fallback is db.candidates[-1]
|
||||
assert fallback.id == _default_calendar_id("alice", 1)
|
||||
assert fallback.owner == "alice"
|
||||
assert db.nested_entries == 2
|
||||
assert db.locking_read is True
|
||||
assert db.winner.owner == "bob"
|
||||
|
||||
|
||||
def test_generic_backend_fallback_keeps_outer_transaction_usable(tmp_path):
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'generic-savepoint-calendar.db'}",
|
||||
poolclass=NullPool,
|
||||
)
|
||||
cdb.Base.metadata.create_all(engine)
|
||||
# SQLite supplies a lightweight local SQL executor here; changing only the
|
||||
# dispatch name exercises the real Session/savepoint branch used by
|
||||
# PostgreSQL-style backends without pretending to validate their dialect.
|
||||
engine.dialect.name = "postgresql"
|
||||
factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
|
||||
|
||||
seed = factory()
|
||||
try:
|
||||
seed.add(CalendarCal(
|
||||
id=_default_calendar_id("alice"),
|
||||
owner="bob",
|
||||
name="Personal",
|
||||
source="local",
|
||||
))
|
||||
seed.commit()
|
||||
finally:
|
||||
seed.close()
|
||||
|
||||
db = factory()
|
||||
try:
|
||||
cal = _ensure_default_calendar(db, "alice")
|
||||
start = datetime(2126, 7, 20, 9)
|
||||
db.add(CalendarEvent(
|
||||
uid="after-fallback",
|
||||
calendar_id=cal.id,
|
||||
summary="Atomic",
|
||||
dtstart=start,
|
||||
dtend=start + timedelta(hours=1),
|
||||
))
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
verify = factory()
|
||||
try:
|
||||
assert [
|
||||
(row.owner, row.id)
|
||||
for row in verify.query(CalendarCal).order_by(CalendarCal.owner).all()
|
||||
] == [
|
||||
("alice", _default_calendar_id("alice", 1)),
|
||||
("bob", _default_calendar_id("alice")),
|
||||
]
|
||||
assert verify.query(CalendarEvent).count() == 1
|
||||
finally:
|
||||
verify.close()
|
||||
engine.dispose()
|
||||
@@ -0,0 +1,203 @@
|
||||
"""Execute the round-aware live model-provenance state helper under Node."""
|
||||
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
_REPO = Path(__file__).resolve().parents[1]
|
||||
_MODULE = (_REPO / "static" / "js" / "chatModelProvenance.js").as_uri()
|
||||
|
||||
|
||||
def test_round_two_fallback_then_provider_alias_does_not_relabel_round_one():
|
||||
if not shutil.which("node"):
|
||||
pytest.skip("node is not installed")
|
||||
|
||||
script = f"""
|
||||
import {{ applyModelRouteEventState }} from {json.dumps(_MODULE)};
|
||||
const round1 = {{ _requestedModel: 'selected-model', _actualModel: 'selected-model' }};
|
||||
const round2 = {{ _requestedModel: 'selected-model', _actualModel: 'selected-model' }};
|
||||
|
||||
const fallbackTarget = applyModelRouteEventState({{
|
||||
type: 'fallback', round: 2,
|
||||
selected_model: 'selected-model', answered_by: 'backup-model'
|
||||
}}, round1, round2, 'selected-model');
|
||||
const aliasTarget = applyModelRouteEventState({{
|
||||
type: 'model_actual', round: 2,
|
||||
requested_model: 'selected-model', model: 'provider-backup-alias'
|
||||
}}, round1, round2, 'selected-model');
|
||||
|
||||
console.log(JSON.stringify({{
|
||||
fallbackIsRound2: fallbackTarget === round2,
|
||||
aliasIsRound2: aliasTarget === round2,
|
||||
round1,
|
||||
round2,
|
||||
}}));
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=script,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=_REPO,
|
||||
timeout=30,
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
state = json.loads(result.stdout)
|
||||
assert state == {
|
||||
"fallbackIsRound2": True,
|
||||
"aliasIsRound2": True,
|
||||
"round1": {
|
||||
"_requestedModel": "selected-model",
|
||||
"_actualModel": "selected-model",
|
||||
},
|
||||
"round2": {
|
||||
"_requestedModel": "selected-model",
|
||||
"_actualModel": "provider-backup-alias",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_next_round_and_final_metrics_preserve_each_agent_round_route():
|
||||
if not shutil.which("node"):
|
||||
pytest.skip("node is not installed")
|
||||
|
||||
script = f"""
|
||||
import {{
|
||||
applyModelMetricsState,
|
||||
applyModelRouteEventState,
|
||||
inheritModelRouteState,
|
||||
}} from {json.dumps(_MODULE)};
|
||||
const round1 = {{ _requestedModel: 'selected-model', _actualModel: 'selected-model' }};
|
||||
const round2 = {{}};
|
||||
inheritModelRouteState(round1, round1, round2, 'selected-model');
|
||||
applyModelRouteEventState({{
|
||||
type: 'fallback', round: 2,
|
||||
selected_model: 'selected-model', answered_by: 'backup-model'
|
||||
}}, round1, round2, 'selected-model');
|
||||
applyModelRouteEventState({{
|
||||
type: 'model_actual', round: 2,
|
||||
requested_model: 'selected-model', model: 'provider-backup-alias'
|
||||
}}, round1, round2, 'selected-model');
|
||||
|
||||
const round3 = {{}};
|
||||
inheritModelRouteState(round1, round2, round3, 'selected-model');
|
||||
const metricsTarget = applyModelMetricsState({{
|
||||
requested_model: 'selected-model',
|
||||
model: 'provider-backup-alias',
|
||||
round_models: ['selected-model', 'provider-backup-alias', 'backup-model'],
|
||||
}}, round1, round3, 'selected-model');
|
||||
|
||||
console.log(JSON.stringify({{
|
||||
metricsIsRound3: metricsTarget === round3,
|
||||
round1,
|
||||
round2,
|
||||
round3,
|
||||
}}));
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=script,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=_REPO,
|
||||
timeout=30,
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert json.loads(result.stdout) == {
|
||||
"metricsIsRound3": True,
|
||||
"round1": {
|
||||
"_requestedModel": "selected-model",
|
||||
"_actualModel": "selected-model",
|
||||
},
|
||||
"round2": {
|
||||
"_requestedModel": "selected-model",
|
||||
"_actualModel": "provider-backup-alias",
|
||||
},
|
||||
"round3": {
|
||||
"_requestedModel": "selected-model",
|
||||
"_actualModel": "backup-model",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_same_model_fallback_preserves_distinct_endpoint_route_state():
|
||||
if not shutil.which("node"):
|
||||
pytest.skip("node is not installed")
|
||||
|
||||
script = f"""
|
||||
import {{ applyModelMetricsState, applyModelRouteEventState }} from {json.dumps(_MODULE)};
|
||||
const holder = {{ _requestedModel: 'same-model', _actualModel: 'same-model' }};
|
||||
applyModelRouteEventState({{
|
||||
type: 'fallback',
|
||||
selected_model: 'same-model', answered_by: 'same-model',
|
||||
selected_endpoint_id: 'account-one', selected_endpoint_label: 'Account one',
|
||||
answered_by_endpoint_id: 'account-two', answered_by_endpoint_label: 'Account two',
|
||||
}}, holder, null, 'same-model');
|
||||
applyModelMetricsState({{
|
||||
requested_model: 'same-model', model: 'same-model',
|
||||
requested_endpoint_id: 'account-one', requested_endpoint_label: 'Account one',
|
||||
endpoint_id: 'account-two', endpoint_label: 'Account two',
|
||||
}}, holder, null, 'same-model');
|
||||
console.log(JSON.stringify(holder));
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=script,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=_REPO,
|
||||
timeout=30,
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert json.loads(result.stdout) == {
|
||||
"_requestedModel": "same-model",
|
||||
"_actualModel": "same-model",
|
||||
"_requestedEndpointId": "account-one",
|
||||
"_requestedEndpointLabel": "Account one",
|
||||
"_actualEndpointId": "account-two",
|
||||
"_actualEndpointLabel": "Account two",
|
||||
}
|
||||
|
||||
|
||||
def test_metrics_preserve_explicitly_unknown_round_endpoint():
|
||||
if not shutil.which("node"):
|
||||
pytest.skip("node is not installed")
|
||||
|
||||
script = f"""
|
||||
import {{ applyModelMetricsState }} from {json.dumps(_MODULE)};
|
||||
const holder = {{
|
||||
_requestedModel: 'same-model',
|
||||
_actualModel: 'same-model',
|
||||
_requestedEndpointId: 'account-one',
|
||||
_requestedEndpointLabel: 'Account one',
|
||||
}};
|
||||
const roundHolder = {{}};
|
||||
applyModelMetricsState({{
|
||||
requested_model: 'same-model', model: 'same-model',
|
||||
requested_endpoint_id: 'account-one', requested_endpoint_label: 'Account one',
|
||||
endpoint_id: 'account-two', endpoint_label: 'Account two',
|
||||
round_endpoint_ids: [null], round_endpoint_labels: [null],
|
||||
}}, holder, roundHolder, 'same-model');
|
||||
console.log(JSON.stringify(roundHolder));
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=script,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=_REPO,
|
||||
timeout=30,
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert json.loads(result.stdout) == {
|
||||
"_requestedModel": "same-model",
|
||||
"_actualModel": "same-model",
|
||||
"_requestedEndpointId": "account-one",
|
||||
"_requestedEndpointLabel": "Account one",
|
||||
"_actualEndpointId": None,
|
||||
"_actualEndpointLabel": None,
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Execute terminal stream-error classification under Node."""
|
||||
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
_REPO = Path(__file__).resolve().parents[1]
|
||||
_MODULE = (_REPO / "static" / "js" / "chatStreamErrors.js").as_uri()
|
||||
|
||||
|
||||
def test_terminal_provider_errors_preserve_text_and_never_auto_retry():
|
||||
if not shutil.which("node"):
|
||||
pytest.skip("node is not installed")
|
||||
|
||||
script = f"""
|
||||
import {{ createTerminalStreamError, isRecoverableStreamError }} from {json.dumps(_MODULE)};
|
||||
const stringError = createTerminalStreamError({{ status: 401, error: 'invalid key' }});
|
||||
const objectError = createTerminalStreamError({{ status: 404, error: {{ message: 'model missing' }} }});
|
||||
console.log(JSON.stringify({{
|
||||
stringMessage: stringError.message,
|
||||
objectMessage: objectError.message,
|
||||
terminalRecoverable: isRecoverableStreamError(stringError),
|
||||
eofRecoverable: isRecoverableStreamError(new Error('Stream closed before completion')),
|
||||
networkRecoverable: isRecoverableStreamError(new TypeError('fetch failed')),
|
||||
}}));
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=script,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=_REPO,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert json.loads(result.stdout) == {
|
||||
"stringMessage": "invalid key",
|
||||
"objectMessage": "model missing",
|
||||
"terminalRecoverable": False,
|
||||
"eofRecoverable": True,
|
||||
"networkRecoverable": True,
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
@@ -13,7 +14,7 @@ def test_stream_render_helpers_are_visible_to_catch_block():
|
||||
assert "let _cancelThinkingTimer = () => {};" in outer_scope
|
||||
assert "let _removeThinkingSpinner = () => {};" in outer_scope
|
||||
|
||||
assert "_renderStream = () => {" in try_body
|
||||
assert re.search(r"(?m)^\s*_renderStream\s*=", try_body)
|
||||
assert "_cancelThinkingTimer = () => {" in try_body
|
||||
assert "_removeThinkingSpinner = () => {" in try_body
|
||||
assert "function _renderStream()" not in try_body
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
"""Regression coverage for authoritative Python CI validation."""
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
_WORKFLOW = (
|
||||
Path(__file__).resolve().parent.parent / ".github" / "workflows" / "ci.yml"
|
||||
)
|
||||
|
||||
|
||||
def _indented_block(text: str, heading: str, indent: int) -> str:
|
||||
pattern = re.compile(
|
||||
rf"(?ms)^{' ' * indent}{re.escape(heading)}:\n"
|
||||
rf"(?P<body>(?:(?:{' ' * (indent + 2)}.*|\s*)\n)*)"
|
||||
)
|
||||
match = pattern.search(text)
|
||||
assert match is not None, f"missing {heading!r} block"
|
||||
return match.group(0)
|
||||
|
||||
|
||||
def test_ci_runs_on_integrated_dev_pushes():
|
||||
workflow = _WORKFLOW.read_text()
|
||||
push = _indented_block(workflow, "push", 2)
|
||||
|
||||
assert re.search(r"(?m)^ branches:\s*\[main,\s*dev\]\s*$", push)
|
||||
assert "paths-ignore:" not in push
|
||||
|
||||
|
||||
def test_python_tests_are_authoritative():
|
||||
workflow = _WORKFLOW.read_text()
|
||||
python_tests = _indented_block(workflow, "python-tests", 2)
|
||||
|
||||
assert "python -m pytest -q" in python_tests
|
||||
assert "continue-on-error:" not in python_tests
|
||||
@@ -81,6 +81,13 @@ def _make_stream_with_save(sink, chunks, *, hang_after=None):
|
||||
return gen()
|
||||
|
||||
|
||||
async def _collect_subscription(session_id, expected_run=None):
|
||||
return [
|
||||
event
|
||||
async for event in agent_runs.subscribe(session_id, expected_run)
|
||||
]
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# agent_runs: detached-run semantics (what NORMAL chat/agent streams use)
|
||||
# --------------------------------------------------------------------------- #
|
||||
@@ -136,7 +143,7 @@ async def test_stop_cancels_detached_run_and_saves_partial_exactly_once():
|
||||
break
|
||||
await sub.aclose()
|
||||
|
||||
stopped = agent_runs.stop(session_id)
|
||||
stopped = agent_runs.stop(session_id, run.run_id)
|
||||
assert stopped is True
|
||||
|
||||
await run.task # propagates promptly — not stuck on the hung await
|
||||
@@ -165,6 +172,172 @@ async def test_normal_completion_saves_exactly_once_not_partial():
|
||||
assert sink.saves == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_detached_run_identity_is_stable_for_replay_and_unique_per_run():
|
||||
session_id = "sess-detached-run-identity"
|
||||
agent_runs._RUNS.pop(session_id, None)
|
||||
|
||||
first = agent_runs.start(session_id, _make_stream_with_save(_FakeSaveSink(), ["one"]))
|
||||
first_id = first.run_id
|
||||
assert agent_runs.get_run_id(session_id) == first_id
|
||||
await first.task
|
||||
assert agent_runs.get_run_id(session_id) == first_id
|
||||
|
||||
second = agent_runs.start(session_id, _make_stream_with_save(_FakeSaveSink(), ["two"]))
|
||||
assert second.run_id != first_id
|
||||
assert agent_runs.get_run_id(session_id) == second.run_id
|
||||
await second.task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lazy_subscription_stays_bound_to_header_run_after_replacement():
|
||||
session_id = "sess-detached-lazy-subscription"
|
||||
agent_runs._RUNS.pop(session_id, None)
|
||||
|
||||
async def stream(label):
|
||||
yield f'data: {{"delta":"{label}"}}\n\n'
|
||||
|
||||
first = agent_runs.start(session_id, stream("first"))
|
||||
await first.task
|
||||
# StreamingResponse does not iterate its body until after construction.
|
||||
# Capture the same exact run object used for its identity header.
|
||||
lazy_body = agent_runs.subscribe(session_id, first)
|
||||
|
||||
second = agent_runs.start(session_id, stream("second"))
|
||||
await second.task
|
||||
|
||||
replayed = [event async for event in lazy_body]
|
||||
assert replayed == ['data: {"delta":"first"}\n\n']
|
||||
assert agent_runs.get_run_id(session_id) == second.run_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_run_identity_cannot_stop_replacement_run():
|
||||
session_id = "sess-detached-stale-stop"
|
||||
agent_runs._RUNS.pop(session_id, None)
|
||||
release = asyncio.Event()
|
||||
|
||||
async def finished():
|
||||
yield 'data: {"delta":"old"}\n\n'
|
||||
|
||||
async def replacement():
|
||||
yield 'data: {"delta":"new"}\n\n'
|
||||
await release.wait()
|
||||
|
||||
first = agent_runs.start(session_id, finished())
|
||||
await first.task
|
||||
second = agent_runs.start(session_id, replacement())
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert agent_runs.stop(session_id) is False
|
||||
assert agent_runs.stop(session_id, first.run_id) is False
|
||||
assert second.task is not None and not second.task.done()
|
||||
assert agent_runs.stop(session_id, second.run_id) is True
|
||||
await second.task
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_triple_replacement_closes_middle_subscriber_and_preserves_save_order():
|
||||
session_id = "sess-detached-triple-replacement"
|
||||
agent_runs._RUNS.pop(session_id, None)
|
||||
first_closing = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
third_started = asyncio.Event()
|
||||
|
||||
async def first_stream():
|
||||
try:
|
||||
yield 'data: {"delta":"first"}\n\n'
|
||||
await asyncio.Event().wait()
|
||||
finally:
|
||||
first_closing.set()
|
||||
await release_first.wait()
|
||||
|
||||
async def middle_stream():
|
||||
yield 'data: {"delta":"middle"}\n\n'
|
||||
|
||||
async def third_stream():
|
||||
third_started.set()
|
||||
yield 'data: {"delta":"third"}\n\n'
|
||||
|
||||
first = agent_runs.start(session_id, first_stream())
|
||||
while not first.buffer:
|
||||
await asyncio.sleep(0)
|
||||
|
||||
middle = agent_runs.start(session_id, middle_stream())
|
||||
await first_closing.wait()
|
||||
assert middle.task is not None and not middle.task.done()
|
||||
|
||||
middle_events_task = asyncio.create_task(
|
||||
_collect_subscription(session_id, middle)
|
||||
)
|
||||
while not middle.subscribers:
|
||||
await asyncio.sleep(0)
|
||||
|
||||
third = agent_runs.start(session_id, third_stream())
|
||||
|
||||
# The superseded middle response closes immediately even though its task
|
||||
# remains as the transitive barrier for the first run's partial save.
|
||||
assert await asyncio.wait_for(middle_events_task, timeout=1) == []
|
||||
assert middle.status == "stopped"
|
||||
assert middle.task is not None and not middle.task.done()
|
||||
assert not third_started.is_set()
|
||||
|
||||
release_first.set()
|
||||
await asyncio.wait_for(first.task, timeout=1)
|
||||
await asyncio.wait_for(middle.task, timeout=1)
|
||||
await asyncio.wait_for(third.task, timeout=1)
|
||||
|
||||
assert first.status == "stopped"
|
||||
assert middle.status == "stopped"
|
||||
assert third.status == "done"
|
||||
assert third_started.is_set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reconnect_replays_pinned_fallback_run_without_restarting_tools():
|
||||
session_id = "sess-detached-fallback-resume"
|
||||
agent_runs._RUNS.pop(session_id, None)
|
||||
release = asyncio.Event()
|
||||
tool_executions = 0
|
||||
fallback = 'data: {"type":"fallback","answered_by":"backup","candidate_index":1}\n\n'
|
||||
tool = 'data: {"type":"tool_output","tool":"bash","output":"ok"}\n\n'
|
||||
|
||||
async def pinned_run():
|
||||
nonlocal tool_executions
|
||||
yield fallback
|
||||
tool_executions += 1
|
||||
yield tool
|
||||
await release.wait()
|
||||
yield 'data: {"delta":"backup finished"}\n\n'
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
run = agent_runs.start(session_id, pinned_run())
|
||||
first = agent_runs.subscribe(session_id)
|
||||
first_events = []
|
||||
async for event in first:
|
||||
first_events.append(event)
|
||||
if len(first_events) == 2:
|
||||
break
|
||||
await first.aclose()
|
||||
|
||||
assert run.status == "running"
|
||||
assert tool_executions == 1
|
||||
assert agent_runs._RUNS[session_id] is run
|
||||
|
||||
resumed_events = []
|
||||
resumed = agent_runs.subscribe(session_id)
|
||||
async for event in resumed:
|
||||
resumed_events.append(event)
|
||||
if len(resumed_events) == 2:
|
||||
release.set()
|
||||
await run.task
|
||||
|
||||
assert resumed_events[:2] == [fallback, tool]
|
||||
assert resumed_events[-1] == "data: [DONE]\n\n"
|
||||
assert tool_executions == 1
|
||||
assert agent_runs._RUNS[session_id] is run
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# chat_stream: Compare panes must NOT be detached, so the Stop button (closing
|
||||
# the SSE) cancels the upstream generator promptly — exercising the same
|
||||
|
||||
@@ -306,3 +306,24 @@ def test_integration_recalls_from_chat_history_dom():
|
||||
)
|
||||
assert proc.returncode == 0, proc.stderr
|
||||
assert json.loads(proc.stdout.strip()) == {"value": "stored prompt", "prevented": True}
|
||||
|
||||
|
||||
def test_prompt_recall_is_not_duplicated_in_app_js():
|
||||
"""Only composerArrowUpRecall.js may own ArrowUp on #message (issue #5862).
|
||||
|
||||
static/app.js once carried a near-verbatim copy of this recall logic, wired
|
||||
as a second capture-phase listener on the same textarea. That copy lacked
|
||||
the draft guard here, and because it called stopImmediatePropagation it won
|
||||
regardless of registration order — so a typed multi-line prompt was replaced
|
||||
by the last sent one instead of the caret moving up a line.
|
||||
"""
|
||||
app_js = (_REPO / "static" / "app.js").read_text(encoding="utf-8")
|
||||
for marker in (
|
||||
"_odysseusPromptRecallCapture",
|
||||
"_readComposerPromptHistory",
|
||||
"odysseusRecallIndex",
|
||||
):
|
||||
assert marker not in app_js, (
|
||||
f"static/app.js reintroduces prompt recall ({marker!r}); "
|
||||
"it belongs to static/js/composerArrowUpRecall.js alone"
|
||||
)
|
||||
|
||||
@@ -63,6 +63,23 @@ class TestSelfSummaryPrompt:
|
||||
|
||||
|
||||
class TestTrimForContext:
|
||||
def test_system_truncation_preserves_internal_route_metadata(self):
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "persona\n\n" + ("agent prompt " * 2000),
|
||||
"_agent_injected": "merged_prompt",
|
||||
"_agent_base_message": {"role": "system", "content": "persona"},
|
||||
},
|
||||
{"role": "user", "content": "latest"},
|
||||
]
|
||||
|
||||
trimmed = trim_for_context(messages, context_length=1024, reserve_tokens=256)
|
||||
|
||||
system = next(message for message in trimmed if message.get("role") == "system")
|
||||
assert system["_agent_injected"] == "merged_prompt"
|
||||
assert system["_agent_base_message"] == {"role": "system", "content": "persona"}
|
||||
|
||||
def test_keeps_current_large_user_message_by_truncating(self):
|
||||
huge = "A" * 20000
|
||||
messages = [
|
||||
@@ -194,6 +211,50 @@ class TestMaybeCompactFourthMessage:
|
||||
assert len(result) == 3 and result[2] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deferred_compaction_persists_only_after_route_commit(monkeypatch):
|
||||
updates = []
|
||||
state = {}
|
||||
messages = [
|
||||
{"role": "system", "content": "system " * 100},
|
||||
{"role": "user", "content": "one"},
|
||||
{"role": "assistant", "content": "two"},
|
||||
{"role": "user", "content": "three"},
|
||||
{"role": "assistant", "content": "four"},
|
||||
{"role": "user", "content": "five"},
|
||||
]
|
||||
|
||||
monkeypatch.setattr(cc, "get_context_length", lambda *args: 100)
|
||||
monkeypatch.setattr(cc, "resolve_endpoint", lambda *args, **kwargs: (None, None, None))
|
||||
|
||||
async def fake_summary(*args, **kwargs):
|
||||
return "route-specific summary"
|
||||
|
||||
monkeypatch.setattr(cc, "llm_call_async", fake_summary)
|
||||
monkeypatch.setattr(
|
||||
cc,
|
||||
"_update_session_history",
|
||||
lambda *args, **kwargs: updates.append((args, kwargs)),
|
||||
)
|
||||
|
||||
_compacted, _context, was_compacted = await cc.maybe_compact(
|
||||
object(),
|
||||
"https://candidate.example/v1",
|
||||
"candidate-model",
|
||||
messages,
|
||||
persist=False,
|
||||
compaction_state=state,
|
||||
)
|
||||
|
||||
assert was_compacted is True
|
||||
assert updates == []
|
||||
assert state["summary"] == "route-specific summary"
|
||||
assert cc.apply_compaction_state(object(), state) is True
|
||||
assert len(updates) == 1
|
||||
assert cc.apply_compaction_state(object(), state) is False
|
||||
assert len(updates) == 1
|
||||
|
||||
|
||||
class TestResearchPrimerPreserved:
|
||||
"""A research-spinoff primer (metadata research_spinoff_from) must never be
|
||||
trimmed away — it is the Discuss chat's sole knowledge base (drift fix)."""
|
||||
|
||||
@@ -723,7 +723,12 @@ def test_local_windows_download_pid_tracks_inner_bash_and_stop_kills_tree():
|
||||
routes_src = (Path(__file__).resolve().parents[1] / "routes" / "cookbook_routes.py").read_text(encoding="utf-8")
|
||||
running_src = (Path(__file__).resolve().parents[1] / "static" / "js" / "cookbookRunning.js").read_text(encoding="utf-8")
|
||||
|
||||
assert 'printf \'%s\\\\n\' \\"$$\\" > {pp}' in routes_src
|
||||
# The Windows-local runner publishes Python's valid Win32 fallback before
|
||||
# allowing Git Bash to replace it with /proc/$$/winpid.
|
||||
assert "_windows_local_pid_record_line(pid_path, pid_ready_path)" in routes_src
|
||||
assert "/proc/$$/winpid" in routes_src
|
||||
assert "pid_ready_path.touch()" in routes_src
|
||||
assert '\\"$$\\" > {pp}' not in routes_src
|
||||
assert "function Stop-Tree([int]$Id)" in running_src
|
||||
assert "('ParentProcessId = ' + $Id)" in running_src
|
||||
assert "Stop-Tree ([int]$p)" in running_src
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
"""Behavioral regression coverage for Windows-local Cookbook PID recording."""
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from routes.cookbook_routes import _windows_local_pid_record_line
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
COOKBOOK_ROUTES = ROOT / "routes" / "cookbook_routes.py"
|
||||
|
||||
|
||||
def _fake_cat(tmp_path: Path, body: str) -> Path:
|
||||
fake_bin = tmp_path / "bin"
|
||||
fake_bin.mkdir()
|
||||
cat = fake_bin / "cat"
|
||||
cat.write_text("#!/bin/sh\n" + body + "\n", encoding="utf-8")
|
||||
cat.chmod(0o755)
|
||||
return fake_bin
|
||||
|
||||
|
||||
def _env_for(fake_bin: Path, **extra: str) -> dict[str, str]:
|
||||
env = dict(os.environ)
|
||||
env["PATH"] = str(fake_bin) + os.pathsep + env.get("PATH", "")
|
||||
env.update(extra)
|
||||
return env
|
||||
|
||||
|
||||
def _run_pid_line(
|
||||
pid_path: Path,
|
||||
ready_path: Path,
|
||||
fake_bin: Path,
|
||||
**extra_env: str,
|
||||
) -> subprocess.CompletedProcess:
|
||||
return subprocess.run(
|
||||
["bash", "-c", _windows_local_pid_record_line(pid_path, ready_path)],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=_env_for(fake_bin, **extra_env),
|
||||
timeout=10,
|
||||
)
|
||||
|
||||
|
||||
def test_windows_local_pid_line_records_numeric_winpid_after_fallback(tmp_path):
|
||||
pid_path = tmp_path / "serve.pid"
|
||||
ready_path = tmp_path / "serve.pid.ready"
|
||||
|
||||
pid_path.write_text("11111", encoding="utf-8")
|
||||
ready_path.touch()
|
||||
|
||||
cat_arg = tmp_path / "cat-arg.txt"
|
||||
fake_bin = _fake_cat(
|
||||
tmp_path,
|
||||
'printf "%s\\n" "$1" > "$FAKE_CAT_ARG"\n'
|
||||
'printf "%s\\n" "$FAKE_WINPID"',
|
||||
)
|
||||
|
||||
result = _run_pid_line(
|
||||
pid_path,
|
||||
ready_path,
|
||||
fake_bin,
|
||||
FAKE_CAT_ARG=str(cat_arg),
|
||||
FAKE_WINPID="42324",
|
||||
)
|
||||
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert pid_path.read_text(encoding="utf-8").strip() == "42324"
|
||||
assert not ready_path.exists()
|
||||
|
||||
proc_path = cat_arg.read_text(encoding="utf-8").strip()
|
||||
parts = proc_path.strip("/").split("/")
|
||||
assert len(parts) == 3
|
||||
assert parts[0] == "proc"
|
||||
assert parts[1].isdigit()
|
||||
assert parts[2] == "winpid"
|
||||
|
||||
|
||||
def test_windows_local_pid_line_waits_for_python_fallback_before_replacing(tmp_path):
|
||||
pid_path = tmp_path / "serve.pid"
|
||||
ready_path = tmp_path / "serve.pid.ready"
|
||||
|
||||
fake_bin = _fake_cat(
|
||||
tmp_path,
|
||||
'printf "%s\\n" "$FAKE_WINPID"',
|
||||
)
|
||||
|
||||
proc = subprocess.Popen(
|
||||
[
|
||||
"bash",
|
||||
"-c",
|
||||
_windows_local_pid_record_line(pid_path, ready_path),
|
||||
],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
env=_env_for(fake_bin, FAKE_WINPID="42324"),
|
||||
)
|
||||
|
||||
# The inner shell has started, but Python has not published its fallback yet.
|
||||
time.sleep(0.05)
|
||||
assert proc.poll() is None
|
||||
assert not pid_path.exists()
|
||||
|
||||
# Simulate the post-Popen Python publication order.
|
||||
pid_path.write_text("31100", encoding="utf-8")
|
||||
ready_path.touch()
|
||||
|
||||
stdout, stderr = proc.communicate(timeout=10)
|
||||
|
||||
assert proc.returncode == 0, stderr or stdout
|
||||
assert pid_path.read_text(encoding="utf-8").strip() == "42324"
|
||||
assert not ready_path.exists()
|
||||
|
||||
|
||||
def test_windows_local_pid_line_preserves_outer_pid_when_mapping_missing(tmp_path):
|
||||
pid_path = tmp_path / "serve.pid"
|
||||
ready_path = tmp_path / "serve.pid.ready"
|
||||
|
||||
pid_path.write_text("31100", encoding="utf-8")
|
||||
ready_path.touch()
|
||||
|
||||
fake_bin = _fake_cat(tmp_path, "exit 1")
|
||||
|
||||
result = _run_pid_line(
|
||||
pid_path,
|
||||
ready_path,
|
||||
fake_bin,
|
||||
)
|
||||
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert pid_path.read_text(encoding="utf-8").strip() == "31100"
|
||||
assert not ready_path.exists()
|
||||
|
||||
|
||||
def test_windows_local_pid_line_rejects_malformed_mapping(tmp_path):
|
||||
pid_path = tmp_path / "serve.pid"
|
||||
ready_path = tmp_path / "serve.pid.ready"
|
||||
|
||||
pid_path.write_text("31100", encoding="utf-8")
|
||||
ready_path.touch()
|
||||
|
||||
fake_bin = _fake_cat(
|
||||
tmp_path,
|
||||
'printf "not-a-win32-pid\\n"',
|
||||
)
|
||||
|
||||
result = _run_pid_line(
|
||||
pid_path,
|
||||
ready_path,
|
||||
fake_bin,
|
||||
)
|
||||
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert pid_path.read_text(encoding="utf-8").strip() == "31100"
|
||||
assert not ready_path.exists()
|
||||
|
||||
|
||||
def test_local_windows_launcher_publishes_fallback_before_releasing_inner_runner():
|
||||
source = COOKBOOK_ROUTES.read_text(encoding="utf-8")
|
||||
start = source.index(" def _launch_local_detached(")
|
||||
end = source.index(
|
||||
' @router.post("/api/model/download")',
|
||||
start,
|
||||
)
|
||||
launcher = source[start:end]
|
||||
|
||||
assert "_windows_local_pid_record_line(pid_path, pid_ready_path)" in launcher
|
||||
assert "pid_ready_path.unlink(missing_ok=True)" in launcher
|
||||
|
||||
fallback = launcher.index(
|
||||
'pid_path.write_text(str(proc.pid), encoding="utf-8")'
|
||||
)
|
||||
release = launcher.index("pid_ready_path.touch()")
|
||||
|
||||
assert fallback < release
|
||||
|
||||
# Never write Git Bash's bare MSYS $$ to the session pid file.
|
||||
assert '\\"$$\\" > {pp}' not in launcher
|
||||
@@ -0,0 +1,522 @@
|
||||
"""Regressions for process-safe email-account default mutations.
|
||||
|
||||
The file-backed SQLite fixture uses a fresh connection for every Session.
|
||||
That exercises the same database lock boundary used by separate web workers,
|
||||
rather than relying on an in-process Python lock.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
import threading
|
||||
import types
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import create_engine, create_mock_engine, text
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy.pool import NullPool
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def account_db(tmp_path, monkeypatch):
|
||||
from core import database as core_db
|
||||
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'accounts.db'}",
|
||||
connect_args={"check_same_thread": False, "timeout": 5},
|
||||
poolclass=NullPool,
|
||||
)
|
||||
core_db.Base.metadata.create_all(engine)
|
||||
factory = sessionmaker(
|
||||
bind=engine,
|
||||
autocommit=False,
|
||||
autoflush=False,
|
||||
)
|
||||
monkeypatch.setattr(core_db, "SessionLocal", factory)
|
||||
yield factory
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def _endpoint(method, path):
|
||||
from routes import email_routes
|
||||
|
||||
with mock.patch.object(email_routes, "_start_poller"):
|
||||
router = email_routes.setup_email_routes()
|
||||
for route in router.routes:
|
||||
if route.path == path and method in getattr(route, "methods", set()):
|
||||
return route.endpoint
|
||||
raise AssertionError(f"email route not found: {method} {path}")
|
||||
|
||||
|
||||
def _named_endpoint(router, name):
|
||||
for route in router.routes:
|
||||
if getattr(getattr(route, "endpoint", None), "__name__", "") == name:
|
||||
return route.endpoint
|
||||
raise AssertionError(f"route not found: {name}")
|
||||
|
||||
|
||||
def _seed_account(factory, account_id, owner, *, is_default=False, enabled=True):
|
||||
from core.database import EmailAccount
|
||||
|
||||
db = factory()
|
||||
try:
|
||||
db.add(
|
||||
EmailAccount(
|
||||
id=account_id,
|
||||
owner=owner,
|
||||
name=account_id,
|
||||
is_default=is_default,
|
||||
enabled=enabled,
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _rows(factory):
|
||||
from core.database import EmailAccount
|
||||
|
||||
db = factory()
|
||||
try:
|
||||
return [
|
||||
(row.id, row.owner, bool(row.is_default))
|
||||
for row in db.query(EmailAccount).order_by(EmailAccount.id).all()
|
||||
]
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _install_lock_pause(monkeypatch, paused_thread_name):
|
||||
"""Pause one worker after acquisition and observe another waiting."""
|
||||
from routes import email_routes
|
||||
|
||||
real_lock = email_routes._lock_email_account_owner_mutation
|
||||
first_acquired = threading.Event()
|
||||
release_first = threading.Event()
|
||||
contender_attempted = threading.Event()
|
||||
contender_acquired = threading.Event()
|
||||
|
||||
def controlled_lock(db, owner):
|
||||
is_first = threading.current_thread().name == paused_thread_name
|
||||
if not is_first:
|
||||
contender_attempted.set()
|
||||
real_lock(db, owner)
|
||||
if is_first:
|
||||
first_acquired.set()
|
||||
assert release_first.wait(5), "timed out releasing first mutation"
|
||||
else:
|
||||
contender_acquired.set()
|
||||
|
||||
monkeypatch.setattr(
|
||||
email_routes,
|
||||
"_lock_email_account_owner_mutation",
|
||||
controlled_lock,
|
||||
)
|
||||
return first_acquired, release_first, contender_attempted, contender_acquired
|
||||
|
||||
|
||||
def test_concurrent_first_account_creates_choose_one_default(account_db, monkeypatch):
|
||||
create_account = _endpoint("POST", "/api/email/accounts")
|
||||
first_acquired, release_first, attempted, acquired = _install_lock_pause(
|
||||
monkeypatch, "first-account"
|
||||
)
|
||||
results = {}
|
||||
|
||||
def create(name):
|
||||
results[name] = asyncio.run(
|
||||
create_account({"name": name, "is_default": False}, owner="alice")
|
||||
)
|
||||
|
||||
first = threading.Thread(target=create, args=("First",), name="first-account")
|
||||
second = threading.Thread(target=create, args=("Second",), name="second-account")
|
||||
first.start()
|
||||
assert first_acquired.wait(5)
|
||||
second.start()
|
||||
assert attempted.wait(5)
|
||||
assert not acquired.wait(0.1), "second session bypassed the database mutation lock"
|
||||
|
||||
release_first.set()
|
||||
first.join(5)
|
||||
second.join(5)
|
||||
|
||||
assert not first.is_alive()
|
||||
assert not second.is_alive()
|
||||
assert results["First"]["ok"] is True
|
||||
assert results["Second"]["ok"] is True
|
||||
defaults = [row for row in _rows(account_db) if row[2]]
|
||||
assert [(row[1], row[2]) for row in defaults] == [("alice", True)]
|
||||
assert len(defaults) == 1
|
||||
|
||||
|
||||
def test_delete_promotion_and_set_default_are_one_serial_transition(
|
||||
account_db, monkeypatch
|
||||
):
|
||||
from sqlalchemy.orm import Session as OrmSession
|
||||
|
||||
_seed_account(account_db, "alice-a", "alice", is_default=True)
|
||||
_seed_account(account_db, "alice-b", "alice")
|
||||
_seed_account(account_db, "alice-c", "alice")
|
||||
_seed_account(account_db, "bob-a", "bob", is_default=True)
|
||||
|
||||
delete_account = _endpoint("DELETE", "/api/email/accounts/{account_id}")
|
||||
set_default = _endpoint("POST", "/api/email/accounts/{account_id}/set-default")
|
||||
first_acquired, release_first, attempted, acquired = _install_lock_pause(
|
||||
monkeypatch, "delete-default"
|
||||
)
|
||||
delete_commit_finished = threading.Event()
|
||||
release_delete_after_commit = threading.Event()
|
||||
real_commit = OrmSession.commit
|
||||
results = {}
|
||||
|
||||
def pause_after_delete_commit(session):
|
||||
real_commit(session)
|
||||
if (
|
||||
threading.current_thread().name == "delete-default"
|
||||
and not delete_commit_finished.is_set()
|
||||
):
|
||||
delete_commit_finished.set()
|
||||
assert release_delete_after_commit.wait(5), (
|
||||
"timed out releasing delete after its first commit"
|
||||
)
|
||||
|
||||
monkeypatch.setattr(OrmSession, "commit", pause_after_delete_commit)
|
||||
|
||||
def delete_old_default():
|
||||
results["delete"] = asyncio.run(
|
||||
delete_account("alice-a", owner="alice")
|
||||
)
|
||||
|
||||
def select_new_default():
|
||||
results["set"] = asyncio.run(
|
||||
set_default("alice-c", owner="alice")
|
||||
)
|
||||
|
||||
delete_thread = threading.Thread(target=delete_old_default, name="delete-default")
|
||||
set_thread = threading.Thread(target=select_new_default, name="set-default")
|
||||
delete_thread.start()
|
||||
assert first_acquired.wait(5)
|
||||
set_thread.start()
|
||||
assert attempted.wait(5)
|
||||
assert not acquired.wait(0.1), "set-default bypassed the delete transaction"
|
||||
|
||||
release_first.set()
|
||||
assert delete_commit_finished.wait(5)
|
||||
# The deletion transaction has committed. Let the contender complete
|
||||
# before the deleting handler can continue: if promotion were still a
|
||||
# second commit, it would now run after set-default and recreate two
|
||||
# defaults deterministically.
|
||||
assert acquired.wait(5)
|
||||
set_thread.join(5)
|
||||
release_delete_after_commit.set()
|
||||
delete_thread.join(5)
|
||||
|
||||
assert not delete_thread.is_alive()
|
||||
assert not set_thread.is_alive()
|
||||
assert results == {"delete": {"ok": True}, "set": {"ok": True}}
|
||||
assert _rows(account_db) == [
|
||||
("alice-b", "alice", False),
|
||||
("alice-c", "alice", True),
|
||||
("bob-a", "bob", True),
|
||||
]
|
||||
|
||||
|
||||
def test_upgrade_normalizes_legacy_defaults_and_installs_unique_index(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
"""A pre-index schema upgrades without requiring newer account columns."""
|
||||
from core import database as core_db
|
||||
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'legacy-accounts.db'}",
|
||||
connect_args={"check_same_thread": False},
|
||||
poolclass=NullPool,
|
||||
)
|
||||
try:
|
||||
with engine.begin() as conn:
|
||||
conn.execute(text("""
|
||||
CREATE TABLE email_accounts (
|
||||
id VARCHAR PRIMARY KEY,
|
||||
owner VARCHAR,
|
||||
name VARCHAR NOT NULL,
|
||||
is_default BOOLEAN NOT NULL,
|
||||
enabled BOOLEAN NOT NULL,
|
||||
created_at DATETIME,
|
||||
updated_at DATETIME
|
||||
)
|
||||
"""))
|
||||
conn.execute(text("""
|
||||
INSERT INTO email_accounts
|
||||
(id, owner, name, is_default, enabled, created_at, updated_at)
|
||||
VALUES
|
||||
('legacy-old', NULL, 'Old', 1, 1, '2024-01-01', '2024-01-01'),
|
||||
('legacy-new', '', 'New', 1, 1, '2025-01-01', '2025-01-01')
|
||||
"""))
|
||||
|
||||
monkeypatch.setattr(core_db, "engine", engine)
|
||||
core_db._migrate_email_account_default_invariant()
|
||||
core_db._migrate_email_account_default_invariant() # idempotent replay
|
||||
|
||||
with engine.connect() as conn:
|
||||
defaults = conn.execute(text("""
|
||||
SELECT id FROM email_accounts
|
||||
WHERE is_default IS TRUE
|
||||
ORDER BY id
|
||||
""")).scalars().all()
|
||||
index_names = {
|
||||
row[1] for row in conn.execute(text("PRAGMA index_list(email_accounts)"))
|
||||
}
|
||||
assert defaults == ["legacy-old"]
|
||||
assert core_db._EMAIL_ACCOUNT_DEFAULT_INDEX in index_names
|
||||
|
||||
with pytest.raises(IntegrityError):
|
||||
with engine.begin() as conn:
|
||||
conn.execute(text("""
|
||||
INSERT INTO email_accounts
|
||||
(id, owner, name, is_default, enabled, created_at, updated_at)
|
||||
VALUES
|
||||
('legacy-third', NULL, 'Third', 1, 1, '2026-01-01', '2026-01-01')
|
||||
"""))
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_concurrent_legacy_seed_is_one_locked_transaction(
|
||||
tmp_path, monkeypatch, caplog
|
||||
):
|
||||
from core import database as core_db
|
||||
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'seed-accounts.db'}",
|
||||
connect_args={"check_same_thread": False, "timeout": 5},
|
||||
poolclass=NullPool,
|
||||
)
|
||||
core_db.Base.metadata.create_all(engine)
|
||||
factory = sessionmaker(bind=engine, autocommit=False, autoflush=False)
|
||||
settings_file = tmp_path / "settings.json"
|
||||
settings_file.write_text(
|
||||
json.dumps({"imap_host": "imap.example.test", "imap_user": "alice"}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
monkeypatch.setattr(core_db, "engine", engine)
|
||||
monkeypatch.setattr(core_db, "SessionLocal", factory)
|
||||
monkeypatch.setattr(core_db, "SETTINGS_FILE", str(settings_file))
|
||||
|
||||
read_barrier = threading.Barrier(2)
|
||||
real_read_text = Path.read_text
|
||||
|
||||
def synchronized_read(path, *args, **kwargs):
|
||||
value = real_read_text(path, *args, **kwargs)
|
||||
if path == settings_file:
|
||||
read_barrier.wait(5)
|
||||
return value
|
||||
|
||||
monkeypatch.setattr(Path, "read_text", synchronized_read)
|
||||
threads = [
|
||||
threading.Thread(target=core_db._migrate_seed_email_account)
|
||||
for _ in range(2)
|
||||
]
|
||||
try:
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
for thread in threads:
|
||||
thread.join(5)
|
||||
assert all(not thread.is_alive() for thread in threads)
|
||||
|
||||
with engine.connect() as conn:
|
||||
rows = conn.execute(text("""
|
||||
SELECT owner, is_default FROM email_accounts
|
||||
ORDER BY id
|
||||
""")).all()
|
||||
assert rows == [(None, 1)]
|
||||
assert "seed email account migration:" not in caplog.text
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_multi_owner_row_locks_are_acquired_in_canonical_order():
|
||||
from core.database import lock_email_account_owner_mutations
|
||||
|
||||
class FakeSession:
|
||||
def __init__(self):
|
||||
self.locked = []
|
||||
|
||||
def get_bind(self):
|
||||
return SimpleNamespace(dialect=SimpleNamespace(name="postgresql"))
|
||||
|
||||
def get(self, _model, owner_key, **kwargs):
|
||||
assert kwargs == {"with_for_update": True}
|
||||
self.locked.append(owner_key)
|
||||
return object()
|
||||
|
||||
db = FakeSession()
|
||||
lock_email_account_owner_mutations(db, "zeta", "", "alpha", "zeta")
|
||||
assert db.locked == ["", "alpha", "zeta"]
|
||||
|
||||
|
||||
def test_postgresql_fresh_schema_emits_default_unique_index():
|
||||
from core import database as core_db
|
||||
|
||||
statements = []
|
||||
engine_holder = {}
|
||||
|
||||
def capture(statement, *_args, **_kwargs):
|
||||
statements.append(
|
||||
str(statement.compile(dialect=engine_holder["engine"].dialect))
|
||||
)
|
||||
|
||||
mock_engine = create_mock_engine("postgresql://", capture)
|
||||
engine_holder["engine"] = mock_engine
|
||||
core_db.EmailAccount.__table__.create(mock_engine)
|
||||
|
||||
assert any(
|
||||
core_db._EMAIL_ACCOUNT_DEFAULT_INDEX in statement
|
||||
and "COALESCE(owner, '')" in statement
|
||||
and "WHERE is_default IS TRUE" in statement
|
||||
for statement in statements
|
||||
)
|
||||
|
||||
|
||||
def test_rename_serializes_old_and_new_owner_and_stale_set_default_fails_closed(
|
||||
account_db, monkeypatch, tmp_path
|
||||
):
|
||||
from core import database as core_db
|
||||
from routes import auth_routes
|
||||
|
||||
_seed_account(account_db, "alice-a", "alice", is_default=True)
|
||||
_seed_account(account_db, "alice-b", "alice")
|
||||
_seed_account(account_db, "bob-a", "bob", is_default=True)
|
||||
|
||||
prefs_module = types.ModuleType("routes.prefs_routes")
|
||||
prefs_module._load = lambda: {}
|
||||
prefs_module._save = lambda _data: None
|
||||
monkeypatch.setitem(sys.modules, "routes.prefs_routes", prefs_module)
|
||||
monkeypatch.setattr(
|
||||
auth_routes, "DEEP_RESEARCH_DIR", str(tmp_path / "deep_research")
|
||||
)
|
||||
monkeypatch.setattr(auth_routes, "MEMORY_FILE", str(tmp_path / "memory.json"))
|
||||
monkeypatch.setattr(auth_routes, "SKILLS_DIR", str(tmp_path / "skills"))
|
||||
|
||||
auth_manager = mock.MagicMock()
|
||||
auth_manager.get_username_for_token.return_value = "admin"
|
||||
auth_manager.is_admin.return_value = True
|
||||
auth_manager.users = {"admin": {}, "alice": {}}
|
||||
auth_manager.rename_user.return_value = True
|
||||
rename_user = _named_endpoint(
|
||||
auth_routes.setup_auth_routes(auth_manager), "rename_user"
|
||||
)
|
||||
set_default = _endpoint("POST", "/api/email/accounts/{account_id}/set-default")
|
||||
|
||||
rename_acquired = threading.Event()
|
||||
release_rename = threading.Event()
|
||||
set_attempted = threading.Event()
|
||||
set_acquired = threading.Event()
|
||||
real_lock = core_db.lock_email_account_owner_mutations
|
||||
|
||||
def controlled_lock(db, *owners):
|
||||
thread_name = threading.current_thread().name
|
||||
if thread_name == "rename-owner":
|
||||
real_lock(db, *owners)
|
||||
rename_acquired.set()
|
||||
assert release_rename.wait(5)
|
||||
return
|
||||
if thread_name == "stale-set-default":
|
||||
set_attempted.set()
|
||||
real_lock(db, *owners)
|
||||
set_acquired.set()
|
||||
return
|
||||
real_lock(db, *owners)
|
||||
|
||||
monkeypatch.setattr(core_db, "lock_email_account_owner_mutations", controlled_lock)
|
||||
request = SimpleNamespace(
|
||||
cookies={"odysseus_session": "admin-token"},
|
||||
app=SimpleNamespace(
|
||||
state=SimpleNamespace(
|
||||
invalidate_token_cache=lambda: None,
|
||||
session_manager=None,
|
||||
research_handler=None,
|
||||
upload_handler=None,
|
||||
personal_docs_manager=None,
|
||||
)
|
||||
),
|
||||
)
|
||||
results = {}
|
||||
|
||||
def rename_owner():
|
||||
results["rename"] = asyncio.run(
|
||||
rename_user("alice", SimpleNamespace(username="bob"), request)
|
||||
)
|
||||
|
||||
def select_stale_default():
|
||||
try:
|
||||
results["set"] = asyncio.run(
|
||||
set_default("alice-b", owner="alice")
|
||||
)
|
||||
except Exception as exc: # asserted below with its HTTP status
|
||||
results["set_error"] = exc
|
||||
|
||||
rename_thread = threading.Thread(target=rename_owner, name="rename-owner")
|
||||
set_thread = threading.Thread(
|
||||
target=select_stale_default, name="stale-set-default"
|
||||
)
|
||||
rename_thread.start()
|
||||
assert rename_acquired.wait(5)
|
||||
set_thread.start()
|
||||
assert set_attempted.wait(5)
|
||||
assert not set_acquired.wait(0.1), "set-default bypassed the rename lock"
|
||||
|
||||
release_rename.set()
|
||||
rename_thread.join(5)
|
||||
set_thread.join(5)
|
||||
|
||||
assert not rename_thread.is_alive()
|
||||
assert not set_thread.is_alive()
|
||||
assert results["rename"]["ok"] is True
|
||||
assert isinstance(results["set_error"], HTTPException)
|
||||
assert results["set_error"].status_code == 404
|
||||
assert _rows(account_db) == [
|
||||
("alice-a", "bob", False),
|
||||
("alice-b", "bob", False),
|
||||
("bob-a", "bob", True),
|
||||
]
|
||||
|
||||
|
||||
def test_demo_teardown_promotes_replacement_in_same_transaction(
|
||||
account_db, monkeypatch
|
||||
):
|
||||
from core.database import EmailAccount
|
||||
from scripts.demo_email import demo_account
|
||||
|
||||
db = account_db()
|
||||
try:
|
||||
db.add_all([
|
||||
EmailAccount(
|
||||
id="real",
|
||||
owner="",
|
||||
name="Real",
|
||||
is_default=False,
|
||||
enabled=True,
|
||||
),
|
||||
EmailAccount(
|
||||
id="demo",
|
||||
owner="",
|
||||
name=demo_account.NAME,
|
||||
imap_user=demo_account.IMAP_USER,
|
||||
is_default=True,
|
||||
enabled=True,
|
||||
),
|
||||
])
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
monkeypatch.setattr(demo_account, "SessionLocal", account_db)
|
||||
monkeypatch.setattr(demo_account, "engine", account_db.kw["bind"])
|
||||
|
||||
assert demo_account.teardown() == 0
|
||||
assert _rows(account_db) == [("real", "", True)]
|
||||
@@ -0,0 +1,414 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
_REPO = Path(__file__).resolve().parents[1]
|
||||
_EMAIL_LIBRARY = _REPO / "static" / "js" / "emailLibrary.js"
|
||||
|
||||
|
||||
def _source() -> str:
|
||||
return _EMAIL_LIBRARY.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
def _function_source(name: str) -> str:
|
||||
"""Return one top-level JS function using balanced braces."""
|
||||
text = _source()
|
||||
markers = (f"function {name}", f"async function {name}", f"export function {name}", f"export async function {name}")
|
||||
starts = [text.find(marker) for marker in markers]
|
||||
starts = [start for start in starts if start >= 0]
|
||||
assert starts, f"missing function {name}"
|
||||
start = min(starts)
|
||||
paren = text.index("(", start)
|
||||
paren_depth = 0
|
||||
quote = None
|
||||
escaped = False
|
||||
for index in range(paren, len(text)):
|
||||
char = text[index]
|
||||
if quote:
|
||||
if escaped:
|
||||
escaped = False
|
||||
elif char == "\\":
|
||||
escaped = True
|
||||
elif char == quote:
|
||||
quote = None
|
||||
continue
|
||||
if char in ("'", '"', "`"):
|
||||
quote = char
|
||||
elif char == "(":
|
||||
paren_depth += 1
|
||||
elif char == ")":
|
||||
paren_depth -= 1
|
||||
if paren_depth == 0:
|
||||
brace = text.index("{", index)
|
||||
break
|
||||
else:
|
||||
raise AssertionError(f"unterminated signature {name}")
|
||||
depth = 0
|
||||
quote = None
|
||||
escaped = False
|
||||
template_depth = 0
|
||||
for index in range(brace, len(text)):
|
||||
char = text[index]
|
||||
if quote:
|
||||
if escaped:
|
||||
escaped = False
|
||||
elif char == "\\":
|
||||
escaped = True
|
||||
elif char == quote and template_depth == 0:
|
||||
quote = None
|
||||
elif quote == "`" and char == "$" and index + 1 < len(text) and text[index + 1] == "{":
|
||||
template_depth += 1
|
||||
elif quote == "`" and char == "}" and template_depth:
|
||||
template_depth -= 1
|
||||
continue
|
||||
if char in ("'", '"', "`"):
|
||||
quote = char
|
||||
elif char == "{":
|
||||
depth += 1
|
||||
elif char == "}":
|
||||
depth -= 1
|
||||
if depth == 0:
|
||||
return text[start:index + 1]
|
||||
raise AssertionError(f"unterminated function {name}")
|
||||
|
||||
|
||||
def _run_scheduler_scenario(scenario: str):
|
||||
node = shutil.which("node")
|
||||
if not node:
|
||||
pytest.skip("node not on PATH")
|
||||
functions = "\n".join(
|
||||
_function_source(name)
|
||||
for name in (
|
||||
"_isChatInteractionBusy",
|
||||
"_canRunEmailPrewarm",
|
||||
"_isEmailPrewarmTemporarilyBlocked",
|
||||
"_settleEmailPrewarm",
|
||||
"_cancelEmailPrewarm",
|
||||
"_scheduleEmailPrewarm",
|
||||
)
|
||||
)
|
||||
script = f"""
|
||||
let now = 0;
|
||||
Date.now = () => now;
|
||||
const state = {{ _libOpen: false, _libLoading: false }};
|
||||
let _libSearchInFlight = false;
|
||||
let _libPrewarmDelayTimer = null;
|
||||
let _libPrewarmIdleHandle = null;
|
||||
let _libPrewarmPromise = null;
|
||||
let _libPrewarmResolve = null;
|
||||
let _libPrewarmAbortController = null;
|
||||
let _libPrewarmDetachPriorityListeners = null;
|
||||
let _libPrewarmGeneration = 0;
|
||||
let nextHandle = 1;
|
||||
const timers = new Map();
|
||||
const idleCallbacks = new Map();
|
||||
let idleRequestCount = 0;
|
||||
function eventTarget(target) {{
|
||||
const listeners = new Map();
|
||||
target.addEventListener = (type, callback) => {{
|
||||
if (!listeners.has(type)) listeners.set(type, new Set());
|
||||
listeners.get(type).add(callback);
|
||||
}};
|
||||
target.removeEventListener = (type, callback) => listeners.get(type)?.delete(callback);
|
||||
target.dispatchEvent = (event) => {{
|
||||
for (const callback of [...(listeners.get(event.type) || [])]) callback(event);
|
||||
}};
|
||||
target.listenerCount = (type) => listeners.get(type)?.size || 0;
|
||||
return target;
|
||||
}}
|
||||
const document = eventTarget({{ visibilityState: 'visible' }});
|
||||
const window = {{
|
||||
__odysseusChatBusy: false,
|
||||
__odysseusChatBusyUntil: 0,
|
||||
requestIdleCallback(callback) {{
|
||||
const handle = nextHandle++;
|
||||
idleRequestCount += 1;
|
||||
idleCallbacks.set(handle, callback);
|
||||
return handle;
|
||||
}},
|
||||
cancelIdleCallback(handle) {{ idleCallbacks.delete(handle); }},
|
||||
}};
|
||||
eventTarget(window);
|
||||
function setTimeout(callback, delay) {{
|
||||
const handle = nextHandle++;
|
||||
timers.set(handle, {{ callback, at: now + Number(delay || 0) }});
|
||||
return handle;
|
||||
}}
|
||||
function clearTimeout(handle) {{ timers.delete(handle); }}
|
||||
async function flushMicrotasks() {{
|
||||
for (let i = 0; i < 6; i += 1) await Promise.resolve();
|
||||
}}
|
||||
async function advanceTo(target) {{
|
||||
while (true) {{
|
||||
const pending = [...timers.entries()]
|
||||
.filter(([, timer]) => timer.at <= target)
|
||||
.sort((a, b) => a[1].at - b[1].at)[0];
|
||||
if (!pending) break;
|
||||
const [handle, timer] = pending;
|
||||
timers.delete(handle);
|
||||
now = timer.at;
|
||||
timer.callback();
|
||||
await flushMicrotasks();
|
||||
}}
|
||||
now = target;
|
||||
await flushMicrotasks();
|
||||
}}
|
||||
async function fireNextIdle(budget = 5) {{
|
||||
const pending = idleCallbacks.entries().next().value;
|
||||
if (!pending) throw new Error('no idle callback pending');
|
||||
const [handle, callback] = pending;
|
||||
idleCallbacks.delete(handle);
|
||||
callback({{ didTimeout: false, timeRemaining: () => budget }});
|
||||
await flushMicrotasks();
|
||||
}}
|
||||
{functions}
|
||||
{scenario}
|
||||
"""
|
||||
proc = subprocess.run(
|
||||
[node, "--input-type=module"],
|
||||
input=script,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=str(_REPO),
|
||||
timeout=30,
|
||||
)
|
||||
assert proc.returncode == 0, proc.stderr
|
||||
return json.loads(proc.stdout.strip())
|
||||
|
||||
|
||||
def test_prewarm_is_genuine_idle_only_and_single_flight():
|
||||
scheduler = _function_source("_scheduleEmailPrewarm")
|
||||
|
||||
assert "if (_libPrewarmPromise) return _libPrewarmPromise;" in scheduler
|
||||
assert "typeof window.requestIdleCallback !== 'function'" in scheduler
|
||||
assert "return Promise.resolve(false);" in scheduler
|
||||
assert "window.requestIdleCallback((deadline)" in scheduler
|
||||
assert "!deadline.didTimeout" in scheduler
|
||||
assert "deadline.timeRemaining() > 0" in scheduler
|
||||
|
||||
idle_callback = scheduler.index("window.requestIdleCallback((deadline)")
|
||||
assert "Promise.resolve()" in scheduler
|
||||
task_start = scheduler.index("task({ signal: controller.signal, generation })")
|
||||
assert idle_callback < task_start, "network work must only be reachable from the idle callback"
|
||||
|
||||
|
||||
def test_temporary_chat_priority_retries_one_single_flight_until_idle():
|
||||
out = _run_scheduler_scenario("""
|
||||
window.__odysseusChatBusyUntil = 10000;
|
||||
let taskCalls = 0;
|
||||
const task = async () => { taskCalls += 1; return true; };
|
||||
const first = _scheduleEmailPrewarm(task, { delay: 1800 });
|
||||
const joined = _scheduleEmailPrewarm(task, { delay: 0 });
|
||||
const samePromise = first === joined;
|
||||
await advanceTo(1800);
|
||||
await fireNextIdle(7);
|
||||
const callsWhileBusy = taskCalls;
|
||||
while (now < 10300) {
|
||||
await advanceTo(now + 500);
|
||||
await fireNextIdle(7);
|
||||
}
|
||||
const result = await first;
|
||||
console.log(JSON.stringify({
|
||||
result, samePromise, callsWhileBusy, taskCalls, idleRequestCount,
|
||||
timers: timers.size, idleCallbacks: idleCallbacks.size,
|
||||
}));
|
||||
""")
|
||||
assert out == {
|
||||
"result": True,
|
||||
"samePromise": True,
|
||||
"callsWhileBusy": 0,
|
||||
"taskCalls": 1,
|
||||
"idleRequestCount": 18,
|
||||
"timers": 0,
|
||||
"idleCallbacks": 0,
|
||||
}
|
||||
|
||||
|
||||
def test_cancelled_prewarm_cannot_issue_a_delayed_duplicate():
|
||||
out = _run_scheduler_scenario("""
|
||||
let taskCalls = 0;
|
||||
const pending = _scheduleEmailPrewarm(async () => { taskCalls += 1; return true; }, { delay: 1800 });
|
||||
await advanceTo(1400);
|
||||
_cancelEmailPrewarm();
|
||||
await advanceTo(12000);
|
||||
const result = await pending;
|
||||
console.log(JSON.stringify({
|
||||
result, taskCalls, idleRequestCount,
|
||||
timers: timers.size, idleCallbacks: idleCallbacks.size,
|
||||
}));
|
||||
""")
|
||||
assert out == {
|
||||
"result": False,
|
||||
"taskCalls": 0,
|
||||
"idleRequestCount": 0,
|
||||
"timers": 0,
|
||||
"idleCallbacks": 0,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("transition", ["busy", "hidden"])
|
||||
def test_active_prewarm_is_aborted_and_retried_once_after_priority_transition(transition):
|
||||
block = (
|
||||
"window.__odysseusChatBusy = true; "
|
||||
"window.dispatchEvent({ type: 'odysseus:chat-busy-change' });"
|
||||
if transition == "busy"
|
||||
else "document.visibilityState = 'hidden'; document.dispatchEvent({ type: 'visibilitychange' });"
|
||||
)
|
||||
unblock = (
|
||||
"window.__odysseusChatBusy = false; window.__odysseusChatBusyUntil = now; "
|
||||
"window.dispatchEvent({ type: 'odysseus:chat-busy-change' });"
|
||||
if transition == "busy"
|
||||
else "document.visibilityState = 'visible'; document.dispatchEvent({ type: 'visibilitychange' });"
|
||||
)
|
||||
out = _run_scheduler_scenario(f"""
|
||||
let taskCalls = 0;
|
||||
let firstSignal = null;
|
||||
let finishFirst;
|
||||
const firstAttempt = new Promise(resolve => {{ finishFirst = resolve; }});
|
||||
const pending = _scheduleEmailPrewarm(async ({{ signal }}) => {{
|
||||
taskCalls += 1;
|
||||
if (taskCalls === 1) {{ firstSignal = signal; return firstAttempt; }}
|
||||
return true;
|
||||
}});
|
||||
await fireNextIdle(7);
|
||||
{block}
|
||||
const aborted = firstSignal.aborted;
|
||||
{unblock}
|
||||
const callsBeforeLateResult = taskCalls;
|
||||
finishFirst(true);
|
||||
await flushMicrotasks();
|
||||
const stillPendingAfterLateResult = _libPrewarmPromise === pending;
|
||||
await advanceTo(now + 500);
|
||||
await fireNextIdle(7);
|
||||
const result = await pending;
|
||||
console.log(JSON.stringify({{
|
||||
result, aborted, callsBeforeLateResult, taskCalls,
|
||||
stillPendingAfterLateResult,
|
||||
timers: timers.size, idleCallbacks: idleCallbacks.size,
|
||||
chatListeners: window.listenerCount('odysseus:chat-busy-change'),
|
||||
visibilityListeners: document.listenerCount('visibilitychange'),
|
||||
}}));
|
||||
""")
|
||||
assert out == {
|
||||
"result": True,
|
||||
"aborted": True,
|
||||
"callsBeforeLateResult": 1,
|
||||
"taskCalls": 2,
|
||||
"stillPendingAfterLateResult": True,
|
||||
"timers": 0,
|
||||
"idleCallbacks": 0,
|
||||
"chatListeners": 0,
|
||||
"visibilityListeners": 0,
|
||||
}
|
||||
|
||||
|
||||
def test_prewarm_skips_hidden_and_foreground_work():
|
||||
guard = _function_source("_canRunEmailPrewarm")
|
||||
|
||||
assert "state._libOpen" in guard
|
||||
assert "state._libLoading" in guard
|
||||
assert "_libSearchInFlight" in guard
|
||||
assert "document.visibilityState !== 'visible'" in guard
|
||||
assert "!_isChatInteractionBusy()" in guard
|
||||
|
||||
|
||||
def test_prewarm_selects_only_last_used_or_default_account():
|
||||
chooser = _function_source("_chooseEmailPrewarmAccountId")
|
||||
prewarm = _function_source("_prewarmEmailViews")
|
||||
|
||||
assert "_rememberedEmailAccountId()" in chooser
|
||||
assert "a.enabled !== false" in chooser
|
||||
assert "a.is_default" in chooser
|
||||
assert "enabled[0]" in chooser
|
||||
|
||||
assert "for (" not in prewarm
|
||||
assert "orderedAccountIds" not in prewarm
|
||||
assert "slice(0, 4)" not in prewarm
|
||||
assert "/api/email/folders" not in prewarm
|
||||
assert "/api/email/unread-state" not in prewarm
|
||||
assert prewarm.count("/api/email/list") == 1
|
||||
|
||||
|
||||
def test_prewarm_account_chooser_rejects_disabled_or_empty_authoritative_inventory():
|
||||
node = shutil.which("node")
|
||||
if not node:
|
||||
pytest.skip("node not on PATH")
|
||||
chooser = _function_source("_chooseEmailPrewarmAccountId")
|
||||
script = f"""
|
||||
const state = {{ _libAccountId: 'disabled-current' }};
|
||||
function _rememberedEmailAccountId() {{ return 'disabled-remembered'; }}
|
||||
{chooser}
|
||||
const onlyDisabled = _chooseEmailPrewarmAccountId([
|
||||
{{ id: 'disabled-remembered', enabled: false, is_default: true }},
|
||||
{{ id: 'disabled-current', enabled: false }},
|
||||
]);
|
||||
const empty = _chooseEmailPrewarmAccountId([]);
|
||||
const mixed = _chooseEmailPrewarmAccountId([
|
||||
{{ id: 'disabled-remembered', enabled: false, is_default: true }},
|
||||
{{ id: 'enabled-default', enabled: true, is_default: true }},
|
||||
]);
|
||||
console.log(JSON.stringify({{ onlyDisabled, empty, mixed }}));
|
||||
"""
|
||||
proc = subprocess.run(
|
||||
[node, "--input-type=module"],
|
||||
input=script,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=str(_REPO),
|
||||
timeout=30,
|
||||
)
|
||||
assert proc.returncode == 0, proc.stderr
|
||||
assert json.loads(proc.stdout.strip()) == {
|
||||
"onlyDisabled": "",
|
||||
"empty": "",
|
||||
"mixed": "enabled-default",
|
||||
}
|
||||
|
||||
ensure_accounts = _function_source("_ensureEmailAccountsForPrewarm")
|
||||
assert "if (!accountId) return null;" in ensure_accounts
|
||||
assert ensure_accounts.index("if (!accountId) return null;") < ensure_accounts.index("_publishActiveAccount();")
|
||||
|
||||
|
||||
def test_prewarm_is_bounded_to_the_interactive_initial_page_size():
|
||||
text = _source()
|
||||
prewarm = _function_source("_prewarmEmailViews")
|
||||
|
||||
assert "const _LIB_INITIAL_PAGE_SIZE = 100;" in text
|
||||
assert "limit: _LIB_INITIAL_PAGE_SIZE" in prewarm
|
||||
assert text.count("limit=${_LIB_INITIAL_PAGE_SIZE}&offset=${offsetAtStart}") == 2
|
||||
assert "limit: 100" not in prewarm
|
||||
|
||||
|
||||
def test_open_cancels_scheduled_or_inflight_prewarm_first():
|
||||
text = _source()
|
||||
cancel = _function_source("_cancelEmailPrewarm")
|
||||
open_library = _function_source("openEmailLibrary")
|
||||
|
||||
assert "clearTimeout(_libPrewarmDelayTimer)" in cancel
|
||||
assert "window.cancelIdleCallback(_libPrewarmIdleHandle)" in cancel
|
||||
assert "_libPrewarmAbortController?.abort()" in cancel
|
||||
assert "_libPrewarmGeneration += 1" in cancel
|
||||
assert open_library.index("_cancelEmailPrewarm();") < open_library.index("state._libOpen = true;")
|
||||
assert "_loadEmailsWhenChatIdle" not in text
|
||||
assert text.count("_loadEmails({ useCache: true });") >= 2
|
||||
|
||||
|
||||
def test_close_cancels_pending_prewarm_cleanup():
|
||||
close_library = _function_source("closeEmailLibrary")
|
||||
|
||||
assert close_library.index("_cancelEmailPrewarm();") < close_library.index("state._libOpen = false;")
|
||||
|
||||
|
||||
def test_unread_warm_joins_the_same_idle_single_flight_gate():
|
||||
unread_entry = _function_source("prewarmUnreadEmails")
|
||||
unread_work = _function_source("_prewarmUnreadEmailsNow")
|
||||
|
||||
assert "_scheduleEmailPrewarm(" in unread_entry
|
||||
assert "fetch(" not in unread_entry
|
||||
assert "_ensureEmailAccountsForPrewarm({ signal, generation })" in unread_work
|
||||
assert "signal" in unread_work
|
||||
assert "Math.min(20" in unread_work
|
||||
+122
-2
@@ -29,6 +29,7 @@ import base64
|
||||
import json
|
||||
import time
|
||||
import unittest.mock as mock
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -272,8 +273,14 @@ def _callback_endpoint():
|
||||
|
||||
|
||||
class _FakeRequest:
|
||||
"""Minimal stand-in for starlette Request — the callback only reads headers."""
|
||||
headers = {"host": "localhost:7000"}
|
||||
"""Minimal stand-in for starlette Request — the callback reads the Host header
|
||||
and the request scheme. Behind a TLS terminator uvicorn's proxy-headers
|
||||
middleware rewrites the scheme from `X-Forwarded-Proto`, so the route sees
|
||||
`https` there and `http` on a plain origin."""
|
||||
|
||||
def __init__(self, scheme="http", host="localhost:7000"):
|
||||
self.headers = {"host": host}
|
||||
self.url = SimpleNamespace(scheme=scheme)
|
||||
|
||||
|
||||
def _location(resp):
|
||||
@@ -415,6 +422,119 @@ async def test_callback_valid_owner_writes_encrypted_tokens_to_intended_account(
|
||||
assert other.oauth_access_token is None, "tokens must only touch the intended account"
|
||||
|
||||
|
||||
# ── Redirect URI scheme ───────────────────────────────────────────
|
||||
#
|
||||
# Google rejects the token exchange unless the callback's `redirect_uri` is
|
||||
# byte-identical to the one the authorize step sent, so both routes have to
|
||||
# agree — including on the scheme. Deriving it from the request keeps HTTPS
|
||||
# deployments working without pinning GOOGLE_OAUTH_REDIRECT_URI by hand;
|
||||
# hardcoding `http://` produced an unusable redirect behind any TLS front.
|
||||
|
||||
def _authorize_endpoint():
|
||||
"""Return the live google_oauth_authorize endpoint from the email router."""
|
||||
from routes.email_routes import setup_email_routes
|
||||
router = setup_email_routes()
|
||||
for route in router.routes:
|
||||
if route.path == "/api/email/oauth/google/authorize" and "GET" in getattr(route, "methods", set()):
|
||||
return route.endpoint
|
||||
raise AssertionError("google_oauth_authorize route not found")
|
||||
|
||||
|
||||
def _posted_redirect_uri(mock_post):
|
||||
"""Pull `redirect_uri` out of the mocked Google token-exchange POST."""
|
||||
return mock_post.call_args.kwargs["data"]["redirect_uri"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scheme", ("http", "https"))
|
||||
async def test_callback_redirect_uri_follows_the_request_scheme(scheme, monkeypatch):
|
||||
"""The token exchange must echo the scheme the request actually arrived on —
|
||||
`https` behind a TLS terminator, `http` on a plain origin."""
|
||||
from routes.email_helpers import make_oauth_state
|
||||
|
||||
monkeypatch.delenv("GOOGLE_OAUTH_REDIRECT_URI", raising=False)
|
||||
|
||||
db, Factory = _make_db()
|
||||
_make_account(db, account_id="acct-s", owner="alice", imap_user="alice@example.com")
|
||||
db.close()
|
||||
|
||||
token_resp = mock.MagicMock()
|
||||
token_resp.raise_for_status = mock.MagicMock()
|
||||
token_resp.json.return_value = {"access_token": "ya29.t", "refresh_token": "1//r", "expires_in": 3600}
|
||||
userinfo_resp = mock.MagicMock()
|
||||
userinfo_resp.is_success = True
|
||||
userinfo_resp.json.return_value = {"email": "alice@example.com", "name": "Alice"}
|
||||
|
||||
state = make_oauth_state("acct-s", "alice")
|
||||
|
||||
with mock.patch("httpx.post", return_value=token_resp) as mock_post, \
|
||||
mock.patch("httpx.get", return_value=userinfo_resp), \
|
||||
mock.patch("core.database.SessionLocal", Factory):
|
||||
callback = _callback_endpoint()
|
||||
await callback(
|
||||
code="4/code", state=state, error=None,
|
||||
request=_FakeRequest(scheme=scheme, host="odysseus.example.ts.net:7443"),
|
||||
)
|
||||
|
||||
assert _posted_redirect_uri(mock_post) == (
|
||||
f"{scheme}://odysseus.example.ts.net:7443/api/email/oauth/google/callback"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_redirect_uri_env_override_still_wins(monkeypatch):
|
||||
"""An explicit GOOGLE_OAUTH_REDIRECT_URI is used verbatim — deriving the
|
||||
scheme must not override a value the operator pinned by hand."""
|
||||
from routes.email_helpers import make_oauth_state
|
||||
|
||||
pinned = "https://mail.example.com/api/email/oauth/google/callback"
|
||||
monkeypatch.setenv("GOOGLE_OAUTH_REDIRECT_URI", pinned)
|
||||
|
||||
db, Factory = _make_db()
|
||||
_make_account(db, account_id="acct-p", owner="alice", imap_user="alice@example.com")
|
||||
db.close()
|
||||
|
||||
token_resp = mock.MagicMock()
|
||||
token_resp.raise_for_status = mock.MagicMock()
|
||||
token_resp.json.return_value = {"access_token": "ya29.t", "refresh_token": "1//r", "expires_in": 3600}
|
||||
userinfo_resp = mock.MagicMock()
|
||||
userinfo_resp.is_success = True
|
||||
userinfo_resp.json.return_value = {"email": "alice@example.com", "name": "Alice"}
|
||||
|
||||
state = make_oauth_state("acct-p", "alice")
|
||||
|
||||
with mock.patch("httpx.post", return_value=token_resp) as mock_post, \
|
||||
mock.patch("httpx.get", return_value=userinfo_resp), \
|
||||
mock.patch("core.database.SessionLocal", Factory):
|
||||
callback = _callback_endpoint()
|
||||
await callback(code="4/code", state=state, error=None, request=_FakeRequest(scheme="http"))
|
||||
|
||||
assert _posted_redirect_uri(mock_post) == pinned
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("scheme", ("http", "https"))
|
||||
async def test_authorize_redirect_uri_follows_the_request_scheme(scheme, monkeypatch):
|
||||
"""The authorize step builds the same redirect_uri the callback will send.
|
||||
`owner=""` is the unconfigured / single-user case, so no DB is touched."""
|
||||
import urllib.parse
|
||||
|
||||
monkeypatch.delenv("GOOGLE_OAUTH_REDIRECT_URI", raising=False)
|
||||
monkeypatch.setenv("GOOGLE_OAUTH_CLIENT_ID", "client-id.apps.googleusercontent.com")
|
||||
|
||||
authorize = _authorize_endpoint()
|
||||
resp = await authorize(
|
||||
account_id="acct-a",
|
||||
request=_FakeRequest(scheme=scheme, host="odysseus.example.ts.net:7443"),
|
||||
owner="",
|
||||
)
|
||||
|
||||
query = urllib.parse.parse_qs(urllib.parse.urlparse(resp.headers["location"]).query)
|
||||
assert query["redirect_uri"] == [
|
||||
f"{scheme}://odysseus.example.ts.net:7443/api/email/oauth/google/callback"
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_rejects_token_for_a_different_mailbox_identity():
|
||||
"""Reconnecting with another Google identity must not replace the token
|
||||
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Focused browser-side regression coverage for authoritative email opens."""
|
||||
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
_REPO = Path(__file__).resolve().parent.parent
|
||||
_INBOX_JS = _REPO / "static" / "js" / "emailInbox.js"
|
||||
_LIBRARY_JS = _REPO / "static" / "js" / "emailLibrary.js"
|
||||
_HAS_NODE = shutil.which("node") is not None
|
||||
|
||||
|
||||
def _extract_between(source: str, signature: str, next_marker: str) -> str:
|
||||
start = source.index(signature)
|
||||
end = source.index(next_marker, start)
|
||||
return source[start:end].rstrip()
|
||||
|
||||
|
||||
def test_library_unread_preview_has_one_authoritative_request_and_rollback():
|
||||
source = _LIBRARY_JS.read_text(encoding="utf-8")
|
||||
function = _extract_between(source, "async function _toggleCardPreview", "\n/**\n * Wrap a probable signature block")
|
||||
|
||||
assert function.count("/api/email/read/") == 1
|
||||
assert "/api/email/mark-read/" not in function
|
||||
assert "&mark_seen=true" in function
|
||||
assert "_syncEmailReadState(uidAtStart, true, readContext)" in function
|
||||
assert "_syncEmailReadState(uidAtStart, false, readContext)" in function
|
||||
assert "openGeneration === _emailCardOpenSeq" in function
|
||||
assert "_emailReadMutations.get(readContextKey)?.generation !== readMutation.generation" in function
|
||||
assert "authoritativeReadSucceeded = true;" in function
|
||||
assert "if (!authoritativeReadSucceeded) restoreUnreadState();" in function
|
||||
assert "if (!isCurrentOpen()) return" in function
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
def test_library_authoritative_success_defeats_newer_rollback_in_either_order():
|
||||
source = _LIBRARY_JS.read_text(encoding="utf-8")
|
||||
function = _extract_between(source, "async function _toggleCardPreview", "\n/**\n * Wrap a probable signature block")
|
||||
settlements = _extract_between(
|
||||
function,
|
||||
" const restoreUnreadState = () => {",
|
||||
"\n\n // Collapse any other expanded card",
|
||||
)
|
||||
|
||||
harness = f"""
|
||||
const _emailReadMutations = new Map();
|
||||
const readContextKey = 'same-mailbox-message';
|
||||
const uidAtStart = '1';
|
||||
const readContext = {{ accountId: 'acct-a', folder: 'INBOX', uid: '1' }};
|
||||
const readUpdates = [];
|
||||
function _syncEmailReadState(uid, isRead, context) {{
|
||||
readUpdates.push({{ uid, isRead, context }});
|
||||
}}
|
||||
function createSettlers(readMutation) {{
|
||||
{settlements}
|
||||
return {{ restoreUnreadState, commitReadState }};
|
||||
}}
|
||||
function runRace(successFirst) {{
|
||||
_emailReadMutations.clear();
|
||||
readUpdates.length = 0;
|
||||
const mutationA = {{ generation: 1, rollbackUnread: true }};
|
||||
_emailReadMutations.set(readContextKey, mutationA);
|
||||
const settlersA = createSettlers(mutationA);
|
||||
const mutationB = {{ generation: 2, rollbackUnread: true }};
|
||||
_emailReadMutations.set(readContextKey, mutationB);
|
||||
const settlersB = createSettlers(mutationB);
|
||||
if (successFirst) {{
|
||||
settlersA.commitReadState();
|
||||
settlersB.restoreUnreadState();
|
||||
}} else {{
|
||||
settlersB.restoreUnreadState();
|
||||
settlersA.commitReadState();
|
||||
}}
|
||||
return {{
|
||||
hasMutation: _emailReadMutations.has(readContextKey),
|
||||
readUpdates: readUpdates.map(update => update.isRead),
|
||||
}};
|
||||
}}
|
||||
console.log(JSON.stringify({{
|
||||
successFirst: runRace(true),
|
||||
failureFirst: runRace(false),
|
||||
}}));
|
||||
"""
|
||||
proc = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=harness,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=str(_REPO),
|
||||
timeout=30,
|
||||
)
|
||||
assert proc.returncode == 0, f"node failed: {proc.stderr}\n---\n{harness}"
|
||||
assert json.loads(proc.stdout.strip()) == {
|
||||
"successFirst": {"hasMutation": False, "readUpdates": [True]},
|
||||
"failureFirst": {"hasMutation": False, "readUpdates": [False, True]},
|
||||
}
|
||||
|
||||
|
||||
def test_library_reply_open_carries_immutable_mailbox_context():
|
||||
library_source = _LIBRARY_JS.read_text(encoding="utf-8")
|
||||
inbox_source = _INBOX_JS.read_text(encoding="utf-8")
|
||||
|
||||
assert "const mailboxGeneration = _emailMailboxGeneration;" in library_source
|
||||
assert "messageFolder = String(options.email?.folder || libraryFolder)" in library_source
|
||||
assert "return onEmailClick({ ...options, mailboxContext });" in library_source
|
||||
assert "mailboxContext?.messageFolder || _currentFolder" in inbox_source
|
||||
assert "mailboxContextIsCurrent()" in inbox_source
|
||||
assert "if (!isCurrentOpen()) return;\n let activeSid = await _createEmailChat" in inbox_source
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
def test_inbox_late_read_response_cannot_apply_after_newer_open():
|
||||
source = _INBOX_JS.read_text(encoding="utf-8")
|
||||
function = _extract_between(source, "async function _openEmail", "\nfunction _showEmailMenu")
|
||||
assert "let _openEmailRequestSeq = 0;" in source
|
||||
|
||||
harness = f"""
|
||||
const realLog = console.log;
|
||||
console.error = () => {{}};
|
||||
const API_BASE = 'https://odysseus.invalid';
|
||||
const window = {{ __odysseusActiveEmailAccount: 'acct-a' }};
|
||||
let _currentFolder = 'INBOX';
|
||||
const _acct = () => '&account_id=acct-a';
|
||||
let _openEmailRequestSeq = 0;
|
||||
let _docModule = null;
|
||||
const spinnerModule = {{ createWhirlpool() {{ throw new Error('spinner should not run'); }} }};
|
||||
const sessionModule = null;
|
||||
let firstResolve;
|
||||
const calls = [];
|
||||
async function fetch(url) {{
|
||||
calls.push(String(url));
|
||||
if (calls.length === 1) {{
|
||||
return await new Promise((resolve) => {{
|
||||
firstResolve = () => resolve({{ json: async () => ({{ uid: '1', subject: 'old' }}) }});
|
||||
}});
|
||||
}}
|
||||
return {{ json: async () => ({{ error: 'newer open completed test' }}) }};
|
||||
}}
|
||||
{function}
|
||||
const oldEmail = {{ uid: '1', is_read: false }};
|
||||
const newerEmail = {{ uid: '2', is_read: false }};
|
||||
const first = _openEmail(oldEmail, null);
|
||||
await Promise.resolve();
|
||||
const second = _openEmail(newerEmail, null);
|
||||
await second;
|
||||
firstResolve();
|
||||
await first;
|
||||
realLog(JSON.stringify({{ calls, oldRead: oldEmail.is_read, newerRead: newerEmail.is_read }}));
|
||||
"""
|
||||
proc = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=harness,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=str(_REPO),
|
||||
timeout=30,
|
||||
)
|
||||
assert proc.returncode == 0, f"node failed: {proc.stderr}\n---\n{harness}"
|
||||
result = json.loads(proc.stdout.strip())
|
||||
assert len(result["calls"]) == 2
|
||||
assert all("mark_seen=true" in url for url in result["calls"])
|
||||
assert result["oldRead"] is False
|
||||
assert result["newerRead"] is False
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
@pytest.mark.parametrize("context_change", ["account", "folder", "library"])
|
||||
def test_inbox_late_read_response_cannot_apply_after_mailbox_switch(context_change):
|
||||
source = _INBOX_JS.read_text(encoding="utf-8")
|
||||
function = _extract_between(source, "async function _openEmail", "\nfunction _showEmailMenu")
|
||||
|
||||
changes = {
|
||||
"account": "window.__odysseusActiveEmailAccount = 'acct-b';",
|
||||
"folder": "_currentFolder = 'Archive';",
|
||||
"library": "libraryCurrent = false;",
|
||||
}
|
||||
change = changes[context_change]
|
||||
open_call = (
|
||||
"_openEmail(email, null, null, 'reply', '', '', mailboxContext)"
|
||||
if context_change == "library"
|
||||
else "_openEmail(email, null)"
|
||||
)
|
||||
harness = f"""
|
||||
const realLog = console.log;
|
||||
console.error = () => {{}};
|
||||
const API_BASE = 'https://odysseus.invalid';
|
||||
const window = {{ __odysseusActiveEmailAccount: 'acct-a' }};
|
||||
let _currentFolder = 'INBOX';
|
||||
const _acct = () => '&account_id=acct-a';
|
||||
let _openEmailRequestSeq = 0;
|
||||
let libraryCurrent = true;
|
||||
const mailboxContext = {{
|
||||
accountId: 'acct-a',
|
||||
messageFolder: 'Archive',
|
||||
isCurrent: () => libraryCurrent,
|
||||
}};
|
||||
let createCalls = 0;
|
||||
let _docModule = {{}};
|
||||
async function _createEmailChat() {{ createCalls += 1; return 'stale-session'; }}
|
||||
const spinnerModule = {{ createWhirlpool() {{ throw new Error('spinner should not run'); }} }};
|
||||
const sessionModule = null;
|
||||
let resolveRead;
|
||||
async function fetch() {{
|
||||
return await new Promise((resolve) => {{
|
||||
resolveRead = () => resolve({{ json: async () => ({{ uid: '1', subject: 'old' }}) }});
|
||||
}});
|
||||
}}
|
||||
{function}
|
||||
const email = {{ uid: '1', is_read: false }};
|
||||
const pending = {open_call};
|
||||
await Promise.resolve();
|
||||
{change}
|
||||
resolveRead();
|
||||
await pending;
|
||||
realLog(JSON.stringify({{ createCalls, isRead: email.is_read }}));
|
||||
"""
|
||||
proc = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=harness,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=str(_REPO),
|
||||
timeout=30,
|
||||
)
|
||||
assert proc.returncode == 0, f"node failed: {proc.stderr}\n---\n{harness}"
|
||||
result = json.loads(proc.stdout.strip())
|
||||
assert result == {"createCalls": 0, "isRead": False}
|
||||
@@ -0,0 +1,278 @@
|
||||
import asyncio
|
||||
from contextlib import contextmanager
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
RAW_EMAIL = (
|
||||
b"From: Sender <sender@example.com>\r\n"
|
||||
b"To: Alice <alice@example.com>\r\n"
|
||||
b"Subject: Single authoritative open\r\n"
|
||||
b"Message-ID: <single-open@example.com>\r\n"
|
||||
b"Date: Tue, 04 Aug 2026 12:00:00 +0000\r\n"
|
||||
b"Content-Type: text/plain; charset=utf-8\r\n"
|
||||
b"\r\n"
|
||||
b"Body"
|
||||
)
|
||||
|
||||
|
||||
def _route_endpoint(router, path: str, method: str):
|
||||
method = method.upper()
|
||||
for route in router.routes:
|
||||
if route.path == path and method in getattr(route, "methods", set()):
|
||||
return route.endpoint
|
||||
raise AssertionError(f"route not found: {method} {path}")
|
||||
|
||||
|
||||
class FakeImap:
|
||||
def __init__(self, store_status="OK", readonly_mailbox=False):
|
||||
self.store_status = store_status
|
||||
# Shared archives and some provider folders reject a read-write SELECT.
|
||||
self.readonly_mailbox = readonly_mailbox
|
||||
self.selects = []
|
||||
self.commands = []
|
||||
|
||||
def select(self, mailbox, readonly=False):
|
||||
self.selects.append((mailbox, readonly))
|
||||
if self.readonly_mailbox and not readonly:
|
||||
raise OSError("[READ-ONLY] Mailbox is read-only")
|
||||
return "OK", [b"1"]
|
||||
|
||||
def uid(self, command, uid, *args):
|
||||
self.commands.append((command, uid, *args))
|
||||
if command == "FETCH":
|
||||
header, body = RAW_EMAIL.split(b"\r\n\r\n", 1)
|
||||
return "OK", [
|
||||
(b"1 (UID 42 BODY[HEADER])", header + b"\r\n\r\n"),
|
||||
(b"1 (UID 42 BODY[TEXT]<0>)", body),
|
||||
]
|
||||
if command == "STORE":
|
||||
# RFC 3501 STORE takes a parenthesized flag-list. GreenMail rejects
|
||||
# the formerly emitted bare ``\Seen`` atom with BAD, so keep the
|
||||
# fake strict enough to catch that provider-compatibility failure.
|
||||
if args != ("+FLAGS", "(\\Seen)"):
|
||||
return "BAD", [b"Expected:'(' found:'\\'"]
|
||||
return self.store_status, []
|
||||
raise AssertionError(f"unexpected IMAP command: {command}")
|
||||
|
||||
|
||||
def _install_fakes(monkeypatch, tmp_path, *, store_status="OK", readonly_mailbox=False):
|
||||
import routes.email_helpers as email_helpers
|
||||
import routes.email_routes as email_routes
|
||||
|
||||
db_path = tmp_path / "email.db"
|
||||
monkeypatch.setattr(email_helpers, "SCHEDULED_DB", db_path)
|
||||
monkeypatch.setattr(email_routes, "SCHEDULED_DB", db_path)
|
||||
email_helpers._init_scheduled_db()
|
||||
|
||||
connections = []
|
||||
indexed_updates = []
|
||||
|
||||
@contextmanager
|
||||
def fake_imap(account_id=None, owner=""):
|
||||
conn = FakeImap(store_status=store_status, readonly_mailbox=readonly_mailbox)
|
||||
connections.append(conn)
|
||||
yield conn
|
||||
|
||||
monkeypatch.setattr(email_routes, "_start_poller", lambda: None)
|
||||
monkeypatch.setattr(email_routes, "_imap", fake_imap)
|
||||
monkeypatch.setattr(email_routes, "_email_preview_cache_get", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(email_routes, "_email_preview_cache_put", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(email_routes, "_email_attachment_meta_cache_get", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(email_routes, "_email_attachment_meta_cache_put", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(
|
||||
email_routes,
|
||||
"_email_index_update_flags",
|
||||
lambda *args, **_kwargs: indexed_updates.append(args),
|
||||
)
|
||||
return email_routes, connections, indexed_updates
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("mark_seen", [True, False])
|
||||
async def test_read_email_seen_contract_uses_one_imap_connection(monkeypatch, tmp_path, mark_seen):
|
||||
email_routes, connections, indexed_updates = _install_fakes(monkeypatch, tmp_path)
|
||||
router = email_routes.setup_email_routes()
|
||||
read_email = _route_endpoint(router, "/api/email/read/{uid}", "GET")
|
||||
|
||||
result = await read_email(
|
||||
"42",
|
||||
folder="INBOX",
|
||||
account_id="acct-a",
|
||||
mark_seen=mark_seen,
|
||||
full=False,
|
||||
owner="alice",
|
||||
)
|
||||
|
||||
assert result["uid"] == "42"
|
||||
assert len(connections) == 1
|
||||
conn = connections[0]
|
||||
assert conn.selects == [(conn.selects[0][0], not mark_seen)]
|
||||
assert [command[0] for command in conn.commands] == (
|
||||
["FETCH", "STORE"] if mark_seen else ["FETCH"]
|
||||
)
|
||||
assert "BODY.PEEK[HEADER]" in conn.commands[0][2]
|
||||
if mark_seen:
|
||||
assert conn.commands[1][2:] == ("+FLAGS", "(\\Seen)")
|
||||
assert indexed_updates == [("alice", "acct-a", "INBOX", "42", "\\Seen", True)]
|
||||
else:
|
||||
assert indexed_updates == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cached_read_awaits_one_seen_store_without_refetch(monkeypatch, tmp_path):
|
||||
email_routes, connections, indexed_updates = _install_fakes(monkeypatch, tmp_path)
|
||||
router = email_routes.setup_email_routes()
|
||||
read_email = _route_endpoint(router, "/api/email/read/{uid}", "GET")
|
||||
|
||||
first = await read_email(
|
||||
"42", folder="INBOX", account_id="acct-a", mark_seen=False, full=False, owner="alice"
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
asyncio,
|
||||
"create_task",
|
||||
lambda *_args, **_kwargs: (_ for _ in ()).throw(
|
||||
AssertionError("cached mark-seen must be awaited, not scheduled")
|
||||
),
|
||||
)
|
||||
second = await read_email(
|
||||
"42", folder="INBOX", account_id="acct-a", mark_seen=True, full=False, owner="alice"
|
||||
)
|
||||
|
||||
assert first["message_id"] == second["message_id"]
|
||||
assert len(connections) == 2
|
||||
assert [command[0] for command in connections[0].commands] == ["FETCH"]
|
||||
assert [command[0] for command in connections[1].commands] == ["STORE"]
|
||||
assert connections[1].commands[0][2:] == ("+FLAGS", "(\\Seen)")
|
||||
assert connections[1].selects[0][1] is False
|
||||
assert indexed_updates == [("alice", "acct-a", "INBOX", "42", "\\Seen", True)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_seen_store_failure_returns_the_body_and_reports_the_failure(monkeypatch, tmp_path):
|
||||
"""A failed STORE must not cost the reader the message.
|
||||
|
||||
The body was fetched successfully before the flag update was attempted, so
|
||||
the response stays a normal read and carries `mark_seen_failed` for the
|
||||
client to roll its optimistic unread marker back.
|
||||
"""
|
||||
email_routes, connections, indexed_updates = _install_fakes(
|
||||
monkeypatch, tmp_path, store_status="NO"
|
||||
)
|
||||
router = email_routes.setup_email_routes()
|
||||
read_email = _route_endpoint(router, "/api/email/read/{uid}", "GET")
|
||||
|
||||
result = await read_email(
|
||||
"42", folder="INBOX", account_id="acct-a", mark_seen=True, full=False, owner="alice"
|
||||
)
|
||||
|
||||
assert "error" not in result
|
||||
assert result["uid"] == "42"
|
||||
assert result["mark_seen_failed"] is True
|
||||
assert len(connections) == 1
|
||||
assert [command[0] for command in connections[0].commands] == ["FETCH", "STORE"]
|
||||
# The local index must not claim a transition the provider rejected.
|
||||
assert indexed_updates == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_read_only_mailbox_serves_the_message_without_marking_seen(monkeypatch, tmp_path):
|
||||
"""A mailbox that refuses a read-write SELECT is still readable.
|
||||
|
||||
Opening the message is the user's actual goal; the \\Seen transition is a
|
||||
side effect of it. A folder that cannot accept flag changes must therefore
|
||||
fall back to a read-only selection rather than failing the open.
|
||||
"""
|
||||
email_routes, connections, indexed_updates = _install_fakes(
|
||||
monkeypatch, tmp_path, readonly_mailbox=True
|
||||
)
|
||||
router = email_routes.setup_email_routes()
|
||||
read_email = _route_endpoint(router, "/api/email/read/{uid}", "GET")
|
||||
|
||||
result = await read_email(
|
||||
"42", folder="Archive", account_id="acct-a", mark_seen=True, full=False, owner="alice"
|
||||
)
|
||||
|
||||
assert "error" not in result
|
||||
assert result["uid"] == "42"
|
||||
assert result["mark_seen_failed"] is True
|
||||
# Read-write attempt first, then the read-only retry on the same connection.
|
||||
assert [readonly for _mailbox, readonly in connections[0].selects] == [False, True]
|
||||
# No STORE is attempted once the mailbox is known to be read-only.
|
||||
assert [command[0] for command in connections[0].commands] == ["FETCH"]
|
||||
assert indexed_updates == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_seen_state_is_not_replayed_from_cache(monkeypatch, tmp_path):
|
||||
"""`mark_seen_failed` describes one request, not the stored message.
|
||||
|
||||
A second read that does not ask to mark seen must come back clean, or every
|
||||
later reader would inherit a STORE failure it never issued.
|
||||
"""
|
||||
email_routes, connections, _ = _install_fakes(monkeypatch, tmp_path, store_status="NO")
|
||||
router = email_routes.setup_email_routes()
|
||||
read_email = _route_endpoint(router, "/api/email/read/{uid}", "GET")
|
||||
|
||||
failed = await read_email(
|
||||
"42", folder="INBOX", account_id="acct-a", mark_seen=True, full=False, owner="alice"
|
||||
)
|
||||
replayed = await read_email(
|
||||
"42", folder="INBOX", account_id="acct-a", mark_seen=False, full=False, owner="alice"
|
||||
)
|
||||
|
||||
assert failed["mark_seen_failed"] is True
|
||||
assert replayed.get("mark_seen_failed", False) is False
|
||||
assert replayed["uid"] == "42"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unparseable_read_does_not_mark_seen(monkeypatch, tmp_path):
|
||||
email_routes, connections, indexed_updates = _install_fakes(monkeypatch, tmp_path)
|
||||
monkeypatch.setattr(
|
||||
email_routes.email_mod,
|
||||
"message_from_bytes",
|
||||
lambda *_args, **_kwargs: (_ for _ in ()).throw(ValueError("malformed message")),
|
||||
)
|
||||
router = email_routes.setup_email_routes()
|
||||
read_email = _route_endpoint(router, "/api/email/read/{uid}", "GET")
|
||||
|
||||
result = await read_email(
|
||||
"42", folder="INBOX", account_id="acct-a", mark_seen=True, full=False, owner="alice"
|
||||
)
|
||||
|
||||
assert result == {"error": "Mail operation failed"}
|
||||
assert len(connections) == 1
|
||||
assert [command[0] for command in connections[0].commands] == ["FETCH"]
|
||||
assert indexed_updates == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cached_seen_store_failure_returns_the_cached_body(monkeypatch, tmp_path):
|
||||
"""A cache hit already holds a complete message; a failed STORE cannot take it away.
|
||||
|
||||
This is the path where withholding the body would be least defensible — the
|
||||
response is served from memory and needed no network at all.
|
||||
"""
|
||||
email_routes, connections, indexed_updates = _install_fakes(
|
||||
monkeypatch, tmp_path, store_status="NO"
|
||||
)
|
||||
router = email_routes.setup_email_routes()
|
||||
read_email = _route_endpoint(router, "/api/email/read/{uid}", "GET")
|
||||
|
||||
first = await read_email(
|
||||
"42", folder="INBOX", account_id="acct-a", mark_seen=False, full=False, owner="alice"
|
||||
)
|
||||
second = await read_email(
|
||||
"42", folder="INBOX", account_id="acct-a", mark_seen=True, full=False, owner="alice"
|
||||
)
|
||||
|
||||
assert first["uid"] == "42"
|
||||
assert "error" not in second
|
||||
assert second["uid"] == "42"
|
||||
assert second["body"] == first["body"]
|
||||
assert second["mark_seen_failed"] is True
|
||||
assert len(connections) == 2
|
||||
assert [command[0] for command in connections[0].commands] == ["FETCH"]
|
||||
assert [command[0] for command in connections[1].commands] == ["STORE"]
|
||||
assert indexed_updates == []
|
||||
@@ -0,0 +1,52 @@
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
_REPO = Path(__file__).resolve().parent.parent
|
||||
_UTILS = (_REPO / "static" / "js" / "emailLibrary" / "utils.js").as_posix()
|
||||
_HAS_NODE = shutil.which("node") is not None
|
||||
|
||||
pytestmark = pytest.mark.skipif(not _HAS_NODE, reason="node binary not on PATH")
|
||||
|
||||
|
||||
def test_email_summary_renderer_ignores_untrusted_provider_error_text():
|
||||
secret = (
|
||||
"endpoint=https://private.example.internal/v1 provider=ollama "
|
||||
"model=private-model response_body=private-response "
|
||||
"Authorization: Bearer token-secret-value"
|
||||
)
|
||||
script = f"""
|
||||
import {{ _renderEmailSummaryError }} from '{_UTILS}';
|
||||
const host = {{
|
||||
ownerDocument: {{
|
||||
createElement() {{ return {{ style: {{}}, textContent: '' }}; }},
|
||||
}},
|
||||
replaceChildren(node) {{ this.child = node; }},
|
||||
}};
|
||||
_renderEmailSummaryError(host, {{
|
||||
error_code: 'email_summary_unavailable',
|
||||
error: {json.dumps(secret)},
|
||||
}});
|
||||
console.log(JSON.stringify({{
|
||||
text: host.child.textContent,
|
||||
color: host.child.style.color,
|
||||
}}));
|
||||
"""
|
||||
|
||||
proc = subprocess.run(
|
||||
["node", "--input-type=module"],
|
||||
input=script,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=str(_REPO),
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
assert proc.returncode == 0, proc.stderr
|
||||
rendered = json.loads(proc.stdout)
|
||||
assert rendered == {"text": "Failed to summarize", "color": "var(--red)"}
|
||||
assert secret not in proc.stdout
|
||||
@@ -0,0 +1,406 @@
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sqlite3
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
_TMP_DATA = Path(tempfile.mkdtemp(prefix="odysseus-email-summary-"))
|
||||
os.environ.setdefault("DATA_DIR", str(_TMP_DATA))
|
||||
os.environ.setdefault("DATABASE_URL", f"sqlite:///{_TMP_DATA / 'app.db'}")
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
||||
if str(PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
|
||||
def _route_endpoint(router, path: str, method: str):
|
||||
method = method.upper()
|
||||
for route in router.routes:
|
||||
if route.path == path and method in getattr(route, "methods", set()):
|
||||
return route.endpoint
|
||||
raise AssertionError(f"route not found: {method} {path}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_email_summary_uses_shared_llm_adapter(monkeypatch):
|
||||
import routes.email_helpers as email_helpers
|
||||
import src.llm_core as llm_core
|
||||
|
||||
calls = {}
|
||||
|
||||
async def fake_llm_call_async(url, model, messages, **kwargs):
|
||||
calls["url"] = url
|
||||
calls["model"] = model
|
||||
calls["messages"] = messages
|
||||
calls["kwargs"] = kwargs
|
||||
return "thinking before marker\n<<<SUMMARY>>>\n- Pay the invoice by Friday.\n<<<END>>>"
|
||||
|
||||
monkeypatch.setattr(llm_core, "llm_call_async", fake_llm_call_async)
|
||||
|
||||
summary = await email_helpers._generate_email_summary(
|
||||
url="https://chatgpt.com/backend-api/codex/responses",
|
||||
model="gpt-5.5",
|
||||
sender="Billing <billing@example.com>",
|
||||
subject="Invoice due",
|
||||
body_for_llm="Please pay invoice 123 by Friday.",
|
||||
headers={"Authorization": "Bearer test"},
|
||||
max_tokens=1234,
|
||||
timeout=45,
|
||||
)
|
||||
|
||||
assert summary == "- Pay the invoice by Friday."
|
||||
assert calls["url"] == "https://chatgpt.com/backend-api/codex/responses"
|
||||
assert calls["model"] == "gpt-5.5"
|
||||
assert calls["kwargs"]["headers"] == {"Authorization": "Bearer test"}
|
||||
assert calls["kwargs"]["temperature"] == 0.3
|
||||
assert calls["kwargs"]["max_tokens"] == 1234
|
||||
assert calls["kwargs"]["timeout"] == 45
|
||||
assert calls["kwargs"]["workload"] == "foreground"
|
||||
assert calls["messages"][0]["role"] == "system"
|
||||
assert calls["messages"][1]["role"] == "user"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scheduled_email_summary_uses_background_fallback_chain(monkeypatch):
|
||||
import routes.email_helpers as email_helpers
|
||||
import src.llm_core as llm_core
|
||||
import src.task_endpoint as task_endpoint
|
||||
|
||||
candidates = [
|
||||
("http://primary.invalid/v1", "primary-model", {"X-Candidate": "primary"}),
|
||||
("http://fallback.invalid/v1", "fallback-model", {"X-Candidate": "fallback"}),
|
||||
]
|
||||
resolve_calls = []
|
||||
wait_calls = []
|
||||
llm_calls = []
|
||||
|
||||
def fake_resolve_task_candidates(**kwargs):
|
||||
resolve_calls.append(kwargs)
|
||||
return candidates
|
||||
|
||||
async def fake_wait_for_interactive_quiet(label):
|
||||
wait_calls.append(label)
|
||||
return False
|
||||
|
||||
async def fake_llm_call_async(url, model, messages, **kwargs):
|
||||
llm_calls.append((url, model, messages, kwargs))
|
||||
if model == "primary-model":
|
||||
raise RuntimeError("primary unavailable")
|
||||
return "<<<SUMMARY>>>\n- Used the fallback model.\n<<<END>>>"
|
||||
|
||||
monkeypatch.setattr(task_endpoint, "resolve_task_candidates", fake_resolve_task_candidates)
|
||||
monkeypatch.setattr(task_endpoint, "wait_for_interactive_quiet", fake_wait_for_interactive_quiet)
|
||||
monkeypatch.setattr(llm_core, "llm_call_async", fake_llm_call_async)
|
||||
|
||||
summary = await email_helpers._generate_scheduled_email_summary(
|
||||
url="http://caller-fallback.invalid/v1",
|
||||
model="caller-fallback-model",
|
||||
sender="Sender <sender@example.com>",
|
||||
subject="Scheduled subject",
|
||||
body_for_llm="Please summarize this scheduled email.",
|
||||
headers={"Authorization": "Bearer test"},
|
||||
owner="alice",
|
||||
max_tokens=321,
|
||||
timeout=54,
|
||||
)
|
||||
|
||||
assert summary == "- Used the fallback model."
|
||||
assert resolve_calls == [{
|
||||
"fallback_url": "http://caller-fallback.invalid/v1",
|
||||
"fallback_model": "caller-fallback-model",
|
||||
"fallback_headers": {"Authorization": "Bearer test"},
|
||||
"owner": "alice",
|
||||
}]
|
||||
assert wait_calls == ["background task LLM"]
|
||||
assert [call[1] for call in llm_calls] == ["primary-model", "fallback-model"]
|
||||
assert all(call[3]["workload"] == "background" for call in llm_calls)
|
||||
assert all(call[3]["max_tokens"] == 321 for call in llm_calls)
|
||||
assert all(call[3]["timeout"] == 54 for call in llm_calls)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scheduled_local_summary_is_preempted_by_foreground_call(monkeypatch):
|
||||
import routes.email_helpers as email_helpers
|
||||
import src.llm_core as llm_core
|
||||
import src.task_endpoint as task_endpoint
|
||||
|
||||
local_url = "http://127.0.0.1:11434/v1/chat/completions"
|
||||
background_started = asyncio.Event()
|
||||
never_release = asyncio.Event()
|
||||
observed_workloads = []
|
||||
|
||||
monkeypatch.setenv("ODYSSEUS_LOCAL_MODEL_GATE", "true")
|
||||
monkeypatch.setenv("BACKGROUND_TASK_FOREGROUND_GATE", "false")
|
||||
monkeypatch.setattr(llm_core, "_LOCAL_MODEL_LOCK", asyncio.Lock())
|
||||
monkeypatch.setattr(llm_core, "_LOCAL_MODEL_CURRENT", {})
|
||||
monkeypatch.setattr(llm_core, "_LOCAL_MODEL_WAITING_FOREGROUND", 0)
|
||||
monkeypatch.setattr(
|
||||
task_endpoint,
|
||||
"resolve_task_candidates",
|
||||
lambda **_kwargs: [(local_url, "scheduled-model", {})],
|
||||
)
|
||||
|
||||
async def fake_wait_for_interactive_quiet(_label):
|
||||
return False
|
||||
|
||||
async def gated_llm_call(url, model, messages, **kwargs):
|
||||
assert messages
|
||||
workload = kwargs.get("workload")
|
||||
observed_workloads.append(workload)
|
||||
async with llm_core._local_model_slot(url, model, workload=workload):
|
||||
background_started.set()
|
||||
await never_release.wait()
|
||||
return "unreachable"
|
||||
|
||||
monkeypatch.setattr(task_endpoint, "wait_for_interactive_quiet", fake_wait_for_interactive_quiet)
|
||||
monkeypatch.setattr(llm_core, "llm_call_async", gated_llm_call)
|
||||
|
||||
background_task = asyncio.create_task(email_helpers._generate_scheduled_email_summary(
|
||||
url=local_url,
|
||||
model="scheduled-model",
|
||||
sender="Sender",
|
||||
subject="Scheduled",
|
||||
body_for_llm="Scheduled body",
|
||||
owner="alice",
|
||||
))
|
||||
foreground_task = None
|
||||
try:
|
||||
await asyncio.wait_for(background_started.wait(), timeout=1)
|
||||
|
||||
async def run_foreground():
|
||||
async with llm_core._local_model_slot(
|
||||
local_url,
|
||||
"interactive-model",
|
||||
workload="foreground",
|
||||
):
|
||||
return True
|
||||
|
||||
foreground_task = asyncio.create_task(run_foreground())
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await asyncio.wait_for(background_task, timeout=1)
|
||||
assert await asyncio.wait_for(foreground_task, timeout=1) is True
|
||||
assert observed_workloads == ["background"]
|
||||
finally:
|
||||
for task in (background_task, foreground_task):
|
||||
if task is not None and not task.done():
|
||||
task.cancel()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_manual_email_summary_uses_shared_helper_and_caches(tmp_path, monkeypatch):
|
||||
import routes.email_helpers as email_helpers
|
||||
import routes.email_routes as email_routes
|
||||
import src.endpoint_resolver as endpoint_resolver
|
||||
|
||||
db_path = tmp_path / "scheduled_emails.db"
|
||||
monkeypatch.setattr(email_helpers, "SCHEDULED_DB", db_path)
|
||||
monkeypatch.setattr(email_routes, "SCHEDULED_DB", db_path)
|
||||
email_helpers._init_scheduled_db()
|
||||
|
||||
resolve_calls = []
|
||||
|
||||
def fake_resolve_endpoint(kind, owner=None):
|
||||
resolve_calls.append((kind, owner))
|
||||
assert kind == "utility"
|
||||
assert owner == "alice"
|
||||
return (
|
||||
"https://chatgpt.com/backend-api/codex/responses",
|
||||
"gpt-5.5",
|
||||
{"Authorization": "Bearer test"},
|
||||
)
|
||||
|
||||
helper_calls = {}
|
||||
|
||||
async def fake_generate_email_summary(**kwargs):
|
||||
helper_calls.update(kwargs)
|
||||
return "- Manual summary"
|
||||
|
||||
monkeypatch.setattr(endpoint_resolver, "resolve_endpoint", fake_resolve_endpoint)
|
||||
monkeypatch.setattr(email_routes, "_generate_email_summary", fake_generate_email_summary)
|
||||
|
||||
router = email_routes.setup_email_routes()
|
||||
summarize = _route_endpoint(router, "/api/email/summarize", "POST")
|
||||
|
||||
result = await summarize(
|
||||
{
|
||||
"body": "This is a long enough email body for manual summary.",
|
||||
"subject": "Manual subject",
|
||||
"from": "Sender <sender@example.com>",
|
||||
"message_id": "<manual@example.com>",
|
||||
"folder": "INBOX",
|
||||
},
|
||||
owner="alice",
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"success": True,
|
||||
"summary": "- Manual summary",
|
||||
"model_used": "gpt-5.5",
|
||||
}
|
||||
assert resolve_calls == [("utility", "alice")]
|
||||
assert helper_calls["url"] == "https://chatgpt.com/backend-api/codex/responses"
|
||||
assert helper_calls["model"] == "gpt-5.5"
|
||||
assert helper_calls["headers"]["Authorization"] == "Bearer test"
|
||||
assert helper_calls["headers"]["Content-Type"] == "application/json"
|
||||
|
||||
conn = sqlite3.connect(db_path)
|
||||
try:
|
||||
row = conn.execute(
|
||||
"SELECT owner, summary, model_used FROM email_summaries WHERE message_id=?",
|
||||
("<manual@example.com>",),
|
||||
).fetchone()
|
||||
finally:
|
||||
conn.close()
|
||||
assert row == ("alice", "- Manual summary", "gpt-5.5")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("exception_kind", ["http", "runtime"])
|
||||
async def test_manual_email_summary_never_exposes_provider_exception(
|
||||
monkeypatch,
|
||||
caplog,
|
||||
exception_kind,
|
||||
):
|
||||
from fastapi import HTTPException
|
||||
import routes.email_routes as email_routes
|
||||
import src.endpoint_resolver as endpoint_resolver
|
||||
|
||||
secret_detail = (
|
||||
"endpoint=https://private.example.internal/v1 provider=ollama "
|
||||
"model=private-model response_body=private-response "
|
||||
"Authorization: Bearer token-secret-value"
|
||||
)
|
||||
|
||||
def fake_resolve_endpoint(kind, owner=None):
|
||||
assert kind == "utility"
|
||||
assert owner == "alice"
|
||||
return (
|
||||
"https://private.example.internal/v1",
|
||||
"private-model",
|
||||
{"Authorization": "Bearer token-secret-value"},
|
||||
)
|
||||
|
||||
async def fail_summary(**_kwargs):
|
||||
if exception_kind == "http":
|
||||
raise HTTPException(status_code=502, detail=secret_detail)
|
||||
raise RuntimeError(secret_detail)
|
||||
|
||||
monkeypatch.setattr(endpoint_resolver, "resolve_endpoint", fake_resolve_endpoint)
|
||||
monkeypatch.setattr(email_routes, "_generate_email_summary", fail_summary)
|
||||
caplog.set_level(logging.WARNING, logger=email_routes.__name__)
|
||||
|
||||
router = email_routes.setup_email_routes()
|
||||
summarize = _route_endpoint(router, "/api/email/summarize", "POST")
|
||||
result = await summarize(
|
||||
{
|
||||
"body": "This email body is long enough to summarize.",
|
||||
"subject": "Sensitive provider failure",
|
||||
"from": "Sender <sender@example.com>",
|
||||
},
|
||||
owner="alice",
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"success": False,
|
||||
"error": "Failed to summarize",
|
||||
"error_code": "email_summary_unavailable",
|
||||
}
|
||||
exposed = json.dumps(result) + caplog.text
|
||||
for marker in (
|
||||
"private.example.internal",
|
||||
"ollama",
|
||||
"private-model",
|
||||
"private-response",
|
||||
"token-secret-value",
|
||||
):
|
||||
assert marker not in exposed
|
||||
assert f"type={'HTTPException' if exception_kind == 'http' else 'RuntimeError'}" in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scheduled_email_summary_uses_shared_helper_and_caches(tmp_path, monkeypatch):
|
||||
import routes.email_helpers as email_helpers
|
||||
import routes.email_pollers as email_pollers
|
||||
|
||||
db_path = tmp_path / "scheduled_emails.db"
|
||||
monkeypatch.setattr(email_helpers, "SCHEDULED_DB", db_path)
|
||||
monkeypatch.setattr(email_pollers, "SCHEDULED_DB", db_path)
|
||||
email_helpers._init_scheduled_db()
|
||||
|
||||
raw_email = (
|
||||
b"From: Sender <sender@example.com>\r\n"
|
||||
b"To: Alice <alice@example.com>\r\n"
|
||||
b"Subject: Scheduled subject\r\n"
|
||||
b"Message-ID: <scheduled@example.com>\r\n"
|
||||
b"Date: Tue, 01 Jan 2026 12:00:00 +0000\r\n"
|
||||
b"Content-Type: text/plain; charset=utf-8\r\n"
|
||||
b"\r\n"
|
||||
+ (b"Please review this scheduled summary email. " * 8)
|
||||
)
|
||||
|
||||
class FakeImap:
|
||||
def __init__(self):
|
||||
self.logout_calls = 0
|
||||
|
||||
def select(self, _folder, readonly=True):
|
||||
return "OK", []
|
||||
|
||||
def uid(self, command, *args):
|
||||
if command == "SEARCH":
|
||||
return "OK", [b"1"]
|
||||
if command == "FETCH":
|
||||
return "OK", [(b"1 (RFC822)", raw_email)]
|
||||
raise AssertionError(f"unexpected uid command: {command!r} {args!r}")
|
||||
|
||||
def logout(self):
|
||||
self.logout_calls += 1
|
||||
|
||||
fake_conn = FakeImap()
|
||||
|
||||
def fake_resolve_task_candidates(owner=None):
|
||||
assert owner == "alice"
|
||||
return [(
|
||||
"https://chatgpt.com/backend-api/codex/responses",
|
||||
"gpt-5.5",
|
||||
{"Authorization": "Bearer test"},
|
||||
)]
|
||||
|
||||
helper_calls = {}
|
||||
|
||||
async def fake_generate_email_summary(**kwargs):
|
||||
helper_calls.update(kwargs)
|
||||
return "- Scheduled summary"
|
||||
|
||||
monkeypatch.setattr(email_pollers, "_load_settings", lambda: {"email_auto_summarize": True})
|
||||
monkeypatch.setattr(email_pollers, "_owner_for_email_account", lambda _account_id: "alice")
|
||||
monkeypatch.setattr(email_pollers, "_imap_connect", lambda account_id=None, owner="": fake_conn)
|
||||
monkeypatch.setattr(email_pollers, "_get_email_config", lambda account_id=None, owner="": {"from_address": "alice@example.com"})
|
||||
monkeypatch.setattr(email_pollers, "resolve_task_candidates", fake_resolve_task_candidates)
|
||||
monkeypatch.setattr(email_pollers, "_generate_scheduled_email_summary", fake_generate_email_summary)
|
||||
|
||||
result = await email_pollers._auto_summarize_pass_single(account_id="acct-alice")
|
||||
|
||||
assert "summarized 1" in result
|
||||
assert "summary failed" not in result
|
||||
assert helper_calls["url"] == "https://chatgpt.com/backend-api/codex/responses"
|
||||
assert helper_calls["model"] == "gpt-5.5"
|
||||
assert helper_calls["headers"]["Authorization"] == "Bearer test"
|
||||
assert helper_calls["headers"]["Content-Type"] == "application/json"
|
||||
assert helper_calls["owner"] == "alice"
|
||||
assert fake_conn.logout_calls == 1
|
||||
|
||||
conn = sqlite3.connect(db_path)
|
||||
try:
|
||||
row = conn.execute(
|
||||
"SELECT owner, summary, model_used FROM email_summaries WHERE message_id=?",
|
||||
("<scheduled@example.com>",),
|
||||
).fetchone()
|
||||
finally:
|
||||
conn.close()
|
||||
assert row == ("alice", "- Scheduled summary", "gpt-5.5")
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -34,6 +34,11 @@ class _FakeSessionManager:
|
||||
self.sessions = {"src-id": source}
|
||||
self.created = None
|
||||
|
||||
def get_session(self, session_id):
|
||||
# Fork looks the source up through get_session — the hydration seam —
|
||||
# so a session only present in the DB still forks a real transcript.
|
||||
return self.sessions[session_id]
|
||||
|
||||
def create_session(self, session_id=None, name=None, endpoint_url=None,
|
||||
model=None, rag=False, owner=None):
|
||||
self.created = _FakeSession(name=name, owner=owner)
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
import routes.gallery_routes as gallery_routes
|
||||
|
||||
|
||||
class _TorchSentinel:
|
||||
float32 = object()
|
||||
float64 = object()
|
||||
|
||||
|
||||
class _FakeTensor:
|
||||
def __init__(self, dtype):
|
||||
self.dtype = dtype
|
||||
self.to_args = None
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
self.to_args = (args, kwargs)
|
||||
return self
|
||||
|
||||
|
||||
def test_model_inputs_to_device_casts_mps_float64_to_float32():
|
||||
float_tensor = _FakeTensor(_TorchSentinel.float64)
|
||||
int_tensor = _FakeTensor("int64")
|
||||
plain_value = object()
|
||||
|
||||
result = gallery_routes._model_inputs_to_device(
|
||||
{"points": float_tensor, "labels": int_tensor, "plain": plain_value},
|
||||
"mps",
|
||||
_TorchSentinel,
|
||||
)
|
||||
|
||||
assert result["points"] is float_tensor
|
||||
assert float_tensor.to_args == ((), {"device": "mps", "dtype": _TorchSentinel.float32})
|
||||
assert int_tensor.to_args == (("mps",), {})
|
||||
assert result["plain"] is plain_value
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
"""Tools whose names collide with harmony built-ins must be aliased for gpt-oss.
|
||||
|
||||
gpt-oss (harmony format) ships BUILT-IN tools named `python` and `browser`,
|
||||
invoked with the raw body as the argument (`to=python` + bare source), while
|
||||
custom functions use `to=functions.NAME` + JSON. Exposing our own tool under a
|
||||
built-in's name makes the model answer with the built-in convention: it emits
|
||||
raw code, the server tries to parse it as JSON, and the request dies with
|
||||
"error parsing tool call: raw='import sys, ...'". Streaming is worse — Ollama
|
||||
truncates the stream instead of reporting it, so the turn looks like an empty
|
||||
response and the agent loop reads it as a stall.
|
||||
|
||||
Measured on gpt-oss:20b via Ollama /v1 with a fixed agentic prompt:
|
||||
python+bash as-is 2/6, python renamed 5/6, both renamed 6/6.
|
||||
|
||||
The aliasing is transport-only and gpt-oss-only: every other model's tool
|
||||
schemas must pass through untouched, and real tool names must come back out.
|
||||
"""
|
||||
from src.llm_core import (
|
||||
_alias_harmony_tools,
|
||||
_unalias_harmony_tool_name,
|
||||
_is_harmony_model,
|
||||
)
|
||||
|
||||
|
||||
def _tools(*names):
|
||||
return [
|
||||
{"type": "function", "function": {"name": n, "parameters": {}}}
|
||||
for n in names
|
||||
]
|
||||
|
||||
|
||||
def _names(tools):
|
||||
return [t["function"]["name"] for t in tools]
|
||||
|
||||
|
||||
def test_gpt_oss_colliding_names_are_aliased():
|
||||
out = _alias_harmony_tools(_tools("python", "bash", "web_search"), "gpt-oss:20b")
|
||||
assert _names(out) == ["run_python_code", "run_shell_command", "web_search"]
|
||||
|
||||
|
||||
def test_non_harmony_models_are_untouched():
|
||||
tools = _tools("python", "bash", "web_search")
|
||||
for model in ("qwen3-coder:30b", "gemma4:12b", "claude-opus-5", "gpt-4o", "llama-3.3"):
|
||||
out = _alias_harmony_tools(tools, model)
|
||||
assert _names(out) == ["python", "bash", "web_search"], model
|
||||
assert out is tools, f"{model} should get the same list object, not a copy"
|
||||
|
||||
|
||||
def test_aliasing_does_not_mutate_the_caller_list():
|
||||
tools = _tools("python")
|
||||
_alias_harmony_tools(tools, "gpt-oss:20b")
|
||||
assert _names(tools) == ["python"], "caller's schema list must not be mutated"
|
||||
|
||||
|
||||
def test_response_names_map_back_for_gpt_oss():
|
||||
assert _unalias_harmony_tool_name("run_python_code", "gpt-oss:20b") == "python"
|
||||
assert _unalias_harmony_tool_name("run_shell_command", "gpt-oss:20b") == "bash"
|
||||
# Unrelated names pass through untouched.
|
||||
assert _unalias_harmony_tool_name("web_search", "gpt-oss:20b") == "web_search"
|
||||
|
||||
|
||||
def test_response_names_untouched_for_other_models():
|
||||
# A non-harmony model that genuinely has a tool called run_python_code
|
||||
# must not have it rewritten to `python`.
|
||||
assert _unalias_harmony_tool_name("run_python_code", "qwen3-coder:30b") == "run_python_code"
|
||||
|
||||
|
||||
def test_harmony_detection():
|
||||
assert _is_harmony_model("gpt-oss:20b") is True
|
||||
assert _is_harmony_model("GPT-OSS:120B") is True
|
||||
assert _is_harmony_model("qwen3-coder:30b") is False
|
||||
assert _is_harmony_model("") is False
|
||||
assert _is_harmony_model(None) is False
|
||||
|
||||
|
||||
def test_empty_and_none_tools_are_safe():
|
||||
assert _alias_harmony_tools(None, "gpt-oss:20b") is None
|
||||
assert _alias_harmony_tools([], "gpt-oss:20b") == []
|
||||
@@ -5,9 +5,9 @@ The in-memory branch skips messages whose metadata has ``hidden`` (e.g.
|
||||
compaction summaries that are kept for AI context but not shown to the user).
|
||||
The DB fallback (taken when the in-memory history is empty, e.g. after a
|
||||
restart) built the client response from every DB row with no such filter, so
|
||||
hidden messages leaked to the client on DB-served sessions. The rebuilt
|
||||
in-memory ``session.history`` must still keep them, though, so only the response
|
||||
is filtered.
|
||||
hidden messages leaked to the client on DB-served sessions. Hydration of
|
||||
``session.history`` belongs to ``get_session``; this fallback only shapes the
|
||||
response, so only the response is filtered.
|
||||
|
||||
get_session_history depends on the DB, the session manager and a FastAPI
|
||||
request, so this pins the regression at the source level (as other route tests
|
||||
|
||||
@@ -0,0 +1,549 @@
|
||||
"""Display pagination must stay separate from full model-context hydration."""
|
||||
|
||||
import json
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import pytest
|
||||
from fastapi import APIRouter, FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
from starlette.requests import Request
|
||||
from sqlalchemy import create_engine, event
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from core.database import Base, ChatMessage as DbChatMessage, Session as DbSession
|
||||
from core.models import ChatMessage, Session
|
||||
from core.session_manager import SessionManager
|
||||
from routes import chat_routes
|
||||
from routes.history import history_routes
|
||||
from routes import session_routes
|
||||
from src.request_models import ChatRequest
|
||||
|
||||
|
||||
def _database():
|
||||
engine = create_engine(
|
||||
"sqlite://",
|
||||
connect_args={"check_same_thread": False},
|
||||
poolclass=StaticPool,
|
||||
)
|
||||
Base.metadata.create_all(
|
||||
engine,
|
||||
tables=[DbSession.__table__, DbChatMessage.__table__],
|
||||
)
|
||||
return engine, sessionmaker(bind=engine, autocommit=False, autoflush=False)
|
||||
|
||||
|
||||
def _seed_session(db_factory, *, session_id="session-1", message_count=6, stored_count=None):
|
||||
"""Seed `message_count` real rows; `stored_count` overrides the denormalized
|
||||
sessions.message_count column so drift can be reproduced."""
|
||||
db = db_factory()
|
||||
try:
|
||||
db.add(
|
||||
DbSession(
|
||||
id=session_id,
|
||||
name="Long chat",
|
||||
endpoint_url="http://model.test/v1",
|
||||
model="test-model",
|
||||
owner="alice",
|
||||
message_count=message_count if stored_count is None else stored_count,
|
||||
)
|
||||
)
|
||||
start = datetime(2026, 1, 1, 12, 0, 0)
|
||||
for index in range(message_count):
|
||||
db.add(
|
||||
DbChatMessage(
|
||||
id=f"message-{index}",
|
||||
session_id=session_id,
|
||||
role="user" if index % 2 == 0 else "assistant",
|
||||
content=f"content-{index}",
|
||||
timestamp=start + timedelta(seconds=index),
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _chat_message_selects(statements):
|
||||
return [
|
||||
" ".join(statement.lower().split())
|
||||
for statement in statements
|
||||
if statement.lstrip().lower().startswith("select")
|
||||
and "chat_messages" in statement.lower()
|
||||
]
|
||||
|
||||
|
||||
def _manager(db_factory, monkeypatch, sessions=None):
|
||||
"""A real SessionManager bound to the temp DB, with load counting."""
|
||||
monkeypatch.setattr("core.session_manager.SessionLocal", db_factory)
|
||||
manager = object.__new__(SessionManager)
|
||||
manager.upload_handler = None
|
||||
manager.sessions = sessions if sessions is not None else {}
|
||||
manager.full_loads = 0
|
||||
|
||||
original_load = manager._load_session_from_db
|
||||
|
||||
def counting_load(session_id):
|
||||
manager.full_loads += 1
|
||||
return original_load(session_id)
|
||||
|
||||
manager._load_session_from_db = counting_load
|
||||
return manager
|
||||
|
||||
|
||||
def test_paginated_history_reads_only_count_and_requested_page(monkeypatch):
|
||||
engine, db_factory = _database()
|
||||
_seed_session(db_factory)
|
||||
|
||||
class DisplayOnlyManager:
|
||||
def get_session(self, _session_id):
|
||||
raise AssertionError("paginated display history must not hydrate model context")
|
||||
|
||||
monkeypatch.setattr(history_routes, "SessionLocal", db_factory)
|
||||
monkeypatch.setattr(history_routes, "_verify_session_owner", lambda *_args: None)
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(history_routes.setup_history_routes(DisplayOnlyManager()))
|
||||
|
||||
statements = []
|
||||
|
||||
def capture_sql(_conn, _cursor, statement, _parameters, _context, _executemany):
|
||||
statements.append(statement)
|
||||
|
||||
event.listen(engine, "before_cursor_execute", capture_sql)
|
||||
try:
|
||||
response = TestClient(app).get("/api/history/session-1?limit=2")
|
||||
finally:
|
||||
event.remove(engine, "before_cursor_execute", capture_sql)
|
||||
engine.dispose()
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = response.json()
|
||||
assert [message["content"] for message in payload["history"]] == [
|
||||
"content-4",
|
||||
"content-5",
|
||||
]
|
||||
assert payload["total"] == 6
|
||||
assert payload["offset"] == 4
|
||||
assert payload["has_more_before"] is True
|
||||
assert payload["has_more_after"] is False
|
||||
|
||||
# One COUNT for the total plus one page read — never a full-transcript
|
||||
# select. The page bounds are asserted through the response above rather
|
||||
# than by matching SQL text.
|
||||
chat_selects = _chat_message_selects(statements)
|
||||
assert len(chat_selects) == 2, chat_selects
|
||||
assert sum("count(" in statement for statement in chat_selects) == 1
|
||||
|
||||
|
||||
def test_production_router_order_reaches_bounded_canonical_history(monkeypatch):
|
||||
"""The assembled app must not shadow canonical history with session routes."""
|
||||
engine, db_factory = _database()
|
||||
_seed_session(db_factory, message_count=1200)
|
||||
|
||||
class DisplayOnlyManager:
|
||||
def get_session(self, _session_id):
|
||||
raise AssertionError("bounded initial history must not hydrate all messages")
|
||||
|
||||
manager = DisplayOnlyManager()
|
||||
monkeypatch.setattr(
|
||||
session_routes,
|
||||
"router",
|
||||
APIRouter(prefix="/api", tags=["sessions"]),
|
||||
)
|
||||
monkeypatch.setattr(history_routes, "SessionLocal", db_factory)
|
||||
monkeypatch.setattr(history_routes, "_verify_session_owner", lambda *_args: None)
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(session_routes.setup_session_routes(manager, {}))
|
||||
app.include_router(history_routes.setup_history_routes(manager))
|
||||
|
||||
try:
|
||||
response = TestClient(app).get("/api/history/session-1?limit=24")
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.request.url.params["limit"] == "24"
|
||||
payload = response.json()
|
||||
displayed = len(payload["history"])
|
||||
assert 0 < displayed <= payload["limit"] <= 100
|
||||
assert payload["total"] >= 1200
|
||||
assert payload["has_more_before"] is True
|
||||
assert displayed < payload["total"]
|
||||
|
||||
|
||||
def test_incomplete_cached_history_hydrates_once_for_model_context(monkeypatch):
|
||||
engine, db_factory = _database()
|
||||
raw_multimodal = json.dumps(
|
||||
[
|
||||
{"type": "text", "text": "look at the source image"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,AAAA"},
|
||||
},
|
||||
]
|
||||
)
|
||||
db = db_factory()
|
||||
try:
|
||||
db.add(
|
||||
DbSession(
|
||||
id="session-1",
|
||||
name="Long chat",
|
||||
endpoint_url="http://model.test/v1",
|
||||
model="test-model",
|
||||
owner="alice",
|
||||
message_count=3,
|
||||
)
|
||||
)
|
||||
start = datetime(2026, 1, 1, 12, 0, 0)
|
||||
db.add_all(
|
||||
[
|
||||
DbChatMessage(
|
||||
id="message-0",
|
||||
session_id="session-1",
|
||||
role="user",
|
||||
content=raw_multimodal,
|
||||
meta_data=json.dumps(
|
||||
{
|
||||
"attachments": [
|
||||
{
|
||||
"id": "upload-1",
|
||||
"filename": "source.png",
|
||||
"content_type": "image/png",
|
||||
}
|
||||
]
|
||||
}
|
||||
),
|
||||
timestamp=start,
|
||||
),
|
||||
DbChatMessage(
|
||||
id="message-1",
|
||||
session_id="session-1",
|
||||
role="assistant",
|
||||
content="answer",
|
||||
timestamp=start + timedelta(seconds=1),
|
||||
),
|
||||
DbChatMessage(
|
||||
id="message-2",
|
||||
session_id="session-1",
|
||||
role="system",
|
||||
content="compaction summary",
|
||||
meta_data=json.dumps({"hidden": True}),
|
||||
timestamp=start + timedelta(seconds=2),
|
||||
),
|
||||
]
|
||||
)
|
||||
db.commit()
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
manager = _manager(
|
||||
db_factory,
|
||||
monkeypatch,
|
||||
sessions={
|
||||
"session-1": Session(
|
||||
id="session-1",
|
||||
name="Long chat",
|
||||
endpoint_url="http://model.test/v1",
|
||||
model="test-model",
|
||||
owner="alice",
|
||||
history=[ChatMessage("user", "stale partial cache")],
|
||||
# Deliberately stale too: get_session must refresh metadata before
|
||||
# checking whether the cached transcript is complete.
|
||||
message_count=1,
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
hydrated = manager.get_session("session-1")
|
||||
first_full_loads = manager.full_loads
|
||||
warm = manager.get_session("session-1")
|
||||
second_full_loads = manager.full_loads
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
assert hydrated is warm
|
||||
assert len(hydrated.history) == 3
|
||||
assert first_full_loads == 1
|
||||
assert second_full_loads == first_full_loads
|
||||
|
||||
context = hydrated.get_context_messages()
|
||||
assert len(context) == 3
|
||||
assert context[0]["content"][1]["image_url"]["url"] == "data:image/png;base64,AAAA"
|
||||
assert context[0]["metadata"]["attachments"] == [
|
||||
{
|
||||
"id": "upload-1",
|
||||
"filename": "source.png",
|
||||
"content_type": "image/png",
|
||||
}
|
||||
]
|
||||
hidden_summary = next(message for message in context if message["role"] == "system")
|
||||
assert hidden_summary["content"] == "compaction summary"
|
||||
assert hidden_summary["metadata"]["hidden"] is True
|
||||
|
||||
|
||||
def test_inflated_message_count_column_does_not_reload_warm_sessions(monkeypatch):
|
||||
"""A drifted-high sessions.message_count must not reload on every read.
|
||||
|
||||
`_persist_message` swallows a failed insert while `add_message` has already
|
||||
appended in memory, so the next successful persist writes rows+1. Keyed on
|
||||
that column, the hydration gate would stay true forever and re-select the
|
||||
whole transcript on every send, edit, delete and truncate.
|
||||
"""
|
||||
engine, db_factory = _database()
|
||||
_seed_session(db_factory, message_count=6, stored_count=8)
|
||||
|
||||
manager = _manager(db_factory, monkeypatch)
|
||||
try:
|
||||
session = manager.get_session("session-1")
|
||||
cold_loads = manager.full_loads
|
||||
for _ in range(3):
|
||||
manager.get_session("session-1")
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
assert len(session.history) == 6
|
||||
assert cold_loads == 1
|
||||
assert manager.full_loads == 1
|
||||
|
||||
|
||||
def test_stale_low_message_count_column_still_hydrates_for_the_model(monkeypatch):
|
||||
"""The other drift direction must not hand the model a truncated transcript.
|
||||
|
||||
`_persist_message` writes message_count = 0 when the session is not cached.
|
||||
A partly-filled cache plus that stale-low column previously left the send
|
||||
path with whatever RAM happened to hold.
|
||||
"""
|
||||
engine, db_factory = _database()
|
||||
_seed_session(db_factory, message_count=6, stored_count=0)
|
||||
|
||||
manager = _manager(
|
||||
db_factory,
|
||||
monkeypatch,
|
||||
sessions={
|
||||
"session-1": Session(
|
||||
id="session-1",
|
||||
name="Long chat",
|
||||
endpoint_url="http://model.test/v1",
|
||||
model="test-model",
|
||||
owner="alice",
|
||||
history=[ChatMessage("user", "content-0")],
|
||||
message_count=0,
|
||||
)
|
||||
},
|
||||
)
|
||||
try:
|
||||
session = manager.get_session("session-1")
|
||||
manager.get_session("session-1")
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
assert [message.content for message in session.history] == [
|
||||
f"content-{index}" for index in range(6)
|
||||
]
|
||||
assert manager.full_loads == 1
|
||||
|
||||
|
||||
def test_fork_after_restart_copies_the_real_transcript(monkeypatch):
|
||||
"""Forking reads source.history, so it must hydrate through get_session.
|
||||
|
||||
Display pagination no longer fills the cache, so a fork taken after a
|
||||
restart used to return HTTP 200 with an empty conversation.
|
||||
"""
|
||||
engine, db_factory = _database()
|
||||
_seed_session(db_factory, message_count=6)
|
||||
|
||||
# Restart state: metadata-only cache entry, exactly what load_sessions seeds.
|
||||
manager = _manager(db_factory, monkeypatch)
|
||||
manager.load_sessions()
|
||||
|
||||
monkeypatch.setattr(history_routes, "SessionLocal", db_factory)
|
||||
monkeypatch.setattr(history_routes, "_verify_session_owner", lambda *_args: None)
|
||||
monkeypatch.setattr("core.models._SESSION_MANAGER_INSTANCE", manager)
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(history_routes.setup_history_routes(manager))
|
||||
client = TestClient(app)
|
||||
|
||||
try:
|
||||
page = client.get("/api/history/session-1?limit=2")
|
||||
assert page.status_code == 200
|
||||
assert len(manager.sessions["session-1"].history) == 0
|
||||
|
||||
response = client.post("/api/session/session-1/fork", json={"keep_count": 4})
|
||||
assert response.status_code == 200
|
||||
payload = response.json()
|
||||
assert payload["kept"] == 4
|
||||
|
||||
forked = manager.get_session(payload["id"])
|
||||
assert [message.content for message in forked.history] == [
|
||||
f"content-{index}" for index in range(4)
|
||||
]
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
class _ContextBuildReached(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _ToolPolicy:
|
||||
block_all_tool_calls = False
|
||||
|
||||
def blocks(self, _tool_name):
|
||||
return False
|
||||
|
||||
|
||||
class _ChatHandler:
|
||||
async def handle_memory_command(self, _session, _message):
|
||||
return None
|
||||
|
||||
|
||||
def _json_request(path, payload):
|
||||
raw = json.dumps(payload).encode()
|
||||
sent = False
|
||||
|
||||
async def receive():
|
||||
nonlocal sent
|
||||
if sent:
|
||||
return {"type": "http.request", "body": b"", "more_body": False}
|
||||
sent = True
|
||||
return {"type": "http.request", "body": raw, "more_body": False}
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0"},
|
||||
"http_version": "1.1",
|
||||
"method": "POST",
|
||||
"scheme": "http",
|
||||
"path": path,
|
||||
"raw_path": path.encode(),
|
||||
"root_path": "",
|
||||
"query_string": b"",
|
||||
"headers": [(b"content-type", b"application/json")],
|
||||
"client": ("127.0.0.1", 1234),
|
||||
"server": ("testserver", 80),
|
||||
}
|
||||
return Request(scope, receive)
|
||||
|
||||
|
||||
def _route_endpoint(router, path):
|
||||
return next(route.endpoint for route in router.routes if route.path == path)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", ["/api/chat", "/api/chat_stream"])
|
||||
async def test_model_send_routes_hydrate_before_context_build(monkeypatch, path):
|
||||
# A real SessionManager over a real (temp) DB — a stub here would only
|
||||
# assert that the stub hydrates, not that SessionManager does.
|
||||
engine, db_factory = _database()
|
||||
_seed_session(db_factory, message_count=6, stored_count=8)
|
||||
manager = _manager(db_factory, monkeypatch)
|
||||
manager.load_sessions() # restart state: metadata only, no messages cached
|
||||
contexts_built = []
|
||||
|
||||
async def assert_complete_context(session, *_args, **_kwargs):
|
||||
contexts_built.append(session)
|
||||
assert [message.content for message in session.history] == [
|
||||
f"content-{index}" for index in range(6)
|
||||
]
|
||||
raise _ContextBuildReached
|
||||
|
||||
monkeypatch.setattr(chat_routes, "_set_user_time_from_request", lambda *_args: None)
|
||||
monkeypatch.setattr(chat_routes, "_verify_session_owner", lambda *_args: None)
|
||||
monkeypatch.setattr(chat_routes, "effective_user", lambda *_args: "alice")
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"_clear_orphaned_session_endpoint",
|
||||
lambda *_args, **_kwargs: False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"_recover_empty_session_model",
|
||||
lambda *_args, **_kwargs: False,
|
||||
)
|
||||
monkeypatch.setattr(chat_routes, "_enforce_chat_privileges", lambda *_args: None)
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"build_effective_tool_policy",
|
||||
lambda **_kwargs: _ToolPolicy(),
|
||||
)
|
||||
monkeypatch.setattr(chat_routes, "build_chat_context", assert_complete_context)
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"_resolve_request_workspace",
|
||||
lambda *_args: (None, False),
|
||||
)
|
||||
monkeypatch.setattr(chat_routes, "_classify_tool_intent", lambda *_args: None)
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"_is_contextual_web_followup",
|
||||
lambda *_args: False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"_is_contextual_browser_followup",
|
||||
lambda *_args: False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"_resolve_workspace_from_message_path",
|
||||
lambda *_args: (None, None),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"_reconcile_selected_route_from_request",
|
||||
lambda *_args, **_kwargs: False,
|
||||
)
|
||||
monkeypatch.setattr(chat_routes, "resolve_session_auth", lambda *_args, **_kwargs: None)
|
||||
monkeypatch.setattr(chat_routes, "get_session_mode", lambda *_args: "chat")
|
||||
monkeypatch.setattr(
|
||||
chat_routes,
|
||||
"_is_image_generation_session",
|
||||
lambda *_args, **_kwargs: False,
|
||||
)
|
||||
monkeypatch.setattr(chat_routes, "web_search_enabled_for_turn", lambda *_args: False)
|
||||
|
||||
router = chat_routes.setup_chat_routes(
|
||||
manager,
|
||||
_ChatHandler(),
|
||||
object(),
|
||||
object(),
|
||||
object(),
|
||||
object(),
|
||||
)
|
||||
endpoint = _route_endpoint(router, path)
|
||||
|
||||
async def send():
|
||||
if path == "/api/chat":
|
||||
await endpoint(
|
||||
_json_request(path, {}),
|
||||
ChatRequest(message="hello", session="session-1"),
|
||||
)
|
||||
else:
|
||||
await endpoint(
|
||||
_json_request(
|
||||
path,
|
||||
{"message": "hello", "session": "session-1"},
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
with pytest.raises(_ContextBuildReached):
|
||||
await send()
|
||||
first_loads = manager.full_loads
|
||||
|
||||
# Second send on the now-warm session: the transcript is complete, so
|
||||
# it must be served from RAM even though sessions.message_count is
|
||||
# still drifted high in the DB.
|
||||
with pytest.raises(_ContextBuildReached):
|
||||
await send()
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
assert len(contexts_built) == 2
|
||||
assert contexts_built[0] is contexts_built[1]
|
||||
assert first_loads == 1
|
||||
assert manager.full_loads == 1
|
||||
@@ -0,0 +1,19 @@
|
||||
"""Regression for idle UI polls that must not count as foreground activity."""
|
||||
|
||||
from src.interactive_gate import should_track_interactive_request
|
||||
|
||||
|
||||
def test_email_unread_state_is_passive_like_urgency_state():
|
||||
assert should_track_interactive_request("/api/email/urgency-state") is False
|
||||
assert should_track_interactive_request("/api/email/unread-state") is False
|
||||
|
||||
|
||||
def test_real_interactive_paths_still_tracked():
|
||||
assert should_track_interactive_request("/api/chat_stream") is True
|
||||
assert should_track_interactive_request("/api/email/messages") is True
|
||||
assert should_track_interactive_request("/api/tasks", method="POST") is True
|
||||
|
||||
|
||||
def test_options_never_tracked():
|
||||
assert should_track_interactive_request("/api/email/unread-state", method="OPTIONS") is False
|
||||
assert should_track_interactive_request("/api/chat_stream", method="OPTIONS") is False
|
||||
@@ -0,0 +1,31 @@
|
||||
"""The retired default fallback editor must not imply active routing."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
|
||||
_REPO = Path(__file__).resolve().parents[1]
|
||||
|
||||
|
||||
def test_legacy_default_fallback_editor_is_absent():
|
||||
soup = BeautifulSoup(
|
||||
(_REPO / "static" / "index.html").read_text(encoding="utf-8"),
|
||||
"html.parser",
|
||||
)
|
||||
editor = soup.find(id="set-defaultFallbacks")
|
||||
|
||||
assert editor is None
|
||||
assert soup.find(id="set-defaultAddFallback") is None
|
||||
|
||||
|
||||
def test_default_model_save_does_not_rewrite_legacy_fallbacks():
|
||||
source = (_REPO / "static" / "js" / "settings.js").read_text(encoding="utf-8")
|
||||
start = source.index("async function initDefaultChat()")
|
||||
end = source.index("/* ── Utility Model ── */", start)
|
||||
default_chat_source = source[start:end]
|
||||
|
||||
assert "settings.default_model_fallbacks" not in default_chat_source
|
||||
assert "default_model_fallbacks:" not in default_chat_source
|
||||
assert "set-defaultFallbacks" not in default_chat_source
|
||||
assert "set-defaultAddFallback" not in default_chat_source
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user