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
This commit is contained in:
mateo-berri 2026-09-28 21:43:01 -07:00
parent 466022fd84
commit e341ab4649
7 changed files with 1402 additions and 10 deletions

View file

@ -149,10 +149,8 @@ def _has_server_tool_result(tool_call_id: str, results: Iterable[object] | None)
return any(isinstance(result, dict) and result.get("tool_use_id") == tool_call_id for result in results or ())
def _tool_call_cache_control_is_forwarded(tool_call: object, message: object) -> bool:
if _attribute_or_key(tool_call, "type") != "function" or not isinstance(
_attribute_or_key(tool_call, "cache_control"), dict
):
def _tool_call_carries_cache_breakpoint(tool_call: object, message: object) -> bool:
if _attribute_or_key(tool_call, "cache_control") is None:
return False
tool_call_id: Final = _attribute_or_key(tool_call, "id")
@ -524,7 +522,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
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_cache_control_is_forwarded(tool_call, message))
sum(1 for tool_call in tool_calls if _tool_call_carries_cache_breakpoint(tool_call, message))
if tool_calls
else 0
)

View file

@ -27,8 +27,8 @@ 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-4-6"
VERTEX_LOCATION: Final[str] = "us-east5"
VERTEX_MODEL: Final[str] = "vertex_ai/claude-sonnet-5"
VERTEX_LOCATION: Final[str] = "global"
def _deployment_params(*, backend: Backend, inject_cache_control: bool) -> LiteLLMParamsBody:
@ -150,7 +150,9 @@ def _post_chat(client: PassthroughClient, key: str, body: ChatBody) -> Result[Ch
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 == "stop", f"{model_name}: unexpected finish reason: {completion.finish_reason}"
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

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,245 @@
import asyncio
import re
import signal
import socket
from collections import Counter
from collections.abc import Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Final
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 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
client_port: int
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)
async with client.stream("POST", path, json=body, headers={"Authorization": f"Bearer {key}"}) as response:
client_port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1])
await response.aread()
return _Sent(
surface,
marker,
response.status_code,
response.text,
response.headers.get("x-litellm-call-id", ""),
client_port,
)
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 _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:
with wire_server(anthropic_peer) 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 = 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))
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30)
victim: Final = psutil.Process(workers[0])
victim.suspend()
victim_ports: Final = frozenset(
connection.raddr.port for connection in victim.net_connections(kind="tcp") if connection.raddr
)
victim.send_signal(signal.SIGKILL)
served: Final = await burst
during: Final = wire.drain()
after: Final = await _fire(owned_url, owned.gateway.key)
after_received: Final = wire.drain()
assert all(len(requests) == 1 for requests in _by_marker(during).values())
_assert_capped(tuple(item for item in served if item.status == 200), during)
_assert_capped(after, after_received)
survivors: Final = tuple(item for item in served if item.client_port not in victim_ports)
assert survivors, [item.client_port for item in served]
for item in (*survivors, *after):
_single_spend_row(item)
for item in served:
assert len(_spend_rows(item.call_id)) <= 1, item.call_id
@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])],
]

View file

@ -0,0 +1,525 @@
import json
from datetime import datetime, timedelta, timezone
from typing import Final
import pytest
from integration._support.client import Gateway, Scenario, eventually, object_value
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,
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 litellm.utils import get_prompt_cache_min_tokens
from pydantic import JsonValue, TypeAdapter
_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 = [*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)]

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

@ -16,6 +16,10 @@ 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, ChatCompletionAssistantToolCall
from litellm.types.utils import ChatCompletionMessageToolCall, Message
@ -1102,7 +1106,7 @@ def test_cache_control_hook_counts_tool_call_cache_controls():
assert AnthropicCacheControlHook.count_request_cache_breakpoints([message]) == 3
def test_cache_control_hook_counts_only_tool_call_marks_forwarded_to_anthropic():
def test_cache_control_hook_counts_tool_call_marks_except_answered_server_tool_calls():
message: Final[AllMessageValues] = {
"role": "assistant",
"content": None,
@ -1156,7 +1160,37 @@ def test_cache_control_hook_counts_only_tool_call_marks_forwarded_to_anthropic()
},
}
assert AnthropicCacheControlHook.count_request_cache_breakpoints([message]) == 1
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():
@ -2127,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())