Merge branch 'simp/r2-state-b' into simp/r2-state

# Conflicts:
#	hermes_state_messages.py
This commit is contained in:
Teknium
2026-09-02 17:35:13 -07:00
4 changed files with 396 additions and 750 deletions

View File

@@ -9,7 +9,7 @@ import json
import logging import logging
import sqlite3 import sqlite3
import time import time
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional, Tuple
from hermes_state_common import _sql_session_last_active, is_automatic_end_reason from hermes_state_common import _sql_session_last_active, is_automatic_end_reason
@@ -27,6 +27,30 @@ def _ended_by_compression(row) -> bool:
return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression" return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression"
def _cooldown_row(exists: bool, cooldown_until, error) -> Dict[str, Any]:
return {
"session_exists": exists,
"cooldown_until": float(cooldown_until) if cooldown_until is not None else None,
"error": error}
def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now: float, expires_at: float,
stale) -> Tuple[bool, Optional[str]]:
"""Single-transaction lease claim: DELETE a stale holder's row (``stale(holder,
expires_at)``), INSERT OR IGNORE ours, then SELECT to confirm ownership (INSERT OR
IGNORE gives no rowcount signal). Returns ``(acquired, reclaimed_holder)``."""
reclaimed_holder = None
row = conn.execute(f"SELECT holder, expires_at FROM {table} WHERE {key_col} = ?", (key,)).fetchone()
if row is not None and stale(row["holder"], row["expires_at"]):
conn.execute(f"DELETE FROM {table} WHERE {key_col} = ? AND holder = ?", (key, row["holder"]))
reclaimed_holder = row["holder"]
conn.execute(
f"INSERT OR IGNORE INTO {table} ({key_col}, holder, acquired_at, expires_at) VALUES (?, ?, ?, ?)",
(key, holder, now, expires_at))
owner = conn.execute(f"SELECT holder FROM {table} WHERE {key_col} = ?", (key,)).fetchone()
return owner is not None and owner["holder"] == holder, reclaimed_holder
class SessionCompressionMixin: class SessionCompressionMixin:
"""Compression lineage, cooldown/streak counters, locks and turn leases.""" """Compression lineage, cooldown/streak counters, locks and turn leases."""
@@ -96,16 +120,13 @@ class SessionCompressionMixin:
deleted = conn.execute( deleted = conn.execute(
"DELETE FROM compression_locks " "DELETE FROM compression_locks "
"WHERE session_id = ? AND holder = ? AND expires_at = ?", "WHERE session_id = ? AND holder = ? AND expires_at = ?",
(session_id, lock_row["holder"], expires_at), (session_id, lock_row["holder"], expires_at))
)
if deleted.rowcount != 1: if deleted.rowcount != 1:
return False return False
updated = conn.execute( updated = conn.execute(
"UPDATE sessions SET ended_at = NULL, end_reason = NULL " "UPDATE sessions SET ended_at = NULL, end_reason = NULL "
"WHERE id = ? AND ended_at IS NOT NULL " "WHERE id = ? AND ended_at IS NOT NULL AND end_reason = 'compression'",
"AND end_reason = 'compression'", (session_id,))
(session_id,),
)
# rowcount==1 is guaranteed by the parent SELECT in this same txn. A False # rowcount==1 is guaranteed by the parent SELECT in this same txn. A False
# return added past this point must raise instead: the lease DELETE above # return added past this point must raise instead: the lease DELETE above
# commits unless _do raises. # commits unless _do raises.
@@ -115,7 +136,10 @@ class SessionCompressionMixin:
def _publish_child_session_row(self, conn, parent, *, parent_session_id, child_session_id, source, def _publish_child_session_row(self, conn, parent, *, parent_session_id, child_session_id, source,
model, model_config, system_prompt, cwd, profile_name) -> None: model, model_config, system_prompt, cwd, profile_name) -> None:
"""INSERT the compression child's ``sessions`` row copied from *parent*.""" """INSERT the compression child's ``sessions`` row copied from *parent*. Same contract
as _insert_session_row's compression-fork backfill: the child stays on the parent's
profile and keeps gateway routing/origin columns; no owner on either side -> this
store's profile."""
system_prompt_hash = self._store_system_prompt(conn, system_prompt) system_prompt_hash = self._store_system_prompt(conn, system_prompt)
conn.execute( conn.execute(
"""INSERT INTO sessions ( """INSERT INTO sessions (
@@ -126,28 +150,12 @@ class SessionCompressionMixin:
thread_id, display_name, origin_json, started_at thread_id, display_name, origin_json, started_at
) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", ) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
( (
child_session_id, child_session_id, source, model, json.dumps(model_config) if model_config else None,
source, system_prompt_hash, parent_session_id, cwd or parent["cwd"], parent["git_branch"],
model,
json.dumps(model_config) if model_config else None,
system_prompt_hash,
parent_session_id,
cwd or parent["cwd"],
parent["git_branch"],
parent["git_repo_root"], parent["git_repo_root"],
# Same contract as _insert_session_row's compression-fork backfill: the
# child stays on the parent's profile and keeps gateway routing/origin
# columns; no owner on either side -> this store's profile.
profile_name or parent["profile_name"] or self._own_profile_name(), profile_name or parent["profile_name"] or self._own_profile_name(),
parent["user_id"], parent["user_id"], parent["session_key"], parent["chat_id"], parent["chat_type"],
parent["session_key"], parent["thread_id"], parent["display_name"], parent["origin_json"], time.time()),
parent["chat_id"],
parent["chat_type"],
parent["thread_id"],
parent["display_name"],
parent["origin_json"],
time.time(),
),
) )
def publish_compression_child( def publish_compression_child(
@@ -180,8 +188,7 @@ class SessionCompressionMixin:
conn.execute( conn.execute(
"UPDATE compression_locks SET expires_at = ? " "UPDATE compression_locks SET expires_at = ? "
"WHERE session_id = ? AND holder = ?", "WHERE session_id = ? AND holder = ?",
(time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder), (time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder))
)
lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone() lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone()
if require_compression_lease and ( if require_compression_lease and (
lock_row is None lock_row is None
@@ -209,10 +216,8 @@ class SessionCompressionMixin:
# Deliberate boundaries still fail closed. # Deliberate boundaries still fail closed.
if is_automatic_end_reason(parent["end_reason"]): if is_automatic_end_reason(parent["end_reason"]):
conn.execute( conn.execute(
"UPDATE sessions SET ended_at = NULL, end_reason = NULL " "UPDATE sessions SET ended_at = NULL, end_reason = NULL WHERE id = ?",
"WHERE id = ?", (parent_session_id,))
(parent_session_id,),
)
else: else:
raise RuntimeError(f"Compression parent already ended: {parent_session_id}") raise RuntimeError(f"Compression parent already ended: {parent_session_id}")
if not messages: if not messages:
@@ -220,23 +225,17 @@ class SessionCompressionMixin:
self._publish_child_session_row( self._publish_child_session_row(
conn, parent, parent_session_id=parent_session_id, child_session_id=child_session_id, conn, parent, parent_session_id=parent_session_id, child_session_id=child_session_id,
source=source, model=model, model_config=model_config, system_prompt=system_prompt, source=source, model=model, model_config=model_config, system_prompt=system_prompt,
cwd=cwd, profile_name=profile_name, cwd=cwd, profile_name=profile_name)
)
total_messages, total_tool_calls = self._insert_message_rows(conn, child_session_id, messages) total_messages, total_tool_calls = self._insert_message_rows(conn, child_session_id, messages)
if watermark is not None: if watermark is not None:
# Clone the parent's concurrent tail into the child after the handoff; # Clone the parent's concurrent tail into the child after the handoff;
# originals stay in the closed parent for lineage recovery. # originals stay in the closed parent for lineage recovery.
_ceiling_clause = "" bounded = watermark_ceiling is not None
_params: list = [parent_session_id, int(watermark)]
if watermark_ceiling is not None:
_ceiling_clause = " AND id <= ?"
_params.append(int(watermark_ceiling))
tail_ids, tail_tool_calls = self._tail_rows_after_watermark( tail_ids, tail_tool_calls = self._tail_rows_after_watermark(
conn, conn, "SELECT id, tool_calls FROM messages "
"SELECT id, tool_calls FROM messages "
"WHERE session_id = ? AND active = 1 AND id > ?" "WHERE session_id = ? AND active = 1 AND id > ?"
f"{_ceiling_clause} ORDER BY id", f"{' AND id <= ?' if bounded else ''} ORDER BY id",
_params, [parent_session_id, int(watermark), *([int(watermark_ceiling)] if bounded else [])],
) )
if tail_ids: if tail_ids:
self._clone_message_rows(conn, tail_ids, session_id=child_session_id) self._clone_message_rows(conn, tail_ids, session_id=child_session_id)
@@ -244,13 +243,10 @@ class SessionCompressionMixin:
total_tool_calls += tail_tool_calls total_tool_calls += tail_tool_calls
conn.execute( conn.execute(
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
(total_messages, total_tool_calls, child_session_id), (total_messages, total_tool_calls, child_session_id))
)
updated = conn.execute( updated = conn.execute(
"UPDATE sessions SET ended_at = ?, end_reason = 'compression' " "UPDATE sessions SET ended_at = ?, end_reason = 'compression' "
"WHERE id = ? AND ended_at IS NULL", "WHERE id = ? AND ended_at IS NULL", (time.time(), parent_session_id))
(time.time(), parent_session_id),
)
if updated.rowcount != 1: if updated.rowcount != 1:
raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}") raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}")
@@ -276,8 +272,7 @@ class SessionCompressionMixin:
" AND compression_failure_cooldown_until > ? " " AND compression_failure_cooldown_until > ? "
"THEN compression_failure_cooldown_until ELSE ? END, " "THEN compression_failure_cooldown_until ELSE ? END, "
"compression_failure_error = ? WHERE id = ?", "compression_failure_error = ? WHERE id = ?",
(cooldown_until, cooldown_until, error, session_id), (cooldown_until, cooldown_until, error, session_id))
)
def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]: def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]:
"""Return the active (unexpired) compression-failure cooldown, or None.""" """Return the active (unexpired) compression-failure cooldown, or None."""
@@ -285,24 +280,17 @@ class SessionCompressionMixin:
return None return None
now = time.time() now = time.time()
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,))
if row is None or row[0] is None: if row is None or row[0] is None or float(row[0]) <= now:
return None return None
cooldown_until = float(row[0]) return {"cooldown_until": float(row[0]), "remaining_seconds": float(row[0]) - now, "error": row[1]}
if cooldown_until <= now:
return None
return {"cooldown_until": cooldown_until, "remaining_seconds": cooldown_until - now, "error": row[1]}
def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]: def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]:
"""Exact stored cooldown columns, no expiry filtering, so compression """Exact stored cooldown columns, no expiry filtering, so compression
cancellation can roll back an expired, partially-null, or absent row exactly.""" cancellation can roll back an expired, partially-null, or absent row exactly."""
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) if session_id else None row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) if session_id else None
if row is None: if row is None:
return {"session_exists": False, "cooldown_until": None, "error": None} return _cooldown_row(False, None, None)
return { return _cooldown_row(True, row[0], row[1])
"session_exists": True,
"cooldown_until": float(row[0]) if row[0] is not None else None,
"error": row[1],
}
def restore_compression_failure_cooldown_row(self, session_id: str, snapshot: Dict[str, Any]) -> None: def restore_compression_failure_cooldown_row(self, session_id: str, snapshot: Dict[str, Any]) -> None:
"""Restore and verify an exact cooldown-row snapshot. Unlike record/clear this """Restore and verify an exact cooldown-row snapshot. Unlike record/clear this
@@ -318,19 +306,12 @@ class SessionCompressionMixin:
def _do(conn): def _do(conn):
cursor = conn.execute( cursor = conn.execute(
"UPDATE sessions SET compression_failure_cooldown_until = ?, " "UPDATE sessions SET compression_failure_cooldown_until = ?, "
"compression_failure_error = ? WHERE id = ?", "compression_failure_error = ? WHERE id = ?", (deadline, error, session_id))
(deadline, error, session_id),
)
if cursor.rowcount != 1: if cursor.rowcount != 1:
raise RuntimeError(f"compression cooldown rollback session missing: {session_id}") raise RuntimeError(f"compression cooldown rollback session missing: {session_id}")
self._execute_write(_do) self._execute_write(_do)
actual = self.get_compression_failure_cooldown_row(session_id) actual = self.get_compression_failure_cooldown_row(session_id)
expected = { expected = _cooldown_row(True, deadline, error)
"session_exists": True,
"cooldown_until": float(deadline) if deadline is not None else None,
"error": error,
}
if actual != expected: if actual != expected:
raise RuntimeError( raise RuntimeError(
f"compression cooldown rollback verification failed: " f"compression cooldown rollback verification failed: "
@@ -344,9 +325,7 @@ class SessionCompressionMixin:
self._write_sql_logged( self._write_sql_logged(
"clear_compression_failure_cooldown", session_id, "clear_compression_failure_cooldown", session_id,
"UPDATE sessions SET compression_failure_cooldown_until = NULL, " "UPDATE sessions SET compression_failure_cooldown_until = NULL, "
"compression_failure_error = NULL WHERE id = ?", "compression_failure_error = NULL WHERE id = ?", (session_id,))
(session_id,),
)
def _read_session_number(self, column: str, session_id: str, cast: type, zero: Any) -> Any: def _read_session_number(self, column: str, session_id: str, cast: type, zero: Any) -> Any:
"""Read one numeric ``sessions`` column clamped at ``zero``; a missing session, """Read one numeric ``sessions`` column clamped at ``zero``; a missing session,
@@ -370,8 +349,7 @@ class SessionCompressionMixin:
if session_id: if session_id:
self._write_sql( self._write_sql(
"UPDATE sessions SET compression_fallback_streak = ? WHERE id = ?", "UPDATE sessions SET compression_fallback_streak = ? WHERE id = ?",
(max(0, int(streak)), session_id), (max(0, int(streak)), session_id))
)
def get_compression_ineffective_count(self, session_id: str) -> int: def get_compression_ineffective_count(self, session_id: str) -> int:
"""Persisted ineffective-compaction strike count — the durable half of the """Persisted ineffective-compaction strike count — the durable half of the
@@ -384,8 +362,7 @@ class SessionCompressionMixin:
if session_id: if session_id:
self._write_sql( self._write_sql(
"UPDATE sessions SET compression_ineffective_count = ? WHERE id = ?", "UPDATE sessions SET compression_ineffective_count = ? WHERE id = ?",
(max(0, int(count)), session_id), (max(0, int(count)), session_id))
)
def get_compression_recovery_deadline(self, session_id: str) -> float: def get_compression_recovery_deadline(self, session_id: str) -> float:
"""Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable """Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable
@@ -402,8 +379,7 @@ class SessionCompressionMixin:
normalized = 0.0 normalized = 0.0
self._write_sql( self._write_sql(
"UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?", "UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?",
(normalized if normalized > 0.0 else None, session_id), (normalized or None, session_id))
)
def refresh_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool: def refresh_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool:
"""Extend the compression lock lease if ``holder`` still owns it. """Extend the compression lock lease if ``holder`` still owns it.
@@ -419,8 +395,7 @@ class SessionCompressionMixin:
expires_at = time.time() + ttl_seconds expires_at = time.time() + ttl_seconds
try: try:
return self._write_rowcount( return self._write_rowcount(
"UPDATE compression_locks SET expires_at = ? " "UPDATE compression_locks SET expires_at = ? WHERE session_id = ? AND holder = ?",
"WHERE session_id = ? AND holder = ?",
(expires_at, session_id, holder), (expires_at, session_id, holder),
) > 0 ) > 0
except sqlite3.Error as exc: except sqlite3.Error as exc:
@@ -442,32 +417,14 @@ class SessionCompressionMixin:
expires_at = now + ttl_seconds expires_at = now + ttl_seconds
def _do(conn): def _do(conn):
reclaimed_holder = None return _claim_lease_row(
row = conn.execute(_LOCK_ROW_SQL, (session_id,)).fetchone() conn, "compression_locks", "session_id", session_id, holder, now, expires_at,
if row is not None: lambda h, e: e < now or _compression_lock_holder_process_is_dead(h))
current_holder, current_expires_at = row[0], row[1]
if current_expires_at < now or _compression_lock_holder_process_is_dead(current_holder):
conn.execute(
"DELETE FROM compression_locks "
"WHERE session_id = ? AND holder = ?",
(session_id, current_holder),
)
reclaimed_holder = current_holder
conn.execute(
"INSERT OR IGNORE INTO compression_locks "
"(session_id, holder, acquired_at, expires_at) "
"VALUES (?, ?, ?, ?)",
(session_id, holder, now, expires_at),
)
row = conn.execute("SELECT holder FROM compression_locks WHERE session_id = ?", (session_id,)).fetchone()
return row is not None and row[0] == holder, reclaimed_holder
try: try:
acquired, reclaimed_holder = self._execute_write(_do) acquired, reclaimed_holder = self._execute_write(_do)
if reclaimed_holder: if reclaimed_holder:
logger.warning( logger.warning("Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder)
"Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder,
)
return bool(acquired) return bool(acquired)
except sqlite3.Error as exc: except sqlite3.Error as exc:
# False makes the caller skip compression — safe when the lock subsystem is broken. # False makes the caller skip compression — safe when the lock subsystem is broken.
@@ -480,10 +437,8 @@ class SessionCompressionMixin:
return return
self._write_sql_logged( self._write_sql_logged(
"release_compression_lock", session_id, "release_compression_lock", session_id,
"DELETE FROM compression_locks " "DELETE FROM compression_locks WHERE session_id = ? AND holder = ?",
"WHERE session_id = ? AND holder = ?", (session_id, holder))
(session_id, holder),
)
def _session_turn_lease_key_on_conn(self, conn, session_id: str) -> str: def _session_turn_lease_key_on_conn(self, conn, session_id: str) -> str:
"""Walk compression parents on ``conn`` to the conversation lease key. """Walk compression parents on ``conn`` to the conversation lease key.
@@ -496,9 +451,7 @@ class SessionCompressionMixin:
def _row(sid: str): def _row(sid: str):
row = conn.execute( row = conn.execute(
"SELECT id, parent_session_id, source, model_config, end_reason " "SELECT id, parent_session_id, source, model_config, end_reason FROM sessions WHERE id = ?", (sid,),
"FROM sessions WHERE id = ?",
(sid,),
).fetchone() ).fetchone()
return dict(row) if row else None return dict(row) if row else None
@@ -537,29 +490,10 @@ class SessionCompressionMixin:
def _do(conn): def _do(conn):
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
row = conn.execute( return _claim_lease_row(
"SELECT holder, expires_at FROM session_turn_leases " conn, "session_turn_leases", "conversation_id", conversation_id, holder, now, expires_at,
"WHERE conversation_id = ?", lambda h, e: float(e) <= now or _compression_lock_holder_process_is_dead(h),
(conversation_id,), )[0]
).fetchone()
if row is not None:
current_holder = row["holder"]
if float(row["expires_at"]) <= now or _compression_lock_holder_process_is_dead(current_holder):
conn.execute(
"DELETE FROM session_turn_leases "
"WHERE conversation_id = ? AND holder = ?",
(conversation_id, current_holder),
)
conn.execute(
"INSERT OR IGNORE INTO session_turn_leases "
"(conversation_id, holder, acquired_at, expires_at) "
"VALUES (?, ?, ?, ?)",
(conversation_id, holder, now, expires_at),
)
owner = conn.execute(
"SELECT holder FROM session_turn_leases WHERE conversation_id = ?", (conversation_id,),
).fetchone()
return owner is not None and owner["holder"] == holder
return bool(self._execute_write(_do, patience_s=patience_s)) return bool(self._execute_write(_do, patience_s=patience_s))
@@ -588,8 +522,7 @@ class SessionCompressionMixin:
logger.debug("session turn lease should_abort callback failed", exc_info=True) logger.debug("session turn lease should_abort callback failed", exc_info=True)
try: try:
if self.try_acquire_session_turn_lease( if self.try_acquire_session_turn_lease(
session_id, holder, ttl_seconds=ttl_seconds, patience_s=acquire_patience_s, session_id, holder, ttl_seconds=ttl_seconds, patience_s=acquire_patience_s):
):
return True return True
except sqlite3.Error as exc: except sqlite3.Error as exc:
# Long holder transactions can exhaust one write-patience budget; keep # Long holder transactions can exhaust one write-patience budget; keep
@@ -620,12 +553,10 @@ class SessionCompressionMixin:
def _do(conn): def _do(conn):
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
cursor = conn.execute( return conn.execute(
"UPDATE session_turn_leases SET expires_at = ? " "UPDATE session_turn_leases SET expires_at = ? "
"WHERE conversation_id = ? AND holder = ?", "WHERE conversation_id = ? AND holder = ?", (expires_at, conversation_id, holder),
(expires_at, conversation_id, holder), ).rowcount > 0
)
return cursor.rowcount > 0
return bool(self._execute_write(_do)) return bool(self._execute_write(_do))
@@ -637,10 +568,8 @@ class SessionCompressionMixin:
def _do(conn): def _do(conn):
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
conn.execute( conn.execute(
"DELETE FROM session_turn_leases " "DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?",
"WHERE conversation_id = ? AND holder = ?", (conversation_id, holder))
(conversation_id, holder),
)
self._execute_write(_do) self._execute_write(_do)
@@ -649,10 +578,8 @@ class SessionCompressionMixin:
if not session_id: if not session_id:
return None return None
row = self._read_one( row = self._read_one(
"SELECT holder FROM compression_locks " "SELECT holder FROM compression_locks WHERE session_id = ? AND expires_at >= ?",
"WHERE session_id = ? AND expires_at >= ?", (session_id, time.time()))
(session_id, time.time()),
)
return None if row is None else row[0] return None if row is None else row[0]
def finalize_orphaned_compression_sessions(self) -> int: def finalize_orphaned_compression_sessions(self) -> int:
@@ -660,10 +587,7 @@ class SessionCompressionMixin:
has messages, no end_reason/ended_at, api_call_count=0, older than 7 days) as has messages, no end_reason/ended_at, api_call_count=0, older than 7 days) as
``orphaned_compression``. Non-destructive.""" ``orphaned_compression``. Non-destructive."""
cutoff = time.time() - 604800 # 7 days cutoff = time.time() - 604800 # 7 days
return self._write_rowcount(
def _do(conn):
now = time.time()
result = conn.execute(
""" """
UPDATE sessions UPDATE sessions
SET ended_at = ?, SET ended_at = ?,
@@ -684,11 +608,8 @@ class SessionCompressionMixin:
WHERE m.session_id = sessions.id WHERE m.session_id = sessions.id
) )
""", """,
(now, cutoff), (time.time(), cutoff),
) ) or 0
return result.rowcount
return self._execute_write(_do) or 0
def get_compression_chain(self, session_id: str) -> List[str]: def get_compression_chain(self, session_id: str) -> List[str]:
"""Walk the compression-continuation chain forward: root-first through the tip """Walk the compression-continuation chain forward: root-first through the tip
@@ -703,7 +624,7 @@ class SessionCompressionMixin:
are still live over stale closed siblings such as ``ws_orphan_reap``.""" are still live over stale closed siblings such as ``ws_orphan_reap``."""
current = session_id current = session_id
chain = [current] if current else [] chain = [current] if current else []
seen = {current} if current else set() seen = set(chain)
for _ in range(100): # defensive bound; chains this deep are pathological for _ in range(100): # defensive bound; chains this deep are pathological
with self._read_ctx() as conn: with self._read_ctx() as conn:
row = conn.execute( row = conn.execute(
@@ -729,9 +650,7 @@ class SessionCompressionMixin:
""", """,
(current,), (current,),
).fetchone() ).fetchone()
if row is None: child_id = row["id"] if row is not None else None
return chain
child_id = row["id"]
if not child_id or child_id in seen: if not child_id or child_id in seen:
return chain return chain
seen.add(child_id) seen.add(child_id)

File diff suppressed because it is too large Load Diff

View File

@@ -79,8 +79,7 @@ class SessionTitlesMixin:
nothing overwrites a user name, re-running the titler on an llm row is a no-op). nothing overwrites a user name, re-running the titler on an llm row is a no-op).
No writer may move a hidden canonical Bot Chat off its title. Read and write are No writer may move a hidden canonical Bot Chat off its title. Read and write are
one compare-and-swap transaction, so a manual ``/title`` racing an in-flight one compare-and-swap transaction, so a manual ``/title`` racing an in-flight
generation is not clobbered. generation is not clobbered."""
"""
title = self.sanitize_title(title) title = self.sanitize_title(title)
is_user = source == self.TITLE_SOURCE_USER is_user = source == self.TITLE_SOURCE_USER
new_rank = self._title_rank(source) if not is_user else None new_rank = self._title_rank(source) if not is_user else None
@@ -166,21 +165,16 @@ class SessionTitlesMixin:
if source not in self._TITLE_SOURCE_RANK: if source not in self._TITLE_SOURCE_RANK:
raise ValueError(f"invalid title source: {source!r}") raise ValueError(f"invalid title source: {source!r}")
return self._write_rowcount( return self._write_rowcount(
"UPDATE sessions SET title_source = ? " "UPDATE sessions SET title_source = ? WHERE id = ? AND title IS NOT NULL",
"WHERE id = ? AND title IS NOT NULL",
(source, session_id), (source, session_id),
) > 0 ) > 0
def get_session_by_title(self, title: str) -> Optional[Dict[str, Any]]: def get_session_by_title(self, title: str) -> Optional[Dict[str, Any]]:
"""Look up a session by exact title. Returns session dict or None.""" """Look up a session by exact title. Returns session dict or None."""
row = self._read_one( row = self._read_one(
"SELECT s.*, " "SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved "
"COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved " "FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
"FROM sessions s " "WHERE s.title = ?", (title,))
"LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
"WHERE s.title = ?",
(title,),
)
return self._session_row_dict(row) if row else None return self._session_row_dict(row) if row else None
def resolve_session_by_title(self, title: str) -> Optional[str]: def resolve_session_by_title(self, title: str) -> Optional[str]:
@@ -191,8 +185,7 @@ class SessionTitlesMixin:
numbered = self._read_all( numbered = self._read_all(
"SELECT id, title, started_at FROM sessions " "SELECT id, title, started_at FROM sessions "
"WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC", "WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC",
(f"{_escape_like(title)} #%",), (f"{_escape_like(title)} #%",))
)
if numbered: if numbered:
return numbered[0]["id"] return numbered[0]["id"]
return exact["id"] if exact else None return exact["id"] if exact else None
@@ -204,13 +197,9 @@ class SessionTitlesMixin:
base = match.group(1) if match else base_title base = match.group(1) if match else base_title
rows = self._read_all( rows = self._read_all(
"SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'", "SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'",
(base, f"{_escape_like(base)} #%"), (base, f"{_escape_like(base)} #%"))
)
if not rows: if not rows:
return base return base
max_num = 1 # the unnumbered original counts as #1 # The unnumbered original counts as #1.
for row in rows: numbers = [int(m.group(1)) for m in (re.match(r'^.* #(\d+)$', row["title"]) for row in rows) if m]
m = re.match(r'^.* #(\d+)$', row["title"]) return f"{base} #{max([1, *numbers]) + 1}"
if m:
max_num = max(max_num, int(m.group(1)))
return f"{base} #{max_num + 1}"

View File

@@ -78,6 +78,14 @@ _MODEL_USAGE_UPSERT_SQL = """INSERT INTO session_model_usage (
last_seen = excluded.last_seen""" last_seen = excluded.last_seen"""
# Kwargs forwarded verbatim from update_token_counts / record_auxiliary_usage into
# _record_model_usage (the per-route attribution row).
_MODEL_USAGE_FIELDS = frozenset((
"model", "billing_provider", "billing_base_url", "billing_mode", "input_tokens", "output_tokens",
"cache_read_tokens", "cache_write_tokens", "reasoning_tokens", "estimated_cost_usd",
"actual_cost_usd", "cost_status", "cost_source", "api_call_count"))
class SessionUsageMixin: class SessionUsageMixin:
"""Coalesced token writer, per-model usage rows, billing route.""" """Coalesced token writer, per-model usage rows, billing route."""
@@ -110,16 +118,16 @@ class SessionUsageMixin:
to the synchronous path and may raise.""" to the synchronous path and may raise."""
with self._token_queue_cond: with self._token_queue_cond:
thread = self._token_writer_thread thread = self._token_writer_thread
writer_stopped = self._token_writer_stop and (thread is None or not thread.is_alive()) writer_alive = thread is not None and thread.is_alive()
writer_stopped = self._token_writer_stop and not writer_alive
if not writer_stopped: if not writer_stopped:
self._token_queue.append((session_id, kwargs)) self._token_queue.append((session_id, kwargs))
if thread is None or not thread.is_alive(): if not writer_alive:
# Daemon so exit never hangs on accounting; the atexit hook drains # Daemon so exit never hangs on accounting; the atexit hook drains
# leftovers. ``not is_alive()`` (not ``is None``) respawns a writer # leftovers. ``not is_alive()`` (not ``is None``) respawns a writer
# that died from an unexpected escape. # that died from an unexpected escape.
thread = threading.Thread( thread = threading.Thread(
target=self._token_writer_loop, name="session-db-token-writer", daemon=True, target=self._token_writer_loop, name="session-db-token-writer", daemon=True)
)
self._token_writer_thread = thread self._token_writer_thread = thread
thread.start() thread.start()
if self._token_atexit_hook is None: if self._token_atexit_hook is None:
@@ -245,9 +253,7 @@ class SessionUsageMixin:
# Writer stuck mid-apply: leave deltas unapplied rather than race it. # Writer stuck mid-apply: leave deltas unapplied rather than race it.
logger.warning( logger.warning(
"async token accounting: writer did not stop within %.0fs; " "async token accounting: writer did not stop within %.0fs; "
"%d queued delta(s) not persisted", "%d queued delta(s) not persisted", join_timeout, len(self._token_queue))
join_timeout, len(self._token_queue),
)
return return
# Writer gone: apply leftovers synchronously under the same busy protocol. Wait # Writer gone: apply leftovers synchronously under the same busy protocol. Wait
# out a flush caller-drain that already claimed busy — close() nulls the # out a flush caller-drain that already claimed busy — close() nulls the
@@ -260,8 +266,7 @@ class SessionUsageMixin:
logger.warning( logger.warning(
"async token accounting: concurrent drain did not " "async token accounting: concurrent drain did not "
"finish within %.0fs; %d queued delta(s) not persisted", "finish within %.0fs; %d queued delta(s) not persisted",
join_timeout, len(self._token_queue), join_timeout, len(self._token_queue))
)
return return
self._token_queue_cond.wait(remaining) self._token_queue_cond.wait(remaining)
# busy BEFORE clearing the queue (same ordering as the writer loop). # busy BEFORE clearing the queue (same ordering as the writer loop).
@@ -290,6 +295,7 @@ class SessionUsageMixin:
"""Update token counters and backfill model if unset. *absolute*=False """Update token counters and backfill model if unset. *absolute*=False
increments (per-API-call deltas, CLI path); *absolute*=True sets directly increments (per-API-call deltas, CLI path); *absolute*=True sets directly
(gateway path, where the cached agent holds cumulative totals).""" (gateway path, where the cached agent holds cumulative totals)."""
usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS}
# Ensure the row exists: under concurrent load create_session() may have failed # Ensure the row exists: under concurrent load create_session() may have failed
# on locking, and the UPDATE would silently affect 0 rows. # on locking, and the UPDATE would silently affect 0 rows.
self._insert_session_row(session_id, "unknown", model=model) self._insert_session_row(session_id, "unknown", model=model)
@@ -305,8 +311,7 @@ class SessionUsageMixin:
billing_provider if has_accounted_usage else None, billing_provider if has_accounted_usage else None,
billing_base_url if has_accounted_usage else None, billing_base_url if has_accounted_usage else None,
billing_mode if has_accounted_usage else None, model if has_accounted_usage else None, billing_mode if has_accounted_usage else None, model if has_accounted_usage else None,
api_call_count, session_id, api_call_count, session_id)
)
# Per-model attribution: the sessions row keeps one (model, provider) pair, so a # Per-model attribution: the sessions row keeps one (model, provider) pair, so a
# mid-session /model switch would attribute every token to the initial model. # mid-session /model switch would attribute every token to the initial model.
# Only the incremental path records here — absolute cumulative updates cannot be # Only the incremental path records here — absolute cumulative updates cannot be
@@ -317,18 +322,14 @@ class SessionUsageMixin:
row = conn.execute( row = conn.execute(
"SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", (session_id,), "SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", (session_id,),
).fetchone() ).fetchone()
existing_model = row["model"] if row is not None else None existing = dict(row) if row is not None else {}
existing_provider = row["billing_provider"] if row is not None else None
existing_api_calls = int((row["api_call_count"] if row is not None else 0) or 0)
# create_session records the requested route before any API call. If that # create_session records the requested route before any API call. If that
# fails and fallback succeeds, the first accounted usage is the authoritative # fails and fallback succeeds, the first accounted usage is the authoritative
# route; after that keep the row as is (one row cannot represent mixed usage). # route; after that keep the row as is (one row cannot represent mixed usage).
first_accounted_route = ( first_accounted_route = (
existing_api_calls == 0 int(existing.get("api_call_count") or 0) == 0 and has_accounted_usage and bool(model)
and has_accounted_usage
and bool(model)
and bool(billing_provider) and bool(billing_provider)
and (existing_model != model or existing_provider != billing_provider) and (existing.get("model") != model or existing.get("billing_provider") != billing_provider)
) )
if first_accounted_route: if first_accounted_route:
conn.execute( conn.execute(
@@ -340,51 +341,39 @@ class SessionUsageMixin:
) )
conn.execute(sql, params) conn.execute(sql, params)
if record_model_usage: if record_model_usage:
self._record_model_usage( self._record_model_usage(conn, session_id, **usage)
conn, session_id, model=model, billing_provider=billing_provider,
billing_base_url=billing_base_url, billing_mode=billing_mode,
input_tokens=input_tokens, output_tokens=output_tokens,
cache_read_tokens=cache_read_tokens, cache_write_tokens=cache_write_tokens,
reasoning_tokens=reasoning_tokens, estimated_cost_usd=estimated_cost_usd,
actual_cost_usd=actual_cost_usd, cost_status=cost_status, cost_source=cost_source,
api_call_count=api_call_count,
)
self._execute_write(_do) self._execute_write(_do)
def _record_model_usage( def _record_model_usage(
self, conn, session_id: str, *, model: Optional[str], billing_provider: Optional[str], self, conn, session_id: str, *, model: Optional[str] = None, billing_provider: Optional[str] = None,
billing_base_url: Optional[str], billing_mode: Optional[str], input_tokens: int, billing_base_url: Optional[str] = None, billing_mode: Optional[str] = None, input_tokens: int = 0,
output_tokens: int, cache_read_tokens: int, cache_write_tokens: int, reasoning_tokens: int, output_tokens: int = 0, cache_read_tokens: int = 0, cache_write_tokens: int = 0,
estimated_cost_usd: Optional[float], actual_cost_usd: Optional[float], reasoning_tokens: int = 0, estimated_cost_usd: Optional[float] = None,
cost_status: Optional[str], cost_source: Optional[str], api_call_count: int, task: str = "", actual_cost_usd: Optional[float] = None, cost_status: Optional[str] = None,
cost_source: Optional[str] = None, api_call_count: int = 0, task: str = "",
) -> None: ) -> None:
"""Accumulate a per-API-call usage delta into session_model_usage, inside the """Accumulate a per-API-call usage delta into session_model_usage, inside the
caller's write txn after the ``sessions`` UPDATE. A missing model/provider falls caller's write txn after the ``sessions`` UPDATE. A missing model/provider falls
back to the session row (same COALESCE behaviour as the summary update) — except back to the session row (same COALESCE behaviour as the summary update) — except
for aux rows (``task`` set), which must NOT inherit the main-loop route (vision for aux rows (``task`` set), which must NOT inherit the main-loop route (vision
on gemini while the main loop runs anthropic): missing info stays 'unknown'/empty. on gemini while the main loop runs anthropic): missing info stays 'unknown'/empty."""
"""
row = conn.execute( row = conn.execute(
"SELECT model, billing_provider, billing_base_url, billing_mode " "SELECT model, billing_provider, billing_base_url, billing_mode "
"FROM sessions WHERE id = ?", "FROM sessions WHERE id = ?", (session_id,),
(session_id,),
).fetchone() ).fetchone()
sess = dict(row) if (row is not None and not task) else {} sess = dict(row) if (row is not None and not task) else {}
eff_model = model or sess.get("model") or "unknown"
eff_provider = billing_provider or sess.get("billing_provider") or ""
eff_base_url = billing_base_url or sess.get("billing_base_url") or ""
eff_billing_mode = billing_mode or sess.get("billing_mode") or ""
counts = [v or 0 for v in (input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens)] counts = [v or 0 for v in (input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens)]
now = time.time() now = time.time()
conn.execute( conn.execute(
_MODEL_USAGE_UPSERT_SQL, _MODEL_USAGE_UPSERT_SQL,
( (
session_id, eff_model, eff_provider, eff_base_url, eff_billing_mode, task or "", session_id, model or sess.get("model") or "unknown",
billing_provider or sess.get("billing_provider") or "",
billing_base_url or sess.get("billing_base_url") or "",
billing_mode or sess.get("billing_mode") or "", task or "",
api_call_count or 0, *counts, api_call_count or 0, *counts,
float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0), float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0),
cost_status, cost_source, now, now, cost_status, cost_source, now, now))
),
)
def record_auxiliary_usage( def record_auxiliary_usage(
self, session_id: str, task: str, *, model: Optional[str] = None, self, session_id: str, task: str, *, model: Optional[str] = None,
@@ -398,22 +387,13 @@ class SessionUsageMixin:
touching the ``sessions`` summary row (the gateway overwrites those counters with touching the ``sessions`` summary row (the gateway overwrites those counters with
absolute main-loop totals). ``api_call_count`` may aggregate N calls. Best-effort: absolute main-loop totals). ``api_call_count`` may aggregate N calls. Best-effort:
callers must never fail an aux call over accounting.""" callers must never fail an aux call over accounting."""
usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS}
if not session_id or not task: if not session_id or not task:
return return
usage["api_call_count"] = 1 if api_call_count is None else int(api_call_count)
# FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts. # FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts.
self._insert_session_row(session_id, "unknown") self._insert_session_row(session_id, "unknown")
self._execute_write(lambda conn: self._record_model_usage(conn, session_id, task=task, **usage))
def _do(conn):
self._record_model_usage(
conn, session_id, model=model, billing_provider=billing_provider,
billing_base_url=billing_base_url, billing_mode=None,
input_tokens=input_tokens or 0, output_tokens=output_tokens or 0,
cache_read_tokens=cache_read_tokens or 0, cache_write_tokens=cache_write_tokens or 0,
reasoning_tokens=reasoning_tokens or 0, estimated_cost_usd=estimated_cost_usd,
actual_cost_usd=None, cost_status=None, cost_source=None,
api_call_count=1 if api_call_count is None else int(api_call_count), task=task,
)
self._execute_write(_do)
def usage_totals(self, *, min_message_count: int = 1, include_archived: bool = False) -> Dict[str, float]: def usage_totals(self, *, min_message_count: int = 1, include_archived: bool = False) -> Dict[str, float]:
"""Tokens and spend across the whole store (one scan), so the sidebar total does """Tokens and spend across the whole store (one scan), so the sidebar total does