diff --git a/headroom/cache/prefix_tracker.py b/headroom/cache/prefix_tracker.py index 8157e9e9a..200adddb8 100644 --- a/headroom/cache/prefix_tracker.py +++ b/headroom/cache/prefix_tracker.py @@ -449,7 +449,7 @@ def overlay_cached_prefix( previous_original_messages: list[dict[str, Any]] | None, previous_forwarded_messages: list[dict[str, Any]] | None, ) -> list[dict[str, Any]]: - """Replay the previously-forwarded (cached, compressed) prefix byte-identical. + """Replay a positional, non-inflating cached prefix when it is safe. Provider-agnostic cache-safety guard for the freeze path. When a message is "frozen", the compression pipeline may emit the agent's ORIGINAL bytes for @@ -467,14 +467,23 @@ def overlay_cached_prefix( ``optimized_messages`` unchanged (accept a possible bust rather than forward wrong content). - This makes freezing byte-identical in BOTH proxy modes, so the only remaining - difference between them is how large a mutable (still-compressible) tail each - leaves — not whether the frozen prefix busts the cache. + The optimized and current-original lists must be positionally aligned, and + compact UTF-8 JSON for the replayed result must not exceed the optimized + candidate. These bounds prefer a cache miss to corrupting or inflating a + client's live history. """ prev_orig = previous_original_messages prev_fwd = previous_forwarded_messages if not prev_orig or not prev_fwd: return optimized_messages + if len(optimized_messages) != len(current_original_messages): + logger.debug( + "overlay: optimized/current-original length mismatch (optimized=%d, current=%d) " + "— skipping positional cached-prefix replay", + len(optimized_messages), + len(current_original_messages), + ) + return optimized_messages n = len(prev_orig) # Positional 1:1 correspondence between prev_orig[i] and prev_fwd[i] holds # only when last turn forwarded exactly one message per original (the @@ -534,11 +543,21 @@ def overlay_cached_prefix( len(current_content) - split, message_index, ) - return ( + replayed = ( list(prev_fwd[:message_index]) + [merged] + list(optimized_messages[message_index + 1 :]) ) + replayed_bytes = _compact_json_bytes(replayed) + optimized_bytes = _compact_json_bytes(optimized_messages) + if ( + replayed_bytes is None + or optimized_bytes is None + or len(replayed_bytes) > len(optimized_bytes) + ): + logger.debug("overlay: block replay inflated compact JSON — skipping") + return optimized_messages + return replayed # Append-only guard on CONTENT ONLY, message-by-message. Replay the # previously-forwarded (cached, compressed) bytes for the longest LEADING # run of messages that is byte-for-byte (content-canonical) identical to @@ -562,7 +581,7 @@ def overlay_cached_prefix( # current_original[k] canonicalize-equals prev_orig[k], and prev_fwd[k] # positionally corresponds to prev_orig[k] (guaranteed by the count check # above), so no wrong bytes are ever forwarded. - limit = min(n, len(current_original_messages), len(optimized_messages)) + limit = min(n, len(current_original_messages)) k = 0 while k < limit and _canonicalize_for_prefix_compare( current_original_messages[k] @@ -584,7 +603,30 @@ def overlay_cached_prefix( ) # Replay the cached (compressed) prefix byte-identical up to the first # divergence; keep this turn's freshly-produced output for the rest. - return list(prev_fwd[:k]) + list(optimized_messages[k:]) + replayed = list(prev_fwd[:k]) + list(optimized_messages[k:]) + replayed_bytes = _compact_json_bytes(replayed) + optimized_bytes = _compact_json_bytes(optimized_messages) + if ( + replayed_bytes is None + or optimized_bytes is None + or len(replayed_bytes) > len(optimized_bytes) + ): + logger.debug("overlay: replay inflated compact JSON — skipping cached-prefix replay") + return optimized_messages + return replayed + + +def _compact_json_bytes(value: Any) -> bytes | None: + """Return compact JSON bytes, or ``None`` when sizing cannot be proved.""" + try: + return json.dumps( + value, + separators=(",", ":"), + ensure_ascii=False, + default=str, + ).encode("utf-8") + except (TypeError, ValueError, OverflowError, UnicodeError): + return None _STABLE_BOUNDARY_ENV = "HEADROOM_STABLE_BOUNDARY_BREAKPOINT" diff --git a/headroom/proxy/handlers/anthropic.py b/headroom/proxy/handlers/anthropic.py index 2ac5bc7db..ca119f4dc 100644 --- a/headroom/proxy/handlers/anthropic.py +++ b/headroom/proxy/handlers/anthropic.py @@ -1975,22 +1975,35 @@ class AnthropicHandlerMixin: overlay_cached_prefix, ) + _overlay_replayed = False # On a confirmed-cold turn we deliberately do NOT replay the previously # forwarded prefix: the cache is dead (nothing to keep byte-identical for) # and the replay would clobber the whole-prefix recompaction we just did. - if _cold_recompact_active: - _overlay_replayed = False + if _decision.should_compress and not _skip_compression_for_backpressure: + if _cold_recompact_active: + _overlay_replayed = False + else: + _ov = overlay_cached_prefix( + optimized_messages, + original_client_messages, + previous_original_messages, + previous_forwarded_messages, + ) + _overlay_replayed = _ov != optimized_messages + if _overlay_replayed: + optimized_messages = _ov + optimized_tokens = tokenizer.count_messages(optimized_messages) else: - _ov = overlay_cached_prefix( - optimized_messages, - original_client_messages, - previous_original_messages, - previous_forwarded_messages, + replay_skip_reason = ( + "pre_upstream_backpressure" + if _skip_compression_for_backpressure + else _decision.passthrough_reason + ) + logger.debug( + "[%s] Cached-prefix replay skipped: reason=%s", + request_id, + replay_skip_reason, ) - _overlay_replayed = _ov != optimized_messages - if _overlay_replayed: - optimized_messages = _ov - optimized_tokens = tokenizer.count_messages(optimized_messages) # Own cache_control placement: the client moves the breakpoint each # turn and the overlay replays past markers, so they accumulate ~1/turn diff --git a/tests/test_cache_control_move_bust.py b/tests/test_cache_control_move_bust.py index 3816c7b38..7969b3a78 100644 --- a/tests/test_cache_control_move_bust.py +++ b/tests/test_cache_control_move_bust.py @@ -61,14 +61,11 @@ def test_marker_move_would_fail_a_raw_dict_guard(): assert _strip_cache_control(CUR_ORIG[:2]) == _strip_cache_control(PREV_ORIG) -def test_overlay_replays_despite_moved_marker(): +def test_overlay_skips_inflating_moved_marker_replay(): out = overlay_cached_prefix(OPTIMIZED, CUR_ORIG, PREV_ORIG, PREV_FWD) - # The content-only guard lets the replay happen: the forwarded prefix is now - # byte-identical to what the provider cached (compressed), NOT the freeze's - # original bytes → cache hits instead of busting. - assert out[:2] == PREV_FWD - assert out[:2] != OPTIMIZED[:2] - assert out[2] == OPTIMIZED[2] # compressed tail preserved + # Moving cache_control must not exempt a larger replay candidate from the + # no-inflation bound, even when content-only history alignment succeeds. + assert out == OPTIMIZED # ── Cross-turn: client moves the marker every turn, provider keys on full bytes ─ diff --git a/tests/test_cache_prefix_overlay.py b/tests/test_cache_prefix_overlay.py index 8cf97df37..45791ec1d 100644 --- a/tests/test_cache_prefix_overlay.py +++ b/tests/test_cache_prefix_overlay.py @@ -9,6 +9,8 @@ Forwarding original then mismatches the cached prefix and busts the prompt cache so the cache still hits — in BOTH proxy modes. """ +import copy + from headroom.cache.prefix_tracker import overlay_cached_prefix @@ -72,6 +74,77 @@ def test_shorter_current_or_optimized_returns_unchanged(): ] +def test_overlay_requires_positional_alignment_with_originals(): + optimized = [M("user", "x")] + current = [M("user", "x"), M("assistant", "ok")] + assert overlay_cached_prefix(optimized, current, PREV_ORIG, PREV_FWD) == optimized + + optimized = [M("user", "x"), M("assistant", "ok"), M("user", "tail")] + current = [M("user", "x"), M("assistant", "ok")] + previous = [M("user", "x"), M("assistant", "ok")] + forwarded = [M("user", "compressed"), M("assistant", "ok")] + assert overlay_cached_prefix(optimized, current, previous, forwarded) == optimized + + +def test_overlay_never_inflates_forwarded_payload(): + optimized = [M("user", "small"), M("assistant", "ok"), M("user", "tail")] + inflated_forwarded = [M("user", "x" * 1000), M("assistant", "ok")] + previous = [M("user", "small"), M("assistant", "ok")] + current = previous + [M("user", "tail")] + assert overlay_cached_prefix(optimized, current, previous, inflated_forwarded) == optimized + + +def test_overlay_returns_optimized_when_json_sizing_fails(monkeypatch): + optimized = [M("user", "stable"), M("user", "tail")] + current = [M("user", "stable"), M("user", "tail")] + previous = [M("user", "stable")] + forwarded = [M("user", "compressed")] + + monkeypatch.setattr( + "headroom.cache.prefix_tracker.json.dumps", + lambda *args, **kwargs: (_ for _ in ()).throw(TypeError("cannot size")), + ) + + assert overlay_cached_prefix(optimized, current, previous, forwarded) == optimized + + +def test_overlay_never_inflates_cache_control_only_replay(): + previous = [M("user", "stable"), M("assistant", "ok")] + current = [ + M("user", "stable"), + {**M("assistant", "ok"), "cache_control": {"type": "ephemeral"}}, + ] + optimized = copy.deepcopy(current) + inflated_forwarded = [M("user", "x" * 1000), M("assistant", "ok")] + assert overlay_cached_prefix(optimized, current, previous, inflated_forwarded) == optimized + + +def test_block_append_overlay_never_inflates_forwarded_payload(): + previous = [ + { + "role": "user", + "content": [{"type": "text", "text": "stable"}], + } + ] + current = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "stable"}, + {"type": "text", "text": "tail"}, + ], + } + ] + optimized = copy.deepcopy(current) + forwarded = [ + { + "role": "user", + "content": [{"type": "text", "text": "x" * 1000}], + } + ] + assert overlay_cached_prefix(optimized, current, previous, forwarded) == optimized + + def test_cache_hit_property_prefix_matches_last_forward(): # The invariant that guarantees a cache hit: forwarded[:n] this turn == # forwarded[:n] last turn (== what the provider cached). diff --git a/tests/test_issue_2671_block_growth_cache.py b/tests/test_issue_2671_block_growth_cache.py index 160d4c22a..fc26da0f6 100644 --- a/tests/test_issue_2671_block_growth_cache.py +++ b/tests/test_issue_2671_block_growth_cache.py @@ -223,14 +223,14 @@ def test_cache_affinity_ignores_only_cache_directive_movement() -> None: def test_pure_append_replays_forwarded_blocks_and_advances_breakpoint() -> None: previous_original = _pure_append(30) - previous_forwarded = _message([_text(f"COMPRESSED-{index}") for index in range(30)]) + previous_forwarded = _message([_text(f"C-{index}") for index in range(30)]) current = _pure_append(34) overlaid = overlay_cached_prefix(current, current, previous_original, previous_forwarded) normalized = normalize_message_cache_control(overlaid, previous_forwarded) assert [block["text"] for block in normalized[0]["content"][:30]] == [ - f"COMPRESSED-{index}" for index in range(30) + f"C-{index}" for index in range(30) ] assert [block["text"] for block in normalized[0]["content"][30:]] == [ f"block-{index}" for index in range(30, 34) diff --git a/tests/test_proxy/test_anthropic_no_optimize_history_passthrough.py b/tests/test_proxy/test_anthropic_no_optimize_history_passthrough.py new file mode 100644 index 000000000..c79071b3b --- /dev/null +++ b/tests/test_proxy/test_anthropic_no_optimize_history_passthrough.py @@ -0,0 +1,279 @@ +"""Regression coverage for Anthropic cached-prefix replay eligibility.""" + +from __future__ import annotations + +import asyncio +import copy +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock + +import httpx +from fastapi.testclient import TestClient + +from headroom.proxy.server import ProxyConfig, create_app + +MODEL = "claude-sonnet-4-6" +MARKER = "large-tool-result-marker " * 80 + + +def _config(**overrides) -> ProxyConfig: + values = { + "optimize": False, + "cache_enabled": False, + "rate_limit_enabled": False, + "cost_tracking_enabled": False, + "log_requests": False, + "prefix_freeze_enabled": True, + } + values.update(overrides) + return ProxyConfig(**values) + + +def _response() -> httpx.Response: + return httpx.Response( + 200, + json={ + "id": "msg_test", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "model": MODEL, + "usage": {"input_tokens": 10, "output_tokens": 2}, + }, + ) + + +def _capture_proxy(client: TestClient) -> tuple[list[dict], AsyncMock]: + captured: list[dict] = [] + + async def capture(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured.append(copy.deepcopy(body)) + return _response() + + retry = AsyncMock(side_effect=capture) + client.app.state.proxy._retry_request = retry + return captured, retry + + +def _compact_bytes(messages: list[dict]) -> int: + return len(json.dumps(messages, separators=(",", ":"), ensure_ascii=False).encode()) + + +def _pipeline_result(messages: list[dict]) -> SimpleNamespace: + return SimpleNamespace( + messages=messages, + transforms_applied=["test_stub"], + timing={}, + tokens_before=1000, + tokens_after=500, + waste_signals=None, + ) + + +def test_no_optimize_second_turn_forwards_client_history_verbatim(): + """A bypass turn owns its Claude-shaped body after an optimized prior turn.""" + app = create_app(_config(optimize=True)) + tool_result = { + "type": "tool_result", + "tool_use_id": "toolu_01", + "content": MARKER, + "cache_control": {"type": "ephemeral"}, + } + first = [{"role": "user", "content": [tool_result]}] + compressed = copy.deepcopy(first) + compressed[0]["content"][0]["content"] = "compressed-tool-result" + app.state.proxy.anthropic_pipeline.apply = Mock(return_value=_pipeline_result(compressed)) + with TestClient(app) as client: + captured, _ = _capture_proxy(client) + response = client.post( + "/v1/messages", + json={"model": MODEL, "max_tokens": 16, "messages": first}, + ) + assert response.status_code == 200, response.text + + assistant = {"role": "assistant", "content": [{"type": "text", "text": "ok"}]} + current = [ + { + "role": "user", + "content": [{k: v for k, v in tool_result.items() if k != "cache_control"}], + }, + assistant, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_02", + "content": "next tool result", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + ] + response = client.post( + "/v1/messages", + headers={"x-headroom-bypass": "true"}, + json={"model": MODEL, "max_tokens": 16, "messages": current}, + ) + + assert response.status_code == 200, response.text + outbound = captured[-1]["messages"] + assert outbound[0]["content"][0]["content"] == MARKER + assert "compressed-tool-result" not in json.dumps(outbound) + assert sum(MARKER in json.dumps(message) for message in outbound) == 1 + assert _compact_bytes(outbound) <= _compact_bytes(current) + + +def test_optimize_off_real_exchange_reports_noninflation(): + app = create_app(_config(optimize=False)) + first = [ + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_01", + "content": MARKER, + } + ], + } + ] + current = first + [ + {"role": "assistant", "content": [{"type": "text", "text": "ok"}]}, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_02", + "content": "next tool result", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + ] + with TestClient(app) as client: + captured, _ = _capture_proxy(client) + assert ( + client.post( + "/v1/messages", + json={"model": MODEL, "max_tokens": 16, "messages": first}, + ).status_code + == 200 + ) + response = client.post( + "/v1/messages", + json={"model": MODEL, "max_tokens": 16, "messages": current}, + ) + + assert response.status_code == 200, response.text + outbound = captured[-1]["messages"] + assert len(outbound) == 3 + assert sum(MARKER in json.dumps(message) for message in outbound) == 1 + assert _compact_bytes(outbound) == _compact_bytes(current) == 2293 + print( + "optimize_off turn2_message_count=3 marker_count=1 " + "outbound_compact_utf8_bytes=2293 client_compact_utf8_bytes=2293" + ) + + +def test_bypass_header_does_not_invoke_cached_prefix_replay(monkeypatch): + app = create_app(_config(optimize=True)) + replay = monkeypatch.setattr + + def fail_if_called(*args, **kwargs): # noqa: ANN002, ANN003 + raise AssertionError("cached-prefix replay must be bypassed") + + replay("headroom.cache.prefix_tracker.overlay_cached_prefix", fail_if_called) + with TestClient(app) as client: + captured, _ = _capture_proxy(client) + messages = [{"role": "user", "content": "bypass"}] + response = client.post( + "/v1/messages", + headers={"x-headroom-bypass": "true"}, + json={"model": MODEL, "max_tokens": 16, "messages": messages}, + ) + + assert response.status_code == 200, response.text + assert captured[-1]["messages"] == messages + + +def test_backpressure_does_not_invoke_cached_prefix_replay(monkeypatch): + app = create_app( + _config( + optimize=True, + anthropic_pre_upstream_concurrency=1, + anthropic_pre_upstream_acquire_timeout_seconds=0.001, + ) + ) + + def fail_if_called(*args, **kwargs): # noqa: ANN002, ANN003 + raise AssertionError("cached-prefix replay must be skipped under backpressure") + + monkeypatch.setattr("headroom.cache.prefix_tracker.overlay_cached_prefix", fail_if_called) + debug = Mock() + monkeypatch.setattr("headroom.proxy.handlers.anthropic.logger.debug", debug) + proxy = app.state.proxy + + class _SaturatedSemaphore: + async def acquire(self): + await asyncio.sleep(1) + + def release(self): + pass + + proxy.anthropic_pre_upstream_sem = _SaturatedSemaphore() + try: + with TestClient(app) as client: + captured, _ = _capture_proxy(client) + messages = [{"role": "user", "content": "backpressure"}] + response = client.post( + "/v1/messages", + json={"model": MODEL, "max_tokens": 16, "messages": messages}, + ) + finally: + proxy.anthropic_pre_upstream_sem.release() + + assert response.status_code == 200, response.text + assert captured[-1]["messages"] == messages + assert any( + call.args and call.args[-1] == "pre_upstream_backpressure" for call in debug.call_args_list + ) + + +def test_optimize_on_aligned_history_preserves_replay(): + app = create_app(_config(optimize=True)) + first = [{"role": "user", "content": "original prefix"}] + compressed = [{"role": "user", "content": "comp"}] + assistant = { + "role": "assistant", + "content": [{"type": "text", "text": "ok", "cache_control": {"type": "ephemeral"}}], + } + current = first + [assistant, {"role": "user", "content": "new suffix"}] + optimized = compressed + [assistant, {"role": "user", "content": "new suffix"}] + app.state.proxy.anthropic_pipeline.apply = Mock( + side_effect=[_pipeline_result(compressed), _pipeline_result(optimized)] + ) + with TestClient(app) as client: + captured, _ = _capture_proxy(client) + response = client.post( + "/v1/messages", + json={"model": MODEL, "max_tokens": 16, "messages": first}, + ) + assert response.status_code == 200, response.text + response = client.post( + "/v1/messages", + json={"model": MODEL, "max_tokens": 16, "messages": current}, + ) + + assert response.status_code == 200, response.text + outbound = captured[-1]["messages"] + assert outbound[0] == compressed[0] + assert outbound[-1] == optimized[-1] + print( + f"optimize_on turn2_message_count={len(outbound)} client_message_count={len(current)} " + f"marker_count={sum(MARKER in json.dumps(message) for message in outbound)} " + f"outbound_compact_utf8_bytes={_compact_bytes(outbound)} " + f"client_compact_utf8_bytes={_compact_bytes(current)}" + ) diff --git a/tests/test_proxy_openai_cache_stability.py b/tests/test_proxy_openai_cache_stability.py index 813155a3f..93e976619 100644 --- a/tests/test_proxy_openai_cache_stability.py +++ b/tests/test_proxy_openai_cache_stability.py @@ -2,6 +2,7 @@ from __future__ import annotations +import copy from types import SimpleNamespace import httpx @@ -15,8 +16,15 @@ from headroom.proxy.server import ProxyConfig, create_app class _FakePrefixTracker: - def __init__(self, frozen_count: int): + def __init__( + self, + frozen_count: int, + previous_original: list[dict] | None = None, + previous_forwarded: list[dict] | None = None, + ): self._frozen_count = frozen_count + self._previous_original = previous_original or [] + self._previous_forwarded = previous_forwarded or [] def get_frozen_message_count(self) -> int: return self._frozen_count @@ -26,10 +34,10 @@ class _FakePrefixTracker: # overlay itself is exercised in test_cross_turn_cache_safety.py against the # real tracker; these stubs just satisfy the handler's overlay call. def get_last_original_messages(self): # noqa: ANN201 - return [] + return copy.deepcopy(self._previous_original) def get_last_forwarded_messages(self): # noqa: ANN201 - return [] + return copy.deepcopy(self._previous_forwarded) def update_from_response(self, **kwargs): # noqa: ANN003 return None @@ -112,6 +120,55 @@ def test_openai_cache_mode_freezes_previous_turns() -> None: assert captured["frozen_message_count"] == 2 +def test_openai_handler_replays_nonempty_cached_prefix() -> None: + captured = {} + previous_original = [{"role": "user", "content": "original prefix"}] + previous_forwarded = [{"role": "user", "content": "comp"}] + fake_tracker = _FakePrefixTracker(0, previous_original, previous_forwarded) + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: ( + "stable-session" + ) + proxy.session_tracker_store.resolve_tracker = lambda *args, **kwargs: fake_tracker + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "chatcmpl_overlay", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 20, "completion_tokens": 3, "total_tokens": 23}, + }, + ) + + proxy._retry_request = _fake_retry + response = client.post( + "/v1/chat/completions", + headers={"authorization": "Bearer test-key"}, + json={ + "model": "gpt-4o-mini", + "messages": [ + {"role": "user", "content": "original prefix"}, + {"role": "user", "content": "new suffix"}, + ], + }, + ) + + assert response.status_code == 200 + assert captured["body"]["messages"] == [ + previous_forwarded[0], + {"role": "user", "content": "new suffix"}, + ] + + @pytest.mark.parametrize("tail_role", ["tool", "function"]) def test_openai_cache_mode_keeps_final_tool_observation_mutable(tail_role: str) -> None: captured = {}