Merge pull request #35172 from BerriAI/litellm_vertex_cache_skip_tool_final

fix(vertex_ai): skip context caching when the cached block ends on a model turn
This commit is contained in:
Mateo Wang 2026-07-29 20:24:42 -07:00 • committed by GitHub
commit 47f1fb394e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 183 additions and 1 deletions

View file

@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
"""
import re
from typing import List, Optional, Tuple, Literal
from typing import List, Optional, Sequence, Tuple, Literal
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai import CachedContentRequestBody
@ -152,6 +152,20 @@ def separate_cached_messages(
return cached_messages, non_cached_messages
def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool:
"""
The cachedContents API rejects contents ending on a model turn, which is how it
classifies both assistant messages and tool results, with HTTP 400
"Requests ending with a model turn are not supported". System messages are
extracted into system_instruction before contents are built, so the terminal
turn is the last non-system message.
"""
non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system")
if not non_system_messages:
return bool(cached_messages)
return non_system_messages[-1].get("role") not in ("assistant", "tool", "function")
def transform_openai_messages_to_gemini_context_caching(
model: str,
messages: List[AllMessageValues],

View file

@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import (
from ..common_utils import VertexAIError, get_vertex_base_url
from ..vertex_llm_base import VertexBase
from .transformation import (
cached_messages_end_on_supported_turn,
separate_cached_messages,
transform_openai_messages_to_gemini_context_caching,
)
@ -308,6 +309,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
if not cached_messages_end_on_supported_turn(cached_messages):
verbose_logger.debug(
"Vertex AI context caching: cached message block ends on a model turn once "
"system messages are extracted, which the cachedContents API rejects. "
"Skipping context caching."
)
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(
@ -459,6 +468,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
if not cached_messages_end_on_supported_turn(cached_messages):
verbose_logger.debug(
"Vertex AI context caching: cached message block ends on a model turn once "
"system messages are extracted, which the cachedContents API rejects. "
"Skipping context caching."
)
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(

View file

@ -1452,6 +1452,157 @@ class TestContextCachingEndpoints:
# Restart the patcher so teardown_method can stop it cleanly
self._token_check_patcher.start()
def _model_turn_final_messages(self, final_cached_role):
tool_call = {
"id": "call_abc123",
"type": "function",
"function": {"name": "get_weather", "arguments": '{"location": "Boston"}'},
}
cached_tail = {
"assistant": [],
"tool": [
{
"role": "tool",
"tool_call_id": "call_abc123",
"content": "72F and sunny",
"cache_control": {"type": "ephemeral"},
}
],
"system": [
{
"role": "system",
"content": "Tool results are authoritative.",
"cache_control": {"type": "ephemeral"},
}
],
}[final_cached_role]
return [
{
"role": "user",
"content": [
{
"type": "text",
"text": "Use the weather tool for every answer.",
"cache_control": {"type": "ephemeral"},
}
],
},
{
"role": "assistant",
"content": "",
"tool_calls": [tool_call],
"cache_control": {"type": "ephemeral"},
},
*cached_tail,
{"role": "user", "content": "What is the weather in Boston?"},
]
@pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
self, final_cached_role
):
"""The cachedContents API rejects contents ending on an assistant or tool turn
with HTTP 400 "Requests ending with a model turn are not supported", so the
request must proceed uncached instead of failing.
"""
all_messages = self._model_turn_final_messages(final_cached_role)
optional_params = self.sample_optional_params.copy()
result = self.context_caching.check_and_create_cache(
messages=all_messages,
optional_params=optional_params,
api_key="test_key",
api_base=None,
model="gemini-3.6-flash",
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",
)
messages, returned_params, returned_cache = result
assert messages == all_messages
assert returned_cache is None
assert "tools" in returned_params
self.mock_client.get.assert_not_called()
self.mock_client.post.assert_not_called()
@pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
@pytest.mark.asyncio
async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
self, final_cached_role
):
"""Async variant: an unsupported terminal turn skips caching instead of failing."""
all_messages = self._model_turn_final_messages(final_cached_role)
optional_params = self.sample_optional_params.copy()
result = await self.context_caching.async_check_and_create_cache(
messages=all_messages,
optional_params=optional_params,
api_key="test_key",
api_base=None,
model="gemini-3.6-flash",
client=self.mock_async_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",
)
messages, returned_params, returned_cache = result
assert messages == all_messages
assert returned_cache is None
assert "tools" in returned_params
self.mock_async_client.get.assert_not_called()
self.mock_async_client.post.assert_not_called()
def test_cached_messages_end_on_supported_turn():
from litellm.llms.vertex_ai.context_caching.transformation import (
cached_messages_end_on_supported_turn,
)
assert (
cached_messages_end_on_supported_turn(
[{"role": "assistant", "content": "hi"}, {"role": "user", "content": "hello"}]
)
is True
)
assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True
assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False
assert (
cached_messages_end_on_supported_turn(
[
{"role": "user", "content": "hello"},
{"role": "assistant", "content": "hi"},
{"role": "system", "content": "be brief"},
]
)
is False
)
assert (
cached_messages_end_on_supported_turn(
[{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}]
)
is True
)
assert (
cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}])
is False
)
assert (
cached_messages_end_on_supported_turn([{"role": "function", "name": "f", "content": "y"}])
is False
)
assert cached_messages_end_on_supported_turn([]) is False
class TestCheckCachePagination:
"""Test pagination logic in check_cache and async_check_cache methods."""