fix(caching): count tool_call cache_control marks in the injection census (#43556)

* fix(caching): count tool_call cache_control marks in the injection census

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(caching): remove cache census casts

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(caching): only skip injection on message or content marks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(caching): skip injection on messages whose tool calls carry marks

Reverts 9f08d8aef8. A default 5m mark injected on assistant text lands
before the client's 1h tool_use mark, which Anthropic rejects with a 400
because a 1h breakpoint must not follow a 5m one. Keeping the full census
in the skip check leaves the client's tool_call breakpoint as the only
one on that message.

* fix(caching): count every client tool_call cache mark in the breakpoint census

The census gated tool_call marks on type function and dict shape, so a client mark on a call without a type or with a string cache_control slipped past the count and injection overflowed the 4 breakpoint cap. Count any non-None tool_call mark, keep the server tool exclusion, and add integration cells for the capped surfaces, the yaml stand-down, Bedrock and Gemini, and router affinity

* test(integration): hold the upstream so the worker kill lands mid-burst

* refactor(caching): reuse the transform's server tool lookup in the breakpoint census

The census now calls the same helper the Anthropic transform uses to decide
whether a marked tool call becomes a server tool block, so the two cannot
drift apart. The owned-proxy burst test waits up to 90 seconds for the burst
to reach the wire before it kills a worker

* refactor(anthropic): move the server tool rebuild check under llms/anthropic

* test(integration): audit the tool call mark census across chat, messages, responses, and chaos

Adds the /audit cells for the breakpoint census on assistant tool_calls marks: the
Responses stream bridge, the OpenAI and Anthropic SDK clients, in-process Pydantic
messages, response cache twins, request-level points, a provider 401 on a capped request,
malformed provider_specific_fields and tool_call ids, null or empty points, a points
update mid-burst, a proxy restart mid-burst, and a worker SIGKILL that picks the worker
holding the burst's upstream connections

* test(integration): close the SDK clients the cache census cells open

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-03 17:33:48 -07:00 • committed by GitHub
parent 6370104c53
commit 9d16412341
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 2405 additions and 16 deletions

View file

@ -28,6 +28,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
from litellm.llms.anthropic.common_utils import (
is_claude_code_one_shot_subagent_request,
supports_anthropic_cache_control,
tool_call_is_rebuilt_as_server_tool_use,
)
from litellm.types.integrations.anthropic_cache_control_hook import (
GATEWAY_INJECTED_CACHE_METADATA_KEY,
@ -122,7 +123,30 @@ def targets_openai_api(api_base: object) -> bool:
def _carries_cache_breakpoint(block: object) -> bool:
return isinstance(block, dict) and any(block.get(key) is not None for key in CACHE_BREAKPOINT_KEYS)
return any(_attribute_or_key(block, key) is not None for key in CACHE_BREAKPOINT_KEYS)
def _attribute_or_key(value: object, key: str) -> object | None:
if hasattr(value, key):
return getattr(value, key)
if isinstance(value, Mapping):
return value.get(key)
return None
def _as_object_list(value: object | None) -> list[object] | None:
if not isinstance(value, list):
return None
return _validated_object_list(value)
def _tool_call_carries_cache_breakpoint(tool_call: object, message: object) -> bool:
if _attribute_or_key(tool_call, "cache_control") is None:
return False
return not tool_call_is_rebuilt_as_server_tool_use(
_attribute_or_key(tool_call, "id"), _attribute_or_key(message, "provider_specific_fields")
)
def _tool_carries_cache_breakpoint(tool: object) -> bool:
@ -471,13 +495,16 @@ class AnthropicCacheControlHook(CustomPromptManagement):
@staticmethod
def _count_cache_control_blocks(message: object) -> int:
if not isinstance(message, dict):
return 0
count = 1 if _carries_cache_breakpoint(message) else 0
content: Final = message.get("content")
if isinstance(content, list):
count += sum(1 for block in content if _carries_cache_breakpoint(block))
return count
message_count: Final = 1 if _carries_cache_breakpoint(message) else 0
content: Final = _as_object_list(_attribute_or_key(message, "content"))
content_count: Final = sum(1 for block in content if _carries_cache_breakpoint(block)) if content else 0
tool_calls: Final = _as_object_list(_attribute_or_key(message, "tool_calls"))
tool_call_count: Final = (
sum(1 for tool_call in tool_calls if _tool_call_carries_cache_breakpoint(tool_call, message))
if tool_calls
else 0
)
return message_count + content_count + tool_call_count
@staticmethod
def _message_has_cache_control(message: AllMessageValues) -> bool:

View file

@ -1729,11 +1729,16 @@ def convert_function_to_anthropic_tool_invoke(
raise e
def _find_server_tool_result(
ANTHROPIC_SERVER_TOOL_USE_ID_PREFIX: Final = "srvtoolu_"
def find_anthropic_server_tool_result(
tool_id: str,
web_search_results: Sequence[object] | None,
tool_results: Sequence[object] | None,
) -> dict[str, object] | None:
if not tool_id.startswith(ANTHROPIC_SERVER_TOOL_USE_ID_PREFIX):
return None
candidates: Final = (*(web_search_results or ()), *(tool_results or ()))
return next(
(result for result in candidates if isinstance(result, dict) and result.get("tool_use_id") == tool_id),
@ -1808,11 +1813,7 @@ def convert_to_anthropic_tool_invoke(
context="Anthropic tool invoke",
)
server_tool_result = (
_find_server_tool_result(tool_id, web_search_results, tool_results)
if tool_id.startswith("srvtoolu_")
else None
)
server_tool_result = find_anthropic_server_tool_result(tool_id, web_search_results, tool_results)
if server_tool_result is not None:
anthropic_tool_invoke.append(
{

View file

@ -27,6 +27,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.litellm_core_utils.prompt_templates.factory import (
THOUGHT_SIGNATURE_SEPARATOR,
find_anthropic_server_tool_result,
)
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import message_field, parts_of
from litellm.llms.anthropic.wif import (
@ -1834,6 +1835,20 @@ def _replayed_server_tool_use(block: object) -> _ReplayedServerToolUse | None:
return None
def tool_call_is_rebuilt_as_server_tool_use(tool_call_id: object, provider_specific_fields: object) -> bool:
fields: Final = _validated_claude_code_mapping(provider_specific_fields)
if not isinstance(tool_call_id, str) or fields is None:
return False
return (
find_anthropic_server_tool_result(
tool_call_id,
_validated_claude_code_list(fields.get("web_search_results")),
_validated_claude_code_list(fields.get("tool_results")),
)
is not None
)
def _render_web_search_results(
query: str, results: tuple[_ReplayedWebSearchResult, ...] | _ReplayedWebSearchToolResultError
) -> str:

View file

@ -0,0 +1,201 @@
from __future__ import annotations
from typing import Final, Literal, TypeAlias
import pytest
from e2e_config import unique_marker
from e2e_http import Result, unwrap
from lifecycle import ResourceManager
from models import (
CacheControl,
CacheControlInjectionPoint,
ChatAssistantTurn,
ChatBody,
ChatMessage,
ChatResponse,
ChatTool,
ChatToolFunction,
ChatToolResultTurn,
LiteLLMParamsBody,
TextContentPart,
ToolCall,
ToolCallFunction,
)
from passthrough_client import PassthroughClient
pytestmark: Final = pytest.mark.e2e
Backend: TypeAlias = Literal["azure_foundry", "vertex"]
AZURE_MODEL: Final[str] = "azure_ai/claude-haiku-4-5"
VERTEX_MODEL: Final[str] = "vertex_ai/claude-sonnet-5"
VERTEX_LOCATION: Final[str] = "global"
def _deployment_params(*, backend: Backend, inject_cache_control: bool) -> LiteLLMParamsBody:
cache_control_injection_points: Final = (
[
CacheControlInjectionPoint(location="message", role="system"),
CacheControlInjectionPoint(location="message", index=-1),
]
if inject_cache_control
else None
)
match backend:
case "azure_foundry":
return LiteLLMParamsBody(
model=AZURE_MODEL,
api_base="os.environ/AZURE_AI_API_BASE",
api_key="os.environ/AZURE_AI_API_KEY",
cache_control_injection_points=cache_control_injection_points,
)
case "vertex":
return LiteLLMParamsBody(
model=VERTEX_MODEL,
vertex_project="os.environ/VERTEXAI_PROJECT",
vertex_location=VERTEX_LOCATION,
vertex_credentials="os.environ/VERTEXAI_CREDENTIALS",
cache_control_injection_points=cache_control_injection_points,
)
def _register_deployment(
client: PassthroughClient,
resources: ResourceManager,
*,
backend: Backend,
marker: str,
inject_cache_control: bool,
) -> str:
model_name: Final[str] = f"e2e-cache-control-tool-calls-{backend}-{marker}"
model_id: Final[str] = client.proxy.create_model(
model_name,
_deployment_params(backend=backend, inject_cache_control=inject_cache_control),
provider_live=True,
)
resources.defer(lambda: client.proxy.delete_model(model_id))
return model_name
def _request(model: str, marker: str) -> ChatBody:
return ChatBody(
model=model,
messages=[
ChatMessage(role="system", content="Use the provided tool results to answer the user."),
ChatMessage(
role="user",
content=[
TextContentPart(
text="Look up the weather in London, Paris, and Tokyo.",
cache_control=CacheControl(),
)
],
),
ChatAssistantTurn(
content="",
tool_calls=[
ToolCall(
id="call_weather_london",
type="function",
function=ToolCallFunction(name="lookup_weather", arguments='{"city":"London"}'),
cache_control=CacheControl(),
),
ToolCall(
id="call_weather_paris",
type="function",
function=ToolCallFunction(name="lookup_weather", arguments='{"city":"Paris"}'),
cache_control=CacheControl(),
),
ToolCall(
id="call_weather_tokyo",
type="function",
function=ToolCallFunction(name="lookup_weather", arguments='{"city":"Tokyo"}'),
cache_control=CacheControl(),
),
],
),
ChatToolResultTurn(tool_call_id="call_weather_london", content="London is sunny."),
ChatToolResultTurn(tool_call_id="call_weather_paris", content="Paris is cloudy."),
ChatToolResultTurn(tool_call_id="call_weather_tokyo", content="Tokyo is rainy."),
ChatMessage(
role="user",
content=f"Summarize the results in one word and do not call another tool. {marker}",
),
],
tools=[
ChatTool(
function=ChatToolFunction(
name="lookup_weather",
description="Look up the weather in a city.",
parameters={
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
)
)
],
max_tokens=64,
)
def _post_chat(client: PassthroughClient, key: str, body: ChatBody) -> Result[ChatResponse]:
return client.proxy.transport.post(
"/v1/chat/completions",
headers=client.proxy.transport.bearer(key),
json=body,
response_type=ChatResponse,
)
def _assert_normal_completion(response: ChatResponse, model_name: str) -> None:
assert response.choices, f"{model_name}: chat completion returned no choices: {response}"
completion: Final = response.choices[0]
assert completion.finish_reason in ("stop", "length"), (
f"{model_name}: unexpected finish reason: {completion.finish_reason}"
)
assert (
completion.message is not None
and completion.message.content is not None
and completion.message.content.strip()
), f"{model_name}: chat completion returned no text: {completion.message}"
@pytest.mark.parametrize(
"backend",
(pytest.param("azure_foundry", id="azure-foundry"), pytest.param("vertex", id="vertex")),
)
@pytest.mark.provider_live
@pytest.mark.covers("llm.chat_completions.azure_foundry.basic.nonstream.works")
@pytest.mark.covers("llm.chat_completions.vertex.basic.nonstream.works")
class TestCacheControlInjectionToolCalls:
def test_injection_points_respect_cap_with_tool_call_marks(
self, client: PassthroughClient, resources: ResourceManager, backend: Backend
) -> None:
marker: Final[str] = unique_marker()
model_name: Final[str] = _register_deployment(
client,
resources,
backend=backend,
marker=marker,
inject_cache_control=True,
)
key: Final[str] = resources.key()
response: Final[ChatResponse] = unwrap(_post_chat(client, key, _request(model_name, marker)))
_assert_normal_completion(response, model_name)
def test_client_tool_call_marks_work_without_injection_points(
self, client: PassthroughClient, resources: ResourceManager, backend: Backend
) -> None:
marker: Final[str] = unique_marker()
model_name: Final[str] = _register_deployment(
client,
resources,
backend=backend,
marker=marker,
inject_cache_control=False,
)
key: Final[str] = resources.key()
response: Final[ChatResponse] = unwrap(_post_chat(client, key, _request(model_name, marker)))
_assert_normal_completion(response, model_name)

View file

@ -294,6 +294,7 @@ class ToolCall(BaseModel):
id: str | None = None
type: str | None = None
function: ToolCallFunction = ToolCallFunction()
cache_control: CacheControl | None = None
class ThinkingBlock(BaseModel):
@ -1233,6 +1234,12 @@ class FineTuningJobsResponse(BaseModel):
# ---------- model management ----------
class CacheControlInjectionPoint(BaseModel):
location: Literal["message"]
role: str | None = None
index: int | None = None
class LiteLLMParamsBody(BaseModel):
"""POST /model/new litellm_params: `model` is the only required field; `api_key`
et al may be an `os.environ/FOO` reference the proxy resolves at call time.
@ -1291,6 +1298,7 @@ class LiteLLMParamsBody(BaseModel):
max_retries: int | None = None
cooldown_time: float | None = None
extra_body: DeploymentExtraBody | None = None
cache_control_injection_points: list[CacheControlInjectionPoint] | None = None
tpm: int | None = None
weight: int | None = None
order: int | None = None

View file

@ -0,0 +1,461 @@
import json
import re
import uuid
from collections.abc import Iterator, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
import yaml
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.wire import Reply, Request
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
# https://platform.claude.com/docs/en/build-with-claude/prompt-caching (read 2026-09-28): at most 4 blocks with cache_control
ANTHROPIC_CACHE_CONTROL_CAP: Final = 4
# https://docs.aws.amazon.com/bedrock/latest/userguide/prompt-caching.html (read 2026-09-28): at most 4 cache checkpoints
BEDROCK_CACHE_CHECKPOINT_CAP: Final = 4
ANTHROPIC_MODEL: Final = "claude-opus-5-5"
BEDROCK_MODEL: Final = "anthropic.claude-opus-5-5"
PROVIDER_KEY: Final = "synthetic-provider-key"
SYSTEM: Final = "Use the provided tool results to answer the user."
ASK: Final = "Look up the weather in London, Paris, and Tokyo."
CITIES: Final = ("London", "Paris", "Tokyo")
EPHEMERAL: Final[dict[str, JsonValue]] = {"type": "ephemeral"}
POINTS: Final[list[JsonValue]] = [
{"location": "message", "role": "system"},
{"location": "message", "index": -1},
]
TOOL: Final[dict[str, JsonValue]] = {
"type": "function",
"function": {
"name": "lookup_weather",
"description": "Look up the weather in a city.",
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
},
}
SYSTEM_LABEL: Final = f"system:{SYSTEM}"
ASK_LABEL: Final = f"user:text:{ASK}"
_MARKER: Final = re.compile(rb"marker-([0-9a-f]{32})")
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
@dataclass(frozen=True, slots=True)
class Mark:
label: str
ttl: str | None
def new_marker() -> str:
return uuid.uuid4().hex
def final_text(marker: str) -> str:
return f"Summarize the results in one word. marker-{marker}"
def final_label(marker: str) -> str:
return f"user:text:{final_text(marker)}"
def call_id(city: str) -> str:
return f"call_weather_{city.lower()}"
def tool_use_label(city: str) -> str:
return f"assistant:tool_use:{call_id(city)}"
def tool_call(city: str, **fields: JsonValue) -> dict[str, JsonValue]:
return {
"id": call_id(city),
"type": "function",
"function": {"name": "lookup_weather", "arguments": json.dumps({"city": city})},
**fields,
}
def marked_calls(mark: JsonValue = EPHEMERAL, cities: Sequence[str] = CITIES) -> list[JsonValue]:
return [tool_call(city, cache_control=mark) for city in cities]
def client_marked() -> list[str]:
return [ASK_LABEL, *(tool_use_label(city) for city in CITIES)]
def ask(*, marked: bool) -> dict[str, JsonValue]:
block: Final[dict[str, JsonValue]] = {"type": "text", "text": ASK}
return {"role": "user", "content": [{**block, "cache_control": EPHEMERAL} if marked else block]}
def tool_results(calls: Sequence[JsonValue]) -> list[JsonValue]:
return [
{"role": "tool", "tool_call_id": call["id"], "content": "sunny."}
for call in calls
if isinstance(call, dict) and str(call.get("id", "")).startswith("call_")
]
def conversation(
marker: str,
calls: Sequence[JsonValue],
*,
ask_marked: bool = True,
assistant: dict[str, JsonValue] | None = None,
) -> list[JsonValue]:
turn: Final[dict[str, JsonValue]] = {
"role": "assistant",
"content": "",
"tool_calls": list(calls),
**(assistant or {}),
}
return [
{"role": "system", "content": SYSTEM},
ask(marked=ask_marked),
turn,
*tool_results(calls),
{"role": "user", "content": final_text(marker)},
]
def chat_body(model: str, messages: Sequence[JsonValue], **fields: JsonValue) -> dict[str, JsonValue]:
return {"model": model, "messages": list(messages), "tools": [TOOL], "max_tokens": 64, **fields}
def messages_body(model: str, marker: str, *, stream: bool) -> dict[str, JsonValue]:
return {
"model": model,
"max_tokens": 64,
"stream": stream,
"system": [{"type": "text", "text": SYSTEM}],
"tools": [
{
"name": "lookup_weather",
"description": "Look up the weather in a city.",
"input_schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
}
],
"messages": [
{"role": "user", "content": [{"type": "text", "text": ASK, "cache_control": EPHEMERAL}]},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": call_id(city),
"name": "lookup_weather",
"input": {"city": city},
"cache_control": EPHEMERAL,
}
for city in CITIES
],
},
{
"role": "user",
"content": [
*({"type": "tool_result", "tool_use_id": call_id(city), "content": "sunny."} for city in CITIES),
{"type": "text", "text": final_text(marker)},
],
},
],
}
def responses_body(model: str, marker: str) -> dict[str, JsonValue]:
items: Final[list[JsonValue]] = [
{"role": "user", "content": [{"type": "input_text", "text": ASK, "cache_control": EPHEMERAL}]},
*(
{
"type": "function_call",
"call_id": call_id(city),
"name": "lookup_weather",
"arguments": json.dumps({"city": city}),
"cache_control": EPHEMERAL,
}
for city in CITIES
),
*({"type": "function_call_output", "call_id": call_id(city), "output": "sunny."} for city in CITIES),
{"role": "user", "content": final_text(marker)},
]
return {
"model": model,
"instructions": SYSTEM,
"max_output_tokens": 64,
"tools": [{"type": "function", "name": "lookup_weather", "parameters": {"type": "object"}}],
"input": items,
}
def marker_of(request: Request) -> str:
found: Final = _MARKER.findall(request.body)
return found[-1].decode() if found else "unmarked"
class _AnthropicBlock(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
type: str
text: str | None = None
id: str | None = None
tool_use_id: str | None = None
name: str | None = None
cache_control: JsonValue = None
class _AnthropicMessage(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
role: str
content: str | tuple[_AnthropicBlock, ...]
class _AnthropicTool(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
name: str = ""
cache_control: JsonValue = None
class _AnthropicBody(BaseModel):
model_config = ConfigDict(extra="allow", frozen=True)
system: str | tuple[_AnthropicBlock, ...] = ()
messages: tuple[_AnthropicMessage, ...] = ()
tools: tuple[_AnthropicTool, ...] = ()
cache_control: JsonValue = None
stream: bool = False
def _ttl(cache_control: JsonValue) -> str | None:
if not isinstance(cache_control, dict):
return None
ttl: Final = cache_control.get("ttl")
return ttl if isinstance(ttl, str) else None
def _block_label(role: str, block: _AnthropicBlock) -> str:
detail: Final = block.text if block.type == "text" else block.id or block.tool_use_id or ""
return f"{role}:{block.type}:{detail}"
def _message_blocks(body: _AnthropicBody) -> Iterator[tuple[str, _AnthropicBlock]]:
for message in body.messages:
if isinstance(message.content, tuple):
yield from ((message.role, block) for block in message.content)
def _anthropic_body(request: Request) -> _AnthropicBody:
return _AnthropicBody.model_validate_json(request.body)
def anthropic_marks(request: Request) -> tuple[Mark, ...]:
body: Final = _anthropic_body(request)
system: Final = body.system if isinstance(body.system, tuple) else ()
return (
*(Mark(f"tool:{tool.name}", _ttl(tool.cache_control)) for tool in body.tools if tool.cache_control is not None),
*(
Mark(f"system:{block.text}", _ttl(block.cache_control))
for block in system
if block.cache_control is not None
),
*(
Mark(_block_label(role, block), _ttl(block.cache_control))
for role, block in _message_blocks(body)
if block.cache_control is not None
),
*((Mark("request", _ttl(body.cache_control)),) if body.cache_control is not None else ()),
)
def anthropic_labels(request: Request) -> list[str]:
return [mark.label for mark in anthropic_marks(request)]
def _anthropic_error(message: str) -> Reply:
return Reply(
status=400,
body=json.dumps({"type": "error", "error": {"type": "invalid_request_error", "message": message}}).encode(),
)
def _ttl_out_of_order(marks: Sequence[Mark]) -> bool:
first_short: Final = next((index for index, mark in enumerate(marks) if mark.ttl != "1h"), len(marks))
return any(mark.ttl == "1h" for mark in marks[first_short:])
def _usage() -> dict[str, JsonValue]:
return {"input_tokens": 12, "output_tokens": 1, "cache_creation_input_tokens": 12, "cache_read_input_tokens": 0}
def _anthropic_message(identity: str) -> dict[str, JsonValue]:
return {
"id": identity,
"type": "message",
"role": "assistant",
"model": ANTHROPIC_MODEL,
"content": [{"type": "text", "text": "sunny"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": _usage(),
}
def _event(name: str, data: dict[str, JsonValue]) -> bytes:
return f"event: {name}\ndata: {json.dumps(data)}\n\n".encode()
def _anthropic_stream(identity: str) -> tuple[bytes, ...]:
return (
_event(
"message_start",
{"type": "message_start", "message": {**_anthropic_message(identity), "content": [], "stop_reason": None}},
),
_event(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
),
_event(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "sunny"}},
),
_event("content_block_stop", {"type": "content_block_stop", "index": 0}),
_event(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 1},
},
),
_event("message_stop", {"type": "message_stop"}),
)
def anthropic_peer(request: Request) -> Reply:
marks: Final = anthropic_marks(request)
if len(marks) > ANTHROPIC_CACHE_CONTROL_CAP:
return _anthropic_error(
f"A maximum of {ANTHROPIC_CACHE_CONTROL_CAP} blocks with cache_control may be provided. Found {len(marks)}."
)
if _ttl_out_of_order(marks):
return _anthropic_error("a ttl='1h' cache_control block must not come after a ttl='5m' cache_control block")
identity: Final = f"msg_{marker_of(request)}"
if _anthropic_body(request).stream:
return Reply(content_type="text/event-stream", chunks=_anthropic_stream(identity))
return Reply(body=json.dumps(_anthropic_message(identity)).encode())
def _bedrock_label(role: str, block: JsonValue) -> str:
if not isinstance(block, dict):
return f"{role}:start"
if isinstance(block.get("text"), str):
return f"{role}:text:{block['text']}"
for kind in ("toolUse", "toolResult"):
inner = block.get(kind)
if isinstance(inner, dict):
return f"{role}:{kind}:{inner.get('toolUseId')}"
spec: Final = block.get("toolSpec")
return f"tool:{spec.get('name')}" if isinstance(spec, dict) else f"{role}:other"
def _cache_points(role: str, blocks: JsonValue) -> Iterator[str]:
listed: Final = blocks if isinstance(blocks, list) else []
for previous, block in zip([None, *listed], listed):
if isinstance(block, dict) and "cachePoint" in block:
yield _bedrock_label(role, previous)
def _bedrock_sections(body: dict[str, JsonValue]) -> Iterator[tuple[str, JsonValue]]:
tool_config: Final = body.get("toolConfig")
yield ("tool", tool_config.get("tools") if isinstance(tool_config, dict) else None)
yield ("system", body.get("system"))
messages: Final = body.get("messages")
for message in messages if isinstance(messages, list) else []:
if isinstance(message, dict):
yield (str(message.get("role")), message.get("content"))
def bedrock_labels(request: Request) -> list[str]:
body: Final = _JSON_OBJECT.validate_json(request.body)
return [label for role, blocks in _bedrock_sections(body) for label in _cache_points(role, blocks)]
def bedrock_peer(request: Request) -> Reply:
found: Final = len(bedrock_labels(request))
if found > BEDROCK_CACHE_CHECKPOINT_CAP:
return Reply(
status=400,
headers={"x-amzn-errortype": "ValidationException"},
body=json.dumps(
{
"message": f"A maximum of {BEDROCK_CACHE_CHECKPOINT_CAP} cache checkpoints may be provided. Found {found}."
}
).encode(),
)
return Reply(
body=json.dumps(
{
"output": {"message": {"role": "assistant", "content": [{"text": "sunny"}]}},
"stopReason": "end_turn",
"usage": {
"inputTokens": 12,
"outputTokens": 1,
"totalTokens": 13,
"cacheWriteInputTokens": 12,
"cacheReadInputTokens": 0,
},
"metrics": {"latencyMs": 1},
}
).encode()
)
def gateway_injected(response_id: str) -> bool:
rows: Final = eventually(
lambda: read_rows(
"""SELECT metadata->>'litellm_gateway_injected_cache' AS injected FROM "LiteLLM_SpendLogs" """
"WHERE request_id=%s",
(response_id,),
),
lambda values: len(values) == 1,
seconds=70,
)
return rows[0]["injected"] is not None
def post_chat(gateway: Gateway, body: dict[str, JsonValue], *, key: str | None = None) -> tuple[int, str, str]:
response: Final = gateway.request("POST", "/v1/chat/completions", body, key=key)
identity: Final = _JSON_OBJECT.validate_json(response.content).get("id") if response.status_code == 200 else None
return response.status_code, str(identity), response.text
def anthropic_deployment(name: str, api_base: str, **fields: JsonValue) -> dict[str, JsonValue]:
return {
"model_name": name,
"litellm_params": {
"model": f"anthropic/{ANTHROPIC_MODEL}",
"api_base": api_base,
"api_key": PROVIDER_KEY,
**fields,
},
}
def owned_config(
directory: Path,
model_list: Sequence[JsonValue],
*,
litellm_settings: Mapping[str, JsonValue] = MappingProxyType({}),
router_settings: Mapping[str, JsonValue] = MappingProxyType({}),
) -> Path:
config: Final = _JSON_OBJECT.validate_python(
yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
)
merged: Final = {
**config,
"model_list": list(model_list),
"litellm_settings": {**object_value(config["litellm_settings"]), **litellm_settings},
"router_settings": {**object_value(config["router_settings"]), "num_retries": 0, **router_settings},
}
path: Final = directory / f"cache-control-marks-{uuid.uuid4().hex}.yaml"
path.write_text(yaml.safe_dump(merged))
return path

View file

@ -0,0 +1,294 @@
import asyncio
import os
import re
import signal
import socket
import threading
from collections import Counter
from collections.abc import Callable, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
from urllib.parse import urlsplit
import httpx
import psutil
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from integration.providers._cache_control_marks_support import (
ASK_LABEL,
CITIES,
POINTS,
SYSTEM_LABEL,
anthropic_deployment,
anthropic_labels,
anthropic_peer,
chat_body,
client_marked,
conversation,
final_label,
marked_calls,
marker_of,
messages_body,
new_marker,
owned_config,
responses_body,
tool_call,
tool_use_label,
)
from pydantic import JsonValue
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
_SURFACES: Final = ("chat", "chat-stream", "chat-unmarked", "messages", "responses")
_MODEL: Final = "capped-claude"
_BURST: Final = 30
_OUTAGE_STATUS: Final = 500
@dataclass(frozen=True, slots=True)
class _Sent:
surface: str
marker: str
status: int
text: str
call_id: str
def _free_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return int(reserve.getsockname()[1])
def _surface_request(surface: str, marker: str) -> tuple[str, dict[str, JsonValue]]:
if surface == "messages":
return "/v1/messages", messages_body(_MODEL, marker, stream=False)
if surface == "responses":
return "/v1/responses", responses_body(_MODEL, marker)
if surface == "chat-unmarked":
unmarked: Final = [tool_call(city) for city in CITIES]
return "/v1/chat/completions", chat_body(_MODEL, conversation(marker, unmarked, ask_marked=False))
return "/v1/chat/completions", chat_body(
_MODEL, conversation(marker, marked_calls()), stream=surface == "chat-stream"
)
def _expected_labels(item: _Sent) -> list[str]:
if item.surface == "responses":
return [SYSTEM_LABEL, ASK_LABEL]
if item.surface == "chat-unmarked":
return [SYSTEM_LABEL, final_label(item.marker)]
return client_marked()
async def _fire(owned_url: str, key: str, *, tolerate_transport_errors: bool = False) -> tuple[_Sent, ...]:
async def one(client: httpx.AsyncClient, index: int) -> _Sent:
surface: Final = _SURFACES[index % len(_SURFACES)]
marker: Final = new_marker()
path, body = _surface_request(surface, marker)
response: Final = await client.post(path, json=body, headers={"Authorization": f"Bearer {key}"})
return _Sent(
surface, marker, response.status_code, response.text, response.headers.get("x-litellm-call-id", "")
)
async with httpx.AsyncClient(base_url=owned_url, timeout=30, trust_env=False) as client:
results: Final = await asyncio.gather(
*(one(client, index) for index in range(_BURST)), return_exceptions=tolerate_transport_errors
)
for result in results:
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, _Sent))
def _held_anthropic_peer(release: threading.Event) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert release.wait(timeout=120), "Held upstream was never released"
return anthropic_peer(request)
return respond
def _held_upstream_connections(pid: int, upstream: str) -> int:
port: Final = urlsplit(upstream).port
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
def _by_marker(received: Sequence[Request]) -> dict[str, tuple[Request, ...]]:
counted: Final = Counter(marker_of(request) for request in received)
return {marker: tuple(request for request in received if marker_of(request) == marker) for marker in counted}
def _assert_capped(served: Sequence[_Sent], received: Sequence[Request]) -> None:
by_marker: Final = _by_marker(received)
for item in served:
assert item.status == 200, (item.surface, item.text)
assert len(by_marker.get(item.marker, ())) == 1, (item.surface, item.marker)
assert anthropic_labels(by_marker[item.marker][0]) == _expected_labels(item), item.surface
def _assert_outage_error(item: _Sent, received: Sequence[Request]) -> None:
assert item.status == _OUTAGE_STATUS, (item.surface, item.status, item.text)
assert '"error"' in item.text, (item.surface, item.text)
assert item.marker not in {marker_of(request) for request in received}, item.surface
def _spend_rows(call_id: str) -> list[dict[str, JsonValue]]:
return read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,))
def _single_spend_row(item: _Sent) -> None:
assert item.call_id, (item.surface, item.status, item.text)
rows: Final = eventually(lambda: _spend_rows(item.call_id), lambda values: len(values) == 1, seconds=70)
assert len(rows) == 1, (item.surface, item.call_id)
@pytest.mark.timeout(240)
async def test_capped_burst_rides_out_a_provider_outage_and_logs_every_request_once(
gateway: Gateway, tmp_path: Path
) -> None:
port: Final = _free_port()
config: Final = owned_config(
tmp_path, [anthropic_deployment(_MODEL, f"http://127.0.0.1:{port}", cache_control_injection_points=POINTS)]
)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
owned_url: Final = str(owned.gateway.client.base_url)
with wire_server(anthropic_peer, port=port) as wire:
healthy: Final = await _fire(owned_url, owned.gateway.key)
_assert_capped(healthy, wire.drain())
racing: Final = asyncio.create_task(_fire(owned_url, owned.gateway.key))
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30)
down: Final = await _fire(owned_url, owned.gateway.key)
raced: Final = await racing
raced_received: Final = wire.drain()
with wire_server(anthropic_peer, port=port) as restarted:
recovered: Final = await _fire(owned_url, owned.gateway.key)
recovered_received: Final = restarted.drain()
assert len(_STARTED_WORKER.findall(owned.log.read_text())) >= 2
assert all(len(requests) == 1 for requests in _by_marker(raced_received).values())
_assert_capped(tuple(item for item in raced if item.status == 200), raced_received)
for item in raced:
if item.status != 200:
_assert_outage_error(item, raced_received)
for item in down:
_assert_outage_error(item, (*raced_received, *recovered_received))
_assert_capped(recovered, recovered_received)
assert {marker_of(request) for request in recovered_received} == {item.marker for item in recovered}
for item in (*healthy, *raced, *down, *recovered):
_single_spend_row(item)
@pytest.mark.timeout(240)
async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_capped_requests(
gateway: Gateway, tmp_path: Path
) -> None:
release: Final = threading.Event()
with wire_server(_held_anthropic_peer(release)) as wire:
config: Final = owned_config(
tmp_path, [anthropic_deployment(_MODEL, wire.url, cache_control_injection_points=POINTS)]
)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
workers: Final[tuple[int, ...]] = eventually(
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
lambda pids: len(pids) == 2,
seconds=30,
)
owned_url: Final = str(owned.gateway.client.base_url)
burst: Final = asyncio.create_task(_fire(owned_url, owned.gateway.key, tolerate_transport_errors=True))
try:
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= _BURST, 90)
held_by: Final = MappingProxyType({pid: _held_upstream_connections(pid, wire.url) for pid in workers})
victim: Final = max(workers, key=held_by.__getitem__)
os.kill(victim, signal.SIGKILL)
finally:
release.set()
served: Final = await burst
during: Final = wire.drain()
after: Final = await _fire(owned_url, owned.gateway.key)
after_received: Final = wire.drain()
assert sum(held_by.values()) == _BURST, held_by
assert held_by[victim] > 0, held_by
assert len(served) == _BURST - held_by[victim], (len(served), held_by)
assert len(after) == _BURST, len(after)
assert all(len(requests) == 1 for requests in _by_marker(during).values())
_assert_capped(served, during)
_assert_capped(after, after_received)
for item in (*served, *after):
_single_spend_row(item)
@pytest.mark.timeout(240)
def test_yaml_auto_caching_stands_down_for_tool_call_marks_and_outranks_a_key_opt_out(
gateway: Gateway, tmp_path: Path
) -> None:
markers: Final = (new_marker(), new_marker(), new_marker())
unmarked: Final = [tool_call(city) for city in CITIES]
with wire_server(anthropic_peer) as wire:
config: Final = owned_config(
tmp_path,
[anthropic_deployment(_MODEL, wire.url)],
litellm_settings={"enable_anthropic_prompt_caching": True},
)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
with owned.gateway.scenario() as scenario:
opted_out: Final = scenario.key(metadata={"enable_prompt_caching": False})
cells: Final = (
(conversation(markers[0], marked_calls(), ask_marked=False), owned.gateway.key),
(conversation(markers[1], unmarked, ask_marked=False), owned.gateway.key),
(conversation(markers[2], unmarked, ask_marked=False), opted_out),
)
responses: Final = tuple(
owned.gateway.request("POST", "/v1/chat/completions", chat_body(_MODEL, messages), key=key)
for messages, key in cells
)
received: Final = _by_marker(wire.drain())
assert [response.status_code for response in responses] == [200, 200, 200], [
response.text for response in responses
]
assert [anthropic_labels(received[marker][0]) for marker in markers] == [
[tool_use_label(city) for city in CITIES],
[SYSTEM_LABEL, final_label(markers[1])],
[SYSTEM_LABEL, final_label(markers[2])],
]
@pytest.mark.timeout(360)
async def test_proxy_restart_mid_burst_serves_capped_requests_after_the_reboot(
gateway: Gateway, tmp_path: Path
) -> None:
release: Final = threading.Event()
with wire_server(_held_anthropic_peer(release)) as wire:
config: Final = owned_config(
tmp_path, [anthropic_deployment(_MODEL, wire.url, cache_control_injection_points=POINTS)]
)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as first:
burst: Final = asyncio.create_task(
_fire(str(first.gateway.client.base_url), first.gateway.key, tolerate_transport_errors=True)
)
try:
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= _BURST, 90)
first.process.terminate()
finally:
release.set()
served: Final = await burst
during: Final = wire.drain()
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as second:
after: Final = await _fire(str(second.gateway.client.base_url), second.gateway.key)
after_received: Final = wire.drain()
completed: Final = tuple(item for item in served if item.status == 200)
lost: Final = tuple(item for item in served if item.status != 200)
assert len(after) == _BURST, len(after)
assert all(len(requests) == 1 for requests in _by_marker(during).values())
_assert_capped(completed, during)
for item in lost:
assert '"error"' in item.text or item.text == "", (item.surface, item.status, item.text)
_assert_capped(after, after_received)
for item in (*completed, *after):
_single_spend_row(item)

View file

@ -0,0 +1,934 @@
import asyncio
import json
from collections.abc import Iterable
from datetime import datetime, timedelta, timezone
from typing import Final, cast
import anthropic
import httpx
import openai
import pytest
from anthropic.types import MessageParam, TextBlockParam, ToolParam
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.providers._cache_control_marks_support import (
ANTHROPIC_MODEL,
ASK,
ASK_LABEL,
BEDROCK_MODEL,
CITIES,
EPHEMERAL,
POINTS,
PROVIDER_KEY,
SYSTEM,
SYSTEM_LABEL,
TOOL,
anthropic_labels,
anthropic_marks,
anthropic_peer,
bedrock_labels,
bedrock_peer,
call_id,
chat_body,
client_marked,
conversation,
final_label,
final_text,
gateway_injected,
marked_calls,
messages_body,
new_marker,
post_chat,
responses_body,
tool_call,
tool_use_label,
)
from openai.types.chat import ChatCompletionMessageParam, ChatCompletionToolParam
from openai.types.responses import ResponseInputParam
from openai.types.responses import ToolParam as ResponsesToolParam
from pydantic import JsonValue, TypeAdapter
from litellm.utils import get_prompt_cache_min_tokens
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_CLIENT_MARKED: Final = client_marked()
_GEMINI_REPLY: Final = json.dumps(
{
"candidates": [{"content": {"role": "model", "parts": [{"text": "sunny"}]}, "finishReason": "STOP"}],
"usageMetadata": {"promptTokenCount": 1600, "candidatesTokenCount": 1, "totalTokenCount": 1601},
}
).encode()
_OPENAI_REPLY: Final = json.dumps(
{
"id": "chatcmpl-cache-census",
"object": "chat.completion",
"created": 1,
"model": "gpt-6-sol",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "sunny"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 12, "completion_tokens": 1, "total_tokens": 13},
}
).encode()
def _anthropic_deployment(scenario: Scenario, wire: Wire, **fields: JsonValue) -> str:
return scenario.model(model=f"anthropic/{ANTHROPIC_MODEL}", api_base=wire.url, api_key=PROVIDER_KEY, **fields)
def _only_request(wire: Wire) -> Request:
received: Final = wire.drain()
assert len(received) == 1, [request.target for request in received]
return received[0]
def _stream(gateway: Gateway, path: str, body: dict[str, JsonValue], *, key: str | None = None) -> tuple[int, str]:
headers: Final = {"Authorization": f"Bearer {gateway.key if key is None else key}"}
with gateway.client.stream("POST", path, json=body, headers=headers) as response:
return response.status_code, response.read().decode()
def _chat_stream_chunks(text: str) -> list[dict[str, JsonValue]]:
return [
_JSON_OBJECT.validate_json(line.removeprefix("data: "))
for line in text.splitlines()
if line.startswith("data: ") and line != "data: [DONE]"
]
def _chat_stream_content(chunks: list[dict[str, JsonValue]]) -> str:
return "".join(
str(object_value(object_value(choice).get("delta") or {}).get("content") or "")
for chunk in chunks
for choice in (chunk.get("choices") if isinstance(chunk.get("choices"), list) else [])
)
@pytest.mark.parametrize("stream", (False, True), ids=("json", "stream"))
def test_chat_points_skip_injection_when_client_marks_fill_the_cap(gateway: Gateway, stream: bool) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
body: Final = chat_body(model, conversation(marker, marked_calls()), stream=stream)
if stream:
status, text = _stream(gateway, "/v1/chat/completions", body)
assert status == 200, text
chunks: Final = _chat_stream_chunks(text)
assert _chat_stream_content(chunks) == "sunny", text
assert "data: [DONE]" in text, text
response_id = str(chunks[0]["id"])
else:
status, response_id, text = post_chat(gateway, body)
assert status == 200, text
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
assert gateway_injected(response_id) is False
def test_chat_points_inject_system_and_last_message_without_client_marks(gateway: Gateway) -> None:
marker: Final = new_marker()
unmarked: Final = [tool_call(city) for city in CITIES]
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, response_id, text = post_chat(
gateway, chat_body(model, conversation(marker, unmarked, ask_marked=False))
)
assert status == 200, text
assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, final_label(marker)]
assert gateway_injected(response_id) is True
def _prompt_caching_rows(gateway: Gateway, cursor: dict[str, str]) -> list[dict[str, JsonValue]]:
now: Final = datetime.now(timezone.utc)
page: Final = gateway.get(
"/cost_optimization/prompt_caching/requests",
{
"start_date": (now - timedelta(minutes=10)).isoformat(),
"end_date": (now + timedelta(minutes=10)).isoformat(),
"page_size": "100",
"filter": "injected",
**cursor,
},
)
rows: Final = [object_value(row) for row in page["requests"]] if isinstance(page["requests"], list) else []
following: Final = page.get("next_cursor")
if not page.get("has_more") or not isinstance(following, dict):
return rows
return [
*rows,
*_prompt_caching_rows(
gateway,
{"cursor_start_time": str(following["start_time"]), "cursor_request_id": str(following["request_id"])},
),
]
def _listed_as_injected(gateway: Gateway, response_id: str) -> bool:
return any(row["request_id"] == response_id for row in _prompt_caching_rows(gateway, {}))
@pytest.mark.parametrize(
("calls", "ask_marked", "expected", "injected"),
(
pytest.param(marked_calls(), False, [tool_use_label(city) for city in CITIES], False, id="three-tool-calls"),
pytest.param([tool_call(city) for city in CITIES], False, None, True, id="no-client-marks"),
pytest.param(
[*marked_calls(cities=CITIES[:2]), tool_call(CITIES[2])],
False,
[tool_use_label(city) for city in CITIES[:2]],
False,
id="two-tool-calls",
),
),
)
def test_auto_prompt_caching_stands_down_when_client_marks_only_tool_calls(
gateway: Gateway,
calls: list[JsonValue],
ask_marked: bool,
expected: list[str] | None,
injected: bool,
) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire)
key: Final = scenario.key(metadata={"enable_prompt_caching": True})
status, response_id, text = post_chat(
gateway, chat_body(model, conversation(marker, calls, ask_marked=ask_marked)), key=key
)
assert status == 200, text
labels: Final = anthropic_labels(_only_request(wire))
assert labels == (expected if expected is not None else [SYSTEM_LABEL, final_label(marker)])
assert gateway_injected(response_id) is injected
assert (
eventually(lambda: _listed_as_injected(gateway, response_id), lambda listed: listed is injected, 70) is injected
)
def test_assistant_point_skips_message_whose_tool_call_carries_a_one_hour_mark(gateway: Gateway) -> None:
marker: Final = new_marker()
hour: Final[dict[str, JsonValue]] = {"type": "ephemeral", "ttl": "1h"}
calls: Final = [tool_call(CITIES[0], cache_control=hour)]
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(
scenario, wire, cache_control_injection_points=[{"location": "message", "role": "assistant"}]
)
status, _, text = post_chat(
gateway,
chat_body(model, conversation(marker, calls, ask_marked=False, assistant={"content": "I will check."})),
)
assert status == 200, text
marks: Final = anthropic_marks(_only_request(wire))
assert [(mark.label, mark.ttl) for mark in marks] == [(tool_use_label(CITIES[0]), "1h")]
@pytest.mark.parametrize(
"mark",
(
pytest.param("ephemeral", id="string"),
pytest.param(1, id="int"),
pytest.param(["ephemeral"], id="list"),
pytest.param("", id="empty-string"),
pytest.param("x" * 5120, id="5kb-string"),
),
)
def test_non_object_tool_call_marks_count_against_the_cap(gateway: Gateway, mark: JsonValue) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls(mark))))
assert status == 200, text
assert anthropic_labels(_only_request(wire)) == [ASK_LABEL]
@pytest.mark.parametrize(
("calls", "expected"),
(
pytest.param(marked_calls({}), _CLIENT_MARKED, id="empty-object"),
pytest.param(marked_calls(None), None, id="null"),
pytest.param(
[tool_call(CITIES[0], cache_control=EPHEMERAL)] * 2,
[SYSTEM_LABEL, ASK_LABEL, tool_use_label(CITIES[0])],
id="same-call-twice",
),
),
)
def test_tool_call_mark_shapes_keep_the_request_within_the_cap(
gateway: Gateway, calls: list[JsonValue], expected: list[str] | None
) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls)))
assert status == 200, text
labels: Final = anthropic_labels(_only_request(wire))
assert labels == (expected if expected is not None else [SYSTEM_LABEL, ASK_LABEL, final_label(marker)])
def _untyped_call(city: str) -> dict[str, JsonValue]:
return {name: value for name, value in tool_call(city, cache_control=EPHEMERAL).items() if name != "type"}
_BEDROCK_CLIENT_MARKED: Final = [f"user:text:{ASK}", *(f"assistant:toolUse:{call_id(city)}" for city in CITIES)]
@pytest.mark.parametrize(
("calls", "expected"),
(
pytest.param(marked_calls(), _BEDROCK_CLIENT_MARKED, id="object-marks"),
pytest.param(marked_calls("ephemeral"), _BEDROCK_CLIENT_MARKED, id="string-marks"),
pytest.param([_untyped_call(city) for city in CITIES], _BEDROCK_CLIENT_MARKED, id="calls-without-type"),
pytest.param(
[tool_call(CITIES[0], cache_control=EPHEMERAL)] * 2,
[f"system:text:{SYSTEM}", ASK_LABEL, f"assistant:toolUse:{call_id(CITIES[0])}", "assistant:other"],
id="same-call-twice",
),
),
)
def test_bedrock_converse_cache_points_stay_within_the_cap(
gateway: Gateway, calls: list[JsonValue], expected: list[str]
) -> None:
marker: Final = new_marker()
with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"bedrock/converse/{BEDROCK_MODEL}",
api_key=PROVIDER_KEY,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint=wire.url,
cache_control_injection_points=POINTS,
)
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls)))
assert status == 200, text
request: Final = _only_request(wire)
assert request.target.endswith("/converse"), request.target
assert bedrock_labels(request) == expected
_SERVER_CALL: Final = tool_call(
"web", cache_control=EPHEMERAL, id="srvtoolu_web", function={"name": "web_search", "arguments": "{}"}
)
_WEB_RESULTS: Final[dict[str, JsonValue]] = {
"provider_specific_fields": {
"web_search_results": [{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_web", "content": []}]
}
}
@pytest.mark.parametrize(
("assistant", "expected"),
(
pytest.param(
_WEB_RESULTS,
[SYSTEM_LABEL, ASK_LABEL, tool_use_label(CITIES[0]), tool_use_label(CITIES[1])],
id="server-tool-with-result",
),
pytest.param(
{},
[ASK_LABEL, tool_use_label(CITIES[0]), tool_use_label(CITIES[1]), "assistant:tool_use:srvtoolu_web"],
id="server-tool-without-result",
),
),
)
def test_server_tool_call_mark_counts_only_when_forwarded(
gateway: Gateway, assistant: dict[str, JsonValue], expected: list[str]
) -> None:
marker: Final = new_marker()
calls: Final[list[JsonValue]] = [*marked_calls(cities=CITIES[:2]), _SERVER_CALL]
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls, assistant=assistant)))
assert status == 200, text
labels: Final = anthropic_labels(_only_request(wire))
assert labels == expected
def test_azure_ai_claude_points_stay_within_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"azure_ai/{ANTHROPIC_MODEL}",
api_base=wire.url,
api_key=PROVIDER_KEY,
cache_control_injection_points=POINTS,
)
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls())))
assert status == 200, text
request: Final = _only_request(wire)
assert request.target == "/anthropic/v1/messages", request.target
assert anthropic_labels(request) == _CLIENT_MARKED
def _gemini_peer(request: Request) -> Reply:
if "cachedContents" in request.target and request.method == "GET":
return Reply(body=b'{"cachedContents":[]}')
if "cachedContents" in request.target:
return Reply(
body=json.dumps(
{
"name": "cachedContents/census",
"model": "models/gemini-3.8-flash",
"expireTime": "2099-01-01T00:00:00Z",
}
).encode()
)
return Reply(body=_GEMINI_REPLY)
_GEMINI: Final = "gemini/gemini-3.8-flash"
_GEMINI_CACHE_WRITE: Final = [
("GET", "/models/gemini-3.8-flash:cachedContents"),
("POST", "/models/gemini-3.8-flash:cachedContents"),
("POST", "/models/gemini-3.8-flash:generateContent"),
]
@pytest.mark.parametrize(
("calls", "expected"),
(
pytest.param(
marked_calls(), [("POST", "/models/gemini-3.8-flash:generateContent")], id="tool-calls-fill-the-cap"
),
pytest.param([tool_call(city) for city in CITIES], _GEMINI_CACHE_WRITE, id="unmarked-tool-calls"),
),
)
def test_gemini_context_cache_follows_the_cap_census(
gateway: Gateway, calls: list[JsonValue], expected: list[tuple[str, str]]
) -> None:
marker: Final = new_marker()
long_system: Final[JsonValue] = {"role": "system", "content": "lorem " * (2 * get_prompt_cache_min_tokens(_GEMINI))}
messages: Final = [long_system, *conversation(marker, calls)[1:]]
with wire_server(_gemini_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=_GEMINI,
api_base=wire.url,
api_key=PROVIDER_KEY,
cache_control_injection_points=POINTS,
)
status, _, text = post_chat(gateway, chat_body(model, messages))
assert status == 200, text
targets: Final = [(request.method, request.target.split("?")[0]) for request in wire.drain()]
assert targets == expected
def _openai_marks(request: Request) -> list[str]:
body: Final = _JSON_OBJECT.validate_json(request.body)
messages: Final = body["messages"] if isinstance(body["messages"], list) else []
return [label for message in messages if isinstance(message, dict) for label in _openai_message_marks(message)]
def _openai_message_marks(message: dict[str, JsonValue]) -> list[str]:
role: Final = str(message.get("role"))
content: Final = message.get("content")
calls: Final = message.get("tool_calls")
blocks: Final = content if isinstance(content, list) else []
return [
*(f"{key}@{role}:message" for key in ("cache_control", "prompt_cache_breakpoint") if key in message),
*(
f"{key}@{role}:text:{block.get('text')}"
for block in blocks
if isinstance(block, dict)
for key in ("cache_control", "prompt_cache_breakpoint")
if key in block
),
*(
f"cache_control@tool_call:{call.get('id')}"
for call in (calls if isinstance(calls, list) else [])
if isinstance(call, dict) and "cache_control" in call
),
]
@pytest.mark.parametrize(
"options",
(
pytest.param({"prompt_cache_options": {"mode": "explicit"}}, id="breakpoint-dialect"),
pytest.param({}, id="plain"),
),
)
def test_openai_points_skip_injection_when_client_marks_fill_the_cap(
gateway: Gateway, options: dict[str, JsonValue]
) -> None:
marker: Final = new_marker()
with wire_server(lambda request: Reply(body=_OPENAI_REPLY)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model="openai/gpt-6-sol",
api_base=f"{wire.url}/v1",
api_key=PROVIDER_KEY,
cache_control_injection_points=POINTS,
**options,
)
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls())))
assert status == 200, text
marks: Final = _openai_marks(_only_request(wire))
assert marks == [f"cache_control@user:text:{ASK}", *(f"cache_control@tool_call:{call_id(city)}" for city in CITIES)]
@pytest.mark.parametrize("stream", (False, True), ids=("json", "stream"))
def test_messages_endpoint_points_skip_injection_when_tool_use_marks_fill_the_cap(
gateway: Gateway, stream: bool
) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, text = _stream(gateway, "/v1/messages", messages_body(model, marker, stream=stream))
assert status == 200, text
assert (
("event: message_stop" in text) if stream else (_JSON_OBJECT.validate_json(text)["id"] == f"msg_{marker}")
)
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
def test_responses_bridge_keeps_system_and_user_marks_within_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
response: Final = gateway.request("POST", "/v1/responses", responses_body(model, marker))
assert response.status_code == 200, response.text
assert _JSON_OBJECT.validate_json(response.content)["status"] == "completed", response.text
assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, ASK_LABEL]
def test_response_cache_serves_the_capped_request_once(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
body: Final = chat_body(model, conversation(marker, marked_calls()))
first: Final = post_chat(gateway, body)
second: Final = post_chat(gateway, body)
received: Final = wire.drain()
assert (first[0], second[0]) == (200, 200), (first[2], second[2])
assert first[1].startswith("chatcmpl-"), first[2]
assert first[1] == second[1], (first[2], second[2])
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED]
def test_unauthenticated_request_never_reaches_the_provider(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, _, text = post_chat(
gateway, chat_body(model, conversation(marker, marked_calls())), key=f"sk-not-a-key-{marker}"
)
received: Final = wire.drain()
assert status == 401, text
assert "error" in _JSON_OBJECT.validate_json(text), text
assert received == ()
@pytest.mark.parametrize(
"calls",
(
pytest.param({"tool_calls": []}, id="empty"),
pytest.param({"tool_calls": None}, id="null"),
pytest.param({}, id="missing"),
),
)
def test_assistant_without_tool_calls_keeps_configured_points(gateway: Gateway, calls: dict[str, JsonValue]) -> None:
marker: Final = new_marker()
messages: Final[list[JsonValue]] = [
{"role": "system", "content": SYSTEM},
{"role": "user", "content": [{"type": "text", "text": ASK, "cache_control": EPHEMERAL}]},
{"role": "assistant", "content": "I will check.", **calls},
{"role": "user", "content": final_text(marker)},
]
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, _, text = post_chat(gateway, chat_body(model, messages))
assert status == 200, text
labels: Final = anthropic_labels(_only_request(wire))
assert labels == [SYSTEM_LABEL, ASK_LABEL, final_label(marker)]
def test_responses_stream_bridge_keeps_system_and_user_marks_within_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, text = _stream(gateway, "/v1/responses", {**responses_body(model, marker), "stream": True})
assert status == 200, text
assert '"type":"response.completed"' in text, text
assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, ASK_LABEL]
def _openai_client(gateway: Gateway) -> openai.OpenAI:
return openai.OpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=gateway.key,
max_retries=0,
http_client=httpx.Client(trust_env=False, timeout=60),
)
def _async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI:
return openai.AsyncOpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=gateway.key,
max_retries=0,
http_client=httpx.AsyncClient(trust_env=False, timeout=60),
)
def _anthropic_client(gateway: Gateway) -> anthropic.Anthropic:
return anthropic.Anthropic(
base_url=str(gateway.client.base_url),
api_key=gateway.key,
max_retries=0,
http_client=httpx.Client(trust_env=False, timeout=60),
)
def _async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic:
return anthropic.AsyncAnthropic(
base_url=str(gateway.client.base_url),
api_key=gateway.key,
max_retries=0,
http_client=httpx.AsyncClient(trust_env=False, timeout=60),
)
def _sdk_messages(marker: str) -> list[ChatCompletionMessageParam]:
return cast(list[ChatCompletionMessageParam], conversation(marker, marked_calls()))
_SDK_TOOLS: Final = cast(list[ChatCompletionToolParam], [TOOL])
def test_openai_sdk_sync_chat_keeps_the_capped_request_within_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
with _openai_client(gateway) as client:
completion: Final = client.chat.completions.create(
model=model, messages=_sdk_messages(marker), tools=_SDK_TOOLS, max_tokens=64
)
assert completion.choices[0].message.content == "sunny", completion.model_dump()
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
assert gateway_injected(completion.id) is False
async def test_openai_sdk_async_chat_stream_keeps_the_capped_request_within_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
async with _async_openai_client(gateway) as client:
stream: Final = await client.chat.completions.create(
model=model, messages=_sdk_messages(marker), tools=_SDK_TOOLS, max_tokens=64, stream=True
)
chunks: Final = [chunk async for chunk in stream]
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == "sunny", chunks
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
assert gateway_injected(chunks[0].id) is False
def _messages_system(model: str, marker: str) -> Iterable[TextBlockParam]:
return cast(Iterable[TextBlockParam], messages_body(model, marker, stream=False)["system"])
def _messages_tools(model: str, marker: str) -> Iterable[ToolParam]:
return cast(Iterable[ToolParam], messages_body(model, marker, stream=False)["tools"])
def _messages_turns(model: str, marker: str) -> Iterable[MessageParam]:
return cast(Iterable[MessageParam], messages_body(model, marker, stream=False)["messages"])
def test_anthropic_sdk_sync_messages_keeps_the_capped_request_within_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
with _anthropic_client(gateway) as client:
message: Final = client.messages.create(
model=model,
max_tokens=64,
system=_messages_system(model, marker),
tools=_messages_tools(model, marker),
messages=_messages_turns(model, marker),
)
assert message.id == f"msg_{marker}", message.model_dump()
assert [block.text for block in message.content if block.type == "text"] == ["sunny"], message.model_dump()
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
assert gateway_injected(message.id) is False
async def test_anthropic_sdk_async_messages_stream_keeps_the_capped_request_within_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
async with _async_anthropic_client(gateway) as client:
stream: Final = await client.messages.create(
model=model,
max_tokens=64,
system=_messages_system(model, marker),
tools=_messages_tools(model, marker),
messages=_messages_turns(model, marker),
stream=True,
)
events: Final = [event async for event in stream]
assert [event.type for event in events][-1] == "message_stop", events
starts: Final = [event for event in events if event.type == "message_start"]
assert [start.message.id for start in starts] == [f"msg_{marker}"], events
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
assert gateway_injected(f"msg_{marker}") is False
async def test_openai_sdk_async_responses_keeps_system_and_user_marks_within_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
body: Final = responses_body(model, marker)
async with _async_openai_client(gateway) as client:
response: Final = await client.responses.create(
model=model,
instructions=SYSTEM,
max_output_tokens=64,
tools=cast(Iterable[ResponsesToolParam], body["tools"]),
input=cast(ResponseInputParam, body["input"]),
)
assert response.status == "completed", response.model_dump()
assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, ASK_LABEL]
def test_request_level_points_skip_injection_when_client_marks_fill_the_cap(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire)
body: Final = chat_body(model, conversation(marker, marked_calls()), cache_control_injection_points=POINTS)
status, response_id, text = post_chat(gateway, body)
assert status == 200, text
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
assert gateway_injected(response_id) is False
_OTHER_CALL_RESULT: Final[dict[str, JsonValue]] = {
"web_search_results": [{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_other", "content": []}]
}
_SERVER_MARK_COUNTED: Final = [
ASK_LABEL,
tool_use_label(CITIES[0]),
tool_use_label(CITIES[1]),
"assistant:tool_use:srvtoolu_web",
]
@pytest.mark.parametrize(
"fields",
(
pytest.param(1, id="int"),
pytest.param(["web_search_results"], id="list"),
pytest.param("", id="empty-string"),
pytest.param("x" * 5120, id="5kb-string"),
pytest.param({"web_search_results": "srvtoolu_web"}, id="results-not-a-list"),
pytest.param({"web_search_results": ["srvtoolu_web", 1]}, id="results-without-objects"),
pytest.param(_OTHER_CALL_RESULT, id="result-for-another-call"),
),
)
def test_malformed_provider_specific_fields_count_the_server_tool_call_mark(
gateway: Gateway, fields: JsonValue
) -> None:
marker: Final = new_marker()
calls: Final[list[JsonValue]] = [*marked_calls(cities=CITIES[:2]), _SERVER_CALL]
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, _, text = post_chat(
gateway, chat_body(model, conversation(marker, calls, assistant={"provider_specific_fields": fields}))
)
assert status == 200, text
labels: Final = anthropic_labels(_only_request(wire))
assert labels == _SERVER_MARK_COUNTED
@pytest.mark.parametrize(
("identity", "forwarded"),
(
pytest.param("", "tool_use_id", id="empty-string"),
pytest.param("x" * 5120, "x" * 5120, id="5kb-string"),
),
)
def test_odd_tool_call_ids_still_count_their_marks(gateway: Gateway, identity: str, forwarded: str) -> None:
marker: Final = new_marker()
calls: Final[list[JsonValue]] = [
*marked_calls(cities=CITIES[:2]),
tool_call(CITIES[2], cache_control=EPHEMERAL, id=identity),
]
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls)))
assert status == 200, text
labels: Final = anthropic_labels(_only_request(wire))
assert labels == [
ASK_LABEL,
tool_use_label(CITIES[0]),
tool_use_label(CITIES[1]),
f"assistant:tool_use:{forwarded}",
]
def _without_id(call: dict[str, JsonValue]) -> dict[str, JsonValue]:
return {key: value for key, value in call.items() if key != "id"}
@pytest.mark.parametrize(
"odd_call",
(
pytest.param(tool_call(CITIES[2], cache_control=EPHEMERAL, id=123), id="int"),
pytest.param(tool_call(CITIES[2], cache_control=EPHEMERAL, id=["call_x"]), id="list"),
pytest.param(tool_call(CITIES[2], cache_control=EPHEMERAL, id=None), id="null"),
pytest.param(_without_id(tool_call(CITIES[2], cache_control=EPHEMERAL)), id="missing"),
),
)
def test_non_string_tool_call_ids_answer_400_without_an_upstream_call(
gateway: Gateway, odd_call: dict[str, JsonValue]
) -> None:
marker: Final = new_marker()
calls: Final[list[JsonValue]] = [*marked_calls(cities=CITIES[:2]), odd_call]
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls)))
received: Final = wire.drain()
assert status == 400, text
assert '"error"' in text, text
assert len(received) == 0, [request.target for request in received]
def _unauthorized_peer(request: Request) -> Reply:
assert len(anthropic_labels(request)) <= 4, anthropic_labels(request)
return Reply(
status=401,
body=json.dumps(
{"type": "error", "error": {"type": "authentication_error", "message": "invalid x-api-key"}}
).encode(),
)
def test_provider_error_on_a_capped_request_reaches_the_caller(gateway: Gateway) -> None:
marker: Final = new_marker()
with wire_server(_unauthorized_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
response: Final = gateway.request(
"POST", "/v1/chat/completions", chat_body(model, conversation(marker, marked_calls()))
)
received: Final = wire.drain()
assert response.status_code == 401, response.text
assert "invalid x-api-key" in response.text, response.text
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED]
@pytest.mark.parametrize("points", (pytest.param(None, id="null"), pytest.param([], id="empty")))
def test_points_null_or_empty_forward_only_the_client_marks(gateway: Gateway, points: JsonValue) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=points)
status, response_id, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls())))
assert status == 200, text
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
assert gateway_injected(response_id) is False
def _cache_hit_rows(response_id: str) -> list[dict[str, JsonValue]]:
return read_rows(
"""SELECT request_id, metadata->>'litellm_gateway_injected_cache' AS injected FROM "LiteLLM_SpendLogs" """
"WHERE request_id LIKE %s",
(f"{response_id}_cache_hit%",),
)
def test_response_cache_hit_records_the_injection_on_the_first_row_only(gateway: Gateway) -> None:
marker: Final = new_marker()
unmarked: Final = [tool_call(city) for city in CITIES]
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
body: Final = chat_body(model, conversation(marker, unmarked, ask_marked=False))
first: Final = post_chat(gateway, body)
second: Final = post_chat(gateway, body)
received: Final = wire.drain()
assert (first[0], second[0]) == (200, 200), (first[2], second[2])
assert first[1] == second[1], (first[2], second[2])
assert [anthropic_labels(request) for request in received] == [[SYSTEM_LABEL, final_label(marker)]]
assert gateway_injected(first[1]) is True
hit_rows: Final = eventually(lambda: _cache_hit_rows(first[1]), lambda rows: len(rows) == 1, seconds=70)
assert hit_rows[0]["injected"] is None, hit_rows
@pytest.mark.parametrize("path", ("/v1/messages", "/v1/responses"))
def test_response_cache_twins_on_messages_and_responses_stay_within_the_cap(gateway: Gateway, path: str) -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
body: Final = (
messages_body(model, marker, stream=False) if path == "/v1/messages" else responses_body(model, marker)
)
responses: Final = tuple(gateway.request("POST", path, body) for _ in range(2))
received: Final = wire.drain()
assert [response.status_code for response in responses] == [200, 200], [response.text for response in responses]
expected: Final = _CLIENT_MARKED if path == "/v1/messages" else [SYSTEM_LABEL, ASK_LABEL]
assert [anthropic_labels(request) for request in received] == [expected] * len(received)
assert len(received) == 1, [request.target for request in received]
def _model_id(gateway: Gateway, model: str) -> str:
entries: Final = gateway.get("/model/info")["data"]
assert isinstance(entries, list), entries
matches: Final = [
string_value(object_value(object_value(entry)["model_info"])["id"])
for entry in entries
if object_value(entry)["model_name"] == model
]
assert len(matches) == 1, matches
return matches[0]
_ASSISTANT_POINT: Final[list[JsonValue]] = [{"location": "message", "role": "assistant"}]
_UPDATE_BURST: Final = 20
async def _capped_burst(gateway: Gateway, model: str) -> tuple[tuple[int, str], ...]:
async def one(client: httpx.AsyncClient) -> tuple[int, str]:
body: Final = chat_body(model, conversation(new_marker(), marked_calls()))
response: Final = await client.post(
"/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {gateway.key}"}
)
return response.status_code, response.text
async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client:
return tuple(await asyncio.gather(*(one(client) for _ in range(_UPDATE_BURST))))
def _probe_assistant_point(gateway: Gateway, wire: Wire, model: str, marker: str) -> list[str]:
messages: Final[list[JsonValue]] = [
{"role": "system", "content": SYSTEM},
{"role": "user", "content": ASK},
{"role": "assistant", "content": "I will check."},
{"role": "user", "content": final_text(marker)},
]
status, _, text = post_chat(gateway, chat_body(model, messages))
assert status == 200, text
return anthropic_labels(wire.drain()[-1])
async def test_points_update_mid_burst_keeps_every_request_within_the_cap(gateway: Gateway) -> None:
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
identity: Final = _model_id(gateway, model)
burst: Final = asyncio.create_task(_capped_burst(gateway, model))
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 3, 30)
await asyncio.to_thread(
gateway.post,
"/model/update",
{
"model_name": model,
"litellm_params": {
"model": f"anthropic/{ANTHROPIC_MODEL}",
"api_base": wire.url,
"api_key": PROVIDER_KEY,
"cache_control_injection_points": _ASSISTANT_POINT,
},
"model_info": {"id": identity},
},
)
sent: Final = await burst
received: Final = wire.drain()
assert [status for status, _ in sent] == [200] * _UPDATE_BURST, [text for _, text in sent]
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED] * _UPDATE_BURST
landed: Final = await asyncio.to_thread(
eventually,
lambda: _probe_assistant_point(gateway, wire, model, new_marker()),
lambda labels: labels == ["assistant:text:I will check."],
60,
)
assert landed == ["assistant:text:I will check."]

View file

@ -0,0 +1,93 @@
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows, scratch_database
from integration._support.process import owned_proxy_process
from integration._support.wire import Wire, wire_server
from integration.providers._cache_control_marks_support import (
CITIES,
anthropic_deployment,
anthropic_peer,
chat_body,
conversation,
final_text,
marked_calls,
marker_of,
new_marker,
owned_config,
post_chat,
tool_call,
)
from pydantic import JsonValue
_MODEL: Final = "affinity-claude"
_FOLLOW_UPS: Final = 24
_LONG_SYSTEM: Final = "Answer from the tool results below. " + "lorem ipsum " * 1500
def _first_turn(session: str, marker: str) -> list[JsonValue]:
calls: Final = [*marked_calls(cities=CITIES[:2]), tool_call(CITIES[2])]
return [
{"role": "system", "content": f"{_LONG_SYSTEM} session {session}"},
*conversation(marker, calls, ask_marked=False)[1:],
]
def _follow_up(first_turn: list[JsonValue], marker: str) -> list[JsonValue]:
return [*first_turn, {"role": "assistant", "content": "sunny"}, {"role": "user", "content": final_text(marker)}]
def _served(wire: Wire) -> frozenset[str]:
return frozenset(marker_of(request) for request in wire.drain())
@pytest.mark.timeout(240)
def test_router_affinity_is_lost_when_auto_caching_stands_down_for_tool_call_marks(
gateway: Gateway, tmp_path: Path
) -> None:
session: Final = new_marker()
first_marker: Final = new_marker()
follow_up_markers: Final = tuple(new_marker() for _ in range(_FOLLOW_UPS))
first_turn: Final = _first_turn(session, first_marker)
with scratch_database() as database_url, wire_server(anthropic_peer) as left, wire_server(anthropic_peer) as right:
config: Final = owned_config(
tmp_path,
[
{**anthropic_deployment(_MODEL, left.url), "model_info": {"id": f"affinity-left-{session}"}},
{**anthropic_deployment(_MODEL, right.url), "model_info": {"id": f"affinity-right-{session}"}},
],
router_settings={"optional_pre_call_checks": ["prompt_caching"]},
)
with (
owned_proxy_process(
gateway,
tmp_path,
{"DATABASE_URL": database_url},
config=config,
remove_environment=("DATABASE_URL_READ_REPLICA",),
workers=2,
) as owned,
owned.gateway.scenario() as scenario,
):
key: Final = scenario.key(metadata={"enable_prompt_caching": True})
status, response_id, text = post_chat(owned.gateway, chat_body(_MODEL, first_turn), key=key)
assert status == 200, text
eventually(
lambda: read_rows(
'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(response_id,),
database_url=database_url,
),
lambda rows: len(rows) == 1,
seconds=70,
)
follow_ups: Final = tuple(
post_chat(owned.gateway, chat_body(_MODEL, _follow_up(first_turn, marker)), key=key)
for marker in follow_up_markers
)
served: Final = {"left": _served(left), "right": _served(right)}
assert [status for status, _, _ in follow_ups] == [200] * _FOLLOW_UPS, [text for _, _, text in follow_ups]
assert {first_marker, *follow_up_markers} == served["left"] | served["right"]
assert all(served[side] & set(follow_up_markers) for side in served), served

View file

@ -0,0 +1,80 @@
import json
from typing import Final
from integration._support.wire import Wire, wire_server
from integration.providers._cache_control_marks_support import (
ANTHROPIC_MODEL,
CITIES,
EPHEMERAL,
POINTS,
PROVIDER_KEY,
SYSTEM,
TOOL,
anthropic_labels,
anthropic_peer,
ask,
call_id,
client_marked,
final_text,
new_marker,
)
from pydantic import JsonValue
import litellm
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.types.utils import ChatCompletionMessageToolCall, Message, ModelResponse
_CLIENT_MARKED: Final = client_marked()
def _marked_call(city: str) -> ChatCompletionMessageToolCall:
return ChatCompletionMessageToolCall(
id=call_id(city),
type="function",
function={"name": "lookup_weather", "arguments": json.dumps({"city": city})},
cache_control=EPHEMERAL,
)
def _pydantic_conversation(marker: str) -> list[JsonValue | Message]:
return [
{"role": "system", "content": SYSTEM},
ask(marked=True),
Message(role="assistant", content="", tool_calls=[_marked_call(city) for city in CITIES]),
*({"role": "tool", "tool_call_id": call_id(city), "content": "sunny."} for city in CITIES),
{"role": "user", "content": final_text(marker)},
]
def _completion_kwargs(wire: Wire, marker: str) -> dict[str, object]:
return {
"model": f"anthropic/{ANTHROPIC_MODEL}",
"messages": _pydantic_conversation(marker),
"tools": [TOOL],
"max_tokens": 64,
"api_base": wire.url,
"api_key": PROVIDER_KEY,
"num_retries": 0,
"cache_control_injection_points": POINTS,
}
def test_pydantic_tool_call_marks_count_against_the_cap_in_process() -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire:
response: Final = litellm.completion(**_completion_kwargs(wire, marker)) # pyright: ignore[reportArgumentType] # Message objects stand in for the dict shapes the signature names
assert isinstance(response, ModelResponse), type(response)
assert response.choices[0].message.content == "sunny", response # pyright: ignore[reportAttributeAccessIssue] # choices are Choices for a non-stream response
received: Final = wire.drain()
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED]
async def test_pydantic_tool_call_marks_count_against_the_cap_in_process_async_stream() -> None:
marker: Final = new_marker()
with wire_server(anthropic_peer) as wire:
stream: Final = await litellm.acompletion(**_completion_kwargs(wire, marker), stream=True) # pyright: ignore[reportArgumentType] # Message objects stand in for the dict shapes the signature names
assert isinstance(stream, CustomStreamWrapper), type(stream)
chunks: Final = [chunk async for chunk in stream]
assert "".join(str(chunk.choices[0].delta.content or "") for chunk in chunks) == "sunny", chunks
received: Final = wire.drain()
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED]

View file

@ -4,7 +4,8 @@ import os
import subprocess
import sys
import textwrap
from typing import Final, List, Optional, Tuple
from collections.abc import Mapping
from typing import Final, List, Optional, Tuple, cast
from unittest.mock import MagicMock, patch
import pytest
@ -15,8 +16,13 @@ from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
supports_openai_prompt_cache_breakpoint,
)
from litellm.litellm_core_utils.prompt_templates.factory import (
_convert_to_bedrock_tool_call_invoke,
convert_to_anthropic_tool_invoke,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.openai import AllMessageValues, ChatCompletionAssistantToolCall
from litellm.types.utils import ChatCompletionMessageToolCall, Message
@pytest.fixture(autouse=True)
@ -1045,6 +1051,36 @@ def _count_cache_control(messages: List[AllMessageValues]) -> int:
return count
def _count_tool_call_cache_controls(message: AllMessageValues) -> int:
message_mapping: Final = cast(Mapping[str, object], message)
tool_calls: Final = message_mapping.get("tool_calls")
tool_call_values: Final = cast(list[object], tool_calls) if isinstance(tool_calls, list) else None
return (
sum(
1
for tool_call in tool_call_values
if isinstance(tool_call, dict) and isinstance(tool_call.get("cache_control"), dict)
)
if tool_call_values is not None
else 0
)
def _marked_function_tool_calls() -> list[ChatCompletionAssistantToolCall]:
return cast(
list[ChatCompletionAssistantToolCall],
[
{
"id": f"call_{index}",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
"cache_control": {"type": "ephemeral"},
}
for index in range(3)
],
)
def _build_injection_points():
return [
{
@ -1060,6 +1096,184 @@ def _build_injection_points():
]
def test_cache_control_hook_counts_tool_call_cache_controls():
message: Final[AllMessageValues] = {
"role": "assistant",
"content": None,
"tool_calls": _marked_function_tool_calls(),
}
assert AnthropicCacheControlHook.count_request_cache_breakpoints([message]) == 3
def test_cache_control_hook_counts_tool_call_marks_except_answered_server_tool_calls():
message: Final[AllMessageValues] = {
"role": "assistant",
"content": None,
"tool_calls": cast(
list[ChatCompletionAssistantToolCall],
[
{
"id": "nested",
"type": "function",
"function": {
"name": "lookup",
"arguments": "{}",
"cache_control": {"type": "ephemeral"},
},
},
{
"id": "non_function",
"type": "custom",
"function": {"name": "lookup", "arguments": "{}"},
"cache_control": {"type": "ephemeral"},
},
{
"id": "prompt_breakpoint",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
"prompt_cache_breakpoint": {"type": "ephemeral"},
},
{
"id": "srvtoolu_web_search",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
"cache_control": {"type": "ephemeral"},
},
{
"id": "srvtoolu_tool_result",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
"cache_control": {"type": "ephemeral"},
},
{
"id": "srvtoolu_unmatched",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
"cache_control": {"type": "ephemeral"},
},
],
),
"provider_specific_fields": {
"web_search_results": [{"tool_use_id": "srvtoolu_web_search"}],
"tool_results": [{"tool_use_id": "srvtoolu_tool_result"}],
},
}
assert AnthropicCacheControlHook.count_request_cache_breakpoints([message]) == 2
@pytest.mark.parametrize("tool_call_type", ["function", "custom", None])
@pytest.mark.parametrize(
"mark",
[{"type": "ephemeral"}, {}, "ephemeral", "", 7, ["ephemeral"], "x" * 5000, None],
ids=["dict", "empty_dict", "string", "empty_string", "int", "list", "5kb_string", "none"],
)
def test_tool_call_census_matches_the_breakpoints_providers_send(mark: object, tool_call_type: str | None):
tool_call: Final[dict[str, object]] = {
"id": "call_0",
"function": {"name": "lookup", "arguments": "{}"},
"cache_control": mark,
**({"type": tool_call_type} if tool_call_type is not None else {}),
}
message: Final = cast(AllMessageValues, {"role": "assistant", "content": None, "tool_calls": [tool_call]})
bedrock_cache_points: Final = sum(
1
for block in _convert_to_bedrock_tool_call_invoke([tool_call], model="anthropic.claude-sonnet-4-5-20250929-v1:0")
if "cachePoint" in block
)
anthropic_marks: Final = sum(
1 for block in convert_to_anthropic_tool_invoke([tool_call]) if block.get("cache_control") is not None
)
census: Final = AnthropicCacheControlHook.count_request_cache_breakpoints([message])
assert census == bedrock_cache_points
assert census >= anthropic_marks
assert census == (0 if mark is None else 1)
def test_cache_control_hook_caps_customer_tool_call_marks_before_injection():
hook = AnthropicCacheControlHook()
messages: Final[list[AllMessageValues]] = [
{"role": "system", "content": "Follow the tool instructions."},
{
"role": "user",
"content": [{"type": "text", "text": "Look up three values.", "cache_control": {"type": "ephemeral"}}],
},
{"role": "assistant", "content": None, "tool_calls": _marked_function_tool_calls()},
{"role": "tool", "tool_call_id": "call_0", "content": "first"},
{"role": "tool", "tool_call_id": "call_1", "content": "second"},
{"role": "tool", "tool_call_id": "call_2", "content": "third"},
{"role": "user", "content": "Summarize the values."},
]
_, processed, _ = hook.get_chat_completion_prompt(
model="bedrock/us.anthropic.claude-opus-4-6-v1:0",
messages=messages,
non_default_params={"cache_control_injection_points": _build_injection_points()},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
)
forwarded_mark_count: Final = _count_cache_control(processed) + sum(
_count_tool_call_cache_controls(message) for message in processed
)
assert forwarded_mark_count <= 4
assert AnthropicCacheControlHook.count_request_cache_breakpoints(processed) == forwarded_mark_count
def test_cache_control_hook_counts_pydantic_message_tool_call_marks():
tool_call: Final = ChatCompletionMessageToolCall(
id="call_1",
type="function",
function={"name": "lookup", "arguments": "{}"},
cache_control={"type": "ephemeral"},
)
message: Final = Message(role="assistant", content=None, tool_calls=[tool_call])
assert AnthropicCacheControlHook.count_request_cache_breakpoints(cast(list[AllMessageValues], [message])) == 1
def test_injection_skips_assistant_whose_tool_call_carries_a_longer_ttl_mark():
hook: Final = AnthropicCacheControlHook()
tool_call_ttl: Final = {"type": "ephemeral", "ttl": "1h"}
messages: Final[list[AllMessageValues]] = [
{"role": "user", "content": "What is the weather in Paris?"},
{
"role": "assistant",
"content": "Let me look that up.",
"tool_calls": [
{
"id": "toolu_01A",
"type": "function",
"function": {"name": "get_weather", "arguments": '{"city": "Paris"}'},
"cache_control": tool_call_ttl,
}
],
},
{"role": "tool", "tool_call_id": "toolu_01A", "content": "18C and sunny"},
]
_, processed, _ = hook.get_chat_completion_prompt(
model="anthropic/claude-haiku-4-5",
messages=messages,
non_default_params={"cache_control_injection_points": [{"location": "message", "role": "assistant"}]},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
)
assistant_message: Final = processed[1]
assistant_tool_calls: Final = assistant_message.get("tool_calls")
assert assistant_message.get("cache_control") is None
assert assistant_message.get("content") == "Let me look that up."
assert isinstance(assistant_tool_calls, list)
assert assistant_tool_calls[0].get("cache_control") == tool_call_ttl
assert AnthropicCacheControlHook.count_request_cache_breakpoints(processed) == 1
def test_cache_control_hook_caps_at_four_blocks_with_client_cache_control():
"""Regression for LIT-3667 / Anthropic 'A maximum of 4 blocks ... Found 5'.
@ -1947,6 +2161,40 @@ class TestEnableAnthropicPromptCaching:
assert result_sys == "sys"
assert result_msgs == messages
@pytest.mark.parametrize(
"tool_call_controls",
[
pytest.param(({"type": "ephemeral"},) * 3, id="three_5m_marks_would_exceed_the_cap"),
pytest.param(({"type": "ephemeral", "ttl": "1h"},), id="1h_mark_would_follow_a_5m_default"),
],
)
def test_seed_stands_down_when_only_assistant_tool_calls_carry_cache_control(self, tool_call_controls):
tool_calls: Final = [
{
"id": f"call_{index}",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
"cache_control": control,
}
for index, control in enumerate(tool_call_controls)
]
messages: Final = [
{"role": "system", "content": "a long system prompt"},
{"role": "user", "content": "weather in three cities"},
{"role": "assistant", "content": "Checking.", "tool_calls": tool_calls},
*({"role": "tool", "tool_call_id": call["id"], "content": "sunny"} for call in tool_calls),
{"role": "user", "content": "summarize"},
]
params: Final[dict] = {}
AnthropicCacheControlHook.maybe_seed_default_injection_points(
non_default_params=params,
messages=cast(List[AllMessageValues], messages),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
enable_prompt_caching=True,
)
assert "cache_control_injection_points" not in params
def test_default_ttl_is_anthropics_five_minute_cache(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
assert all(p["control"] == {"type": "ephemeral"} for p in self._points())

View file

@ -4121,3 +4121,30 @@ def test_tool_changes_beta_requires_system_tool_reference(role: str, content: ob
)
assert ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER not in headers.get("anthropic-beta", "").split(",")
@pytest.mark.parametrize(
("tool_call_id", "provider_specific_fields", "rebuilt"),
(
("srvtoolu_search", {"web_search_results": [{"tool_use_id": "srvtoolu_search"}]}, True),
("srvtoolu_code", {"tool_results": [{"tool_use_id": "srvtoolu_code"}]}, True),
(
"srvtoolu_code",
{"web_search_results": "srvtoolu_code", "tool_results": [{"tool_use_id": "srvtoolu_code"}]},
True,
),
("srvtoolu_other", {"web_search_results": [{"tool_use_id": "srvtoolu_search"}]}, False),
("call_client", {"tool_results": [{"tool_use_id": "call_client"}]}, False),
("srvtoolu_search", {"web_search_results": ["srvtoolu_search"]}, False),
("srvtoolu_search", {}, False),
("srvtoolu_search", None, False),
("srvtoolu_search", [{"tool_use_id": "srvtoolu_search"}], False),
(None, {"web_search_results": [{"tool_use_id": None}]}, False),
),
)
def test_tool_call_is_rebuilt_as_server_tool_use_only_with_a_stored_result(
tool_call_id: object, provider_specific_fields: object, rebuilt: bool
) -> None:
from litellm.llms.anthropic.common_utils import tool_call_is_rebuilt_as_server_tool_use
assert tool_call_is_rebuilt_as_server_tool_use(tool_call_id, provider_specific_fields) is rebuilt