From 2625789a2885263bd279ef6e1d35f38845af9fa8 Mon Sep 17 00:00:00 2001 From: JerrettDavis Date: Sat, 4 Apr 2026 14:36:29 -0500 Subject: [PATCH] Harden cache-mode immutability for OpenAI and fix stats mode reporting --- headroom/proxy/handlers/openai.py | 57 +++++++- headroom/proxy/server.py | 2 +- tests/test_proxy_modes.py | 29 ++++ tests/test_proxy_openai_cache_stability.py | 157 +++++++++++++++++++++ 4 files changed, 243 insertions(+), 2 deletions(-) create mode 100644 tests/test_proxy_openai_cache_stability.py diff --git a/headroom/proxy/handlers/openai.py b/headroom/proxy/handlers/openai.py index 3f9b70d38..070d12a0e 100644 --- a/headroom/proxy/handlers/openai.py +++ b/headroom/proxy/handlers/openai.py @@ -6,6 +6,7 @@ Contains all OpenAI Chat Completions, Responses API, and passthrough handlers. from __future__ import annotations import asyncio +import copy import contextlib import json import logging @@ -26,6 +27,43 @@ logger = logging.getLogger("headroom.proxy") class OpenAIHandlerMixin: """Mixin providing OpenAI API handler methods for HeadroomProxy.""" + @staticmethod + def _strict_previous_turn_frozen_count( + messages: list[dict[str, Any]], + base_frozen_count: int, + ) -> int: + """Freeze all prior turns in cache mode; only final user turn is mutable.""" + if not messages: + return base_frozen_count + final_idx = len(messages) - 1 + if messages[final_idx].get("role") == "user": + return max(base_frozen_count, final_idx) + return len(messages) + + @staticmethod + def _restore_frozen_prefix( + original_messages: list[dict[str, Any]], + candidate_messages: list[dict[str, Any]], + *, + frozen_message_count: int, + ) -> tuple[list[dict[str, Any]], int]: + """Force frozen prefix bytes to match original request exactly.""" + if frozen_message_count <= 0 or not original_messages: + return candidate_messages, 0 + + frozen = min(frozen_message_count, len(original_messages)) + restored = list(candidate_messages) + + if len(restored) < frozen: + return list(original_messages[:frozen]) + restored, frozen + + changed = 0 + for idx in range(frozen): + if restored[idx] != original_messages[idx]: + restored[idx] = original_messages[idx] + changed += 1 + return restored, changed + async def handle_openai_chat( self, request: Request, @@ -41,7 +79,7 @@ class OpenAIHandlerMixin: MAX_REQUEST_BODY_SIZE, _read_request_json, ) - from headroom.proxy.modes import is_token_mode + from headroom.proxy.modes import is_cache_mode, is_token_mode from headroom.tokenizers import get_tokenizer from headroom.utils import extract_user_query @@ -78,6 +116,7 @@ class OpenAIHandlerMixin: ) model = body.get("model", "unknown") messages = body.get("messages", []) + original_client_messages = copy.deepcopy(messages) # Validate message array size if len(messages) > MAX_MESSAGE_ARRAY_LENGTH: @@ -184,6 +223,11 @@ class OpenAIHandlerMixin: openai_session_id, "openai" ) openai_frozen_count = openai_prefix_tracker.get_frozen_message_count() + if is_cache_mode(self.config.mode): + openai_frozen_count = self._strict_previous_turn_frozen_count( + original_client_messages, + openai_frozen_count, + ) _compression_failed = False original_messages = messages # Preserve for 400-retry fallback @@ -304,6 +348,17 @@ class OpenAIHandlerMixin: ) # Query Echo: disabled — hurts prefix caching in long conversations. + if is_cache_mode(self.config.mode): + optimized_messages, restored_count = self._restore_frozen_prefix( + original_client_messages, + optimized_messages, + frozen_message_count=openai_frozen_count, + ) + if restored_count > 0: + logger.warning( + f"[{request_id}] Restored {restored_count} frozen prefix message(s) " + "to preserve cache stability (openai)" + ) body["messages"] = optimized_messages if tools is not None: diff --git a/headroom/proxy/server.py b/headroom/proxy/server.py index 08dd7d1da..87215aa61 100644 --- a/headroom/proxy/server.py +++ b/headroom/proxy/server.py @@ -1048,7 +1048,7 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI: "total_tokens_saved": total_tokens_saved, } else: - compression_cache_stats = {"mode": PROXY_MODE_TOKEN} + compression_cache_stats = {"mode": proxy.config.mode} # Build unified savings summary (all layers) compression_tokens = m.tokens_saved_total diff --git a/tests/test_proxy_modes.py b/tests/test_proxy_modes.py index a3fa06544..a361fcba9 100644 --- a/tests/test_proxy_modes.py +++ b/tests/test_proxy_modes.py @@ -1,5 +1,7 @@ """Tests for proxy token/cache mode normalization.""" +import pytest + from headroom.proxy.modes import ( PROXY_MODE_CACHE, PROXY_MODE_TOKEN, @@ -28,3 +30,30 @@ def test_proxy_mode_invalid_falls_back_to_default() -> None: def test_proxy_mode_predicates() -> None: assert is_token_mode("token_headroom") is True assert is_cache_mode("cost_savings") is True + + +def test_stats_reports_configured_mode_for_compression_cache() -> None: + pytest.importorskip("fastapi") + from fastapi.testclient import TestClient + + from headroom.proxy.server import ProxyConfig, create_app + + app = create_app( + ProxyConfig( + mode="cache", + optimize=False, + cache_enabled=False, + rate_limit_enabled=False, + cost_tracking_enabled=False, + log_requests=False, + ccr_inject_tool=False, + ccr_handle_responses=False, + ccr_context_tracking=False, + ) + ) + + with TestClient(app) as client: + response = client.get("/stats") + assert response.status_code == 200 + data = response.json() + assert data["compression_cache"]["mode"] == "cache" diff --git a/tests/test_proxy_openai_cache_stability.py b/tests/test_proxy_openai_cache_stability.py new file mode 100644 index 000000000..a016fb9c0 --- /dev/null +++ b/tests/test_proxy_openai_cache_stability.py @@ -0,0 +1,157 @@ +"""Regression tests for OpenAI cache-mode stability in proxy mode.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import httpx +import pytest + +pytest.importorskip("fastapi") + +from fastapi.testclient import TestClient + +from headroom.proxy.server import ProxyConfig, create_app + + +class _FakePrefixTracker: + def __init__(self, frozen_count: int): + self._frozen_count = frozen_count + + def get_frozen_message_count(self) -> int: + return self._frozen_count + + def update_from_response(self, **kwargs): # noqa: ANN003 + return None + + +def _make_proxy_client() -> TestClient: + config = ProxyConfig( + optimize=False, + cache_enabled=False, + rate_limit_enabled=False, + cost_tracking_enabled=False, + log_requests=False, + ccr_inject_tool=False, + ccr_handle_responses=False, + ccr_context_tracking=False, + image_optimize=False, + ) + app = create_app(config) + return TestClient(app) + + +def test_openai_cache_mode_freezes_previous_turns() -> None: + captured = {} + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.mode = "cache" + + fake_tracker = _FakePrefixTracker(frozen_count=0) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: "stable-session" + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + def _fake_apply(**kwargs): + captured["frozen_message_count"] = kwargs.get("frozen_message_count") + return SimpleNamespace( + messages=kwargs["messages"], + transforms_applied=[], + timing={}, + tokens_before=60, + tokens_after=60, + waste_signals=None, + ) + + proxy.openai_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001 + return httpx.Response( + 200, + json={ + "id": "chatcmpl_1", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 60, "completion_tokens": 3, "total_tokens": 63}, + }, + ) + + 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": "turn1"}, + {"role": "assistant", "content": "turn1-assistant"}, + {"role": "user", "content": "current turn"}, + ], + }, + ) + + assert response.status_code == 200 + assert captured["frozen_message_count"] == 2 + + +def test_openai_cache_mode_restores_mutated_frozen_prefix() -> None: + captured = {} + with _make_proxy_client() as client: + proxy = client.app.state.proxy + proxy.config.optimize = True + proxy.config.mode = "cache" + + fake_tracker = _FakePrefixTracker(frozen_count=0) + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: "stable-session" + proxy.session_tracker_store.get_or_create = lambda session_id, provider: fake_tracker + + original_messages = [ + {"role": "user", "content": "turn1"}, + {"role": "assistant", "content": "turn1-assistant"}, + {"role": "user", "content": "current turn"}, + ] + + def _fake_apply(**kwargs): + mutated = list(kwargs["messages"]) + mutated[0] = {**mutated[0], "content": "MUTATED_PREFIX"} + return SimpleNamespace( + messages=mutated, + transforms_applied=["fake:mutated"], + timing={}, + tokens_before=70, + tokens_after=65, + waste_signals=None, + ) + + proxy.openai_pipeline.apply = _fake_apply + + async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001 + captured["body"] = body + return httpx.Response( + 200, + json={ + "id": "chatcmpl_2", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 65, "completion_tokens": 3, "total_tokens": 68}, + }, + ) + + proxy._retry_request = _fake_retry + + response = client.post( + "/v1/chat/completions", + headers={"authorization": "Bearer test-key"}, + json={ + "model": "gpt-4o-mini", + "messages": original_messages, + }, + ) + + assert response.status_code == 200 + sent_messages = captured["body"]["messages"] + assert sent_messages[0] == original_messages[0] + assert sent_messages[1] == original_messages[1]