mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
47f1fb394e
3 changed files with 183 additions and 1 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue