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'.