Harden cache-mode immutability for OpenAI and fix stats mode reporting

This commit is contained in:
JerrettDavis 2026-04-04 14:36:29 -05:00
parent 54419ad8b8
commit 2625789a28
4 changed files with 243 additions and 2 deletions

View file

@ -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:

View file

@ -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

View file

@ -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"

View file

@ -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]