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>
This commit is contained in:
Devin AI 2026-09-28 11:12:24 +00:00
parent 90e4962c81
commit 39a69561db
4 changed files with 412 additions and 10 deletions

View file

@ -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:

View file

@ -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)

View file

@ -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

View file

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