diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index c54b42bf0ca..9e95678af93 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -439,6 +439,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): """ used_blocks = AnthropicCacheControlHook.count_request_cache_breakpoints(messages) + taken: set[int] = set() limit_reached = False for point in points: if used_blocks >= max_blocks: @@ -449,7 +450,9 @@ class AnthropicCacheControlHook(CustomPromptManagement): type="ephemeral" ) - for target_index in AnthropicCacheControlHook._resolve_target_indices(point=point, messages=messages): + for target_index in AnthropicCacheControlHook._resolve_target_indices( + point=point, messages=messages, taken=taken + ): if used_blocks >= max_blocks: limit_reached = True break @@ -463,6 +466,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): ) if AnthropicCacheControlHook._message_has_cache_control(messages[target_index]): used_blocks += 1 + taken.add(target_index) if limit_reached: break @@ -477,9 +481,16 @@ class AnthropicCacheControlHook(CustomPromptManagement): @staticmethod def _resolve_target_indices( - point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues] + point: CacheControlMessageInjectionPoint, messages: list[AllMessageValues], taken: set[int] | None = None ) -> list[int]: - """Resolve which message indices an injection point targets.""" + """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. + """ + already_marked: Final = taken or set() _targetted_index: Final[int | str | None] = point.get("index", None) targetted_index: int | None = None if isinstance(_targetted_index, str): @@ -527,7 +538,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): position = targetted_index while position >= 0: index = candidates[position] - if _message_accepts_cache_control(messages[index]): + if index not in already_marked and _message_accepts_cache_control(messages[index]): return [index] position -= 1 return [] diff --git a/tests/unit/integrations/test_anthropic_cache_control_hook.py b/tests/unit/integrations/test_anthropic_cache_control_hook.py index c9af51e0392..ed0ef328fe4 100644 --- a/tests/unit/integrations/test_anthropic_cache_control_hook.py +++ b/tests/unit/integrations/test_anthropic_cache_control_hook.py @@ -4061,3 +4061,30 @@ async def test_anthropic_cache_control_hook_walks_back_off_a_thinking_only_turn( "content": [{"type": "text", "text": "a1", "cache_control": {"type": "ephemeral"}}], } assert marked[3] == {"role": "assistant", "content": [{"type": "thinking", "thinking": "t", "signature": "s"}]} + + +@pytest.mark.asyncio +async def test_anthropic_cache_control_hook_two_points_stay_two_breakpoints(monkeypatch: pytest.MonkeyPatch): + """ + The second point walks past the message the first one marked. Arriving on it, + finding it marked and dropping the point turns two configured breakpoints into one, + and the four exist so a prefix that stops matching at one can match at an earlier one. + """ + marked = await _marked_messages( + [ + {"role": "user", "content": "open"}, + {"role": "assistant", "content": "ack"}, + {"role": "user", "content": ""}, + ], + [{"location": "message", "index": -1}, {"location": "message", "index": -2}], + monkeypatch, + ) + + assert marked == [ + {"role": "user", "content": [{"type": "text", "text": "open", "cache_control": {"type": "ephemeral"}}]}, + {"role": "assistant", "content": [{"type": "text", "text": "ack", "cache_control": {"type": "ephemeral"}}]}, + { + "role": "user", + "content": [{"type": "text", "text": "[System: Empty message content sanitised to satisfy protocol]"}], + }, + ]