mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(anthropic): keep cache_control for Gemini targets on /v1/messages and normalize Anthropic ttl units
This commit is contained in:
parent
6019e451ce
commit
3289e22834
8 changed files with 123 additions and 426 deletions
|
|
@ -384,7 +384,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
cache_control: Final = (
|
||||
source.get("cache_control") if isinstance(source, dict) else getattr(source, "cache_control", None)
|
||||
)
|
||||
if cache_control and model and (self.is_anthropic_claude_model(model) or self.is_bedrock_arn_model(model)):
|
||||
if cache_control and model and self.target_consumes_cache_control(model):
|
||||
# TypedDict objects support dict operations at runtime
|
||||
# Use type ignore consistent with codebase pattern (see anthropic/chat/transformation.py:432)
|
||||
if isinstance(target, dict):
|
||||
|
|
@ -677,6 +677,10 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
model_lower: Final = model.lower()
|
||||
return "arn:" in model_lower and ":bedrock:" in model_lower
|
||||
|
||||
@classmethod
|
||||
def target_consumes_cache_control(cls, model: str) -> bool:
|
||||
return cls.is_anthropic_claude_model(model) or cls.is_bedrock_arn_model(model) or "gemini" in model.lower()
|
||||
|
||||
@staticmethod
|
||||
def translate_thinking_for_model(
|
||||
thinking: AnthropicThinkingParam,
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ Why separate file? Make it easy to see how transformation works
|
|||
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
|
@ -57,145 +58,56 @@ def extract_ttl_from_cached_messages(messages: list[AllMessageValues]) -> str |
|
|||
messages: List of messages to extract TTL from
|
||||
|
||||
Returns:
|
||||
Optional[str]: TTL string in format "3600s" or None if not found/invalid
|
||||
Optional[str]: TTL normalized to Gemini's "<seconds>s" form, or None if not found/invalid
|
||||
"""
|
||||
for message in messages:
|
||||
# Check message-level cache_control first
|
||||
msg_cache_control = (
|
||||
message.get("cache_control") if isinstance(message, dict) else getattr(message, "cache_control", None)
|
||||
)
|
||||
if msg_cache_control is not None:
|
||||
cc_type = (
|
||||
msg_cache_control.get("type")
|
||||
if isinstance(msg_cache_control, dict)
|
||||
else getattr(msg_cache_control, "type", None)
|
||||
)
|
||||
if cc_type == "ephemeral":
|
||||
ttl = (
|
||||
msg_cache_control.get("ttl")
|
||||
if isinstance(msg_cache_control, dict)
|
||||
else getattr(msg_cache_control, "ttl", None)
|
||||
)
|
||||
normalized = _normalize_ttl_to_seconds(ttl)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
if not is_cached_message(message):
|
||||
continue
|
||||
|
||||
content = message.get("content") if isinstance(message, dict) else getattr(message, "content", None)
|
||||
if not isinstance(content, list):
|
||||
content = message.get("content")
|
||||
if not content or isinstance(content, str):
|
||||
continue
|
||||
|
||||
for content_item in content:
|
||||
# Check if content_item is dict or object model
|
||||
if isinstance(content_item, dict):
|
||||
cache_control = content_item.get("cache_control")
|
||||
item_type = content_item.get("type")
|
||||
else:
|
||||
cache_control = getattr(content_item, "cache_control", None)
|
||||
item_type = getattr(content_item, "type", None)
|
||||
# Type check to ensure content_item is a dictionary before calling .get()
|
||||
if not isinstance(content_item, dict):
|
||||
continue
|
||||
|
||||
if item_type == "text" and cache_control is not None:
|
||||
cc_type = (
|
||||
cache_control.get("type")
|
||||
if isinstance(cache_control, dict)
|
||||
else getattr(cache_control, "type", None)
|
||||
)
|
||||
if cc_type == "ephemeral":
|
||||
ttl = (
|
||||
cache_control.get("ttl")
|
||||
if isinstance(cache_control, dict)
|
||||
else getattr(cache_control, "ttl", None)
|
||||
)
|
||||
normalized = _normalize_ttl_to_seconds(ttl)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
cache_control = content_item.get("cache_control")
|
||||
if not cache_control or not isinstance(cache_control, dict):
|
||||
continue
|
||||
|
||||
if cache_control.get("type") != "ephemeral":
|
||||
continue
|
||||
|
||||
normalized_ttl = _normalize_ttl_to_seconds(cache_control.get("ttl"))
|
||||
if normalized_ttl is not None:
|
||||
return normalized_ttl
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _is_valid_ttl_format(ttl: str) -> bool:
|
||||
"""
|
||||
Validate TTL format. Should be a string ending with 's' for seconds.
|
||||
Examples: "3600s", "7200s", "1.5s"
|
||||
|
||||
Args:
|
||||
ttl: TTL string to validate
|
||||
|
||||
Returns:
|
||||
bool: True if valid format, False otherwise
|
||||
"""
|
||||
if not isinstance(ttl, str):
|
||||
return False
|
||||
|
||||
# TTL should end with 's' and contain a valid number before it
|
||||
pattern: Final = r"^([0-9]*\.?[0-9]+)s$"
|
||||
match: Final = re.match(pattern, ttl)
|
||||
|
||||
if not match:
|
||||
return False
|
||||
|
||||
try:
|
||||
# Ensure the numeric part is valid and positive
|
||||
numeric_part: Final = float(match.group(1))
|
||||
return numeric_part > 0
|
||||
except ValueError:
|
||||
return False
|
||||
_TTL_PATTERN: Final = re.compile(r"^([0-9]*\.?[0-9]+)([smh])$")
|
||||
_TTL_UNIT_SECONDS: Final = MappingProxyType({"s": 1, "m": 60, "h": 3600})
|
||||
|
||||
|
||||
def _normalize_ttl_to_seconds(ttl: object) -> str | None:
|
||||
"""
|
||||
Normalize a cache_control TTL into Gemini's "<seconds>s" format.
|
||||
|
||||
Accepts Gemini-native seconds (e.g. "3600s", "1.5s") and Anthropic-style
|
||||
minute/hour units (e.g. "5m", "1h") that Claude Code and the Anthropic
|
||||
/v1/messages spec use. Caps the requested TTL at 24 hours (86400s) to
|
||||
prevent unbounded persistent storage costs. Returns None for missing or
|
||||
unparseable values so Gemini falls back to its own default TTL.
|
||||
Gemini's cachedContents API only takes a TTL as "<seconds>s", while Anthropic clients
|
||||
(Claude Code among them) send the minute and hour units the Anthropic API defines, "5m"
|
||||
and "1h". Returns the Gemini form for any of the three units, or None for a missing,
|
||||
non-positive, or unparseable value so the cache falls back to Gemini's default TTL.
|
||||
"""
|
||||
if not isinstance(ttl, str):
|
||||
return None
|
||||
|
||||
match = re.match(r"^([0-9]*\.?[0-9]+)(s|m|h)$", ttl)
|
||||
if not match:
|
||||
match: Final = _TTL_PATTERN.match(ttl)
|
||||
if match is None:
|
||||
return None
|
||||
|
||||
value = float(match.group(1))
|
||||
|
||||
value: Final = float(match.group(1))
|
||||
if value <= 0:
|
||||
return None
|
||||
|
||||
multiplier = {"s": 1, "m": 60, "h": 3600}[match.group(2)]
|
||||
seconds = value * multiplier
|
||||
|
||||
# Cap explicit caches to 24 hours to prevent unbounded billing costs
|
||||
seconds = min(seconds, 86400.0)
|
||||
|
||||
# Google Protobuf Duration requires up to 9 fractional digits
|
||||
seconds = round(seconds, 9)
|
||||
return f"{int(seconds)}s" if seconds.is_integer() else f"{seconds}s"
|
||||
|
||||
|
||||
def get_gemini_context_caching_min_tokens(model: str) -> int:
|
||||
"""
|
||||
Minimum input token count required to create an explicit Gemini context cache.
|
||||
|
||||
Looks up the `cache_creation_min_tokens` property from model_prices_and_context_window.json.
|
||||
Defaults to string-matching fallbacks for unknown models.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
try:
|
||||
model_info = litellm.get_model_info(model=model)
|
||||
if model_info and "cache_creation_min_tokens" in model_info:
|
||||
return int(model_info["cache_creation_min_tokens"])
|
||||
except Exception: # noqa: BLE001 # fallback to string-matching heuristic if model lookup fails
|
||||
pass
|
||||
|
||||
model_lower = model.lower()
|
||||
if "gemini-2.5" in model_lower or "gemini-2-5" in model_lower:
|
||||
return 2048
|
||||
if "gemini-3" in model_lower:
|
||||
return 4096
|
||||
return 32768
|
||||
seconds: Final = round(value * _TTL_UNIT_SECONDS[match.group(2)], 9)
|
||||
return f"{seconds:.9f}".rstrip("0").rstrip(".") + "s"
|
||||
|
||||
|
||||
def separate_cached_messages(
|
||||
|
|
|
|||
|
|
@ -25066,6 +25066,7 @@
|
|||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
@ -25925,6 +25926,7 @@
|
|||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
@ -27137,6 +27139,7 @@
|
|||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
@ -27895,6 +27898,7 @@
|
|||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
|
|||
|
|
@ -25066,6 +25066,7 @@
|
|||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
@ -25925,6 +25926,7 @@
|
|||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
@ -27137,6 +27139,7 @@
|
|||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
@ -27895,6 +27898,7 @@
|
|||
"supports_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"prompt_cache_min_tokens": 2048,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
|
|
|
|||
|
|
@ -2145,21 +2145,6 @@ def test_should_add_cache_control_for_gemini_model():
|
|||
assert target.get("cache_control") == cache_control
|
||||
|
||||
|
||||
def test_cache_control_fallback_setattr():
|
||||
"""Verify cache_control is safely assigned to non-dict target objects using setattr."""
|
||||
adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
cache_control = {"type": "ephemeral"}
|
||||
|
||||
class MockTarget:
|
||||
pass
|
||||
|
||||
target = MockTarget()
|
||||
adapter._add_cache_control_if_applicable(
|
||||
{"cache_control": cache_control}, target, "claude-3-opus-20240229"
|
||||
)
|
||||
assert getattr(target, "cache_control", None) == cache_control
|
||||
|
||||
|
||||
def test_cache_control_preserved_in_text_content_for_gemini():
|
||||
"""cache_control must survive message translation for a Gemini target."""
|
||||
anthropic_messages = [
|
||||
|
|
|
|||
|
|
@ -1,143 +1,77 @@
|
|||
import pytest
|
||||
from litellm.llms.vertex_ai.context_caching.transformation import (
|
||||
extract_ttl_from_cached_messages,
|
||||
get_gemini_context_caching_min_tokens,
|
||||
_is_valid_ttl_format,
|
||||
_normalize_ttl_to_seconds,
|
||||
extract_ttl_from_cached_messages,
|
||||
transform_openai_messages_to_gemini_context_caching,
|
||||
)
|
||||
|
||||
|
||||
class TestGeminiContextCachingMinTokens:
|
||||
"""Per-model floor for explicit Gemini context cache creation."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, expected",
|
||||
[
|
||||
("gemini-1.5-pro", 32768),
|
||||
("gemini-1.5-flash", 32768),
|
||||
("vertex_ai/gemini-1.5-pro-001", 32768),
|
||||
("gemini-2.5-flash", 2048),
|
||||
("gemini-2.5-pro", 2048),
|
||||
("gemini/gemini-2.5-pro", 2048),
|
||||
("vertex_ai/gemini-2.5-flash", 2048),
|
||||
("gemini-3.5-flash", 4096),
|
||||
("gemini-3.1-pro-preview", 4096),
|
||||
("gemini/gemini-3.5-flash", 4096),
|
||||
("gemini-unknown-future-model", 32768),
|
||||
],
|
||||
)
|
||||
def test_min_tokens_by_model(self, model, expected):
|
||||
assert get_gemini_context_caching_min_tokens(model) == expected
|
||||
|
||||
def test_min_tokens_from_model_info(self, monkeypatch):
|
||||
"""Should prefer cache_creation_min_tokens from model_info if present."""
|
||||
import litellm
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"get_model_info",
|
||||
lambda model, **kwargs: {"cache_creation_min_tokens": 12345}
|
||||
)
|
||||
assert get_gemini_context_caching_min_tokens("gemini-1.5-pro") == 12345
|
||||
|
||||
|
||||
class TestTTLValidation:
|
||||
"""Test TTL format validation"""
|
||||
|
||||
def test_valid_ttl_formats(self):
|
||||
"""Test various valid TTL formats"""
|
||||
valid_ttls = ["3600s", "1s", "7200s", "1.5s", "0.1s", "86400s", "123.456s"]
|
||||
|
||||
for ttl in valid_ttls:
|
||||
assert _is_valid_ttl_format(ttl), f"TTL {ttl} should be valid"
|
||||
|
||||
def test_invalid_ttl_formats(self):
|
||||
"""Test various invalid TTL formats"""
|
||||
invalid_ttls = [
|
||||
"3600", # missing 's'
|
||||
"s", # missing number
|
||||
"-1s", # negative number
|
||||
"0s", # zero
|
||||
"3600m", # wrong unit
|
||||
"abc.s", # invalid number
|
||||
"", # empty string
|
||||
"3600.s", # invalid decimal
|
||||
"3600 s", # space
|
||||
"3600ss", # extra 's'
|
||||
None, # None
|
||||
123, # not a string
|
||||
]
|
||||
|
||||
for ttl in invalid_ttls:
|
||||
assert not _is_valid_ttl_format(ttl), f"TTL {ttl} should be invalid"
|
||||
|
||||
|
||||
class TestTTLNormalization:
|
||||
"""Normalization of anthropic-style TTL units into Gemini's seconds format."""
|
||||
"""Gemini only takes "<seconds>s"; Anthropic clients send "5m" and "1h" too"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ttl, expected",
|
||||
[
|
||||
("3600s", "3600s"),
|
||||
("1s", "1s"),
|
||||
("1.5s", "1.5s"),
|
||||
("0.1s", "0.1s"),
|
||||
("123.456s", "123.456s"),
|
||||
("1.3333333333333333s", "1.333333333s"),
|
||||
("5m", "300s"),
|
||||
("90m", "5400s"),
|
||||
("1h", "3600s"),
|
||||
("2h", "7200s"),
|
||||
("0.5h", "1800s"),
|
||||
("48h", "86400s"),
|
||||
("1500m", "86400s"),
|
||||
("1000000s", "86400s"),
|
||||
("48h", "172800s"),
|
||||
],
|
||||
)
|
||||
def test_normalizes_units_to_seconds(self, ttl, expected):
|
||||
def test_normalizes_supported_units_to_seconds(self, ttl, expected):
|
||||
assert _normalize_ttl_to_seconds(ttl) == expected
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ttl",
|
||||
["invalid", "", "0m", "0h", "-1h", "5d", "1 h", "m", None, 123, 3600],
|
||||
[
|
||||
"3600",
|
||||
"s",
|
||||
"-1s",
|
||||
"0s",
|
||||
"0m",
|
||||
"0h",
|
||||
"5d",
|
||||
"abc.s",
|
||||
"",
|
||||
"3600.s",
|
||||
"3600 s",
|
||||
"3600ss",
|
||||
"1 h",
|
||||
None,
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_rejects_unparseable_ttl(self, ttl):
|
||||
assert _normalize_ttl_to_seconds(ttl) is None
|
||||
|
||||
def test_extract_ttl_normalizes_anthropic_hour_unit(self):
|
||||
"""Claude Code / Anthropic send "1h"; Gemini must receive "3600s"."""
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "cached",
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h"},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
assert extract_ttl_from_cached_messages(messages) == "3600s"
|
||||
|
||||
def test_extract_ttl_normalizes_anthropic_minute_unit(self):
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "cached",
|
||||
"cache_control": {"type": "ephemeral", "ttl": "5m"},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
assert extract_ttl_from_cached_messages(messages) == "300s"
|
||||
|
||||
|
||||
class TestTTLExtraction:
|
||||
"""Test TTL extraction from cached messages"""
|
||||
|
||||
@pytest.mark.parametrize("ttl, expected", [("1h", "3600s"), ("5m", "300s")])
|
||||
def test_extract_ttl_normalizes_anthropic_units(self, ttl, expected):
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "cached",
|
||||
"cache_control": {"type": "ephemeral", "ttl": ttl},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
assert extract_ttl_from_cached_messages(messages) == expected
|
||||
|
||||
def test_extract_ttl_from_single_message(self):
|
||||
"""Test extracting TTL from a single cached message"""
|
||||
messages = [
|
||||
|
|
@ -189,7 +123,9 @@ class TestTTLExtraction:
|
|||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "text", "text": "Regular message without cache control"}],
|
||||
"content": [
|
||||
{"type": "text", "text": "Regular message without cache control"}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
|
|
@ -271,7 +207,9 @@ class TestTTLExtraction:
|
|||
class TestTransformationWithTTL:
|
||||
"""Test the complete transformation with TTL support"""
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
def test_transform_with_valid_ttl(self, custom_llm_provider):
|
||||
"""Test transformation includes TTL when provided"""
|
||||
messages = [
|
||||
|
|
@ -312,7 +250,9 @@ class TestTransformationWithTTL:
|
|||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
def test_transform_without_ttl(self, custom_llm_provider):
|
||||
"""Test transformation without TTL"""
|
||||
messages = [
|
||||
|
|
@ -352,7 +292,9 @@ class TestTransformationWithTTL:
|
|||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
def test_transform_with_invalid_ttl(self, custom_llm_provider):
|
||||
"""Test transformation with invalid TTL (should be ignored)"""
|
||||
messages = [
|
||||
|
|
@ -391,7 +333,9 @@ class TestTransformationWithTTL:
|
|||
|
||||
assert result["displayName"] == "test-cache-key"
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
def test_transform_with_system_message_and_ttl(self, custom_llm_provider):
|
||||
"""Test transformation with system message and TTL"""
|
||||
messages = [
|
||||
|
|
@ -476,143 +420,6 @@ class TestEdgeCases:
|
|||
assert isinstance(ttl, str)
|
||||
assert ttl == "3600s"
|
||||
|
||||
def test_cache_control_preserved_for_object_content_items(self):
|
||||
"""Test that cache_control is preserved when content items are real Pydantic models."""
|
||||
from pydantic import BaseModel, Field
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
class MockContentBlock:
|
||||
def __init__(self):
|
||||
self.type = "text"
|
||||
self.text = "hello"
|
||||
self.cache_control = {"type": "ephemeral"}
|
||||
|
||||
class RealPydanticV2Block(BaseModel):
|
||||
type: str = "text"
|
||||
text: str = "hello v2"
|
||||
cache_control: dict = Field(default_factory=lambda: {"type": "ephemeral"})
|
||||
|
||||
class MockBlockWithNoneCacheControl:
|
||||
def __init__(self):
|
||||
self.type = "text"
|
||||
self.text = "hello none"
|
||||
self.cache_control = None
|
||||
|
||||
content = [
|
||||
MockContentBlock(),
|
||||
RealPydanticV2Block(),
|
||||
MockBlockWithNoneCacheControl(),
|
||||
]
|
||||
result = LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content(content)
|
||||
assert result == [
|
||||
{"type": "text", "text": "hello", "cache_control": {"type": "ephemeral"}},
|
||||
{"type": "text", "text": "hello v2", "cache_control": {"type": "ephemeral"}},
|
||||
{"type": "text", "text": "hello none"},
|
||||
]
|
||||
|
||||
def test_is_cached_message_for_object_message_and_content_item(self):
|
||||
"""Test is_cached_message on custom objects / models."""
|
||||
from litellm.utils import is_cached_message
|
||||
|
||||
# Test message level cache_control object
|
||||
class MockCacheControl:
|
||||
def __init__(self):
|
||||
self.type = "ephemeral"
|
||||
|
||||
class MockMessageLevelObj:
|
||||
def __init__(self):
|
||||
self.role = "system"
|
||||
self.content = "hello"
|
||||
self.cache_control = MockCacheControl()
|
||||
|
||||
msg = MockMessageLevelObj()
|
||||
assert is_cached_message(msg) is True
|
||||
|
||||
# Test content level cache_control object
|
||||
class MockContentItem:
|
||||
def __init__(self):
|
||||
self.type = "text"
|
||||
self.text = "hello"
|
||||
self.cache_control = MockCacheControl()
|
||||
|
||||
class MockContentLevelObj:
|
||||
def __init__(self):
|
||||
self.role = "system"
|
||||
self.content = [MockContentItem()]
|
||||
|
||||
msg = MockContentLevelObj()
|
||||
assert is_cached_message(msg) is True
|
||||
|
||||
def test_extract_ttl_from_cached_messages_for_object_models(self):
|
||||
"""Test extract_ttl_from_cached_messages with object-based messages and content items."""
|
||||
|
||||
class MockCacheControl:
|
||||
def __init__(self):
|
||||
self.type = "ephemeral"
|
||||
self.ttl = "3600s"
|
||||
|
||||
class MockContentItem:
|
||||
def __init__(self):
|
||||
self.type = "text"
|
||||
self.text = "hello"
|
||||
self.cache_control = MockCacheControl()
|
||||
|
||||
class MockMessageObj:
|
||||
def __init__(self):
|
||||
self.role = "system"
|
||||
self.content = [MockContentItem()]
|
||||
|
||||
messages = [MockMessageObj()]
|
||||
ttl = extract_ttl_from_cached_messages(messages)
|
||||
assert ttl == "3600s"
|
||||
|
||||
def test_extract_ttl_from_cached_messages_with_message_level_object_cache_control(self):
|
||||
"""Test extract_ttl_from_cached_messages with message-level object cache_control."""
|
||||
|
||||
class MockCacheControl:
|
||||
def __init__(self):
|
||||
self.type = "ephemeral"
|
||||
self.ttl = "7200s"
|
||||
|
||||
class MockMessageObj:
|
||||
def __init__(self):
|
||||
self.role = "system"
|
||||
self.content = "hello"
|
||||
self.cache_control = MockCacheControl()
|
||||
|
||||
messages = [MockMessageObj()]
|
||||
ttl = extract_ttl_from_cached_messages(messages)
|
||||
assert ttl == "7200s"
|
||||
|
||||
def test_is_cached_message_for_dict_message_with_dict_content_items(self):
|
||||
"""Test is_cached_message with dict message and dict content list items."""
|
||||
from litellm.utils import is_cached_message
|
||||
|
||||
# Dictionary message without content should return False
|
||||
assert is_cached_message({"role": "user"}) is False
|
||||
|
||||
msg = {
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "hello", "cache_control": {"type": "ephemeral"}}
|
||||
],
|
||||
}
|
||||
assert is_cached_message(msg) is True
|
||||
|
||||
def test_normalize_responses_api_object_to_dict_pydantic_v1(self):
|
||||
"""Test _normalize_responses_api_object_to_dict with Pydantic v1 dict fallback."""
|
||||
from litellm.responses.litellm_completion_transformation.transformation import LiteLLMCompletionResponsesConfig
|
||||
|
||||
class MockPydanticV1Model:
|
||||
def dict(self):
|
||||
return {"type": "text", "text": "hello", "cache_control": {"type": "ephemeral"}}
|
||||
|
||||
item = MockPydanticV1Model()
|
||||
res = LiteLLMCompletionResponsesConfig._normalize_responses_api_object_to_dict(item)
|
||||
assert res == {"type": "text", "text": "hello", "cache_control": {"type": "ephemeral"}}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
|
|
|
|||
|
|
@ -1396,62 +1396,44 @@ class TestContextCachingEndpoints:
|
|||
# Restart the patcher so teardown_method can stop it cleanly
|
||||
self._token_check_patcher.start()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, expected_min",
|
||||
[
|
||||
("gemini-3.5-flash", 4096),
|
||||
("gemini/gemini-3.5-flash", 4096),
|
||||
("gemini-3.1-pro-preview", 4096),
|
||||
("gemini-1.5-pro", 32768),
|
||||
("gemini-2.5-flash", 2048),
|
||||
("gemini-2.5-pro", 2048),
|
||||
],
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
def test_check_and_create_cache_uses_model_specific_min_tokens(
|
||||
self, mock_separate, model, expected_min
|
||||
@pytest.mark.parametrize("model", ["gemini-2.5-flash", "gemini-2.5-pro"])
|
||||
def test_check_and_create_cache_skips_between_default_and_gemini_2_5_minimum(
|
||||
self, model, local_model_cost_map
|
||||
):
|
||||
"""The Gemini per-model floor must be forwarded to the token-count guard.
|
||||
"""Gemini 2.5 Flash and Pro need 2048 cached tokens, twice the provider-agnostic default.
|
||||
|
||||
A flat 1024 floor let content between 1024 and the real minimum (2048 for
|
||||
2.5, 4096 for 3.x) reach Gemini and 400. Assert the model-derived floor is
|
||||
passed so the guard skips instead of erroring.
|
||||
Content between the two used to reach Google's cachedContents endpoint and 400.
|
||||
"""
|
||||
self._token_check_patcher.stop()
|
||||
|
||||
cached_messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "cached",
|
||||
"content": " ".join(["word"] * 1500),
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
]
|
||||
non_cached_messages = [{"role": "user", "content": "Hello"}]
|
||||
mock_separate.return_value = (cached_messages, non_cached_messages)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.is_prompt_caching_valid_prompt",
|
||||
return_value=False,
|
||||
) as mock_valid:
|
||||
self.context_caching.check_and_create_cache(
|
||||
messages=cached_messages + non_cached_messages,
|
||||
optional_params=self.sample_optional_params.copy(),
|
||||
api_key="test_key",
|
||||
api_base=None,
|
||||
model=model,
|
||||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
cached_content=None,
|
||||
custom_llm_provider="gemini",
|
||||
vertex_project="test_project",
|
||||
vertex_location="us-central1",
|
||||
vertex_auth_header="test_token",
|
||||
)
|
||||
messages, _, returned_cache = self.context_caching.check_and_create_cache(
|
||||
messages=cached_messages + non_cached_messages,
|
||||
optional_params=self.sample_optional_params.copy(),
|
||||
api_key="test_key",
|
||||
api_base=None,
|
||||
model=model,
|
||||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
cached_content=None,
|
||||
custom_llm_provider="gemini",
|
||||
vertex_project="test_project",
|
||||
vertex_location="us-central1",
|
||||
vertex_auth_header="test_token",
|
||||
)
|
||||
|
||||
assert mock_valid.call_args.kwargs["min_token_count"] == expected_min
|
||||
assert messages == cached_messages + non_cached_messages
|
||||
assert returned_cache is None
|
||||
self.mock_client.post.assert_not_called()
|
||||
|
||||
self._token_check_patcher.start()
|
||||
|
||||
|
|
|
|||
|
|
@ -686,7 +686,6 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"supports_computer_use": {"type": "boolean"},
|
||||
"cache_creation_input_audio_token_cost": {"type": "number"},
|
||||
"cache_creation_input_token_cost": {"type": "number"},
|
||||
"cache_creation_min_tokens": {"type": "number"},
|
||||
"cache_creation_input_token_cost_above_1hr": {"type": "number"},
|
||||
"cache_creation_input_token_cost_above_128k_tokens": {"type": "number"},
|
||||
"cache_creation_input_token_cost_above_200k_tokens": {"type": "number"},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue