headroom/tests/test_cross_turn_cache_safety.py
Tejas Chopra 248ae0f3e0
fix(proxy): freeze must forward cached (compressed) prefix byte-identical — stop token-mode cache busting (#1850)
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. -->
2026-07-06 14:54:39 -07:00

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