mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
466022fd84
commit
e341ab4649
7 changed files with 1402 additions and 10 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
461
tests/integration/providers/_cache_control_marks_support.py
Normal file
461
tests/integration/providers/_cache_control_marks_support.py
Normal 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
|
||||
|
|
@ -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])],
|
||||
]
|
||||
|
|
@ -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)]
|
||||
|
|
@ -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
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue