mirror of
https://github.com/NousResearch/hermes-agent.git
synced 2026-07-31 19:16:29 +00:00
perf(state): apply per-call token accounting on a background single-writer queue
Every API call in the tool loop persisted its token/cost delta by calling SessionDB.update_token_counts() synchronously on the turn thread — a BEGIN IMMEDIATE sessions UPDATE plus a session_model_usage upsert, measured in production at p50 3.3ms / p95 70.4ms per call and up to 299ms against a cold multi-GB state.db. The tool loop stalls for that long between calls, N times per multi-tool turn. SessionDB gains queue_token_counts(): same signature and semantics as update_token_counts(), but the critical path is a deque append plus a condvar notify. A lazily started daemon thread applies deltas in enqueue order through the existing update_token_counts -> _execute_write path, so the established self._lock / BEGIN IMMEDIATE / jitter-retry discipline is unchanged. When a backlog forms, adjacent same-route incremental deltas coalesce into one UPDATE: token and api-call fields sum, cost fields sum None-preservingly (an all-None run stays None so COALESCE keeps the stored value), and absolute=True deltas never merge and act as ordering barriers. Route equality is required for a merge because those fields feed COALESCE backfill, the last-non-None-wins status fields, and the per-model usage attribution key — a merged apply is row-equivalent to sequential applies. Correctness and durability: - flush_token_counts() gives read-your-writes to token/cost readers (get_session, list_sessions_rich, _get_session_rich_row, list_gateway_sessions, InsightsEngine.generate) — a plain attribute check when nothing is queued. The writer sets its busy flag before popping the queue so the lock-free fast path can never miss an in-flight batch. - update_session_model / update_session_billing_route / update_session_meta write the sessions row synchronously, bypassing the queue, so they flush it first: a still-queued first-of-session delta carries the pre-switch route, and applying it after the switch UPDATE would trip the first_accounted_route branch (api_call_count == 0 plus a route mismatch) and resurrect the old model/provider. - AIAgent._persist_session flushes at turn finalize and every error-exit persist point; close() stops and drains the writer before the WAL checkpoint; an atexit hook (registered on first enqueue, unregistered on close so closed instances are not pinned until interpreter exit) drains at shutdown. Worst-case crash loss is the in-flight call's delta — the same window as the old inline write. - A flush trusts a live stop-flagged writer (its loop drains before exiting) and only drains on the caller's thread when the writer is dead or never started, claiming the same busy flag so concurrent flushes wait instead of racing an in-flight batch. - After close() has stopped the writer, queue_token_counts applies the delta inline instead of parking it on a queue nothing will drain; a closed-connection failure then raises at the call site, which already guards for it, exactly like the old synchronous path. - Writer apply failures are logged and never raise into a turn; the writer thread survives and keeps applying. Call sites switched to the queue: the per-call site in agent/conversation_loop.py and both codex app-server sites in agent/codex_runtime.py. In-memory per-turn counters (agent.session_estimated_cost_usd etc.) stay synchronous, so live turn displays never see the queue. Tests: tests/agent/test_async_token_accounting.py (19 tests: enqueue ordering, absolute-as-barrier, backlog coalescing with exact sums, coalesced-vs-sequential row equivalence, merge unit rules, None-cost preservation, read-your-writes, flush vs stop-flagged/concurrent drains, inline apply after writer stop, close/atexit durability, _persist_session drain, writer failure isolation); tests/run_agent/test_token_persistence_non_cli.py updated to the queue_token_counts contract.
This commit is contained in:
parent
b5dc471152
commit
174ad45939
7 changed files with 881 additions and 15 deletions
|
|
@ -74,7 +74,10 @@ def _record_codex_app_server_usage(agent, turn) -> dict[str, Any]:
|
|||
try:
|
||||
if not agent._session_db_created:
|
||||
agent._ensure_db_session()
|
||||
agent._session_db.update_token_counts(
|
||||
# Enqueued for the SessionDB background writer — keeps the
|
||||
# per-call accounting write off the turn thread (see
|
||||
# conversation_loop's queue_token_counts call).
|
||||
agent._session_db.queue_token_counts(
|
||||
agent.session_id,
|
||||
model=agent.model,
|
||||
billing_provider=agent.provider,
|
||||
|
|
@ -154,7 +157,8 @@ def _record_codex_app_server_usage(agent, turn) -> dict[str, Any]:
|
|||
try:
|
||||
if not agent._session_db_created:
|
||||
agent._ensure_db_session()
|
||||
agent._session_db.update_token_counts(
|
||||
# Enqueued for the SessionDB background writer (see above).
|
||||
agent._session_db.queue_token_counts(
|
||||
agent.session_id,
|
||||
input_tokens=canonical_usage.input_tokens,
|
||||
output_tokens=canonical_usage.output_tokens,
|
||||
|
|
|
|||
|
|
@ -3217,7 +3217,12 @@ def run_conversation(
|
|||
_cost_delta = (_cost_delta or 0.0) + float(_moa_ref_cost)
|
||||
except (TypeError, ValueError): # pragma: no cover
|
||||
pass
|
||||
agent._session_db.update_token_counts(
|
||||
# Enqueued, not written: the background writer
|
||||
# applies the delta off the turn thread (a cold
|
||||
# state.db UPDATE here stalled the tool loop for
|
||||
# up to hundreds of ms per API call). Drained at
|
||||
# turn finalize via _persist_session.
|
||||
agent._session_db.queue_token_counts(
|
||||
agent.session_id,
|
||||
input_tokens=canonical_usage.input_tokens,
|
||||
output_tokens=canonical_usage.output_tokens,
|
||||
|
|
|
|||
|
|
@ -113,6 +113,13 @@ class InsightsEngine:
|
|||
"""
|
||||
cutoff = time.time() - (days * 86400)
|
||||
|
||||
# Token/cost totals may still sit on the SessionDB's async
|
||||
# accounting queue; drain so the report reflects exact counters.
|
||||
# (self.db may be a raw sqlite3 connection in tests — guard.)
|
||||
flush = getattr(self.db, "flush_token_counts", None)
|
||||
if callable(flush):
|
||||
flush()
|
||||
|
||||
# Gather raw data
|
||||
sessions = self._get_sessions(cutoff, source)
|
||||
tool_usage = self._get_tool_usage(cutoff, source)
|
||||
|
|
|
|||
283
hermes_state.py
283
hermes_state.py
|
|
@ -15,6 +15,7 @@ Key design decisions:
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import atexit
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
|
|
@ -25,6 +26,7 @@ import sqlite3
|
|||
import sys
|
||||
import threading
|
||||
import time
|
||||
from collections import deque
|
||||
from pathlib import Path
|
||||
|
||||
from agent.memory_manager import sanitize_context
|
||||
|
|
@ -2015,6 +2017,14 @@ class SessionDB:
|
|||
self._fts_cjk_available = False
|
||||
self._fts_unavailable_warned = False
|
||||
self._conn = None
|
||||
# Async token accounting (see queue_token_counts). The condition
|
||||
# guards queue + writer state; it is distinct from self._lock so
|
||||
# enqueue/flush bookkeeping never contends with SQLite writes.
|
||||
self._token_queue: deque = deque()
|
||||
self._token_queue_cond = threading.Condition(threading.Lock())
|
||||
self._token_writer_thread: Optional[threading.Thread] = None
|
||||
self._token_writer_stop = False
|
||||
self._token_writer_busy = False
|
||||
try:
|
||||
if read_only:
|
||||
# Read-only attach for cross-profile aggregation: SELECT-only,
|
||||
|
|
@ -2660,9 +2670,17 @@ class SessionDB:
|
|||
def close(self):
|
||||
"""Close the database connection.
|
||||
|
||||
Attempts a TRUNCATE WAL checkpoint first so that exiting processes
|
||||
help shrink the WAL file.
|
||||
Drains queued token deltas first (the background writer needs the
|
||||
connection), then attempts a TRUNCATE WAL checkpoint so that
|
||||
exiting processes help shrink the WAL file.
|
||||
"""
|
||||
self._stop_token_writer()
|
||||
# The atexit hook holds a strong reference to this instance (bound
|
||||
# method); without unregistering, every closed SessionDB stays
|
||||
# reachable until interpreter exit. Bound methods compare equal by
|
||||
# (instance, function), so this removes exactly our registration;
|
||||
# no-op when the writer never started.
|
||||
atexit.unregister(self._drain_token_queue_at_exit)
|
||||
with self._lock:
|
||||
if self._conn:
|
||||
try:
|
||||
|
|
@ -4099,6 +4117,9 @@ class SessionDB:
|
|||
filters on ``source``; ``active_only`` restricts to sessions that
|
||||
have not ended.
|
||||
"""
|
||||
# Full rows carry token/cost totals (MCP listings, /status) — drain
|
||||
# queued async accounting deltas so consumers see exact counters.
|
||||
self.flush_token_counts()
|
||||
query = """
|
||||
SELECT sessions.*,
|
||||
COALESCE(
|
||||
|
|
@ -4938,6 +4959,9 @@ class SessionDB:
|
|||
column unchanged. Routes through _execute_write for the standard
|
||||
BEGIN IMMEDIATE + jitter-retry + lock guarantee.
|
||||
"""
|
||||
# Barrier against queued token deltas — see update_session_model.
|
||||
self.flush_token_counts()
|
||||
|
||||
def _do(conn):
|
||||
conn.execute(
|
||||
"UPDATE sessions SET model_config = ?, model = COALESCE(?, model) WHERE id = ?",
|
||||
|
|
@ -4965,6 +4989,14 @@ class SessionDB:
|
|||
switch explicitly replaces any confirmed Browser runtime lock while
|
||||
preserving unrelated lineage markers in ``model_config``.
|
||||
"""
|
||||
# This write bypasses the token queue, so deltas enqueued before the
|
||||
# switch must land first: a still-queued first delta carries the
|
||||
# pre-switch route, and applying it after this UPDATE would trip the
|
||||
# first_accounted_route overwrite in update_token_counts (row sees
|
||||
# api_call_count == 0 + a route mismatch) and resurrect the old
|
||||
# model/provider. Flushing here restores the pre-queue ordering.
|
||||
self.flush_token_counts()
|
||||
|
||||
def _do(conn):
|
||||
conn.execute(
|
||||
"""UPDATE sessions SET
|
||||
|
|
@ -5053,6 +5085,9 @@ class SessionDB:
|
|||
stale ``Model:`` / ``Provider:`` header) is rebuilt — matching the
|
||||
behavior of ``update_session_model`` (see #48173, #48248).
|
||||
"""
|
||||
# Barrier against queued token deltas — see update_session_model.
|
||||
self.flush_token_counts()
|
||||
|
||||
def _do(conn):
|
||||
conn.execute(
|
||||
"""UPDATE sessions SET
|
||||
|
|
@ -5065,6 +5100,241 @@ class SessionDB:
|
|||
)
|
||||
self._execute_write(_do)
|
||||
|
||||
# ── Async token accounting ──
|
||||
# update_token_counts() runs a sessions UPDATE (plus a per-model usage
|
||||
# upsert) inside BEGIN IMMEDIATE; against a cold multi-GB state.db one
|
||||
# call can stall the turn thread for tens to hundreds of ms, and the
|
||||
# tool loop pays it after EVERY API call (measured p50 3.3ms / p95 70ms
|
||||
# per call in production). queue_token_counts() reduces the critical
|
||||
# path to a deque append: a dedicated single-writer thread applies
|
||||
# deltas in enqueue order, coalescing consecutive same-route deltas
|
||||
# into one UPDATE when a backlog forms. Readers that need exact
|
||||
# mid-turn totals (get_session and friends) call flush_token_counts()
|
||||
# first — a plain attribute check when nothing is queued.
|
||||
|
||||
# Delta fields summed when coalescing. Route fields must be equal for
|
||||
# two deltas to merge: model/billing_* feed COALESCE backfill and the
|
||||
# per-model usage attribution key, and cost_status/cost_source are
|
||||
# last-non-None-wins — equality makes the merged UPDATE byte-for-byte
|
||||
# equivalent to applying the deltas sequentially.
|
||||
_TOKEN_DELTA_SUM_FIELDS = (
|
||||
"input_tokens", "output_tokens", "cache_read_tokens",
|
||||
"cache_write_tokens", "reasoning_tokens", "api_call_count",
|
||||
)
|
||||
_TOKEN_DELTA_COST_FIELDS = ("estimated_cost_usd", "actual_cost_usd")
|
||||
_TOKEN_DELTA_ROUTE_FIELDS = (
|
||||
"model", "cost_status", "cost_source", "pricing_version",
|
||||
"billing_provider", "billing_base_url", "billing_mode",
|
||||
)
|
||||
|
||||
def queue_token_counts(self, session_id: str, **kwargs) -> None:
|
||||
"""Enqueue a token/cost delta for the background writer.
|
||||
|
||||
Accepts the same keyword arguments as :meth:`update_token_counts`
|
||||
and applies them asynchronously with identical semantics. Cheap
|
||||
(append + notify) — safe to call on the turn thread after every
|
||||
API call. After close() has stopped the writer, falls back to the
|
||||
synchronous path and may raise like :meth:`update_token_counts`.
|
||||
"""
|
||||
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()
|
||||
)
|
||||
if not writer_stopped:
|
||||
self._token_queue.append((session_id, kwargs))
|
||||
if thread is None:
|
||||
# Daemon so process exit never hangs on accounting; the
|
||||
# atexit hook drains anything still queued at interpreter
|
||||
# shutdown (registered once per instance, on first use).
|
||||
thread = threading.Thread(
|
||||
target=self._token_writer_loop,
|
||||
name="session-db-token-writer",
|
||||
daemon=True,
|
||||
)
|
||||
self._token_writer_thread = thread
|
||||
thread.start()
|
||||
atexit.register(self._drain_token_queue_at_exit)
|
||||
self._token_queue_cond.notify_all()
|
||||
if writer_stopped:
|
||||
# Writer permanently stopped (close() ran; a stop-flagged but
|
||||
# still-live writer keeps accepting — its loop drains before
|
||||
# exiting). Enqueueing now would drop the delta silently: no
|
||||
# writer will run and close() already unregistered the atexit
|
||||
# hook. Apply inline instead so a closed-connection failure
|
||||
# raises at the call site, exactly like the old synchronous
|
||||
# update_token_counts path these call sites still guard for.
|
||||
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.
|
||||
|
||||
Returns True when the queue is fully drained, False on timeout
|
||||
(callers then read totals that are stale by the still-queued
|
||||
deltas — no worse than reading before the flush existed).
|
||||
Never raises: apply failures are logged by the writer.
|
||||
"""
|
||||
# Fast path — nothing queued, nothing in flight.
|
||||
if not self._token_queue and not self._token_writer_busy:
|
||||
return True
|
||||
batch = None
|
||||
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
|
||||
# (close() in progress): its loop drains the queue before
|
||||
# exiting, and draining here instead would race its
|
||||
# in-flight batch — newer deltas committing before older
|
||||
# ones breaks the last-non-None-wins / first-accounted-
|
||||
# route / COALESCE-backfill fields. Only when the writer is
|
||||
# dead (or never started for these deltas) does the caller
|
||||
# take the leftovers. Re-checked each wakeup: the writer
|
||||
# can exit mid-wait with deltas enqueued after its final
|
||||
# empty-queue check. busy is claimed while draining (same
|
||||
# protocol as the writer) so a concurrent flush cannot
|
||||
# report drained — or pop a newer delta — while this batch
|
||||
# is still unapplied; a claimed busy therefore also 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
|
||||
):
|
||||
self._token_writer_busy = True
|
||||
batch = list(self._token_queue)
|
||||
self._token_queue.clear()
|
||||
break
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
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()
|
||||
return True
|
||||
|
||||
def _token_writer_loop(self) -> None:
|
||||
while True:
|
||||
with self._token_queue_cond:
|
||||
while not self._token_queue and not self._token_writer_stop:
|
||||
self._token_queue_cond.wait()
|
||||
if not self._token_queue:
|
||||
return # stop requested and fully drained
|
||||
# busy is set BEFORE the queue is cleared: the lock-free
|
||||
# fast path in flush_token_counts() reads queue-then-busy,
|
||||
# so this order guarantees it can never observe an empty
|
||||
# queue while the popped batch is still 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()
|
||||
|
||||
def _apply_token_batch(self, batch: List[Tuple[str, Dict[str, Any]]]) -> None:
|
||||
"""Apply queued deltas in order, coalescing where safe. Never raises."""
|
||||
for session_id, kwargs in self._coalesce_token_deltas(batch):
|
||||
try:
|
||||
self.update_token_counts(session_id, **kwargs)
|
||||
except Exception as exc:
|
||||
# Same contract as the old inline call sites: accounting
|
||||
# loss is logged, never raised into a turn.
|
||||
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 consecutive incremental deltas with an identical route.
|
||||
|
||||
Only adjacent deltas merge, so ordering across sessions and across
|
||||
a mid-session /model switch is preserved exactly. absolute=True
|
||||
deltas (cumulative overwrites) never merge.
|
||||
"""
|
||||
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
|
||||
)
|
||||
if groups and key is not None and groups[-1][0] == key:
|
||||
merged = groups[-1][2]
|
||||
for f in self._TOKEN_DELTA_SUM_FIELDS:
|
||||
merged[f] = merged.get(f, 0) + kwargs.get(f, 0)
|
||||
for f in self._TOKEN_DELTA_COST_FIELDS:
|
||||
value = kwargs.get(f)
|
||||
if value is not None:
|
||||
# None-preserving sum: an all-None run must stay
|
||||
# None so COALESCE keeps the stored value untouched.
|
||||
merged[f] = (merged.get(f) or 0.0) + value
|
||||
else:
|
||||
groups.append((key, session_id, dict(kwargs)))
|
||||
return [(sid, kw) for _, sid, kw in groups]
|
||||
|
||||
def _stop_token_writer(self, join_timeout: float = 10.0) -> None:
|
||||
"""Stop the writer thread and drain remaining deltas. Never raises."""
|
||||
with self._token_queue_cond:
|
||||
self._token_writer_stop = True
|
||||
self._token_queue_cond.notify_all()
|
||||
thread = self._token_writer_thread
|
||||
if thread is not None and thread.is_alive():
|
||||
thread.join(timeout=join_timeout)
|
||||
if thread.is_alive():
|
||||
# Writer stuck mid-apply (pathological lock contention).
|
||||
# Leave any queued deltas unapplied rather than racing the
|
||||
# stuck apply and misordering/double-counting.
|
||||
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 exited (or never started) — apply leftovers synchronously.
|
||||
# Claim busy like the writer/flush drains do, so a concurrent
|
||||
# flush_token_counts cannot fast-path True while this batch is
|
||||
# still being applied; conversely, 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:
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
logger.warning(
|
||||
"async token accounting: concurrent drain did not "
|
||||
"finish within %.0fs; %d queued delta(s) not persisted",
|
||||
join_timeout, len(self._token_queue),
|
||||
)
|
||||
return
|
||||
self._token_queue_cond.wait(remaining)
|
||||
batch = list(self._token_queue)
|
||||
self._token_queue.clear()
|
||||
if batch:
|
||||
self._token_writer_busy = True
|
||||
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()
|
||||
|
||||
def _drain_token_queue_at_exit(self) -> None:
|
||||
try:
|
||||
self._stop_token_writer()
|
||||
except Exception:
|
||||
pass # Best effort — never fatal at interpreter shutdown.
|
||||
|
||||
def update_token_counts(
|
||||
self,
|
||||
session_id: str,
|
||||
|
|
@ -5477,6 +5747,10 @@ class SessionDB:
|
|||
|
||||
def get_session(self, session_id: str) -> Optional[Dict[str, Any]]:
|
||||
"""Get a session by ID."""
|
||||
# Cost/usage readers (/status, /usage, gateway endpoints) reach the
|
||||
# row through here; drain queued token deltas so they see exact
|
||||
# totals. No-op attribute check when nothing is queued.
|
||||
self.flush_token_counts()
|
||||
with self._lock:
|
||||
cursor = self._conn.execute(
|
||||
"SELECT * FROM sessions WHERE id = ?", (session_id,)
|
||||
|
|
@ -6028,6 +6302,9 @@ class SessionDB:
|
|||
significant I/O saving on large databases where the blob routinely
|
||||
runs to tens of kilobytes per row.
|
||||
"""
|
||||
# Rows carry token/cost totals — drain queued deltas first so
|
||||
# listings (sidebar, /resume, dashboards) show exact counters.
|
||||
self.flush_token_counts()
|
||||
where_clauses = []
|
||||
params = []
|
||||
|
||||
|
|
@ -6327,6 +6604,8 @@ class SessionDB:
|
|||
Pass ``compact_rows=True`` to omit the ``system_prompt`` blob (see
|
||||
``list_sessions_rich`` for details).
|
||||
"""
|
||||
# Same read-your-writes guarantee as list_sessions_rich.
|
||||
self.flush_token_counts()
|
||||
_sel = self._compact_session_cols() if compact_rows else "s.*"
|
||||
query = f"""
|
||||
SELECT {_sel},
|
||||
|
|
|
|||
17
run_agent.py
17
run_agent.py
|
|
@ -1804,20 +1804,25 @@ class AIAgent:
|
|||
from agent.agent_runtime_helpers import note_turn_persisted
|
||||
|
||||
persist_lock = getattr(self, "_session_persist_lock", None)
|
||||
if persist_lock is None:
|
||||
|
||||
def _persist_and_drain() -> None:
|
||||
self._drop_trailing_empty_response_scaffolding(messages)
|
||||
self._session_messages = messages
|
||||
self._save_session_log(messages)
|
||||
self._flush_messages_to_session_db(messages, conversation_history)
|
||||
# Drain async token-accounting deltas at every persist point (turn
|
||||
# finalize + error exits) so a crash after this line loses at most
|
||||
# the in-flight API call's delta. Cheap no-op when nothing queued.
|
||||
if self._session_db is not None:
|
||||
self._session_db.flush_token_counts()
|
||||
note_turn_persisted(self)
|
||||
|
||||
if persist_lock is None:
|
||||
_persist_and_drain()
|
||||
return
|
||||
|
||||
with persist_lock:
|
||||
self._drop_trailing_empty_response_scaffolding(messages)
|
||||
self._session_messages = messages
|
||||
self._save_session_log(messages)
|
||||
self._flush_messages_to_session_db(messages, conversation_history)
|
||||
note_turn_persisted(self)
|
||||
_persist_and_drain()
|
||||
|
||||
def _drop_trailing_empty_response_scaffolding(self, messages: List[Dict]) -> None:
|
||||
"""Remove private empty-response retry/failure scaffolding from transcript tails.
|
||||
|
|
|
|||
564
tests/agent/test_async_token_accounting.py
Normal file
564
tests/agent/test_async_token_accounting.py
Normal file
|
|
@ -0,0 +1,564 @@
|
|||
"""Async token accounting — SessionDB background writer queue.
|
||||
|
||||
queue_token_counts() must take the per-call sessions UPDATE off the turn
|
||||
thread while preserving update_token_counts() semantics exactly:
|
||||
|
||||
1. Deltas apply in enqueue order.
|
||||
2. Coalescing consecutive same-route deltas is sum-equivalent to applying
|
||||
them one by one (sessions row AND session_model_usage breakdown).
|
||||
3. flush_token_counts() gives readers read-your-writes (get_session and
|
||||
friends call it), and turn finalize / close() drain the queue.
|
||||
4. A failing apply is logged by the writer and never raises into a turn.
|
||||
"""
|
||||
|
||||
import sqlite3
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from hermes_state import SessionDB
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
db_path = tmp_path / "test_state.db"
|
||||
session_db = SessionDB(db_path=db_path)
|
||||
yield session_db
|
||||
session_db.close()
|
||||
|
||||
|
||||
def _totals(db, session_id):
|
||||
"""Read token totals via raw SQL — bypasses get_session's flush so the
|
||||
read observes only what the writer has actually persisted."""
|
||||
with db._lock:
|
||||
row = db._conn.execute(
|
||||
"SELECT input_tokens, output_tokens, cache_read_tokens,"
|
||||
" cache_write_tokens, reasoning_tokens, api_call_count,"
|
||||
" estimated_cost_usd, actual_cost_usd, model, cost_status"
|
||||
" FROM sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
return dict(row) if row is not None else None
|
||||
|
||||
|
||||
def _model_usage(db, session_id):
|
||||
with db._lock:
|
||||
rows = db._conn.execute(
|
||||
"SELECT model, input_tokens, output_tokens, api_call_count,"
|
||||
" estimated_cost_usd FROM session_model_usage"
|
||||
" WHERE session_id = ? ORDER BY model",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Ordering
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestOrdering:
|
||||
def test_deltas_apply_in_enqueue_order(self, db):
|
||||
"""The writer applies deltas strictly in enqueue order, including
|
||||
across sessions (which never coalesce with each other)."""
|
||||
db.create_session("s-a", "test")
|
||||
db.create_session("s-b", "test")
|
||||
|
||||
applied = []
|
||||
original = db.update_token_counts
|
||||
|
||||
def recording(session_id, **kwargs):
|
||||
applied.append((session_id, kwargs.get("input_tokens", 0)))
|
||||
return original(session_id, **kwargs)
|
||||
|
||||
db.update_token_counts = recording
|
||||
try:
|
||||
expected = []
|
||||
for i in range(1, 7):
|
||||
sid = "s-a" if i % 2 else "s-b"
|
||||
db.queue_token_counts(sid, input_tokens=i, api_call_count=1)
|
||||
expected.append((sid, i))
|
||||
assert db.flush_token_counts()
|
||||
finally:
|
||||
db.update_token_counts = original
|
||||
|
||||
# Alternating sessions defeats coalescing, so every delta must be
|
||||
# applied individually, in order.
|
||||
assert applied == expected
|
||||
assert _totals(db, "s-a")["input_tokens"] == 1 + 3 + 5
|
||||
assert _totals(db, "s-b")["input_tokens"] == 2 + 4 + 6
|
||||
|
||||
def test_absolute_delta_is_an_ordering_barrier(self, db):
|
||||
"""incremental → absolute → incremental applies in order: the
|
||||
absolute overwrite wins over earlier increments, later increments
|
||||
stack on top of it."""
|
||||
db.create_session("s-abs", "test")
|
||||
db.queue_token_counts("s-abs", input_tokens=100, api_call_count=1)
|
||||
db.queue_token_counts(
|
||||
"s-abs", input_tokens=500, output_tokens=50,
|
||||
api_call_count=3, absolute=True,
|
||||
)
|
||||
db.queue_token_counts("s-abs", input_tokens=7, api_call_count=1)
|
||||
assert db.flush_token_counts()
|
||||
|
||||
totals = _totals(db, "s-abs")
|
||||
assert totals["input_tokens"] == 507
|
||||
assert totals["output_tokens"] == 50
|
||||
assert totals["api_call_count"] == 4
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Coalescing
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestCoalescing:
|
||||
def test_backlog_coalesces_and_sums_match(self, db):
|
||||
"""When a backlog forms, same-route deltas merge into fewer applies
|
||||
while totals stay exact."""
|
||||
db.create_session("s-c", "test")
|
||||
|
||||
apply_calls = []
|
||||
first_apply_started = threading.Event()
|
||||
release_first_apply = threading.Event()
|
||||
original = db.update_token_counts
|
||||
|
||||
def gated(session_id, **kwargs):
|
||||
apply_calls.append(kwargs)
|
||||
if len(apply_calls) == 1:
|
||||
first_apply_started.set()
|
||||
# Hold the writer inside its first apply so the remaining
|
||||
# enqueues pile up into one batch.
|
||||
assert release_first_apply.wait(timeout=10)
|
||||
return original(session_id, **kwargs)
|
||||
|
||||
db.update_token_counts = gated
|
||||
try:
|
||||
n = 20
|
||||
db.queue_token_counts(
|
||||
"s-c", input_tokens=1, output_tokens=1,
|
||||
estimated_cost_usd=0.001, model="m1",
|
||||
billing_provider="p1", api_call_count=1,
|
||||
)
|
||||
assert first_apply_started.wait(timeout=10)
|
||||
for _ in range(n - 1):
|
||||
db.queue_token_counts(
|
||||
"s-c", input_tokens=1, output_tokens=1,
|
||||
estimated_cost_usd=0.001, model="m1",
|
||||
billing_provider="p1", api_call_count=1,
|
||||
)
|
||||
release_first_apply.set()
|
||||
assert db.flush_token_counts()
|
||||
finally:
|
||||
db.update_token_counts = original
|
||||
|
||||
# The backlog collapses into far fewer UPDATEs than enqueues.
|
||||
assert len(apply_calls) < n
|
||||
totals = _totals(db, "s-c")
|
||||
assert totals["input_tokens"] == n
|
||||
assert totals["output_tokens"] == n
|
||||
assert totals["api_call_count"] == n
|
||||
assert totals["estimated_cost_usd"] == pytest.approx(0.001 * n)
|
||||
# Per-model attribution must also see the full sum.
|
||||
usage = _model_usage(db, "s-c")
|
||||
assert len(usage) == 1
|
||||
assert usage[0]["input_tokens"] == n
|
||||
assert usage[0]["api_call_count"] == n
|
||||
|
||||
def test_coalesced_apply_equals_sequential_apply(self, db, tmp_path):
|
||||
"""Applying a coalesced batch produces byte-identical session and
|
||||
per-model rows to applying the same deltas one at a time."""
|
||||
batch = [
|
||||
("s-eq", dict(input_tokens=10, output_tokens=2, model="m1",
|
||||
billing_provider="p1", estimated_cost_usd=0.01,
|
||||
cost_status="estimated", api_call_count=1)),
|
||||
("s-eq", dict(input_tokens=20, output_tokens=4, model="m1",
|
||||
billing_provider="p1", estimated_cost_usd=0.02,
|
||||
cost_status="estimated", api_call_count=1)),
|
||||
# /model switch mid-session — must not merge with the m1 run.
|
||||
("s-eq", dict(input_tokens=5, output_tokens=1, model="m2",
|
||||
billing_provider="p1", estimated_cost_usd=0.005,
|
||||
cost_status="estimated", api_call_count=1)),
|
||||
]
|
||||
|
||||
db.create_session("s-eq", "test")
|
||||
db._apply_token_batch(list(batch))
|
||||
|
||||
seq_db = SessionDB(db_path=tmp_path / "sequential.db")
|
||||
try:
|
||||
seq_db.create_session("s-eq", "test")
|
||||
for sid, kwargs in batch:
|
||||
seq_db.update_token_counts(sid, **kwargs)
|
||||
assert _totals(db, "s-eq") == _totals(seq_db, "s-eq")
|
||||
assert _model_usage(db, "s-eq") == _model_usage(seq_db, "s-eq")
|
||||
finally:
|
||||
seq_db.close()
|
||||
|
||||
def test_coalesce_unit_rules(self, db):
|
||||
"""_coalesce_token_deltas merge rules: same route merges, session /
|
||||
route changes and absolute deltas do not."""
|
||||
inc = dict(model="m1", billing_provider="p1")
|
||||
out = db._coalesce_token_deltas([
|
||||
("a", dict(input_tokens=1, api_call_count=1, **inc)),
|
||||
("a", dict(input_tokens=2, api_call_count=1, **inc)),
|
||||
("b", dict(input_tokens=4, api_call_count=1, **inc)),
|
||||
("a", dict(input_tokens=8, api_call_count=1, **inc)),
|
||||
("a", dict(input_tokens=16, api_call_count=1, model="m2",
|
||||
billing_provider="p1")),
|
||||
("a", dict(input_tokens=32, absolute=True)),
|
||||
("a", dict(input_tokens=64, absolute=True)),
|
||||
])
|
||||
assert [(sid, kw.get("input_tokens")) for sid, kw in out] == [
|
||||
("a", 3), # merged 1+2
|
||||
("b", 4), # session change
|
||||
("a", 8), # session change back
|
||||
("a", 16), # model change
|
||||
("a", 32), # absolute never merges
|
||||
("a", 64),
|
||||
]
|
||||
assert out[0][1]["api_call_count"] == 2
|
||||
|
||||
def test_coalesce_cost_none_preserved(self, db):
|
||||
"""An all-None cost run stays None after merging (COALESCE in the
|
||||
UPDATE must keep the stored value untouched)."""
|
||||
out = db._coalesce_token_deltas([
|
||||
("a", dict(input_tokens=1, estimated_cost_usd=None)),
|
||||
("a", dict(input_tokens=1, estimated_cost_usd=None)),
|
||||
])
|
||||
assert len(out) == 1
|
||||
assert out[0][1]["estimated_cost_usd"] is None
|
||||
|
||||
out = db._coalesce_token_deltas([
|
||||
("a", dict(input_tokens=1, estimated_cost_usd=None)),
|
||||
("a", dict(input_tokens=1, estimated_cost_usd=0.5)),
|
||||
])
|
||||
assert out[0][1]["estimated_cost_usd"] == pytest.approx(0.5)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Read-your-writes
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestReaderFlush:
|
||||
def test_get_session_sees_queued_deltas(self, db):
|
||||
"""get_session drains the queue first, so readers observe exact
|
||||
totals even while the writer is mid-backlog."""
|
||||
db.create_session("s-r", "test")
|
||||
|
||||
original = db.update_token_counts
|
||||
|
||||
def slow(session_id, **kwargs):
|
||||
time.sleep(0.05) # keep the writer visibly behind the reader
|
||||
return original(session_id, **kwargs)
|
||||
|
||||
db.update_token_counts = slow
|
||||
try:
|
||||
for i in range(1, 5):
|
||||
sid_tokens = i
|
||||
# Alternate models to defeat coalescing — four real applies.
|
||||
db.queue_token_counts(
|
||||
"s-r", input_tokens=sid_tokens,
|
||||
model=f"m{i % 2}", api_call_count=1,
|
||||
)
|
||||
row = db.get_session("s-r")
|
||||
finally:
|
||||
db.update_token_counts = original
|
||||
|
||||
assert row["input_tokens"] == 1 + 2 + 3 + 4
|
||||
assert row["api_call_count"] == 4
|
||||
|
||||
def test_flush_empty_queue_is_cheap_noop(self, db):
|
||||
assert db.flush_token_counts()
|
||||
# No writer thread was ever started by a bare flush.
|
||||
assert db._token_writer_thread is None
|
||||
|
||||
def test_flush_after_close_drains_on_caller_thread(self, db):
|
||||
"""After close() stops the writer, a late flush still drains queued
|
||||
deltas synchronously instead of losing them."""
|
||||
db.create_session("s-late", "test")
|
||||
db.flush_token_counts()
|
||||
db._stop_token_writer() # simulate a stopped writer with the conn open
|
||||
db._token_queue.append(("s-late", dict(input_tokens=9, api_call_count=1)))
|
||||
assert db.flush_token_counts()
|
||||
assert _totals(db, "s-late")["input_tokens"] == 9
|
||||
|
||||
def test_flush_waits_for_stop_flagged_live_writer(self, db):
|
||||
"""A stop-flagged but still-running writer owns the queue: flush must
|
||||
wait for it (its loop drains before exiting), never drain on the
|
||||
caller's thread — that would commit newer deltas before the writer's
|
||||
in-flight older batch and could return True with that batch
|
||||
unapplied."""
|
||||
db.create_session("s-stop", "test")
|
||||
|
||||
applied = []
|
||||
gate = threading.Event()
|
||||
first_apply_started = threading.Event()
|
||||
original = db.update_token_counts
|
||||
|
||||
def gated(session_id, **kwargs):
|
||||
applied.append(kwargs.get("input_tokens"))
|
||||
if len(applied) == 1:
|
||||
first_apply_started.set()
|
||||
assert gate.wait(timeout=10)
|
||||
return original(session_id, **kwargs)
|
||||
|
||||
db.update_token_counts = gated
|
||||
try:
|
||||
db.queue_token_counts("s-stop", input_tokens=1, api_call_count=1)
|
||||
assert first_apply_started.wait(timeout=10)
|
||||
# close() has set the stop flag but the writer is mid-apply.
|
||||
db._token_writer_stop = True
|
||||
db._token_queue.append(
|
||||
("s-stop", dict(input_tokens=2, api_call_count=1))
|
||||
)
|
||||
# The writer is alive, so flush waits — timing out, NOT applying
|
||||
# the newer delta on this thread ahead of the in-flight batch.
|
||||
assert db.flush_token_counts(timeout=0.3) is False
|
||||
assert applied == [1]
|
||||
gate.set()
|
||||
# Once released, the stop-flagged writer drains the queue itself
|
||||
# before exiting, preserving enqueue order.
|
||||
assert db.flush_token_counts()
|
||||
finally:
|
||||
db.update_token_counts = original
|
||||
|
||||
assert applied == [1, 2]
|
||||
assert _totals(db, "s-stop")["input_tokens"] == 3
|
||||
|
||||
def test_concurrent_flush_waits_for_caller_drain(self, db):
|
||||
"""The dead-writer caller-drain claims busy: a second flush must not
|
||||
report drained (fast path or locked path) while the first flusher's
|
||||
popped batch is still being applied outside the condition lock."""
|
||||
db.create_session("s-cc", "test")
|
||||
db.flush_token_counts()
|
||||
db._stop_token_writer() # writer dead, connection still open
|
||||
|
||||
applied = threading.Event()
|
||||
gate = threading.Event()
|
||||
original = db.update_token_counts
|
||||
|
||||
def gated(session_id, **kwargs):
|
||||
applied.set()
|
||||
assert gate.wait(timeout=10)
|
||||
return original(session_id, **kwargs)
|
||||
|
||||
db.update_token_counts = gated
|
||||
try:
|
||||
db._token_queue.append(
|
||||
("s-cc", dict(input_tokens=4, api_call_count=1))
|
||||
)
|
||||
results = {}
|
||||
t_a = threading.Thread(
|
||||
target=lambda: results.__setitem__(
|
||||
"a", db.flush_token_counts()
|
||||
)
|
||||
)
|
||||
t_a.start()
|
||||
assert applied.wait(timeout=10)
|
||||
# Flusher A is mid-apply with the queue already popped: B must
|
||||
# wait on the claimed busy flag, not return True.
|
||||
assert db.flush_token_counts(timeout=0.3) is False
|
||||
gate.set()
|
||||
t_a.join(timeout=10)
|
||||
assert results.get("a") is True
|
||||
assert db.flush_token_counts()
|
||||
finally:
|
||||
db.update_token_counts = original
|
||||
|
||||
assert _totals(db, "s-cc")["input_tokens"] == 4
|
||||
|
||||
def test_enqueue_after_writer_stop_applies_synchronously(self, db):
|
||||
"""Once the writer is stopped for good, queue_token_counts falls back
|
||||
to the synchronous path instead of parking deltas on a queue no
|
||||
writer will ever drain."""
|
||||
db.create_session("s-sync", "test")
|
||||
db.queue_token_counts("s-sync", input_tokens=1, api_call_count=1)
|
||||
db._stop_token_writer() # writer dead, connection still open
|
||||
|
||||
db.queue_token_counts("s-sync", input_tokens=2, api_call_count=1)
|
||||
|
||||
# Applied inline — nothing queued, no writer restarted.
|
||||
assert not db._token_queue
|
||||
assert db._token_writer_thread is None or not db._token_writer_thread.is_alive()
|
||||
totals = _totals(db, "s-sync")
|
||||
assert totals["input_tokens"] == 3
|
||||
assert totals["api_call_count"] == 2
|
||||
|
||||
def test_enqueue_after_close_raises_at_call_site(self, tmp_path):
|
||||
"""After close() the synchronous fallback surfaces the failure to the
|
||||
caller (whose try/except logs it) — the pre-queue contract — instead
|
||||
of silently dropping the delta."""
|
||||
db = SessionDB(db_path=tmp_path / "closed.db")
|
||||
db.create_session("s-closed", "test")
|
||||
db.queue_token_counts("s-closed", input_tokens=1, api_call_count=1)
|
||||
db.close()
|
||||
|
||||
with pytest.raises(Exception):
|
||||
db.queue_token_counts("s-closed", input_tokens=2, api_call_count=1)
|
||||
assert not db._token_queue # not parked on a dead queue either
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Ordering vs synchronous route writes (/model switch)
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestRouteSwitchBarrier:
|
||||
def test_model_switch_applies_queued_deltas_first(self, db):
|
||||
"""update_session_model / update_session_billing_route bypass the
|
||||
queue, so they must flush it first: a still-queued first delta
|
||||
carries the pre-switch route, and applying it after the switch
|
||||
UPDATE trips first_accounted_route (api_call_count == 0 + route
|
||||
mismatch) and resurrects the old model/provider on the row."""
|
||||
db.create_session("s-sw", "test")
|
||||
# First delta of the session, queued but not yet applied (writer
|
||||
# not started — same state as a backlogged writer).
|
||||
db._token_queue.append(("s-sw", dict(
|
||||
input_tokens=10, model="m1", billing_provider="p1",
|
||||
api_call_count=1,
|
||||
)))
|
||||
|
||||
db.update_session_model("s-sw", "m2")
|
||||
db.update_session_billing_route(
|
||||
"s-sw", provider="p2", base_url="https://p2.example"
|
||||
)
|
||||
|
||||
totals = _totals(db, "s-sw")
|
||||
# The switch wins on the session row…
|
||||
assert totals["model"] == "m2"
|
||||
assert _model_usage(db, "s-sw")[0]["model"] == "m1"
|
||||
with db._lock:
|
||||
row = db._conn.execute(
|
||||
"SELECT billing_provider FROM sessions WHERE id = ?",
|
||||
("s-sw",),
|
||||
).fetchone()
|
||||
assert row["billing_provider"] == "p2"
|
||||
# …and the queued delta was applied (before it), not dropped.
|
||||
assert totals["input_tokens"] == 10
|
||||
assert totals["api_call_count"] == 1
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Durability
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestDurability:
|
||||
def test_close_drains_queue(self, tmp_path):
|
||||
"""close() drains queued deltas before closing the connection, so a
|
||||
clean shutdown loses nothing."""
|
||||
db_path = tmp_path / "drain.db"
|
||||
db = SessionDB(db_path=db_path)
|
||||
db.create_session("s-d", "test")
|
||||
for i in range(5):
|
||||
db.queue_token_counts("s-d", input_tokens=10, api_call_count=1)
|
||||
db.close()
|
||||
|
||||
reopened = SessionDB(db_path=db_path)
|
||||
try:
|
||||
totals = _totals(reopened, "s-d")
|
||||
assert totals["input_tokens"] == 50
|
||||
assert totals["api_call_count"] == 5
|
||||
finally:
|
||||
reopened.close()
|
||||
|
||||
def test_atexit_drain_is_idempotent_and_never_raises(self, db):
|
||||
db.create_session("s-x", "test")
|
||||
db.queue_token_counts("s-x", input_tokens=3, api_call_count=1)
|
||||
db._drain_token_queue_at_exit()
|
||||
db._drain_token_queue_at_exit() # second call: writer already stopped
|
||||
assert _totals(db, "s-x")["input_tokens"] == 3
|
||||
|
||||
def test_close_unregisters_atexit_hook(self, tmp_path):
|
||||
"""close() must unregister the atexit drain hook: it holds a strong
|
||||
reference (bound method) that would otherwise pin every closed
|
||||
SessionDB — and its sqlite connection object — until interpreter
|
||||
exit in multi-open/close processes."""
|
||||
import gc
|
||||
import weakref
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "atexit.db")
|
||||
db.create_session("s-gc", "test")
|
||||
db.queue_token_counts("s-gc", input_tokens=1, api_call_count=1)
|
||||
db.close()
|
||||
|
||||
ref = weakref.ref(db)
|
||||
del db
|
||||
gc.collect()
|
||||
assert ref() is None
|
||||
|
||||
def test_persist_session_drains_queue(self, tmp_path, monkeypatch):
|
||||
"""Turn finalize (_persist_session) flushes the accounting queue —
|
||||
the crash window is at most the in-flight call's delta."""
|
||||
import os
|
||||
monkeypatch.setitem(os.environ, "OPENROUTER_API_KEY", "test-key")
|
||||
from run_agent import AIAgent
|
||||
|
||||
db = SessionDB(db_path=tmp_path / "finalize.db")
|
||||
try:
|
||||
agent = AIAgent(
|
||||
api_key="test-key",
|
||||
base_url="https://openrouter.ai/api/v1",
|
||||
model="test/model",
|
||||
quiet_mode=True,
|
||||
session_db=db,
|
||||
session_id="s-fin",
|
||||
skip_context_files=True,
|
||||
skip_memory=True,
|
||||
)
|
||||
agent._ensure_db_session()
|
||||
|
||||
db.queue_token_counts(
|
||||
"s-fin", input_tokens=11, output_tokens=2, api_call_count=1
|
||||
)
|
||||
agent._persist_session(
|
||||
[{"role": "user", "content": "q"}],
|
||||
[],
|
||||
)
|
||||
# Raw read: the flush happened inside _persist_session itself.
|
||||
totals = _totals(db, "s-fin")
|
||||
assert totals["input_tokens"] == 11
|
||||
assert totals["api_call_count"] == 1
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Failure isolation
|
||||
# =========================================================================
|
||||
|
||||
|
||||
class TestWriterFailure:
|
||||
def test_apply_failure_logs_and_does_not_raise(self, db, caplog):
|
||||
"""A failing UPDATE is logged by the writer; enqueue/flush never
|
||||
raise into the turn, and the writer survives to apply later deltas."""
|
||||
db.create_session("s-f", "test")
|
||||
|
||||
original = db.update_token_counts
|
||||
boom = {"raise": True}
|
||||
|
||||
def flaky(session_id, **kwargs):
|
||||
if boom["raise"]:
|
||||
raise sqlite3.OperationalError("database is locked")
|
||||
return original(session_id, **kwargs)
|
||||
|
||||
db.update_token_counts = flaky
|
||||
try:
|
||||
with caplog.at_level("WARNING", logger="hermes_state"):
|
||||
db.queue_token_counts("s-f", input_tokens=5, api_call_count=1)
|
||||
assert db.flush_token_counts()
|
||||
assert any(
|
||||
"async token accounting" in rec.getMessage()
|
||||
for rec in caplog.records
|
||||
)
|
||||
|
||||
# Writer thread survived the failure and keeps applying.
|
||||
boom["raise"] = False
|
||||
db.queue_token_counts("s-f", input_tokens=7, api_call_count=1)
|
||||
assert db.flush_token_counts()
|
||||
finally:
|
||||
db.update_token_counts = original
|
||||
|
||||
assert _totals(db, "s-f")["input_tokens"] == 7
|
||||
|
|
@ -50,8 +50,10 @@ def test_run_conversation_persists_tokens_for_telegram_sessions():
|
|||
result = agent.run_conversation("hello")
|
||||
|
||||
assert result["final_response"] == "done"
|
||||
session_db.update_token_counts.assert_called_once()
|
||||
assert session_db.update_token_counts.call_args.args[0] == "telegram-session"
|
||||
# Per-call deltas are enqueued for the SessionDB background writer
|
||||
# (queue_token_counts) rather than written inline on the turn thread.
|
||||
session_db.queue_token_counts.assert_called_once()
|
||||
assert session_db.queue_token_counts.call_args.args[0] == "telegram-session"
|
||||
|
||||
|
||||
def test_run_conversation_persists_tokens_for_cron_sessions():
|
||||
|
|
@ -61,8 +63,8 @@ def test_run_conversation_persists_tokens_for_cron_sessions():
|
|||
result = agent.run_conversation("hello")
|
||||
|
||||
assert result["final_response"] == "done"
|
||||
session_db.update_token_counts.assert_called_once()
|
||||
assert session_db.update_token_counts.call_args.args[0] == "cron-session"
|
||||
session_db.queue_token_counts.assert_called_once()
|
||||
assert session_db.queue_token_counts.call_args.args[0] == "cron-session"
|
||||
|
||||
|
||||
def test_session_search_lazily_opens_db_when_entrypoint_did_not_pass_one(monkeypatch):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue