From 39a69561db81f164320c0bc09ada3ec525aaeece Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 28 Sep 2026 11:12:24 +0000 Subject: [PATCH 1/6] 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> --- .../anthropic_cache_control_hook.py | 69 +++++- ..._cache_control_injection_tool_calls_e2e.py | 199 ++++++++++++++++++ tests/e2e/models.py | 8 + .../test_anthropic_cache_control_hook.py | 146 ++++++++++++- 4 files changed, 412 insertions(+), 10 deletions(-) create mode 100644 tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 1db144b5fdc..96d48b79edf 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -122,7 +122,57 @@ 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 cast(object, getattr(value, key)) + if isinstance(value, Mapping): + value_mapping: Final = cast(Mapping[str, object], value) + return value_mapping.get(key) + return None + + +def _as_object_list(value: object | None) -> list[object] | None: + if not isinstance(value, list): + return None + return cast(list[object], value) + + +def _as_object_iterable(value: object | None) -> Iterable[object] | None: + if not isinstance(value, Iterable): + return None + return cast(Iterable[object], value) + + +def _has_server_tool_result(tool_call_id: str, results: Iterable[object] | None) -> bool: + return any(isinstance(result, dict) and result.get("tool_use_id") == tool_call_id for result in results or ()) + + +def _tool_call_cache_control_is_forwarded(tool_call: object, message: object) -> bool: + if _attribute_or_key(tool_call, "type") != "function" or not isinstance( + _attribute_or_key(tool_call, "cache_control"), dict + ): + return False + + tool_call_id: Final = _attribute_or_key(tool_call, "id") + if not isinstance(tool_call_id, str) or not tool_call_id.startswith("srvtoolu_"): + return True + + provider_specific_fields: Final = _attribute_or_key(message, "provider_specific_fields") + if not isinstance(provider_specific_fields, dict): + return True + + provider_fields_mapping: Final = cast(Mapping[str, object], provider_specific_fields) + server_tool_result_keys: Final = ("web_search_results", "tool_results") + return not any( + _has_server_tool_result( + tool_call_id, + _as_object_iterable(provider_fields_mapping.get(result_key)), + ) + for result_key in server_tool_result_keys + ) def _tool_carries_cache_breakpoint(tool: object) -> bool: @@ -471,13 +521,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_cache_control_is_forwarded(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: diff --git a/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py b/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py new file mode 100644 index 00000000000..4c06c7469eb --- /dev/null +++ b/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py @@ -0,0 +1,199 @@ +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-4-6" +VERTEX_LOCATION: Final[str] = "us-east5" + + +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 == "stop", 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) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 6e66529ec8f..d73c7a538d7 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -294,6 +294,7 @@ class ToolCall(BaseModel): id: str | None = None type: str | None = None function: ToolCallFunction = ToolCallFunction() + cache_control: CacheControl | None = None class ChatAssistantTurn(BaseModel): @@ -1203,6 +1204,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. @@ -1261,6 +1268,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 diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index 1d70a21af7b..de2836a1a79 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -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 @@ -16,7 +17,8 @@ from litellm.integrations.anthropic_cache_control_hook import ( supports_openai_prompt_cache_breakpoint, ) 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 +1047,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 +1092,116 @@ 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_only_tool_call_marks_forwarded_to_anthropic(): + 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]) == 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_cache_control_hook_caps_at_four_blocks_with_client_cache_control(): """Regression for LIT-3667 / Anthropic 'A maximum of 4 blocks ... Found 5'. From 8f05215cee03c84dc861a6efc8709242ce754bc1 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 28 Sep 2026 11:52:12 +0000 Subject: [PATCH 2/6] fix(caching): remove cache census casts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/anthropic_cache_control_hook.py | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 96d48b79edf..0b1343d196b 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -127,23 +127,22 @@ def _carries_cache_breakpoint(block: object) -> bool: def _attribute_or_key(value: object, key: str) -> object | None: if hasattr(value, key): - return cast(object, getattr(value, key)) + return getattr(value, key) if isinstance(value, Mapping): - value_mapping: Final = cast(Mapping[str, object], value) - return value_mapping.get(key) + return value.get(key) return None def _as_object_list(value: object | None) -> list[object] | None: if not isinstance(value, list): return None - return cast(list[object], value) + return _validated_object_list(value) def _as_object_iterable(value: object | None) -> Iterable[object] | None: if not isinstance(value, Iterable): return None - return cast(Iterable[object], value) + return value def _has_server_tool_result(tool_call_id: str, results: Iterable[object] | None) -> bool: @@ -164,12 +163,11 @@ def _tool_call_cache_control_is_forwarded(tool_call: object, message: object) -> if not isinstance(provider_specific_fields, dict): return True - provider_fields_mapping: Final = cast(Mapping[str, object], provider_specific_fields) server_tool_result_keys: Final = ("web_search_results", "tool_results") return not any( _has_server_tool_result( tool_call_id, - _as_object_iterable(provider_fields_mapping.get(result_key)), + _as_object_iterable(provider_specific_fields.get(result_key)), ) for result_key in server_tool_result_keys ) From 9f08d8aef84dd7b37960669c48113c38a079c6be Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 28 Sep 2026 22:55:39 +0000 Subject: [PATCH 3/6] fix(caching): only skip injection on message or content marks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../anthropic_cache_control_hook.py | 12 ++++--- .../test_anthropic_cache_control_hook.py | 36 +++++++++++++++++++ 2 files changed, 44 insertions(+), 4 deletions(-) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 0b1343d196b..d7b5619c2d9 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -518,22 +518,26 @@ class AnthropicCacheControlHook(CustomPromptManagement): return [] @staticmethod - def _count_cache_control_blocks(message: object) -> int: + def _count_message_and_content_breakpoints(message: object) -> int: 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 + return message_count + content_count + + @staticmethod + def _count_cache_control_blocks(message: object) -> int: tool_calls: Final = _as_object_list(_attribute_or_key(message, "tool_calls")) tool_call_count: Final = ( sum(1 for tool_call in tool_calls if _tool_call_cache_control_is_forwarded(tool_call, message)) if tool_calls else 0 ) - return message_count + content_count + tool_call_count + return AnthropicCacheControlHook._count_message_and_content_breakpoints(message) + tool_call_count @staticmethod def _message_has_cache_control(message: AllMessageValues) -> bool: - """Return True if the message already carries any cache_control.""" - return AnthropicCacheControlHook._count_cache_control_blocks(message) > 0 + """Return True if message-level or content-block cache_control is present.""" + return AnthropicCacheControlHook._count_message_and_content_breakpoints(message) > 0 @staticmethod def _safe_insert_cache_control_in_message( diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index de2836a1a79..181fb810470 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -1202,6 +1202,42 @@ def test_cache_control_hook_counts_pydantic_message_tool_call_marks(): assert AnthropicCacheControlHook.count_request_cache_breakpoints(cast(list[AllMessageValues], [message])) == 1 +def test_injection_adds_message_mark_without_overwriting_tool_call_ttl(): + hook: Final = AnthropicCacheControlHook() + tool_call_ttl: Final = {"type": "ephemeral", "ttl": "1h"} + messages: Final = [ + { + "role": "assistant", + "content": "ok", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + "cache_control": tool_call_ttl, + } + ], + } + ] + + _, 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": [{"location": "message", "index": -1}]}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + + assistant_message: Final = processed[0] + assistant_tool_calls: Final = assistant_message.get("tool_calls") + assert assistant_message.get("cache_control") == {"type": "ephemeral"} + assert isinstance(assistant_tool_calls, list) + tool_call: Final = assistant_tool_calls[0] + assert tool_call.get("cache_control") == tool_call_ttl + assert AnthropicCacheControlHook.count_request_cache_breakpoints(processed) == 2 + + 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'. From 466022fd848247ada2fe1d74443760ec1e1b9a65 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 17:09:06 -0700 Subject: [PATCH 4/6] 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. --- .../anthropic_cache_control_hook.py | 12 +++----- .../test_anthropic_cache_control_hook.py | 28 ++++++++++--------- 2 files changed, 19 insertions(+), 21 deletions(-) diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index d7b5619c2d9..0b1343d196b 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -518,26 +518,22 @@ class AnthropicCacheControlHook(CustomPromptManagement): return [] @staticmethod - def _count_message_and_content_breakpoints(message: object) -> int: + def _count_cache_control_blocks(message: object) -> int: 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 - return message_count + content_count - - @staticmethod - def _count_cache_control_blocks(message: object) -> int: tool_calls: Final = _as_object_list(_attribute_or_key(message, "tool_calls")) tool_call_count: Final = ( sum(1 for tool_call in tool_calls if _tool_call_cache_control_is_forwarded(tool_call, message)) if tool_calls else 0 ) - return AnthropicCacheControlHook._count_message_and_content_breakpoints(message) + tool_call_count + return message_count + content_count + tool_call_count @staticmethod def _message_has_cache_control(message: AllMessageValues) -> bool: - """Return True if message-level or content-block cache_control is present.""" - return AnthropicCacheControlHook._count_message_and_content_breakpoints(message) > 0 + """Return True if the message already carries any cache_control.""" + return AnthropicCacheControlHook._count_cache_control_blocks(message) > 0 @staticmethod def _safe_insert_cache_control_in_message( diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index 181fb810470..5e29ec7d700 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -1202,40 +1202,42 @@ def test_cache_control_hook_counts_pydantic_message_tool_call_marks(): assert AnthropicCacheControlHook.count_request_cache_breakpoints(cast(list[AllMessageValues], [message])) == 1 -def test_injection_adds_message_mark_without_overwriting_tool_call_ttl(): +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 = [ + messages: Final[list[AllMessageValues]] = [ + {"role": "user", "content": "What is the weather in Paris?"}, { "role": "assistant", - "content": "ok", + "content": "Let me look that up.", "tool_calls": [ { - "id": "call_1", + "id": "toolu_01A", "type": "function", - "function": {"name": "lookup", "arguments": "{}"}, + "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="bedrock/us.anthropic.claude-opus-4-6-v1:0", + model="anthropic/claude-haiku-4-5", messages=messages, - non_default_params={"cache_control_injection_points": [{"location": "message", "index": -1}]}, + non_default_params={"cache_control_injection_points": [{"location": "message", "role": "assistant"}]}, prompt_id=None, prompt_variables=None, dynamic_callback_params={}, ) - assistant_message: Final = processed[0] + assistant_message: Final = processed[1] assistant_tool_calls: Final = assistant_message.get("tool_calls") - assert assistant_message.get("cache_control") == {"type": "ephemeral"} + assert assistant_message.get("cache_control") is None + assert assistant_message.get("content") == "Let me look that up." assert isinstance(assistant_tool_calls, list) - tool_call: Final = assistant_tool_calls[0] - assert tool_call.get("cache_control") == tool_call_ttl - assert AnthropicCacheControlHook.count_request_cache_breakpoints(processed) == 2 + 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(): From e341ab46491e729fbb560ec4fad2450580d642d2 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:43:01 -0700 Subject: [PATCH 5/6] 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 --- .../anthropic_cache_control_hook.py | 8 +- ..._cache_control_injection_tool_calls_e2e.py | 8 +- .../providers/_cache_control_marks_support.py | 461 +++++++++++++++ ...che_control_tool_call_marks_owned_proxy.py | 245 ++++++++ ...test_cache_control_tool_call_marks_wire.py | 525 ++++++++++++++++++ ...prompt_caching_affinity_tool_call_marks.py | 93 ++++ .../test_anthropic_cache_control_hook.py | 72 ++- 7 files changed, 1402 insertions(+), 10 deletions(-) create mode 100644 tests/integration/providers/_cache_control_marks_support.py create mode 100644 tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py create mode 100644 tests/integration/providers/test_cache_control_tool_call_marks_wire.py create mode 100644 tests/integration/routing/test_prompt_caching_affinity_tool_call_marks.py diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 0b1343d196b..d103de4bd21 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -149,10 +149,8 @@ def _has_server_tool_result(tool_call_id: str, results: Iterable[object] | None) return any(isinstance(result, dict) and result.get("tool_use_id") == tool_call_id for result in results or ()) -def _tool_call_cache_control_is_forwarded(tool_call: object, message: object) -> bool: - if _attribute_or_key(tool_call, "type") != "function" or not isinstance( - _attribute_or_key(tool_call, "cache_control"), dict - ): +def _tool_call_carries_cache_breakpoint(tool_call: object, message: object) -> bool: + if _attribute_or_key(tool_call, "cache_control") is None: return False tool_call_id: Final = _attribute_or_key(tool_call, "id") @@ -524,7 +522,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): content_count: Final = sum(1 for block in content if _carries_cache_breakpoint(block)) if content else 0 tool_calls: Final = _as_object_list(_attribute_or_key(message, "tool_calls")) tool_call_count: Final = ( - sum(1 for tool_call in tool_calls if _tool_call_cache_control_is_forwarded(tool_call, message)) + sum(1 for tool_call in tool_calls if _tool_call_carries_cache_breakpoint(tool_call, message)) if tool_calls else 0 ) diff --git a/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py b/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py index 4c06c7469eb..1481185a601 100644 --- a/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py +++ b/tests/e2e/llm_translation/test_cache_control_injection_tool_calls_e2e.py @@ -27,8 +27,8 @@ pytestmark: Final = pytest.mark.e2e Backend: TypeAlias = Literal["azure_foundry", "vertex"] AZURE_MODEL: Final[str] = "azure_ai/claude-haiku-4-5" -VERTEX_MODEL: Final[str] = "vertex_ai/claude-sonnet-4-6" -VERTEX_LOCATION: Final[str] = "us-east5" +VERTEX_MODEL: Final[str] = "vertex_ai/claude-sonnet-5" +VERTEX_LOCATION: Final[str] = "global" def _deployment_params(*, backend: Backend, inject_cache_control: bool) -> LiteLLMParamsBody: @@ -150,7 +150,9 @@ def _post_chat(client: PassthroughClient, key: str, body: ChatBody) -> Result[Ch def _assert_normal_completion(response: ChatResponse, model_name: str) -> None: assert response.choices, f"{model_name}: chat completion returned no choices: {response}" completion: Final = response.choices[0] - assert completion.finish_reason == "stop", f"{model_name}: unexpected finish reason: {completion.finish_reason}" + assert completion.finish_reason in ("stop", "length"), ( + f"{model_name}: unexpected finish reason: {completion.finish_reason}" + ) assert ( completion.message is not None and completion.message.content is not None diff --git a/tests/integration/providers/_cache_control_marks_support.py b/tests/integration/providers/_cache_control_marks_support.py new file mode 100644 index 00000000000..5194f7d7ca3 --- /dev/null +++ b/tests/integration/providers/_cache_control_marks_support.py @@ -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 diff --git a/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py b/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py new file mode 100644 index 00000000000..46d6e9a40cb --- /dev/null +++ b/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py @@ -0,0 +1,245 @@ +import asyncio +import re +import signal +import socket +from collections import Counter +from collections.abc import Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.wire import Request, wire_server +from integration.providers._cache_control_marks_support import ( + ASK_LABEL, + CITIES, + POINTS, + SYSTEM_LABEL, + anthropic_deployment, + anthropic_labels, + anthropic_peer, + chat_body, + client_marked, + conversation, + final_label, + marked_calls, + marker_of, + messages_body, + new_marker, + owned_config, + responses_body, + tool_call, + tool_use_label, +) +from pydantic import JsonValue + +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_SURFACES: Final = ("chat", "chat-stream", "chat-unmarked", "messages", "responses") +_MODEL: Final = "capped-claude" +_BURST: Final = 30 +_OUTAGE_STATUS: Final = 500 + + +@dataclass(frozen=True, slots=True) +class _Sent: + surface: str + marker: str + status: int + text: str + call_id: str + client_port: int + + +def _free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return int(reserve.getsockname()[1]) + + +def _surface_request(surface: str, marker: str) -> tuple[str, dict[str, JsonValue]]: + if surface == "messages": + return "/v1/messages", messages_body(_MODEL, marker, stream=False) + if surface == "responses": + return "/v1/responses", responses_body(_MODEL, marker) + if surface == "chat-unmarked": + unmarked: Final = [tool_call(city) for city in CITIES] + return "/v1/chat/completions", chat_body(_MODEL, conversation(marker, unmarked, ask_marked=False)) + return "/v1/chat/completions", chat_body( + _MODEL, conversation(marker, marked_calls()), stream=surface == "chat-stream" + ) + + +def _expected_labels(item: _Sent) -> list[str]: + if item.surface == "responses": + return [SYSTEM_LABEL, ASK_LABEL] + if item.surface == "chat-unmarked": + return [SYSTEM_LABEL, final_label(item.marker)] + return client_marked() + + +async def _fire(owned_url: str, key: str, *, tolerate_transport_errors: bool = False) -> tuple[_Sent, ...]: + async def one(client: httpx.AsyncClient, index: int) -> _Sent: + surface: Final = _SURFACES[index % len(_SURFACES)] + marker: Final = new_marker() + path, body = _surface_request(surface, marker) + async with client.stream("POST", path, json=body, headers={"Authorization": f"Bearer {key}"}) as response: + client_port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1]) + await response.aread() + return _Sent( + surface, + marker, + response.status_code, + response.text, + response.headers.get("x-litellm-call-id", ""), + client_port, + ) + + async with httpx.AsyncClient(base_url=owned_url, timeout=30, trust_env=False) as client: + results: Final = await asyncio.gather( + *(one(client, index) for index in range(_BURST)), return_exceptions=tolerate_transport_errors + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Sent)) + + +def _by_marker(received: Sequence[Request]) -> dict[str, tuple[Request, ...]]: + counted: Final = Counter(marker_of(request) for request in received) + return {marker: tuple(request for request in received if marker_of(request) == marker) for marker in counted} + + +def _assert_capped(served: Sequence[_Sent], received: Sequence[Request]) -> None: + by_marker: Final = _by_marker(received) + for item in served: + assert item.status == 200, (item.surface, item.text) + assert len(by_marker.get(item.marker, ())) == 1, (item.surface, item.marker) + assert anthropic_labels(by_marker[item.marker][0]) == _expected_labels(item), item.surface + + +def _assert_outage_error(item: _Sent, received: Sequence[Request]) -> None: + assert item.status == _OUTAGE_STATUS, (item.surface, item.status, item.text) + assert '"error"' in item.text, (item.surface, item.text) + assert item.marker not in {marker_of(request) for request in received}, item.surface + + +def _spend_rows(call_id: str) -> list[dict[str, JsonValue]]: + return read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', (call_id,)) + + +def _single_spend_row(item: _Sent) -> None: + assert item.call_id, (item.surface, item.status, item.text) + rows: Final = eventually(lambda: _spend_rows(item.call_id), lambda values: len(values) == 1, seconds=70) + assert len(rows) == 1, (item.surface, item.call_id) + + +@pytest.mark.timeout(240) +async def test_capped_burst_rides_out_a_provider_outage_and_logs_every_request_once( + gateway: Gateway, tmp_path: Path +) -> None: + port: Final = _free_port() + config: Final = owned_config( + tmp_path, [anthropic_deployment(_MODEL, f"http://127.0.0.1:{port}", cache_control_injection_points=POINTS)] + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + owned_url: Final = str(owned.gateway.client.base_url) + with wire_server(anthropic_peer, port=port) as wire: + healthy: Final = await _fire(owned_url, owned.gateway.key) + _assert_capped(healthy, wire.drain()) + racing: Final = asyncio.create_task(_fire(owned_url, owned.gateway.key)) + await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30) + down: Final = await _fire(owned_url, owned.gateway.key) + raced: Final = await racing + raced_received: Final = wire.drain() + with wire_server(anthropic_peer, port=port) as restarted: + recovered: Final = await _fire(owned_url, owned.gateway.key) + recovered_received: Final = restarted.drain() + assert len(_STARTED_WORKER.findall(owned.log.read_text())) >= 2 + assert all(len(requests) == 1 for requests in _by_marker(raced_received).values()) + _assert_capped(tuple(item for item in raced if item.status == 200), raced_received) + for item in raced: + if item.status != 200: + _assert_outage_error(item, raced_received) + for item in down: + _assert_outage_error(item, (*raced_received, *recovered_received)) + _assert_capped(recovered, recovered_received) + assert {marker_of(request) for request in recovered_received} == {item.marker for item in recovered} + for item in (*healthy, *raced, *down, *recovered): + _single_spend_row(item) + + +@pytest.mark.timeout(240) +async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_capped_requests( + gateway: Gateway, tmp_path: Path +) -> None: + with wire_server(anthropic_peer) as wire: + config: Final = owned_config( + tmp_path, [anthropic_deployment(_MODEL, wire.url, cache_control_injection_points=POINTS)] + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + owned_url: Final = str(owned.gateway.client.base_url) + burst: Final = asyncio.create_task(_fire(owned_url, owned.gateway.key, tolerate_transport_errors=True)) + await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30) + victim: Final = psutil.Process(workers[0]) + victim.suspend() + victim_ports: Final = frozenset( + connection.raddr.port for connection in victim.net_connections(kind="tcp") if connection.raddr + ) + victim.send_signal(signal.SIGKILL) + served: Final = await burst + during: Final = wire.drain() + after: Final = await _fire(owned_url, owned.gateway.key) + after_received: Final = wire.drain() + assert all(len(requests) == 1 for requests in _by_marker(during).values()) + _assert_capped(tuple(item for item in served if item.status == 200), during) + _assert_capped(after, after_received) + survivors: Final = tuple(item for item in served if item.client_port not in victim_ports) + assert survivors, [item.client_port for item in served] + for item in (*survivors, *after): + _single_spend_row(item) + for item in served: + assert len(_spend_rows(item.call_id)) <= 1, item.call_id + + +@pytest.mark.timeout(240) +def test_yaml_auto_caching_stands_down_for_tool_call_marks_and_outranks_a_key_opt_out( + gateway: Gateway, tmp_path: Path +) -> None: + markers: Final = (new_marker(), new_marker(), new_marker()) + unmarked: Final = [tool_call(city) for city in CITIES] + with wire_server(anthropic_peer) as wire: + config: Final = owned_config( + tmp_path, + [anthropic_deployment(_MODEL, wire.url)], + litellm_settings={"enable_anthropic_prompt_caching": True}, + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + with owned.gateway.scenario() as scenario: + opted_out: Final = scenario.key(metadata={"enable_prompt_caching": False}) + cells: Final = ( + (conversation(markers[0], marked_calls(), ask_marked=False), owned.gateway.key), + (conversation(markers[1], unmarked, ask_marked=False), owned.gateway.key), + (conversation(markers[2], unmarked, ask_marked=False), opted_out), + ) + responses: Final = tuple( + owned.gateway.request("POST", "/v1/chat/completions", chat_body(_MODEL, messages), key=key) + for messages, key in cells + ) + received: Final = _by_marker(wire.drain()) + assert [response.status_code for response in responses] == [200, 200, 200], [ + response.text for response in responses + ] + assert [anthropic_labels(received[marker][0]) for marker in markers] == [ + [tool_use_label(city) for city in CITIES], + [SYSTEM_LABEL, final_label(markers[1])], + [SYSTEM_LABEL, final_label(markers[2])], + ] diff --git a/tests/integration/providers/test_cache_control_tool_call_marks_wire.py b/tests/integration/providers/test_cache_control_tool_call_marks_wire.py new file mode 100644 index 00000000000..76c00536107 --- /dev/null +++ b/tests/integration/providers/test_cache_control_tool_call_marks_wire.py @@ -0,0 +1,525 @@ +import json +from datetime import datetime, timedelta, timezone +from typing import Final + +import pytest +from integration._support.client import Gateway, Scenario, eventually, object_value +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.providers._cache_control_marks_support import ( + ANTHROPIC_MODEL, + ASK, + ASK_LABEL, + BEDROCK_MODEL, + CITIES, + EPHEMERAL, + POINTS, + PROVIDER_KEY, + SYSTEM, + SYSTEM_LABEL, + anthropic_labels, + anthropic_marks, + anthropic_peer, + bedrock_labels, + bedrock_peer, + call_id, + chat_body, + client_marked, + conversation, + final_label, + final_text, + gateway_injected, + marked_calls, + messages_body, + new_marker, + post_chat, + responses_body, + tool_call, + tool_use_label, +) +from litellm.utils import get_prompt_cache_min_tokens +from pydantic import JsonValue, TypeAdapter + +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_CLIENT_MARKED: Final = client_marked() +_GEMINI_REPLY: Final = json.dumps( + { + "candidates": [{"content": {"role": "model", "parts": [{"text": "sunny"}]}, "finishReason": "STOP"}], + "usageMetadata": {"promptTokenCount": 1600, "candidatesTokenCount": 1, "totalTokenCount": 1601}, + } +).encode() +_OPENAI_REPLY: Final = json.dumps( + { + "id": "chatcmpl-cache-census", + "object": "chat.completion", + "created": 1, + "model": "gpt-6-sol", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "sunny"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 12, "completion_tokens": 1, "total_tokens": 13}, + } +).encode() + + +def _anthropic_deployment(scenario: Scenario, wire: Wire, **fields: JsonValue) -> str: + return scenario.model(model=f"anthropic/{ANTHROPIC_MODEL}", api_base=wire.url, api_key=PROVIDER_KEY, **fields) + + +def _only_request(wire: Wire) -> Request: + received: Final = wire.drain() + assert len(received) == 1, [request.target for request in received] + return received[0] + + +def _stream(gateway: Gateway, path: str, body: dict[str, JsonValue], *, key: str | None = None) -> tuple[int, str]: + headers: Final = {"Authorization": f"Bearer {gateway.key if key is None else key}"} + with gateway.client.stream("POST", path, json=body, headers=headers) as response: + return response.status_code, response.read().decode() + + +def _chat_stream_chunks(text: str) -> list[dict[str, JsonValue]]: + return [ + _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ] + + +def _chat_stream_content(chunks: list[dict[str, JsonValue]]) -> str: + return "".join( + str(object_value(object_value(choice).get("delta") or {}).get("content") or "") + for chunk in chunks + for choice in (chunk.get("choices") if isinstance(chunk.get("choices"), list) else []) + ) + + +@pytest.mark.parametrize("stream", (False, True), ids=("json", "stream")) +def test_chat_points_skip_injection_when_client_marks_fill_the_cap(gateway: Gateway, stream: bool) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + body: Final = chat_body(model, conversation(marker, marked_calls()), stream=stream) + if stream: + status, text = _stream(gateway, "/v1/chat/completions", body) + assert status == 200, text + chunks: Final = _chat_stream_chunks(text) + assert _chat_stream_content(chunks) == "sunny", text + assert "data: [DONE]" in text, text + response_id = str(chunks[0]["id"]) + else: + status, response_id, text = post_chat(gateway, body) + assert status == 200, text + assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED + assert gateway_injected(response_id) is False + + +def test_chat_points_inject_system_and_last_message_without_client_marks(gateway: Gateway) -> None: + marker: Final = new_marker() + unmarked: Final = [tool_call(city) for city in CITIES] + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + status, response_id, text = post_chat( + gateway, chat_body(model, conversation(marker, unmarked, ask_marked=False)) + ) + assert status == 200, text + assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, final_label(marker)] + assert gateway_injected(response_id) is True + + +def _prompt_caching_rows(gateway: Gateway, cursor: dict[str, str]) -> list[dict[str, JsonValue]]: + now: Final = datetime.now(timezone.utc) + page: Final = gateway.get( + "/cost_optimization/prompt_caching/requests", + { + "start_date": (now - timedelta(minutes=10)).isoformat(), + "end_date": (now + timedelta(minutes=10)).isoformat(), + "page_size": "100", + "filter": "injected", + **cursor, + }, + ) + rows: Final = [object_value(row) for row in page["requests"]] if isinstance(page["requests"], list) else [] + following: Final = page.get("next_cursor") + if not page.get("has_more") or not isinstance(following, dict): + return rows + return [ + *rows, + *_prompt_caching_rows( + gateway, + {"cursor_start_time": str(following["start_time"]), "cursor_request_id": str(following["request_id"])}, + ), + ] + + +def _listed_as_injected(gateway: Gateway, response_id: str) -> bool: + return any(row["request_id"] == response_id for row in _prompt_caching_rows(gateway, {})) + + +@pytest.mark.parametrize( + ("calls", "ask_marked", "expected", "injected"), + ( + pytest.param(marked_calls(), False, [tool_use_label(city) for city in CITIES], False, id="three-tool-calls"), + pytest.param([tool_call(city) for city in CITIES], False, None, True, id="no-client-marks"), + pytest.param( + [*marked_calls(cities=CITIES[:2]), tool_call(CITIES[2])], + False, + [tool_use_label(city) for city in CITIES[:2]], + False, + id="two-tool-calls", + ), + ), +) +def test_auto_prompt_caching_stands_down_when_client_marks_only_tool_calls( + gateway: Gateway, + calls: list[JsonValue], + ask_marked: bool, + expected: list[str] | None, + injected: bool, +) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire) + key: Final = scenario.key(metadata={"enable_prompt_caching": True}) + status, response_id, text = post_chat( + gateway, chat_body(model, conversation(marker, calls, ask_marked=ask_marked)), key=key + ) + assert status == 200, text + labels: Final = anthropic_labels(_only_request(wire)) + assert labels == (expected if expected is not None else [SYSTEM_LABEL, final_label(marker)]) + assert gateway_injected(response_id) is injected + assert ( + eventually(lambda: _listed_as_injected(gateway, response_id), lambda listed: listed is injected, 70) is injected + ) + + +def test_assistant_point_skips_message_whose_tool_call_carries_a_one_hour_mark(gateway: Gateway) -> None: + marker: Final = new_marker() + hour: Final[dict[str, JsonValue]] = {"type": "ephemeral", "ttl": "1h"} + calls: Final = [tool_call(CITIES[0], cache_control=hour)] + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment( + scenario, wire, cache_control_injection_points=[{"location": "message", "role": "assistant"}] + ) + status, _, text = post_chat( + gateway, + chat_body(model, conversation(marker, calls, ask_marked=False, assistant={"content": "I will check."})), + ) + assert status == 200, text + marks: Final = anthropic_marks(_only_request(wire)) + assert [(mark.label, mark.ttl) for mark in marks] == [(tool_use_label(CITIES[0]), "1h")] + + +@pytest.mark.parametrize( + "mark", + ( + pytest.param("ephemeral", id="string"), + pytest.param(1, id="int"), + pytest.param(["ephemeral"], id="list"), + pytest.param("", id="empty-string"), + pytest.param("x" * 5120, id="5kb-string"), + ), +) +def test_non_object_tool_call_marks_count_against_the_cap(gateway: Gateway, mark: JsonValue) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls(mark)))) + assert status == 200, text + assert anthropic_labels(_only_request(wire)) == [ASK_LABEL] + + +@pytest.mark.parametrize( + ("calls", "expected"), + ( + pytest.param(marked_calls({}), _CLIENT_MARKED, id="empty-object"), + pytest.param(marked_calls(None), None, id="null"), + pytest.param( + [tool_call(CITIES[0], cache_control=EPHEMERAL)] * 2, + [SYSTEM_LABEL, ASK_LABEL, tool_use_label(CITIES[0])], + id="same-call-twice", + ), + ), +) +def test_tool_call_mark_shapes_keep_the_request_within_the_cap( + gateway: Gateway, calls: list[JsonValue], expected: list[str] | None +) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls))) + assert status == 200, text + labels: Final = anthropic_labels(_only_request(wire)) + assert labels == (expected if expected is not None else [SYSTEM_LABEL, ASK_LABEL, final_label(marker)]) + + +def _untyped_call(city: str) -> dict[str, JsonValue]: + return {name: value for name, value in tool_call(city, cache_control=EPHEMERAL).items() if name != "type"} + + +_BEDROCK_CLIENT_MARKED: Final = [f"user:text:{ASK}", *(f"assistant:toolUse:{call_id(city)}" for city in CITIES)] + + +@pytest.mark.parametrize( + ("calls", "expected"), + ( + pytest.param(marked_calls(), _BEDROCK_CLIENT_MARKED, id="object-marks"), + pytest.param(marked_calls("ephemeral"), _BEDROCK_CLIENT_MARKED, id="string-marks"), + pytest.param([_untyped_call(city) for city in CITIES], _BEDROCK_CLIENT_MARKED, id="calls-without-type"), + pytest.param( + [tool_call(CITIES[0], cache_control=EPHEMERAL)] * 2, + [f"system:text:{SYSTEM}", ASK_LABEL, f"assistant:toolUse:{call_id(CITIES[0])}", "assistant:other"], + id="same-call-twice", + ), + ), +) +def test_bedrock_converse_cache_points_stay_within_the_cap( + gateway: Gateway, calls: list[JsonValue], expected: list[str] +) -> None: + marker: Final = new_marker() + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"bedrock/converse/{BEDROCK_MODEL}", + api_key=PROVIDER_KEY, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint=wire.url, + cache_control_injection_points=POINTS, + ) + status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls))) + assert status == 200, text + request: Final = _only_request(wire) + assert request.target.endswith("/converse"), request.target + assert bedrock_labels(request) == expected + + +_SERVER_CALL: Final = tool_call("web", cache_control=EPHEMERAL) | { + "id": "srvtoolu_web", + "function": {"name": "web_search", "arguments": "{}"}, +} +_WEB_RESULTS: Final[dict[str, JsonValue]] = { + "provider_specific_fields": { + "web_search_results": [{"type": "web_search_tool_result", "tool_use_id": "srvtoolu_web", "content": []}] + } +} + + +@pytest.mark.parametrize( + ("assistant", "expected"), + ( + pytest.param( + _WEB_RESULTS, + [SYSTEM_LABEL, ASK_LABEL, tool_use_label(CITIES[0]), tool_use_label(CITIES[1])], + id="server-tool-with-result", + ), + pytest.param( + {}, + [ASK_LABEL, tool_use_label(CITIES[0]), tool_use_label(CITIES[1]), "assistant:tool_use:srvtoolu_web"], + id="server-tool-without-result", + ), + ), +) +def test_server_tool_call_mark_counts_only_when_forwarded( + gateway: Gateway, assistant: dict[str, JsonValue], expected: list[str] +) -> None: + marker: Final = new_marker() + calls: Final = [*marked_calls(cities=CITIES[:2]), _SERVER_CALL] + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + status, _, text = post_chat(gateway, chat_body(model, conversation(marker, calls, assistant=assistant))) + assert status == 200, text + labels: Final = anthropic_labels(_only_request(wire)) + assert labels == expected + + +def test_azure_ai_claude_points_stay_within_the_cap(gateway: Gateway) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"azure_ai/{ANTHROPIC_MODEL}", + api_base=wire.url, + api_key=PROVIDER_KEY, + cache_control_injection_points=POINTS, + ) + status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls()))) + assert status == 200, text + request: Final = _only_request(wire) + assert request.target == "/anthropic/v1/messages", request.target + assert anthropic_labels(request) == _CLIENT_MARKED + + +def _gemini_peer(request: Request) -> Reply: + if "cachedContents" in request.target and request.method == "GET": + return Reply(body=b'{"cachedContents":[]}') + if "cachedContents" in request.target: + return Reply( + body=json.dumps( + { + "name": "cachedContents/census", + "model": "models/gemini-3.8-flash", + "expireTime": "2099-01-01T00:00:00Z", + } + ).encode() + ) + return Reply(body=_GEMINI_REPLY) + + +_GEMINI: Final = "gemini/gemini-3.8-flash" +_GEMINI_CACHE_WRITE: Final = [ + ("GET", "/models/gemini-3.8-flash:cachedContents"), + ("POST", "/models/gemini-3.8-flash:cachedContents"), + ("POST", "/models/gemini-3.8-flash:generateContent"), +] + + +@pytest.mark.parametrize( + ("calls", "expected"), + ( + pytest.param( + marked_calls(), [("POST", "/models/gemini-3.8-flash:generateContent")], id="tool-calls-fill-the-cap" + ), + pytest.param([tool_call(city) for city in CITIES], _GEMINI_CACHE_WRITE, id="unmarked-tool-calls"), + ), +) +def test_gemini_context_cache_follows_the_cap_census( + gateway: Gateway, calls: list[JsonValue], expected: list[tuple[str, str]] +) -> None: + marker: Final = new_marker() + long_system: Final[JsonValue] = {"role": "system", "content": "lorem " * (2 * get_prompt_cache_min_tokens(_GEMINI))} + messages: Final = [long_system, *conversation(marker, calls)[1:]] + with wire_server(_gemini_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=_GEMINI, + api_base=wire.url, + api_key=PROVIDER_KEY, + cache_control_injection_points=POINTS, + ) + status, _, text = post_chat(gateway, chat_body(model, messages)) + assert status == 200, text + targets: Final = [(request.method, request.target.split("?")[0]) for request in wire.drain()] + assert targets == expected + + +def _openai_marks(request: Request) -> list[str]: + body: Final = _JSON_OBJECT.validate_json(request.body) + messages: Final = body["messages"] if isinstance(body["messages"], list) else [] + return [label for message in messages if isinstance(message, dict) for label in _openai_message_marks(message)] + + +def _openai_message_marks(message: dict[str, JsonValue]) -> list[str]: + role: Final = str(message.get("role")) + content: Final = message.get("content") + calls: Final = message.get("tool_calls") + blocks: Final = content if isinstance(content, list) else [] + return [ + *(f"{key}@{role}:message" for key in ("cache_control", "prompt_cache_breakpoint") if key in message), + *( + f"{key}@{role}:text:{block.get('text')}" + for block in blocks + if isinstance(block, dict) + for key in ("cache_control", "prompt_cache_breakpoint") + if key in block + ), + *( + f"cache_control@tool_call:{call.get('id')}" + for call in (calls if isinstance(calls, list) else []) + if isinstance(call, dict) and "cache_control" in call + ), + ] + + +@pytest.mark.parametrize( + "options", + ( + pytest.param({"prompt_cache_options": {"mode": "explicit"}}, id="breakpoint-dialect"), + pytest.param({}, id="plain"), + ), +) +def test_openai_points_skip_injection_when_client_marks_fill_the_cap( + gateway: Gateway, options: dict[str, JsonValue] +) -> None: + marker: Final = new_marker() + with wire_server(lambda request: Reply(body=_OPENAI_REPLY)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/gpt-6-sol", + api_base=f"{wire.url}/v1", + api_key=PROVIDER_KEY, + cache_control_injection_points=POINTS, + **options, + ) + status, _, text = post_chat(gateway, chat_body(model, conversation(marker, marked_calls()))) + assert status == 200, text + marks: Final = _openai_marks(_only_request(wire)) + assert marks == [f"cache_control@user:text:{ASK}", *(f"cache_control@tool_call:{call_id(city)}" for city in CITIES)] + + +@pytest.mark.parametrize("stream", (False, True), ids=("json", "stream")) +def test_messages_endpoint_points_skip_injection_when_tool_use_marks_fill_the_cap( + gateway: Gateway, stream: bool +) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + status, text = _stream(gateway, "/v1/messages", messages_body(model, marker, stream=stream)) + assert status == 200, text + assert ( + ("event: message_stop" in text) if stream else (_JSON_OBJECT.validate_json(text)["id"] == f"msg_{marker}") + ) + assert anthropic_labels(_only_request(wire)) == _CLIENT_MARKED + + +def test_responses_bridge_keeps_system_and_user_marks_within_the_cap(gateway: Gateway) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + response: Final = gateway.request("POST", "/v1/responses", responses_body(model, marker)) + assert response.status_code == 200, response.text + assert _JSON_OBJECT.validate_json(response.content)["status"] == "completed", response.text + assert anthropic_labels(_only_request(wire)) == [SYSTEM_LABEL, ASK_LABEL] + + +def test_response_cache_serves_the_capped_request_once(gateway: Gateway) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + body: Final = chat_body(model, conversation(marker, marked_calls())) + first: Final = post_chat(gateway, body) + second: Final = post_chat(gateway, body) + received: Final = wire.drain() + assert (first[0], second[0]) == (200, 200), (first[2], second[2]) + assert first[1].startswith("chatcmpl-"), first[2] + assert first[1] == second[1], (first[2], second[2]) + assert [anthropic_labels(request) for request in received] == [_CLIENT_MARKED] + + +def test_unauthenticated_request_never_reaches_the_provider(gateway: Gateway) -> None: + marker: Final = new_marker() + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + status, _, text = post_chat( + gateway, chat_body(model, conversation(marker, marked_calls())), key=f"sk-not-a-key-{marker}" + ) + received: Final = wire.drain() + assert status == 401, text + assert "error" in _JSON_OBJECT.validate_json(text), text + assert received == () + + +@pytest.mark.parametrize( + "calls", + ( + pytest.param({"tool_calls": []}, id="empty"), + pytest.param({"tool_calls": None}, id="null"), + pytest.param({}, id="missing"), + ), +) +def test_assistant_without_tool_calls_keeps_configured_points(gateway: Gateway, calls: dict[str, JsonValue]) -> None: + marker: Final = new_marker() + messages: Final[list[JsonValue]] = [ + {"role": "system", "content": SYSTEM}, + {"role": "user", "content": [{"type": "text", "text": ASK, "cache_control": EPHEMERAL}]}, + {"role": "assistant", "content": "I will check.", **calls}, + {"role": "user", "content": final_text(marker)}, + ] + with wire_server(anthropic_peer) as wire, gateway.scenario() as scenario: + model: Final = _anthropic_deployment(scenario, wire, cache_control_injection_points=POINTS) + status, _, text = post_chat(gateway, chat_body(model, messages)) + assert status == 200, text + labels: Final = anthropic_labels(_only_request(wire)) + assert labels == [SYSTEM_LABEL, ASK_LABEL, final_label(marker)] diff --git a/tests/integration/routing/test_prompt_caching_affinity_tool_call_marks.py b/tests/integration/routing/test_prompt_caching_affinity_tool_call_marks.py new file mode 100644 index 00000000000..9d4bf62adf5 --- /dev/null +++ b/tests/integration/routing/test_prompt_caching_affinity_tool_call_marks.py @@ -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 diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index 5e29ec7d700..4c9fd3cfdf2 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -16,6 +16,10 @@ from litellm.integrations.anthropic_cache_control_hook import ( AnthropicCacheControlHook, supports_openai_prompt_cache_breakpoint, ) +from litellm.litellm_core_utils.prompt_templates.factory import ( + _convert_to_bedrock_tool_call_invoke, + convert_to_anthropic_tool_invoke, +) from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.llms.openai import AllMessageValues, ChatCompletionAssistantToolCall from litellm.types.utils import ChatCompletionMessageToolCall, Message @@ -1102,7 +1106,7 @@ def test_cache_control_hook_counts_tool_call_cache_controls(): assert AnthropicCacheControlHook.count_request_cache_breakpoints([message]) == 3 -def test_cache_control_hook_counts_only_tool_call_marks_forwarded_to_anthropic(): +def test_cache_control_hook_counts_tool_call_marks_except_answered_server_tool_calls(): message: Final[AllMessageValues] = { "role": "assistant", "content": None, @@ -1156,7 +1160,37 @@ def test_cache_control_hook_counts_only_tool_call_marks_forwarded_to_anthropic() }, } - assert AnthropicCacheControlHook.count_request_cache_breakpoints([message]) == 1 + assert AnthropicCacheControlHook.count_request_cache_breakpoints([message]) == 2 + + +@pytest.mark.parametrize("tool_call_type", ["function", "custom", None]) +@pytest.mark.parametrize( + "mark", + [{"type": "ephemeral"}, {}, "ephemeral", "", 7, ["ephemeral"], "x" * 5000, None], + ids=["dict", "empty_dict", "string", "empty_string", "int", "list", "5kb_string", "none"], +) +def test_tool_call_census_matches_the_breakpoints_providers_send(mark: object, tool_call_type: str | None): + tool_call: Final[dict[str, object]] = { + "id": "call_0", + "function": {"name": "lookup", "arguments": "{}"}, + "cache_control": mark, + **({"type": tool_call_type} if tool_call_type is not None else {}), + } + message: Final = cast(AllMessageValues, {"role": "assistant", "content": None, "tool_calls": [tool_call]}) + bedrock_cache_points: Final = sum( + 1 + for block in _convert_to_bedrock_tool_call_invoke([tool_call], model="anthropic.claude-sonnet-4-5-20250929-v1:0") + if "cachePoint" in block + ) + anthropic_marks: Final = sum( + 1 for block in convert_to_anthropic_tool_invoke([tool_call]) if block.get("cache_control") is not None + ) + + census: Final = AnthropicCacheControlHook.count_request_cache_breakpoints([message]) + + assert census == bedrock_cache_points + assert census >= anthropic_marks + assert census == (0 if mark is None else 1) def test_cache_control_hook_caps_customer_tool_call_marks_before_injection(): @@ -2127,6 +2161,40 @@ class TestEnableAnthropicPromptCaching: assert result_sys == "sys" assert result_msgs == messages + @pytest.mark.parametrize( + "tool_call_controls", + [ + pytest.param(({"type": "ephemeral"},) * 3, id="three_5m_marks_would_exceed_the_cap"), + pytest.param(({"type": "ephemeral", "ttl": "1h"},), id="1h_mark_would_follow_a_5m_default"), + ], + ) + def test_seed_stands_down_when_only_assistant_tool_calls_carry_cache_control(self, tool_call_controls): + tool_calls: Final = [ + { + "id": f"call_{index}", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + "cache_control": control, + } + for index, control in enumerate(tool_call_controls) + ] + messages: Final = [ + {"role": "system", "content": "a long system prompt"}, + {"role": "user", "content": "weather in three cities"}, + {"role": "assistant", "content": "Checking.", "tool_calls": tool_calls}, + *({"role": "tool", "tool_call_id": call["id"], "content": "sunny"} for call in tool_calls), + {"role": "user", "content": "summarize"}, + ] + params: Final[dict] = {} + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, + messages=cast(List[AllMessageValues], messages), + model="claude-sonnet-4-5", + custom_llm_provider="anthropic", + enable_prompt_caching=True, + ) + assert "cache_control_injection_points" not in params + def test_default_ttl_is_anthropics_five_minute_cache(self, monkeypatch): monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) assert all(p["control"] == {"type": "ephemeral"} for p in self._points()) From 4ef93e4cde83bfe0ce137cd07d3f07729e693d78 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 22:31:13 -0700 Subject: [PATCH 6/6] test(integration): hold the upstream so the worker kill lands mid-burst --- ...che_control_tool_call_marks_owned_proxy.py | 51 +++++++++---------- 1 file changed, 24 insertions(+), 27 deletions(-) diff --git a/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py b/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py index 46d6e9a40cb..2d8a50a2d33 100644 --- a/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py +++ b/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py @@ -1,20 +1,21 @@ import asyncio +import os import re import signal import socket +import threading from collections import Counter -from collections.abc import Sequence +from collections.abc import Callable, Sequence from dataclasses import dataclass from pathlib import Path from typing import Final import httpx -import psutil import pytest from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.process import owned_proxy_process -from integration._support.wire import Request, wire_server +from integration._support.wire import Reply, Request, wire_server from integration.providers._cache_control_marks_support import ( ASK_LABEL, CITIES, @@ -52,7 +53,6 @@ class _Sent: status: int text: str call_id: str - client_port: int def _free_port() -> int: @@ -87,16 +87,9 @@ async def _fire(owned_url: str, key: str, *, tolerate_transport_errors: bool = F surface: Final = _SURFACES[index % len(_SURFACES)] marker: Final = new_marker() path, body = _surface_request(surface, marker) - async with client.stream("POST", path, json=body, headers={"Authorization": f"Bearer {key}"}) as response: - client_port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1]) - await response.aread() + 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", ""), - client_port, + 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: @@ -108,6 +101,14 @@ async def _fire(owned_url: str, key: str, *, tolerate_transport_errors: bool = F return tuple(result for result in results if isinstance(result, _Sent)) +def _held_anthropic_peer(release: threading.Event) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert release.wait(timeout=120), "Held upstream was never released" + return anthropic_peer(request) + + return respond + + def _by_marker(received: Sequence[Request]) -> dict[str, tuple[Request, ...]]: counted: Final = Counter(marker_of(request) for request in received) return {marker: tuple(request for request in received if marker_of(request) == marker) for marker in counted} @@ -176,7 +177,8 @@ async def test_capped_burst_rides_out_a_provider_outage_and_logs_every_request_o async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_capped_requests( gateway: Gateway, tmp_path: Path ) -> None: - with wire_server(anthropic_peer) as wire: + 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)] ) @@ -188,26 +190,21 @@ async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_capped_reques ) owned_url: Final = str(owned.gateway.client.base_url) burst: Final = asyncio.create_task(_fire(owned_url, owned.gateway.key, tolerate_transport_errors=True)) - await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30) - victim: Final = psutil.Process(workers[0]) - victim.suspend() - victim_ports: Final = frozenset( - connection.raddr.port for connection in victim.net_connections(kind="tcp") if connection.raddr - ) - victim.send_signal(signal.SIGKILL) + try: + await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= _BURST, 20) + os.kill(workers[0], signal.SIGKILL) + finally: + release.set() served: Final = await burst during: Final = wire.drain() after: Final = await _fire(owned_url, owned.gateway.key) after_received: Final = wire.drain() + assert 0 < len(served) < _BURST, len(served) assert all(len(requests) == 1 for requests in _by_marker(during).values()) - _assert_capped(tuple(item for item in served if item.status == 200), during) + _assert_capped(served, during) _assert_capped(after, after_received) - survivors: Final = tuple(item for item in served if item.client_port not in victim_ports) - assert survivors, [item.client_port for item in served] - for item in (*survivors, *after): + for item in (*served, *after): _single_spend_row(item) - for item in served: - assert len(_spend_rows(item.call_id)) <= 1, item.call_id @pytest.mark.timeout(240)