diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 24932591f70..733937039c5 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -12,7 +12,7 @@ Supported for both `v1/chat/completions` (via the prompt-management hook) and import copy import os import re -from collections.abc import Container, Iterable, Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, cast from urllib.parse import urlparse @@ -67,10 +67,12 @@ _GPT_VERSION_PATTERN: Final = re.compile(r"^gpt-(\d+)(?:\.(\d+))?") OPENAI_PROMPT_CACHE_BREAKPOINT_BLOCK_TYPES: Final = frozenset( {"text", "image", "image_url", "file", "input_audio", "input_text", "input_image", "input_file"} ) -# Anthropic lists the block types a cache_control marker may sit on: text, image, -# tool_use, tool_result and document. A list of refused types rather than accepted -# ones so a block type this code has not been told about still takes a marker. -ANTHROPIC_BLOCK_TYPES_WITHOUT_CACHE_CONTROL: Final = frozenset({"thinking", "redacted_thinking"}) +# Block types a marker never reaches the provider on. Anthropic accepts one on text, +# image, tool_use, tool_result and document, so a thinking block is refused there; a +# tool_reference is rebuilt from its type and tool_name alone by the tool_result +# conversion, which drops everything else. A list of refused types rather than +# accepted ones, so a block type this code has not been told about still takes one. +ANTHROPIC_BLOCK_TYPES_WITHOUT_CACHE_CONTROL: Final = frozenset({"thinking", "redacted_thinking", "tool_reference"}) OPENAI_API_HOST: Final = "api.openai.com" OPENAI_API_BASE_ENV_VARS: Final = ("OPENAI_BASE_URL", "OPENAI_API_BASE") _OBJECT_MAPPING_ADAPTER: Final = TypeAdapter(dict[object, object]) @@ -144,14 +146,10 @@ def _accepts_prompt_cache_breakpoint(block: object) -> bool: def _index_of_block_accepting_cache_control(content: list[object], on_a_tool_message: bool) -> int | None: - """Position of the last block a cache_control marker can be written on, or None. + """Last block a marker written here reaches the provider on, searched from the end. - Searched from the end: a marker caches everything up to and including its own - block, so the last one caches the most. - - An empty text block is refused because the rewrites this hook runs ahead of drop - it, and the marker goes with it. ``on_a_tool_message`` lifts that refusal: a tool - message keeps its empty block, nested in the tool_result the rewrite builds. + An empty text block is replaced by a placeholder unless it sits on a tool message, + where it goes out inside the tool_result. """ for index in range(len(content) - 1, -1, -1): block = content[index] @@ -166,13 +164,7 @@ def _index_of_block_accepting_cache_control(content: list[object], on_a_tool_mes def _message_accepts_cache_control(message: object) -> bool: - """Whether a cache_control marker written on this message reaches the provider. - - Reads the same block rule as `_safe_insert_cache_control_in_message`, so the two - cannot disagree about where a marker goes. A message whose content is empty -- - an assistant turn that said everything through ``tool_calls`` -- has nowhere to - put one. - """ + """Whether a marker written on this message reaches the provider.""" if not isinstance(message, dict): return False on_a_tool_message: Final = message.get("role") == "tool" @@ -439,7 +431,6 @@ class AnthropicCacheControlHook(CustomPromptManagement): """ used_blocks = AnthropicCacheControlHook.count_request_cache_breakpoints(messages) - taken: set[int] = set() # mutable-ok: the messages this request has already marked limit_reached = False for point in points: if used_blocks >= max_blocks: @@ -450,9 +441,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): type="ephemeral" ) - for target_index in AnthropicCacheControlHook._resolve_target_indices( - point=point, messages=messages, taken=taken - ): + for target_index in AnthropicCacheControlHook._resolve_target_indices(point=point, messages=messages): if used_blocks >= max_blocks: limit_reached = True break @@ -466,7 +455,6 @@ class AnthropicCacheControlHook(CustomPromptManagement): ) if AnthropicCacheControlHook._message_has_cache_control(messages[target_index]): used_blocks += 1 - taken.add(target_index) if limit_reached: break @@ -481,15 +469,9 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def _resolve_target_indices( - point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues], taken: Container[int] = frozenset() + point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues] ) -> list[int]: - """Resolve which message indices an injection point targets. - - ``taken`` is the messages an earlier point in the same request already marked. - The walk goes past them: arriving on one, finding it marked and dropping the - point turns two configured breakpoints into one, and the four exist so that a - prefix which stops matching at one can still match at an earlier one. - """ + """Resolve which message indices an injection point targets.""" _targetted_index: Final[int | str | None] = point.get("index", None) targetted_index: int | None = None if isinstance(_targetted_index, str): @@ -500,47 +482,36 @@ class AnthropicCacheControlHook(CustomPromptManagement): else: targetted_index = _targetted_index - # The messages the point names. A point with a role counts its index within - # that role's turns, so {role: assistant, index: -1} is the last assistant - # turn rather than the last message. targetted_role: Final = point.get("role", None) candidates: Final = tuple( index for index, message in enumerate(messages) if targetted_role is None or message.get("role") == targetted_role ) + free: Final = frozenset( + index + for index in candidates + if not AnthropicCacheControlHook._message_has_cache_control(messages[index]) + and _message_accepts_cache_control(messages[index]) + ) # Case 1: Target by role alone if targetted_index is None: - if targetted_role is None: - return [] - return [index for index in candidates if _message_accepts_cache_control(messages[index])] + return [] if targetted_role is None else [index for index in candidates if index in free] - # Case 2: Target by index, within the role the point named - original_index: Final = targetted_index - if targetted_index < 0: - targetted_index += len(candidates) - - if not 0 <= targetted_index < len(candidates): + # Case 2: Target by index, counted within the role the point named + position: Final = targetted_index + len(candidates) if targetted_index < 0 else targetted_index + if not 0 <= position < len(candidates): verbose_logger.warning( "AnthropicCacheControlHook: Provided index %s is out of bounds for message list of length %s. Targeted index was %s. Skipping cache control injection for this point.", - original_index, - len(candidates), targetted_index, + len(candidates), + position, ) return [] - # The message it arrives at may have nowhere to write a marker -- an assistant - # turn that said everything through tool_calls is the common one -- so walk back - # to the nearest earlier turn that does. Stopping there would spend the point - # and write nothing. - position = targetted_index - while position >= 0: - index = candidates[position] - if index not in taken and _message_accepts_cache_control(messages[index]): - return [index] - position -= 1 - return [] + landing: Final = next((candidates[step] for step in range(position, -1, -1) if candidates[step] in free), None) + return [] if landing is None else [landing] @staticmethod def _count_cache_control_blocks(message: object) -> int: diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index ed0ef328fe4..13842332a5c 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -16,6 +16,7 @@ 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.integrations.anthropic_cache_control_hook import CacheControlMessageInjectionPoint from litellm.types.llms.openai import AllMessageValues @@ -3819,7 +3820,11 @@ def _anthropic_response_mock() -> MagicMock: return mock_response -async def _marked_messages(messages, points, monkeypatch: pytest.MonkeyPatch): +async def _marked_messages( + messages: list[AllMessageValues], + points: list[CacheControlMessageInjectionPoint], + monkeypatch: pytest.MonkeyPatch, +) -> list[AllMessageValues]: monkeypatch.setenv("ANTHROPIC_API_KEY", "fake_anthropic_key") monkeypatch.setattr(litellm, "callbacks", [AnthropicCacheControlHook()]) client = AsyncHTTPHandler() @@ -4088,3 +4093,47 @@ async def test_anthropic_cache_control_hook_two_points_stay_two_breakpoints(monk "content": [{"type": "text", "text": "[System: Empty message content sanitised to satisfy protocol]"}], }, ] + + +@pytest.mark.asyncio +async def test_anthropic_cache_control_hook_skips_a_tool_reference_block(monkeypatch: pytest.MonkeyPatch): + """ + The tool_result conversion rebuilds a tool_reference from its type and tool_name + alone, so a marker written on one never reaches the provider. + """ + marked = await _marked_messages( + [ + {"role": "user", "content": "find a tool"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "tool_search", "arguments": "{}"}} + ], + }, + { + "role": "tool", + "tool_call_id": "call_1", + "content": [ + {"type": "text", "text": "found"}, + {"type": "tool_reference", "tool_name": "get_weather"}, + ], + }, + ], + [{"location": "message", "index": 2}], + monkeypatch, + ) + + assert marked[-1] == { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "call_1", + "content": [ + {"type": "text", "text": "found", "cache_control": {"type": "ephemeral"}}, + {"type": "tool_reference", "tool_name": "get_weather"}, + ], + } + ], + }