fix(state): close leaked SessionDB connections on exception paths (#83226)

SessionDB could leave native SQLite handles open when construction failed
partway through schema/pragma/FTS/repair/lock/interrupt handling. Other
short-lived callers (MCP reads/polling, session search, reactions, trace
upload, insights, shutdown recovery) opened temporary SessionDB handles
without a complete ownership boundary. API-server profile caches and
RetainDB shutdown had similar late-close races. Under sustained load this
exhausted file descriptors (EMFILE).

- Close partially initialized SessionDB connections on every constructor
  exception path via a finally block guarded by an initialization-complete
  flag.
- Close temporary/cross-profile SessionDB handles in finally blocks across
  CLI, MCP, search, trace, reactions, insights, and recovery paths.
- Add API-server per-profile cache ownership and disconnect cleanup.
- Make RetainDB writer-queue shutdown exception-safe: track connections per
  thread, close on worker exit, reject new enqueues after shutdown starts,
  and sweep any connections left by short-lived threads.
- Add regression coverage for constructor failures, worker-thread readers,
  API disconnect failures, shutdown recovery, RetainDB late enqueue, and
  foreign-loop async clients.

Salvage notes: the original PR's per-thread WAL-reader ownership changes
were superseded by main's read-connection pool (permits + checkout/return);
its cron timeout-abandon fix is credited separately to #72822's earlier
identical fix.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
joaomarcos
2026-08-10 23:12:01 -03:00
committed by Teknium
parent 0d91ab8889
commit 39e480c051
25 changed files with 1035 additions and 138 deletions
+98 -16
View File
@@ -7523,17 +7523,71 @@ def _force_close_async_httpx(client: Any) -> None:
pass
def _close_cached_client(client: Any) -> None:
"""Apply the canonical best-effort close policy to one cached client."""
if client is None:
return
_force_close_async_httpx(client)
def _schedule_async_close(close_result: Any, client: Any) -> None:
"""Finish an async close without leaking an unawaited coroutine."""
async def _await_close() -> None:
try:
close_fn = getattr(client, "close", None)
if callable(close_fn) and not inspect.iscoroutinefunction(close_fn):
close_fn()
await close_result
except Exception:
pass
finally:
_force_close_async_httpx(client)
runner = _await_close()
try:
import asyncio as _aio
try:
loop = _aio.get_running_loop()
except RuntimeError:
_aio.run(runner)
else:
task = loop.create_task(runner)
def _consume(completed_task) -> None:
try:
completed_task.exception()
except BaseException:
pass
task.add_done_callback(_consume)
runner = None
except Exception:
if runner is not None:
try:
runner.close()
except Exception:
pass
_force_close_async_httpx(client)
def _close_cached_client(client: Any, *, close_async: bool = False) -> None:
"""Close one cached client, awaiting async transports only when safe."""
if client is None:
return
close_fn = getattr(client, "close", None)
if not callable(close_fn):
_force_close_async_httpx(client)
return
try:
close_result = close_fn()
except Exception:
_force_close_async_httpx(client)
return
if inspect.isawaitable(close_result):
if close_async:
_schedule_async_close(close_result, client)
else:
# Do not await a client owned by another live event loop.
# Closing the coroutine avoids an unawaited-coroutine warning;
# the transport is still neutered for safe eventual GC.
try:
close_result.close()
except Exception:
pass
_force_close_async_httpx(client)
return
_force_close_async_httpx(client)
def shutdown_cached_clients() -> None:
@@ -7541,14 +7595,34 @@ def shutdown_cached_clients() -> None:
Call this during CLI shutdown, *before* the event loop is closed, to
avoid ``AsyncHttpxClientWrapper.__del__`` raising on a dead loop.
Snapshot and clear the cache under the lock, then close transports outside
it. Async transport shutdown may block while an owner loop drains; holding
the global cache lock during that wait stalls unrelated auxiliary callers
and can turn teardown into a process-wide lock convoy.
"""
with _client_cache_lock:
for key, entry in list(_client_cache.items()):
client = entry[0]
if client is None:
continue
_close_cached_client(client)
clients = [
(entry[0], entry[2])
for entry in _client_cache.values()
if entry[0] is not None
]
_client_cache.clear()
try:
import asyncio as _aio
running_loop = _aio.get_running_loop()
except RuntimeError:
running_loop = None
for client, owner_loop in clients:
# A live foreign loop owns its async transport. Calling its coroutine
# on this thread can bind/close sockets from the wrong loop; neuter it
# and let that owner finish teardown. Closed loops are safe to drain
# locally, and the current loop can await its own client.
close_async = owner_loop is not None and (
owner_loop.is_closed() or owner_loop is running_loop
)
_close_cached_client(client, close_async=close_async)
def cleanup_stale_async_clients() -> None:
@@ -7559,15 +7633,18 @@ def cleanup_stale_async_clients() -> None:
This is defense-in-depth — the primary fix is ``neuter_async_httpx_del``
which disables ``__del__`` entirely.
"""
stale_clients = []
with _client_cache_lock:
stale_keys = []
for key, entry in _client_cache.items():
client, _default, cached_loop = entry
if cached_loop is not None and cached_loop.is_closed():
_force_close_async_httpx(client)
stale_keys.append(key)
stale_clients.append(client)
for key in stale_keys:
del _client_cache[key]
for client in stale_clients:
_close_cached_client(client, close_async=True)
def _is_openrouter_client(client: Any) -> bool:
@@ -7660,7 +7737,12 @@ def _get_cached_client(
effective = _compat_model(cached_client, model, cached_default)
return cached_client, effective
# Stale — evict and fall through to create a new client.
_force_close_async_httpx(cached_client)
# Only a client whose owner loop is closed may be awaited from
# this thread; a live foreign loop remains force-neutered.
owner_loop_closed = (
cached_loop is not None and cached_loop.is_closed()
)
_close_cached_client(cached_client, close_async=owner_loop_closed)
del _client_cache[cache_key]
else:
effective = _compat_model(cached_client, model, cached_default)
@@ -7710,7 +7792,7 @@ def _get_cached_client(
client, default_model, _ = _client_cache[cache_key]
# This concurrently built loser was never exposed to a caller,
# so it is safe to close immediately.
_close_cached_client(built_client)
_close_cached_client(built_client, close_async=async_mode)
return client, model or default_model
+6
View File
@@ -336,10 +336,16 @@ def load_session_messages(
"""
from hermes_state import SessionDB
db = SessionDB(db_path=db_path) if db_path else SessionDB()
try:
resolved = db.resolve_session_id(session_id) or session_id
meta = db.get_session(resolved) or {}
messages = db.get_messages_as_conversation(resolved)
return messages, meta
finally:
try:
db.close()
except Exception:
logger.debug("Failed to close trace-upload SessionDB", exc_info=True)
def upload_session_trace(
+2
View File
@@ -12510,9 +12510,11 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin):
from agent.insights import InsightsEngine
db = SessionDB()
try:
engine = InsightsEngine(db)
report = engine.generate(days=days, source=source)
print(engine.format_terminal(report))
finally:
db.close()
except Exception as e:
print(f" Error generating insights: {e}")
+37 -12
View File
@@ -1437,6 +1437,9 @@ class APIServerAdapter(BasePlatformAdapter):
# in-flight run by run_id.
self._run_approval_sessions: Dict[str, str] = {}
self._session_db: Optional[Any] = None # Lazy-init SessionDB for session continuity
self._session_dbs: Dict[str, Any] = {}
self._session_db_cache_lock = threading.Lock()
self._session_db_cache_closed = False
# Last-known-good resolved model per session (keyed by gateway_session_key
# ONLY — never session_id, which rotates/is ephemeral for one-off API
# server requests; "*" is the process-wide fallback), mirroring
@@ -2182,16 +2185,30 @@ class APIServerAdapter(BasePlatformAdapter):
from hermes_state import SessionDB
key = str(home)
cache = getattr(self, "_session_dbs", None)
if cache is None:
cache = {}
self._session_dbs = cache
db = cache.get(key)
with self._session_db_cache_lock:
if self._session_db_cache_closed:
return None
db = self._session_dbs.get(key)
if db is None:
db = SessionDB(db_path=home / "state.db")
cache[key] = db
self._session_dbs[key] = db
return db
def _close_cached_session_dbs(self) -> None:
"""Close SessionDB handles owned by this adapter's profile cache."""
with self._session_db_cache_lock:
self._session_db_cache_closed = True
cached = list(self._session_dbs.values())
self._session_dbs.clear()
shared_db = getattr(self, "_session_db", None)
for db in cached:
if db is shared_db:
continue
try:
db.close()
except Exception:
logger.debug("Failed to close API-server SessionDB", exc_info=True)
def _ensure_session_db(self):
"""Lazily initialise and return the SessionDB for the active profile home.
@@ -2232,15 +2249,17 @@ class APIServerAdapter(BasePlatformAdapter):
home = get_hermes_home()
key = str(home)
cache = getattr(self, "_session_dbs", None)
if cache is not None and cache.get(key) is not None:
return cache[key]
with self._session_db_cache_lock:
cached = self._session_dbs.get(key)
if cached is not None:
return cached
if self._session_db_lock is None:
self._session_db_lock = asyncio.Lock()
async with self._session_db_lock:
cache = getattr(self, "_session_dbs", None)
if cache is not None and cache.get(key) is not None:
return cache[key]
with self._session_db_cache_lock:
cached = self._session_dbs.get(key)
if cached is not None:
return cached
return await asyncio.to_thread(self._open_and_cache_session_db, home)
except Exception as e:
logger.debug("SessionDB unavailable for API server: %s", e)
@@ -7368,6 +7387,9 @@ class APIServerAdapter(BasePlatformAdapter):
logger.warning("[%s] aiohttp not installed", self.name)
return False
with self._session_db_cache_lock:
self._session_db_cache_closed = False
if not self._api_key_passes_startup_guard():
# A rejected API_SERVER_KEY is a configuration error, not a
# transient blip — the key will not become valid on its own. A
@@ -7538,12 +7560,15 @@ class APIServerAdapter(BasePlatformAdapter):
logger.debug(
"Failed to close response store for %s", self.name, exc_info=True,
)
try:
if self._site:
await self._site.stop()
self._site = None
if self._runner:
await self._runner.cleanup()
self._runner = None
finally:
self._close_cached_session_dbs()
self._app = None
logger.info("[%s] API server stopped", self.name)
+13 -5
View File
@@ -316,6 +316,14 @@ def recover_pending_to_db(
session_db = SessionDB()
own_db = True
def _close_owned_db() -> None:
if not own_db:
return
try:
session_db.close()
except Exception:
pass
recovered = 0
for path in flush_files:
try:
@@ -389,6 +397,10 @@ def recover_pending_to_db(
)
recovered += 1
path.unlink(missing_ok=True)
except BaseException:
# Shutdown cancellation/interrupt must not strand an owned DB.
_close_owned_db()
raise
except Exception as exc:
logger.warning(
"Failed to recover pending message from %s: %s",
@@ -396,11 +408,7 @@ def recover_pending_to_db(
)
# Leave the file for next startup retry.
if own_db:
try:
session_db.close()
except Exception:
pass
_close_owned_db()
if recovered:
logger.info(
+23 -3
View File
@@ -1558,6 +1558,7 @@ def _resolve_session_by_name_or_id(name_or_id: str) -> Optional[str]:
from an exit summary printed before the bug fix, or from notes) get
resumed at the live tip instead of a stale parent with no messages.
"""
db = None
try:
from hermes_state import SessionDB
@@ -1580,10 +1581,15 @@ def _resolve_session_by_name_or_id(name_or_id: str) -> Optional[str]:
except Exception:
pass
db.close()
return resolved_id
except Exception:
pass
finally:
if db is not None:
try:
db.close()
except Exception:
pass
return None
@@ -2681,10 +2687,12 @@ def cmd_chat(args):
and not getattr(args, "no_restore_cwd", False)
and not getattr(args, "worktree", False)
):
_resume_db = None
try:
from hermes_state import SessionDB
_saved_cwd = ((SessionDB().get_session(args.resume) or {}).get("cwd") or "").strip()
_resume_db = SessionDB()
_saved_cwd = ((_resume_db.get_session(args.resume) or {}).get("cwd") or "").strip()
if _saved_cwd and not os.path.isdir(_saved_cwd):
print(f"⚠ session's recorded dir is gone ({_saved_cwd}); staying in {os.getcwd()}")
elif _saved_cwd and os.path.realpath(_saved_cwd) != os.path.realpath(os.getcwd()):
@@ -2692,6 +2700,12 @@ def cmd_chat(args):
print(f"↪ restored workspace dir: {_saved_cwd}")
except Exception:
pass # never let cwd-restore break a resume
finally:
if _resume_db is not None:
try:
_resume_db.close()
except Exception:
pass
# xAI retirement warning — one-shot, non-blocking, never fails startup
try:
@@ -11472,6 +11486,7 @@ def cmd_tools(args):
def cmd_insights(args):
db = None
try:
from hermes_state import SessionDB
from agent.insights import InsightsEngine
@@ -11480,9 +11495,14 @@ def cmd_insights(args):
engine = InsightsEngine(db)
report = engine.generate(days=args.days, source=args.source)
print(engine.format_terminal(report))
db.close()
except Exception as e:
print(f"Error generating insights: {e}")
finally:
if db is not None:
try:
db.close()
except Exception:
pass
def cmd_monitoring(args):
+22 -4
View File
@@ -2836,6 +2836,16 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
data["system_prompt"] = resolved
return data
@staticmethod
def _close_connection_quietly(conn: Optional[sqlite3.Connection]) -> None:
"""Close a partially initialized connection without masking its error."""
if conn is None:
return
try:
conn.close()
except Exception:
logger.debug("Could not close a SessionDB connection", exc_info=True)
def __init__(self, db_path: Path = None, read_only: bool = False):
self.db_path = db_path or _default_db_path()
# Fail hard (before any connection/pragma/mkdir) if a pytest-context
@@ -2933,6 +2943,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
self._token_writer_thread: Optional[threading.Thread] = None
self._token_writer_stop = False
self._token_writer_busy = False
initialization_complete = False
try:
if read_only:
# Read-only attach for cross-profile aggregation: SELECT-only,
@@ -2957,8 +2968,9 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
# only so read-only search keeps its FTS and trigram paths.
# Close the connection on ANY probe failure (e.g. malformed
# schema raises DatabaseError, not the OperationalError the
# probe handles): the outer except re-raises without cleanup,
# and a leaked tracked connection blocks _backup_db_file's
# probe handles). The constructor's outer finally also covers
# failures before this probe and BaseException paths, so a
# leaked tracked connection cannot block _backup_db_file's
# raw-copy for the rest of the process — the writable heal
# that follows would then repair WITHOUT its forensic backup.
try:
@@ -2982,6 +2994,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
except Exception:
pass
raise
initialization_complete = True
return
self.db_path.parent.mkdir(parents=True, exist_ok=True)
@@ -3115,6 +3128,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
# racing session lifecycle and the surprise disk/latency cost on
# an unattended open. (An interrupted optimize resumes when the
# user re-runs the command.)
initialization_complete = True
except Exception as exc:
# Capture the cause so /resume and friends can surface WHY the
# session DB is unavailable instead of a bare "Session database
@@ -3130,6 +3144,10 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
# ``hermes_state._set_last_init_error(None)`` explicitly.
_set_last_init_error(f"{type(exc).__name__}: {exc}")
raise
finally:
if not initialization_complete:
conn, self._conn = self._conn, None
self._close_connection_quietly(conn)
# ── Read-path split ──
@@ -4001,8 +4019,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
"WAL checkpoint (PASSIVE) at close failed: %s",
exc,
)
self._conn.close()
self._conn = None
conn, self._conn = self._conn, None
self._close_connection_quietly(conn)
# ── Chunked FTS rebuild engine (v23 opt-in optimize) ──
#
+40 -20
View File
@@ -79,6 +79,22 @@ def _get_session_db():
return None
def _load_session_messages(session_id: str):
"""Read one session and close the temporary database handle."""
db = _get_session_db()
if db is None:
return None, "Session database unavailable"
try:
return db.get_messages(session_id), None
except Exception as e:
return None, f"Failed to read messages: {e}"
finally:
try:
db.close()
except Exception:
logger.debug("Failed to close MCP SessionDB", exc_info=True)
def _load_sessions_index() -> dict:
"""Load the gateway session routing index.
@@ -448,6 +464,18 @@ class EventBridge:
self._new_event.set()
def _establish_baseline(self) -> None:
db = _get_session_db()
if not db:
return
try:
self._establish_baseline_with_db(db)
finally:
try:
db.close()
except Exception:
logger.debug("Failed to close MCP baseline SessionDB", exc_info=True)
def _establish_baseline_with_db(self, db) -> None:
"""Record the latest per-session message timestamp and the current
state.db mtime WITHOUT emitting events, so startup does not replay
history (#13414).
@@ -457,9 +485,6 @@ class EventBridge:
last_seen=0.0 in _poll_once, so a brand-new conversation's first
message is still delivered on its state.db-change tick.
"""
db = _get_session_db()
if not db:
return
try:
from hermes_constants import get_hermes_home
db_file = get_hermes_home() / "state.db"
@@ -486,7 +511,6 @@ class EventBridge:
latest = max(all_ts)
if latest > 0.0:
self._last_poll_timestamps[session_key] = latest
def _poll_loop(self):
"""Background loop: poll SessionDB for new messages."""
db = _get_session_db()
@@ -494,12 +518,18 @@ class EventBridge:
logger.warning("EventBridge: SessionDB unavailable, event polling disabled")
return
try:
while self._running:
try:
self._poll_once(db)
except Exception as e:
logger.debug("EventBridge poll error: %s", e)
time.sleep(POLL_INTERVAL)
finally:
try:
db.close()
except Exception:
logger.debug("Failed to close MCP polling SessionDB", exc_info=True)
def _poll_once(self, db):
"""Check for new messages across all sessions.
@@ -722,14 +752,9 @@ def create_mcp_server(event_bridge: Optional[EventBridge] = None) -> "FastMCP":
if not session_id:
return json.dumps({"error": "No session ID for this conversation"})
db = _get_session_db()
if not db:
return json.dumps({"error": "Session database unavailable"})
try:
all_messages = db.get_messages(session_id)
except Exception as e:
return json.dumps({"error": f"Failed to read messages: {e}"})
all_messages, error = _load_session_messages(session_id)
if error:
return json.dumps({"error": error})
filtered = []
for msg in all_messages:
@@ -778,14 +803,9 @@ def create_mcp_server(event_bridge: Optional[EventBridge] = None) -> "FastMCP":
if not session_id:
return json.dumps({"error": "No session ID for this conversation"})
db = _get_session_db()
if not db:
return json.dumps({"error": "Session database unavailable"})
try:
all_messages = db.get_messages(session_id)
except Exception as e:
return json.dumps({"error": f"Failed to read messages: {e}"})
all_messages, error = _load_session_messages(session_id)
if error:
return json.dumps({"error": error})
# Find the target message
target_msg = None
+59 -3
View File
@@ -364,6 +364,10 @@ class _WriteQueue:
self._db_path.parent.mkdir(parents=True, exist_ok=True)
# Thread-local connection cache — one connection per thread, reused.
self._local = threading.local()
self._connections: set[sqlite3.Connection] = set()
self._connections_lock = threading.Lock()
self._shutdown_lock = threading.Lock()
self._shutdown = False
self._init_db()
self._thread.start()
# Replay any rows left from a previous crash
@@ -374,11 +378,38 @@ class _WriteQueue:
"""Return a cached connection for the current thread."""
conn = getattr(self._local, "conn", None)
if conn is None:
conn = sqlite3.connect(str(self._db_path), timeout=30)
conn = sqlite3.connect(
str(self._db_path), timeout=30, check_same_thread=False
)
conn.row_factory = sqlite3.Row
self._local.conn = conn
with self._connections_lock:
self._connections.add(conn)
return conn
def _close_thread_conn(self) -> None:
conn = getattr(self._local, "conn", None)
if conn is None:
return
self._local.conn = None
with self._connections_lock:
self._connections.discard(conn)
try:
conn.close()
except Exception:
pass
def _close_all_connections(self) -> None:
"""Close tracked connections left by short-lived worker threads."""
with self._connections_lock:
connections = list(self._connections)
self._connections.clear()
for conn in connections:
try:
conn.close()
except Exception:
pass
def _init_db(self) -> None:
conn = self._get_conn()
conn.execute("""CREATE TABLE IF NOT EXISTS pending (
@@ -394,6 +425,9 @@ class _WriteQueue:
def enqueue(self, user_id: str, session_id: str, messages: list) -> None:
now = datetime.now(timezone.utc).isoformat()
with self._shutdown_lock:
if self._shutdown:
return
conn = self._get_conn()
cur = conn.execute(
"INSERT INTO pending (user_id, session_id, messages_json, created_at) VALUES (?,?,?,?)",
@@ -417,6 +451,7 @@ class _WriteQueue:
time.sleep(2)
def _loop(self) -> None:
try:
while True:
try:
item = self._q.get(timeout=5)
@@ -427,10 +462,24 @@ class _WriteQueue:
continue
except Exception as exc:
logger.error("RetainDB writer error: %s", exc)
finally:
# sqlite3 connections must close on their owning thread.
self._close_thread_conn()
def shutdown(self) -> None:
with self._shutdown_lock:
if self._shutdown:
return
self._shutdown = True
self._q.put(_ASYNC_SHUTDOWN)
# Caller thread owns connection opened by _init_db/_pending_rows.
self._close_thread_conn()
self._thread.join(timeout=10)
if not self._thread.is_alive():
# MemoryManager's executor may have opened a connection on a
# worker that has already exited; check_same_thread=False lets
# shutdown close that tracked handle deterministically.
self._close_all_connections()
# ---------------------------------------------------------------------------
@@ -581,6 +630,9 @@ class RetainDBMemoryProvider(MemoryProvider):
# Prevents thread accumulation if turns fire faster than prefetches complete.
for t in self._prefetch_threads:
t.join(timeout=2.0)
if any(t.is_alive() for t in self._prefetch_threads):
logger.debug("RetainDB prefetch still running; skipping new batch")
return
threads = [
threading.Thread(target=self._prefetch_context, args=(query,), name="retaindb-ctx", daemon=True),
threading.Thread(target=self._prefetch_dialectic, args=(query,), name="retaindb-dialectic", daemon=True),
@@ -795,8 +847,12 @@ class RetainDBMemoryProvider(MemoryProvider):
def shutdown(self) -> None:
for t in self._prefetch_threads:
t.join(timeout=3.0)
if self._queue:
self._queue.shutdown()
self._prefetch_threads = []
queue_obj = self._queue
self._queue = None
if queue_obj:
queue_obj.shutdown()
self._client = None
def register(ctx) -> None:
+17 -5
View File
@@ -4208,13 +4208,14 @@ class AIAgent:
)
def shutdown_memory_provider(self, messages: list = None) -> None:
"""Shut down the memory provider and context engine — call at actual session boundaries.
"""Shut down the memory provider and context engine at session end.
This calls on_session_end() then shutdown_all() on the memory
manager, and on_session_end() on the context engine.
NOT called per-turn — only at CLI exit, /reset, gateway
session expiry, etc.
Idempotent: gateway cleanup and AIAgent.close() may share this
ownership boundary.
"""
if getattr(self, "_memory_provider_shutdown", False):
return
self._memory_provider_shutdown = True
if self._memory_manager:
try:
self._memory_manager.on_session_end(messages or [])
@@ -4399,6 +4400,17 @@ class AIAgent:
Safe to call multiple times (idempotent). Each cleanup step is
independently guarded so a failure in one does not prevent the rest.
"""
# AIAgent.close() is the hard owner boundary. Gateway cleanup may
# call shutdown_memory_provider() first; its idempotence prevents
# duplicate extraction while direct callers cannot skip provider close.
try:
session_messages = getattr(self, "_session_messages", None)
self.shutdown_memory_provider(
session_messages if isinstance(session_messages, list) else None
)
except Exception:
pass
task_id = getattr(self, "session_id", None) or ""
# 1. Kill background processes for this task
+12
View File
@@ -13,6 +13,7 @@ import pytest
from agent import trace_upload
from agent.trace_upload import (
build_trace_jsonl,
load_session_messages,
upload_session_trace,
_resolve_hf_token,
_do_upload,
@@ -97,6 +98,17 @@ def test_converter_keeps_secrets_when_redact_disabled():
assert secret in jsonl
def test_load_session_messages_closes_database_on_failure(monkeypatch):
db = MagicMock()
db.resolve_session_id.side_effect = RuntimeError("read failed")
monkeypatch.setattr("hermes_state.SessionDB", lambda: db)
with pytest.raises(RuntimeError, match="read failed"):
load_session_messages("s1")
db.close.assert_called_once()
+12
View File
@@ -1,6 +1,8 @@
from unittest.mock import MagicMock, patch
from types import SimpleNamespace
from cli import HermesCLI
from hermes_cli.main import cmd_insights
class _InsightsEngineStub:
@@ -41,3 +43,13 @@ def test_cli_insights_keeps_days_flag_and_source(capsys):
assert calls == [{"days": 14, "source": "discord"}]
db.close.assert_called_once()
assert "days=14 source=discord" in capsys.readouterr().out
def test_subcommand_insights_closes_database_when_generation_fails(capsys):
db = MagicMock()
with patch("hermes_state.SessionDB", return_value=db), \
patch("agent.insights.InsightsEngine", side_effect=RuntimeError("boom")):
cmd_insights(SimpleNamespace(days=30, source=None))
db.close.assert_called_once()
assert "Error generating insights: boom" in capsys.readouterr().out
@@ -21,6 +21,7 @@ this file would have caught the regression and now pins the fix.
from __future__ import annotations
import asyncio
import threading
import time
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
@@ -254,9 +255,43 @@ class TestAPIServerDisconnectClosesResponseStore:
adapter._runner = None
adapter._app = None
adapter._response_store = store
adapter._session_dbs = {}
adapter._session_db_cache_lock = threading.Lock()
adapter._session_db_cache_closed = False
adapter.platform = Platform.API_SERVER
return adapter
@pytest.mark.asyncio
async def test_disconnect_closes_cached_session_dbs(self):
"""Disconnect must release per-profile SessionDB cache handles."""
store = MagicMock()
adapter = self._build_adapter_with_store(store)
first_db = MagicMock()
second_db = MagicMock()
adapter._session_dbs = {"default": first_db, "work": second_db}
await adapter.disconnect()
first_db.close.assert_called_once_with()
second_db.close.assert_called_once_with()
assert adapter._session_dbs == {}
@pytest.mark.asyncio
async def test_disconnect_closes_cached_session_dbs_when_runner_cleanup_fails(
self,
):
"""Runner teardown errors must not strand cached DB handles."""
adapter = self._build_adapter_with_store(MagicMock())
cached_db = MagicMock()
adapter._session_dbs = {"default": cached_db}
adapter._runner = MagicMock()
adapter._runner.cleanup = AsyncMock(side_effect=RuntimeError("boom"))
with pytest.raises(RuntimeError, match="boom"):
await adapter.disconnect()
cached_db.close.assert_called_once_with()
@pytest.mark.asyncio
async def test_disconnect_closes_response_store(self, tmp_path):
"""Closing the adapter's ResponseStore releases its SQLite connection.
+36
View File
@@ -97,6 +97,42 @@ def test_recover_inserts_via_append_message_and_deletes_file(tmp_path, monkeypat
assert not flush_file.exists()
def test_recover_closes_owned_db_when_unexpected_exception_escapes(
tmp_path, monkeypatch
):
"""Owned SessionDB must close even when recovery is interrupted."""
flush_dir = _make_flush_dir(tmp_path)
monkeypatch.setattr(
"gateway.shutdown_flush._get_flush_dir", lambda: flush_dir
)
(flush_dir / "pending.json").write_text(
json.dumps(
{
"session_key": "agent:main:telegram:123",
"data": {"text": "message", "session_id": "sid"},
}
),
encoding="utf-8",
)
class InterruptingDB:
closed = False
def append_message(self, **_kwargs):
raise KeyboardInterrupt
def close(self):
self.closed = True
db = InterruptingDB()
monkeypatch.setattr("hermes_state.SessionDB", lambda: db)
with pytest.raises(KeyboardInterrupt):
recover_pending_to_db()
assert db.closed is True
def test_serialise_object_with_text():
obj = MagicMock()
obj.text = "msg"
@@ -1,12 +1,65 @@
from __future__ import annotations
import sqlite3
from unittest.mock import MagicMock
import agent.file_safety as fs
import pytest
import plugins.memory.retaindb as retaindb
from plugins.memory.retaindb import RetainDBMemoryProvider
def test_write_queue_closes_owner_connection(tmp_path):
queue = retaindb._WriteQueue(object(), tmp_path / "retaindb.db")
owner_conn = queue._local.conn
worker = retaindb.threading.Thread(target=queue._get_conn)
worker.start()
worker.join()
queue.shutdown()
assert not queue._connections
with pytest.raises(sqlite3.ProgrammingError):
owner_conn.execute("SELECT 1")
def test_write_queue_ignores_enqueue_after_shutdown(tmp_path):
queue = retaindb._WriteQueue(object(), tmp_path / "retaindb.db")
queue.shutdown()
queue.enqueue("user", "session", [])
assert not queue._connections
def test_prefetch_does_not_spawn_when_previous_batch_is_alive(monkeypatch):
provider = RetainDBMemoryProvider()
provider._client = object()
class _RunningThread:
def join(self, timeout):
pass
def is_alive(self):
return True
previous = _RunningThread()
provider._prefetch_threads = [previous]
created = []
class _Thread:
def __init__(self, *args, **kwargs):
created.append((args, kwargs))
def start(self):
pass
monkeypatch.setattr(retaindb.threading, "Thread", _Thread)
provider.queue_prefetch("query")
assert provider._prefetch_threads == [previous]
assert not created
def test_upload_file_rejects_hermes_credential_store(tmp_path, monkeypatch):
hermes_home = tmp_path / "hermes_home"
hermes_home.mkdir()
@@ -100,6 +100,7 @@ class TestCleanupStaleAsyncClients:
try:
cleanup_stale_async_clients()
mock_client.close.assert_called_once()
with _client_cache_lock:
assert key not in _client_cache, "Stale entry should be removed"
finally:
@@ -107,6 +108,102 @@ class TestCleanupStaleAsyncClients:
with _client_cache_lock:
_client_cache.pop(key, None)
def test_awaits_async_close_for_closed_loop(self):
from agent.auxiliary_client import (
_client_cache,
_client_cache_lock,
cleanup_stale_async_clients,
)
class AsyncClient:
def __init__(self):
self._client = MagicMock()
self._client.is_closed = False
self.closed = False
async def close(self):
self.closed = True
loop = asyncio.new_event_loop()
loop.close()
client = AsyncClient()
key = ("test_async_close", True, "", "", "", (), False)
with _client_cache_lock:
_client_cache[key] = (client, "test-model", loop)
try:
cleanup_stale_async_clients()
assert client.closed
finally:
with _client_cache_lock:
_client_cache.pop(key, None)
def test_shutdown_closes_outside_cache_lock(self):
from agent.auxiliary_client import (
_client_cache,
_client_cache_lock,
shutdown_cached_clients,
)
lock_observations = []
class Client:
_client = None
def close(self):
acquired = _client_cache_lock.acquire(blocking=False)
lock_observations.append(acquired)
if acquired:
_client_cache_lock.release()
key = ("test_shutdown_lock", False, "", "", "", (), False)
with _client_cache_lock:
previous = dict(_client_cache)
_client_cache.clear()
_client_cache[key] = (Client(), "test-model", None)
try:
shutdown_cached_clients()
finally:
with _client_cache_lock:
_client_cache.clear()
_client_cache.update(previous)
assert lock_observations == [True]
def test_shutdown_does_not_await_live_foreign_loop_client(self):
from agent.auxiliary_client import (
_client_cache,
_client_cache_lock,
shutdown_cached_clients,
)
owner_loop = asyncio.new_event_loop()
class Client:
def __init__(self):
self.awaited = False
async def close(self):
self.awaited = True
client = Client()
key = ("test_shutdown_foreign_loop", True, "", "", "", (), False)
with _client_cache_lock:
previous = dict(_client_cache)
_client_cache.clear()
_client_cache[key] = (client, "test-model", owner_loop)
try:
shutdown_cached_clients()
assert client.awaited is False
finally:
owner_loop.close()
with _client_cache_lock:
_client_cache.clear()
_client_cache.update(previous)
def test_keeps_live_entries(self):
"""Entries with an open loop should be preserved."""
from agent.auxiliary_client import (
@@ -25,6 +25,24 @@ class RecordingMemoryProvider:
pass
def test_shutdown_memory_provider_is_idempotent():
from unittest.mock import MagicMock
from run_agent import AIAgent
manager = MagicMock()
agent = object.__new__(AIAgent)
agent._memory_manager = manager
agent.context_compressor = None
agent.session_id = "session-1"
agent.shutdown_memory_provider([{"role": "user", "content": "one"}])
agent.shutdown_memory_provider([{"role": "user", "content": "two"}])
manager.on_session_end.assert_called_once()
manager.shutdown_all.assert_called_once()
def test_blank_memory_provider_does_not_auto_enable_honcho():
"""Blank memory.provider should remain opt-out even if Honcho fallback looks configured."""
cfg = {"memory": {"provider": ""}, "agent": {}}
@@ -59,6 +77,22 @@ def test_blank_memory_provider_does_not_auto_enable_honcho():
save_config.assert_not_called()
def test_close_shuts_down_memory_provider():
from unittest.mock import MagicMock
from run_agent import AIAgent
agent = object.__new__(AIAgent)
agent._memory_manager = MagicMock()
agent.context_compressor = None
agent.session_id = ""
agent._session_messages = []
agent.close()
agent._memory_manager.shutdown_all.assert_called_once()
def test_aiagent_forwards_user_id_alt_to_memory_provider():
provider = RecordingMemoryProvider()
cfg = {"memory": {"provider": "recording"}, "agent": {}}
+109
View File
@@ -3,6 +3,7 @@
import sqlite3
import time
import json
import threading
from pathlib import Path
from unittest import mock
@@ -103,6 +104,114 @@ def _no_fts_rebuild_throttle(monkeypatch):
class TestConnectionLifecycle:
def test_failed_writable_open_does_not_leak_tracked_connection(
self, tmp_path, monkeypatch
):
"""A failed schema init must close the connection opened before it."""
from hermes_cli.sqlite_safe_read import has_live_connection
db_path = tmp_path / "state.db"
opened = []
real_connect = hermes_state._connect_tracked_db
def capture_connect(*args, **kwargs):
conn = real_connect(*args, **kwargs)
opened.append(conn)
return conn
monkeypatch.setattr(hermes_state, "_connect_tracked_db", capture_connect)
monkeypatch.setattr(
SessionDB,
"_init_schema",
mock.Mock(side_effect=RuntimeError("schema init failed")),
)
try:
with pytest.raises(RuntimeError, match="schema init failed"):
SessionDB(db_path=db_path)
assert has_live_connection(db_path) is False
finally:
for conn in opened:
try:
conn.close()
except Exception:
pass
def test_failed_wal_read_open_does_not_leak_tracked_connection(
self, tmp_path, monkeypatch
):
"""A post-open read setup failure must close its unregistered conn."""
from hermes_cli import sqlite_safe_read
db_path = tmp_path / "state.db"
db = SessionDB(db_path=db_path)
opened = []
real_connect = hermes_state._connect_tracked_db
real_pragmas = hermes_state.apply_database_pragmas
def capture_connect(*args, **kwargs):
conn = real_connect(*args, **kwargs)
opened.append(conn)
return conn
def fail_pragmas(*args, **kwargs):
raise RuntimeError("read setup failed")
monkeypatch.setattr(hermes_state, "_connect_tracked_db", capture_connect)
monkeypatch.setattr(hermes_state, "apply_database_pragmas", fail_pragmas)
before = dict(sqlite_safe_read._live_connections)
db._wal_active = True
try:
with pytest.raises(RuntimeError, match="read setup failed"):
db._get_read_conn()
assert sqlite_safe_read._live_connections == before
finally:
monkeypatch.setattr(
hermes_state, "apply_database_pragmas", real_pragmas
)
for conn in opened:
try:
conn.close()
except Exception:
pass
db.close()
def test_close_closes_wal_read_connection_created_on_worker_thread(
self, tmp_path
):
"""SessionDB.close() must drain read conns created by other threads."""
from hermes_cli.sqlite_safe_read import has_live_connection
db_path = tmp_path / "state.db"
db = SessionDB(db_path=db_path)
db._wal_active = True
opened = threading.Event()
release = threading.Event()
errors = []
def open_read_connection():
try:
assert db._get_read_conn() is not None
opened.set()
release.wait(timeout=10)
except BaseException as exc:
errors.append(exc)
opened.set()
worker = threading.Thread(target=open_read_connection)
worker.start()
assert opened.wait(timeout=10)
assert not errors
db.close()
assert has_live_connection(db_path) is False
release.set()
worker.join(timeout=10)
assert not worker.is_alive()
assert not errors
def test_read_only_close_never_requests_wal_checkpoint(self, tmp_path):
db_path = tmp_path / "state.db"
writable = SessionDB(db_path=db_path)
+13
View File
@@ -272,6 +272,19 @@ class TestImports:
class TestHelpers:
def test_load_session_messages_closes_database_on_error(self, monkeypatch):
import mcp_serve
db = MagicMock()
db.get_messages.side_effect = RuntimeError("read failed")
monkeypatch.setattr(mcp_serve, "_get_session_db", lambda: db)
messages, error = mcp_serve._load_session_messages("s1")
assert messages is None
assert "read failed" in error
db.close.assert_called_once()
def test_get_sessions_dir(self, tmp_path):
from mcp_serve import _get_sessions_dir
result = _get_sessions_dir()
+89
View File
@@ -10348,6 +10348,95 @@ def test_prompt_submit_merges_on_model_switch_marker(monkeypatch):
server._sessions.pop("sid", None)
def test_prompt_submit_merges_on_personality_pivot_marker(monkeypatch):
"""A personality pivot injected mid-turn must merge like a model switch.
`/personality` applies immediately — there is no deferred queue for it the
way `pending_model_switch` defers a mid-turn model change — so choosing a
personality while a turn is running bumps `history_version` from the RPC
thread. The mid-turn reconciliation only recognized the model-switch
marker, so the pivot read as a genuine desync and the finished turn was
dropped from session history: the user saw the reply and it was never
stored (#82756).
"""
session_ref: dict[str, dict | None] = {"s": None}
class _PivotAgent:
def run_conversation(
self, prompt, conversation_history=None, stream_callback=None, **_kwargs
):
# Real injection point, mid-turn, exactly as the personality RPC
# would reach it from the other thread.
server._apply_personality_to_session(
"sid", session_ref["s"], "Answer tersely.", "terse"
)
return {
"final_response": "agent reply",
"messages": list(conversation_history)
+ [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "agent reply"},
],
}
class _ImmediateThread:
def __init__(self, target=None, daemon=None):
self._target = target
def start(self):
self._target()
server._sessions["sid"] = _session(
agent=_PivotAgent(),
history=[{"role": "user", "content": "hello"}],
)
session_ref["s"] = server._sessions["sid"]
emits: list[tuple] = []
try:
monkeypatch.setattr(server.threading, "Thread", _ImmediateThread)
monkeypatch.setattr(server, "_get_usage", lambda _a: {})
monkeypatch.setattr(server, "render_message", lambda _t, _c: "")
monkeypatch.setattr(server, "_session_info", lambda *a, **k: {})
monkeypatch.setattr(server, "_emit", lambda *a: emits.append(a))
resp = server.handle_request(
{
"id": "1",
"method": "prompt.submit",
"params": {"session_id": "sid", "text": "hi"},
}
)
assert resp.get("result"), f"got error: {resp.get('error')}"
final_history = server._sessions["sid"]["history"]
assistant_msgs = [
e
for e in final_history
if isinstance(e, dict)
and e.get("role") == "assistant"
and e.get("content") == "agent reply"
]
assert len(assistant_msgs) == 1, (
"the personality pivot discarded the finished turn instead of "
f"merging it (got {len(assistant_msgs)} assistant replies)"
)
pivots = [
e
for e in final_history
if isinstance(e, dict) and e.get("display_kind") == "personality_switch"
]
assert len(pivots) == 1, f"expected exactly 1 pivot, got {len(pivots)}"
complete_calls = [a for a in emits if a[0] == "message.complete"]
assert len(complete_calls) == 1
_, _, payload = complete_calls[0]
assert "warning" not in payload, "merge path should not surface a warning"
finally:
server._sessions.pop("sid", None)
def test_prompt_submit_sanitizes_bracketed_paste_before_agent(monkeypatch):
"""prompt.submit must sanitize corrupted user text before run_conversation."""
captured: dict[str, str] = {}
+22
View File
@@ -0,0 +1,22 @@
"""Ownership tests for desktop message reactions."""
from unittest.mock import MagicMock
from tools import react_to_message_tool as reactions
def test_reaction_database_closes_when_write_fails(monkeypatch):
db = MagicMock()
db.latest_message_row_id.return_value = 42
db.set_message_reaction.side_effect = RuntimeError("write failed")
monkeypatch.setattr(reactions, "_open_session_db", lambda: db)
monkeypatch.setattr(
reactions,
"get_session_env",
lambda _name, _default="": "session-1",
)
result = reactions.react_to_message_tool("👍")
assert "write failed" in result
db.close.assert_called_once()
+44
View File
@@ -94,6 +94,50 @@ class TestFormatTimestamp:
# =========================================================================
class TestBrowseShape:
def test_lazy_database_is_closed_after_search(self, monkeypatch):
class _DB:
closed = 0
def list_sessions_rich(self, **_kwargs):
return []
def close(self):
self.closed += 1
db = _DB()
monkeypatch.setattr("hermes_state.SessionDB", lambda: db)
result = json.loads(session_search())
assert result["success"] is True
assert db.closed == 1
def test_cross_profile_database_is_closed_but_shared_database_is_not(
self, monkeypatch
):
class _DB:
def __init__(self):
self.closed = 0
def list_sessions_rich(self, **_kwargs):
return []
def close(self):
self.closed += 1
shared_db = _DB()
profile_db = _DB()
monkeypatch.setattr(
"tools.session_search_tool._resolve_profile_db",
lambda _profile: profile_db,
)
result = json.loads(session_search(db=shared_db, profile="work"))
assert result["success"] is True
assert profile_db.closed == 1
assert shared_db.closed == 0
def test_no_args_returns_recent_sessions(self, db):
_seed_modpack_sessions(db)
result = json.loads(session_search(db=db))
+37 -10
View File
@@ -31,20 +31,18 @@ def _open_session_db():
return None
def react_to_message_tool(emoji: str, message_row_id=None, messages_back=None) -> str:
def _react_to_message_with_db(
emoji: str,
message_row_id=None,
messages_back=None,
*,
db,
session_key: str,
) -> str:
"""Attach (or with an empty ``emoji`` retract) the agent's reaction."""
emoji = (emoji or "").strip()
session_key = get_session_env("HERMES_SESSION_KEY", "") or get_session_env(
"HERMES_SESSION_ID", ""
)
if not session_key:
return tool_error("No active session — reactions need a persisted conversation.")
db = _open_session_db()
if db is None:
return tool_error("Session storage is unavailable.")
row_id = message_row_id
target_role = "user"
if row_id is None:
@@ -89,6 +87,35 @@ def react_to_message_tool(emoji: str, message_row_id=None, messages_back=None) -
)
def react_to_message_tool(emoji: str, message_row_id=None, messages_back=None) -> str:
"""Attach (or with an empty ``emoji`` retract) the agent's reaction."""
emoji = (emoji or "").strip()
session_key = get_session_env("HERMES_SESSION_KEY", "") or get_session_env(
"HERMES_SESSION_ID", ""
)
if not session_key:
return tool_error("No active session — reactions need a persisted conversation.")
db = _open_session_db()
if db is None:
return tool_error("Session storage is unavailable.")
try:
return _react_to_message_with_db(
emoji,
message_row_id,
messages_back,
db=db,
session_key=session_key,
)
finally:
try:
db.close()
except Exception:
pass
def check_react_requirements() -> bool:
"""Opt-in feature flag — surface eligibility is the toolset's job.
+56 -10
View File
@@ -916,7 +916,7 @@ def _discover(
return json.dumps(_final_payload, ensure_ascii=False)
def session_search(
def _session_search_impl(
query: str = "",
role_filter: str = None,
limit: int = 3,
@@ -930,6 +930,8 @@ def session_search(
sort: str = None,
# Cross-profile (any shape)
profile: str = None,
*,
_owned_dbs: Optional[List[Any]] = None,
) -> str:
"""Single-shape tool. Mode inferred from which args are set.
@@ -942,15 +944,6 @@ def session_search(
``@session:<profile>/<id>`` link). Scroll wins over read/discovery when an
anchor is set — the agent has asked for a specific slice.
"""
if db is None:
try:
from hermes_state import SessionDB
db = SessionDB()
except Exception:
logging.debug("SessionDB unavailable for session_search", exc_info=True)
from hermes_state import format_session_db_unavailable
return tool_error(format_session_db_unavailable(), success=False)
# Normalise a raw `@session:<profile>/<id>` link value passed as session_id.
# Session ids never contain "/", so a slash unambiguously means profile/id —
# always strip the prefix off the id, and adopt the embedded profile only
@@ -973,6 +966,8 @@ def session_search(
return tool_error(f"profile '{profile}': {e}", success=False)
if profile_db is not None:
db = profile_db
if _owned_dbs is not None:
_owned_dbs.append(profile_db)
current_session_id = None
# Scroll shape takes precedence — explicit anchor beats any query.
@@ -1041,6 +1036,57 @@ def session_search(
)
def session_search(
query: str = "",
role_filter: str = None,
limit: int = 3,
db=None,
current_session_id: str = None,
# Scroll shape
session_id: str = None,
around_message_id: int = None,
window: int = 5,
# Discovery shape
sort: str = None,
# Cross-profile (any shape)
profile: str = None,
) -> str:
"""Run session search and close databases opened by this invocation."""
owned_dbs: List[Any] = []
if db is None:
try:
from hermes_state import SessionDB
db = SessionDB()
owned_dbs.append(db)
except Exception:
logging.debug("SessionDB unavailable for session_search", exc_info=True)
from hermes_state import format_session_db_unavailable
return tool_error(format_session_db_unavailable(), success=False)
try:
return _session_search_impl(
query=query,
role_filter=role_filter,
limit=limit,
db=db,
current_session_id=current_session_id,
session_id=session_id,
around_message_id=around_message_id,
window=window,
sort=sort,
profile=profile,
_owned_dbs=owned_dbs,
)
finally:
for owned_db in reversed(owned_dbs):
try:
owned_db.close()
except Exception:
logging.debug("Failed to close session_search SessionDB", exc_info=True)
def check_session_search_requirements() -> bool:
"""Requires the SQLite state database."""
try:
+28 -9
View File
@@ -4128,6 +4128,20 @@ def _is_model_switch_marker(entry: Any) -> bool:
return isinstance(content, str) and content.startswith(_MODEL_SWITCH_MARKER_PREFIX)
def _is_pivot_marker(entry: Any) -> bool:
"""Whether a history entry is a marker the gateway splices in mid-turn.
Model switches and personality changes both inject a ``role=user`` pivot
into the live history from the RPC thread while a turn may be running, so
either one can be the sole reason turn-start and current history differ.
Only the model-switch marker is self-replacing, which is why the dedup in
:func:`_append_model_switch_marker` stays narrower than this.
"""
if _is_model_switch_marker(entry):
return True
return isinstance(entry, dict) and entry.get("display_kind") == "personality_switch"
def _append_model_switch_marker(session: dict | None, *, model: str, provider: str) -> None:
"""Record a real system-history pivot after a live model switch.
@@ -10632,10 +10646,15 @@ def _run_prompt_submit(
session["history_version"] = history_version + 1
else:
# History mutated externally during the turn.
# Check if the only mutation was a model-switch
# marker inserted mid-turn (#76870). If so the
# agent output is still valid — merge it into the
# current history that now contains the marker.
# Check if the only mutation was a pivot marker
# the gateway itself inserted mid-turn (#76870).
# If so the agent output is still valid — merge it
# into the current history that now contains the
# marker. A personality change counts here too:
# unlike a model switch it has no pending queue, so
# `/personality` during a running turn lands
# immediately and used to read as a genuine desync,
# dropping the finished turn (#82756).
#
# _append_model_switch_marker strips prior markers
# in-place then appends a new one, so the delta
@@ -10643,19 +10662,19 @@ def _run_prompt_submit(
# content, not indices.
current_history = list(session["history"])
history_no_markers = [
e for e in history if not _is_model_switch_marker(e)
e for e in history if not _is_pivot_marker(e)
]
current_no_markers = [
e for e in current_history if not _is_model_switch_marker(e)
e for e in current_history if not _is_pivot_marker(e)
]
model_switch_only = (
pivot_only = (
current_no_markers == history_no_markers
and any(
_is_model_switch_marker(e)
_is_pivot_marker(e)
for e in current_history
)
)
if model_switch_only:
if pivot_only:
# The agent's new messages start after the
# turn-start history. Guard against
# auto-compression making result["messages"]