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:
Victor Uceda 2026-06-03 11:48:47 +02:00
parent d45e9e4d56
commit 82ce13addf
2 changed files with 83 additions and 23 deletions

View file

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

View file

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