mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(vertex_ai): use REP host for context caching on eu/us multi-region endpoints
Context caching built the cachedContents URL as https://{location}-aiplatform.googleapis.com, which is an invalid host for the eu/us multi-region endpoints and returns 404. The inference path already resolves these to the REP host (https://aiplatform.{geo}.rep.googleapis.com) via get_vertex_base_url(); reuse that helper in _get_token_and_url_context_caching so caching uses the same host as inference. Adds tests covering the eu/us multi-region cachedContents URLs (v1 and v1beta1). Fixes #29571
This commit is contained in:
parent
d45e9e4d56
commit
82ce13addf
2 changed files with 83 additions and 23 deletions
|
|
@ -19,7 +19,7 @@ from litellm.types.llms.vertex_ai import (
|
|||
VertexAICachedContentResponseObject,
|
||||
)
|
||||
|
||||
from ..common_utils import VertexAIError
|
||||
from ..common_utils import VertexAIError, get_vertex_base_url
|
||||
from ..vertex_llm_base import VertexBase
|
||||
from .transformation import (
|
||||
separate_cached_messages,
|
||||
|
|
@ -69,17 +69,13 @@ class ContextCachingEndpoints(VertexBase):
|
|||
elif custom_llm_provider == "vertex_ai":
|
||||
auth_header = vertex_auth_header
|
||||
endpoint = "cachedContents"
|
||||
if vertex_location == "global":
|
||||
url = f"https://aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
else:
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
base_url = get_vertex_base_url(vertex_location)
|
||||
url = f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
else:
|
||||
auth_header = vertex_auth_header
|
||||
endpoint = "cachedContents"
|
||||
if vertex_location == "global":
|
||||
url = f"https://aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
else:
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
base_url = get_vertex_base_url(vertex_location)
|
||||
url = f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
|
||||
|
||||
return self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
|
|
|
|||
|
|
@ -86,7 +86,7 @@ class TestContextCachingEndpoints:
|
|||
cached_content = "cached_content_123"
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -129,7 +129,7 @@ class TestContextCachingEndpoints:
|
|||
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"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -177,7 +177,7 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -251,7 +251,7 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -321,7 +321,7 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute and Assert
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
|
|
@ -361,7 +361,7 @@ class TestContextCachingEndpoints:
|
|||
cached_content = "cached_content_123"
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -401,7 +401,7 @@ class TestContextCachingEndpoints:
|
|||
mock_separate.return_value = ([], self.sample_messages)
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -450,7 +450,7 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -529,7 +529,7 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -600,7 +600,7 @@ class TestContextCachingEndpoints:
|
|||
|
||||
optional_params = self.sample_optional_params.copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute and Assert
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
|
|
@ -642,7 +642,7 @@ class TestContextCachingEndpoints:
|
|||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Mock the check_cache to return existing cache so we don't make HTTP calls
|
||||
with patch.object(
|
||||
|
|
@ -688,7 +688,7 @@ class TestContextCachingEndpoints:
|
|||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
|
|
@ -729,7 +729,7 @@ class TestContextCachingEndpoints:
|
|||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Execute
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
|
|
@ -772,7 +772,7 @@ class TestContextCachingEndpoints:
|
|||
optional_params = self.sample_optional_params.copy()
|
||||
original_tools = optional_params["tools"].copy()
|
||||
test_project = "test_project"
|
||||
test_location = "test_location"
|
||||
test_location = "us-central1"
|
||||
|
||||
# Mock the async_check_cache to return existing cache so we don't make HTTP calls
|
||||
with patch.object(
|
||||
|
|
@ -1335,6 +1335,70 @@ class TestVertexAIGlobalLocation:
|
|||
"global-aiplatform" not in url
|
||||
), "URL should not contain 'global-aiplatform' prefix"
|
||||
|
||||
@pytest.mark.parametrize("vertex_location", ["eu", "us"])
|
||||
def test_multi_region_location_url_construction_v1(self, vertex_location):
|
||||
"""Multi-region locations (eu/us) must use the REP host for v1 API.
|
||||
|
||||
Regression test: previously context caching built
|
||||
``https://{location}-aiplatform.googleapis.com`` for every non-global
|
||||
location, producing the invalid host ``eu-aiplatform.googleapis.com``
|
||||
(404) for the multi-region endpoints. Inference already resolved these
|
||||
to ``aiplatform.{geo}.rep.googleapis.com`` via ``get_vertex_base_url``;
|
||||
context caching now uses the same helper.
|
||||
"""
|
||||
caching = ContextCachingEndpoints()
|
||||
|
||||
with patch.object(
|
||||
caching,
|
||||
"_check_custom_proxy",
|
||||
side_effect=lambda **kwargs: (kwargs.get("auth_header"), kwargs.get("url")),
|
||||
):
|
||||
auth_header, url = caching._get_token_and_url_context_caching(
|
||||
gemini_api_key=None,
|
||||
custom_llm_provider="vertex_ai",
|
||||
api_base=None,
|
||||
vertex_project="test-project",
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header="Bearer test-token",
|
||||
)
|
||||
|
||||
expected_url = (
|
||||
f"https://aiplatform.{vertex_location}.rep.googleapis.com"
|
||||
f"/v1/projects/test-project/locations/{vertex_location}/cachedContents"
|
||||
)
|
||||
assert url == expected_url, f"Expected {expected_url}, got {url}"
|
||||
assert (
|
||||
f"{vertex_location}-aiplatform" not in url
|
||||
), "URL must not use the invalid '<location>-aiplatform' host for multi-region"
|
||||
|
||||
@pytest.mark.parametrize("vertex_location", ["eu", "us"])
|
||||
def test_multi_region_location_url_construction_v1beta1(self, vertex_location):
|
||||
"""Multi-region locations (eu/us) must use the REP host for v1beta1 API."""
|
||||
caching = ContextCachingEndpoints()
|
||||
|
||||
with patch.object(
|
||||
caching,
|
||||
"_check_custom_proxy",
|
||||
side_effect=lambda **kwargs: (kwargs.get("auth_header"), kwargs.get("url")),
|
||||
):
|
||||
auth_header, url = caching._get_token_and_url_context_caching(
|
||||
gemini_api_key=None,
|
||||
custom_llm_provider="vertex_ai_beta",
|
||||
api_base=None,
|
||||
vertex_project="test-project",
|
||||
vertex_location=vertex_location,
|
||||
vertex_auth_header="Bearer test-token",
|
||||
)
|
||||
|
||||
expected_url = (
|
||||
f"https://aiplatform.{vertex_location}.rep.googleapis.com"
|
||||
f"/v1beta1/projects/test-project/locations/{vertex_location}/cachedContents"
|
||||
)
|
||||
assert url == expected_url, f"Expected {expected_url}, got {url}"
|
||||
assert (
|
||||
f"{vertex_location}-aiplatform" not in url
|
||||
), "URL must not use the invalid '<location>-aiplatform' host for multi-region"
|
||||
|
||||
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.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue