diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index e9235bc80a7..7f78b16ec74 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -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, diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index 1563fb80d1b..ef415dfa19c 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -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 "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 "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 "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( diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index f8c585873de..fcdf6c4baa3 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index f8c585873de..fcdf6c4baa3 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index d0b92c3e073..471c09153c0 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -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 = [ diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_context_caching_ttl.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_context_caching_ttl.py index 4896a75ade7..b2da8da4cc5 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_context_caching_ttl.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_context_caching_ttl.py @@ -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 "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"]) diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 226c6441516..5c33b9a995b 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -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() diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index c1c4613f7a9..2fda5dfc490 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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"},