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).
This commit is contained in:
+1
-1
@@ -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,
|
||||
|
||||
@@ -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", {
|
||||
|
||||
@@ -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 []
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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'),
|
||||
|
||||
@@ -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]] = {}
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+2
-4
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user