mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
Harden Anthropic prefix cache stability across proxy and batch paths
This commit is contained in:
parent
e8ab444f09
commit
7d02829f02
6 changed files with 4559 additions and 3816 deletions
|
|
@ -20,6 +20,7 @@ Each provider has different result formats, but the logic is the same:
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
|
@ -342,7 +343,24 @@ class BatchResultProcessor:
|
|||
"max_tokens": request_context.extras.get("max_tokens", 4096),
|
||||
}
|
||||
if tools:
|
||||
body["tools"] = tools
|
||||
def _tool_sort_key(tool: dict[str, Any]) -> tuple[str, str]:
|
||||
name = (
|
||||
str(tool.get("name", ""))
|
||||
or str(tool.get("function", {}).get("name", ""))
|
||||
or str(tool.get("type", ""))
|
||||
)
|
||||
try:
|
||||
canonical = json.dumps(
|
||||
tool,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
ensure_ascii=False,
|
||||
)
|
||||
except Exception:
|
||||
canonical = str(tool)
|
||||
return (name, canonical)
|
||||
|
||||
body["tools"] = sorted(tools, key=_tool_sort_key)
|
||||
|
||||
response = await self.http_client.post(
|
||||
url,
|
||||
|
|
|
|||
|
|
@ -25,6 +25,110 @@ logger = logging.getLogger("headroom.proxy")
|
|||
class AnthropicHandlerMixin:
|
||||
"""Mixin providing Anthropic API handler methods for HeadroomProxy."""
|
||||
|
||||
@staticmethod
|
||||
def _tool_sort_key(tool: dict[str, Any]) -> tuple[str, str]:
|
||||
"""Deterministic sort key for Anthropic/OpenAI-style tool definitions."""
|
||||
name = (
|
||||
str(tool.get("name", ""))
|
||||
or str(tool.get("function", {}).get("name", ""))
|
||||
or str(tool.get("type", ""))
|
||||
)
|
||||
try:
|
||||
canonical = json.dumps(tool, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
|
||||
except Exception:
|
||||
canonical = str(tool)
|
||||
return (name, canonical)
|
||||
|
||||
@classmethod
|
||||
def _sort_tools_deterministically(
|
||||
cls, tools: list[dict[str, Any]] | None
|
||||
) -> list[dict[str, Any]] | None:
|
||||
"""Return tools in deterministic order to preserve prompt-cache stability."""
|
||||
if not tools:
|
||||
return tools
|
||||
return sorted(tools, key=cls._tool_sort_key)
|
||||
|
||||
@staticmethod
|
||||
def _compress_latest_user_turn_images_cache_safe(
|
||||
messages: list[dict[str, Any]],
|
||||
*,
|
||||
frozen_message_count: int,
|
||||
compressor: Any,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Compress images only in the latest non-frozen user turn.
|
||||
|
||||
This avoids rewriting historical image bytes that may already be in the
|
||||
provider prefix cache.
|
||||
"""
|
||||
if not messages:
|
||||
return messages
|
||||
|
||||
target_idx = len(messages) - 1
|
||||
if target_idx < frozen_message_count:
|
||||
return messages
|
||||
target_msg = messages[target_idx]
|
||||
if target_msg.get("role") != "user":
|
||||
return messages
|
||||
content = target_msg.get("content")
|
||||
if not isinstance(content, list):
|
||||
return messages
|
||||
if not any(isinstance(block, dict) and block.get("type") == "image" for block in content):
|
||||
return messages
|
||||
|
||||
compressed_one = compressor.compress([target_msg], provider="anthropic")
|
||||
if not compressed_one:
|
||||
return messages
|
||||
|
||||
if compressed_one[0] == target_msg:
|
||||
return messages
|
||||
|
||||
updated = list(messages)
|
||||
updated[target_idx] = compressed_one[0]
|
||||
return updated
|
||||
|
||||
@staticmethod
|
||||
def _append_context_to_latest_non_frozen_user_turn(
|
||||
messages: list[dict[str, Any]],
|
||||
context_text: str,
|
||||
*,
|
||||
frozen_message_count: int,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Append context only to the latest non-frozen user text turn.
|
||||
|
||||
Returns input unchanged if no eligible user text turn exists.
|
||||
"""
|
||||
if not messages or not context_text:
|
||||
return messages
|
||||
|
||||
i = len(messages) - 1
|
||||
if i < frozen_message_count:
|
||||
return messages
|
||||
msg = messages[i]
|
||||
if msg.get("role") != "user":
|
||||
return messages
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, str):
|
||||
updated = list(messages)
|
||||
updated[i] = {**msg, "content": content + "\n\n" + context_text}
|
||||
return updated
|
||||
return messages
|
||||
|
||||
@staticmethod
|
||||
def _strict_previous_turn_frozen_count(
|
||||
messages: list[dict[str, Any]],
|
||||
base_frozen_count: int,
|
||||
) -> int:
|
||||
"""Freeze all prior turns; only the final turn is mutable.
|
||||
|
||||
If the final message is not a user turn, freeze everything.
|
||||
"""
|
||||
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)
|
||||
|
||||
async def handle_anthropic_messages(
|
||||
self,
|
||||
request: Request,
|
||||
|
|
@ -107,19 +211,6 @@ class AnthropicHandlerMixin:
|
|||
if _bypass:
|
||||
logger.info(f"[{request_id}] Bypass: skipping compression (header)")
|
||||
|
||||
# Image compression (before text optimization)
|
||||
if self.config.image_optimize and messages and not _bypass:
|
||||
compressor = _get_image_compressor()
|
||||
if compressor and compressor.has_images(messages):
|
||||
messages = compressor.compress(messages, provider="anthropic")
|
||||
if compressor.last_result:
|
||||
logger.info(
|
||||
f"Image compression: {compressor.last_result.technique.value} "
|
||||
f"({compressor.last_result.savings_percent:.0f}% saved, "
|
||||
f"{compressor.last_result.original_tokens} -> "
|
||||
f"{compressor.last_result.compressed_tokens} tokens)"
|
||||
)
|
||||
|
||||
# Extract headers and tags
|
||||
headers = dict(request.headers.items())
|
||||
headers.pop("host", None)
|
||||
|
|
@ -213,6 +304,28 @@ class AnthropicHandlerMixin:
|
|||
session_id = self.session_tracker_store.compute_session_id(request, model, messages)
|
||||
prefix_tracker = self.session_tracker_store.get_or_create(session_id, "anthropic")
|
||||
frozen_message_count = prefix_tracker.get_frozen_message_count()
|
||||
frozen_message_count = self._strict_previous_turn_frozen_count(
|
||||
messages,
|
||||
frozen_message_count,
|
||||
)
|
||||
|
||||
# Image compression (cache-safe): only compress latest non-frozen user turn.
|
||||
# Rewriting historical image bytes can invalidate Anthropic prompt caches.
|
||||
if self.config.image_optimize and messages and not _bypass:
|
||||
compressor = _get_image_compressor()
|
||||
if compressor and compressor.has_images(messages):
|
||||
messages = self._compress_latest_user_turn_images_cache_safe(
|
||||
messages,
|
||||
frozen_message_count=frozen_message_count,
|
||||
compressor=compressor,
|
||||
)
|
||||
if compressor.last_result:
|
||||
logger.info(
|
||||
f"Image compression: {compressor.last_result.technique.value} "
|
||||
f"({compressor.last_result.savings_percent:.0f}% saved, "
|
||||
f"{compressor.last_result.original_tokens} -> "
|
||||
f"{compressor.last_result.compressed_tokens} tokens)"
|
||||
)
|
||||
|
||||
_compression_failed = False
|
||||
original_messages = messages # Preserve for 400-retry fallback
|
||||
|
|
@ -235,7 +348,13 @@ class AnthropicHandlerMixin:
|
|||
working_messages = comp_cache.apply_cached(messages)
|
||||
|
||||
# Re-freeze boundary: consecutive stable messages from start
|
||||
frozen_message_count = comp_cache.compute_frozen_count(messages)
|
||||
# Safety: never freeze beyond provider-confirmed cached prefix.
|
||||
cache_frozen_count = comp_cache.compute_frozen_count(messages)
|
||||
frozen_message_count = min(frozen_message_count, cache_frozen_count)
|
||||
frozen_message_count = self._strict_previous_turn_frozen_count(
|
||||
messages,
|
||||
frozen_message_count,
|
||||
)
|
||||
|
||||
result = await asyncio.wait_for(
|
||||
asyncio.to_thread(
|
||||
|
|
@ -323,11 +442,18 @@ class AnthropicHandlerMixin:
|
|||
if (
|
||||
self.config.ccr_inject_tool or self.config.ccr_inject_system_instructions
|
||||
) and not _bypass:
|
||||
inject_system_instructions = self.config.ccr_inject_system_instructions
|
||||
if inject_system_instructions and frozen_message_count > 0:
|
||||
logger.info(
|
||||
f"[{request_id}] CCR: skipping system instruction injection "
|
||||
f"(frozen prefix={frozen_message_count}) to preserve cache"
|
||||
)
|
||||
inject_system_instructions = False
|
||||
# Create fresh injector to avoid state leakage between requests
|
||||
injector = CCRToolInjector(
|
||||
provider="anthropic",
|
||||
inject_tool=self.config.ccr_inject_tool,
|
||||
inject_system_instructions=self.config.ccr_inject_system_instructions,
|
||||
inject_system_instructions=inject_system_instructions,
|
||||
)
|
||||
optimized_messages, tools, was_injected = injector.process_request(
|
||||
optimized_messages, tools
|
||||
|
|
@ -392,15 +518,11 @@ class AnthropicHandlerMixin:
|
|||
f"[{request_id}] CCR: Proactively expanded {len(expansions)} context(s) "
|
||||
f"based on query relevance"
|
||||
)
|
||||
# Append to the last user message
|
||||
if optimized_messages and optimized_messages[-1].get("role") == "user":
|
||||
last_msg = optimized_messages[-1]
|
||||
content = last_msg.get("content", "")
|
||||
if isinstance(content, str):
|
||||
optimized_messages[-1] = {
|
||||
**last_msg,
|
||||
"content": content + "\n\n" + expansion_text,
|
||||
}
|
||||
optimized_messages = self._append_context_to_latest_non_frozen_user_turn(
|
||||
optimized_messages,
|
||||
expansion_text,
|
||||
frozen_message_count=frozen_message_count,
|
||||
)
|
||||
|
||||
# Traffic Learner: Extract patterns from inbound tool results
|
||||
if self.traffic_learner:
|
||||
|
|
@ -440,12 +562,23 @@ class AnthropicHandlerMixin:
|
|||
memory_user_id, optimized_messages
|
||||
)
|
||||
if memory_context:
|
||||
optimized_messages = self._inject_system_context(
|
||||
optimized_messages, memory_context, body=body
|
||||
)
|
||||
logger.info(
|
||||
f"[{request_id}] Memory: Injected {len(memory_context)} chars of context"
|
||||
)
|
||||
if frozen_message_count > 0:
|
||||
optimized_messages = self._append_context_to_latest_non_frozen_user_turn(
|
||||
optimized_messages,
|
||||
memory_context,
|
||||
frozen_message_count=frozen_message_count,
|
||||
)
|
||||
logger.info(
|
||||
f"[{request_id}] Memory: Appended {len(memory_context)} chars "
|
||||
f"to latest non-frozen user turn (prefix cache-safe)"
|
||||
)
|
||||
else:
|
||||
optimized_messages = self._inject_system_context(
|
||||
optimized_messages, memory_context, body=body
|
||||
)
|
||||
logger.info(
|
||||
f"[{request_id}] Memory: Injected {len(memory_context)} chars of context"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"[{request_id}] Memory: Context injection failed: {e}")
|
||||
|
||||
|
|
@ -482,6 +615,7 @@ class AnthropicHandlerMixin:
|
|||
# Update body
|
||||
body["messages"] = optimized_messages
|
||||
if tools is not None:
|
||||
tools = self._sort_tools_deterministically(tools)
|
||||
body["tools"] = tools
|
||||
|
||||
# Forward request - use Bedrock backend if configured, otherwise direct API
|
||||
|
|
@ -1079,12 +1213,21 @@ class AnthropicHandlerMixin:
|
|||
for batch_req in requests_list:
|
||||
custom_id = batch_req.get("custom_id", "")
|
||||
params = batch_req.get("params", {})
|
||||
canonical_params = dict(params)
|
||||
canonical_tools = canonical_params.get("tools")
|
||||
if canonical_tools is not None:
|
||||
canonical_params["tools"] = self._sort_tools_deterministically(canonical_tools)
|
||||
messages = params.get("messages", [])
|
||||
model = params.get("model", "unknown")
|
||||
|
||||
if not messages or not self.config.optimize:
|
||||
# No messages or optimization disabled - pass through unchanged
|
||||
compressed_requests.append(batch_req)
|
||||
compressed_requests.append(
|
||||
{
|
||||
"custom_id": custom_id,
|
||||
"params": canonical_params,
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
# Apply optimization
|
||||
|
|
@ -1095,6 +1238,7 @@ class AnthropicHandlerMixin:
|
|||
model=model,
|
||||
model_limit=context_limit,
|
||||
context=extract_user_query(messages),
|
||||
frozen_message_count=self._strict_previous_turn_frozen_count(messages, 0),
|
||||
)
|
||||
|
||||
optimized_messages = result.messages
|
||||
|
|
@ -1109,7 +1253,7 @@ class AnthropicHandlerMixin:
|
|||
total_tokens_saved += tokens_saved
|
||||
|
||||
# CCR Tool Injection: Inject retrieval tool if compression occurred
|
||||
tools = params.get("tools")
|
||||
tools = canonical_params.get("tools")
|
||||
if self.config.ccr_inject_tool and tokens_saved > 0:
|
||||
injector = CCRToolInjector(
|
||||
provider="anthropic",
|
||||
|
|
@ -1127,7 +1271,7 @@ class AnthropicHandlerMixin:
|
|||
# Create compressed batch request
|
||||
compressed_params = {**params, "messages": optimized_messages}
|
||||
if tools is not None:
|
||||
compressed_params["tools"] = tools
|
||||
compressed_params["tools"] = self._sort_tools_deterministically(tools)
|
||||
compressed_requests.append(
|
||||
{
|
||||
"custom_id": custom_id,
|
||||
|
|
|
|||
|
|
@ -102,7 +102,7 @@ reports = [
|
|||
]
|
||||
# any-llm multi-provider backend (requires Python 3.11+)
|
||||
anyllm = [
|
||||
"any-llm-sdk>=1.0.0",
|
||||
"any-llm-sdk>=1.0.0; python_version >= '3.11'",
|
||||
]
|
||||
# LangChain integration
|
||||
langchain = [
|
||||
|
|
|
|||
|
|
@ -765,6 +765,8 @@ class TestContinuationCalls:
|
|||
call_args = http_client.post.call_args
|
||||
assert "tools" in call_args.kwargs["json"]
|
||||
assert len(call_args.kwargs["json"]["tools"]) == 2
|
||||
sent_names = [tool.get("name") for tool in call_args.kwargs["json"]["tools"]]
|
||||
assert sent_names == sorted(sent_names)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_make_continuation_call_unknown_provider(self):
|
||||
|
|
|
|||
568
tests/test_proxy_anthropic_cache_stability.py
Normal file
568
tests/test_proxy_anthropic_cache_stability.py
Normal file
|
|
@ -0,0 +1,568 @@
|
|||
"""Regression tests for Anthropic prefix-cache stability in proxy mode."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("fastapi")
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from headroom.proxy.handlers.anthropic import AnthropicHandlerMixin
|
||||
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
|
||||
|
||||
|
||||
class _FakeImageCompressor:
|
||||
def __init__(self):
|
||||
self.last_result = None
|
||||
|
||||
def has_images(self, messages): # noqa: ANN001
|
||||
return True
|
||||
|
||||
def compress(self, messages, provider="anthropic"): # noqa: ANN001
|
||||
assert provider == "anthropic"
|
||||
assert len(messages) == 1
|
||||
msg = messages[0]
|
||||
content = msg["content"]
|
||||
updated_content = []
|
||||
for block in content:
|
||||
if isinstance(block, dict) and block.get("type") == "image":
|
||||
src = block.get("source", {})
|
||||
updated_content.append(
|
||||
{
|
||||
"type": "image",
|
||||
"source": {**src, "data": "COMPRESSED_IMAGE_BYTES"},
|
||||
}
|
||||
)
|
||||
else:
|
||||
updated_content.append(block)
|
||||
return [{**msg, "content": updated_content}]
|
||||
|
||||
|
||||
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=True,
|
||||
)
|
||||
app = create_app(config)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_anthropic_tools_sorted_deterministically_before_forward() -> None:
|
||||
captured = {}
|
||||
with _make_proxy_client() as client:
|
||||
proxy = client.app.state.proxy
|
||||
|
||||
async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001
|
||||
captured["body"] = body
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"usage": {
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 3,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
proxy._retry_request = _fake_retry
|
||||
|
||||
response = client.post(
|
||||
"/v1/messages",
|
||||
headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"},
|
||||
json={
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 128,
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"tools": [
|
||||
{"name": "zeta", "description": "z", "input_schema": {"type": "object"}},
|
||||
{"name": "alpha", "description": "a", "input_schema": {"type": "object"}},
|
||||
{"name": "mu", "description": "m", "input_schema": {"type": "object"}},
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
sent_tools = captured["body"]["tools"]
|
||||
assert [t["name"] for t in sent_tools] == ["alpha", "mu", "zeta"]
|
||||
|
||||
|
||||
def test_image_compression_only_applies_to_latest_non_frozen_user_turn() -> None:
|
||||
fake_compressor = _FakeImageCompressor()
|
||||
|
||||
old_image = {
|
||||
"type": "image",
|
||||
"source": {"type": "base64", "media_type": "image/png", "data": "OLD_IMAGE_BYTES"},
|
||||
}
|
||||
new_image = {
|
||||
"type": "image",
|
||||
"source": {"type": "base64", "media_type": "image/png", "data": "NEW_IMAGE_BYTES"},
|
||||
}
|
||||
messages = [
|
||||
{"role": "user", "content": [old_image, {"type": "text", "text": "old image turn"}]},
|
||||
{"role": "assistant", "content": "ack"},
|
||||
{"role": "user", "content": [new_image, {"type": "text", "text": "new image turn"}]},
|
||||
]
|
||||
|
||||
result = AnthropicHandlerMixin._compress_latest_user_turn_images_cache_safe(
|
||||
messages,
|
||||
frozen_message_count=1,
|
||||
compressor=fake_compressor,
|
||||
)
|
||||
|
||||
# Frozen prefix must remain byte-identical.
|
||||
assert result[0]["content"][0]["source"]["data"] == "OLD_IMAGE_BYTES"
|
||||
# Latest non-frozen user turn is eligible for compression.
|
||||
assert result[2]["content"][0]["source"]["data"] == "COMPRESSED_IMAGE_BYTES"
|
||||
|
||||
|
||||
def test_image_compression_does_not_touch_previous_turns_if_last_message_not_user() -> None:
|
||||
fake_compressor = _FakeImageCompressor()
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "image",
|
||||
"source": {"type": "base64", "media_type": "image/png", "data": "OLD_IMAGE_BYTES"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "assistant", "content": "last turn is assistant"},
|
||||
]
|
||||
result = AnthropicHandlerMixin._compress_latest_user_turn_images_cache_safe(
|
||||
messages,
|
||||
frozen_message_count=0,
|
||||
compressor=fake_compressor,
|
||||
)
|
||||
assert result[0]["content"][0]["source"]["data"] == "OLD_IMAGE_BYTES"
|
||||
|
||||
|
||||
def test_anthropic_batch_tools_sorted_deterministically_before_forward() -> None:
|
||||
captured = {}
|
||||
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)
|
||||
|
||||
with TestClient(app) as client:
|
||||
proxy = client.app.state.proxy
|
||||
|
||||
async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001
|
||||
captured["body"] = body
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "msgbatch_1",
|
||||
"type": "message_batch",
|
||||
"processing_status": "in_progress",
|
||||
"request_counts": {"processing": 1, "succeeded": 0, "errored": 0, "canceled": 0},
|
||||
},
|
||||
)
|
||||
|
||||
proxy._retry_request = _fake_retry
|
||||
|
||||
response = client.post(
|
||||
"/v1/messages/batches",
|
||||
headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"},
|
||||
json={
|
||||
"requests": [
|
||||
{
|
||||
"custom_id": "req-1",
|
||||
"params": {
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 128,
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"tools": [
|
||||
{"name": "zeta", "description": "z", "input_schema": {"type": "object"}},
|
||||
{"name": "alpha", "description": "a", "input_schema": {"type": "object"}},
|
||||
{"name": "mu", "description": "m", "input_schema": {"type": "object"}},
|
||||
],
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
sent_tools = captured["body"]["requests"][0]["params"]["tools"]
|
||||
assert [t["name"] for t in sent_tools] == ["alpha", "mu", "zeta"]
|
||||
|
||||
|
||||
def test_append_context_targets_latest_non_frozen_user_turn() -> None:
|
||||
messages = [
|
||||
{"role": "user", "content": "frozen prefix"},
|
||||
{"role": "assistant", "content": "ack"},
|
||||
{"role": "user", "content": "active turn"},
|
||||
]
|
||||
result = AnthropicHandlerMixin._append_context_to_latest_non_frozen_user_turn(
|
||||
messages,
|
||||
"CTX",
|
||||
frozen_message_count=1,
|
||||
)
|
||||
assert result[0]["content"] == "frozen prefix"
|
||||
assert result[2]["content"].endswith("CTX")
|
||||
|
||||
|
||||
def test_append_context_does_not_touch_previous_turns_if_last_message_not_user() -> None:
|
||||
messages = [
|
||||
{"role": "user", "content": "previous user turn"},
|
||||
{"role": "assistant", "content": "assistant last"},
|
||||
]
|
||||
result = AnthropicHandlerMixin._append_context_to_latest_non_frozen_user_turn(
|
||||
messages,
|
||||
"CTX",
|
||||
frozen_message_count=0,
|
||||
)
|
||||
assert result[0]["content"] == "previous user turn"
|
||||
|
||||
|
||||
def test_token_headroom_freeze_is_capped_by_prefix_tracker() -> None:
|
||||
captured = {}
|
||||
with _make_proxy_client() as client:
|
||||
proxy = client.app.state.proxy
|
||||
proxy.config.optimize = True
|
||||
proxy.config.mode = "token_headroom"
|
||||
proxy.config.image_optimize = False
|
||||
|
||||
fake_tracker = _FakePrefixTracker(frozen_count=1)
|
||||
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
|
||||
|
||||
class _FakeCompressionCache:
|
||||
def apply_cached(self, messages): # noqa: ANN001
|
||||
return messages
|
||||
|
||||
def compute_frozen_count(self, messages): # noqa: ANN001
|
||||
return 99
|
||||
|
||||
def update_from_result(self, originals, compressed): # noqa: ANN001
|
||||
return None
|
||||
|
||||
proxy._get_compression_cache = lambda session_id: _FakeCompressionCache()
|
||||
|
||||
def _fake_apply(**kwargs):
|
||||
captured["frozen_message_count"] = kwargs.get("frozen_message_count")
|
||||
return SimpleNamespace(
|
||||
messages=kwargs["messages"],
|
||||
transforms_applied=[],
|
||||
timing={},
|
||||
tokens_before=50,
|
||||
tokens_after=50,
|
||||
waste_signals=None,
|
||||
)
|
||||
|
||||
proxy.anthropic_pipeline.apply = _fake_apply
|
||||
|
||||
async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "msg_tc_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"usage": {
|
||||
"input_tokens": 50,
|
||||
"output_tokens": 3,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
proxy._retry_request = _fake_retry
|
||||
|
||||
response = client.post(
|
||||
"/v1/messages",
|
||||
headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"},
|
||||
json={
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured["frozen_message_count"] == 1
|
||||
|
||||
|
||||
def test_memory_context_avoids_system_mutation_when_prefix_frozen() -> None:
|
||||
captured = {}
|
||||
with _make_proxy_client() as client:
|
||||
proxy = client.app.state.proxy
|
||||
proxy.config.optimize = False
|
||||
proxy.config.image_optimize = False
|
||||
proxy.config.ccr_proactive_expansion = False
|
||||
|
||||
fake_tracker = _FakePrefixTracker(frozen_count=1)
|
||||
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
|
||||
|
||||
proxy.memory_handler = SimpleNamespace(
|
||||
config=SimpleNamespace(inject_context=True, inject_tools=False),
|
||||
search_and_format_context=AsyncMock(return_value="MEMCTX"),
|
||||
has_memory_tool_calls=lambda resp, provider: False,
|
||||
)
|
||||
|
||||
async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001
|
||||
captured["body"] = body
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "msg_mem_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"usage": {
|
||||
"input_tokens": 20,
|
||||
"output_tokens": 3,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
proxy._retry_request = _fake_retry
|
||||
|
||||
response = client.post(
|
||||
"/v1/messages",
|
||||
headers={
|
||||
"x-api-key": "test-key",
|
||||
"anthropic-version": "2023-06-01",
|
||||
"x-headroom-user-id": "u1",
|
||||
},
|
||||
json={
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 64,
|
||||
"system": "base system",
|
||||
"messages": [
|
||||
{"role": "user", "content": "frozen prefix"},
|
||||
{"role": "assistant", "content": "ack"},
|
||||
{"role": "user", "content": "latest user"},
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
sent = captured["body"]
|
||||
assert sent["system"] == "base system"
|
||||
assert sent["messages"][2]["content"].endswith("MEMCTX")
|
||||
|
||||
|
||||
def test_ccr_system_instruction_injection_disabled_when_prefix_frozen(monkeypatch) -> None:
|
||||
captured = {"inject_system": None}
|
||||
with _make_proxy_client() as client:
|
||||
proxy = client.app.state.proxy
|
||||
proxy.config.optimize = False
|
||||
proxy.config.image_optimize = False
|
||||
proxy.config.ccr_inject_tool = False
|
||||
proxy.config.ccr_inject_system_instructions = True
|
||||
|
||||
fake_tracker = _FakePrefixTracker(frozen_count=1)
|
||||
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
|
||||
|
||||
class _FakeInjector:
|
||||
def __init__(
|
||||
self,
|
||||
provider, # noqa: ANN001
|
||||
inject_tool, # noqa: ANN001
|
||||
inject_system_instructions, # noqa: ANN001
|
||||
):
|
||||
captured["inject_system"] = inject_system_instructions
|
||||
self.has_compressed_content = False
|
||||
self.detected_hashes = []
|
||||
|
||||
def process_request(self, messages, tools): # noqa: ANN001
|
||||
return messages, tools, False
|
||||
|
||||
monkeypatch.setattr("headroom.ccr.CCRToolInjector", _FakeInjector)
|
||||
|
||||
async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "msg_ccr_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"usage": {
|
||||
"input_tokens": 20,
|
||||
"output_tokens": 3,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
proxy._retry_request = _fake_retry
|
||||
|
||||
response = client.post(
|
||||
"/v1/messages",
|
||||
headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"},
|
||||
json={
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured["inject_system"] is False
|
||||
|
||||
|
||||
def test_previous_turns_always_frozen_only_final_turn_mutable() -> None:
|
||||
captured = {}
|
||||
with _make_proxy_client() as client:
|
||||
proxy = client.app.state.proxy
|
||||
proxy.config.optimize = True
|
||||
proxy.config.mode = "cost_savings"
|
||||
proxy.config.image_optimize = False
|
||||
|
||||
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=80,
|
||||
tokens_after=80,
|
||||
waste_signals=None,
|
||||
)
|
||||
|
||||
proxy.anthropic_pipeline.apply = _fake_apply
|
||||
|
||||
async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "msg_frz_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "text", "text": "ok"}],
|
||||
"usage": {
|
||||
"input_tokens": 80,
|
||||
"output_tokens": 3,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_creation_input_tokens": 0,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
proxy._retry_request = _fake_retry
|
||||
|
||||
response = client.post(
|
||||
"/v1/messages",
|
||||
headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"},
|
||||
json={
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 64,
|
||||
"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_batch_optimization_freezes_previous_turns_only() -> None:
|
||||
captured = {}
|
||||
with _make_proxy_client() as client:
|
||||
proxy = client.app.state.proxy
|
||||
proxy.config.optimize = True
|
||||
proxy.config.image_optimize = False
|
||||
proxy.config.ccr_inject_tool = False
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
proxy.anthropic_pipeline.apply = _fake_apply
|
||||
|
||||
async def _fake_retry(method, url, headers, body, stream=False): # noqa: ANN001
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "msgbatch_2",
|
||||
"type": "message_batch",
|
||||
"processing_status": "in_progress",
|
||||
"request_counts": {"processing": 1, "succeeded": 0, "errored": 0, "canceled": 0},
|
||||
},
|
||||
)
|
||||
|
||||
proxy._retry_request = _fake_retry
|
||||
|
||||
response = client.post(
|
||||
"/v1/messages/batches",
|
||||
headers={"x-api-key": "test-key", "anthropic-version": "2023-06-01"},
|
||||
json={
|
||||
"requests": [
|
||||
{
|
||||
"custom_id": "req-1",
|
||||
"params": {
|
||||
"model": "claude-sonnet-4-6",
|
||||
"max_tokens": 128,
|
||||
"messages": [
|
||||
{"role": "user", "content": "old turn"},
|
||||
{"role": "assistant", "content": "old assistant"},
|
||||
{"role": "user", "content": "current turn"},
|
||||
],
|
||||
},
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured["frozen_message_count"] == 2
|
||||
Loading…
Add table
Add a link
Reference in a new issue