refactor(agent): compact token estimation helpers (shadow, fingerprint, tools cache)

This commit is contained in:
Teknium
2026-09-02 21:49:03 -07:00
parent 11e0f3862d
commit f42fc3c5c7

View File

@@ -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