From 251d56ba5e3db403e7e48d346543558ef59bcb28 Mon Sep 17 00:00:00 2001 From: Otavio Brito Date: Sun, 5 Oct 2025 23:47:24 -0300 Subject: [PATCH] update docs --- docs/my-website/docs/providers/vertex.md | 156 +++++++++++++++- .../vertex_ai_context_caching.py | 5 +- .../test_vertex_ai_context_caching.py | 166 ++++++++++++++++-- 3 files changed, 309 insertions(+), 18 deletions(-) diff --git a/docs/my-website/docs/providers/vertex.md b/docs/my-website/docs/providers/vertex.md index 3f4d1068958..e4a9c1b6b80 100644 --- a/docs/my-website/docs/providers/vertex.md +++ b/docs/my-website/docs/providers/vertex.md @@ -191,7 +191,7 @@ print(json.loads(completion.choices[0].message.content)) model_list: - model_name: gemini-2.5-pro litellm_params: - model: vertex_ai/gemini-1.5-pro + model: vertex_ai/gemini-2.5-pro vertex_project: "project-id" vertex_location: "us-central1" vertex_credentials: "/path/to/service_account.json" # [OPTIONAL] Do this OR `!gcloud auth application-default login` - run this to add vertex credentials to your env @@ -264,7 +264,7 @@ except JSONSchemaValidationError as e: model_list: - model_name: gemini-2.5-pro litellm_params: - model: vertex_ai/gemini-1.5-pro + model: vertex_ai/gemini-2.5-pro vertex_project: "project-id" vertex_location: "us-central1" vertex_credentials: "/path/to/service_account.json" # [OPTIONAL] Do this OR `!gcloud auth application-default login` - run this to add vertex credentials to your env @@ -811,11 +811,155 @@ curl http://0.0.0.0:4000/v1/chat/completions \ ### **Context Caching** -Use Vertex AI context caching is supported by calling provider api directly. (Unified Endpoint support coming soon.). +#### Unified Endpoint + +Use Vertex AI context caching in the same way as [**Google AI Studio - Context Caching**](../providers/gemini.md#context-caching) + + +##### Example usage + + + + +```python +from litellm import completion + +for _ in range(2): + resp = completion( + model="vertex_ai/gemini-2.5-pro", + messages=[ + # System Message + { + "role": "system", + "content": [ + { + "type": "text", + "text": "Here is the full text of a complex legal agreement" * 4000, + "cache_control": {"type": "ephemeral"}, # 👈 KEY CHANGE + } + ], + }, + # marked for caching with the cache_control parameter, so that this checkpoint can read from the previous cache. + { + "role": "user", + "content": [ + { + "type": "text", + "text": "What are the key terms and conditions in this agreement?", + "cache_control": {"type": "ephemeral"}, + } + ], + }] + ) + + print(resp.usage) # 👈 2nd usage block will be less, since cached tokens used +``` + + + + +```python +from litellm import completion + +# Cache for 2 hours (7200 seconds) +resp = completion( + model="vertex_ai/gemini-2.5-pro", + messages=[ + { + "role": "system", + "content": [ + { + "type": "text", + "text": "Here is the full text of a complex legal agreement" * 4000, + "cache_control": { + "type": "ephemeral", + "ttl": "7200s" # 👈 Cache for 2 hours + }, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "text", + "text": "What are the key terms and conditions in this agreement?", + "cache_control": { + "type": "ephemeral", + "ttl": "3600s" # 👈 This TTL will be ignored (first one is used) + }, + } + ], + } + ] +) + +print(resp.usage) +``` + + + + +1. Setup config.yaml + +```yaml +model_list: + - model_name: gemini-2.5-pro + litellm_params: + model: vertex_ai/gemini-2.5-pro + vertex_project: "project-id" + vertex_location: "us-central1" + vertex_credentials: "/path/to/service_account.json" +``` + +2. Start proxy + +```bash +litellm --config /path/to/config.yaml +``` + +3. Test it! + +```bash + +curl -X POST 'http://0.0.0.0:4000/chat/completions' \ +-H 'Content-Type: application/json' \ +-H 'Authorization: Bearer sk-1234' \ +-d '{ + "model": "gemini-2.5-flash", + "messages": [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": "Long cache message (must be >= 1025 tokens)", + "cache_control": { + "type": "ephemeral", + "ttl": "7200s" + } + } + ] + }, + { + "role": "user", + "content": [ + { + "type": "text", + "text": "What is the text about?" + } + ] + } + ] +}' + +``` + +#### Calling provider api directly [**Go straight to provider**](../pass_through/vertex_ai.md#context-caching) -#### 1. Create the Cache +##### 1. Create the Cache First, create the cache by sending a `POST` request to the `cachedContents` endpoint via the LiteLLM proxy. @@ -841,7 +985,7 @@ curl http://0.0.0.0:4000/vertex_ai/v1/projects/{project_id}/locations/{location} -#### 2. Get the Cache Name from the Response +##### 2. Get the Cache Name from the Response Vertex AI will return a response containing the `name` of the cached content. This name is the identifier for your cached data. @@ -860,7 +1004,7 @@ Vertex AI will return a response containing the `name` of the cached content. Th } ``` -#### 3. Use the Cached Content +##### 3. Use the Cached Content Use the `name` from the response as `cachedContent` or `cached_content` in subsequent API calls to reuse the cached information. This is passed in the body of your request to `/chat/completions`. 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 1747f2f67ab..70b068b5a4d 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 @@ -66,7 +66,10 @@ class ContextCachingEndpoints(VertexBase): endpoint = "cachedContents" url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}" else: - raise NotImplementedError + auth_header = vertex_auth_header + endpoint = "cachedContents" + url = f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}" + return self._check_custom_proxy( api_base=api_base, 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 ef3404353ae..0320092c7a3 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 @@ -55,6 +55,9 @@ class TestContextCachingEndpoints: self.sample_optional_params = {"tools": self.sample_tools.copy()} + @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" ) @@ -62,12 +65,14 @@ class TestContextCachingEndpoints: "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj" ) def test_check_and_create_cache_with_cached_content( - self, mock_cache_obj, mock_separate + self, mock_cache_obj, mock_separate, custom_llm_provider ): """Test check_and_create_cache when cached_content is provided""" # Setup cached_content = "cached_content_123" optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute result = self.context_caching.check_and_create_cache( @@ -80,6 +85,10 @@ class TestContextCachingEndpoints: timeout=30.0, logging_obj=self.mock_logging, cached_content=cached_content, + custom_llm_provider=custom_llm_provider, + vertex_project=test_project, + vertex_location=test_location, + vertex_auth_header="vertext_test_token", ) # Assert @@ -92,14 +101,21 @@ class TestContextCachingEndpoints: mock_separate.assert_not_called() mock_cache_obj.get_cache_key.assert_not_called() + @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_no_cached_messages(self, mock_separate): + def test_check_and_create_cache_no_cached_messages( + self, mock_separate, custom_llm_provider + ): """Test check_and_create_cache when no cached messages are found""" # Setup mock_separate.return_value = ([], self.sample_messages) # No cached messages optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute result = self.context_caching.check_and_create_cache( @@ -111,6 +127,10 @@ class TestContextCachingEndpoints: 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", ) # Assert @@ -119,6 +139,9 @@ class TestContextCachingEndpoints: assert returned_params == optional_params assert returned_cache is None + @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" ) @@ -127,7 +150,7 @@ class TestContextCachingEndpoints: ) @patch.object(ContextCachingEndpoints, "check_cache") def test_check_and_create_cache_existing_cache_found( - self, mock_check_cache, mock_cache_obj, mock_separate + self, mock_check_cache, mock_cache_obj, mock_separate, custom_llm_provider ): """Test check_and_create_cache when existing cache is found""" # Setup @@ -139,6 +162,8 @@ class TestContextCachingEndpoints: mock_check_cache.return_value = "existing_cache_name" optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute result = self.context_caching.check_and_create_cache( @@ -150,6 +175,10 @@ class TestContextCachingEndpoints: 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", ) # Assert @@ -163,6 +192,9 @@ class TestContextCachingEndpoints: messages=cached_messages, tools=self.sample_tools ) + @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" ) @@ -181,6 +213,7 @@ class TestContextCachingEndpoints: mock_transform, mock_cache_obj, mock_separate, + custom_llm_provider, ): """Test check_and_create_cache when creating new cache""" # Setup @@ -203,6 +236,8 @@ class TestContextCachingEndpoints: self.mock_client.post.return_value = mock_response optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute result = self.context_caching.check_and_create_cache( @@ -214,6 +249,10 @@ class TestContextCachingEndpoints: 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", ) # Assert @@ -228,6 +267,9 @@ class TestContextCachingEndpoints: assert "tools" in call_args.kwargs["json"] assert call_args.kwargs["json"]["tools"] == self.sample_tools + @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" ) @@ -237,7 +279,12 @@ class TestContextCachingEndpoints: @patch.object(ContextCachingEndpoints, "check_cache") @patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching") def test_check_and_create_cache_http_error( - self, mock_get_token_url, mock_check_cache, mock_cache_obj, mock_separate + self, + mock_get_token_url, + mock_check_cache, + mock_cache_obj, + mock_separate, + custom_llm_provider, ): """Test check_and_create_cache handles HTTP errors properly""" # Setup @@ -259,6 +306,8 @@ class TestContextCachingEndpoints: self.mock_client.post.side_effect = http_error optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute and Assert with pytest.raises(VertexAIError) as exc_info: @@ -271,12 +320,19 @@ class TestContextCachingEndpoints: 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", ) assert exc_info.value.status_code == 400 assert "Bad Request" in str(exc_info.value.message) @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" ) @@ -284,12 +340,14 @@ class TestContextCachingEndpoints: "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj" ) async def test_async_check_and_create_cache_with_cached_content( - self, mock_cache_obj, mock_separate + self, mock_cache_obj, mock_separate, custom_llm_provider ): """Test async_check_and_create_cache when cached_content is provided""" # Setup cached_content = "cached_content_123" optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -302,6 +360,10 @@ class TestContextCachingEndpoints: timeout=30.0, logging_obj=self.mock_logging, cached_content=cached_content, + custom_llm_provider=custom_llm_provider, + vertex_project=test_project, + vertex_location=test_location, + vertex_auth_header="vertext_test_token", ) # Assert @@ -311,14 +373,21 @@ class TestContextCachingEndpoints: assert returned_cache == cached_content @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" ) - async def test_async_check_and_create_cache_no_cached_messages(self, mock_separate): + async def test_async_check_and_create_cache_no_cached_messages( + self, mock_separate, custom_llm_provider + ): """Test async_check_and_create_cache when no cached messages are found""" # Setup mock_separate.return_value = ([], self.sample_messages) optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -330,6 +399,10 @@ class TestContextCachingEndpoints: 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", ) # Assert @@ -339,6 +412,9 @@ class TestContextCachingEndpoints: assert returned_cache is None @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" ) @@ -347,7 +423,7 @@ class TestContextCachingEndpoints: ) @patch.object(ContextCachingEndpoints, "async_check_cache") async def test_async_check_and_create_cache_existing_cache_found( - self, mock_async_check_cache, mock_cache_obj, mock_separate + self, mock_async_check_cache, mock_cache_obj, mock_separate, custom_llm_provider ): """Test async_check_and_create_cache when existing cache is found""" # Setup @@ -359,6 +435,8 @@ class TestContextCachingEndpoints: mock_async_check_cache.return_value = "existing_cache_name" optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -370,6 +448,10 @@ class TestContextCachingEndpoints: 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", ) # Assert @@ -384,6 +466,9 @@ class TestContextCachingEndpoints: ) @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" ) @@ -406,6 +491,7 @@ class TestContextCachingEndpoints: mock_transform, mock_cache_obj, mock_separate, + custom_llm_provider, ): """Test async_check_and_create_cache when creating new cache""" # Setup @@ -428,6 +514,8 @@ class TestContextCachingEndpoints: self.mock_async_client.post = AsyncMock(return_value=mock_response) optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -439,6 +527,10 @@ class TestContextCachingEndpoints: 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", ) # Assert @@ -454,6 +546,9 @@ class TestContextCachingEndpoints: assert call_args.kwargs["json"]["tools"] == self.sample_tools @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" ) @@ -472,6 +567,7 @@ class TestContextCachingEndpoints: mock_async_check_cache, mock_cache_obj, mock_separate, + custom_llm_provider, ): """Test async_check_and_create_cache handles timeout errors properly""" # Setup @@ -489,6 +585,8 @@ class TestContextCachingEndpoints: ) optional_params = self.sample_optional_params.copy() + test_project = "test_project" + test_location = "test_location" # Execute and Assert with pytest.raises(VertexAIError) as exc_info: @@ -501,12 +599,21 @@ class TestContextCachingEndpoints: 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", ) assert exc_info.value.status_code == 408 assert "Timeout error occurred" in str(exc_info.value.message) - def test_check_and_create_cache_tools_popped_from_optional_params(self): + @pytest.mark.parametrize( + "custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"] + ) + def test_check_and_create_cache_tools_popped_from_optional_params( + self, custom_llm_provider + ): """Test that tools are properly popped from optional_params when there are cached messages""" with patch( "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" @@ -520,6 +627,8 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() original_tools = optional_params["tools"].copy() + test_project = "test_project" + test_location = "test_location" # Mock the check_cache to return existing cache so we don't make HTTP calls with patch.object( @@ -535,6 +644,10 @@ class TestContextCachingEndpoints: 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", ) # Assert tools were popped from optional_params @@ -543,7 +656,12 @@ class TestContextCachingEndpoints: # But original tools should still be available for comparison assert original_tools == self.sample_tools - def test_check_and_create_cache_tools_not_popped_when_no_cached_messages(self): + @pytest.mark.parametrize( + "custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"] + ) + def test_check_and_create_cache_tools_not_popped_when_no_cached_messages( + self, custom_llm_provider + ): """Test that tools are NOT popped from optional_params when there are no cached messages""" with patch( "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" @@ -555,6 +673,8 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() original_tools = optional_params["tools"].copy() + test_project = "test_project" + test_location = "test_location" # Execute result = self.context_caching.check_and_create_cache( @@ -566,6 +686,10 @@ class TestContextCachingEndpoints: 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", ) # Assert tools were NOT popped from optional_params (early return) @@ -573,8 +697,11 @@ class TestContextCachingEndpoints: assert optional_params["tools"] == original_tools @pytest.mark.asyncio + @pytest.mark.parametrize( + "custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"] + ) async def test_async_check_and_create_cache_tools_not_popped_when_no_cached_messages( - self, + self, custom_llm_provider ): """Test that tools are NOT popped from optional_params in async version when there are no cached messages""" with patch( @@ -587,6 +714,8 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() original_tools = optional_params["tools"].copy() + test_project = "test_project" + test_location = "test_location" # Execute result = await self.context_caching.async_check_and_create_cache( @@ -598,6 +727,10 @@ class TestContextCachingEndpoints: 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", ) # Assert tools were NOT popped from optional_params (early return) @@ -605,7 +738,12 @@ class TestContextCachingEndpoints: assert optional_params["tools"] == original_tools @pytest.mark.asyncio - async def test_async_check_and_create_cache_tools_popped_from_optional_params(self): + @pytest.mark.parametrize( + "custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"] + ) + async def test_async_check_and_create_cache_tools_popped_from_optional_params( + self, custom_llm_provider + ): """Test that tools are properly popped from optional_params in async version when there are cached messages""" with patch( "litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages" @@ -619,6 +757,8 @@ class TestContextCachingEndpoints: optional_params = self.sample_optional_params.copy() original_tools = optional_params["tools"].copy() + test_project = "test_project" + test_location = "test_location" # Mock the async_check_cache to return existing cache so we don't make HTTP calls with patch.object( @@ -634,6 +774,10 @@ class TestContextCachingEndpoints: 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", ) # Assert tools were popped from optional_params