From 464ee4ab2699a6fd50e3d34c6dfbc86c8fa0a96c Mon Sep 17 00:00:00 2001 From: Tan Nguyen Date: Wed, 23 Sep 2026 14:19:15 +0700 Subject: [PATCH] fix(anthropic): refuse a marker on a tool_reference block, and drop the mutable state The tool_result conversion rebuilds a tool_reference from its type and tool_name alone, so a marker written on one is dropped before the request goes out. Refuse that block type the way thinking blocks are refused. Resolve an injection point against the messages that carry no breakpoint yet instead of carrying a mutable set of the ones this request marked, so the walk needs no reassignment and the caller's own markers are walked past too. --- .../anthropic_cache_control_hook.py | 85 ++++++------------- .../test_anthropic_cache_control_hook.py | 51 ++++++++++- 2 files changed, 78 insertions(+), 58 deletions(-) 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"}, + ], + } + ], + }