refactor(state): resume — verified messages/compression/titles/usage simplification
This commit is contained in:
+225
-450
@@ -16,30 +16,29 @@ from hermes_state_common import _sql_session_last_active, is_automatic_end_reaso
|
||||
# Log-record parity with the origin module (caplog tests pin "hermes_state").
|
||||
logger = logging.getLogger("hermes_state")
|
||||
|
||||
_ENDED_ROW_SQL = "SELECT ended_at, end_reason FROM sessions WHERE id = ?"
|
||||
_LOCK_ROW_SQL = "SELECT holder, expires_at FROM compression_locks WHERE session_id = ?"
|
||||
_COOLDOWN_ROW_SQL = (
|
||||
"SELECT compression_failure_cooldown_until, compression_failure_error FROM sessions WHERE id = ?"
|
||||
)
|
||||
|
||||
|
||||
def _ended_by_compression(row) -> bool:
|
||||
return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression"
|
||||
|
||||
|
||||
class SessionCompressionMixin:
|
||||
"""Compression lineage, cooldown/streak counters, locks and turn leases."""
|
||||
|
||||
def find_live_compression_child(
|
||||
self, parent_session_id: str
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Return the unique live direct child of a compression-ended session.
|
||||
def find_live_compression_child(self, parent_session_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""The unique live direct child of a compression-ended session, else None.
|
||||
|
||||
A stale agent whose parent was rotated elsewhere may recover only when the
|
||||
lineage names exactly one live continuation; more than one fails closed
|
||||
rather than guessing which transcript owns later messages."""
|
||||
lineage names exactly one live continuation; more than one fails closed."""
|
||||
if not parent_session_id:
|
||||
return None
|
||||
with self._read_ctx() as conn:
|
||||
parent = conn.execute(
|
||||
"SELECT ended_at, end_reason FROM sessions WHERE id = ?",
|
||||
(parent_session_id,),
|
||||
).fetchone()
|
||||
if (
|
||||
parent is None
|
||||
or parent["ended_at"] is None
|
||||
or parent["end_reason"] != "compression"
|
||||
):
|
||||
if not _ended_by_compression(conn.execute(_ENDED_ROW_SQL, (parent_session_id,)).fetchone()):
|
||||
return None
|
||||
rows = conn.execute(
|
||||
"""
|
||||
@@ -61,28 +60,16 @@ class SessionCompressionMixin:
|
||||
return self._session_row_dict(rows[0]) if len(rows) == 1 else None
|
||||
|
||||
def reopen_orphaned_compression_session(self, session_id: str) -> bool:
|
||||
"""Reopen a compression parent only when no continuation was published.
|
||||
|
||||
Publication is atomic now, but older builds could leave a closed parent
|
||||
after an interrupted handoff. Conservative by design: an active lease or
|
||||
any canonical child means another path owns the lineage — fail closed."""
|
||||
"""Reopen a compression parent only when no continuation was published (older
|
||||
builds could leave a closed parent after an interrupted handoff). Conservative:
|
||||
an active lease or any canonical child means another path owns the lineage."""
|
||||
if not session_id:
|
||||
return False
|
||||
|
||||
def _do(conn):
|
||||
parent = conn.execute(
|
||||
"SELECT ended_at, end_reason FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if (
|
||||
parent is None
|
||||
or parent["ended_at"] is None
|
||||
or parent["end_reason"] != "compression"
|
||||
):
|
||||
if not _ended_by_compression(conn.execute(_ENDED_ROW_SQL, (session_id,)).fetchone()):
|
||||
return False
|
||||
|
||||
# Any non-branch/non-delegate/non-tool child is a continuation, ended
|
||||
# or not; reopening past it could give one lineage a second live head.
|
||||
# Any non-branch/non-delegate/non-tool child is a continuation, ended or not.
|
||||
child = conn.execute(
|
||||
"""
|
||||
SELECT 1
|
||||
@@ -97,17 +84,11 @@ class SessionCompressionMixin:
|
||||
).fetchone()
|
||||
if child is not None:
|
||||
return False
|
||||
|
||||
# refresh_compression_lock() lets an owner revive its own expired
|
||||
# row, so reclaim it inside this write transaction before reopening:
|
||||
# refresh-first makes the lease active and aborts recovery;
|
||||
# recovery-first deletes the holder so a later refresh can't resurrect it.
|
||||
# refresh_compression_lock() lets an owner revive its own expired row, so
|
||||
# reclaim it inside this write txn: refresh-first makes the lease active and
|
||||
# aborts recovery; recovery-first deletes the holder so a refresh can't resurrect it.
|
||||
now = time.time()
|
||||
lock_row = conn.execute(
|
||||
"SELECT holder, expires_at FROM compression_locks "
|
||||
"WHERE session_id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
lock_row = conn.execute(_LOCK_ROW_SQL, (session_id,)).fetchone()
|
||||
if lock_row is not None:
|
||||
expires_at = lock_row["expires_at"]
|
||||
if expires_at is None or float(expires_at) >= now:
|
||||
@@ -119,20 +100,56 @@ class SessionCompressionMixin:
|
||||
)
|
||||
if deleted.rowcount != 1:
|
||||
return False
|
||||
|
||||
updated = conn.execute(
|
||||
"UPDATE sessions SET ended_at = NULL, end_reason = NULL "
|
||||
"WHERE id = ? AND ended_at IS NOT NULL "
|
||||
"AND end_reason = 'compression'",
|
||||
(session_id,),
|
||||
)
|
||||
# rowcount==1 is guaranteed by the parent SELECT in this same BEGIN
|
||||
# IMMEDIATE transaction. If a False return is ever added past this
|
||||
# point, raise instead: _execute_write commits the lease DELETE above unless _do raises.
|
||||
# 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
|
||||
# commits unless _do raises.
|
||||
return updated.rowcount == 1
|
||||
|
||||
return bool(self._execute_write(_do))
|
||||
|
||||
def _publish_child_session_row(self, conn, parent, *, parent_session_id, child_session_id, source,
|
||||
model, model_config, system_prompt, cwd, profile_name) -> None:
|
||||
"""INSERT the compression child's ``sessions`` row copied from *parent*."""
|
||||
system_prompt_hash = self._store_system_prompt(conn, system_prompt)
|
||||
conn.execute(
|
||||
"""INSERT INTO sessions (
|
||||
id, source, model, model_config, system_prompt,
|
||||
system_prompt_hash,
|
||||
parent_session_id, cwd, git_branch, git_repo_root,
|
||||
profile_name, user_id, session_key, chat_id, chat_type,
|
||||
thread_id, display_name, origin_json, started_at
|
||||
) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
||||
(
|
||||
child_session_id,
|
||||
source,
|
||||
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"],
|
||||
# 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(),
|
||||
parent["user_id"],
|
||||
parent["session_key"],
|
||||
parent["chat_id"],
|
||||
parent["chat_type"],
|
||||
parent["thread_id"],
|
||||
parent["display_name"],
|
||||
parent["origin_json"],
|
||||
time.time(),
|
||||
),
|
||||
)
|
||||
|
||||
def publish_compression_child(
|
||||
self,
|
||||
*,
|
||||
@@ -154,35 +171,29 @@ class SessionCompressionMixin:
|
||||
) -> None:
|
||||
"""Atomically close a parent and publish its durable compression child.
|
||||
|
||||
Closure, child row, and handoff commit in one transaction: readers see
|
||||
the live parent or a complete child, never an ended parent with a
|
||||
missing/empty child.
|
||||
Closure, child row, and handoff commit in one transaction: readers see the live
|
||||
parent or a complete child, never an ended parent with a missing/empty child.
|
||||
|
||||
*watermark* (parent's ``get_active_message_watermark`` at compression
|
||||
start): parent rows with ``id > watermark`` — appends landed during the
|
||||
slow summary call — are column-cloned into the child AFTER the handoff
|
||||
so they survive rotation. *watermark_ceiling* bounds the clone: the
|
||||
rotation path flushes its OWN transcript to the parent just before
|
||||
publishing and those rows are already in the handoff, so the caller
|
||||
captures ``MAX(id)`` right BEFORE that flush and only
|
||||
``(watermark, watermark_ceiling]`` is foreign tail. ``None`` = unbounded.
|
||||
*watermark* (parent's ``get_active_message_watermark`` at compression start):
|
||||
parent rows with ``id > watermark`` — appends landed during the slow summary —
|
||||
are column-cloned into the child AFTER the handoff. *watermark_ceiling* bounds
|
||||
the clone: the rotation path flushes its OWN transcript to the parent just
|
||||
before publishing and those rows are already in the handoff, so only
|
||||
``(watermark, watermark_ceiling]`` is foreign tail (``None`` = unbounded).
|
||||
|
||||
*require_lease_refresh* + *compression_lock_holder* refreshes the lease
|
||||
on the same ``conn`` before the expiry check (no TOCTOU window), so a
|
||||
refresher that died on transient DB errors gets one last chance."""
|
||||
*require_lease_refresh* + *compression_lock_holder* refreshes the lease on the
|
||||
same ``conn`` before the expiry check (no TOCTOU window), so a refresher that
|
||||
died on transient DB errors gets one last chance."""
|
||||
from hermes_state import CompressionSessionBusyError
|
||||
|
||||
def _do(conn):
|
||||
if require_lease_refresh and compression_lock_holder:
|
||||
conn.execute(
|
||||
"UPDATE compression_locks SET expires_at = ? "
|
||||
"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(
|
||||
"SELECT holder, expires_at FROM compression_locks WHERE session_id = ?",
|
||||
(parent_session_id,),
|
||||
).fetchone()
|
||||
lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone()
|
||||
if require_compression_lease and (
|
||||
lock_row is None
|
||||
or not compression_lock_holder
|
||||
@@ -202,14 +213,11 @@ class SessionCompressionMixin:
|
||||
if parent is None:
|
||||
raise RuntimeError(f"Compression parent not found: {parent_session_id}")
|
||||
if parent["ended_at"] is not None:
|
||||
# An ended stamp from AUTOMATIC cleanup (tui_shutdown, ws_disconnect,
|
||||
# orphan reap, idle/LRU evict) is stale by construction — this lease
|
||||
# holder is still continuing the conversation. Left alone it wedges
|
||||
# rotation forever (every attempt aborts here; each pre-publish flush
|
||||
# re-grows the parent until the provider rejects it). Clear it; the
|
||||
# closure UPDATE below re-stamps end_reason='compression'. Deliberate
|
||||
# boundaries (compression, session_reset, explicit close) still fail
|
||||
# closed — another path owns the lineage.
|
||||
# An AUTOMATIC end stamp (tui_shutdown, ws_disconnect, orphan reap,
|
||||
# idle/LRU evict) is stale by construction — this lease holder is still
|
||||
# continuing the conversation, and left alone it wedges rotation forever.
|
||||
# Clear it; the closure UPDATE below re-stamps end_reason='compression'.
|
||||
# Deliberate boundaries still fail closed.
|
||||
if is_automatic_end_reason(parent["end_reason"]):
|
||||
conn.execute(
|
||||
"UPDATE sessions SET ended_at = NULL, end_reason = NULL "
|
||||
@@ -217,90 +225,34 @@ class SessionCompressionMixin:
|
||||
(parent_session_id,),
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Compression parent already ended: {parent_session_id}"
|
||||
)
|
||||
raise RuntimeError(f"Compression parent already ended: {parent_session_id}")
|
||||
if not messages:
|
||||
raise RuntimeError("Compression child handoff must not be empty")
|
||||
system_prompt_hash = self._store_system_prompt(conn, system_prompt)
|
||||
|
||||
conn.execute(
|
||||
"""INSERT INTO sessions (
|
||||
id, source, model, model_config, system_prompt,
|
||||
system_prompt_hash,
|
||||
parent_session_id, cwd, git_branch, git_repo_root,
|
||||
profile_name, user_id, session_key, chat_id, chat_type,
|
||||
thread_id, display_name, origin_json, started_at
|
||||
) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
||||
(
|
||||
child_session_id,
|
||||
source,
|
||||
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"],
|
||||
# Same contract as _insert_session_row's compression-fork backfill:
|
||||
# child stays on the parent's profile and keeps gateway routing/
|
||||
# origin columns so peer recovery works after a boundary crash. No
|
||||
# owner on either side (legacy NULL parent) → stamp this store's
|
||||
# profile so the child doesn't extend the unowned lineage.
|
||||
profile_name
|
||||
or parent["profile_name"]
|
||||
or self._own_profile_name(),
|
||||
parent["user_id"],
|
||||
parent["session_key"],
|
||||
parent["chat_id"],
|
||||
parent["chat_type"],
|
||||
parent["thread_id"],
|
||||
parent["display_name"],
|
||||
parent["origin_json"],
|
||||
time.time(),
|
||||
),
|
||||
)
|
||||
total_messages, total_tool_calls = self._insert_message_rows(
|
||||
conn, child_session_id, messages
|
||||
self._publish_child_session_row(
|
||||
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,
|
||||
cwd=cwd, profile_name=profile_name,
|
||||
)
|
||||
total_messages, total_tool_calls = self._insert_message_rows(conn, child_session_id, messages)
|
||||
if watermark is not None:
|
||||
# Clone the parent's concurrent tail (see docstring) into the
|
||||
# child after the handoff: column-exact except id/session_id;
|
||||
# Clone the parent's concurrent tail into the child after the handoff;
|
||||
# originals stay in the closed parent for lineage recovery.
|
||||
_ceiling_clause = ""
|
||||
_params: list = [parent_session_id, int(watermark)]
|
||||
if watermark_ceiling is not None:
|
||||
_ceiling_clause = " AND id <= ?"
|
||||
_params.append(int(watermark_ceiling))
|
||||
tail_rows = conn.execute(
|
||||
tail_ids, tail_tool_calls = self._tail_rows_after_watermark(
|
||||
conn,
|
||||
"SELECT id, tool_calls FROM messages "
|
||||
"WHERE session_id = ? AND active = 1 AND id > ?"
|
||||
f"{_ceiling_clause} ORDER BY id",
|
||||
_params,
|
||||
).fetchall()
|
||||
if tail_rows:
|
||||
tail_ids = [int(r["id"]) for r in tail_rows]
|
||||
placeholders = ",".join("?" for _ in tail_ids)
|
||||
clone_cols = [
|
||||
c for c in self._message_column_names(conn)
|
||||
if c not in ("id", "session_id", "active", "compacted")
|
||||
]
|
||||
col_list = ", ".join(clone_cols)
|
||||
conn.execute(
|
||||
f"INSERT INTO messages ({col_list}, session_id, active, compacted) "
|
||||
f"SELECT {col_list}, ?, 1, 0 FROM messages "
|
||||
f"WHERE id IN ({placeholders}) ORDER BY id",
|
||||
[child_session_id, *tail_ids],
|
||||
)
|
||||
)
|
||||
if tail_ids:
|
||||
self._clone_message_rows(conn, tail_ids, session_id=child_session_id)
|
||||
total_messages += len(tail_ids)
|
||||
for r in tail_rows:
|
||||
raw = r["tool_calls"]
|
||||
if raw:
|
||||
try:
|
||||
parsed = json.loads(raw) if isinstance(raw, str) else raw
|
||||
total_tool_calls += len(parsed) if isinstance(parsed, list) else 0
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
total_tool_calls += tail_tool_calls
|
||||
conn.execute(
|
||||
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
|
||||
(total_messages, total_tool_calls, child_session_id),
|
||||
@@ -311,111 +263,66 @@ class SessionCompressionMixin:
|
||||
(time.time(), parent_session_id),
|
||||
)
|
||||
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}")
|
||||
|
||||
self._execute_write(_do)
|
||||
|
||||
def record_compression_failure_cooldown(
|
||||
self,
|
||||
session_id: str,
|
||||
cooldown_until: float,
|
||||
error: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Persist the active compression-failure cooldown for a session."""
|
||||
def _write_sql_logged(self, op: str, session_id: str, sql: str, params) -> None:
|
||||
"""``_write_sql`` that logs (never raises) on ``sqlite3.Error``."""
|
||||
try:
|
||||
self._write_sql(sql, params)
|
||||
except sqlite3.Error as exc:
|
||||
logger.warning("%s(%s) failed: %s", op, session_id, exc)
|
||||
|
||||
def record_compression_failure_cooldown(self, session_id: str, cooldown_until: float, error: Optional[str] = None) -> None:
|
||||
"""Persist the active compression-failure cooldown. Merge-max with any longer
|
||||
live deadline so a later shorter write can't reopen the thrash window; error
|
||||
always takes the latest diagnostic."""
|
||||
if not session_id:
|
||||
return
|
||||
self._write_sql_logged(
|
||||
"record_compression_failure_cooldown", session_id,
|
||||
"UPDATE sessions SET compression_failure_cooldown_until = CASE "
|
||||
"WHEN compression_failure_cooldown_until IS NOT NULL "
|
||||
" AND compression_failure_cooldown_until > ? "
|
||||
"THEN compression_failure_cooldown_until ELSE ? END, "
|
||||
"compression_failure_error = ? WHERE id = ?",
|
||||
(cooldown_until, cooldown_until, error, session_id),
|
||||
)
|
||||
|
||||
try:
|
||||
# Merge-max with any longer live deadline so a later shorter write
|
||||
# can't reopen the thrash window; error always takes the latest diagnostic.
|
||||
self._write_sql(
|
||||
"UPDATE sessions SET compression_failure_cooldown_until = CASE "
|
||||
"WHEN compression_failure_cooldown_until IS NOT NULL "
|
||||
" AND compression_failure_cooldown_until > ? "
|
||||
"THEN compression_failure_cooldown_until ELSE ? END, "
|
||||
"compression_failure_error = ? WHERE id = ?",
|
||||
(cooldown_until, cooldown_until, error, session_id),
|
||||
)
|
||||
except sqlite3.Error as exc:
|
||||
logger.warning(
|
||||
"record_compression_failure_cooldown(%s) failed: %s",
|
||||
session_id, exc,
|
||||
)
|
||||
|
||||
def get_compression_failure_cooldown(
|
||||
self,
|
||||
session_id: str,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Return the active compression-failure cooldown for ``session_id``."""
|
||||
def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""Return the active (unexpired) compression-failure cooldown, or None."""
|
||||
if not session_id:
|
||||
return None
|
||||
now = time.time()
|
||||
row = self._read_one(
|
||||
"SELECT compression_failure_cooldown_until, compression_failure_error "
|
||||
"FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
)
|
||||
if row is None:
|
||||
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,))
|
||||
if row is None or row[0] is None:
|
||||
return None
|
||||
cooldown_until = row[0]
|
||||
if cooldown_until is None:
|
||||
return None
|
||||
cooldown_until = float(cooldown_until)
|
||||
cooldown_until = float(row[0])
|
||||
if cooldown_until <= now:
|
||||
return None
|
||||
error = row[1]
|
||||
return {
|
||||
"cooldown_until": cooldown_until,
|
||||
"remaining_seconds": cooldown_until - now,
|
||||
"error": error,
|
||||
}
|
||||
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]:
|
||||
"""Exact stored cooldown columns, no expiry filtering. Compression
|
||||
cancellation uses this under its session lease so rollback preserves an
|
||||
expired, partially-null, or absent row exactly instead of coercing it
|
||||
through the active-cooldown API."""
|
||||
if not session_id:
|
||||
return {"session_exists": False, "cooldown_until": None, "error": None}
|
||||
row = self._read_one(
|
||||
"SELECT compression_failure_cooldown_until, compression_failure_error "
|
||||
"FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
)
|
||||
def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]:
|
||||
"""Exact stored cooldown columns, no expiry filtering, so compression
|
||||
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
|
||||
if row is None:
|
||||
return {"session_exists": False, "cooldown_until": None, "error": None}
|
||||
cooldown_until = row[0]
|
||||
error = row[1]
|
||||
return {
|
||||
"session_exists": True,
|
||||
"cooldown_until": (
|
||||
float(cooldown_until) if cooldown_until is not None else None
|
||||
),
|
||||
"error": error,
|
||||
"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:
|
||||
"""Restore and verify an exact cooldown-row snapshot. Unlike record/clear,
|
||||
this rollback API propagates write and verification failures: cancellation
|
||||
must not be reported mutation-free when compensation failed."""
|
||||
expected_exists = bool(snapshot.get("session_exists", False))
|
||||
if not expected_exists:
|
||||
actual = self.get_compression_failure_cooldown_row(session_id)
|
||||
if actual.get("session_exists", False):
|
||||
raise RuntimeError(
|
||||
"cannot restore absent compression cooldown row: session now exists"
|
||||
)
|
||||
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
|
||||
rollback API propagates write and verification failures: cancellation must not
|
||||
be reported mutation-free when compensation failed."""
|
||||
if not snapshot.get("session_exists", False):
|
||||
if self.get_compression_failure_cooldown_row(session_id).get("session_exists", False):
|
||||
raise RuntimeError("cannot restore absent compression cooldown row: session now exists")
|
||||
return
|
||||
|
||||
deadline = snapshot.get("cooldown_until")
|
||||
error = snapshot.get("error")
|
||||
|
||||
@@ -426,9 +333,7 @@ class SessionCompressionMixin:
|
||||
(deadline, error, session_id),
|
||||
)
|
||||
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)
|
||||
actual = self.get_compression_failure_cooldown_row(session_id)
|
||||
@@ -447,27 +352,19 @@ class SessionCompressionMixin:
|
||||
"""Clear any persisted compression-failure cooldown for a session."""
|
||||
if not session_id:
|
||||
return
|
||||
|
||||
try:
|
||||
self._write_sql(
|
||||
"UPDATE sessions SET compression_failure_cooldown_until = NULL, "
|
||||
"compression_failure_error = NULL WHERE id = ?",
|
||||
(session_id,),
|
||||
)
|
||||
except sqlite3.Error as exc:
|
||||
logger.warning(
|
||||
"clear_compression_failure_cooldown(%s) failed: %s",
|
||||
session_id, exc,
|
||||
)
|
||||
self._write_sql_logged(
|
||||
"clear_compression_failure_cooldown", session_id,
|
||||
"UPDATE sessions SET compression_failure_cooldown_until = NULL, "
|
||||
"compression_failure_error = NULL WHERE id = ?",
|
||||
(session_id,),
|
||||
)
|
||||
|
||||
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, NULL, or unparsable value also reads as ``zero``."""
|
||||
"""Read one numeric ``sessions`` column clamped at ``zero``; a missing session,
|
||||
NULL, or unparsable value also reads as ``zero``."""
|
||||
if not session_id:
|
||||
return zero
|
||||
row = self._read_one(
|
||||
f"SELECT {column} FROM sessions WHERE id = ?", (session_id,)
|
||||
)
|
||||
row = self._read_one(f"SELECT {column} FROM sessions WHERE id = ?", (session_id,))
|
||||
if row is None:
|
||||
return zero
|
||||
try:
|
||||
@@ -488,9 +385,9 @@ class SessionCompressionMixin:
|
||||
)
|
||||
|
||||
def get_compression_ineffective_count(self, session_id: str) -> int:
|
||||
"""Persisted ineffective-compaction strike count: the durable half of
|
||||
the built-in compressor's anti-thrash guard, so a fresh compressor bound
|
||||
to a resumed session inherits an armed/tripped guard across restarts."""
|
||||
"""Persisted ineffective-compaction strike count — the durable half of the
|
||||
built-in compressor's anti-thrash guard, so a fresh compressor bound to a resumed
|
||||
session inherits an armed/tripped guard across restarts."""
|
||||
return self._read_session_number("compression_ineffective_count", session_id, int, 0)
|
||||
|
||||
def set_compression_ineffective_count(self, session_id: str, count: int) -> None:
|
||||
@@ -502,10 +399,8 @@ class SessionCompressionMixin:
|
||||
)
|
||||
|
||||
def get_compression_recovery_deadline(self, session_id: str) -> float:
|
||||
"""Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed).
|
||||
Durable because the gateway rebuilds the compressor every turn / cache
|
||||
eviction: a process-local deadline restarted on each rebuild, so a
|
||||
tripped session never earned its probe."""
|
||||
"""Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable
|
||||
because the gateway rebuilds the compressor every turn / cache eviction."""
|
||||
return self._read_session_number("compression_recovery_deadline", session_id, float, 0.0)
|
||||
|
||||
def set_compression_recovery_deadline(self, session_id: str, deadline: float) -> None:
|
||||
@@ -516,39 +411,23 @@ class SessionCompressionMixin:
|
||||
normalized = max(0.0, float(deadline or 0.0))
|
||||
except (TypeError, ValueError):
|
||||
normalized = 0.0
|
||||
stored = normalized if normalized > 0.0 else None
|
||||
|
||||
self._write_sql(
|
||||
"UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?",
|
||||
(stored, session_id),
|
||||
(normalized if normalized > 0.0 else 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.
|
||||
|
||||
Ownership is decided by ``holder`` alone, deliberately NOT ``expires_at``:
|
||||
a live owner whose refresher stalled past its TTL (GC pause, loaded CI
|
||||
runner, slow write escaping ``_execute_write``'s retry budget) must be
|
||||
able to revive its still-unclaimed row. Requiring ``expires_at >= now``
|
||||
made such a stall permanent — every later refresh matched 0 rows and the
|
||||
owner kept compressing/rotating with no lease, exactly the window in
|
||||
which a competing path can fork the lineage.
|
||||
|
||||
It cannot resurrect a lock someone else took: SQLite serialises writes,
|
||||
so :meth:`try_acquire_compression_lock`'s reclaim (DELETE-expired +
|
||||
INSERT-or-IGNORE) never interleaves with this UPDATE. Reclaim-first
|
||||
replaces ``holder`` and this matches nothing; refresh-first pushes
|
||||
``expires_at`` forward and the reclaimer's DELETE matches nothing."""
|
||||
Ownership is decided by ``holder`` alone, deliberately NOT ``expires_at``: a live
|
||||
owner whose refresher stalled past its TTL must be able to revive its still-
|
||||
unclaimed row, otherwise it keeps compressing with no lease — the window in
|
||||
which a competing path can fork the lineage. It cannot resurrect a lock someone
|
||||
else took: SQLite serialises writes, so the reclaim (DELETE-expired + INSERT OR
|
||||
IGNORE) never interleaves with this UPDATE."""
|
||||
if not session_id or not holder:
|
||||
return False
|
||||
now = time.time()
|
||||
expires_at = now + ttl_seconds
|
||||
|
||||
expires_at = time.time() + ttl_seconds
|
||||
try:
|
||||
return self._write_rowcount(
|
||||
"UPDATE compression_locks SET expires_at = ? "
|
||||
@@ -556,27 +435,17 @@ class SessionCompressionMixin:
|
||||
(expires_at, session_id, holder),
|
||||
) > 0
|
||||
except sqlite3.Error as exc:
|
||||
logger.warning(
|
||||
"refresh_compression_lock(%s) failed: %s",
|
||||
session_id, exc,
|
||||
)
|
||||
logger.warning("refresh_compression_lock(%s) failed: %s", session_id, exc)
|
||||
return False
|
||||
|
||||
def try_acquire_compression_lock(
|
||||
self,
|
||||
session_id: str,
|
||||
holder: str,
|
||||
ttl_seconds: float = 300.0,
|
||||
) -> bool:
|
||||
def try_acquire_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool:
|
||||
"""Try to atomically acquire the compression lock for ``session_id``.
|
||||
|
||||
``True``: caller owns the lock and must :meth:`release_compression_lock`.
|
||||
``False``: another holder owns a live lock and the caller MUST NOT
|
||||
compress — its rotation would race the holder's and split the lineage.
|
||||
Expired locks and structured holders whose local ``pid=`` is dead are
|
||||
reclaimed transparently, so a gateway killed mid-compression doesn't
|
||||
stall its replacement for the full TTL. Single-transaction DELETE-expired
|
||||
+ INSERT-or-IGNORE + SELECT-to-confirm; SQLite serialises writes, so it's atomic."""
|
||||
``False``: another holder owns a live lock and the caller MUST NOT compress (its
|
||||
rotation would split the lineage). Expired locks and structured holders whose
|
||||
local ``pid=`` is dead are reclaimed transparently. Single-transaction DELETE-
|
||||
expired + INSERT OR IGNORE + SELECT-to-confirm (INSERT OR IGNORE gives no
|
||||
rowcount signal)."""
|
||||
from hermes_state import _compression_lock_holder_process_is_dead
|
||||
if not session_id:
|
||||
return False
|
||||
@@ -585,90 +454,54 @@ class SessionCompressionMixin:
|
||||
|
||||
def _do(conn):
|
||||
reclaimed_holder = None
|
||||
row = conn.execute(
|
||||
"SELECT holder, expires_at FROM compression_locks "
|
||||
"WHERE session_id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
row = conn.execute(_LOCK_ROW_SQL, (session_id,)).fetchone()
|
||||
if row is not None:
|
||||
current_holder = (
|
||||
row[0]
|
||||
)
|
||||
current_expires_at = (
|
||||
row[1]
|
||||
)
|
||||
if (
|
||||
current_expires_at < now
|
||||
or _compression_lock_holder_process_is_dead(current_holder)
|
||||
):
|
||||
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
|
||||
# INSERT OR IGNORE gives no rowcount signal — verify ownership via SELECT.
|
||||
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()
|
||||
acquired = row is not None and (
|
||||
row[0]
|
||||
) == holder
|
||||
return acquired, reclaimed_holder
|
||||
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:
|
||||
acquired, reclaimed_holder = self._execute_write(_do)
|
||||
if reclaimed_holder:
|
||||
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)
|
||||
except sqlite3.Error as exc:
|
||||
logger.warning(
|
||||
"try_acquire_compression_lock(%s) failed: %s",
|
||||
session_id, exc,
|
||||
)
|
||||
# False makes the caller skip compression — the safe behaviour
|
||||
# when the lock subsystem is broken.
|
||||
# False makes the caller skip compression — safe when the lock subsystem is broken.
|
||||
logger.warning("try_acquire_compression_lock(%s) failed: %s", session_id, exc)
|
||||
return False
|
||||
|
||||
def release_compression_lock(self, session_id: str, holder: str) -> None:
|
||||
"""Release the compression lock for ``session_id`` iff we own it. Idempotent
|
||||
when the lock is gone or reclaimed; the ``holder`` check stops a late
|
||||
compressor clobbering someone else's fresh lock."""
|
||||
"""Release the compression lock iff we own it; idempotent when gone/reclaimed."""
|
||||
if not session_id:
|
||||
return
|
||||
|
||||
try:
|
||||
self._write_sql(
|
||||
"DELETE FROM compression_locks "
|
||||
"WHERE session_id = ? AND holder = ?",
|
||||
(session_id, holder),
|
||||
)
|
||||
except sqlite3.Error as exc:
|
||||
logger.warning(
|
||||
"release_compression_lock(%s) failed: %s",
|
||||
session_id, exc,
|
||||
)
|
||||
self._write_sql_logged(
|
||||
"release_compression_lock", session_id,
|
||||
"DELETE FROM compression_locks "
|
||||
"WHERE session_id = ? AND holder = ?",
|
||||
(session_id, holder),
|
||||
)
|
||||
|
||||
def _session_turn_lease_key_on_conn(self, conn, session_id: str) -> str:
|
||||
"""Walk compression parents on ``conn`` to the conversation lease key.
|
||||
|
||||
Must share the connection of the lease INSERT/UPDATE/DELETE: a failed
|
||||
``get_session`` must not yield a child id the write then persists
|
||||
(refresh would walk to the parent and fail-close). Markers bind to
|
||||
``parent_session_id`` (as in ``_NON_CONTINUATION_CHILD_FILTER_SQL``).
|
||||
Lock errors propagate so ``_execute_write`` / ``acquire_session_turn_lease`` can retry."""
|
||||
Must share the connection of the lease INSERT/UPDATE/DELETE: a failed lookup
|
||||
must not yield a child id the write then persists. Markers bind to
|
||||
``parent_session_id``. Lock errors propagate so ``_execute_write`` can retry."""
|
||||
if not session_id:
|
||||
return session_id
|
||||
|
||||
@@ -684,11 +517,7 @@ class SessionCompressionMixin:
|
||||
seen = {session_id}
|
||||
while current:
|
||||
parent_id = current.get("parent_session_id")
|
||||
if (
|
||||
not parent_id
|
||||
or parent_id in seen
|
||||
or self._is_explicit_fork_child_row(current)
|
||||
):
|
||||
if not parent_id or parent_id in seen or self._is_explicit_fork_child_row(current):
|
||||
break
|
||||
parent = _row(parent_id)
|
||||
if not parent or parent.get("end_reason") != "compression":
|
||||
@@ -698,29 +527,19 @@ class SessionCompressionMixin:
|
||||
return str(current.get("id") or session_id) if current else session_id
|
||||
|
||||
def _session_turn_lease_key(self, session_id: str) -> str:
|
||||
"""Return the stable serialization key for every compression segment.
|
||||
|
||||
Acquire/refresh/release resolve this inside their write transaction; this
|
||||
is for tests/diagnostics. It does not swallow lock errors — a swallowed
|
||||
walk plus a later successful write was the fail-open that replayed the
|
||||
post-rotation refresh miss."""
|
||||
"""Stable serialization key for every compression segment (tests/diagnostics;
|
||||
the write paths resolve it inside their own txn). Does not swallow lock errors."""
|
||||
if not session_id:
|
||||
return session_id
|
||||
with self._read_ctx() as conn:
|
||||
return self._session_turn_lease_key_on_conn(conn, session_id)
|
||||
|
||||
def try_acquire_session_turn_lease(
|
||||
self,
|
||||
session_id: str,
|
||||
holder: str,
|
||||
*,
|
||||
ttl_seconds: float = 300.0,
|
||||
patience_s: Optional[float] = None,
|
||||
self, session_id: str, holder: str, *, ttl_seconds: float = 300.0, patience_s: Optional[float] = None,
|
||||
) -> bool:
|
||||
"""Atomically acquire the cross-process turn lease for a conversation.
|
||||
Compression rotates a session into child segments, so the durable key is
|
||||
the lineage root, not the current segment id. The walk, the INSERT, and
|
||||
reclaim of expired or dead-local-PID leases share one write transaction."""
|
||||
"""Atomically acquire the cross-process turn lease for a conversation (keyed by
|
||||
the lineage root). The walk, the INSERT, and reclaim of expired or dead-local-PID
|
||||
leases share one write transaction."""
|
||||
from hermes_state import _compression_lock_holder_process_is_dead
|
||||
if not session_id or not holder:
|
||||
return False
|
||||
@@ -736,10 +555,7 @@ class SessionCompressionMixin:
|
||||
).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)
|
||||
):
|
||||
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 = ?",
|
||||
@@ -752,8 +568,7 @@ class SessionCompressionMixin:
|
||||
(conversation_id, holder, now, expires_at),
|
||||
)
|
||||
owner = conn.execute(
|
||||
"SELECT holder FROM session_turn_leases WHERE conversation_id = ?",
|
||||
(conversation_id,),
|
||||
"SELECT holder FROM session_turn_leases WHERE conversation_id = ?", (conversation_id,),
|
||||
).fetchone()
|
||||
return owner is not None and owner["holder"] == holder
|
||||
|
||||
@@ -774,10 +589,9 @@ class SessionCompressionMixin:
|
||||
) -> bool:
|
||||
"""Wait for a cross-process turn lease without holding a SQLite lock.
|
||||
|
||||
``on_wait(elapsed)`` is best-effort: called when the first attempt fails
|
||||
(elapsed ~0) and about every ``wait_notice_interval_seconds`` after, so
|
||||
UIs can show another process holds the conversation. ``should_abort()``
|
||||
True (e.g. ``/stop``) returns False at once, not after ``wait_seconds``."""
|
||||
``on_wait(elapsed)`` is best-effort: called when the first attempt fails and
|
||||
about every ``wait_notice_interval_seconds`` after. ``should_abort()`` True
|
||||
(e.g. ``/stop``) returns False at once."""
|
||||
from hermes_state import classify_persistence_error
|
||||
deadline = time.monotonic() + max(0.0, float(wait_seconds))
|
||||
wait_started = None
|
||||
@@ -789,22 +603,15 @@ class SessionCompressionMixin:
|
||||
if should_abort():
|
||||
return False
|
||||
except Exception:
|
||||
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:
|
||||
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
|
||||
except sqlite3.Error as exc:
|
||||
# Long holder transactions (compression publish, large flushes)
|
||||
# can exhaust one write-patience budget; keep polling until
|
||||
# wait_seconds or should_abort.
|
||||
# Long holder transactions can exhaust one write-patience budget; keep
|
||||
# polling until wait_seconds or should_abort.
|
||||
if classify_persistence_error(exc) != "locked":
|
||||
raise
|
||||
now = time.monotonic()
|
||||
@@ -814,27 +621,16 @@ class SessionCompressionMixin:
|
||||
if wait_started is None:
|
||||
wait_started = now
|
||||
if on_wait is not None and (
|
||||
last_notice_at is None
|
||||
or notice_every == 0.0
|
||||
or (now - last_notice_at) >= notice_every
|
||||
last_notice_at is None or notice_every == 0.0 or (now - last_notice_at) >= notice_every
|
||||
):
|
||||
try:
|
||||
on_wait(max(0.0, now - wait_started))
|
||||
except Exception:
|
||||
logger.debug(
|
||||
"session turn lease on_wait callback failed",
|
||||
exc_info=True,
|
||||
)
|
||||
logger.debug("session turn lease on_wait callback failed", exc_info=True)
|
||||
last_notice_at = now
|
||||
time.sleep(min(max(0.01, float(poll_interval_seconds)), remaining))
|
||||
|
||||
def refresh_session_turn_lease(
|
||||
self,
|
||||
session_id: str,
|
||||
holder: str,
|
||||
*,
|
||||
ttl_seconds: float = 300.0,
|
||||
) -> bool:
|
||||
def refresh_session_turn_lease(self, session_id: str, holder: str, *, ttl_seconds: float = 300.0) -> bool:
|
||||
"""Extend a turn lease only while ``holder`` still owns it."""
|
||||
if not session_id or not holder:
|
||||
return False
|
||||
@@ -867,24 +663,20 @@ class SessionCompressionMixin:
|
||||
self._execute_write(_do)
|
||||
|
||||
def get_compression_lock_holder(self, session_id: str) -> Optional[str]:
|
||||
"""Return the current (non-expired) holder for ``session_id``, or None.
|
||||
Diagnostic only — not part of the locking protocol."""
|
||||
"""Current (non-expired) holder for ``session_id``, or None. Diagnostic only."""
|
||||
if not session_id:
|
||||
return None
|
||||
now = time.time()
|
||||
row = self._read_one(
|
||||
"SELECT holder FROM compression_locks "
|
||||
"WHERE session_id = ? AND expires_at >= ?",
|
||||
(session_id, now),
|
||||
(session_id, time.time()),
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
return row[0]
|
||||
return None if row is None else row[0]
|
||||
|
||||
def finalize_orphaned_compression_sessions(self) -> int:
|
||||
"""Mark orphaned compression continuations (parent ended by compression;
|
||||
child has messages, no end_reason/ended_at, api_call_count=0) as
|
||||
``orphaned_compression``. Non-destructive: messages are preserved."""
|
||||
"""Mark orphaned compression continuations (parent ended by compression; child
|
||||
has messages, no end_reason/ended_at, api_call_count=0, older than 7 days) as
|
||||
``orphaned_compression``. Non-destructive."""
|
||||
cutoff = time.time() - 604800 # 7 days
|
||||
|
||||
def _do(conn):
|
||||
@@ -917,29 +709,22 @@ class SessionCompressionMixin:
|
||||
return self._execute_write(_do) or 0
|
||||
|
||||
def get_compression_chain(self, session_id: str) -> List[str]:
|
||||
"""Walk the compression-continuation chain forward and return every id.
|
||||
"""Walk the compression-continuation chain forward: root-first through the tip
|
||||
(``[session_id]`` when no continuation). ``get_compression_tip`` is this walk's
|
||||
last element.
|
||||
|
||||
Root-first, ending at the tip; ``[session_id]`` when no continuation
|
||||
exists. ``get_compression_tip`` is this walk's last element — one
|
||||
implementation so the two can never disagree.
|
||||
|
||||
A continuation is a child of a session with ``end_reason='compression'``.
|
||||
Older builds also required ``child.started_at >= parent.ended_at``;
|
||||
too brittle — gateway + compression races can insert the real
|
||||
continuation before the parent's ``ended_at`` is written while a stale
|
||||
websocket later creates a sibling that passes the timestamp test, so
|
||||
desktop resume followed the sibling and recent messages looked "lost".
|
||||
Instead: follow only children of compression-ended parents, exclude
|
||||
explicit branch/delegate/tool children, and prefer children that continue
|
||||
the chain (``end_reason='compression'``) or are still live over stale
|
||||
closed siblings such as ``ws_orphan_reap``."""
|
||||
A continuation is a child of a session with ``end_reason='compression'``. The
|
||||
old ``child.started_at >= parent.ended_at`` test was too brittle (gateway +
|
||||
compression races insert the real continuation before ``ended_at`` is written,
|
||||
while a stale websocket later creates a sibling that passes it). Instead exclude
|
||||
branch/delegate/tool children and prefer children that continue the chain or
|
||||
are still live over stale closed siblings such as ``ws_orphan_reap``."""
|
||||
current = session_id
|
||||
chain = [current] if current else []
|
||||
seen = {current} if current else set()
|
||||
# Defensive bound; chains this deep are pathological.
|
||||
for _ in range(100):
|
||||
for _ in range(100): # defensive bound; chains this deep are pathological
|
||||
with self._read_ctx() as conn:
|
||||
cursor = conn.execute(
|
||||
row = conn.execute(
|
||||
f"""
|
||||
SELECT child.id
|
||||
FROM sessions parent
|
||||
@@ -961,8 +746,7 @@ class SessionCompressionMixin:
|
||||
LIMIT 1
|
||||
""",
|
||||
(current,),
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return chain
|
||||
child_id = row["id"]
|
||||
@@ -974,8 +758,8 @@ class SessionCompressionMixin:
|
||||
return chain
|
||||
|
||||
def get_compression_tip(self, session_id: str) -> Optional[str]:
|
||||
"""Live tip of a compression chain (walk semantics: ``get_compression_chain``);
|
||||
the input id when no continuation exists."""
|
||||
"""Live tip of a compression chain (``get_compression_chain`` semantics); the
|
||||
input id when no continuation exists."""
|
||||
chain = self.get_compression_chain(session_id)
|
||||
return chain[-1] if chain else session_id
|
||||
|
||||
@@ -991,7 +775,6 @@ class SessionCompressionMixin:
|
||||
session = self.get_session(session_id)
|
||||
if not session or self._is_explicit_fork_child_row(session):
|
||||
return [session_id] if session else []
|
||||
|
||||
root = session
|
||||
ancestors = {root["id"]}
|
||||
while self._is_compression_child_row(root):
|
||||
@@ -1000,7 +783,6 @@ class SessionCompressionMixin:
|
||||
break
|
||||
root = parent
|
||||
ancestors.add(root["id"])
|
||||
|
||||
lineage = [root["id"]]
|
||||
seen = {root["id"]}
|
||||
current = root
|
||||
@@ -1013,18 +795,11 @@ class SessionCompressionMixin:
|
||||
""",
|
||||
(current["id"],),
|
||||
)
|
||||
next_child = None
|
||||
for row in rows:
|
||||
candidate = dict(row)
|
||||
if self._is_compression_child_row(candidate):
|
||||
next_child = candidate
|
||||
break
|
||||
next_child = next((dict(row) for row in rows if self._is_compression_child_row(dict(row))), None)
|
||||
if not next_child or next_child["id"] in seen:
|
||||
break
|
||||
lineage.append(next_child["id"])
|
||||
seen.add(next_child["id"])
|
||||
current = next_child
|
||||
if current["id"] == session_id:
|
||||
# Later tips are included only when the requested session itself was compacted.
|
||||
continue
|
||||
# Later tips are included only when the requested session itself was compacted.
|
||||
return lineage if session_id in lineage else [session_id]
|
||||
|
||||
+739
-1680
File diff suppressed because it is too large
Load Diff
+80
-169
@@ -13,66 +13,44 @@ from hermes_state_common import _COMPRESSION_CHILD_SQL, escape_like as _escape_l
|
||||
# caplog tests pin the "hermes_state" logger name.
|
||||
logger = logging.getLogger("hermes_state")
|
||||
|
||||
# ASCII controls (keeping \t \n \r for the whitespace collapse), then zero-width,
|
||||
# bidi override, object-replacement and interlinear-annotation code points.
|
||||
_TITLE_CONTROL_RE = re.compile(r'[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]')
|
||||
_TITLE_INVISIBLE_RE = re.compile(r'[\u200b-\u200f\u2028-\u202e\u2060-\u2069\ufeff\ufffc\ufff9-\ufffb]')
|
||||
_NUMBERED_TITLE_RE = re.compile(r'^(.*?) #(\d+)$')
|
||||
|
||||
|
||||
class SessionTitlesMixin:
|
||||
"""Sanitizing, ranking auto/user titles, lineage-aware lookups."""
|
||||
|
||||
@classmethod
|
||||
def _title_rank(cls, source: Optional[str]) -> int:
|
||||
"""Rank a stored title_source.
|
||||
|
||||
NULL (pre-provenance rows) is indistinguishable from a manual ``/title``
|
||||
of that era, so it ranks as ``user``: auto-titling only ever fills
|
||||
genuinely empty legacy titles.
|
||||
"""
|
||||
"""Rank a stored title_source. NULL (pre-provenance rows) is indistinguishable
|
||||
from a manual ``/title`` of that era, so it ranks as ``user``."""
|
||||
if source is None:
|
||||
return cls._TITLE_SOURCE_RANK[cls.TITLE_SOURCE_USER]
|
||||
return cls._TITLE_SOURCE_RANK.get(str(source), 0)
|
||||
|
||||
@staticmethod
|
||||
def sanitize_title(title: Optional[str]) -> Optional[str]:
|
||||
"""Strip control/zero-width/bidi chars, collapse whitespace, normalize
|
||||
empty to None. Raises ValueError if longer than MAX_TITLE_LENGTH
|
||||
after cleaning."""
|
||||
"""Strip control/zero-width/bidi chars (and lone surrogates sqlite3 cannot
|
||||
bind), collapse whitespace, normalize empty to None. Raises ValueError if
|
||||
longer than MAX_TITLE_LENGTH after cleaning."""
|
||||
from hermes_state import SessionDB
|
||||
if not title:
|
||||
return None
|
||||
|
||||
# Lone surrogates cannot be bound by sqlite3 (UnicodeEncodeError).
|
||||
title = _sanitize_surrogates(title)
|
||||
|
||||
# ASCII controls, keeping \t \n \r so the whitespace collapse below
|
||||
# turns them into spaces.
|
||||
cleaned = re.sub(r'[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]', '', title)
|
||||
|
||||
# Zero-width, bidi override, object-replacement, interlinear annotation.
|
||||
cleaned = re.sub(
|
||||
r'[\u200b-\u200f\u2028-\u202e\u2060-\u2069\ufeff\ufffc\ufff9-\ufffb]',
|
||||
'', cleaned,
|
||||
)
|
||||
|
||||
cleaned = _TITLE_INVISIBLE_RE.sub('', _TITLE_CONTROL_RE.sub('', _sanitize_surrogates(title)))
|
||||
cleaned = re.sub(r'\s+', ' ', cleaned).strip()
|
||||
|
||||
if not cleaned:
|
||||
return None
|
||||
|
||||
if len(cleaned) > SessionDB.MAX_TITLE_LENGTH:
|
||||
raise ValueError(
|
||||
f"Title too long ({len(cleaned)} chars, max {SessionDB.MAX_TITLE_LENGTH})"
|
||||
)
|
||||
|
||||
raise ValueError(f"Title too long ({len(cleaned)} chars, max {SessionDB.MAX_TITLE_LENGTH})")
|
||||
return cleaned
|
||||
|
||||
def _is_compression_ancestor(
|
||||
self, conn, *, ancestor_id: str, descendant_id: str
|
||||
) -> bool:
|
||||
"""True if *ancestor_id* is a compression predecessor of *descendant_id*.
|
||||
|
||||
Uses the canonical continuation edge ``_COMPRESSION_CHILD_SQL`` (parent
|
||||
ended with ``end_reason = 'compression'`` and child started at/after its
|
||||
``ended_at``), which excludes delegate/branch children that also carry
|
||||
``parent_session_id``. One recursive CTE so the edge is defined once.
|
||||
"""
|
||||
def _is_compression_ancestor(self, conn, *, ancestor_id: str, descendant_id: str) -> bool:
|
||||
"""True if *ancestor_id* is a compression predecessor of *descendant_id*, via the
|
||||
canonical continuation edge ``_COMPRESSION_CHILD_SQL`` (excludes delegate/branch
|
||||
children that also carry ``parent_session_id``)."""
|
||||
if not ancestor_id or not descendant_id or ancestor_id == descendant_id:
|
||||
return False
|
||||
edge = _COMPRESSION_CHILD_SQL.format(a="child")
|
||||
@@ -93,23 +71,15 @@ class SessionTitlesMixin:
|
||||
).fetchone()
|
||||
return row is not None
|
||||
|
||||
def _set_session_title(
|
||||
self,
|
||||
session_id: str,
|
||||
title: str,
|
||||
*,
|
||||
source: str,
|
||||
) -> bool:
|
||||
def _set_session_title(self, session_id: str, title: str, *, source: str) -> bool:
|
||||
"""Write a title, enforcing provenance precedence.
|
||||
|
||||
A ``user`` write always lands. ``derived``/``llm`` land only when the
|
||||
row is untitled or holds strictly lower authority, so derived upgrades
|
||||
to llm exactly once, nothing overwrites a user name, and re-running the
|
||||
titler on an llm row is a no-op (stops sessions renaming themselves).
|
||||
No writer may move a hidden canonical Bot Chat off its title.
|
||||
|
||||
Read and write are one compare-and-swap in a single transaction, so a
|
||||
manual ``/title`` racing an in-flight generation is not clobbered.
|
||||
A ``user`` write always lands. ``derived``/``llm`` land only when the row is
|
||||
untitled or holds strictly lower authority (derived upgrades to llm exactly once,
|
||||
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
|
||||
one compare-and-swap transaction, so a manual ``/title`` racing an in-flight
|
||||
generation is not clobbered.
|
||||
"""
|
||||
title = self.sanitize_title(title)
|
||||
is_user = source == self.TITLE_SOURCE_USER
|
||||
@@ -117,18 +87,14 @@ class SessionTitlesMixin:
|
||||
|
||||
def _do(conn):
|
||||
current = conn.execute(
|
||||
"SELECT title, title_source, hidden FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
"SELECT title, title_source, hidden FROM sessions WHERE id = ?", (session_id,),
|
||||
).fetchone()
|
||||
if current is None:
|
||||
return 0
|
||||
# The canonical Bot Chat's NAME is its identity: Bot Mode resolves it
|
||||
# by exact-title lookup on every open, so a rename orphans the whole
|
||||
# conversation (next open mints an empty replacement and UNIQUE(title)
|
||||
# blocks renaming back). Refuse here, the single write path every
|
||||
# surface funnels through. Hidden is the discriminator: canonical
|
||||
# chats are born hidden; a visible session merely named "Bot Chat"
|
||||
# stays renameable. Provenance-blind so the auto-titler no-ops too.
|
||||
# The canonical Bot Chat's NAME is its identity (Bot Mode resolves it by
|
||||
# exact-title lookup on every open), so a rename orphans the conversation.
|
||||
# Hidden is the discriminator: canonical chats are born hidden; a visible
|
||||
# session merely named "Bot Chat" stays renameable. Provenance-blind.
|
||||
if (
|
||||
(current["title"] or "") == self.CANONICAL_BOT_CHAT_TITLE
|
||||
and bool(current["hidden"])
|
||||
@@ -141,100 +107,64 @@ class SessionTitlesMixin:
|
||||
"To start fresh, create a new bot instead."
|
||||
)
|
||||
return 0
|
||||
if not is_user and current["title"] is not None:
|
||||
if self._title_rank(current["title_source"]) >= new_rank:
|
||||
return 0
|
||||
|
||||
if not is_user and current["title"] is not None and self._title_rank(current["title_source"]) >= new_rank:
|
||||
return 0
|
||||
if title:
|
||||
cursor = conn.execute(
|
||||
"SELECT id FROM sessions WHERE title = ? AND id != ?",
|
||||
(title, session_id),
|
||||
)
|
||||
conflict = cursor.fetchone()
|
||||
conflict = conn.execute(
|
||||
"SELECT id FROM sessions WHERE title = ? AND id != ?", (title, session_id),
|
||||
).fetchone()
|
||||
if conflict:
|
||||
conflict_id = conflict["id"]
|
||||
# If the conflicting holder is a hidden compressed ancestor
|
||||
# of this continuation, the user cannot free the title, so
|
||||
# transfer it onto the tip. Uniqueness and lineage are kept.
|
||||
if self._is_compression_ancestor(
|
||||
conn, ancestor_id=conflict_id, descendant_id=session_id
|
||||
):
|
||||
conn.execute(
|
||||
"UPDATE sessions SET title = NULL WHERE id = ?",
|
||||
(conflict_id,),
|
||||
)
|
||||
# A hidden compressed ancestor holding the title cannot be freed by
|
||||
# the user, so transfer it onto the tip (uniqueness + lineage kept).
|
||||
if self._is_compression_ancestor(conn, ancestor_id=conflict_id, descendant_id=session_id):
|
||||
conn.execute("UPDATE sessions SET title = NULL WHERE id = ?", (conflict_id,))
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Title '{title}' is already in use by session {conflict_id}"
|
||||
)
|
||||
# CAS on the values just read (``IS`` is NULL-safe): a concurrent
|
||||
# write between the SELECT and here loses instead of being overwritten.
|
||||
raise ValueError(f"Title '{title}' is already in use by session {conflict_id}")
|
||||
# CAS on the values just read (``IS`` is NULL-safe): a concurrent write
|
||||
# between the SELECT and here loses instead of being overwritten.
|
||||
cursor = conn.execute(
|
||||
"UPDATE sessions SET title = ?, title_source = ? "
|
||||
"WHERE id = ? AND title IS ? AND title_source IS ?",
|
||||
(
|
||||
title,
|
||||
source if title else None,
|
||||
session_id,
|
||||
current["title"],
|
||||
current["title_source"],
|
||||
),
|
||||
(title, source if title else None, session_id, current["title"], current["title_source"]),
|
||||
)
|
||||
return cursor.rowcount
|
||||
|
||||
rowcount = self._execute_write(_do)
|
||||
return rowcount > 0
|
||||
return self._execute_write(_do) > 0
|
||||
|
||||
def set_session_title(self, session_id: str, title: str) -> bool:
|
||||
"""Set a title on the user's behalf (``user`` provenance; auto-titling
|
||||
never replaces it). Empty clears the title. Raises ValueError on a
|
||||
title conflict or validation failure. Automatic callers must use
|
||||
:meth:`set_auto_title`."""
|
||||
return self._set_session_title(
|
||||
session_id, title, source=self.TITLE_SOURCE_USER
|
||||
)
|
||||
"""Set a title on the user's behalf (``user`` provenance). Empty clears it.
|
||||
Raises ValueError on conflict or validation failure."""
|
||||
return self._set_session_title(session_id, title, source=self.TITLE_SOURCE_USER)
|
||||
|
||||
def set_auto_title(self, session_id: str, title: str, *, source: str) -> bool:
|
||||
"""Set an automatic title; False (untouched) when a higher-authority
|
||||
title already holds the row."""
|
||||
"""Set an automatic title; False (untouched) when a higher-authority title
|
||||
already holds the row."""
|
||||
if source not in (self.TITLE_SOURCE_DERIVED, self.TITLE_SOURCE_LLM):
|
||||
raise ValueError(f"invalid automatic title source: {source!r}")
|
||||
return self._set_session_title(session_id, title, source=source)
|
||||
|
||||
def set_auto_title_if_empty(self, session_id: str, title: str) -> bool:
|
||||
"""Back-compat shim (third-party plugins reference it by name); new
|
||||
code calls :meth:`set_auto_title` with an explicit source."""
|
||||
return self.set_auto_title(
|
||||
session_id, title, source=self.TITLE_SOURCE_LLM
|
||||
)
|
||||
"""Back-compat shim (third-party plugins reference it by name)."""
|
||||
return self.set_auto_title(session_id, title, source=self.TITLE_SOURCE_LLM)
|
||||
|
||||
def get_session_title(self, session_id: str) -> Optional[str]:
|
||||
"""Get the title for a session, or None."""
|
||||
with self._read_ctx() as conn:
|
||||
cursor = conn.execute(
|
||||
"SELECT title FROM sessions WHERE id = ?", (session_id,)
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
row = self._read_one("SELECT title FROM sessions WHERE id = ?", (session_id,))
|
||||
return row["title"] if row else None
|
||||
|
||||
def get_session_title_source(self, session_id: str) -> Optional[str]:
|
||||
"""Get the provenance of a session's title, or None when untitled."""
|
||||
with self._read_ctx() as conn:
|
||||
cursor = conn.execute(
|
||||
"SELECT title, title_source FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
row = self._read_one("SELECT title, title_source FROM sessions WHERE id = ?", (session_id,))
|
||||
if not row or row["title"] is None:
|
||||
return None
|
||||
return row["title_source"]
|
||||
|
||||
def set_session_title_source(self, session_id: str, source: str) -> bool:
|
||||
"""Overwrite a title's provenance without touching the text: a title
|
||||
copied across a compression rotation keeps the original's authority."""
|
||||
"""Overwrite a title's provenance without touching the text (a title copied
|
||||
across a compression rotation keeps the original's authority)."""
|
||||
if source not in self._TITLE_SOURCE_RANK:
|
||||
raise ValueError(f"invalid title source: {source!r}")
|
||||
|
||||
return self._write_rowcount(
|
||||
"UPDATE sessions SET title_source = ? "
|
||||
"WHERE id = ? AND title IS NOT NULL",
|
||||
@@ -243,63 +173,44 @@ class SessionTitlesMixin:
|
||||
|
||||
def get_session_by_title(self, title: str) -> Optional[Dict[str, Any]]:
|
||||
"""Look up a session by exact title. Returns session dict or None."""
|
||||
with self._read_ctx() as conn:
|
||||
cursor = conn.execute(
|
||||
"SELECT s.*, "
|
||||
"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 "
|
||||
"WHERE s.title = ?",
|
||||
(title,),
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
row = self._read_one(
|
||||
"SELECT s.*, "
|
||||
"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 "
|
||||
"WHERE s.title = ?",
|
||||
(title,),
|
||||
)
|
||||
return self._session_row_dict(row) if row else None
|
||||
|
||||
def resolve_session_by_title(self, title: str) -> Optional[str]:
|
||||
"""Resolve a title to a session ID, preferring the latest "title #N"
|
||||
continuation over the exact match."""
|
||||
exact = self.get_session_by_title(title)
|
||||
|
||||
# Escape LIKE wildcards so "%"/"_" in titles cannot false-match.
|
||||
escaped = _escape_like(title)
|
||||
with self._read_ctx() as conn:
|
||||
cursor = conn.execute(
|
||||
"SELECT id, title, started_at FROM sessions "
|
||||
"WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC",
|
||||
(f"{escaped} #%",),
|
||||
)
|
||||
numbered = cursor.fetchall()
|
||||
|
||||
numbered = self._read_all(
|
||||
"SELECT id, title, started_at FROM sessions "
|
||||
"WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC",
|
||||
(f"{_escape_like(title)} #%",),
|
||||
)
|
||||
if numbered:
|
||||
return numbered[0]["id"]
|
||||
elif exact:
|
||||
return exact["id"]
|
||||
return None
|
||||
return exact["id"] if exact else None
|
||||
|
||||
def get_next_title_in_lineage(self, base_title: str) -> str:
|
||||
"""Next title in a lineage ("my session" → "my session #2"): strip any
|
||||
" #N" suffix, then increment the highest existing number."""
|
||||
match = re.match(r'^(.*?) #(\d+)$', base_title)
|
||||
if match:
|
||||
base = match.group(1)
|
||||
else:
|
||||
base = base_title
|
||||
|
||||
escaped = _escape_like(base)
|
||||
with self._read_ctx() as conn:
|
||||
cursor = conn.execute(
|
||||
"SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'",
|
||||
(base, f"{escaped} #%"),
|
||||
)
|
||||
existing = [row["title"] for row in cursor.fetchall()]
|
||||
|
||||
if not existing:
|
||||
"""Next title in a lineage ("my session" -> "my session #2"): strip any " #N"
|
||||
suffix, then increment the highest existing number."""
|
||||
match = _NUMBERED_TITLE_RE.match(base_title)
|
||||
base = match.group(1) if match else base_title
|
||||
rows = self._read_all(
|
||||
"SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'",
|
||||
(base, f"{_escape_like(base)} #%"),
|
||||
)
|
||||
if not rows:
|
||||
return base
|
||||
|
||||
max_num = 1 # The unnumbered original counts as #1
|
||||
for t in existing:
|
||||
m = re.match(r'^.* #(\d+)$', t)
|
||||
max_num = 1 # the unnumbered original counts as #1
|
||||
for row in rows:
|
||||
m = re.match(r'^.* #(\d+)$', row["title"])
|
||||
if m:
|
||||
max_num = max(max_num, int(m.group(1)))
|
||||
|
||||
return f"{base} #{max_num + 1}"
|
||||
|
||||
+166
-292
@@ -14,24 +14,79 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
# caplog tests pin the "hermes_state" logger name.
|
||||
logger = logging.getLogger("hermes_state")
|
||||
|
||||
_TOKEN_UPDATE_ABSOLUTE_SQL = """UPDATE sessions SET
|
||||
input_tokens = ?,
|
||||
output_tokens = ?,
|
||||
cache_read_tokens = ?,
|
||||
cache_write_tokens = ?,
|
||||
reasoning_tokens = ?,
|
||||
estimated_cost_usd = COALESCE(?, 0),
|
||||
actual_cost_usd = CASE
|
||||
WHEN ? IS NULL THEN actual_cost_usd
|
||||
ELSE ?
|
||||
END,
|
||||
cost_status = COALESCE(?, cost_status),
|
||||
cost_source = COALESCE(?, cost_source),
|
||||
pricing_version = COALESCE(?, pricing_version),
|
||||
billing_provider = COALESCE(billing_provider, ?),
|
||||
billing_base_url = COALESCE(billing_base_url, ?),
|
||||
billing_mode = COALESCE(billing_mode, ?),
|
||||
model = COALESCE(model, ?),
|
||||
api_call_count = ?
|
||||
WHERE id = ?"""
|
||||
|
||||
_TOKEN_UPDATE_DELTA_SQL = """UPDATE sessions SET
|
||||
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 = COALESCE(estimated_cost_usd, 0) + COALESCE(?, 0),
|
||||
actual_cost_usd = CASE
|
||||
WHEN ? IS NULL THEN actual_cost_usd
|
||||
ELSE COALESCE(actual_cost_usd, 0) + ?
|
||||
END,
|
||||
cost_status = COALESCE(?, cost_status),
|
||||
cost_source = COALESCE(?, cost_source),
|
||||
pricing_version = COALESCE(?, pricing_version),
|
||||
billing_provider = COALESCE(billing_provider, ?),
|
||||
billing_base_url = COALESCE(billing_base_url, ?),
|
||||
billing_mode = COALESCE(billing_mode, ?),
|
||||
model = COALESCE(model, ?),
|
||||
api_call_count = COALESCE(api_call_count, 0) + ?
|
||||
WHERE id = ?"""
|
||||
|
||||
_MODEL_USAGE_UPSERT_SQL = """INSERT INTO session_model_usage (
|
||||
session_id, model, billing_provider, billing_base_url, billing_mode,
|
||||
task, api_call_count, input_tokens, output_tokens,
|
||||
cache_read_tokens, cache_write_tokens, reasoning_tokens,
|
||||
estimated_cost_usd, actual_cost_usd, cost_status, cost_source,
|
||||
first_seen, last_seen
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(session_id, model, billing_provider, billing_base_url, billing_mode, task)
|
||||
DO UPDATE SET
|
||||
api_call_count = api_call_count + excluded.api_call_count,
|
||||
input_tokens = input_tokens + excluded.input_tokens,
|
||||
output_tokens = output_tokens + excluded.output_tokens,
|
||||
cache_read_tokens = cache_read_tokens + excluded.cache_read_tokens,
|
||||
cache_write_tokens = cache_write_tokens + excluded.cache_write_tokens,
|
||||
reasoning_tokens = reasoning_tokens + excluded.reasoning_tokens,
|
||||
estimated_cost_usd = estimated_cost_usd + excluded.estimated_cost_usd,
|
||||
actual_cost_usd = actual_cost_usd + excluded.actual_cost_usd,
|
||||
cost_status = COALESCE(excluded.cost_status, cost_status),
|
||||
cost_source = COALESCE(excluded.cost_source, cost_source),
|
||||
last_seen = excluded.last_seen"""
|
||||
|
||||
|
||||
class SessionUsageMixin:
|
||||
"""Coalesced token writer, per-model usage rows, billing route."""
|
||||
|
||||
def update_session_billing_route(
|
||||
self,
|
||||
session_id: str,
|
||||
*,
|
||||
provider: str,
|
||||
base_url: str,
|
||||
billing_mode: Optional[str] = None,
|
||||
self, session_id: str, *, provider: str, base_url: str, billing_mode: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Unconditionally set the billing route (``update_token_counts`` only
|
||||
COALESCE-fills NULLs) so the dashboard reflects the latest /model switch.
|
||||
|
||||
Also nulls ``system_prompt`` so the cached snapshot (stale ``Model:`` /
|
||||
``Provider:`` header) is rebuilt, like ``update_session_model``.
|
||||
"""
|
||||
COALESCE-fills NULLs) so the dashboard reflects the latest /model switch. Also
|
||||
nulls ``system_prompt`` so the cached snapshot header is rebuilt."""
|
||||
# Barrier against queued token deltas — see update_session_model.
|
||||
self.flush_token_counts()
|
||||
|
||||
@@ -50,29 +105,20 @@ class SessionUsageMixin:
|
||||
self._execute_write(_do)
|
||||
|
||||
def queue_token_counts(self, session_id: str, **kwargs) -> None:
|
||||
"""Enqueue a token/cost delta for the background writer.
|
||||
|
||||
Same kwargs and semantics as :meth:`update_token_counts`, applied
|
||||
asynchronously; cheap enough for the turn thread. After close() has
|
||||
stopped the writer, falls back to the synchronous path and may raise.
|
||||
"""
|
||||
"""Enqueue a token/cost delta for the background writer (same kwargs as
|
||||
:meth:`update_token_counts`). After close() has stopped the writer, falls back
|
||||
to the synchronous path and may raise."""
|
||||
with self._token_queue_cond:
|
||||
thread = self._token_writer_thread
|
||||
writer_stopped = self._token_writer_stop and (
|
||||
thread is None or not thread.is_alive()
|
||||
)
|
||||
writer_stopped = self._token_writer_stop and (thread is None or not thread.is_alive())
|
||||
if not writer_stopped:
|
||||
self._token_queue.append((session_id, kwargs))
|
||||
if thread is None or not thread.is_alive():
|
||||
# Daemon so exit never hangs on accounting; the atexit hook
|
||||
# (registered once per instance) drains leftovers. Checking
|
||||
# ``not is_alive()`` rather than ``is None`` respawns a writer
|
||||
# that died from an unexpected escape, otherwise deltas
|
||||
# would pile up until a reader's flush drained them.
|
||||
# Daemon so exit never hangs on accounting; the atexit hook drains
|
||||
# leftovers. ``not is_alive()`` (not ``is None``) respawns a writer
|
||||
# that died from an unexpected escape.
|
||||
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
|
||||
thread.start()
|
||||
@@ -88,18 +134,21 @@ class SessionUsageMixin:
|
||||
atexit.register(_drain_at_exit)
|
||||
self._token_queue_cond.notify_all()
|
||||
if writer_stopped:
|
||||
# close() ran (a stop-flagged but live writer still accepts; its
|
||||
# loop drains before exiting). Enqueueing now would drop the delta
|
||||
# silently — no writer, atexit hook gone — so apply inline and let a
|
||||
# closed-connection failure raise at the call site.
|
||||
# close() ran: enqueueing would drop the delta silently, so apply inline.
|
||||
self.update_token_counts(session_id, **kwargs)
|
||||
|
||||
def flush_token_counts(self, timeout: float = 5.0) -> bool:
|
||||
"""Block until every queued token delta has been applied.
|
||||
def _apply_claimed_batch(self, batch) -> None:
|
||||
"""Apply a batch whose ``busy`` flag the caller already claimed, then release."""
|
||||
try:
|
||||
self._apply_token_batch(batch)
|
||||
finally:
|
||||
with self._token_queue_cond:
|
||||
self._token_writer_busy = False
|
||||
self._token_queue_cond.notify_all()
|
||||
|
||||
False on timeout (callers then read totals stale by the queued deltas).
|
||||
Never raises: apply failures are logged by the writer.
|
||||
"""
|
||||
def flush_token_counts(self, timeout: float = 5.0) -> bool:
|
||||
"""Block until every queued token delta has been applied. False on timeout
|
||||
(callers then read totals stale by the queued deltas). Never raises."""
|
||||
# Lock-free fast path: reads queue-then-busy (see ordering notes below).
|
||||
if not self._token_queue and not self._token_writer_busy:
|
||||
return True
|
||||
@@ -107,20 +156,13 @@ class SessionUsageMixin:
|
||||
with self._token_queue_cond:
|
||||
deadline = time.monotonic() + timeout
|
||||
while self._token_queue or self._token_writer_busy:
|
||||
# A live writer is authoritative even when stop-flagged: draining
|
||||
# here would race its in-flight batch, and newer deltas committing
|
||||
# before older ones breaks last-non-None-wins / first-accounted-
|
||||
# route / COALESCE-backfill fields. Only a dead writer lets the
|
||||
# caller take leftovers; re-checked each wakeup because the writer
|
||||
# can exit mid-wait with deltas enqueued after its final check.
|
||||
# busy is claimed while draining so a concurrent flush cannot
|
||||
# report drained or pop a newer delta while this batch is
|
||||
# unapplied: a claimed busy means "wait", never "drain alongside".
|
||||
# A live writer is authoritative even when stop-flagged: draining here
|
||||
# would race its in-flight batch and reorder deltas (breaking last-non-
|
||||
# None-wins / first-accounted-route / COALESCE-backfill fields). Only a
|
||||
# dead writer lets the caller take leftovers; a claimed busy means
|
||||
# "wait", never "drain alongside".
|
||||
thread = self._token_writer_thread
|
||||
if (
|
||||
(thread is None or not thread.is_alive())
|
||||
and not self._token_writer_busy
|
||||
):
|
||||
if (thread is None or not thread.is_alive()) and not self._token_writer_busy:
|
||||
self._token_writer_busy = True
|
||||
batch = list(self._token_queue)
|
||||
self._token_queue.clear()
|
||||
@@ -130,12 +172,7 @@ class SessionUsageMixin:
|
||||
return False
|
||||
self._token_queue_cond.wait(remaining)
|
||||
if batch:
|
||||
try:
|
||||
self._apply_token_batch(batch)
|
||||
finally:
|
||||
with self._token_queue_cond:
|
||||
self._token_writer_busy = False
|
||||
self._token_queue_cond.notify_all()
|
||||
self._apply_claimed_batch(batch)
|
||||
return True
|
||||
|
||||
def _token_writer_loop(self) -> None:
|
||||
@@ -145,62 +182,44 @@ class SessionUsageMixin:
|
||||
while not self._token_queue and not self._token_writer_stop:
|
||||
remaining = idle_deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
# Retire under the same lock queue_token_counts() uses to
|
||||
# decide to spawn, so no delta strands behind an exiting worker.
|
||||
# Retire under the same lock queue_token_counts() uses to decide
|
||||
# to spawn, so no delta strands behind an exiting worker.
|
||||
self._token_writer_thread = None
|
||||
return
|
||||
self._token_queue_cond.wait(remaining)
|
||||
if not self._token_queue:
|
||||
self._token_writer_thread = None
|
||||
return # stop requested and fully drained
|
||||
# busy BEFORE clearing the queue: flush's lock-free fast path
|
||||
# reads queue-then-busy and must never see "empty and idle"
|
||||
# while a popped batch is unapplied.
|
||||
# busy BEFORE clearing the queue: flush's lock-free fast path must never
|
||||
# see "empty and idle" while a popped batch is unapplied.
|
||||
self._token_writer_busy = True
|
||||
batch = list(self._token_queue)
|
||||
self._token_queue.clear()
|
||||
try:
|
||||
self._apply_token_batch(batch)
|
||||
finally:
|
||||
with self._token_queue_cond:
|
||||
self._token_writer_busy = False
|
||||
self._token_queue_cond.notify_all()
|
||||
self._apply_claimed_batch(batch)
|
||||
|
||||
def _apply_token_batch(self, batch: List[Tuple[str, Dict[str, Any]]]) -> None:
|
||||
"""Apply queued deltas in order, coalescing where safe. Never raises."""
|
||||
try:
|
||||
coalesced = self._coalesce_token_deltas(batch)
|
||||
except Exception as exc:
|
||||
# Coalescing must never kill the writer (callers cannot observe a
|
||||
# dead one); the merge is only an optimization.
|
||||
logger.warning(
|
||||
"async token accounting: coalesce failed, applying raw "
|
||||
"batch: %s", exc,
|
||||
)
|
||||
# Coalescing must never kill the writer; the merge is only an optimization.
|
||||
logger.warning("async token accounting: coalesce failed, applying raw batch: %s", exc)
|
||||
coalesced = batch
|
||||
for session_id, kwargs in coalesced:
|
||||
try:
|
||||
self.update_token_counts(session_id, **kwargs)
|
||||
except Exception as exc:
|
||||
# Accounting loss is logged, never raised into a turn.
|
||||
logger.warning(
|
||||
"async token accounting: apply failed (session=%s): %s",
|
||||
session_id, exc,
|
||||
)
|
||||
logger.warning("async token accounting: apply failed (session=%s): %s", session_id, exc)
|
||||
|
||||
def _coalesce_token_deltas(
|
||||
self, batch: List[Tuple[str, Dict[str, Any]]]
|
||||
) -> List[Tuple[str, Dict[str, Any]]]:
|
||||
"""Merge adjacent incremental deltas with an identical route, so
|
||||
ordering across sessions and /model switches is preserved exactly.
|
||||
absolute=True deltas never merge."""
|
||||
def _coalesce_token_deltas(self, batch: List[Tuple[str, Dict[str, Any]]]) -> List[Tuple[str, Dict[str, Any]]]:
|
||||
"""Merge adjacent incremental deltas with an identical route, so ordering across
|
||||
sessions and /model switches is preserved exactly. absolute=True never merges."""
|
||||
groups: List[Tuple[Optional[tuple], str, Dict[str, Any]]] = []
|
||||
for session_id, kwargs in batch:
|
||||
key = None
|
||||
if not kwargs.get("absolute"):
|
||||
key = (session_id,) + tuple(
|
||||
kwargs.get(f) for f in self._TOKEN_DELTA_ROUTE_FIELDS
|
||||
)
|
||||
key = (session_id,) + tuple(kwargs.get(f) for f in self._TOKEN_DELTA_ROUTE_FIELDS)
|
||||
if groups and key is not None and groups[-1][0] == key:
|
||||
merged = groups[-1][2]
|
||||
for f in self._TOKEN_DELTA_SUM_FIELDS:
|
||||
@@ -223,18 +242,16 @@ class SessionUsageMixin:
|
||||
if thread is not None and thread.is_alive():
|
||||
thread.join(timeout=join_timeout)
|
||||
if thread.is_alive():
|
||||
# Writer stuck mid-apply: leave deltas unapplied rather than
|
||||
# race it and misorder/double-count.
|
||||
# Writer stuck mid-apply: leave deltas unapplied rather than race it.
|
||||
logger.warning(
|
||||
"async token accounting: writer did not stop within %.0fs; "
|
||||
"%d queued delta(s) not persisted",
|
||||
join_timeout, len(self._token_queue),
|
||||
)
|
||||
return
|
||||
# Writer gone: apply leftovers synchronously under the same busy
|
||||
# protocol. Wait out a flush caller-drain that already claimed busy —
|
||||
# close() nulls the connection right after this returns and must not
|
||||
# yank it mid-batch.
|
||||
# Writer gone: apply leftovers synchronously under the same busy protocol. Wait
|
||||
# out a flush caller-drain that already claimed busy — close() nulls the
|
||||
# connection right after this returns and must not yank it mid-batch.
|
||||
with self._token_queue_cond:
|
||||
deadline = time.monotonic() + join_timeout
|
||||
while self._token_writer_busy:
|
||||
@@ -247,19 +264,13 @@ class SessionUsageMixin:
|
||||
)
|
||||
return
|
||||
self._token_queue_cond.wait(remaining)
|
||||
# busy BEFORE clearing the queue (same ordering as the writer loop),
|
||||
# or flush's lock-free fast path could see "empty and idle".
|
||||
# busy BEFORE clearing the queue (same ordering as the writer loop).
|
||||
batch = list(self._token_queue)
|
||||
if batch:
|
||||
self._token_writer_busy = True
|
||||
self._token_queue.clear()
|
||||
if batch:
|
||||
try:
|
||||
self._apply_token_batch(batch)
|
||||
finally:
|
||||
with self._token_queue_cond:
|
||||
self._token_writer_busy = False
|
||||
self._token_queue_cond.notify_all()
|
||||
self._apply_claimed_batch(batch)
|
||||
|
||||
def _drain_token_queue_at_exit(self) -> None:
|
||||
try:
|
||||
@@ -287,75 +298,21 @@ class SessionUsageMixin:
|
||||
api_call_count: int = 0,
|
||||
absolute: bool = False,
|
||||
) -> None:
|
||||
"""Update token counters and backfill model if unset.
|
||||
|
||||
*absolute*=False increments (per-API-call deltas, CLI path);
|
||||
*absolute*=True sets directly (gateway path, where the cached agent
|
||||
holds cumulative totals).
|
||||
"""
|
||||
# Ensure the row exists: under concurrent load the initial
|
||||
# create_session() may have failed on SQLite locking, and the UPDATE
|
||||
# would silently affect 0 rows.
|
||||
"""Update token counters and backfill model if unset. *absolute*=False
|
||||
increments (per-API-call deltas, CLI path); *absolute*=True sets directly
|
||||
(gateway path, where the cached agent holds cumulative totals)."""
|
||||
# Ensure the row exists: under concurrent load create_session() may have failed
|
||||
# on locking, and the UPDATE would silently affect 0 rows.
|
||||
self._insert_session_row(session_id, "unknown", model=model)
|
||||
if absolute:
|
||||
sql = """UPDATE sessions SET
|
||||
input_tokens = ?,
|
||||
output_tokens = ?,
|
||||
cache_read_tokens = ?,
|
||||
cache_write_tokens = ?,
|
||||
reasoning_tokens = ?,
|
||||
estimated_cost_usd = COALESCE(?, 0),
|
||||
actual_cost_usd = CASE
|
||||
WHEN ? IS NULL THEN actual_cost_usd
|
||||
ELSE ?
|
||||
END,
|
||||
cost_status = COALESCE(?, cost_status),
|
||||
cost_source = COALESCE(?, cost_source),
|
||||
pricing_version = COALESCE(?, pricing_version),
|
||||
billing_provider = COALESCE(billing_provider, ?),
|
||||
billing_base_url = COALESCE(billing_base_url, ?),
|
||||
billing_mode = COALESCE(billing_mode, ?),
|
||||
model = COALESCE(model, ?),
|
||||
api_call_count = ?
|
||||
WHERE id = ?"""
|
||||
else:
|
||||
sql = """UPDATE sessions SET
|
||||
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 = COALESCE(estimated_cost_usd, 0) + COALESCE(?, 0),
|
||||
actual_cost_usd = CASE
|
||||
WHEN ? IS NULL THEN actual_cost_usd
|
||||
ELSE COALESCE(actual_cost_usd, 0) + ?
|
||||
END,
|
||||
cost_status = COALESCE(?, cost_status),
|
||||
cost_source = COALESCE(?, cost_source),
|
||||
pricing_version = COALESCE(?, pricing_version),
|
||||
billing_provider = COALESCE(billing_provider, ?),
|
||||
billing_base_url = COALESCE(billing_base_url, ?),
|
||||
billing_mode = COALESCE(billing_mode, ?),
|
||||
model = COALESCE(model, ?),
|
||||
api_call_count = COALESCE(api_call_count, 0) + ?
|
||||
WHERE id = ?"""
|
||||
has_accounted_usage = bool(
|
||||
sql = _TOKEN_UPDATE_ABSOLUTE_SQL if absolute else _TOKEN_UPDATE_DELTA_SQL
|
||||
has_usage = bool(
|
||||
input_tokens or output_tokens or cache_read_tokens
|
||||
or cache_write_tokens or reasoning_tokens or api_call_count
|
||||
or estimated_cost_usd or actual_cost_usd
|
||||
or cache_write_tokens or reasoning_tokens or api_call_count or estimated_cost_usd
|
||||
)
|
||||
has_accounted_usage = bool(has_usage or actual_cost_usd)
|
||||
params = (
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
cache_read_tokens,
|
||||
cache_write_tokens,
|
||||
reasoning_tokens,
|
||||
estimated_cost_usd,
|
||||
actual_cost_usd,
|
||||
actual_cost_usd,
|
||||
cost_status,
|
||||
cost_source,
|
||||
pricing_version,
|
||||
input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens,
|
||||
estimated_cost_usd, actual_cost_usd, actual_cost_usd, cost_status, cost_source, pricing_version,
|
||||
billing_provider if has_accounted_usage else None,
|
||||
billing_base_url if has_accounted_usage else None,
|
||||
billing_mode if has_accounted_usage else None,
|
||||
@@ -363,31 +320,22 @@ class SessionUsageMixin:
|
||||
api_call_count,
|
||||
session_id,
|
||||
)
|
||||
# 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. Each delta carries the route active at call time
|
||||
# and is recorded into session_model_usage keyed by it. Only the
|
||||
# incremental path records here: absolute cumulative updates cannot be
|
||||
# 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.
|
||||
# Only the incremental path records here — absolute cumulative updates cannot be
|
||||
# split back into routes; Insights reconciles the residual instead.
|
||||
record_model_usage = (not absolute) and (
|
||||
input_tokens or output_tokens or cache_read_tokens
|
||||
or cache_write_tokens or reasoning_tokens or api_call_count
|
||||
or estimated_cost_usd
|
||||
)
|
||||
record_model_usage = (not absolute) and has_usage
|
||||
|
||||
def _do(conn):
|
||||
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()
|
||||
existing_model = row["model"] if row is not None else None
|
||||
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 fails and fallback succeeds, the first accounted usage is
|
||||
# the authoritative route; after that keep the row as is (one row
|
||||
# cannot represent mixed-provider usage).
|
||||
# create_session records the requested route before any API call. If that
|
||||
# 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).
|
||||
first_accounted_route = (
|
||||
existing_api_calls == 0
|
||||
and has_accounted_usage
|
||||
@@ -406,21 +354,12 @@ class SessionUsageMixin:
|
||||
conn.execute(sql, params)
|
||||
if record_model_usage:
|
||||
self._record_model_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,
|
||||
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)
|
||||
@@ -446,77 +385,31 @@ class SessionUsageMixin:
|
||||
api_call_count: int,
|
||||
task: str = "",
|
||||
) -> None:
|
||||
"""Accumulate a per-API-call usage delta into session_model_usage.
|
||||
|
||||
Runs inside the caller's write transaction, after the ``sessions``
|
||||
UPDATE, so per-model rows stay consistent with the summary row. A
|
||||
missing model/provider falls back to the session row (same COALESCE
|
||||
behaviour as the summary update). ``task`` is ``''`` for the main loop;
|
||||
auxiliary calls record their task name via :meth:`record_auxiliary_usage`.
|
||||
"""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
|
||||
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
|
||||
on gemini while the main loop runs anthropic): missing info stays 'unknown'/empty.
|
||||
"""
|
||||
row = conn.execute(
|
||||
"SELECT model, billing_provider, billing_base_url, billing_mode "
|
||||
"FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
sess_model = row["model"] if row is not None else None
|
||||
sess_provider = row["billing_provider"] if row is not None else None
|
||||
sess_base_url = row["billing_base_url"] if row is not None else None
|
||||
sess_billing_mode = row["billing_mode"] if row is not None else None
|
||||
|
||||
# Aux rows must NOT inherit the main-loop route (vision on gemini while
|
||||
# the main loop runs anthropic); missing info stays 'unknown'/empty.
|
||||
if task:
|
||||
eff_model = model or "unknown"
|
||||
eff_provider = billing_provider or ""
|
||||
eff_base_url = billing_base_url or ""
|
||||
eff_billing_mode = billing_mode or ""
|
||||
else:
|
||||
eff_model = model or sess_model or "unknown"
|
||||
eff_provider = billing_provider or sess_provider or ""
|
||||
eff_base_url = billing_base_url or sess_base_url or ""
|
||||
eff_billing_mode = billing_mode or sess_billing_mode or ""
|
||||
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)]
|
||||
now = time.time()
|
||||
conn.execute(
|
||||
"""INSERT INTO session_model_usage (
|
||||
session_id, model, billing_provider, billing_base_url, billing_mode,
|
||||
task, api_call_count, input_tokens, output_tokens,
|
||||
cache_read_tokens, cache_write_tokens, reasoning_tokens,
|
||||
estimated_cost_usd, actual_cost_usd, cost_status, cost_source,
|
||||
first_seen, last_seen
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(session_id, model, billing_provider, billing_base_url, billing_mode, task)
|
||||
DO UPDATE SET
|
||||
api_call_count = api_call_count + excluded.api_call_count,
|
||||
input_tokens = input_tokens + excluded.input_tokens,
|
||||
output_tokens = output_tokens + excluded.output_tokens,
|
||||
cache_read_tokens = cache_read_tokens + excluded.cache_read_tokens,
|
||||
cache_write_tokens = cache_write_tokens + excluded.cache_write_tokens,
|
||||
reasoning_tokens = reasoning_tokens + excluded.reasoning_tokens,
|
||||
estimated_cost_usd = estimated_cost_usd + excluded.estimated_cost_usd,
|
||||
actual_cost_usd = actual_cost_usd + excluded.actual_cost_usd,
|
||||
cost_status = COALESCE(excluded.cost_status, cost_status),
|
||||
cost_source = COALESCE(excluded.cost_source, cost_source),
|
||||
last_seen = excluded.last_seen""",
|
||||
_MODEL_USAGE_UPSERT_SQL,
|
||||
(
|
||||
session_id,
|
||||
eff_model,
|
||||
eff_provider,
|
||||
eff_base_url,
|
||||
eff_billing_mode,
|
||||
task or "",
|
||||
api_call_count or 0,
|
||||
input_tokens or 0,
|
||||
output_tokens or 0,
|
||||
cache_read_tokens or 0,
|
||||
cache_write_tokens or 0,
|
||||
reasoning_tokens or 0,
|
||||
float(estimated_cost_usd or 0.0),
|
||||
float(actual_cost_usd or 0.0),
|
||||
cost_status,
|
||||
cost_source,
|
||||
now,
|
||||
now,
|
||||
session_id, eff_model, eff_provider, eff_base_url, eff_billing_mode, task or "",
|
||||
api_call_count or 0, *counts,
|
||||
float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0),
|
||||
cost_status, cost_source, now, now,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -536,16 +429,11 @@ class SessionUsageMixin:
|
||||
estimated_cost_usd: Optional[float] = None,
|
||||
api_call_count: int = 1,
|
||||
) -> None:
|
||||
"""Record an auxiliary LLM call's usage (vision, compression, title
|
||||
generation, ...) against *session_id*.
|
||||
|
||||
Writes a per-(model, provider, task) delta into ``session_model_usage``
|
||||
WITHOUT touching the ``sessions`` summary row: the gateway overwrites
|
||||
session counters with absolute main-loop totals, so aux tokens there
|
||||
would be clobbered or double-counted. Insights read the union.
|
||||
``api_call_count`` may aggregate N calls (background-review forks).
|
||||
Best-effort: callers must never fail an aux call over accounting.
|
||||
"""
|
||||
"""Record an auxiliary LLM call's usage (vision, compression, title generation,
|
||||
...) as a per-(model, provider, task) delta in ``session_model_usage`` WITHOUT
|
||||
touching the ``sessions`` summary row (the gateway overwrites those counters with
|
||||
absolute main-loop totals). ``api_call_count`` may aggregate N calls. Best-effort:
|
||||
callers must never fail an aux call over accounting."""
|
||||
if not session_id or not task:
|
||||
return
|
||||
# FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts.
|
||||
@@ -553,37 +441,24 @@ class SessionUsageMixin:
|
||||
|
||||
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)
|
||||
),
|
||||
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]:
|
||||
"""Tokens and spend across the whole store (one scan), so the sidebar
|
||||
total does not shrink with paging. Spend prefers the billed figure over
|
||||
the estimate, the same precedence a single row renders."""
|
||||
"""Tokens and spend across the whole store (one scan), so the sidebar total does
|
||||
not shrink with paging. Spend prefers the billed figure over the estimate."""
|
||||
where = ["parent_session_id IS NULL", "message_count >= ?"]
|
||||
params: List[Any] = [min_message_count]
|
||||
if not include_archived:
|
||||
where.append("COALESCE(archived, 0) = 0")
|
||||
|
||||
row = self._read_one(
|
||||
f"""
|
||||
SELECT COALESCE(SUM(COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0)), 0),
|
||||
@@ -593,5 +468,4 @@ class SessionUsageMixin:
|
||||
""",
|
||||
params,
|
||||
)
|
||||
|
||||
return {"tokens": int(row[0] or 0), "cost_usd": float(row[1] or 0.0)}
|
||||
|
||||
Reference in New Issue
Block a user