From c0aaa238f6583293b420d059a79f14e3d72489cd Mon Sep 17 00:00:00 2001 From: 686f6c61 Date: Mon, 31 Aug 2026 18:51:09 +0200 Subject: [PATCH] feat(compression): usage anchor survives DB reloads and process restarts (salvage #99585) The usage anchor (real usage.prompt_tokens + delta estimate of what was appended since) identified the priced transcript by id() of the last message, so it was None on EVERY gateway turn (history is re-read from the DB each turn) and in every fresh process (--resume, desktop per-turn serve). Those are exactly the surfaces where the bytes/4 estimate then fired local compression against payloads the provider priced far under threshold (#99421, #104462). - agent/usage_anchor.py owns the anchor: content fingerprint instead of id(), persisted on the session row (model_config._usage_anchor) via set_usage_anchor(), restored on the first resumed turn while the durable transcript still matches, cleared with the row on compaction / codex-native rewrite / session reset. - Callers repointed from model_metadata (the compat table follows). Design and persistence slot from #99585 by @686f6c61; re-authored against the Sep 2026 layout (the branch predates the model_metadata / agent_init split). --- agent/agent_init.py | 2 +- agent/codex_runtime.py | 8 +- agent/context_breakdown.py | 3 +- agent/conversation_compression.py | 4 +- agent/conversation_loop.py | 4 +- agent/model_metadata.py | 48 ------ agent/turn_context.py | 8 +- agent/turn_request_assembly.py | 2 +- agent/turn_usage.py | 6 +- agent/usage_anchor.py | 165 +++++++++++++++++++ hermes_cli/cli_status_bar_mixin.py | 2 +- tests/agent/test_turn_base_display_anchor.py | 34 ++-- tests/agent/test_usage_anchor.py | 66 +++++--- 13 files changed, 246 insertions(+), 106 deletions(-) create mode 100644 agent/usage_anchor.py diff --git a/agent/agent_init.py b/agent/agent_init.py index 7e1e4239e0..5a5f22c1d8 100644 --- a/agent/agent_init.py +++ b/agent/agent_init.py @@ -2132,7 +2132,7 @@ def _init_usage_state(agent): _USAGE_STATE: Dict[str, Any] = { "_user_turn_count": 0, "_is_user_initiated_turn": False, # Copilot x-initiator: first call of a user turn = "user" - # Usage anchors (agent/model_metadata.py): last response's exact usage + transcript + # Usage anchors (agent/usage_anchor.py): last response's exact usage + transcript # snapshot; invalidated on compaction/session switch so stale anchors never suppress compression. "_usage_anchor": None, "_turn_base_usage_anchor": None, diff --git a/agent/codex_runtime.py b/agent/codex_runtime.py index eaf50c3467..c90e2d3852 100644 --- a/agent/codex_runtime.py +++ b/agent/codex_runtime.py @@ -14,6 +14,7 @@ from types import SimpleNamespace from typing import Any, Callable, Dict, List from agent.stream_single_writer import claim_stream_writer, stream_writer_is_current +from agent.usage_anchor import set_usage_anchor logger = logging.getLogger(__name__) _codex_watchdog_state_var: contextvars.ContextVar[Any | None] = contextvars.ContextVar( @@ -121,11 +122,11 @@ def _record_codex_app_server_usage(agent, turn, messages=None) -> dict[str, Any] except Exception: logger.debug("codex app-server usage update failed", exc_info=True) if isinstance(messages, list): - from agent.model_metadata import capture_usage_anchor + from agent.usage_anchor import capture_usage_anchor, set_usage_anchor anchor = capture_usage_anchor(prompt_tokens, canonical_usage.output_tokens, messages) if anchor is not None: - agent._usage_anchor = anchor + set_usage_anchor(agent, anchor) for key, value in usage_dict.items(): setattr(agent, f"session_{key}", getattr(agent, f"session_{key}") + value) cost_result = estimate_usage_cost( @@ -170,8 +171,7 @@ def _record_codex_app_server_compaction(agent, turn, *, approx_tokens: int | Non compressor.last_prompt_tokens, compressor.last_completion_tokens = -1, 0 compressor.awaiting_real_usage_after_compression = True # Provider-side context was rewritten; the usage anchor's transcript snapshot no longer matches. - agent._usage_anchor = None - agent._turn_base_usage_anchor = None + set_usage_anchor(agent, None) agent._last_compaction_in_place = False _call_guarded(getattr(agent, "event_callback", None) or None, "event_callback error on codex session:compress", args=("session:compress", { diff --git a/agent/context_breakdown.py b/agent/context_breakdown.py index 597bfc8218..2626bb612e 100644 --- a/agent/context_breakdown.py +++ b/agent/context_breakdown.py @@ -92,7 +92,8 @@ def _glyph(cat: Dict[str, Any]) -> str: def compute_session_context_breakdown(agent: Any, messages: Optional[List[dict]] = None) -> Dict[str, Any]: """Return a Cursor-style context usage breakdown for one live agent.""" - from agent.model_metadata import anchored_context_tokens, estimate_messages_tokens_rough + from agent.model_metadata import estimate_messages_tokens_rough + from agent.usage_anchor import anchored_context_tokens from agent.system_prompt import build_system_prompt_parts messages = messages or [] diff --git a/agent/conversation_compression.py b/agent/conversation_compression.py index 7cabe0f574..7a82999937 100644 --- a/agent/conversation_compression.py +++ b/agent/conversation_compression.py @@ -32,6 +32,7 @@ from agent.context_engine import automatic_compaction_status_message, sanitize_m from agent.memory_provider import PRE_COMPRESS_CHECKPOINT_API_VERSION from agent.model_metadata import estimate_messages_tokens_rough, estimate_request_tokens_rough from agent.session_activity import ActivityProvenance, normalize_activity_provenance +from agent.usage_anchor import set_usage_anchor logger = logging.getLogger(__name__) @@ -3114,8 +3115,7 @@ def _finish_compaction_boundary( compressor.awaiting_real_usage_after_compression = True # Transcript rewritten: invalidate the usage anchor's base snapshot explicitly # (its structural check would fail closed anyway); estimate until re-anchored. - agent._usage_anchor = None - agent._turn_base_usage_anchor = None + set_usage_anchor(agent, None) # Arm the effectiveness verdict only after a completed rewrite crosses the # boundary so later usage isn't charged to an attempt that changed nothing. if compression_made_progress: diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index 5bd41cc751..2e88c10327 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -1558,9 +1558,9 @@ _PLUGIN_COMPAT_LAZY = { 'PARTIAL_STREAM_STUB_ID': ('hermes_constants', 'PARTIAL_STREAM_STUB_ID'), 'PRE_API_COMPRESSION_STATUS_TEMPLATE': ('agent.conversation_compression', 'PRE_API_COMPRESSION_STATUS_TEMPLATE'), 'adaptive_rate_limit_backoff': ('agent.retry_utils', 'adaptive_rate_limit_backoff'), - 'anchored_context_tokens': ('agent.model_metadata', 'anchored_context_tokens'), + 'anchored_context_tokens': ('agent.usage_anchor', 'anchored_context_tokens'), 'automatic_compaction_status_message': ('agent.context_engine', 'automatic_compaction_status_message'), - 'capture_usage_anchor': ('agent.model_metadata', 'capture_usage_anchor'), + 'capture_usage_anchor': ('agent.usage_anchor', 'capture_usage_anchor'), 'classify_api_error': ('agent.error_classifier', 'classify_api_error'), 'close_interrupted_tool_sequence': ('agent.message_sanitization', 'close_interrupted_tool_sequence'), 'coalesce_tool_call_id': ('agent.message_sanitization', 'coalesce_tool_call_id'), diff --git a/agent/model_metadata.py b/agent/model_metadata.py index b5742d822d..1a13b76df7 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -2133,54 +2133,6 @@ def estimate_request_tokens_rough( return total -# Usage-anchored accounting: ``usage.prompt_tokens`` is EXACT for everything sent on that request, so -# anchoring shrinks chars/4 estimation to the messages appended since. Fields: prompt_tokens / -# completion_tokens (provider usage at capture); base_count (len(messages) at capture — the reply is -# not yet appended and is covered by completion_tokens, so the delta walk skips it at index base_count); -# base_last_id / base_last_role (identity of the last message; compaction/splices replace it -> full estimation). - - -def capture_usage_anchor(prompt_tokens: Any, completion_tokens: Any, messages: List[Dict[str, Any]]) -> Optional[Dict[str, Any]]: - """Build a usage anchor from provider-reported usage, or None.""" - try: - pt = int(prompt_tokens or 0) - ct = int(completion_tokens or 0) - except (TypeError, ValueError): - return None - if pt <= 0 or not isinstance(messages, list): - return None # no usable usage (some endpoints omit it) — caller keeps its anchor - last = messages[-1] if messages else None - return { - "prompt_tokens": pt, - "completion_tokens": max(0, ct), - "base_count": len(messages), - "base_last_id": id(last) if last is not None else None, - "base_last_role": last.get("role") if isinstance(last, dict) else None, - } - - -def anchored_context_tokens(messages: List[Dict[str, Any]], anchor: Optional[Dict[str, Any]], *, charge_stale_thinking: bool = True) -> Optional[int]: - """Anchored prompt+completion tokens plus a rough estimate of ONLY the messages appended since; - None when the anchor is missing or stale. The anchored response's own reply is skipped (already - in completion_tokens). ``charge_stale_thinking`` is forwarded to the delta estimate.""" - if not isinstance(anchor, dict) or not isinstance(messages, list): - return None - base_count = anchor.get("base_count") or 0 - if base_count <= 0 or len(messages) < base_count: - return None - base_msg = messages[base_count - 1] - base_role = base_msg.get("role") if isinstance(base_msg, dict) else None - if id(base_msg) != anchor.get("base_last_id") or base_role != anchor.get("base_last_role"): - return None - total = int(anchor["prompt_tokens"]) + int(anchor.get("completion_tokens") or 0) - delta = messages[base_count:] - if delta and isinstance(delta[0], dict) and delta[0].get("role") == "assistant": - delta = delta[1:] - if delta: - total += estimate_messages_tokens_rough(delta, charge_stale_thinking=charge_stale_thinking) - return total - - # Keyed by ``id(tools)``; bounded, oldest-first eviction. Repeated ``str(tools)`` on # large schemas stalls GUI event loops under GIL pressure. _TOOLS_TOKENS_CACHE: dict[int, Tuple[int, str, str, int]] = {} diff --git a/agent/turn_context.py b/agent/turn_context.py index 82bf522cb4..7d85aaf226 100644 --- a/agent/turn_context.py +++ b/agent/turn_context.py @@ -22,9 +22,8 @@ from agent.iteration_budget import IterationBudget from agent.memory_manager import build_memory_context_block from agent.memory_provider import is_trivial_prompt from agent.message_metadata import append_message, stamp_message_timestamp -from agent.model_metadata import ( - anchored_context_tokens, estimate_messages_tokens_rough, estimate_request_tokens_rough -) +from agent.model_metadata import estimate_messages_tokens_rough, estimate_request_tokens_rough +from agent.usage_anchor import anchored_context_tokens, restore_usage_anchor logger = logging.getLogger(__name__) @@ -536,6 +535,9 @@ def _hydrate_from_history(agent: Any, conversation_history: Optional[List[Any]]) # exact route/issuer/replay filtering and tolerate plugin compressors without the # optional hook. if agent._user_turn_count == 0: + # A fresh process has no in-memory anchor; the persisted one is honored only while the + # restored transcript still carries the priced prefix (see agent/usage_anchor.py). + restore_usage_anchor(agent, conversation_history) note_checkpoint = getattr( getattr(agent, "context_compressor", None), "note_native_compaction_checkpoint", diff --git a/agent/turn_request_assembly.py b/agent/turn_request_assembly.py index 8eb9beadb9..f090d7e81d 100644 --- a/agent/turn_request_assembly.py +++ b/agent/turn_request_assembly.py @@ -14,7 +14,7 @@ import logging from typing import Any from agent.message_sanitization import _sanitize_messages_surrogates -from agent.model_metadata import anchored_context_tokens +from agent.usage_anchor import anchored_context_tokens from agent.prompt_caching import build_prompt_cache_plan, effective_cache_ttl from agent.turn_context import build_api_messages diff --git a/agent/turn_usage.py b/agent/turn_usage.py index 63a507aeda..3c9e0e9494 100644 --- a/agent/turn_usage.py +++ b/agent/turn_usage.py @@ -15,7 +15,7 @@ from contextlib import suppress from dataclasses import dataclass from typing import Any, Dict, List -from agent.model_metadata import capture_usage_anchor +from agent.usage_anchor import capture_usage_anchor, set_usage_anchor from agent.usage_pricing import estimate_usage_cost, normalize_usage logger = logging.getLogger("agent.conversation_loop") @@ -120,9 +120,7 @@ def record_response_usage( aggregator_usage.prompt_tokens, aggregator_usage.output_tokens, messages ) if _new_anchor is not None: - agent._usage_anchor = _new_anchor - if api_call_count == 1: - agent._turn_base_usage_anchor = _new_anchor + set_usage_anchor(agent, _new_anchor, turn_base=api_call_count == 1) _compression_threshold = int(getattr(compressor, "threshold_tokens", 0) or 0) if _loop_mod()._should_rearm_compression_budget( compression_attempts, completed_compaction_pending=_completed_compaction_pending, diff --git a/agent/usage_anchor.py b/agent/usage_anchor.py new file mode 100644 index 0000000000..d93930c720 --- /dev/null +++ b/agent/usage_anchor.py @@ -0,0 +1,165 @@ +"""Usage-anchored token accounting: the provider's real ``usage.prompt_tokens`` is the only +authoritative context size; the local ``bytes/4`` estimate covers ONLY messages appended since. + +An anchor = provider usage at capture + a snapshot of the transcript position it priced: +``base_count`` (len(messages) at capture; the reply is not yet appended and is covered by +``completion_tokens``, so the delta walk skips an assistant row at that index), ``base_last_role`` +and ``base_last_fp`` (content fingerprint of the last priced message; compaction, splices and +rewinds replace it → anchor fails closed → full estimation until the next real reading). + +The fingerprint (not ``id()``) is the identity: the gateway re-reads the transcript from the DB +every turn and a resumed session runs in a fresh process, so object identity is never stable +across the surfaces where the estimate mattered most (#99421, #104462). The anchor also persists +on the session row (``model_config._usage_anchor``) so a restarted process can restore it; a +restored anchor is honored only while the durable transcript still matches its fingerprint. +""" + +from __future__ import annotations + +import hashlib +import json +import logging +from typing import Any, Dict, List, Optional + +logger = logging.getLogger(__name__) + +USAGE_ANCHOR_MODEL_CONFIG_KEY = "_usage_anchor" + +# Identity of a priced message = the provider-visible fields that round-trip the session DB +# byte-for-byte. Display/persistence metadata (timestamps, row ids, display kinds) is rewritten +# on reload and would only ever fail the match closed. +_FINGERPRINT_KEYS = ("role", "content", "api_content", "tool_call_id", "tool_calls") + + +def message_fingerprint(msg: Any) -> Optional[str]: + """Stable digest of one transcript message over its provider-visible, persisted fields.""" + if not isinstance(msg, dict): + return None + payload = {k: msg.get(k) for k in _FINGERPRINT_KEYS if msg.get(k) is not None} + try: + raw = json.dumps(payload, sort_keys=True, default=str, ensure_ascii=True, separators=(",", ":")) + except (TypeError, ValueError): + raw = repr(sorted(payload.items())) + return hashlib.sha256(raw.encode("utf-8", "replace")).hexdigest() + + +def capture_usage_anchor(prompt_tokens: Any, completion_tokens: Any, messages: List[Dict[str, Any]]) -> Optional[Dict[str, Any]]: + """Build a usage anchor from provider-reported usage, or None when usage is unusable.""" + try: + pt = int(prompt_tokens or 0) + ct = int(completion_tokens or 0) + except (TypeError, ValueError): + return None + if pt <= 0 or not isinstance(messages, list) or not messages: + return None # some endpoints omit usage — caller keeps its anchor + last = messages[-1] + return { + "prompt_tokens": pt, + "completion_tokens": max(0, ct), + "base_count": len(messages), + "base_last_role": last.get("role") if isinstance(last, dict) else None, + "base_last_fp": message_fingerprint(last), + } + + +def _anchor_matches(messages: List[Dict[str, Any]], anchor: Dict[str, Any]) -> bool: + try: + base_count = int(anchor.get("base_count") or 0) + except (TypeError, ValueError): + return False + if base_count <= 0 or len(messages) < base_count: + return False + base_msg = messages[base_count - 1] + if not isinstance(base_msg, dict) or base_msg.get("role") != anchor.get("base_last_role"): + return False + fp = anchor.get("base_last_fp") + return isinstance(fp, str) and bool(fp) and message_fingerprint(base_msg) == fp + + +def anchored_context_tokens(messages: List[Dict[str, Any]], anchor: Optional[Dict[str, Any]], *, charge_stale_thinking: bool = True) -> Optional[int]: + """Anchored prompt+completion tokens plus a rough estimate of ONLY the messages appended since; + None when the anchor is missing or stale. The anchored response's own reply is skipped (already + in completion_tokens). ``charge_stale_thinking`` is forwarded to the delta estimate.""" + if not isinstance(anchor, dict) or not isinstance(messages, list) or not _anchor_matches(messages, anchor): + return None + from agent.model_metadata import estimate_messages_tokens_rough + + total = int(anchor["prompt_tokens"]) + int(anchor.get("completion_tokens") or 0) + delta = messages[int(anchor["base_count"]):] + if delta and isinstance(delta[0], dict) and delta[0].get("role") == "assistant": + delta = delta[1:] + if delta: + total += estimate_messages_tokens_rough(delta, charge_stale_thinking=charge_stale_thinking) + return total + + +def _serialize(anchor: Any) -> Optional[Dict[str, Any]]: + if not isinstance(anchor, dict): + return None + try: + pt, ct, base_count = (int(anchor.get(k) or 0) for k in ("prompt_tokens", "completion_tokens", "base_count")) + except (TypeError, ValueError): + return None + fp, role = anchor.get("base_last_fp"), anchor.get("base_last_role") + if pt <= 0 or base_count <= 0 or not isinstance(fp, str) or not fp: + return None + return {"prompt_tokens": pt, "completion_tokens": max(0, ct), "base_count": base_count, + "base_last_role": role if isinstance(role, str) else None, "base_last_fp": fp} + + +def persist_usage_anchor(agent: Any, anchor: Optional[Dict[str, Any]]) -> None: + """Write (or clear, ``None``) the session row's anchor blob. Best-effort: the row may not exist yet.""" + if getattr(agent, "_persist_disabled", False): + return + session_id = getattr(agent, "session_id", None) + patcher = getattr(getattr(agent, "_session_db", None), "patch_session_model_config", None) + if not session_id or not callable(patcher): + return + try: + patcher(session_id, {USAGE_ANCHOR_MODEL_CONFIG_KEY: _serialize(anchor)}) + except Exception: + logger.debug("usage anchor persist failed", exc_info=True) + + +def set_usage_anchor(agent: Any, anchor: Optional[Dict[str, Any]], *, turn_base: bool = False) -> None: + """Install ``anchor`` on the agent (``None`` clears) and mirror it to the session row.""" + agent._usage_anchor = anchor + if turn_base or anchor is None: + agent._turn_base_usage_anchor = anchor + persist_usage_anchor(agent, anchor) + + +def restore_usage_anchor(agent: Any, conversation_history: Optional[List[Dict[str, Any]]]) -> None: + """On a resumed session, adopt the persisted anchor when ``conversation_history`` still carries + the priced prefix; otherwise clear the stale blob so it can never suppress compression.""" + if getattr(agent, "_usage_anchor", None) is not None or getattr(agent, "_persist_disabled", False): + return + session_id = getattr(agent, "session_id", None) + getter = getattr(getattr(agent, "_session_db", None), "get_session_model_config_value", None) + if not session_id or not callable(getter) or not isinstance(conversation_history, list): + return + try: + anchor = _serialize(getter(session_id, USAGE_ANCHOR_MODEL_CONFIG_KEY, None)) + except Exception: + logger.debug("usage anchor load failed", exc_info=True) + return + if anchor is None: + return + if _anchor_matches(conversation_history, anchor): + agent._usage_anchor = anchor + else: + persist_usage_anchor(agent, None) + + +def persisted_anchor_tokens(session_db: Any, session_id: Any, messages: Any) -> Optional[int]: + """Anchored token figure from the session row's persisted anchor, for callers without a live + agent (gateway hygiene); None when absent, unreadable, or stale against ``messages``.""" + getter = getattr(session_db, "get_session_model_config_value", None) + if not session_id or not callable(getter) or not isinstance(messages, list): + return None + try: + anchor = _serialize(getter(session_id, USAGE_ANCHOR_MODEL_CONFIG_KEY, None)) + except Exception: + logger.debug("usage anchor load failed", exc_info=True) + return None + return anchored_context_tokens(messages, anchor) if anchor else None diff --git a/hermes_cli/cli_status_bar_mixin.py b/hermes_cli/cli_status_bar_mixin.py index 8950ec82cd..78e2914b53 100644 --- a/hermes_cli/cli_status_bar_mixin.py +++ b/hermes_cli/cli_status_bar_mixin.py @@ -279,7 +279,7 @@ class CLIStatusBarMixin: # Anchor on the turn's FIRST response plus a delta estimate of appended messages. # The compression trigger keeps using real last-request usage. try: - from agent.model_metadata import anchored_context_tokens + from agent.usage_anchor import anchored_context_tokens _msgs = getattr(agent, "_session_messages", None) _anchored = anchored_context_tokens( diff --git a/tests/agent/test_turn_base_display_anchor.py b/tests/agent/test_turn_base_display_anchor.py index ac78ab4184..942b721848 100644 --- a/tests/agent/test_turn_base_display_anchor.py +++ b/tests/agent/test_turn_base_display_anchor.py @@ -20,11 +20,8 @@ Covers: from types import SimpleNamespace -from agent.model_metadata import ( - anchored_context_tokens, - capture_usage_anchor, - estimate_messages_tokens_rough, -) +from agent.model_metadata import estimate_messages_tokens_rough +from agent.usage_anchor import anchored_context_tokens, capture_usage_anchor def _msg(role, content, **extra): @@ -177,21 +174,20 @@ class TestContextBreakdownPrefersTurnBaseAnchor: class TestInvalidationSitesClearTurnBaseAnchor: - def test_compression_invalidation_clears_both(self): - import inspect - from agent import conversation_compression + def test_clearing_the_anchor_clears_the_turn_base_too(self): + """Compaction and the codex-native rewrite clear via set_usage_anchor(None): both the + last-response anchor and the display turn-base anchor go; a later same-turn capture + (turn_base=False) leaves the turn-base untouched.""" + from agent.usage_anchor import set_usage_anchor - src = inspect.getsource(conversation_compression) - block = src.split("agent._usage_anchor = None", 1)[1][:200] - assert "_turn_base_usage_anchor = None" in block - - def test_codex_native_invalidation_clears_both(self): - import inspect - from agent import codex_runtime - - src = inspect.getsource(codex_runtime) - block = src.split("agent._usage_anchor = None", 1)[1][:200] - assert "_turn_base_usage_anchor = None" in block + messages = [_msg("user", "a"), _msg("assistant", "b")] + agent = SimpleNamespace(_usage_anchor=None, _turn_base_usage_anchor=None, _session_db=None, session_id=None) + first = capture_usage_anchor(1_000, 10, messages) + set_usage_anchor(agent, first, turn_base=True) + set_usage_anchor(agent, capture_usage_anchor(2_000, 10, messages)) + assert agent._turn_base_usage_anchor is first + set_usage_anchor(agent, None) + assert agent._usage_anchor is None and agent._turn_base_usage_anchor is None def test_agent_init_defines_turn_base_anchor(self): import inspect diff --git a/tests/agent/test_usage_anchor.py b/tests/agent/test_usage_anchor.py index 71645e3055..12ac3585a6 100644 --- a/tests/agent/test_usage_anchor.py +++ b/tests/agent/test_usage_anchor.py @@ -1,4 +1,4 @@ -"""Usage-anchored context accounting (agent/model_metadata.py). +"""Usage-anchored context accounting (agent/usage_anchor.py). Context-size checks anchor on the provider-reported ``usage.prompt_tokens`` of the last main-loop response and estimate ONLY the messages appended @@ -9,8 +9,10 @@ since. These tests cover: heuristic vs provider truth); * fallback to full estimation when no anchor exists (first request, usage-less providers); - * invalidation when compaction rewrites the transcript (structural - id/index check fails closed) and on explicit reset sites; + * invalidation when compaction rewrites the transcript (content fingerprint + fails closed) while a DB-reloaded transcript with the same content still + matches (the gateway re-reads history every turn); + * persistence on the session row and restore in a fresh process; * the preflight consumer (_preflight_request_tokens) preferring the anchor, plus a sabotage check proving the anchored path (not the heuristic) produces the number. @@ -20,12 +22,14 @@ from types import SimpleNamespace import pytest -from agent.model_metadata import ( +from agent.model_metadata import estimate_messages_tokens_rough +from agent.turn_context import _preflight_request_tokens +from agent.usage_anchor import ( anchored_context_tokens, capture_usage_anchor, - estimate_messages_tokens_rough, + restore_usage_anchor, + set_usage_anchor, ) -from agent.turn_context import _preflight_request_tokens def _msg(role, content): @@ -47,6 +51,10 @@ def _image_msg(): } +def _plain_history(): + return [_msg("user", "start"), _msg("assistant", "hello"), _msg("user", "do the thing"), _msg("assistant", "done")] + + def _history_with_images(n_images=10): msgs = [_msg("user", "start")] for i in range(n_images): @@ -124,25 +132,43 @@ class TestAnchorInvalidation: spliced = messages[:1] + [_msg("assistant", "[marker]")] + messages[5:] assert anchored_context_tokens(spliced, anchor) is None - def test_same_length_different_objects_fails_closed(self): + def test_reloaded_transcript_with_same_content_still_matches(self): + """The gateway re-reads history from the DB every turn (fresh dicts, extra + persistence keys); identity must survive that or every gateway turn falls + back to the whole-history estimate.""" messages = _history_with_images(4) anchor = capture_usage_anchor(30_000, 50, messages) - rebuilt = [dict(m) for m in messages] # fresh dicts, same values - assert anchored_context_tokens(rebuilt, anchor) is None + reloaded = [dict(m, timestamp=1.0, _row_id=i) for i, m in enumerate(messages)] + assert anchored_context_tokens(reloaded, anchor) == 30_050 + edited = [dict(m) for m in messages] + edited[-1] = dict(edited[-1], content="different last message") + assert anchored_context_tokens(edited, anchor) is None - def test_explicit_invalidation_sites(self): - """The compaction + session-reset sites null agent._usage_anchor.""" - import inspect + def test_persist_and_restore_across_processes(self, tmp_path): + """A fresh agent (desktop per-turn ``serve``, ``--resume``) adopts the persisted anchor + while the durable transcript still matches, and clears it once it does not.""" + from hermes_state import SessionDB - import agent.conversation_compression as cc - import agent.codex_runtime as cr - import run_agent + db = SessionDB(db_path=tmp_path / "state.db") + sid = "anchor-restore" + db.create_session(sid, source="cli") + messages = _plain_history() + for m in messages: + db.append_message(sid, m["role"], m["content"]) + durable = db.get_messages_as_conversation(sid) + live = SimpleNamespace(session_id=sid, _session_db=db, _persist_disabled=False, _usage_anchor=None) + set_usage_anchor(live, capture_usage_anchor(10_000, 20, durable)) - assert "agent._usage_anchor = None" in inspect.getsource(cc) - assert "agent._usage_anchor = None" in inspect.getsource(cr) - assert "self._usage_anchor = None" in inspect.getsource( - run_agent.AIAgent.reset_session_state - ) + fresh = SimpleNamespace(session_id=sid, _session_db=db, _persist_disabled=False, _usage_anchor=None) + restore_usage_anchor(fresh, db.get_messages_as_conversation(sid)) + assert fresh._usage_anchor is not None + assert anchored_context_tokens(db.get_messages_as_conversation(sid), fresh._usage_anchor) == 10_020 + + set_usage_anchor(live, None) # compaction / reset clears the row too + stale = SimpleNamespace(session_id=sid, _session_db=db, _persist_disabled=False, _usage_anchor=None) + restore_usage_anchor(stale, durable) + assert stale._usage_anchor is None + db.close() class TestPreflightConsumer: