diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 2d35dd9b480..63725af49d3 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -5,7 +5,6 @@ import httpx import litellm from litellm._logging import verbose_logger from litellm.caching.caching import Cache, LiteLLMCacheType -from litellm.constants import MINIMUM_PROMPT_CACHE_TOKEN_COUNT from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -17,7 +16,7 @@ from litellm.types.llms.vertex_ai import ( CachedContentListAllResponseBody, VertexAICachedContentResponseObject, ) -from litellm.utils import is_prompt_caching_valid_prompt +from litellm.utils import get_prompt_cache_min_tokens, is_prompt_caching_valid_prompt from ..common_utils import VertexAIError, get_vertex_base_url from ..vertex_llm_base import VertexBase @@ -32,6 +31,23 @@ local_cache_obj: Final = Cache(type=LiteLLMCacheType.LOCAL) # only used for cal MAX_PAGINATION_PAGES: Final = 100 # Reasonable upper bound for pagination +def _is_cached_content_too_small_error(err: httpx.HTTPStatusError) -> bool: + if err.response.status_code != 400: + return False + error_text: Final = (err.response.text or "").lower() + return "minimum token count to start explicit caching" in error_text or "cached content is too small" in error_text + + +def _raise_unless_cached_content_too_small(err: httpx.HTTPStatusError) -> None: + if not _is_cached_content_too_small_error(err): + raise VertexAIError(status_code=err.response.status_code, message=err.response.text) + verbose_logger.debug( + "Vertex AI context caching: server rejected cached content as below " + "minimum token threshold (%s). Falling back to uncached request.", + err.response.text, + ) + + class ContextCachingEndpoints(VertexBase): """ Covers context caching endpoints for Vertex AI + Google AI Studio @@ -317,7 +333,6 @@ class ContextCachingEndpoints(VertexBase): ) return messages, optional_params, None - # Gemini requires a minimum of 1024 tokens for context caching. # Skip caching if the cached content is too small to avoid API errors. if not is_prompt_caching_valid_prompt( model=model, @@ -328,10 +343,11 @@ class ContextCachingEndpoints(VertexBase): verbose_logger.debug( "Vertex AI context caching: cached content is below minimum token " "count (%d). Skipping context caching.", - MINIMUM_PROMPT_CACHE_TOKEN_COUNT, + get_prompt_cache_min_tokens(model=model), ) return messages, optional_params, None + fallback_optional_params: Final = optional_params.copy() tools: Final = optional_params.pop("tools", None) tool_choice: Final = optional_params.pop("tool_choice", None) @@ -419,8 +435,8 @@ class ContextCachingEndpoints(VertexBase): ) response.raise_for_status() except httpx.HTTPStatusError as err: - error_code: Final = err.response.status_code - raise VertexAIError(status_code=error_code, message=err.response.text) + _raise_unless_cached_content_too_small(err) + return messages, fallback_optional_params, None except httpx.TimeoutException: raise VertexAIError(status_code=408, message="Timeout error occurred.") @@ -477,7 +493,6 @@ class ContextCachingEndpoints(VertexBase): ) return messages, optional_params, None - # Gemini requires a minimum of 1024 tokens for context caching. # Skip caching if the cached content is too small to avoid API errors. if not is_prompt_caching_valid_prompt( model=model, @@ -488,10 +503,11 @@ class ContextCachingEndpoints(VertexBase): verbose_logger.debug( "Vertex AI context caching: cached content is below minimum token " "count (%d). Skipping context caching.", - MINIMUM_PROMPT_CACHE_TOKEN_COUNT, + get_prompt_cache_min_tokens(model=model), ) return messages, optional_params, None + fallback_optional_params: Final = optional_params.copy() tools: Final = optional_params.pop("tools", None) tool_choice: Final = optional_params.pop("tool_choice", None) @@ -575,8 +591,8 @@ class ContextCachingEndpoints(VertexBase): ) response.raise_for_status() except httpx.HTTPStatusError as err: - error_code: Final = err.response.status_code - raise VertexAIError(status_code=error_code, message=err.response.text) + _raise_unless_cached_content_too_small(err) + return messages, fallback_optional_params, None except httpx.TimeoutException: raise VertexAIError(status_code=408, message="Timeout error occurred.") diff --git a/litellm/utils.py b/litellm/utils.py index a7d7447d7f3..9d37c54b4a9 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5490,7 +5490,13 @@ def get_max_tokens(model: str) -> int | None: def _strip_stable_vertex_version(model_name) -> str: - return re.sub(r"-\d+$", "", model_name) + stripped_region: Final = re.sub( + r"(^|/)(?:[a-z0-9_-]+\.)+(?=gem(?:ini|ma)-)", + r"\1", + model_name, + flags=re.IGNORECASE, + ) + return re.sub(r"-\d+$", "", stripped_region) _DATED_SNAPSHOT_SUFFIX: Final = re.compile(r"-\d{4}-\d{2}-\d{2}$") @@ -10058,6 +10064,25 @@ def should_use_cohere_v1_client(api_base: str | None, present_version_params: li return api_base.endswith("/v1/rerank") or (uses_v1_params and not api_base.endswith("/v2/rerank")) +_GEMINI_4096_PROMPT_CACHE_MIN_TOKENS: Final = 4096 + + +def _is_gemini_4096_cache_min_model(model: str) -> bool: + return bool( + re.search( + r"gemini-(?:3\.\d+-flash|2\.5-pro|3(?:\.\d+)?-pro)", + model.lower(), + ) + ) + + +def _lookup_model_prompt_cache_min_tokens(model: str) -> int | None: + try: + return get_model_info(model=model).get("prompt_cache_min_tokens") + except Exception: + return None + + def get_prompt_cache_min_tokens(model: str) -> int: """ Returns the smallest prefix `model` will actually cache. @@ -10073,10 +10098,15 @@ def get_prompt_cache_min_tokens(model: str) -> int: """ if MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE is not None: return MINIMUM_PROMPT_CACHE_TOKEN_COUNT_OVERRIDE - try: - min_tokens: Final = get_model_info(model=model).get("prompt_cache_min_tokens") - except Exception: - return DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT + normalized_model: Final = re.sub( + r"(^|/)(?:[a-z0-9_-]+\.)+(?=gem(?:ini|ma)-)", + r"\1", + model, + flags=re.IGNORECASE, + ) + min_tokens: Final = _lookup_model_prompt_cache_min_tokens(normalized_model) + if _is_gemini_4096_cache_min_model(normalized_model): + return max(min_tokens or 0, _GEMINI_4096_PROMPT_CACHE_MIN_TOKENS) if min_tokens is None: return DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT return min_tokens diff --git a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 67d78d6030e..26067076bdb 100644 --- a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1452,6 +1452,184 @@ class TestContextCachingEndpoints: self._token_check_patcher.start() + @pytest.mark.parametrize( + "model", + [ + "gemini-3.5-flash", + "au.gemini-3.5-flash", + "eu.gemini-3.5-flash", + "us.gemini-3.5-flash", + "global.gemini-3.5-flash", + "us-central1.gemini-3.5-flash", + "europe-west4.gemini-3.5-flash", + "australia-southeast1.gemini-3.5-flash", + "gemini-2.5-pro", + "au.gemini-2.5-pro", + "gemini-3.1-pro", + ], + ) + def test_check_and_create_cache_skips_below_4096_for_gemini_35_flash_and_25_pro_all_regions( + self, local_model_cost_map, model: str + ): + self._token_check_patcher.stop() + + cached_messages: Final = [ + { + "role": "system", + "content": " ".join(["word"] * 2200), + "cache_control": {"type": "ephemeral"}, + } + ] + non_cached_messages: Final = [{"role": "user", "content": "Hello"}] + all_messages: Final = cached_messages + non_cached_messages + + messages, _, returned_cache = self.context_caching.check_and_create_cache( + messages=all_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="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + assert messages == all_messages + assert returned_cache is None + self.mock_client.post.assert_not_called() + + self._token_check_patcher.start() + + @pytest.mark.parametrize( + "error_message", + [ + ( + "Cached content is too small. Labeller: tokens_count=1722. " + "The minimum token count to start explicit caching is 4096." + ), + "INVALID_ARGUMENT: Cached content is too small.", + ], + ) + def test_check_and_create_cache_falls_back_gracefully_on_400_cached_content_too_small( + self, + error_message: str, + ): + cached_messages: Final = [self.sample_messages[0]] + non_cached_messages: Final = [self.sample_messages[1]] + + mock_response: Final = MagicMock() + mock_response.status_code = 400 + mock_response.text = error_message + self.mock_client.post.side_effect = httpx.HTTPStatusError( + "Error", request=MagicMock(), response=mock_response + ) + + optional_params: Final = { + **self.sample_optional_params, + "tool_choice": {"functionCallingConfig": {"mode": "AUTO"}}, + } + + with ( + patch( + "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages", + return_value=(cached_messages, non_cached_messages), + ), + patch.object( + ContextCachingEndpoints, + "check_cache", + return_value=None, + ), + patch.object( + ContextCachingEndpoints, + "_get_token_and_url_context_caching", + return_value=("token", "https://test-url.com"), + ), + ): + messages, returned_params, returned_cache = self.context_caching.check_and_create_cache( + messages=self.sample_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-3.5-flash", + client=self.mock_client, + timeout=30.0, + logging_obj=self.mock_logging, + custom_llm_provider="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="vertex_test_token", + ) + + assert messages == self.sample_messages + assert returned_cache is None + assert returned_params.get("tools") == self.sample_tools + assert returned_params.get("tool_choice") == {"functionCallingConfig": {"mode": "AUTO"}} + + @pytest.mark.asyncio + async def test_async_check_and_create_cache_falls_back_gracefully_on_400_cached_content_too_small( + self, + ): + cached_messages: Final = [self.sample_messages[0]] + non_cached_messages: Final = [self.sample_messages[1]] + + mock_response: Final = MagicMock() + mock_response.status_code = 400 + mock_response.text = ( + "Cached content is too small. Labeller: tokens_count=1722. " + "The minimum token count to start explicit caching is 4096." + ) + self.mock_async_client.post.side_effect = httpx.HTTPStatusError( + "Error", request=MagicMock(), response=mock_response + ) + + optional_params: Final = { + **self.sample_optional_params, + "tool_choice": {"functionCallingConfig": {"mode": "AUTO"}}, + } + + with ( + patch( + "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages", + return_value=(cached_messages, non_cached_messages), + ), + patch.object( + ContextCachingEndpoints, + "async_check_cache", + return_value=None, + ), + patch.object( + ContextCachingEndpoints, + "_get_token_and_url_context_caching", + return_value=("token", "https://test-url.com"), + ), + ): + messages, returned_params, returned_cache = ( + await self.context_caching.async_check_and_create_cache( + messages=self.sample_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-3.5-flash", + client=self.mock_async_client, + timeout=30.0, + logging_obj=self.mock_logging, + custom_llm_provider="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="vertex_test_token", + ) + ) + + assert messages == self.sample_messages + assert returned_cache is None + assert returned_params.get("tools") == self.sample_tools + assert returned_params.get("tool_choice") == {"functionCallingConfig": {"mode": "AUTO"}} + @pytest.mark.parametrize( "custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"] ) diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index a72aec07766..a664da64fbc 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -4079,6 +4079,67 @@ def test_gemini_3_flash_and_31_pro_preview_resolve_4096_cache_minimum(local_mode assert not wrong, f"prompt_cache_min_tokens must be 4096: {wrong}" +@pytest.mark.parametrize("provider_prefix", ["", "gemini/", "vertex_ai/"]) +@pytest.mark.parametrize( + "region_prefix", + [ + "", + "au.", + "eu.", + "us.", + "ca.", + "jp.", + "uk.", + "in.", + "sg.", + "kr.", + "global.", + "apac.", + "us-central1.", + "europe-west4.", + "australia-southeast1.", + "asia-northeast1.", + ], +) +@pytest.mark.parametrize( + "base_model", + [ + "gemini-3.5-flash", + "gemini-3.5-flash-preview", + "gemini-2.5-pro", + "gemini-3-pro", + "gemini-3-pro-preview", + "gemini-3.1-pro", + "gemini-3.5-pro", + ], +) +def test_gemini_4096_cache_minimum_across_all_regions_and_pro_variants( + provider_prefix: str, + region_prefix: str, + base_model: str, + local_model_cost_map: None, +) -> None: + model: Final = f"{provider_prefix}{region_prefix}{base_model}" + assert get_prompt_cache_min_tokens(model=model) == 4096 + + +@pytest.mark.parametrize( + "model", + [ + "gemini-2.5-flash", + "au.gemini-2.5-flash", + "eu.gemini-2.5-flash", + "us.gemini-2.5-flash", + "us-central1.gemini-2.5-flash", + "vertex_ai/au.gemini-2.5-flash", + ], +) +def test_gemini_25_flash_resolves_1024_cache_minimum_all_regions( + model: str, local_model_cost_map: None +) -> None: + assert get_prompt_cache_min_tokens(model=model) == 1024 + + def test_get_prompt_cache_min_tokens_unmapped_model_falls_back_to_default(local_model_cost_map: None) -> None: """get_model_info raises for a model it has no entry for. The resolver must swallow that and fall back to the default, otherwise the raise reaches callers that would read it as @@ -4092,7 +4153,7 @@ def test_is_prompt_caching_valid_prompt_uses_per_model_minimum(local_model_cost_ the flat-1024 check reported claude-opus-4-6 as cacheable and the cache write was rejected upstream. Both assertions must live together: is_prompt_caching_valid_prompt returns False on any internal error, so the True case is what proves the False case isn't a swallowed exception.""" - token_count = litellm.token_counter( + token_count: Final = litellm.token_counter( model="claude-opus-4-6", messages=PROMPT_CACHE_MESSAGES, use_default_image_token_count=True ) assert 1024 <= token_count < 4096, ( @@ -4102,6 +4163,12 @@ def test_is_prompt_caching_valid_prompt_uses_per_model_minimum(local_model_cost_ assert is_prompt_caching_valid_prompt(model="claude-opus-4-6", messages=PROMPT_CACHE_MESSAGES) is False assert is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=PROMPT_CACHE_MESSAGES) is True + assert is_prompt_caching_valid_prompt(model="gemini-3.5-flash", messages=PROMPT_CACHE_MESSAGES) is False + assert is_prompt_caching_valid_prompt(model="au.gemini-3.5-flash", messages=PROMPT_CACHE_MESSAGES) is False + assert is_prompt_caching_valid_prompt(model="us-central1.gemini-3.5-flash", messages=PROMPT_CACHE_MESSAGES) is False + assert is_prompt_caching_valid_prompt(model="gemini-2.5-pro", messages=PROMPT_CACHE_MESSAGES) is False + assert is_prompt_caching_valid_prompt(model="gemini-2.5-flash", messages=PROMPT_CACHE_MESSAGES) is True + assert is_prompt_caching_valid_prompt(model="au.gemini-2.5-flash", messages=PROMPT_CACHE_MESSAGES) is True def test_is_prompt_caching_valid_prompt_explicit_min_token_count_overrides_model(local_model_cost_map: None) -> None: