Skip to content

Commit 7c2f0ea

Browse files
authored
feat(cache): provider-agnostic cache-mode delta + cc-agnostic prefix comparison (headroomlabs-ai#1868)
## Description <!-- Briefly explain the change and why it is needed. --> Closes # ## Type of Change - [ ] Bug fix (non-breaking change that fixes an issue) - [ ] New feature (non-breaking change that adds functionality) - [ ] Breaking change (fix or feature that would cause existing functionality to change) - [ ] Documentation update - [ ] Performance improvement - [ ] Code refactoring (no functional changes) ## Changes Made - ## Testing <!-- Check what you actually ran, then paste the real command output below. --> - [ ] Unit tests pass (`pytest`) - [ ] Linting passes (`ruff check .`) - [ ] Type checking passes (`mypy headroom`) - [ ] New tests added for new functionality - [ ] Manual testing performed ### Test Output ```text # Paste relevant command output or artifact links here ``` ## Real Behavior Proof - Environment: - Exact command / steps: - Observed result: - Not tested: ## Review Readiness - [ ] I have performed a self-review - [ ] This PR is ready for human review ## Checklist - [ ] My code follows the project's style guidelines - [ ] I have performed a self-review of my code - [ ] I have commented my code, particularly in hard-to-understand areas - [ ] I have made corresponding changes to the documentation - [ ] My changes generate no new warnings - [ ] I have added tests that prove my fix is effective or that my feature works - [ ] New and existing unit tests pass locally with my changes - [ ] I have updated the CHANGELOG.md if applicable ## Screenshots (if applicable) Add screenshots to help explain your changes. ## Additional Notes <!-- Mention any N/A checklist items, tradeoffs, follow-ups, or maintainer context. -->
1 parent 3807488 commit 7c2f0ea

21 files changed

Lines changed: 2203 additions & 84 deletions

‎headroom/agent_savings.py‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from __future__ import annotations
44

55
import logging
6+
import os
67
from dataclasses import dataclass, replace
78
from typing import Protocol
89

@@ -278,6 +279,19 @@ def proxy_pipeline_kwargs(config: object) -> dict[str, object]:
278279
if smart_crusher_with_compaction is not None:
279280
kwargs["smart_crusher_with_compaction"] = bool(smart_crusher_with_compaction)
280281

282+
# Lower the block-compression char floor (default 500) so modest tool outputs
283+
# are eligible for the LOSSY path too. Matters in cache mode, where the only
284+
# compressible content each turn is a single (often small) delta observation;
285+
# a 500-char floor buckets most of them as "small" (skipped). Env-gated so it
286+
# only changes behavior when explicitly set; lossless folding has no floor and
287+
# is unaffected.
288+
_min_chars_block = os.environ.get("HEADROOM_MIN_CHARS_FOR_BLOCK")
289+
if _min_chars_block:
290+
try:
291+
kwargs["min_chars_for_block_compression"] = int(_min_chars_block)
292+
except ValueError:
293+
pass
294+
281295
return kwargs
282296

283297

‎headroom/cache/prefix_tracker.py‎

Lines changed: 155 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -128,6 +128,134 @@ def _strip_cache_control(obj: Any) -> Any:
128128
return obj
129129

130130

131+
# Keys that carry NO semantic payload for the model — transport / caching-directive
132+
# / telemetry / client-routing annotations that clients attach and vary turn-to-turn.
133+
# Grounded in provider API docs (Anthropic Messages, OpenAI Chat+Responses, Bedrock
134+
# Converse) + client-library field inventories (litellm, Vercel AI SDK, opencode,
135+
# Claude Code, Cline). Dropped from the cross-turn prefix-equality key ONLY.
136+
#
137+
# NOTE ON SAFETY: this projection is a COMPARISON KEY, never a source to rebuild
138+
# forwarded bytes — the cache-stable-delta path always forwards the previously
139+
# forwarded bytes + the raw appended delta. So dropping these can't deprive the
140+
# model. What we must NOT do is drop a *semantic* field (that would mask a real
141+
# divergence and replay a stale prefix), which is why: (1) reasoning SIGNATURES are
142+
# NOT in this set (Anthropic 400s if a thinking block is altered/missing, and a
143+
# present/absent flip is a real divergence we want to detect); (2) tool inputs /
144+
# arguments / json payloads are treated as OPAQUE and compared verbatim (see
145+
# _OPAQUE_PAYLOAD_KEYS) so a user key that happens to be named "index"/"state" is
146+
# never stripped from inside a tool call.
147+
_NON_SEMANTIC_KEYS = frozenset(
148+
{
149+
# cache-breakpoint markers (moved to the newest block every turn)
150+
"cache_control", # Anthropic (per-block)
151+
"cachePoint", # Bedrock (per-block content block)
152+
# litellm unified-message / tool annotations
153+
"caller", # litellm programmatic-tool tag on tool_use
154+
"provider_specific_fields",
155+
"reasoning_content", # litellm display echo (the paired signature is separate)
156+
"reasoning_items",
157+
"annotations", # citation/display metadata
158+
# OpenAI response echoes that can ride on assistant messages
159+
"system_fingerprint",
160+
"service_tier",
161+
# Vercel AI SDK / opencode part transport
162+
"providerMetadata",
163+
"providerOptions",
164+
"callProviderMetadata",
165+
"state",
166+
"providerExecuted",
167+
"synthetic",
168+
"ignored",
169+
# streaming-assembly artifact
170+
"index",
171+
}
172+
)
173+
174+
# Values under these keys are opaque semantic payloads (tool-call input, OpenAI
175+
# stringified arguments, Bedrock tool_result json). They are compared VERBATIM — we
176+
# never recurse into them to strip "noise" keys, because arbitrary user data there
177+
# may legitimately contain keys that collide with _NON_SEMANTIC_KEYS (e.g. an
178+
# `input` of {"state": "CA", "index": 3}). Recursing would corrupt the comparison.
179+
_OPAQUE_PAYLOAD_KEYS = frozenset({"input", "arguments", "json"})
180+
181+
182+
def _canonicalize_for_prefix_compare(obj: Any) -> Any:
183+
"""Representation-agnostic canonical form for cross-turn prefix equality.
184+
185+
Providers accept several *equivalent* encodings for the same message, and real
186+
clients vary them turn-to-turn; a raw-dict prefix compare then fails spuriously
187+
and drops cache mode to raw (uncompressed) forwarding. This normalizes ONLY
188+
representation:
189+
* drops non-semantic annotation / cache-directive / telemetry keys
190+
(_NON_SEMANTIC_KEYS) at any message/block level;
191+
* wraps a bare string ``content`` into ``[{"type": "text", "text": ...}]``
192+
(Anthropic's string sugar, which litellm flips per turn);
193+
* leaves tool ``input`` / ``arguments`` / ``json`` payloads verbatim
194+
(_OPAQUE_PAYLOAD_KEYS) so user data is never corrupted;
195+
* KEEPS all real content (text, tool name/input, tool_result content, reasoning
196+
signatures, ids) so two messages canonicalize-equal iff they are semantically
197+
identical.
198+
199+
Used ONLY as a comparison key for the cache-stable delta path; the original,
200+
unmodified messages are always what gets forwarded.
201+
"""
202+
if isinstance(obj, dict):
203+
out: dict[str, Any] = {}
204+
for key, value in obj.items():
205+
if key in _NON_SEMANTIC_KEYS:
206+
continue
207+
if key in _OPAQUE_PAYLOAD_KEYS:
208+
out[key] = value # verbatim — do not recurse into user payloads
209+
elif key == "content" and isinstance(value, str):
210+
out[key] = [{"type": "text", "text": value}]
211+
else:
212+
out[key] = _canonicalize_for_prefix_compare(value)
213+
return out
214+
if isinstance(obj, list):
215+
canon = [_canonicalize_for_prefix_compare(value) for value in obj]
216+
# Drop blocks that projected to {} — a pure cache-directive content block
217+
# (e.g. Bedrock {"cachePoint": {...}}) whose only key was non-semantic. Left
218+
# in place it would be an empty-dict entry, so a directive block moving
219+
# position across turns would spuriously fail the length/order compare.
220+
return [value for value in canon if value != {}]
221+
return obj
222+
223+
224+
def extract_cache_stable_delta(
225+
current_messages: list[dict[str, Any]],
226+
previous_original_messages: list[dict[str, Any]] | None,
227+
previous_forwarded_messages: list[dict[str, Any]] | None,
228+
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]] | None:
229+
"""Return ``(stable_forwarded_prefix, appended_delta_messages)`` when the current
230+
request append-only-extends the previous one, else ``None``.
231+
232+
Provider-agnostic delta engine for cache mode. "Append-only" is decided by comparing
233+
the *canonicalized* prefix (:func:`_canonicalize_for_prefix_compare`, which ignores
234+
per-turn transport / cache-directive / client-annotation noise across
235+
Anthropic / OpenAI / Bedrock and the common clients), so a moved cache marker or
236+
shape churn does not spuriously collapse cache mode to raw forwarding. On a match the
237+
caller replays the byte-identical previously-forwarded prefix and compresses ONLY the
238+
appended delta.
239+
240+
This is a COMPARISON + slice only: the returned prefix is the previously-forwarded
241+
bytes verbatim and the delta is the raw appended messages — never a rebuild from the
242+
canonical projection — so the projection dropping non-semantic fields is safe.
243+
"""
244+
if not previous_original_messages or previous_forwarded_messages is None:
245+
return None
246+
prefix_len = len(previous_original_messages)
247+
if len(current_messages) < prefix_len:
248+
return None
249+
if _canonicalize_for_prefix_compare(
250+
current_messages[:prefix_len]
251+
) != _canonicalize_for_prefix_compare(previous_original_messages):
252+
return None
253+
return (
254+
copy.deepcopy(previous_forwarded_messages),
255+
copy.deepcopy(current_messages[prefix_len:]),
256+
)
257+
258+
131259
def overlay_cached_prefix(
132260
optimized_messages: list[dict[str, Any]],
133261
current_original_messages: list[dict[str, Any]],
@@ -168,12 +296,17 @@ def overlay_cached_prefix(
168296
if len(current_original_messages) < n or len(optimized_messages) < n:
169297
return optimized_messages
170298
# Append-only guard on CONTENT ONLY: the frozen region must be the same
171-
# messages we cached. Compare with cache_control stripped — clients move that
172-
# breakpoint to the newest message each turn, so a raw dict compare would
173-
# spuriously fail whenever a marker lands in the frozen prefix, skip the
174-
# replay, and bust the cache (the residual busts observed after the first
175-
# fix). Content stability is what the provider's prefix cache actually keys on.
176-
if _strip_cache_control(current_original_messages[:n]) != _strip_cache_control(prev_orig):
299+
# messages we cached. Compare with the shared canonicalizer (not just
300+
# cache_control-stripping) so the guard is robust to ALL per-turn transport /
301+
# annotation churn — cache_control movement (Anthropic), litellm `caller`,
302+
# provider_specific_fields, streaming `index`, string<->block content shape,
303+
# etc. — across providers/clients. Content stability is what the provider's
304+
# prefix cache actually keys on; a coarser cache_control-only strip let other
305+
# clients' noise spuriously fail the guard, skip the replay, and bust. This
306+
# helps every handler that shares overlay_cached_prefix (Anthropic + OpenAI).
307+
if _canonicalize_for_prefix_compare(
308+
current_original_messages[:n]
309+
) != _canonicalize_for_prefix_compare(prev_orig):
177310
return optimized_messages
178311
# Replay the cached (compressed) prefix byte-identical; keep this turn's tail.
179312
return list(prev_fwd) + list(optimized_messages[n:])
@@ -563,6 +696,22 @@ def _estimate_message_tokens(messages: list[dict[str, Any]]) -> list[int]:
563696
chars += len(text)
564697
else:
565698
chars = 0
699+
# OpenAI function-calling: the assistant's command lives in the
700+
# top-level `tool_calls` (or legacy `function_call`) field, NOT in
701+
# `content` (which is empty/None on a tool-call turn). Anthropic puts
702+
# the equivalent in a `tool_use` content BLOCK (counted above), but
703+
# the OpenAI shape was never counted here. That under-counted every
704+
# tool-based assistant turn to ~0, so the frozen-prefix estimate
705+
# overshot the real cache boundary and froze the NEWEST delta — which
706+
# is why OpenAI/Kimi (fireworks) tool harnesses got ~zero compression
707+
# while text/back-tick harnesses (command in `content`) compressed.
708+
for tc in msg.get("tool_calls") or []:
709+
if isinstance(tc, dict):
710+
fn = tc.get("function") or {}
711+
chars += len(str(fn.get("name", ""))) + len(str(fn.get("arguments", "")))
712+
fc = msg.get("function_call")
713+
if isinstance(fc, dict):
714+
chars += len(str(fc.get("name", ""))) + len(str(fc.get("arguments", "")))
566715
# Add overhead for role, block structure, etc.
567716
chars += 20
568717
counts.append(max(1, int(chars / 3.5)))

‎headroom/proxy/cost.py‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -738,6 +738,21 @@ def record_tokens(
738738
uncached_tokens: Non-cached input tokens from API response usage.
739739
output_tokens: Output tokens from API response usage.
740740
"""
741+
# Post-guard invariant (all providers): Headroom never forwards a request
742+
# larger than the original (handlers revert any inflation before sending),
743+
# so compression savings are >= 0 by construction. A negative here is an
744+
# intermediate/hook token-count artifact that never reached the model;
745+
# clamp it so `total_tokens_removed` reflects actually-forwarded bytes
746+
# instead of surfacing spurious negatives (verified clean on the wire).
747+
if tokens_saved < 0:
748+
import logging as _lg
749+
750+
_lg.getLogger(__name__).debug(
751+
"record_tokens: clamping negative tokens_saved=%d to 0 for %s (artifact; wire not inflated)",
752+
tokens_saved,
753+
model,
754+
)
755+
tokens_saved = 0
741756
self._tokens_saved_by_model[model] = (
742757
self._tokens_saved_by_model.get(model, 0) + tokens_saved
743758
)

‎headroom/proxy/handlers/anthropic.py‎

Lines changed: 59 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -362,17 +362,20 @@ def _extract_cache_stable_delta(
362362
Safe means the prior original request is an exact message-prefix of the
363363
current original request. This lets us replay the exact forwarded bytes
364364
for historical context and only transform newly appended message suffixes.
365+
366+
The append-only check ignores per-turn transport / cache-directive / client
367+
annotation noise (cache_control moved to the newest block, litellm caller,
368+
provider_specific_fields, streaming index, string<->block content shape, …) via
369+
the shared canonicalizer, so that churn doesn't spuriously drop cache mode to raw
370+
forwarding. Delegates to the provider-agnostic engine in prefix_tracker so
371+
OpenAI / Bedrock share one implementation.
365372
"""
366-
if not previous_original_messages or previous_forwarded_messages is None:
367-
return None
368-
prefix_len = len(previous_original_messages)
369-
if len(current_messages) < prefix_len:
370-
return None
371-
if current_messages[:prefix_len] != previous_original_messages:
372-
return None
373-
return (
374-
copy.deepcopy(previous_forwarded_messages),
375-
copy.deepcopy(current_messages[prefix_len:]),
373+
from headroom.cache.prefix_tracker import extract_cache_stable_delta
374+
375+
return extract_cache_stable_delta(
376+
current_messages,
377+
previous_original_messages,
378+
previous_forwarded_messages,
376379
)
377380

378381
@staticmethod
@@ -1342,21 +1345,63 @@ class _DeferredCompressionResult:
13421345
optimized_messages = messages
13431346
optimized_tokens = tokenizer.count_messages(optimized_messages)
13441347
else:
1348+
# Compress the delta, with two cache-mode adjustments:
1349+
#
1350+
# fix-5: strip the client's transient cache_control marker so
1351+
# the router's per-block "never compress an explicit cache
1352+
# key" guard (content_router.py:4006) doesn't skip the ONLY
1353+
# compressible content every turn (route_counts had
1354+
# cache_control_protected == the whole delta -> 0%). In cache
1355+
# mode that marker is NOT the real forwarded breakpoint: the
1356+
# compressed delta is frozen + replayed verbatim next turn and
1357+
# normalize_message_cache_control (AFTER compression, below)
1358+
# owns the single forwarded breakpoint. Cache-safety is
1359+
# enforced post-compression, not by protecting the delta.
1360+
#
1361+
# fix-6: the delta is a lone tool_result whose tool_use (tool
1362+
# NAME + call args) lives in the frozen prefix. Passing only
1363+
# the delta to the router leaves tool_name="" so
1364+
# _bash_search_fold (lossless grep/rg folding, no size floor),
1365+
# per-tool bias, and relevance-query enrichment all degrade.
1366+
# Pass the FULL current messages with frozen_message_count =
1367+
# prefix length: _build_tool_name_map scans ALL messages (the
1368+
# delta resolves its tool_name from the prefix's tool_use) but
1369+
# the compression loop only touches indices >= frozen count,
1370+
# so ONLY the delta is compressed. Splice the compressed delta
1371+
# onto the byte-stable forwarded prefix.
1372+
from headroom.cache.prefix_tracker import _strip_cache_control
1373+
1374+
# Compression context = the EXACT forwarded (cached) prefix
1375+
# + the stripped delta, with the prefix frozen. Using the
1376+
# forwarded prefix (not the original) keeps _build_tool_name_map
1377+
# AND cross-turn dedup consistent with what is actually cached:
1378+
# dedup can only reference bytes that are truly present in the
1379+
# forwarded context, so no pointer can dangle. The prefix is
1380+
# frozen (never compressed) and we discard the router's copy of
1381+
# it below, so the forwarded prefix stays byte-identical to last
1382+
# turn -> append-only -> no bust.
1383+
prefix_n = len(stable_forwarded_prefix)
1384+
compression_input = list(stable_forwarded_prefix) + list(
1385+
_strip_cache_control(delta_messages)
1386+
)
13451387
result = await self._run_compression_in_executor(
13461388
lambda: self.anthropic_pipeline.apply(
1347-
messages=delta_messages,
1389+
messages=compression_input,
13481390
model=model,
13491391
model_limit=context_limit,
1350-
context=extract_user_query(delta_messages),
1351-
frozen_message_count=0,
1392+
context=extract_user_query(compression_input),
1393+
frozen_message_count=prefix_n,
13521394
biases=biases,
13531395
request_id=request_id,
13541396
compression_policy=compression_policy,
13551397
**proxy_pipeline_kwargs(self.config),
13561398
),
13571399
timeout=COMPRESSION_TIMEOUT_SECONDS,
13581400
)
1359-
optimized_messages = stable_forwarded_prefix + result.messages
1401+
# Only the delta was eligible for compression (prefix frozen);
1402+
# forward the byte-identical cached prefix + the compressed delta.
1403+
compressed_delta = result.messages[prefix_n:]
1404+
optimized_messages = stable_forwarded_prefix + compressed_delta
13601405
transforms_applied = result.transforms_applied
13611406
pipeline_timing = result.timing
13621407
optimized_tokens = tokenizer.count_messages(optimized_messages)

0 commit comments

Comments
 (0)