mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
The freeze path (both providers) emits the agent's ORIGINAL bytes for a frozen message, but the provider cached whatever we FORWARDED last turn (the compressed form). Forwarding original then mismatches the cached prefix and busts it from that point — re-creating the whole suffix. Measured on a real SWE-bench run: 100% of attributed misses were prefix_change, ~56% of ALL cache-writes were bust-induced (2.8M tokens), driving cache_create +150% and cost +41% vs baseline. Cache mode already avoided this via _extract_cache_stable_delta (replay the previously-forwarded prefix, compress only the delta). Token mode called apply(frozen_count) directly, which forwards original for the frozen region. Fix: add a shared, provider-agnostic overlay_cached_prefix() that replays the previously-forwarded (cached, compressed) prefix byte-identical, append-only guarded and idempotent, and apply it in BOTH the Anthropic and OpenAI handlers right before forwarding. This makes freezing byte-identical in every mode, so the only remaining difference between "token" and "cache" mode is how large a mutable (still-compressible) tail each leaves — not whether the frozen prefix busts the cache. Tests: - test_cache_prefix_overlay.py: the helper (replay, append-only guard, idempotence). - test_cross_turn_cache_safety.py: the invariant that was missing — drive the REAL tracker + freeze + overlay over multiple append-only turns against a simulated provider prefix cache and assert the forwarded prefix stays byte-identical turn-over-turn. Load-bearing: it fails (detects the bust) without the overlay. ## 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. -->
140 lines
5.8 KiB
Python
140 lines
5.8 KiB
Python
"""Cross-turn cache-safety invariant — the test class that catches cache busts.
|
|
|
|
Why the +150%-cache_create / +41%-cost bug slipped through: every prior cache
|
|
test was SINGLE-turn and used a fake tracker, so nobody exercised the real
|
|
multi-turn invariant that actually governs prompt-cache cost:
|
|
|
|
Across append-only turns, the forwarded prefix must stay BYTE-IDENTICAL to
|
|
what was forwarded (and cached) last turn — otherwise the provider re-creates
|
|
the whole suffix (a cache bust) instead of reading it.
|
|
|
|
This simulates the provider's prefix cache (longest byte-identical leading run of
|
|
messages = cache_read; the rest = cache_create) and drives the REAL
|
|
``PrefixCacheTracker`` + the freeze model + ``overlay_cached_prefix`` over several
|
|
turns. It asserts the invariant directly, and proves the guard is load-bearing:
|
|
WITHOUT the overlay the freeze forwards the agent's original bytes and busts every
|
|
turn; WITH it the prefix stays stable.
|
|
"""
|
|
|
|
from headroom.cache.prefix_tracker import (
|
|
PrefixCacheTracker,
|
|
PrefixFreezeConfig,
|
|
overlay_cached_prefix,
|
|
)
|
|
|
|
|
|
def _toklen(m) -> int:
|
|
return max(1, len(str(m.get("content", ""))))
|
|
|
|
|
|
def _compress(m):
|
|
"""Deterministic stand-in for a real compressor (kompress is deterministic
|
|
per content via the result cache): shrink the content by half."""
|
|
c = str(m.get("content", ""))
|
|
return {**m, "content": c[: max(1, len(c) // 2)]}
|
|
|
|
|
|
def _apply_freeze(original, frozen_count):
|
|
"""Faithful model of pipeline.apply()'s freeze: the frozen prefix is
|
|
forwarded as the agent's ORIGINAL bytes; everything else is compressed.
|
|
(Mirrors content_router.py: `result_slots[i] = message` for i < frozen.)"""
|
|
return [
|
|
(original[i] if i < frozen_count else _compress(original[i])) for i in range(len(original))
|
|
]
|
|
|
|
|
|
def _provider_cache_read(forwarded, prev_forwarded):
|
|
"""Longest byte-identical leading run of messages the provider can serve from
|
|
cache, in tokens. A single differing message breaks the prefix (bust)."""
|
|
if not prev_forwarded:
|
|
return 0
|
|
matched = 0
|
|
for a, b in zip(forwarded, prev_forwarded):
|
|
if a == b:
|
|
matched += _toklen(a)
|
|
else:
|
|
break
|
|
return matched
|
|
|
|
|
|
def _drive_turns(*, use_overlay: bool, turns: int = 5):
|
|
"""Return per-turn (expected_cache_read, actual_cache_read). A bust is any
|
|
turn where actual < expected (the previously-cached prefix wasn't reused)."""
|
|
# min_cached_tokens=0 so freeze activates from turn 2 regardless of size.
|
|
tracker = PrefixCacheTracker("anthropic", PrefixFreezeConfig(min_cached_tokens=0))
|
|
convo: list[dict] = []
|
|
prev_forwarded: list[dict] | None = None
|
|
out = []
|
|
for t in range(1, turns + 1):
|
|
# Append-only growth: one new large tool output per turn.
|
|
convo = convo + [{"role": "user", "content": f"tool-output-turn-{t}:" + "X" * 400}]
|
|
|
|
frozen = tracker.get_frozen_message_count()
|
|
forwarded = _apply_freeze(convo, frozen)
|
|
if use_overlay:
|
|
forwarded = overlay_cached_prefix(
|
|
forwarded,
|
|
convo,
|
|
tracker.get_last_original_messages(),
|
|
tracker.get_last_forwarded_messages(),
|
|
)
|
|
|
|
expected_read = sum(_toklen(m) for m in prev_forwarded) if prev_forwarded else 0
|
|
actual_read = _provider_cache_read(forwarded, prev_forwarded)
|
|
out.append((expected_read, actual_read))
|
|
|
|
counts = [_toklen(m) for m in forwarded]
|
|
write = sum(counts) - actual_read
|
|
tracker.update_from_response(
|
|
actual_read, write, forwarded, message_token_counts=counts, original_messages=convo
|
|
)
|
|
prev_forwarded = forwarded
|
|
return out
|
|
|
|
|
|
def test_freeze_busts_cache_every_turn_without_overlay():
|
|
"""Proves the test is load-bearing: the raw freeze path busts the cache."""
|
|
results = _drive_turns(use_overlay=False)
|
|
# From turn 2 on, a hit was expected but the prefix broke (actual < expected).
|
|
busts = [exp > act for (exp, act) in results[1:]]
|
|
assert any(busts), "expected the un-fixed freeze path to bust the prefix cache"
|
|
|
|
|
|
def test_overlay_keeps_prefix_byte_identical_no_bust():
|
|
"""The fix: every turn reuses the full previously-cached prefix — no bust."""
|
|
results = _drive_turns(use_overlay=True)
|
|
for exp, act in results[1:]:
|
|
assert act >= exp, (
|
|
f"cache bust: expected to read {exp} cached tokens but only read {act} "
|
|
"— forwarded prefix diverged from last turn"
|
|
)
|
|
|
|
|
|
def test_cache_create_stays_bounded_to_the_delta_with_overlay():
|
|
"""Cost proxy: with the fix, per-turn cache_create ≈ the new delta only, not
|
|
the whole re-created prefix (which is what drove +150% cache_create)."""
|
|
tracker = PrefixCacheTracker("anthropic", PrefixFreezeConfig(min_cached_tokens=0))
|
|
convo: list[dict] = []
|
|
prev_forwarded: list[dict] | None = None
|
|
creates = []
|
|
for t in range(1, 6):
|
|
convo = convo + [{"role": "user", "content": f"turn-{t}:" + "X" * 400}]
|
|
frozen = tracker.get_frozen_message_count()
|
|
forwarded = overlay_cached_prefix(
|
|
_apply_freeze(convo, frozen),
|
|
convo,
|
|
tracker.get_last_original_messages(),
|
|
tracker.get_last_forwarded_messages(),
|
|
)
|
|
read = _provider_cache_read(forwarded, prev_forwarded)
|
|
counts = [_toklen(m) for m in forwarded]
|
|
create = sum(counts) - read
|
|
creates.append(create)
|
|
tracker.update_from_response(
|
|
read, create, forwarded, message_token_counts=counts, original_messages=convo
|
|
)
|
|
prev_forwarded = forwarded
|
|
# Steady-state cache_create per turn should be ~one delta message, NOT growing
|
|
# with conversation length. Assert the last turn creates no more than the
|
|
# first (which had no cache to reuse).
|
|
assert creates[-1] <= creates[0] + 1
|