update docs

This commit is contained in:
Otavio Brito 2025-10-05 23:47:24 -03:00
parent 80edf70206
commit 251d56ba5e
3 changed files with 309 additions and 18 deletions

View file

@ -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
<Tabs>
<TabItem value="sdk" label="SDK">
```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
```
</TabItem>
<TabItem value="sdk-ttl" label="SDK with Custom TTL">
```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)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
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}
</TabItem>
</Tabs>
#### 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`.

View file

@ -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,

View file

@ -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