mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
90e4962c81
commit
39a69561db
4 changed files with 412 additions and 10 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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'.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue