mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
update docs
This commit is contained in:
parent
80edf70206
commit
251d56ba5e
3 changed files with 309 additions and 18 deletions
|
|
@ -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`.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue