mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(vertex_ai): use cachedContents collection endpoint for gemini custom api_base
This commit is contained in:
parent
c274cf321c
commit
0c084ff0ab
2 changed files with 121 additions and 14 deletions
|
|
@ -1,3 +1,4 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import List, Literal, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
|
@ -50,20 +51,24 @@ class ContextCachingEndpoints(VertexBase):
|
|||
vertex_location: Optional[str],
|
||||
vertex_auth_header: Optional[str],
|
||||
model: Optional[str] = None,
|
||||
) -> Tuple[Optional[str], str]:
|
||||
) -> Tuple[Union[str, Mapping[str, str], None], str]:
|
||||
"""
|
||||
Internal function. Returns the token and url for the call.
|
||||
|
||||
Handles logic if it's google ai studio vs. vertex ai.
|
||||
|
||||
For Google AI Studio the credential is a header mapping, since the API key is sent as
|
||||
`x-goog-api-key` instead of a bearer token.
|
||||
|
||||
Returns
|
||||
token, url
|
||||
"""
|
||||
auth_header: Optional[str]
|
||||
if custom_llm_provider == "gemini":
|
||||
auth_header = {"x-goog-api-key": gemini_api_key} # type: ignore[assignment]
|
||||
endpoint = "cachedContents"
|
||||
url = "https://generativelanguage.googleapis.com/v1beta/{}".format(endpoint)
|
||||
gemini_auth_header = {"x-goog-api-key": gemini_api_key} if gemini_api_key is not None else None
|
||||
if api_base:
|
||||
return gemini_auth_header, "{}/cachedContents".format(api_base.rstrip("/"))
|
||||
return gemini_auth_header, "https://generativelanguage.googleapis.com/v1beta/cachedContents"
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
auth_header = vertex_auth_header
|
||||
endpoint = "cachedContents"
|
||||
|
|
@ -339,7 +344,7 @@ class ContextCachingEndpoints(VertexBase):
|
|||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
if isinstance(token, dict):
|
||||
if isinstance(token, Mapping):
|
||||
headers.update(token)
|
||||
elif token is not None:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
|
|
@ -490,7 +495,7 @@ class ContextCachingEndpoints(VertexBase):
|
|||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
if isinstance(token, dict):
|
||||
if isinstance(token, Mapping):
|
||||
headers.update(token)
|
||||
elif token is not None:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
|
|
|
|||
|
|
@ -1881,26 +1881,30 @@ class TestVertexAIGlobalLocation:
|
|||
"global-aiplatform" not in url
|
||||
), "URL should not contain 'global-aiplatform' prefix"
|
||||
|
||||
def test_gemini_context_caching_with_custom_api_base_passes_model(self):
|
||||
"""Gemini context caching with custom api_base must pass model to _check_custom_proxy.
|
||||
@pytest.mark.parametrize(
|
||||
"api_base", ["https://my-proxy.example.com/genai/v1beta", "https://my-proxy.example.com/genai/v1beta/"]
|
||||
)
|
||||
def test_gemini_context_caching_with_custom_api_base_uses_collection_endpoint(self, api_base):
|
||||
"""Gemini context caching with custom api_base must hit the `cachedContents` collection.
|
||||
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/23846
|
||||
Previously model was hardcoded to None, causing ValueError when api_base was set.
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/34872 (and #23846, which
|
||||
required the call to not raise when api_base is set). `cachedContents` is not a model
|
||||
action, so the model must not be spliced into the url.
|
||||
"""
|
||||
caching = ContextCachingEndpoints()
|
||||
|
||||
auth_header, url = caching._get_token_and_url_context_caching(
|
||||
gemini_api_key="test-key",
|
||||
custom_llm_provider="gemini",
|
||||
api_base="https://my-proxy.example.com",
|
||||
api_base=api_base,
|
||||
vertex_project=None,
|
||||
vertex_location=None,
|
||||
vertex_auth_header=None,
|
||||
model="gemini-1.5-pro",
|
||||
model="gemini-2.5-flash",
|
||||
)
|
||||
|
||||
assert "models/gemini-1.5-pro" in url
|
||||
assert url.startswith("https://my-proxy.example.com/")
|
||||
assert url == "https://my-proxy.example.com/genai/v1beta/cachedContents"
|
||||
assert auth_header == {"x-goog-api-key": "test-key"}
|
||||
|
||||
def test_gemini_context_caching_without_api_base_ignores_model(self):
|
||||
"""Without custom api_base, model param is not needed (default URL is used)."""
|
||||
|
|
@ -1971,3 +1975,101 @@ class TestContextCachingMultiRegionUrls:
|
|||
|
||||
assert url.startswith("https://aiplatform.googleapis.com/")
|
||||
assert "/locations/global/cachedContents" in url
|
||||
|
||||
|
||||
class TestGeminiCustomApiBaseCacheRequests:
|
||||
"""Regression coverage for #34872: with a custom Gemini `api_base`, cache list/create
|
||||
requests must go to `{api_base}/cachedContents` with the api key in `x-goog-api-key`,
|
||||
never a stringified dict in an `Authorization` header."""
|
||||
|
||||
EXPECTED_URL = "https://my-proxy.example.com/genai/v1beta/cachedContents"
|
||||
|
||||
def setup_method(self):
|
||||
self.caching = ContextCachingEndpoints()
|
||||
self.logging_obj = MagicMock(spec=Logging)
|
||||
self.messages = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Stable cacheable prefix. " * 800,
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "Reply with ok."},
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _empty_list_response():
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"cachedContents": []}
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _create_response():
|
||||
response = MagicMock()
|
||||
response.json.return_value = {
|
||||
"name": "cachedContents/abc123",
|
||||
"model": "models/gemini-2.5-flash",
|
||||
}
|
||||
return response
|
||||
|
||||
def _assert_requests(self, get_call, post_call):
|
||||
assert get_call.kwargs["url"] == self.EXPECTED_URL
|
||||
assert post_call.kwargs["url"] == self.EXPECTED_URL
|
||||
assert post_call.kwargs["json"]["model"] == "models/gemini-2.5-flash"
|
||||
for headers in (get_call.kwargs["headers"], post_call.kwargs["headers"]):
|
||||
assert headers["x-goog-api-key"] == "sk-gemini-key"
|
||||
assert headers["x-custom"] == "1"
|
||||
assert "Authorization" not in headers
|
||||
|
||||
def test_sync_cache_lifecycle_uses_collection_url_and_api_key_header(self):
|
||||
client = MagicMock(spec=HTTPHandler)
|
||||
client.get.return_value = self._empty_list_response()
|
||||
client.post.return_value = self._create_response()
|
||||
|
||||
_, _, cache_name = self.caching.check_and_create_cache(
|
||||
messages=self.messages,
|
||||
optional_params={},
|
||||
api_key="sk-gemini-key",
|
||||
api_base="https://my-proxy.example.com/genai/v1beta",
|
||||
model="gemini-2.5-flash",
|
||||
client=client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.logging_obj,
|
||||
custom_llm_provider="gemini",
|
||||
vertex_project=None,
|
||||
vertex_location=None,
|
||||
vertex_auth_header=None,
|
||||
extra_headers={"x-custom": "1"},
|
||||
)
|
||||
|
||||
assert cache_name == "cachedContents/abc123"
|
||||
self._assert_requests(client.get.call_args, client.post.call_args)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_cache_lifecycle_uses_collection_url_and_api_key_header(self):
|
||||
client = MagicMock(spec=AsyncHTTPHandler)
|
||||
client.get = AsyncMock(return_value=self._empty_list_response())
|
||||
client.post = AsyncMock(return_value=self._create_response())
|
||||
|
||||
_, _, cache_name = await self.caching.async_check_and_create_cache(
|
||||
messages=self.messages,
|
||||
optional_params={},
|
||||
api_key="sk-gemini-key",
|
||||
api_base="https://my-proxy.example.com/genai/v1beta",
|
||||
model="gemini-2.5-flash",
|
||||
client=client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.logging_obj,
|
||||
custom_llm_provider="gemini",
|
||||
vertex_project=None,
|
||||
vertex_location=None,
|
||||
vertex_auth_header=None,
|
||||
extra_headers={"x-custom": "1"},
|
||||
)
|
||||
|
||||
assert cache_name == "cachedContents/abc123"
|
||||
self._assert_requests(client.get.call_args, client.post.call_args)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue