From f42fc3c5c7bedcc873cccbe4e7187da8551e04ae Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 21:49:03 -0700 Subject: [PATCH] refactor(agent): compact token estimation helpers (shadow, fingerprint, tools cache) --- agent/model_metadata.py | 190 ++++++++++++++-------------------------- 1 file changed, 67 insertions(+), 123 deletions(-) diff --git a/agent/model_metadata.py b/agent/model_metadata.py index 9c494b2673..10ada00b60 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -1924,64 +1924,50 @@ def _is_cjk_token_dense_char(ch: str) -> bool: def estimate_tokens_rough(text: str) -> int: - """Rough token estimate: ceil(chars/4), CJK/Hangul/Kana codepoints ~1 token each. - Ceiling keeps short texts from estimating 0. Runs on every preflight walk, so the - all-ASCII case stays O(1) (``str.isascii()`` is a flag check on CPython).""" + """Rough token estimate: ceil(chars/4), CJK/Hangul/Kana codepoints ~1 token each. Ceiling keeps + short texts from estimating 0. Runs on every preflight walk, so the all-ASCII case stays O(1).""" if not text: return 0 text = str(text) - if text.isascii(): - return (len(text) + 3) // 4 - dense = len(text) - len(_CJK_DENSE_RE.sub("", text)) - if not dense: # non-ASCII but no CJK (accents, Cyrillic, emoji) - return (len(text) + 3) // 4 + # ``str.isascii()`` is a flag check on CPython; non-ASCII without CJK (accents, Cyrillic, emoji) also gets chars/4. + dense = 0 if text.isascii() else len(text) - len(_CJK_DENSE_RE.sub("", text)) return dense + ((len(text) - dense + 3) // 4) def estimate_messages_tokens_rough(messages: List[Dict[str, Any]], *, charge_stale_thinking: bool = True) -> int: - """Rough token estimate for a message list (pre-flight only). Images cost a flat - ~1500 tokens each rather than their base64 length (a 1MB screenshot ≈ 250K). - ``charge_stale_thinking=False`` mirrors the tail-budget walk - (``context_compressor._estimate_msg_budget_tokens``): on non-echo routes stale - reasoning rides the wire only for the NEWEST assistant turn, so excluding it - keeps the compaction TRIGGER in the same size class as the walk — otherwise - reasoning-heavy sessions fire preflight forever while the walk finds nothing. - """ + """Rough token estimate for a message list (pre-flight only). Images cost a flat ~1500 tokens + each rather than their base64 length. ``charge_stale_thinking=False`` mirrors the tail-budget + walk (``context_compressor._estimate_msg_budget_tokens``): on non-echo routes stale reasoning + rides the wire only for the NEWEST assistant turn, so excluding it keeps the compaction TRIGGER + in the same size class as the walk — otherwise reasoning-heavy sessions fire preflight forever.""" _IMAGE_TOKEN_COST = 1500 if not charge_stale_thinking: messages = _strip_stale_thinking_for_estimate(messages) return sum(_estimate_message_tokens_cached(msg, _IMAGE_TOKEN_COST) for msg in messages) -# Generic thinking-text keys replayed for at most the newest assistant turn -# on non-echo routes — must stay in lockstep with -# ``context_compressor._NEWEST_TURN_ONLY_BUDGET_KEYS``. +# Thinking-text keys replayed for at most the newest assistant turn on non-echo routes — must stay +# in lockstep with ``context_compressor._NEWEST_TURN_ONLY_BUDGET_KEYS``. _STALE_THINKING_ESTIMATE_KEYS = ("reasoning", "reasoning_content") def _strip_stale_thinking_for_estimate(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: - """Copy of ``messages`` with stale thinking keys removed (newest kept). - Shallow stripped copies share the original value objects, so the - per-message memo still hits for the stripped shape on subsequent walks. - """ - newest = next( - (i for i in range(len(messages) - 1, -1, -1) - if isinstance(messages[i], dict) and messages[i].get("role") == "assistant"), - -1, - ) - out: List[Dict[str, Any]] = [] - for i, m in enumerate(messages): - if i != newest and isinstance(m, dict) and m.get("role") == "assistant" and any(m.get(k) for k in _STALE_THINKING_ESTIMATE_KEYS): - m = {k: v for k, v in m.items() if k not in _STALE_THINKING_ESTIMATE_KEYS} - out.append(m) - return out + """Copy of ``messages`` with stale thinking keys removed (newest kept). Shallow stripped copies + share the original value objects, so the per-message memo still hits for the stripped shape.""" + def _is_assistant(m: Any) -> bool: + return isinstance(m, dict) and m.get("role") == "assistant" + newest = next((i for i in range(len(messages) - 1, -1, -1) if _is_assistant(messages[i])), -1) + return [ + {k: v for k, v in m.items() if k not in _STALE_THINKING_ESTIMATE_KEYS} + if i != newest and _is_assistant(m) and any(m.get(k) for k in _STALE_THINKING_ESTIMATE_KEYS) else m + for i, m in enumerate(messages) + ] -# Per-message token-estimate memo keyed by an exact value fingerprint: strings by -# ``id()`` AND pinned (strong ref in the entry, so the id can't be reused and -# immutability makes id-equality value-equality); numbers/bools/None by value; -# dicts/lists structurally in key order (``str(shadow)`` depends on it); any other -# type aborts the memo. api_messages shallow-copies dicts but shares the strings. +# Per-message token-estimate memo keyed by an exact value fingerprint: strings by ``id()`` AND +# pinned (strong ref in the entry, so the id can't be reused and immutability makes id-equality +# value-equality); numbers/bools/None by value; dicts/lists structurally in key order (``str(shadow)`` +# depends on it); any other type aborts the memo. api_messages shallow-copies dicts but shares the strings. _MSG_TOKENS_CACHE: Dict[Any, Tuple[list, int]] = {} _MSG_TOKENS_CACHE_MAX = 4096 @@ -1997,10 +1983,8 @@ def _msg_fingerprint(value: Any, pins: list) -> Any: return ("n", t.__name__, value) if t is dict: return ("d", tuple((_msg_fingerprint(k, pins), _msg_fingerprint(v, pins)) for k, v in value.items())) - if t is list: - return ("l", tuple(_msg_fingerprint(v, pins) for v in value)) - if t is tuple: - return ("t", tuple(_msg_fingerprint(v, pins) for v in value)) + if t is list or t is tuple: + return ("l" if t is list else "t", tuple(_msg_fingerprint(v, pins) for v in value)) raise ValueError("unfingerprintable message value") @@ -2027,9 +2011,7 @@ def _estimate_message_tokens_cached(msg: Any, image_cost: int) -> int: def _count_parts(parts: Any, types: set) -> int: - if not isinstance(parts, list): - return 0 - return sum(1 for part in parts if isinstance(part, dict) and part.get("type") in types) + return sum(1 for part in parts if isinstance(part, dict) and part.get("type") in types) if isinstance(parts, list) else 0 def _count_image_tokens(msg: Dict[str, Any], cost_per_image: int) -> int: @@ -2047,39 +2029,33 @@ def _count_image_tokens(msg: Dict[str, Any], cost_per_image: int) -> int: def _wire_message_shadow(msg: Dict[str, Any]) -> Dict[str, Any]: """Shadow of a message holding only what the provider actually receives. - * ``api_content`` SUBSTITUTES ``content`` (mirrors ``turn_context.substitute_api_content`` - exactly): only a non-empty STRING sidecar on a user/assistant row displaces - content; substituting any other shape would UNDERcount — the dangerous direction. + * ``api_content`` SUBSTITUTES ``content`` (mirrors ``turn_context.substitute_api_content`` exactly): + only a non-empty STRING sidecar on a user/assistant row displaces content; substituting any + other shape would UNDERcount — the dangerous direction. * Base64 images become a placeholder; ``_count_image_tokens`` charges them flat. - * ``reasoning`` never ships as-is (request builds pop it after optionally promoting - it into ``reasoning_content``); counting both inflated estimates up to +53%. - """ + * ``reasoning`` never ships as-is (request builds pop it after optionally promoting it into + ``reasoning_content``); counting both inflated estimates up to +53%.""" sidecar = msg.get("api_content") sidecar_wins = isinstance(sidecar, str) and bool(sidecar) and msg.get("role") in ("user", "assistant") _rc = msg.get("reasoning_content") drop_reasoning_dup = isinstance(_rc, str) and bool(_rc.strip()) shadow: Dict[str, Any] = {} for k, v in msg.items(): - if k in ("_anthropic_content_blocks", "reasoning_details") or k in PERSISTENCE_ONLY_MESSAGE_FIELDS: - continue - if k == "reasoning" and drop_reasoning_dup: + if k in ("_anthropic_content_blocks", "reasoning_details") or k in PERSISTENCE_ONLY_MESSAGE_FIELDS or (k == "reasoning" and drop_reasoning_dup): continue if k == "api_content": if sidecar_wins: shadow["content"] = v - elif k == "content": - if sidecar_wins: - continue - if isinstance(v, list): - shadow[k] = [ - {"type": part.get("type"), "image": "[stripped]"} - if isinstance(part, dict) and part.get("type") in {"image", "image_url", "input_image"} else part - for part in v - ] - elif isinstance(v, dict) and v.get("_multimodal"): - shadow[k] = v.get("text_summary", "") - else: - shadow[k] = v + elif k == "content" and sidecar_wins: + continue + elif k == "content" and isinstance(v, list): + shadow[k] = [ + {"type": part.get("type"), "image": "[stripped]"} + if isinstance(part, dict) and part.get("type") in {"image", "image_url", "input_image"} else part + for part in v + ] + elif k == "content" and isinstance(v, dict) and v.get("_multimodal"): + shadow[k] = v.get("text_summary", "") else: shadow[k] = v return shadow @@ -2087,44 +2063,29 @@ def _wire_message_shadow(msg: Dict[str, Any]) -> Dict[str, Any]: def _estimate_message_tokens_without_images(msg: Dict[str, Any]) -> int: """Token estimate for a message shadow with image payloads stripped.""" - if not isinstance(msg, dict): - return estimate_tokens_rough(str(msg)) - return estimate_tokens_rough(str(_wire_message_shadow(msg))) + return estimate_tokens_rough(str(_wire_message_shadow(msg) if isinstance(msg, dict) else msg)) def estimate_request_tokens_rough( - messages: List[Dict[str, Any]], - *, - system_prompt: str = "", - tools: Optional[List[Dict[str, Any]]] = None, - charge_stale_thinking: bool = True, + messages: List[Dict[str, Any]], *, system_prompt: str = "", tools: Optional[List[Dict[str, Any]]] = None, charge_stale_thinking: bool = True, ) -> int: - """Rough token estimate for a full request: system prompt + messages + tool - schemas (50+ tools add 20-30K on their own). ``charge_stale_thinking`` - is forwarded — pass False when the route provably strips stale thinking - (``message_sanitization.stale_thinking_reaches_wire``).""" - total = 0 - if system_prompt: - total += estimate_tokens_rough(system_prompt) + """Rough token estimate for a full request: system prompt + messages + tool schemas (50+ tools + add 20-30K on their own). ``charge_stale_thinking`` is forwarded — pass False when the route + provably strips stale thinking (``message_sanitization.stale_thinking_reaches_wire``).""" + total = estimate_tokens_rough(system_prompt) if system_prompt else 0 if messages: - if charge_stale_thinking: - # Positional call: test seams and plugin engines monkeypatch - # estimate_messages_tokens_rough with (messages)-only signatures. - total += estimate_messages_tokens_rough(messages) - else: - total += estimate_messages_tokens_rough(messages, charge_stale_thinking=False) + # Positional call: test seams and plugin engines monkeypatch estimate_messages_tokens_rough with (messages)-only signatures. + total += estimate_messages_tokens_rough(messages) if charge_stale_thinking else estimate_messages_tokens_rough(messages, charge_stale_thinking=False) if tools: total += _estimate_tools_tokens_rough(tools) 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 and fall back to full estimation). +# 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]]: @@ -2146,26 +2107,18 @@ def capture_usage_anchor(prompt_tokens: Any, completion_tokens: Any, messages: L } -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.""" +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] - if id(base_msg) != anchor.get("base_last_id"): - return None base_role = base_msg.get("role") if isinstance(base_msg, dict) else None - if base_role != anchor.get("base_last_role"): + 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:] @@ -2187,9 +2140,7 @@ def _tool_name_for_cache(tool: Any) -> str: return "" fn = tool.get("function") name = fn.get("name") if isinstance(fn, dict) else None - if isinstance(name, str): - return name - name = tool.get("name") + name = name if isinstance(name, str) else tool.get("name") return name if isinstance(name, str) else "" @@ -2197,11 +2148,9 @@ def _estimate_tools_tokens_rough(tools: List[Dict[str, Any]]) -> int: if not tools: return 0 key = id(tools) - n = len(tools) - first = _tool_name_for_cache(tools[0]) - last = _tool_name_for_cache(tools[-1]) + signature = (len(tools), _tool_name_for_cache(tools[0]), _tool_name_for_cache(tools[-1])) cached = _TOOLS_TOKENS_CACHE.get(key) - if cached is not None and cached[:3] == (n, first, last): + if cached is not None and cached[:3] == signature: return cached[3] # Sum the major schema fields (descriptions + parameters dominate). total_chars = 0 @@ -2210,13 +2159,8 @@ def _estimate_tools_tokens_rough(tools: List[Dict[str, Any]]) -> int: continue fn = tool.get("function") src = fn if isinstance(fn, dict) else tool - name = src.get("name") or "" - desc = src.get("description") or "" params = src.get("parameters") or {} - if isinstance(name, str): - total_chars += len(name) - if isinstance(desc, str): - total_chars += len(desc) + total_chars += sum(len(v) for v in (src.get("name") or "", src.get("description") or "") if isinstance(v, str)) try: # JSON is closer to wire size than repr() total_chars += len(json.dumps(params, ensure_ascii=False, separators=(",", ":"))) except Exception: @@ -2224,5 +2168,5 @@ def _estimate_tools_tokens_rough(tools: List[Dict[str, Any]]) -> int: tokens = (total_chars + 3) // 4 if len(_TOOLS_TOKENS_CACHE) >= _TOOLS_TOKENS_CACHE_MAX: _TOOLS_TOKENS_CACHE.pop(next(iter(_TOOLS_TOKENS_CACHE)), None) - _TOOLS_TOKENS_CACHE[key] = (n, first, last, tokens) + _TOOLS_TOKENS_CACHE[key] = (*signature, tokens) return tokens