mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(caching): count tool_call cache_control marks in the injection census (#43556)
* fix(caching): count tool_call cache_control marks in the injection census
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(caching): remove cache census casts
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(caching): only skip injection on message or content marks
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(caching): skip injection on messages whose tool calls carry marks
Reverts 9f08d8aef8. A default 5m mark injected on assistant text lands
before the client's 1h tool_use mark, which Anthropic rejects with a 400
because a 1h breakpoint must not follow a 5m one. Keeping the full census
in the skip check leaves the client's tool_call breakpoint as the only
one on that message.
* fix(caching): count every client tool_call cache mark in the breakpoint census
The census gated tool_call marks on type function and dict shape, so a client mark on a call without a type or with a string cache_control slipped past the count and injection overflowed the 4 breakpoint cap. Count any non-None tool_call mark, keep the server tool exclusion, and add integration cells for the capped surfaces, the yaml stand-down, Bedrock and Gemini, and router affinity
* test(integration): hold the upstream so the worker kill lands mid-burst
* refactor(caching): reuse the transform's server tool lookup in the breakpoint census
The census now calls the same helper the Anthropic transform uses to decide
whether a marked tool call becomes a server tool block, so the two cannot
drift apart. The owned-proxy burst test waits up to 90 seconds for the burst
to reach the wire before it kills a worker
* refactor(anthropic): move the server tool rebuild check under llms/anthropic
* test(integration): audit the tool call mark census across chat, messages, responses, and chaos
Adds the /audit cells for the breakpoint census on assistant tool_calls marks: the
Responses stream bridge, the OpenAI and Anthropic SDK clients, in-process Pydantic
messages, response cache twins, request-level points, a provider 401 on a capped request,
malformed provider_specific_fields and tool_call ids, null or empty points, a points
update mid-burst, a proxy restart mid-burst, and a worker SIGKILL that picks the worker
holding the burst's upstream connections
* test(integration): close the SDK clients the cache census cells open
---------
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
6370104c53
commit
9d16412341
12 changed files with 2405 additions and 16 deletions
|
|
@ -28,6 +28,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
from litellm.llms.anthropic.common_utils import (
|
||||
is_claude_code_one_shot_subagent_request,
|
||||
supports_anthropic_cache_control,
|
||||
tool_call_is_rebuilt_as_server_tool_use,
|
||||
)
|
||||
from litellm.types.integrations.anthropic_cache_control_hook import (
|
||||
GATEWAY_INJECTED_CACHE_METADATA_KEY,
|
||||
|
|
@ -122,7 +123,30 @@ def targets_openai_api(api_base: object) -> bool:
|
|||
|
||||
|
||||
def _carries_cache_breakpoint(block: object) -> bool:
|
||||
return isinstance(block, dict) and any(block.get(key) is not None for key in CACHE_BREAKPOINT_KEYS)
|
||||
return any(_attribute_or_key(block, key) is not None for key in CACHE_BREAKPOINT_KEYS)
|
||||
|
||||
|
||||
def _attribute_or_key(value: object, key: str) -> object | None:
|
||||
if hasattr(value, key):
|
||||
return getattr(value, key)
|
||||
if isinstance(value, Mapping):
|
||||
return value.get(key)
|
||||
return None
|
||||
|
||||
|
||||
def _as_object_list(value: object | None) -> list[object] | None:
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
return _validated_object_list(value)
|
||||
|
||||
|
||||
def _tool_call_carries_cache_breakpoint(tool_call: object, message: object) -> bool:
|
||||
if _attribute_or_key(tool_call, "cache_control") is None:
|
||||
return False
|
||||
|
||||
return not tool_call_is_rebuilt_as_server_tool_use(
|
||||
_attribute_or_key(tool_call, "id"), _attribute_or_key(message, "provider_specific_fields")
|
||||
)
|
||||
|
||||
|
||||
def _tool_carries_cache_breakpoint(tool: object) -> bool:
|
||||
|
|
@ -471,13 +495,16 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
@staticmethod
|
||||
def _count_cache_control_blocks(message: object) -> int:
|
||||
if not isinstance(message, dict):
|
||||
return 0
|
||||
count = 1 if _carries_cache_breakpoint(message) else 0
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, list):
|
||||
count += sum(1 for block in content if _carries_cache_breakpoint(block))
|
||||
return count
|
||||
message_count: Final = 1 if _carries_cache_breakpoint(message) else 0
|
||||
content: Final = _as_object_list(_attribute_or_key(message, "content"))
|
||||
content_count: Final = sum(1 for block in content if _carries_cache_breakpoint(block)) if content else 0
|
||||
tool_calls: Final = _as_object_list(_attribute_or_key(message, "tool_calls"))
|
||||
tool_call_count: Final = (
|
||||
sum(1 for tool_call in tool_calls if _tool_call_carries_cache_breakpoint(tool_call, message))
|
||||
if tool_calls
|
||||
else 0
|
||||
)
|
||||
return message_count + content_count + tool_call_count
|
||||
|
||||
@staticmethod
|
||||
def _message_has_cache_control(message: AllMessageValues) -> bool:
|
||||
|
|
|
|||
|
|
@ -1729,11 +1729,16 @@ def convert_function_to_anthropic_tool_invoke(
|
|||
raise e
|
||||
|
||||
|
||||
def _find_server_tool_result(
|
||||
ANTHROPIC_SERVER_TOOL_USE_ID_PREFIX: Final = "srvtoolu_"
|
||||
|
||||
|
||||
def find_anthropic_server_tool_result(
|
||||
tool_id: str,
|
||||
web_search_results: Sequence[object] | None,
|
||||
tool_results: Sequence[object] | None,
|
||||
) -> dict[str, object] | None:
|
||||
if not tool_id.startswith(ANTHROPIC_SERVER_TOOL_USE_ID_PREFIX):
|
||||
return None
|
||||
candidates: Final = (*(web_search_results or ()), *(tool_results or ()))
|
||||
return next(
|
||||
(result for result in candidates if isinstance(result, dict) and result.get("tool_use_id") == tool_id),
|
||||
|
|
@ -1808,11 +1813,7 @@ def convert_to_anthropic_tool_invoke(
|
|||
context="Anthropic tool invoke",
|
||||
)
|
||||
|
||||
server_tool_result = (
|
||||
_find_server_tool_result(tool_id, web_search_results, tool_results)
|
||||
if tool_id.startswith("srvtoolu_")
|
||||
else None
|
||||
)
|
||||
server_tool_result = find_anthropic_server_tool_result(tool_id, web_search_results, tool_results)
|
||||
if server_tool_result is not None:
|
||||
anthropic_tool_invoke.append(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
THOUGHT_SIGNATURE_SEPARATOR,
|
||||
find_anthropic_server_tool_result,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import message_field, parts_of
|
||||
from litellm.llms.anthropic.wif import (
|
||||
|
|
@ -1834,6 +1835,20 @@ def _replayed_server_tool_use(block: object) -> _ReplayedServerToolUse | None:
|
|||
return None
|
||||
|
||||
|
||||
def tool_call_is_rebuilt_as_server_tool_use(tool_call_id: object, provider_specific_fields: object) -> bool:
|
||||
fields: Final = _validated_claude_code_mapping(provider_specific_fields)
|
||||
if not isinstance(tool_call_id, str) or fields is None:
|
||||
return False
|
||||
return (
|
||||
find_anthropic_server_tool_result(
|
||||
tool_call_id,
|
||||
_validated_claude_code_list(fields.get("web_search_results")),
|
||||
_validated_claude_code_list(fields.get("tool_results")),
|
||||
)
|
||||
is not None
|
||||
)
|
||||
|
||||
|
||||
def _render_web_search_results(
|
||||
query: str, results: tuple[_ReplayedWebSearchResult, ...] | _ReplayedWebSearchToolResultError
|
||||
) -> str:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -294,6 +294,7 @@ class ToolCall(BaseModel):
|
|||
id: str | None = None
|
||||
type: str | None = None
|
||||
function: ToolCallFunction = ToolCallFunction()
|
||||
cache_control: CacheControl | None = None
|
||||
|
||||
|
||||
class ThinkingBlock(BaseModel):
|
||||
|
|
@ -1233,6 +1234,12 @@ class FineTuningJobsResponse(BaseModel):
|
|||
# ---------- model management ----------
|
||||
|
||||
|
||||
class CacheControlInjectionPoint(BaseModel):
|
||||
location: Literal["message"]
|
||||
role: str | None = None
|
||||
index: int | None = None
|
||||
|
||||
|
||||
class LiteLLMParamsBody(BaseModel):
|
||||
"""POST /model/new litellm_params: `model` is the only required field; `api_key`
|
||||
et al may be an `os.environ/FOO` reference the proxy resolves at call time.
|
||||
|
|
@ -1291,6 +1298,7 @@ class LiteLLMParamsBody(BaseModel):
|
|||
max_retries: int | None = None
|
||||
cooldown_time: float | None = None
|
||||
extra_body: DeploymentExtraBody | None = None
|
||||
cache_control_injection_points: list[CacheControlInjectionPoint] | None = None
|
||||
tpm: int | None = None
|
||||
weight: int | None = None
|
||||
order: int | None = None
|
||||
|
|
|
|||
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,294 @@
|
|||
import asyncio
|
||||
import os
|
||||
import re
|
||||
import signal
|
||||
import socket
|
||||
import threading
|
||||
from collections import Counter
|
||||
from collections.abc import Callable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from integration.providers._cache_control_marks_support import (
|
||||
ASK_LABEL,
|
||||
CITIES,
|
||||
POINTS,
|
||||
SYSTEM_LABEL,
|
||||
anthropic_deployment,
|
||||
anthropic_labels,
|
||||
anthropic_peer,
|
||||
chat_body,
|
||||
client_marked,
|
||||
conversation,
|
||||
final_label,
|
||||
marked_calls,
|
||||
marker_of,
|
||||
messages_body,
|
||||
new_marker,
|
||||
owned_config,
|
||||
responses_body,
|
||||
tool_call,
|
||||
tool_use_label,
|
||||
)
|
||||
from pydantic import JsonValue
|
||||
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
_SURFACES: Final = ("chat", "chat-stream", "chat-unmarked", "messages", "responses")
|
||||
_MODEL: Final = "capped-claude"
|
||||
_BURST: Final = 30
|
||||
_OUTAGE_STATUS: Final = 500
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Sent:
|
||||
surface: str
|
||||
marker: str
|
||||
status: int
|
||||
text: str
|
||||
call_id: str
|
||||
|
||||
|
||||
def _free_port() -> int:
|
||||
with socket.socket() as reserve:
|
||||
reserve.bind(("127.0.0.1", 0))
|
||||
return int(reserve.getsockname()[1])
|
||||
|
||||
|
||||
def _surface_request(surface: str, marker: str) -> tuple[str, dict[str, JsonValue]]:
|
||||
if surface == "messages":
|
||||
return "/v1/messages", messages_body(_MODEL, marker, stream=False)
|
||||
if surface == "responses":
|
||||
return "/v1/responses", responses_body(_MODEL, marker)
|
||||
if surface == "chat-unmarked":
|
||||
unmarked: Final = [tool_call(city) for city in CITIES]
|
||||
return "/v1/chat/completions", chat_body(_MODEL, conversation(marker, unmarked, ask_marked=False))
|
||||
return "/v1/chat/completions", chat_body(
|
||||
_MODEL, conversation(marker, marked_calls()), stream=surface == "chat-stream"
|
||||
)
|
||||
|
||||
|
||||
def _expected_labels(item: _Sent) -> list[str]:
|
||||
if item.surface == "responses":
|
||||
return [SYSTEM_LABEL, ASK_LABEL]
|
||||
if item.surface == "chat-unmarked":
|
||||
return [SYSTEM_LABEL, final_label(item.marker)]
|
||||
return client_marked()
|
||||
|
||||
|
||||
async def _fire(owned_url: str, key: str, *, tolerate_transport_errors: bool = False) -> tuple[_Sent, ...]:
|
||||
async def one(client: httpx.AsyncClient, index: int) -> _Sent:
|
||||
surface: Final = _SURFACES[index % len(_SURFACES)]
|
||||
marker: Final = new_marker()
|
||||
path, body = _surface_request(surface, marker)
|
||||
response: Final = await client.post(path, json=body, headers={"Authorization": f"Bearer {key}"})
|
||||
return _Sent(
|
||||
surface, marker, response.status_code, response.text, response.headers.get("x-litellm-call-id", "")
|
||||
)
|
||||
|
||||
async with httpx.AsyncClient(base_url=owned_url, timeout=30, trust_env=False) as client:
|
||||
results: Final = await asyncio.gather(
|
||||
*(one(client, index) for index in range(_BURST)), return_exceptions=tolerate_transport_errors
|
||||
)
|
||||
for result in results:
|
||||
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
|
||||
return tuple(result for result in results if isinstance(result, _Sent))
|
||||
|
||||
|
||||
def _held_anthropic_peer(release: threading.Event) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert release.wait(timeout=120), "Held upstream was never released"
|
||||
return anthropic_peer(request)
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _held_upstream_connections(pid: int, upstream: str) -> int:
|
||||
port: Final = urlsplit(upstream).port
|
||||
return sum(
|
||||
1
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
|
||||
)
|
||||
|
||||
|
||||
def _by_marker(received: Sequence[Request]) -> dict[str, tuple[Request, ...]]:
|
||||
counted: Final = Counter(marker_of(request) for request in received)
|
||||
return {marker: tuple(request for request in received if marker_of(request) == marker) for marker in counted}
|
||||
|
||||
|
||||
def _assert_capped(served: Sequence[_Sent], received: Sequence[Request]) -> None:
|
||||
by_marker: Final = _by_marker(received)
|
||||
for item in served:
|
||||
assert item.status == 200, (item.surface, item.text)
|
||||
assert len(by_marker.get(item.marker, ())) == 1, (item.surface, item.marker)
|
||||
assert anthropic_labels(by_marker[item.marker][0]) == _expected_labels(item), item.surface
|
||||
|
||||
|
||||
def _assert_outage_error(item: _Sent, received: Sequence[Request]) -> None:
|
||||
assert item.status == _OUTAGE_STATUS, (item.surface, item.status, item.text)
|
||||
assert '"error"' in item.text, (item.surface, item.text)
|
||||
assert item.marker not in {marker_of(request) for request in received}, item.surface
|
||||
|
||||
|
||||
def _spend_rows(call_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,))
|
||||
|
||||
|
||||
def _single_spend_row(item: _Sent) -> None:
|
||||
assert item.call_id, (item.surface, item.status, item.text)
|
||||
rows: Final = eventually(lambda: _spend_rows(item.call_id), lambda values: len(values) == 1, seconds=70)
|
||||
assert len(rows) == 1, (item.surface, item.call_id)
|
||||
|
||||
|
||||
@pytest.mark.timeout(240)
|
||||
async def test_capped_burst_rides_out_a_provider_outage_and_logs_every_request_once(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
port: Final = _free_port()
|
||||
config: Final = owned_config(
|
||||
tmp_path, [anthropic_deployment(_MODEL, f"http://127.0.0.1:{port}", cache_control_injection_points=POINTS)]
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
|
||||
owned_url: Final = str(owned.gateway.client.base_url)
|
||||
with wire_server(anthropic_peer, port=port) as wire:
|
||||
healthy: Final = await _fire(owned_url, owned.gateway.key)
|
||||
_assert_capped(healthy, wire.drain())
|
||||
racing: Final = asyncio.create_task(_fire(owned_url, owned.gateway.key))
|
||||
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30)
|
||||
down: Final = await _fire(owned_url, owned.gateway.key)
|
||||
raced: Final = await racing
|
||||
raced_received: Final = wire.drain()
|
||||
with wire_server(anthropic_peer, port=port) as restarted:
|
||||
recovered: Final = await _fire(owned_url, owned.gateway.key)
|
||||
recovered_received: Final = restarted.drain()
|
||||
assert len(_STARTED_WORKER.findall(owned.log.read_text())) >= 2
|
||||
assert all(len(requests) == 1 for requests in _by_marker(raced_received).values())
|
||||
_assert_capped(tuple(item for item in raced if item.status == 200), raced_received)
|
||||
for item in raced:
|
||||
if item.status != 200:
|
||||
_assert_outage_error(item, raced_received)
|
||||
for item in down:
|
||||
_assert_outage_error(item, (*raced_received, *recovered_received))
|
||||
_assert_capped(recovered, recovered_received)
|
||||
assert {marker_of(request) for request in recovered_received} == {item.marker for item in recovered}
|
||||
for item in (*healthy, *raced, *down, *recovered):
|
||||
_single_spend_row(item)
|
||||
|
||||
|
||||
@pytest.mark.timeout(240)
|
||||
async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_capped_requests(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
release: Final = threading.Event()
|
||||
with wire_server(_held_anthropic_peer(release)) as wire:
|
||||
config: Final = owned_config(
|
||||
tmp_path, [anthropic_deployment(_MODEL, wire.url, cache_control_injection_points=POINTS)]
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
|
||||
workers: Final[tuple[int, ...]] = eventually(
|
||||
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
|
||||
lambda pids: len(pids) == 2,
|
||||
seconds=30,
|
||||
)
|
||||
owned_url: Final = str(owned.gateway.client.base_url)
|
||||
burst: Final = asyncio.create_task(_fire(owned_url, owned.gateway.key, tolerate_transport_errors=True))
|
||||
try:
|
||||
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= _BURST, 90)
|
||||
held_by: Final = MappingProxyType({pid: _held_upstream_connections(pid, wire.url) for pid in workers})
|
||||
victim: Final = max(workers, key=held_by.__getitem__)
|
||||
os.kill(victim, signal.SIGKILL)
|
||||
finally:
|
||||
release.set()
|
||||
served: Final = await burst
|
||||
during: Final = wire.drain()
|
||||
after: Final = await _fire(owned_url, owned.gateway.key)
|
||||
after_received: Final = wire.drain()
|
||||
assert sum(held_by.values()) == _BURST, held_by
|
||||
assert held_by[victim] > 0, held_by
|
||||
assert len(served) == _BURST - held_by[victim], (len(served), held_by)
|
||||
assert len(after) == _BURST, len(after)
|
||||
assert all(len(requests) == 1 for requests in _by_marker(during).values())
|
||||
_assert_capped(served, during)
|
||||
_assert_capped(after, after_received)
|
||||
for item in (*served, *after):
|
||||
_single_spend_row(item)
|
||||
|
||||
|
||||
@pytest.mark.timeout(240)
|
||||
def test_yaml_auto_caching_stands_down_for_tool_call_marks_and_outranks_a_key_opt_out(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
markers: Final = (new_marker(), new_marker(), new_marker())
|
||||
unmarked: Final = [tool_call(city) for city in CITIES]
|
||||
with wire_server(anthropic_peer) as wire:
|
||||
config: Final = owned_config(
|
||||
tmp_path,
|
||||
[anthropic_deployment(_MODEL, wire.url)],
|
||||
litellm_settings={"enable_anthropic_prompt_caching": True},
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
opted_out: Final = scenario.key(metadata={"enable_prompt_caching": False})
|
||||
cells: Final = (
|
||||
(conversation(markers[0], marked_calls(), ask_marked=False), owned.gateway.key),
|
||||
(conversation(markers[1], unmarked, ask_marked=False), owned.gateway.key),
|
||||
(conversation(markers[2], unmarked, ask_marked=False), opted_out),
|
||||
)
|
||||
responses: Final = tuple(
|
||||
owned.gateway.request("POST", "/v1/chat/completions", chat_body(_MODEL, messages), key=key)
|
||||
for messages, key in cells
|
||||
)
|
||||
received: Final = _by_marker(wire.drain())
|
||||
assert [response.status_code for response in responses] == [200, 200, 200], [
|
||||
response.text for response in responses
|
||||
]
|
||||
assert [anthropic_labels(received[marker][0]) for marker in markers] == [
|
||||
[tool_use_label(city) for city in CITIES],
|
||||
[SYSTEM_LABEL, final_label(markers[1])],
|
||||
[SYSTEM_LABEL, final_label(markers[2])],
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.timeout(360)
|
||||
async def test_proxy_restart_mid_burst_serves_capped_requests_after_the_reboot(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
release: Final = threading.Event()
|
||||
with wire_server(_held_anthropic_peer(release)) as wire:
|
||||
config: Final = owned_config(
|
||||
tmp_path, [anthropic_deployment(_MODEL, wire.url, cache_control_injection_points=POINTS)]
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as first:
|
||||
burst: Final = asyncio.create_task(
|
||||
_fire(str(first.gateway.client.base_url), first.gateway.key, tolerate_transport_errors=True)
|
||||
)
|
||||
try:
|
||||
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= _BURST, 90)
|
||||
first.process.terminate()
|
||||
finally:
|
||||
release.set()
|
||||
served: Final = await burst
|
||||
during: Final = wire.drain()
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as second:
|
||||
after: Final = await _fire(str(second.gateway.client.base_url), second.gateway.key)
|
||||
after_received: Final = wire.drain()
|
||||
completed: Final = tuple(item for item in served if item.status == 200)
|
||||
lost: Final = tuple(item for item in served if item.status != 200)
|
||||
assert len(after) == _BURST, len(after)
|
||||
assert all(len(requests) == 1 for requests in _by_marker(during).values())
|
||||
_assert_capped(completed, during)
|
||||
for item in lost:
|
||||
assert '"error"' in item.text or item.text == "", (item.surface, item.status, item.text)
|
||||
_assert_capped(after, after_received)
|
||||
for item in (*completed, *after):
|
||||
_single_spend_row(item)
|
||||
|
|
@ -0,0 +1,934 @@
|
|||
import asyncio
|
||||
import json
|
||||
from collections.abc import Iterable
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final, cast
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
from anthropic.types import MessageParam, TextBlockParam, ToolParam
|
||||
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from integration.providers._cache_control_marks_support import (
|
||||
ANTHROPIC_MODEL,
|
||||
ASK,
|
||||
ASK_LABEL,
|
||||
BEDROCK_MODEL,
|
||||
CITIES,
|
||||
EPHEMERAL,
|
||||
POINTS,
|
||||
PROVIDER_KEY,
|
||||
SYSTEM,
|
||||
SYSTEM_LABEL,
|
||||
TOOL,
|
||||
anthropic_labels,
|
||||
anthropic_marks,
|
||||
anthropic_peer,
|
||||
bedrock_labels,
|
||||
bedrock_peer,
|
||||
call_id,
|
||||
chat_body,
|
||||
client_marked,
|
||||
conversation,
|
||||
final_label,
|
||||
final_text,
|
||||
gateway_injected,
|
||||
marked_calls,
|
||||
messages_body,
|
||||
new_marker,
|
||||
post_chat,
|
||||
responses_body,
|
||||
tool_call,
|
||||
tool_use_label,
|
||||
)
|
||||
from openai.types.chat import ChatCompletionMessageParam, ChatCompletionToolParam
|
||||
from openai.types.responses import ResponseInputParam
|
||||
from openai.types.responses import ToolParam as ResponsesToolParam
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
from litellm.utils import get_prompt_cache_min_tokens
|
||||
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_CLIENT_MARKED: Final = client_marked()
|
||||
_GEMINI_REPLY: Final = json.dumps(
|
||||
{
|
||||
"candidates": [{"content": {"role": "model", "parts": [{"text": "sunny"}]}, "finishReason": "STOP"}],
|
||||
"usageMetadata": {"promptTokenCount": 1600, "candidatesTokenCount": 1, "totalTokenCount": 1601},
|
||||
}
|
||||
).encode()
|
||||
_OPENAI_REPLY: Final = json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-cache-census",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-6-sol",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "sunny"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 12, "completion_tokens": 1, "total_tokens": 13},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _anthropic_deployment(scenario: Scenario, wire: Wire, **fields: JsonValue) -> str:
|
||||
return scenario.model(model=f"anthropic/{ANTHROPIC_MODEL}", api_base=wire.url, api_key=PROVIDER_KEY, **fields)
|
||||
|
||||
|
||||
def _only_request(wire: Wire) -> Request:
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 1, [request.target for request in received]
|
||||
return received[0]
|
||||
|
||||
|
||||
def _stream(gateway: Gateway, path: str, body: dict[str, JsonValue], *, key: str | None = None) -> tuple[int, str]:
|
||||
headers: Final = {"Authorization": f"Bearer {gateway.key if key is None else key}"}
|
||||
with gateway.client.stream("POST", path, json=body, headers=headers) as response:
|
||||
return response.status_code, response.read().decode()
|
||||
|
||||
|
||||
def _chat_stream_chunks(text: str) -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
_JSON_OBJECT.validate_json(line.removeprefix("data: "))
|
||||
for line in text.splitlines()
|
||||
if line.startswith("data: ") and line != "data: [DONE]"
|
||||
]
|
||||
|
||||
|
||||
def _chat_stream_content(chunks: list[dict[str, JsonValue]]) -> str:
|
||||
return "".join(
|
||||
str(object_value(object_value(choice).get("delta") or {}).get("content") or "")
|
||||
for chunk in chunks
|
||||
for choice in (chunk.get("choices") if isinstance(chunk.get("choices"), list) else [])
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", (False, True), ids=("json", "stream"))
|
||||
def test_chat_points_skip_injection_when_client_marks_fill_the_cap(gateway: Gateway, stream: bool) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
body: Final = chat_body(model, conversation(marker, marked_calls()), stream=stream)
|
||||
if stream:
|
||||
status, text = _stream(gateway, "/v1/chat/completions", body)
|
||||
assert status == 200, text
|
||||
chunks: Final = _chat_stream_chunks(text)
|
||||
assert _chat_stream_content(chunks) == "sunny", text
|
||||
assert "data: [DONE]" in text, text
|
||||
response_id = str(chunks[0]["id"])
|
||||
else:
|
||||
status, response_id, text = post_chat(gateway, body)
|
||||
assert status == 200, text
|
||||
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
|
||||
assert gateway_injected(response_id) is False
|
||||
|
||||
|
||||
def test_chat_points_inject_system_and_last_message_without_client_marks(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
unmarked: Final = [tool_call(city) for city in CITIES]
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, response_id, text = post_chat(
|
||||
gateway, chat_body(model, conversation(marker, unmarked, ask_marked=False))
|
||||
)
|
||||
assert status == 200, text
|
||||
assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, final_label(marker)]
|
||||
assert gateway_injected(response_id) is True
|
||||
|
||||
|
||||
def _prompt_caching_rows(gateway: Gateway, cursor: dict[str, str]) -> list[dict[str, JsonValue]]:
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
page: Final = gateway.get(
|
||||
"/cost_optimization/prompt_caching/requests",
|
||||
{
|
||||
"start_date": (now - timedelta(minutes=10)).isoformat(),
|
||||
"end_date": (now + timedelta(minutes=10)).isoformat(),
|
||||
"page_size": "100",
|
||||
"filter": "injected",
|
||||
**cursor,
|
||||
},
|
||||
)
|
||||
rows: Final = [object_value(row) for row in page["requests"]] if isinstance(page["requests"], list) else []
|
||||
following: Final = page.get("next_cursor")
|
||||
if not page.get("has_more") or not isinstance(following, dict):
|
||||
return rows
|
||||
return [
|
||||
*rows,
|
||||
*_prompt_caching_rows(
|
||||
gateway,
|
||||
{"cursor_start_time": str(following["start_time"]), "cursor_request_id": str(following["request_id"])},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _listed_as_injected(gateway: Gateway, response_id: str) -> bool:
|
||||
return any(row["request_id"] == response_id for row in _prompt_caching_rows(gateway, {}))
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("calls", "ask_marked", "expected", "injected"),
|
||||
(
|
||||
pytest.param(marked_calls(), False, [tool_use_label(city) for city in CITIES], False, id="three-tool-calls"),
|
||||
pytest.param([tool_call(city) for city in CITIES], False, None, True, id="no-client-marks"),
|
||||
pytest.param(
|
||||
[*marked_calls(cities=CITIES[:2]), tool_call(CITIES[2])],
|
||||
False,
|
||||
[tool_use_label(city) for city in CITIES[:2]],
|
||||
False,
|
||||
id="two-tool-calls",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_auto_prompt_caching_stands_down_when_client_marks_only_tool_calls(
|
||||
gateway: Gateway,
|
||||
calls: list[JsonValue],
|
||||
ask_marked: bool,
|
||||
expected: list[str] | None,
|
||||
injected: bool,
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire)
|
||||
key: Final = scenario.key(metadata={"enable_prompt_caching": True})
|
||||
status, response_id, text = post_chat(
|
||||
gateway, chat_body(model, conversation(marker, calls, ask_marked=ask_marked)), key=key
|
||||
)
|
||||
assert status == 200, text
|
||||
labels: Final = anthropic_labels(_only_request(wire))
|
||||
assert labels == (expected if expected is not None else [SYSTEM_LABEL, final_label(marker)])
|
||||
assert gateway_injected(response_id) is injected
|
||||
assert (
|
||||
eventually(lambda: _listed_as_injected(gateway, response_id), lambda listed: listed is injected, 70) is injected
|
||||
)
|
||||
|
||||
|
||||
def test_assistant_point_skips_message_whose_tool_call_carries_a_one_hour_mark(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
hour: Final[dict[str, JsonValue]] = {"type": "ephemeral", "ttl": "1h"}
|
||||
calls: Final = [tool_call(CITIES[0], cache_control=hour)]
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(
|
||||
scenario, wire, cache_control_injection_points=[{"location": "message", "role": "assistant"}]
|
||||
)
|
||||
status, _, text = post_chat(
|
||||
gateway,
|
||||
chat_body(model, conversation(marker, calls, ask_marked=False, assistant={"content": "I will check."})),
|
||||
)
|
||||
assert status == 200, text
|
||||
marks: Final = anthropic_marks(_only_request(wire))
|
||||
assert [(mark.label, mark.ttl) for mark in marks] == [(tool_use_label(CITIES[0]), "1h")]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"mark",
|
||||
(
|
||||
pytest.param("ephemeral", id="string"),
|
||||
pytest.param(1, id="int"),
|
||||
pytest.param(["ephemeral"], id="list"),
|
||||
pytest.param("", id="empty-string"),
|
||||
pytest.param("x" * 5120, id="5kb-string"),
|
||||
),
|
||||
)
|
||||
def test_non_object_tool_call_marks_count_against_the_cap(gateway: Gateway, mark: JsonValue) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls(mark))))
|
||||
assert status == 200, text
|
||||
assert anthropic_labels(_only_request(wire)) == [ASK_LABEL]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("calls", "expected"),
|
||||
(
|
||||
pytest.param(marked_calls({}), _CLIENT_MARKED, id="empty-object"),
|
||||
pytest.param(marked_calls(None), None, id="null"),
|
||||
pytest.param(
|
||||
[tool_call(CITIES[0], cache_control=EPHEMERAL)] * 2,
|
||||
[SYSTEM_LABEL, ASK_LABEL, tool_use_label(CITIES[0])],
|
||||
id="same-call-twice",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_tool_call_mark_shapes_keep_the_request_within_the_cap(
|
||||
gateway: Gateway, calls: list[JsonValue], expected: list[str] | None
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls)))
|
||||
assert status == 200, text
|
||||
labels: Final = anthropic_labels(_only_request(wire))
|
||||
assert labels == (expected if expected is not None else [SYSTEM_LABEL, ASK_LABEL, final_label(marker)])
|
||||
|
||||
|
||||
def _untyped_call(city: str) -> dict[str, JsonValue]:
|
||||
return {name: value for name, value in tool_call(city, cache_control=EPHEMERAL).items() if name != "type"}
|
||||
|
||||
|
||||
_BEDROCK_CLIENT_MARKED: Final = [f"user:text:{ASK}", *(f"assistant:toolUse:{call_id(city)}" for city in CITIES)]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("calls", "expected"),
|
||||
(
|
||||
pytest.param(marked_calls(), _BEDROCK_CLIENT_MARKED, id="object-marks"),
|
||||
pytest.param(marked_calls("ephemeral"), _BEDROCK_CLIENT_MARKED, id="string-marks"),
|
||||
pytest.param([_untyped_call(city) for city in CITIES], _BEDROCK_CLIENT_MARKED, id="calls-without-type"),
|
||||
pytest.param(
|
||||
[tool_call(CITIES[0], cache_control=EPHEMERAL)] * 2,
|
||||
[f"system:text:{SYSTEM}", ASK_LABEL, f"assistant:toolUse:{call_id(CITIES[0])}", "assistant:other"],
|
||||
id="same-call-twice",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_bedrock_converse_cache_points_stay_within_the_cap(
|
||||
gateway: Gateway, calls: list[JsonValue], expected: list[str]
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model=f"bedrock/converse/{BEDROCK_MODEL}",
|
||||
api_key=PROVIDER_KEY,
|
||||
aws_region_name="us-east-1",
|
||||
aws_bedrock_runtime_endpoint=wire.url,
|
||||
cache_control_injection_points=POINTS,
|
||||
)
|
||||
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls)))
|
||||
assert status == 200, text
|
||||
request: Final = _only_request(wire)
|
||||
assert request.target.endswith("/converse"), request.target
|
||||
assert bedrock_labels(request) == expected
|
||||
|
||||
|
||||
_SERVER_CALL: Final = tool_call(
|
||||
"web", cache_control=EPHEMERAL, id="srvtoolu_web", function={"name": "web_search", "arguments": "{}"}
|
||||
)
|
||||
_WEB_RESULTS: Final[dict[str, JsonValue]] = {
|
||||
"provider_specific_fields": {
|
||||
"web_search_results": [{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_web", "content": []}]
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("assistant", "expected"),
|
||||
(
|
||||
pytest.param(
|
||||
_WEB_RESULTS,
|
||||
[SYSTEM_LABEL, ASK_LABEL, tool_use_label(CITIES[0]), tool_use_label(CITIES[1])],
|
||||
id="server-tool-with-result",
|
||||
),
|
||||
pytest.param(
|
||||
{},
|
||||
[ASK_LABEL, tool_use_label(CITIES[0]), tool_use_label(CITIES[1]), "assistant:tool_use:srvtoolu_web"],
|
||||
id="server-tool-without-result",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_server_tool_call_mark_counts_only_when_forwarded(
|
||||
gateway: Gateway, assistant: dict[str, JsonValue], expected: list[str]
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
calls: Final[list[JsonValue]] = [*marked_calls(cities=CITIES[:2]), _SERVER_CALL]
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls, assistant=assistant)))
|
||||
assert status == 200, text
|
||||
labels: Final = anthropic_labels(_only_request(wire))
|
||||
assert labels == expected
|
||||
|
||||
|
||||
def test_azure_ai_claude_points_stay_within_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model=f"azure_ai/{ANTHROPIC_MODEL}",
|
||||
api_base=wire.url,
|
||||
api_key=PROVIDER_KEY,
|
||||
cache_control_injection_points=POINTS,
|
||||
)
|
||||
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls())))
|
||||
assert status == 200, text
|
||||
request: Final = _only_request(wire)
|
||||
assert request.target == "/anthropic/v1/messages", request.target
|
||||
assert anthropic_labels(request) == _CLIENT_MARKED
|
||||
|
||||
|
||||
def _gemini_peer(request: Request) -> Reply:
|
||||
if "cachedContents" in request.target and request.method == "GET":
|
||||
return Reply(body=b'{"cachedContents":[]}')
|
||||
if "cachedContents" in request.target:
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"name": "cachedContents/census",
|
||||
"model": "models/gemini-3.8-flash",
|
||||
"expireTime": "2099-01-01T00:00:00Z",
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
return Reply(body=_GEMINI_REPLY)
|
||||
|
||||
|
||||
_GEMINI: Final = "gemini/gemini-3.8-flash"
|
||||
_GEMINI_CACHE_WRITE: Final = [
|
||||
("GET", "/models/gemini-3.8-flash:cachedContents"),
|
||||
("POST", "/models/gemini-3.8-flash:cachedContents"),
|
||||
("POST", "/models/gemini-3.8-flash:generateContent"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("calls", "expected"),
|
||||
(
|
||||
pytest.param(
|
||||
marked_calls(), [("POST", "/models/gemini-3.8-flash:generateContent")], id="tool-calls-fill-the-cap"
|
||||
),
|
||||
pytest.param([tool_call(city) for city in CITIES], _GEMINI_CACHE_WRITE, id="unmarked-tool-calls"),
|
||||
),
|
||||
)
|
||||
def test_gemini_context_cache_follows_the_cap_census(
|
||||
gateway: Gateway, calls: list[JsonValue], expected: list[tuple[str, str]]
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
long_system: Final[JsonValue] = {"role": "system", "content": "lorem " * (2 * get_prompt_cache_min_tokens(_GEMINI))}
|
||||
messages: Final = [long_system, *conversation(marker, calls)[1:]]
|
||||
with wire_server(_gemini_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model=_GEMINI,
|
||||
api_base=wire.url,
|
||||
api_key=PROVIDER_KEY,
|
||||
cache_control_injection_points=POINTS,
|
||||
)
|
||||
status, _, text = post_chat(gateway, chat_body(model, messages))
|
||||
assert status == 200, text
|
||||
targets: Final = [(request.method, request.target.split("?")[0]) for request in wire.drain()]
|
||||
assert targets == expected
|
||||
|
||||
|
||||
def _openai_marks(request: Request) -> list[str]:
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
messages: Final = body["messages"] if isinstance(body["messages"], list) else []
|
||||
return [label for message in messages if isinstance(message, dict) for label in _openai_message_marks(message)]
|
||||
|
||||
|
||||
def _openai_message_marks(message: dict[str, JsonValue]) -> list[str]:
|
||||
role: Final = str(message.get("role"))
|
||||
content: Final = message.get("content")
|
||||
calls: Final = message.get("tool_calls")
|
||||
blocks: Final = content if isinstance(content, list) else []
|
||||
return [
|
||||
*(f"{key}@{role}:message" for key in ("cache_control", "prompt_cache_breakpoint") if key in message),
|
||||
*(
|
||||
f"{key}@{role}:text:{block.get('text')}"
|
||||
for block in blocks
|
||||
if isinstance(block, dict)
|
||||
for key in ("cache_control", "prompt_cache_breakpoint")
|
||||
if key in block
|
||||
),
|
||||
*(
|
||||
f"cache_control@tool_call:{call.get('id')}"
|
||||
for call in (calls if isinstance(calls, list) else [])
|
||||
if isinstance(call, dict) and "cache_control" in call
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"options",
|
||||
(
|
||||
pytest.param({"prompt_cache_options": {"mode": "explicit"}}, id="breakpoint-dialect"),
|
||||
pytest.param({}, id="plain"),
|
||||
),
|
||||
)
|
||||
def test_openai_points_skip_injection_when_client_marks_fill_the_cap(
|
||||
gateway: Gateway, options: dict[str, JsonValue]
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(lambda request: Reply(body=_OPENAI_REPLY)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-6-sol",
|
||||
api_base=f"{wire.url}/v1",
|
||||
api_key=PROVIDER_KEY,
|
||||
cache_control_injection_points=POINTS,
|
||||
**options,
|
||||
)
|
||||
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls())))
|
||||
assert status == 200, text
|
||||
marks: Final = _openai_marks(_only_request(wire))
|
||||
assert marks == [f"cache_control@user:text:{ASK}", *(f"cache_control@tool_call:{call_id(city)}" for city in CITIES)]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", (False, True), ids=("json", "stream"))
|
||||
def test_messages_endpoint_points_skip_injection_when_tool_use_marks_fill_the_cap(
|
||||
gateway: Gateway, stream: bool
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, text = _stream(gateway, "/v1/messages", messages_body(model, marker, stream=stream))
|
||||
assert status == 200, text
|
||||
assert (
|
||||
("event: message_stop" in text) if stream else (_JSON_OBJECT.validate_json(text)["id"] == f"msg_{marker}")
|
||||
)
|
||||
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
|
||||
|
||||
|
||||
def test_responses_bridge_keeps_system_and_user_marks_within_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
response: Final = gateway.request("POST", "/v1/responses", responses_body(model, marker))
|
||||
assert response.status_code == 200, response.text
|
||||
assert _JSON_OBJECT.validate_json(response.content)["status"] == "completed", response.text
|
||||
assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, ASK_LABEL]
|
||||
|
||||
|
||||
def test_response_cache_serves_the_capped_request_once(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
body: Final = chat_body(model, conversation(marker, marked_calls()))
|
||||
first: Final = post_chat(gateway, body)
|
||||
second: Final = post_chat(gateway, body)
|
||||
received: Final = wire.drain()
|
||||
assert (first[0], second[0]) == (200, 200), (first[2], second[2])
|
||||
assert first[1].startswith("chatcmpl-"), first[2]
|
||||
assert first[1] == second[1], (first[2], second[2])
|
||||
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED]
|
||||
|
||||
|
||||
def test_unauthenticated_request_never_reaches_the_provider(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, _, text = post_chat(
|
||||
gateway, chat_body(model, conversation(marker, marked_calls())), key=f"sk-not-a-key-{marker}"
|
||||
)
|
||||
received: Final = wire.drain()
|
||||
assert status == 401, text
|
||||
assert "error" in _JSON_OBJECT.validate_json(text), text
|
||||
assert received == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"calls",
|
||||
(
|
||||
pytest.param({"tool_calls": []}, id="empty"),
|
||||
pytest.param({"tool_calls": None}, id="null"),
|
||||
pytest.param({}, id="missing"),
|
||||
),
|
||||
)
|
||||
def test_assistant_without_tool_calls_keeps_configured_points(gateway: Gateway, calls: dict[str, JsonValue]) -> None:
|
||||
marker: Final = new_marker()
|
||||
messages: Final[list[JsonValue]] = [
|
||||
{"role": "system", "content": SYSTEM},
|
||||
{"role": "user", "content": [{"type": "text", "text": ASK, "cache_control": EPHEMERAL}]},
|
||||
{"role": "assistant", "content": "I will check.", **calls},
|
||||
{"role": "user", "content": final_text(marker)},
|
||||
]
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, _, text = post_chat(gateway, chat_body(model, messages))
|
||||
assert status == 200, text
|
||||
labels: Final = anthropic_labels(_only_request(wire))
|
||||
assert labels == [SYSTEM_LABEL, ASK_LABEL, final_label(marker)]
|
||||
|
||||
|
||||
def test_responses_stream_bridge_keeps_system_and_user_marks_within_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, text = _stream(gateway, "/v1/responses", {**responses_body(model, marker), "stream": True})
|
||||
assert status == 200, text
|
||||
assert '"type":"response.completed"' in text, text
|
||||
assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, ASK_LABEL]
|
||||
|
||||
|
||||
def _openai_client(gateway: Gateway) -> openai.OpenAI:
|
||||
return openai.OpenAI(
|
||||
base_url=f"{gateway.client.base_url}/v1",
|
||||
api_key=gateway.key,
|
||||
max_retries=0,
|
||||
http_client=httpx.Client(trust_env=False, timeout=60),
|
||||
)
|
||||
|
||||
|
||||
def _async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI:
|
||||
return openai.AsyncOpenAI(
|
||||
base_url=f"{gateway.client.base_url}/v1",
|
||||
api_key=gateway.key,
|
||||
max_retries=0,
|
||||
http_client=httpx.AsyncClient(trust_env=False, timeout=60),
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_client(gateway: Gateway) -> anthropic.Anthropic:
|
||||
return anthropic.Anthropic(
|
||||
base_url=str(gateway.client.base_url),
|
||||
api_key=gateway.key,
|
||||
max_retries=0,
|
||||
http_client=httpx.Client(trust_env=False, timeout=60),
|
||||
)
|
||||
|
||||
|
||||
def _async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic:
|
||||
return anthropic.AsyncAnthropic(
|
||||
base_url=str(gateway.client.base_url),
|
||||
api_key=gateway.key,
|
||||
max_retries=0,
|
||||
http_client=httpx.AsyncClient(trust_env=False, timeout=60),
|
||||
)
|
||||
|
||||
|
||||
def _sdk_messages(marker: str) -> list[ChatCompletionMessageParam]:
|
||||
return cast(list[ChatCompletionMessageParam], conversation(marker, marked_calls()))
|
||||
|
||||
|
||||
_SDK_TOOLS: Final = cast(list[ChatCompletionToolParam], [TOOL])
|
||||
|
||||
|
||||
def test_openai_sdk_sync_chat_keeps_the_capped_request_within_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
with _openai_client(gateway) as client:
|
||||
completion: Final = client.chat.completions.create(
|
||||
model=model, messages=_sdk_messages(marker), tools=_SDK_TOOLS, max_tokens=64
|
||||
)
|
||||
assert completion.choices[0].message.content == "sunny", completion.model_dump()
|
||||
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
|
||||
assert gateway_injected(completion.id) is False
|
||||
|
||||
|
||||
async def test_openai_sdk_async_chat_stream_keeps_the_capped_request_within_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
async with _async_openai_client(gateway) as client:
|
||||
stream: Final = await client.chat.completions.create(
|
||||
model=model, messages=_sdk_messages(marker), tools=_SDK_TOOLS, max_tokens=64, stream=True
|
||||
)
|
||||
chunks: Final = [chunk async for chunk in stream]
|
||||
assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == "sunny", chunks
|
||||
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
|
||||
assert gateway_injected(chunks[0].id) is False
|
||||
|
||||
|
||||
def _messages_system(model: str, marker: str) -> Iterable[TextBlockParam]:
|
||||
return cast(Iterable[TextBlockParam], messages_body(model, marker, stream=False)["system"])
|
||||
|
||||
|
||||
def _messages_tools(model: str, marker: str) -> Iterable[ToolParam]:
|
||||
return cast(Iterable[ToolParam], messages_body(model, marker, stream=False)["tools"])
|
||||
|
||||
|
||||
def _messages_turns(model: str, marker: str) -> Iterable[MessageParam]:
|
||||
return cast(Iterable[MessageParam], messages_body(model, marker, stream=False)["messages"])
|
||||
|
||||
|
||||
def test_anthropic_sdk_sync_messages_keeps_the_capped_request_within_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
with _anthropic_client(gateway) as client:
|
||||
message: Final = client.messages.create(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
system=_messages_system(model, marker),
|
||||
tools=_messages_tools(model, marker),
|
||||
messages=_messages_turns(model, marker),
|
||||
)
|
||||
assert message.id == f"msg_{marker}", message.model_dump()
|
||||
assert [block.text for block in message.content if block.type == "text"] == ["sunny"], message.model_dump()
|
||||
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
|
||||
assert gateway_injected(message.id) is False
|
||||
|
||||
|
||||
async def test_anthropic_sdk_async_messages_stream_keeps_the_capped_request_within_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
async with _async_anthropic_client(gateway) as client:
|
||||
stream: Final = await client.messages.create(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
system=_messages_system(model, marker),
|
||||
tools=_messages_tools(model, marker),
|
||||
messages=_messages_turns(model, marker),
|
||||
stream=True,
|
||||
)
|
||||
events: Final = [event async for event in stream]
|
||||
assert [event.type for event in events][-1] == "message_stop", events
|
||||
starts: Final = [event for event in events if event.type == "message_start"]
|
||||
assert [start.message.id for start in starts] == [f"msg_{marker}"], events
|
||||
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
|
||||
assert gateway_injected(f"msg_{marker}") is False
|
||||
|
||||
|
||||
async def test_openai_sdk_async_responses_keeps_system_and_user_marks_within_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
body: Final = responses_body(model, marker)
|
||||
async with _async_openai_client(gateway) as client:
|
||||
response: Final = await client.responses.create(
|
||||
model=model,
|
||||
instructions=SYSTEM,
|
||||
max_output_tokens=64,
|
||||
tools=cast(Iterable[ResponsesToolParam], body["tools"]),
|
||||
input=cast(ResponseInputParam, body["input"]),
|
||||
)
|
||||
assert response.status == "completed", response.model_dump()
|
||||
assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, ASK_LABEL]
|
||||
|
||||
|
||||
def test_request_level_points_skip_injection_when_client_marks_fill_the_cap(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire)
|
||||
body: Final = chat_body(model, conversation(marker, marked_calls()), cache_control_injection_points=POINTS)
|
||||
status, response_id, text = post_chat(gateway, body)
|
||||
assert status == 200, text
|
||||
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
|
||||
assert gateway_injected(response_id) is False
|
||||
|
||||
|
||||
_OTHER_CALL_RESULT: Final[dict[str, JsonValue]] = {
|
||||
"web_search_results": [{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_other", "content": []}]
|
||||
}
|
||||
_SERVER_MARK_COUNTED: Final = [
|
||||
ASK_LABEL,
|
||||
tool_use_label(CITIES[0]),
|
||||
tool_use_label(CITIES[1]),
|
||||
"assistant:tool_use:srvtoolu_web",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"fields",
|
||||
(
|
||||
pytest.param(1, id="int"),
|
||||
pytest.param(["web_search_results"], id="list"),
|
||||
pytest.param("", id="empty-string"),
|
||||
pytest.param("x" * 5120, id="5kb-string"),
|
||||
pytest.param({"web_search_results": "srvtoolu_web"}, id="results-not-a-list"),
|
||||
pytest.param({"web_search_results": ["srvtoolu_web", 1]}, id="results-without-objects"),
|
||||
pytest.param(_OTHER_CALL_RESULT, id="result-for-another-call"),
|
||||
),
|
||||
)
|
||||
def test_malformed_provider_specific_fields_count_the_server_tool_call_mark(
|
||||
gateway: Gateway, fields: JsonValue
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
calls: Final[list[JsonValue]] = [*marked_calls(cities=CITIES[:2]), _SERVER_CALL]
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, _, text = post_chat(
|
||||
gateway, chat_body(model, conversation(marker, calls, assistant={"provider_specific_fields": fields}))
|
||||
)
|
||||
assert status == 200, text
|
||||
labels: Final = anthropic_labels(_only_request(wire))
|
||||
assert labels == _SERVER_MARK_COUNTED
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("identity", "forwarded"),
|
||||
(
|
||||
pytest.param("", "tool_use_id", id="empty-string"),
|
||||
pytest.param("x" * 5120, "x" * 5120, id="5kb-string"),
|
||||
),
|
||||
)
|
||||
def test_odd_tool_call_ids_still_count_their_marks(gateway: Gateway, identity: str, forwarded: str) -> None:
|
||||
marker: Final = new_marker()
|
||||
calls: Final[list[JsonValue]] = [
|
||||
*marked_calls(cities=CITIES[:2]),
|
||||
tool_call(CITIES[2], cache_control=EPHEMERAL, id=identity),
|
||||
]
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls)))
|
||||
assert status == 200, text
|
||||
labels: Final = anthropic_labels(_only_request(wire))
|
||||
assert labels == [
|
||||
ASK_LABEL,
|
||||
tool_use_label(CITIES[0]),
|
||||
tool_use_label(CITIES[1]),
|
||||
f"assistant:tool_use:{forwarded}",
|
||||
]
|
||||
|
||||
|
||||
def _without_id(call: dict[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return {key: value for key, value in call.items() if key != "id"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"odd_call",
|
||||
(
|
||||
pytest.param(tool_call(CITIES[2], cache_control=EPHEMERAL, id=123), id="int"),
|
||||
pytest.param(tool_call(CITIES[2], cache_control=EPHEMERAL, id=["call_x"]), id="list"),
|
||||
pytest.param(tool_call(CITIES[2], cache_control=EPHEMERAL, id=None), id="null"),
|
||||
pytest.param(_without_id(tool_call(CITIES[2], cache_control=EPHEMERAL)), id="missing"),
|
||||
),
|
||||
)
|
||||
def test_non_string_tool_call_ids_answer_400_without_an_upstream_call(
|
||||
gateway: Gateway, odd_call: dict[str, JsonValue]
|
||||
) -> None:
|
||||
marker: Final = new_marker()
|
||||
calls: Final[list[JsonValue]] = [*marked_calls(cities=CITIES[:2]), odd_call]
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls)))
|
||||
received: Final = wire.drain()
|
||||
assert status == 400, text
|
||||
assert '"error"' in text, text
|
||||
assert len(received) == 0, [request.target for request in received]
|
||||
|
||||
|
||||
def _unauthorized_peer(request: Request) -> Reply:
|
||||
assert len(anthropic_labels(request)) <= 4, anthropic_labels(request)
|
||||
return Reply(
|
||||
status=401,
|
||||
body=json.dumps(
|
||||
{"type": "error", "error": {"type": "authentication_error", "message": "invalid x-api-key"}}
|
||||
).encode(),
|
||||
)
|
||||
|
||||
|
||||
def test_provider_error_on_a_capped_request_reaches_the_caller(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(_unauthorized_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
response: Final = gateway.request(
|
||||
"POST", "/v1/chat/completions", chat_body(model, conversation(marker, marked_calls()))
|
||||
)
|
||||
received: Final = wire.drain()
|
||||
assert response.status_code == 401, response.text
|
||||
assert "invalid x-api-key" in response.text, response.text
|
||||
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("points", (pytest.param(None, id="null"), pytest.param([], id="empty")))
|
||||
def test_points_null_or_empty_forward_only_the_client_marks(gateway: Gateway, points: JsonValue) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=points)
|
||||
status, response_id, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls())))
|
||||
assert status == 200, text
|
||||
assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED
|
||||
assert gateway_injected(response_id) is False
|
||||
|
||||
|
||||
def _cache_hit_rows(response_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
"""SELECT request_id, metadata->>'litellm_gateway_injected_cache' AS injected FROM "LiteLLM_SpendLogs" """
|
||||
"WHERE request_id LIKE %s",
|
||||
(f"{response_id}_cache_hit%",),
|
||||
)
|
||||
|
||||
|
||||
def test_response_cache_hit_records_the_injection_on_the_first_row_only(gateway: Gateway) -> None:
|
||||
marker: Final = new_marker()
|
||||
unmarked: Final = [tool_call(city) for city in CITIES]
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
body: Final = chat_body(model, conversation(marker, unmarked, ask_marked=False))
|
||||
first: Final = post_chat(gateway, body)
|
||||
second: Final = post_chat(gateway, body)
|
||||
received: Final = wire.drain()
|
||||
assert (first[0], second[0]) == (200, 200), (first[2], second[2])
|
||||
assert first[1] == second[1], (first[2], second[2])
|
||||
assert [anthropic_labels(request) for request in received] == [[SYSTEM_LABEL, final_label(marker)]]
|
||||
assert gateway_injected(first[1]) is True
|
||||
hit_rows: Final = eventually(lambda: _cache_hit_rows(first[1]), lambda rows: len(rows) == 1, seconds=70)
|
||||
assert hit_rows[0]["injected"] is None, hit_rows
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ("/v1/messages", "/v1/responses"))
|
||||
def test_response_cache_twins_on_messages_and_responses_stay_within_the_cap(gateway: Gateway, path: str) -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
body: Final = (
|
||||
messages_body(model, marker, stream=False) if path == "/v1/messages" else responses_body(model, marker)
|
||||
)
|
||||
responses: Final = tuple(gateway.request("POST", path, body) for _ in range(2))
|
||||
received: Final = wire.drain()
|
||||
assert [response.status_code for response in responses] == [200, 200], [response.text for response in responses]
|
||||
expected: Final = _CLIENT_MARKED if path == "/v1/messages" else [SYSTEM_LABEL, ASK_LABEL]
|
||||
assert [anthropic_labels(request) for request in received] == [expected] * len(received)
|
||||
assert len(received) == 1, [request.target for request in received]
|
||||
|
||||
|
||||
def _model_id(gateway: Gateway, model: str) -> str:
|
||||
entries: Final = gateway.get("/model/info")["data"]
|
||||
assert isinstance(entries, list), entries
|
||||
matches: Final = [
|
||||
string_value(object_value(object_value(entry)["model_info"])["id"])
|
||||
for entry in entries
|
||||
if object_value(entry)["model_name"] == model
|
||||
]
|
||||
assert len(matches) == 1, matches
|
||||
return matches[0]
|
||||
|
||||
|
||||
_ASSISTANT_POINT: Final[list[JsonValue]] = [{"location": "message", "role": "assistant"}]
|
||||
_UPDATE_BURST: Final = 20
|
||||
|
||||
|
||||
async def _capped_burst(gateway: Gateway, model: str) -> tuple[tuple[int, str], ...]:
|
||||
async def one(client: httpx.AsyncClient) -> tuple[int, str]:
|
||||
body: Final = chat_body(model, conversation(new_marker(), marked_calls()))
|
||||
response: Final = await client.post(
|
||||
"/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {gateway.key}"}
|
||||
)
|
||||
return response.status_code, response.text
|
||||
|
||||
async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client:
|
||||
return tuple(await asyncio.gather(*(one(client) for _ in range(_UPDATE_BURST))))
|
||||
|
||||
|
||||
def _probe_assistant_point(gateway: Gateway, wire: Wire, model: str, marker: str) -> list[str]:
|
||||
messages: Final[list[JsonValue]] = [
|
||||
{"role": "system", "content": SYSTEM},
|
||||
{"role": "user", "content": ASK},
|
||||
{"role": "assistant", "content": "I will check."},
|
||||
{"role": "user", "content": final_text(marker)},
|
||||
]
|
||||
status, _, text = post_chat(gateway, chat_body(model, messages))
|
||||
assert status == 200, text
|
||||
return anthropic_labels(wire.drain()[-1])
|
||||
|
||||
|
||||
async def test_points_update_mid_burst_keeps_every_request_within_the_cap(gateway: Gateway) -> None:
|
||||
with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS)
|
||||
identity: Final = _model_id(gateway, model)
|
||||
burst: Final = asyncio.create_task(_capped_burst(gateway, model))
|
||||
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 3, 30)
|
||||
await asyncio.to_thread(
|
||||
gateway.post,
|
||||
"/model/update",
|
||||
{
|
||||
"model_name": model,
|
||||
"litellm_params": {
|
||||
"model": f"anthropic/{ANTHROPIC_MODEL}",
|
||||
"api_base": wire.url,
|
||||
"api_key": PROVIDER_KEY,
|
||||
"cache_control_injection_points": _ASSISTANT_POINT,
|
||||
},
|
||||
"model_info": {"id": identity},
|
||||
},
|
||||
)
|
||||
sent: Final = await burst
|
||||
received: Final = wire.drain()
|
||||
assert [status for status, _ in sent] == [200] * _UPDATE_BURST, [text for _, text in sent]
|
||||
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED] * _UPDATE_BURST
|
||||
landed: Final = await asyncio.to_thread(
|
||||
eventually,
|
||||
lambda: _probe_assistant_point(gateway, wire, model, new_marker()),
|
||||
lambda labels: labels == ["assistant:text:I will check."],
|
||||
60,
|
||||
)
|
||||
assert landed == ["assistant:text:I will check."]
|
||||
|
|
@ -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
|
||||
|
|
@ -0,0 +1,80 @@
|
|||
import json
|
||||
from typing import Final
|
||||
|
||||
from integration._support.wire import Wire, wire_server
|
||||
from integration.providers._cache_control_marks_support import (
|
||||
ANTHROPIC_MODEL,
|
||||
CITIES,
|
||||
EPHEMERAL,
|
||||
POINTS,
|
||||
PROVIDER_KEY,
|
||||
SYSTEM,
|
||||
TOOL,
|
||||
anthropic_labels,
|
||||
anthropic_peer,
|
||||
ask,
|
||||
call_id,
|
||||
client_marked,
|
||||
final_text,
|
||||
new_marker,
|
||||
)
|
||||
from pydantic import JsonValue
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Message, ModelResponse
|
||||
|
||||
_CLIENT_MARKED: Final = client_marked()
|
||||
|
||||
|
||||
def _marked_call(city: str) -> ChatCompletionMessageToolCall:
|
||||
return ChatCompletionMessageToolCall(
|
||||
id=call_id(city),
|
||||
type="function",
|
||||
function={"name": "lookup_weather", "arguments": json.dumps({"city": city})},
|
||||
cache_control=EPHEMERAL,
|
||||
)
|
||||
|
||||
|
||||
def _pydantic_conversation(marker: str) -> list[JsonValue | Message]:
|
||||
return [
|
||||
{"role": "system", "content": SYSTEM},
|
||||
ask(marked=True),
|
||||
Message(role="assistant", content="", tool_calls=[_marked_call(city) for city in CITIES]),
|
||||
*({"role": "tool", "tool_call_id": call_id(city), "content": "sunny."} for city in CITIES),
|
||||
{"role": "user", "content": final_text(marker)},
|
||||
]
|
||||
|
||||
|
||||
def _completion_kwargs(wire: Wire, marker: str) -> dict[str, object]:
|
||||
return {
|
||||
"model": f"anthropic/{ANTHROPIC_MODEL}",
|
||||
"messages": _pydantic_conversation(marker),
|
||||
"tools": [TOOL],
|
||||
"max_tokens": 64,
|
||||
"api_base": wire.url,
|
||||
"api_key": PROVIDER_KEY,
|
||||
"num_retries": 0,
|
||||
"cache_control_injection_points": POINTS,
|
||||
}
|
||||
|
||||
|
||||
def test_pydantic_tool_call_marks_count_against_the_cap_in_process() -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire:
|
||||
response: Final = litellm.completion(**_completion_kwargs(wire, marker)) # pyright: ignore[reportArgumentType] # Message objects stand in for the dict shapes the signature names
|
||||
assert isinstance(response, ModelResponse), type(response)
|
||||
assert response.choices[0].message.content == "sunny", response # pyright: ignore[reportAttributeAccessIssue] # choices are Choices for a non-stream response
|
||||
received: Final = wire.drain()
|
||||
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED]
|
||||
|
||||
|
||||
async def test_pydantic_tool_call_marks_count_against_the_cap_in_process_async_stream() -> None:
|
||||
marker: Final = new_marker()
|
||||
with wire_server(anthropic_peer) as wire:
|
||||
stream: Final = await litellm.acompletion(**_completion_kwargs(wire, marker), stream=True) # pyright: ignore[reportArgumentType] # Message objects stand in for the dict shapes the signature names
|
||||
assert isinstance(stream, CustomStreamWrapper), type(stream)
|
||||
chunks: Final = [chunk async for chunk in stream]
|
||||
assert "".join(str(chunk.choices[0].delta.content or "") for chunk in chunks) == "sunny", chunks
|
||||
received: Final = wire.drain()
|
||||
assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED]
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -4121,3 +4121,30 @@ def test_tool_changes_beta_requires_system_tool_reference(role: str, content: ob
|
|||
)
|
||||
|
||||
assert ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER not in headers.get("anthropic-beta", "").split(",")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("tool_call_id", "provider_specific_fields", "rebuilt"),
|
||||
(
|
||||
("srvtoolu_search", {"web_search_results": [{"tool_use_id": "srvtoolu_search"}]}, True),
|
||||
("srvtoolu_code", {"tool_results": [{"tool_use_id": "srvtoolu_code"}]}, True),
|
||||
(
|
||||
"srvtoolu_code",
|
||||
{"web_search_results": "srvtoolu_code", "tool_results": [{"tool_use_id": "srvtoolu_code"}]},
|
||||
True,
|
||||
),
|
||||
("srvtoolu_other", {"web_search_results": [{"tool_use_id": "srvtoolu_search"}]}, False),
|
||||
("call_client", {"tool_results": [{"tool_use_id": "call_client"}]}, False),
|
||||
("srvtoolu_search", {"web_search_results": ["srvtoolu_search"]}, False),
|
||||
("srvtoolu_search", {}, False),
|
||||
("srvtoolu_search", None, False),
|
||||
("srvtoolu_search", [{"tool_use_id": "srvtoolu_search"}], False),
|
||||
(None, {"web_search_results": [{"tool_use_id": None}]}, False),
|
||||
),
|
||||
)
|
||||
def test_tool_call_is_rebuilt_as_server_tool_use_only_with_a_stored_result(
|
||||
tool_call_id: object, provider_specific_fields: object, rebuilt: bool
|
||||
) -> None:
|
||||
from litellm.llms.anthropic.common_utils import tool_call_is_rebuilt_as_server_tool_use
|
||||
|
||||
assert tool_call_is_rebuilt_as_server_tool_use(tool_call_id, provider_specific_fields) is rebuilt
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue