mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(vertex_ai): enforce per-model cache token minimums and handle 400 fallback
- Add get_minimum_prompt_cache_token_count() in litellm/utils.py to enforce the 4,096-token minimum for explicit context caching on Gemini 3.5 Flash, Gemini 2.5 Pro, and Gemini 3.x Pro models (defaulting to 1,024 for other models unless overridden via MINIMUM_PROMPT_CACHE_TOKEN_COUNT). - Update is_prompt_caching_valid_prompt() and Vertex AI context caching debug logging to use the model-aware minimum. - Add a graceful fallback in check_and_create_cache and async_check_and_create_cache when Vertex AI returns HTTP 400 because cached content is below the server-side minimum token count, restoring cache_control markers so standard generateContent succeeds without raising VertexAIError. - Add unit tests in tests/unit/test_utils.py and tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py. Fixes #17696
This commit is contained in:
parent
6a8e0a270a
commit
0cb2bb6503
4 changed files with 307 additions and 16 deletions
|
|
@ -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.")
|
||||
|
||||
|
|
|
|||
|
|
@ -5471,7 +5471,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}$")
|
||||
|
|
@ -10057,6 +10063,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.
|
||||
|
|
@ -10072,10 +10097,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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4075,6 +4075,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
|
||||
|
|
@ -4088,7 +4149,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, (
|
||||
|
|
@ -4098,6 +4159,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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue