refactor(state): _cooldown_row shape helper; per-model usage kwargs via field set; small folds

This commit is contained in:
Teknium
2026-09-02 17:17:16 -07:00
parent d77f3ad79b
commit 1991695511
3 changed files with 50 additions and 70 deletions
+17 -24
View File
@@ -27,6 +27,14 @@ def _ended_by_compression(row) -> bool:
return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression" return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression"
def _cooldown_row(exists: bool, cooldown_until, error) -> Dict[str, Any]:
return {
"session_exists": exists,
"cooldown_until": float(cooldown_until) if cooldown_until is not None else None,
"error": error,
}
def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now: float, expires_at: float, def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now: float, expires_at: float,
stale) -> Tuple[bool, Optional[str]]: stale) -> Tuple[bool, Optional[str]]:
"""Single-transaction lease claim: DELETE a stale holder's row (``stale(holder, """Single-transaction lease claim: DELETE a stale holder's row (``stale(holder,
@@ -286,24 +294,17 @@ class SessionCompressionMixin:
return None return None
now = time.time() now = time.time()
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,))
if row is None or row[0] is None: if row is None or row[0] is None or float(row[0]) <= now:
return None return None
cooldown_until = float(row[0]) return {"cooldown_until": float(row[0]), "remaining_seconds": float(row[0]) - now, "error": row[1]}
if cooldown_until <= now:
return None
return {"cooldown_until": cooldown_until, "remaining_seconds": cooldown_until - now, "error": row[1]}
def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]: def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]:
"""Exact stored cooldown columns, no expiry filtering, so compression """Exact stored cooldown columns, no expiry filtering, so compression
cancellation can roll back an expired, partially-null, or absent row exactly.""" cancellation can roll back an expired, partially-null, or absent row exactly."""
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) if session_id else None row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) if session_id else None
if row is None: if row is None:
return {"session_exists": False, "cooldown_until": None, "error": None} return _cooldown_row(False, None, None)
return { return _cooldown_row(True, row[0], row[1])
"session_exists": True,
"cooldown_until": float(row[0]) if row[0] is not None else None,
"error": row[1],
}
def restore_compression_failure_cooldown_row(self, session_id: str, snapshot: Dict[str, Any]) -> None: def restore_compression_failure_cooldown_row(self, session_id: str, snapshot: Dict[str, Any]) -> None:
"""Restore and verify an exact cooldown-row snapshot. Unlike record/clear this """Restore and verify an exact cooldown-row snapshot. Unlike record/clear this
@@ -323,14 +324,9 @@ class SessionCompressionMixin:
) )
if cursor.rowcount != 1: if cursor.rowcount != 1:
raise RuntimeError(f"compression cooldown rollback session missing: {session_id}") raise RuntimeError(f"compression cooldown rollback session missing: {session_id}")
self._execute_write(_do) self._execute_write(_do)
actual = self.get_compression_failure_cooldown_row(session_id) actual = self.get_compression_failure_cooldown_row(session_id)
expected = { expected = _cooldown_row(True, deadline, error)
"session_exists": True,
"cooldown_until": float(deadline) if deadline is not None else None,
"error": error,
}
if actual != expected: if actual != expected:
raise RuntimeError( raise RuntimeError(
f"compression cooldown rollback verification failed: " f"compression cooldown rollback verification failed: "
@@ -582,11 +578,10 @@ class SessionCompressionMixin:
def _do(conn): def _do(conn):
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
cursor = conn.execute( return conn.execute(
"UPDATE session_turn_leases SET expires_at = ? " "UPDATE session_turn_leases SET expires_at = ? "
"WHERE conversation_id = ? AND holder = ?", (expires_at, conversation_id, holder), "WHERE conversation_id = ? AND holder = ?", (expires_at, conversation_id, holder),
) ).rowcount > 0
return cursor.rowcount > 0
return bool(self._execute_write(_do)) return bool(self._execute_write(_do))
@@ -656,7 +651,7 @@ class SessionCompressionMixin:
are still live over stale closed siblings such as ``ws_orphan_reap``.""" are still live over stale closed siblings such as ``ws_orphan_reap``."""
current = session_id current = session_id
chain = [current] if current else [] chain = [current] if current else []
seen = {current} if current else set() seen = set(chain)
for _ in range(100): # defensive bound; chains this deep are pathological for _ in range(100): # defensive bound; chains this deep are pathological
with self._read_ctx() as conn: with self._read_ctx() as conn:
row = conn.execute( row = conn.execute(
@@ -682,9 +677,7 @@ class SessionCompressionMixin:
""", """,
(current,), (current,),
).fetchone() ).fetchone()
if row is None: child_id = row["id"] if row is not None else None
return chain
child_id = row["id"]
if not child_id or child_id in seen: if not child_id or child_id in seen:
return chain return chain
seen.add(child_id) seen.add(child_id)
+3 -6
View File
@@ -203,9 +203,6 @@ class SessionTitlesMixin:
) )
if not rows: if not rows:
return base return base
max_num = 1 # the unnumbered original counts as #1 # The unnumbered original counts as #1.
for row in rows: numbers = [int(m.group(1)) for m in (re.match(r'^.* #(\d+)$', row["title"]) for row in rows) if m]
m = re.match(r'^.* #(\d+)$', row["title"]) return f"{base} #{max([1, *numbers]) + 1}"
if m:
max_num = max(max_num, int(m.group(1)))
return f"{base} #{max_num + 1}"
+30 -40
View File
@@ -78,6 +78,15 @@ _MODEL_USAGE_UPSERT_SQL = """INSERT INTO session_model_usage (
last_seen = excluded.last_seen""" last_seen = excluded.last_seen"""
# Kwargs forwarded verbatim from update_token_counts / record_auxiliary_usage into
# _record_model_usage (the per-route attribution row).
_MODEL_USAGE_FIELDS = frozenset((
"model", "billing_provider", "billing_base_url", "billing_mode", "input_tokens", "output_tokens",
"cache_read_tokens", "cache_write_tokens", "reasoning_tokens", "estimated_cost_usd",
"actual_cost_usd", "cost_status", "cost_source", "api_call_count",
))
class SessionUsageMixin: class SessionUsageMixin:
"""Coalesced token writer, per-model usage rows, billing route.""" """Coalesced token writer, per-model usage rows, billing route."""
@@ -110,10 +119,11 @@ class SessionUsageMixin:
to the synchronous path and may raise.""" to the synchronous path and may raise."""
with self._token_queue_cond: with self._token_queue_cond:
thread = self._token_writer_thread thread = self._token_writer_thread
writer_stopped = self._token_writer_stop and (thread is None or not thread.is_alive()) writer_alive = thread is not None and thread.is_alive()
writer_stopped = self._token_writer_stop and not writer_alive
if not writer_stopped: if not writer_stopped:
self._token_queue.append((session_id, kwargs)) self._token_queue.append((session_id, kwargs))
if thread is None or not thread.is_alive(): if not writer_alive:
# Daemon so exit never hangs on accounting; the atexit hook drains # Daemon so exit never hangs on accounting; the atexit hook drains
# leftovers. ``not is_alive()`` (not ``is None``) respawns a writer # leftovers. ``not is_alive()`` (not ``is None``) respawns a writer
# that died from an unexpected escape. # that died from an unexpected escape.
@@ -289,6 +299,7 @@ class SessionUsageMixin:
"""Update token counters and backfill model if unset. *absolute*=False """Update token counters and backfill model if unset. *absolute*=False
increments (per-API-call deltas, CLI path); *absolute*=True sets directly increments (per-API-call deltas, CLI path); *absolute*=True sets directly
(gateway path, where the cached agent holds cumulative totals).""" (gateway path, where the cached agent holds cumulative totals)."""
usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS}
# Ensure the row exists: under concurrent load create_session() may have failed # Ensure the row exists: under concurrent load create_session() may have failed
# on locking, and the UPDATE would silently affect 0 rows. # on locking, and the UPDATE would silently affect 0 rows.
self._insert_session_row(session_id, "unknown", model=model) self._insert_session_row(session_id, "unknown", model=model)
@@ -316,18 +327,14 @@ class SessionUsageMixin:
row = conn.execute( row = conn.execute(
"SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", (session_id,), "SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", (session_id,),
).fetchone() ).fetchone()
existing_model = row["model"] if row is not None else None existing = dict(row) if row is not None else {}
existing_provider = row["billing_provider"] if row is not None else None
existing_api_calls = int((row["api_call_count"] if row is not None else 0) or 0)
# create_session records the requested route before any API call. If that # create_session records the requested route before any API call. If that
# fails and fallback succeeds, the first accounted usage is the authoritative # fails and fallback succeeds, the first accounted usage is the authoritative
# route; after that keep the row as is (one row cannot represent mixed usage). # route; after that keep the row as is (one row cannot represent mixed usage).
first_accounted_route = ( first_accounted_route = (
existing_api_calls == 0 int(existing.get("api_call_count") or 0) == 0 and has_accounted_usage and bool(model)
and has_accounted_usage
and bool(model)
and bool(billing_provider) and bool(billing_provider)
and (existing_model != model or existing_provider != billing_provider) and (existing.get("model") != model or existing.get("billing_provider") != billing_provider)
) )
if first_accounted_route: if first_accounted_route:
conn.execute( conn.execute(
@@ -339,23 +346,16 @@ class SessionUsageMixin:
) )
conn.execute(sql, params) conn.execute(sql, params)
if record_model_usage: if record_model_usage:
self._record_model_usage( self._record_model_usage(conn, session_id, **usage)
conn, session_id, model=model, billing_provider=billing_provider,
billing_base_url=billing_base_url, billing_mode=billing_mode,
input_tokens=input_tokens, output_tokens=output_tokens,
cache_read_tokens=cache_read_tokens, cache_write_tokens=cache_write_tokens,
reasoning_tokens=reasoning_tokens, estimated_cost_usd=estimated_cost_usd,
actual_cost_usd=actual_cost_usd, cost_status=cost_status, cost_source=cost_source,
api_call_count=api_call_count,
)
self._execute_write(_do) self._execute_write(_do)
def _record_model_usage( def _record_model_usage(
self, conn, session_id: str, *, model: Optional[str], billing_provider: Optional[str], self, conn, session_id: str, *, model: Optional[str] = None, billing_provider: Optional[str] = None,
billing_base_url: Optional[str], billing_mode: Optional[str], input_tokens: int, billing_base_url: Optional[str] = None, billing_mode: Optional[str] = None, input_tokens: int = 0,
output_tokens: int, cache_read_tokens: int, cache_write_tokens: int, reasoning_tokens: int, output_tokens: int = 0, cache_read_tokens: int = 0, cache_write_tokens: int = 0,
estimated_cost_usd: Optional[float], actual_cost_usd: Optional[float], reasoning_tokens: int = 0, estimated_cost_usd: Optional[float] = None,
cost_status: Optional[str], cost_source: Optional[str], api_call_count: int, task: str = "", actual_cost_usd: Optional[float] = None, cost_status: Optional[str] = None,
cost_source: Optional[str] = None, api_call_count: int = 0, task: str = "",
) -> None: ) -> None:
"""Accumulate a per-API-call usage delta into session_model_usage, inside the """Accumulate a per-API-call usage delta into session_model_usage, inside the
caller's write txn after the ``sessions`` UPDATE. A missing model/provider falls caller's write txn after the ``sessions`` UPDATE. A missing model/provider falls
@@ -367,16 +367,15 @@ class SessionUsageMixin:
"FROM sessions WHERE id = ?", (session_id,), "FROM sessions WHERE id = ?", (session_id,),
).fetchone() ).fetchone()
sess = dict(row) if (row is not None and not task) else {} sess = dict(row) if (row is not None and not task) else {}
eff_model = model or sess.get("model") or "unknown"
eff_provider = billing_provider or sess.get("billing_provider") or ""
eff_base_url = billing_base_url or sess.get("billing_base_url") or ""
eff_billing_mode = billing_mode or sess.get("billing_mode") or ""
counts = [v or 0 for v in (input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens)] counts = [v or 0 for v in (input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens)]
now = time.time() now = time.time()
conn.execute( conn.execute(
_MODEL_USAGE_UPSERT_SQL, _MODEL_USAGE_UPSERT_SQL,
( (
session_id, eff_model, eff_provider, eff_base_url, eff_billing_mode, task or "", session_id, model or sess.get("model") or "unknown",
billing_provider or sess.get("billing_provider") or "",
billing_base_url or sess.get("billing_base_url") or "",
billing_mode or sess.get("billing_mode") or "", task or "",
api_call_count or 0, *counts, api_call_count or 0, *counts,
float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0), float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0),
cost_status, cost_source, now, now, cost_status, cost_source, now, now,
@@ -395,22 +394,13 @@ class SessionUsageMixin:
touching the ``sessions`` summary row (the gateway overwrites those counters with touching the ``sessions`` summary row (the gateway overwrites those counters with
absolute main-loop totals). ``api_call_count`` may aggregate N calls. Best-effort: absolute main-loop totals). ``api_call_count`` may aggregate N calls. Best-effort:
callers must never fail an aux call over accounting.""" callers must never fail an aux call over accounting."""
usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS}
if not session_id or not task: if not session_id or not task:
return return
usage["api_call_count"] = 1 if api_call_count is None else int(api_call_count)
# FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts. # FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts.
self._insert_session_row(session_id, "unknown") self._insert_session_row(session_id, "unknown")
self._execute_write(lambda conn: self._record_model_usage(conn, session_id, task=task, **usage))
def _do(conn):
self._record_model_usage(
conn, session_id, model=model, billing_provider=billing_provider,
billing_base_url=billing_base_url, billing_mode=None,
input_tokens=input_tokens or 0, output_tokens=output_tokens or 0,
cache_read_tokens=cache_read_tokens or 0, cache_write_tokens=cache_write_tokens or 0,
reasoning_tokens=reasoning_tokens or 0, estimated_cost_usd=estimated_cost_usd,
actual_cost_usd=None, cost_status=None, cost_source=None,
api_call_count=1 if api_call_count is None else int(api_call_count), task=task,
)
self._execute_write(_do)
def usage_totals(self, *, min_message_count: int = 1, include_archived: bool = False) -> Dict[str, float]: def usage_totals(self, *, min_message_count: int = 1, include_archived: bool = False) -> Dict[str, float]:
"""Tokens and spend across the whole store (one scan), so the sidebar total does """Tokens and spend across the whole store (one scan), so the sidebar total does