This commit is contained in:
Zach Chrystall 2026-10-05 16:04:16 +03:00 • committed by GitHub
commit b53f429bce
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 307 additions and 16 deletions

View file

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

View file

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

View file

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

View file

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