diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index d1b15a4e036..bd2aace185f 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -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, diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 7b575638c44..26a8d6d4dad 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -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):