mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
fix(gemini): include tool_choice in cache key computation
The cache key for context caching was computed from messages, tools, and model, but did not include tool_choice. This meant two requests with the same messages/tools/model but different tool_choice values (e.g. 'auto' vs 'required') would share the same cache key, silently reusing cached content created with a different toolConfig. Add tool_choice to the get_cache_key() call in both sync and async paths so that different tool_choice values produce distinct cache entries. Add tests verifying: - tool_choice is passed to cache key computation - toolConfig appears in the HTTP request body when creating new cached content (sync and async)
This commit is contained in:
parent
501b8fa445
commit
a655c53600
2 changed files with 165 additions and 6 deletions
|
|
@ -358,7 +358,9 @@ class ContextCachingEndpoints(VertexBase):
|
|||
client = client
|
||||
|
||||
## CHECK IF CACHED ALREADY
|
||||
generated_cache_key = local_cache_obj.get_cache_key(messages=cached_messages, tools=tools, model=model)
|
||||
generated_cache_key = local_cache_obj.get_cache_key(
|
||||
messages=cached_messages, tools=tools, tool_choice=tool_choice, model=model
|
||||
)
|
||||
google_cache_name = self.check_cache(
|
||||
cache_key=generated_cache_key,
|
||||
client=client,
|
||||
|
|
@ -500,7 +502,9 @@ class ContextCachingEndpoints(VertexBase):
|
|||
client = client
|
||||
|
||||
## CHECK IF CACHED ALREADY
|
||||
generated_cache_key = local_cache_obj.get_cache_key(messages=cached_messages, tools=tools, model=model)
|
||||
generated_cache_key = local_cache_obj.get_cache_key(
|
||||
messages=cached_messages, tools=tools, tool_choice=tool_choice, model=model
|
||||
)
|
||||
google_cache_name = await self.async_check_cache(
|
||||
cache_key=generated_cache_key,
|
||||
client=client,
|
||||
|
|
|
|||
|
|
@ -178,9 +178,9 @@ class TestContextCachingEndpoints:
|
|||
assert returned_params == optional_params
|
||||
assert returned_cache == "existing_cache_name"
|
||||
|
||||
# Verify cache key was generated with tools and model
|
||||
# Verify cache key was generated with tools, tool_choice, and model
|
||||
mock_cache_obj.get_cache_key.assert_called_once_with(
|
||||
messages=cached_messages, tools=self.sample_tools, model="gemini-1.5-pro"
|
||||
messages=cached_messages, tools=self.sample_tools, tool_choice=None, model="gemini-1.5-pro"
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
|
|
@ -419,9 +419,9 @@ class TestContextCachingEndpoints:
|
|||
assert returned_params == optional_params
|
||||
assert returned_cache == "existing_cache_name"
|
||||
|
||||
# Verify cache key was generated with tools and model
|
||||
# Verify cache key was generated with tools, tool_choice, and model
|
||||
mock_cache_obj.get_cache_key.assert_called_once_with(
|
||||
messages=cached_messages, tools=self.sample_tools, model="gemini-1.5-pro"
|
||||
messages=cached_messages, tools=self.sample_tools, tool_choice=None, model="gemini-1.5-pro"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -831,6 +831,161 @@ class TestContextCachingEndpoints:
|
|||
assert "tool_choice" in optional_params
|
||||
assert optional_params["tool_choice"] == "auto"
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
@patch("litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages")
|
||||
@patch("litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj")
|
||||
@patch.object(ContextCachingEndpoints, "check_cache")
|
||||
def test_cache_key_includes_tool_choice(self, mock_check_cache, mock_cache_obj, mock_separate, custom_llm_provider):
|
||||
"""Test that tool_choice is included in the cache key computation.
|
||||
|
||||
Different tool_choice values (e.g. 'auto' vs 'required') must produce
|
||||
different cache keys, otherwise the second call silently reuses cached
|
||||
content created with a different toolConfig.
|
||||
"""
|
||||
cached_messages = [self.sample_messages[0]]
|
||||
non_cached_messages = [self.sample_messages[1]]
|
||||
mock_separate.return_value = (cached_messages, non_cached_messages)
|
||||
|
||||
mock_cache_obj.get_cache_key.return_value = "test_cache_key"
|
||||
mock_check_cache.return_value = "existing_cache"
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
optional_params["tool_choice"] = "required"
|
||||
|
||||
self.context_caching.check_and_create_cache(
|
||||
messages=self.sample_messages,
|
||||
optional_params=optional_params,
|
||||
api_key="test_key",
|
||||
api_base=None,
|
||||
model="gemini-1.5-pro",
|
||||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project="test_project",
|
||||
vertex_location="test_location",
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
mock_cache_obj.get_cache_key.assert_called_once_with(
|
||||
messages=cached_messages, tools=self.sample_tools, tool_choice="required", model="gemini-1.5-pro"
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
@patch("litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages")
|
||||
@patch("litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj")
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.transform_openai_messages_to_gemini_context_caching"
|
||||
)
|
||||
@patch.object(ContextCachingEndpoints, "check_cache")
|
||||
@patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching")
|
||||
def test_create_new_cache_includes_tool_config_in_body(
|
||||
self,
|
||||
mock_get_token_url,
|
||||
mock_check_cache,
|
||||
mock_transform,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test that toolConfig is included in the HTTP request body when creating new cached content."""
|
||||
cached_messages = [self.sample_messages[0]]
|
||||
non_cached_messages = [self.sample_messages[1]]
|
||||
mock_separate.return_value = (cached_messages, non_cached_messages)
|
||||
|
||||
mock_cache_obj.get_cache_key.return_value = "test_cache_key"
|
||||
mock_check_cache.return_value = None
|
||||
mock_get_token_url.return_value = ("token", "https://test-url.com")
|
||||
mock_transform.return_value = {"model": "gemini-1.5-pro", "contents": []}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"name": "new_cache_name",
|
||||
"model": "gemini-1.5-pro",
|
||||
}
|
||||
self.mock_client.post.return_value = mock_response
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
optional_params["tool_choice"] = "required"
|
||||
|
||||
self.context_caching.check_and_create_cache(
|
||||
messages=self.sample_messages,
|
||||
optional_params=optional_params,
|
||||
api_key="test_key",
|
||||
api_base=None,
|
||||
model="gemini-1.5-pro",
|
||||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project="test_project",
|
||||
vertex_location="test_location",
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
call_args = self.mock_client.post.call_args
|
||||
assert call_args.kwargs["json"]["tools"] == self.sample_tools
|
||||
assert call_args.kwargs["json"]["toolConfig"] == "required"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
@patch("litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages")
|
||||
@patch("litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj")
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.transform_openai_messages_to_gemini_context_caching"
|
||||
)
|
||||
@patch.object(ContextCachingEndpoints, "async_check_cache")
|
||||
@patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching")
|
||||
@patch("litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.get_async_httpx_client")
|
||||
async def test_async_create_new_cache_includes_tool_config_in_body(
|
||||
self,
|
||||
mock_get_client,
|
||||
mock_get_token_url,
|
||||
mock_async_check_cache,
|
||||
mock_transform,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test that toolConfig is included in the HTTP request body when creating new cached content (async)."""
|
||||
cached_messages = [self.sample_messages[0]]
|
||||
non_cached_messages = [self.sample_messages[1]]
|
||||
mock_separate.return_value = (cached_messages, non_cached_messages)
|
||||
|
||||
mock_cache_obj.get_cache_key.return_value = "test_cache_key"
|
||||
mock_async_check_cache.return_value = None
|
||||
mock_get_token_url.return_value = ("token", "https://test-url.com")
|
||||
mock_transform.return_value = {"model": "gemini-1.5-pro", "contents": []}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"name": "new_cache_name",
|
||||
"model": "gemini-1.5-pro",
|
||||
}
|
||||
self.mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
optional_params["tool_choice"] = "required"
|
||||
|
||||
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-1.5-pro",
|
||||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project="test_project",
|
||||
vertex_location="test_location",
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
call_args = self.mock_async_client.post.call_args
|
||||
assert call_args.kwargs["json"]["tools"] == self.sample_tools
|
||||
assert call_args.kwargs["json"]["toolConfig"] == "required"
|
||||
|
||||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
|
||||
@patch("litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages")
|
||||
def test_check_and_create_cache_skips_when_below_min_tokens(self, mock_separate, custom_llm_provider):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue