mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
Harden cache-mode immutability for OpenAI and fix stats mode reporting
This commit is contained in:
parent
54419ad8b8
commit
2625789a28
4 changed files with 243 additions and 2 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
157
tests/test_proxy_openai_cache_stability.py
Normal file
157
tests/test_proxy_openai_cache_stability.py
Normal 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]
|
||||
Loading…
Add table
Add a link
Reference in a new issue