This commit is contained in:
devin-ai-integration[bot] 2026-09-30 11:10:40 -07:00 • committed by GitHub
commit 86d692a6ef
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 1837 additions and 10 deletions

View file

@ -122,7 +122,53 @@ 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 _as_object_iterable(value: object | None) -> Iterable[object] | None:
if not isinstance(value, Iterable):
return None
return value
def _has_server_tool_result(tool_call_id: str, results: Iterable[object] | None) -> bool:
return any(isinstance(result, dict) and result.get("tool_use_id") == tool_call_id for result in results or ())
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")
if not isinstance(tool_call_id, str) or not tool_call_id.startswith("srvtoolu_"):
return True
provider_specific_fields: Final = _attribute_or_key(message, "provider_specific_fields")
if not isinstance(provider_specific_fields, dict):
return True
server_tool_result_keys: Final = ("web_search_results", "tool_results")
return not any(
_has_server_tool_result(
tool_call_id,
_as_object_iterable(provider_specific_fields.get(result_key)),
)
for result_key in server_tool_result_keys
)
def _tool_carries_cache_breakpoint(tool: object) -> bool:
@ -471,13 +517,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

@ -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 ChatAssistantTurn(BaseModel):
@ -1214,6 +1215,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.
@ -1272,6 +1279,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,242 @@
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 typing import Final
import httpx
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 _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 = 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, 20)
os.kill(workers[0], 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 0 < len(served) < _BURST, len(served)
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])],
]

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

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