diff --git a/hermes_state_compression.py b/hermes_state_compression.py index d6eae97865..0f1f816bba 100644 --- a/hermes_state_compression.py +++ b/hermes_state_compression.py @@ -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" +def _cooldown_row(exists: bool, cooldown_until, error) -> Dict[str, Any]: + return { + "session_exists": exists, + "cooldown_until": float(cooldown_until) if cooldown_until is not None else None, + "error": error, + } + + def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now: float, expires_at: float, stale) -> Tuple[bool, Optional[str]]: """Single-transaction lease claim: DELETE a stale holder's row (``stale(holder, @@ -286,24 +294,17 @@ class SessionCompressionMixin: return None now = time.time() 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 - cooldown_until = float(row[0]) - if cooldown_until <= now: - return None - return {"cooldown_until": cooldown_until, "remaining_seconds": cooldown_until - now, "error": row[1]} + return {"cooldown_until": float(row[0]), "remaining_seconds": float(row[0]) - now, "error": row[1]} 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} - return { - "session_exists": True, - "cooldown_until": float(row[0]) if row[0] is not None else None, - "error": row[1], - } + return _cooldown_row(False, None, None) + return _cooldown_row(True, row[0], 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 @@ -323,14 +324,9 @@ class SessionCompressionMixin: ) if cursor.rowcount != 1: raise RuntimeError(f"compression cooldown rollback session missing: {session_id}") - self._execute_write(_do) actual = self.get_compression_failure_cooldown_row(session_id) - expected = { - "session_exists": True, - "cooldown_until": float(deadline) if deadline is not None else None, - "error": error, - } + expected = _cooldown_row(True, deadline, error) if actual != expected: raise RuntimeError( f"compression cooldown rollback verification failed: " @@ -582,11 +578,10 @@ class SessionCompressionMixin: def _do(conn): 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 = ? " "WHERE conversation_id = ? AND holder = ?", (expires_at, conversation_id, holder), - ) - return cursor.rowcount > 0 + ).rowcount > 0 return bool(self._execute_write(_do)) @@ -656,7 +651,7 @@ class SessionCompressionMixin: 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() + seen = set(chain) for _ in range(100): # defensive bound; chains this deep are pathological with self._read_ctx() as conn: row = conn.execute( @@ -682,9 +677,7 @@ class SessionCompressionMixin: """, (current,), ).fetchone() - if row is None: - return chain - child_id = row["id"] + child_id = row["id"] if row is not None else None if not child_id or child_id in seen: return chain seen.add(child_id) diff --git a/hermes_state_titles.py b/hermes_state_titles.py index 9aad634895..a66080c58d 100644 --- a/hermes_state_titles.py +++ b/hermes_state_titles.py @@ -203,9 +203,6 @@ class SessionTitlesMixin: ) if not rows: return base - 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}" + # The unnumbered original counts as #1. + numbers = [int(m.group(1)) for m in (re.match(r'^.* #(\d+)$', row["title"]) for row in rows) if m] + return f"{base} #{max([1, *numbers]) + 1}" diff --git a/hermes_state_usage.py b/hermes_state_usage.py index 90f32aded3..ae1e20790e 100644 --- a/hermes_state_usage.py +++ b/hermes_state_usage.py @@ -78,6 +78,15 @@ _MODEL_USAGE_UPSERT_SQL = """INSERT INTO session_model_usage ( 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: """Coalesced token writer, per-model usage rows, billing route.""" @@ -110,10 +119,11 @@ class SessionUsageMixin: 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_alive = thread is not None and thread.is_alive() + writer_stopped = self._token_writer_stop and not writer_alive if not writer_stopped: 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 # leftovers. ``not is_alive()`` (not ``is None``) respawns a writer # that died from an unexpected escape. @@ -289,6 +299,7 @@ class SessionUsageMixin: """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).""" + 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 # on locking, and the UPDATE would silently affect 0 rows. self._insert_session_row(session_id, "unknown", model=model) @@ -316,18 +327,14 @@ class SessionUsageMixin: row = conn.execute( "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) + existing = dict(row) if row is not None else {} # 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 - and bool(model) + int(existing.get("api_call_count") or 0) == 0 and has_accounted_usage and bool(model) 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: conn.execute( @@ -339,23 +346,16 @@ 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, - api_call_count=api_call_count, - ) + self._record_model_usage(conn, session_id, **usage) self._execute_write(_do) def _record_model_usage( - self, conn, session_id: str, *, model: Optional[str], billing_provider: Optional[str], - billing_base_url: Optional[str], billing_mode: Optional[str], input_tokens: int, - output_tokens: int, cache_read_tokens: int, cache_write_tokens: int, reasoning_tokens: int, - estimated_cost_usd: Optional[float], actual_cost_usd: Optional[float], - cost_status: Optional[str], cost_source: Optional[str], api_call_count: int, task: str = "", + self, conn, session_id: str, *, model: Optional[str] = None, billing_provider: Optional[str] = None, + billing_base_url: Optional[str] = None, billing_mode: Optional[str] = None, input_tokens: int = 0, + output_tokens: int = 0, cache_read_tokens: int = 0, cache_write_tokens: int = 0, + reasoning_tokens: int = 0, estimated_cost_usd: Optional[float] = None, + actual_cost_usd: Optional[float] = None, cost_status: Optional[str] = None, + cost_source: Optional[str] = None, api_call_count: int = 0, task: str = "", ) -> None: """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 @@ -367,16 +367,15 @@ class SessionUsageMixin: "FROM sessions WHERE id = ?", (session_id,), ).fetchone() 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( _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, float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0), cost_status, cost_source, now, now, @@ -395,22 +394,13 @@ class SessionUsageMixin: 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.""" + usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS} if not session_id or not task: 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. self._insert_session_row(session_id, "unknown") - - 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) + self._execute_write(lambda conn: self._record_model_usage(conn, session_id, task=task, **usage)) 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