test(vertex_ai): type the CMEK test helpers and stop mutating caller collections

The cachedContents fakes took a caller-owned list and appended to it,
which mutates a function parameter, and none of the new helpers carried
annotations. Replace both transports with a `_FakeCachedContents` class
that owns its recording and grows an immutable tuple of create bodies,
and annotate every new helper and test.

The kwargs bag in the gemini test is gone: its only precise annotation
would be a mutable collection or Any, so the two transform calls now
pass their arguments explicitly instead.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Ondra Zahradnik 2026-09-18 10:56:13 +02:00
parent 5e2fa50724
commit 1ead3ae91f
2 changed files with 145 additions and 121 deletions

View file

@ -1,5 +1,6 @@
import json
from typing import List
from collections.abc import Mapping, Sequence
from typing import Final, List
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -15,6 +16,34 @@ from litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching import (
)
CREATED_CACHE_NAME: Final = "freshly_created_cache"
class _FakeCachedContents:
"""Fake Google: GET lists the caches in `already_cached`, POST records the create body and succeeds."""
def __init__(self, already_cached: Sequence[str] = ()) -> None:
self._already_cached: Final = tuple(already_cached)
self.created: tuple[Mapping[str, object], ...] = ()
@property
def transport(self) -> httpx.MockTransport:
return httpx.MockTransport(self._handle)
def _handle(self, request: httpx.Request) -> httpx.Response:
if request.method == "GET":
return httpx.Response(
200,
json={
"cachedContents": [
{"name": f"cache-for-{name}", "displayName": name} for name in self._already_cached
]
},
)
self.created = (*self.created, json.loads(request.content))
return httpx.Response(200, json={"name": CREATED_CACHE_NAME, "model": "gemini-2.5-pro"})
class TestContextCachingEndpoints:
"""Test class for ContextCachingEndpoints methods"""
@ -1561,7 +1590,7 @@ class TestContextCachingEndpoints:
KMS_KEY_NAME = "projects/test_project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache"
def _cmek_call_kwargs(self, custom_llm_provider):
def _cmek_call_kwargs(self, custom_llm_provider: str) -> Mapping[str, object]:
return {
"messages": [
{
@ -1588,95 +1617,69 @@ class TestContextCachingEndpoints:
"vertex_auth_header": "test_token",
}
def _cached_contents_transport(self, posted_bodies):
"""Fake Google: the cache list is empty (GET) and every create (POST) succeeds."""
def handle(request: httpx.Request) -> httpx.Response:
if request.method == "GET":
return httpx.Response(200, json={})
posted_bodies.append(json.loads(request.content))
return httpx.Response(200, json={"name": "new_cache_name", "model": "gemini-2.5-pro"})
return httpx.MockTransport(handle)
def _assert_kms_key_only_adds_encryption_spec(self, posted_bodies):
def _assert_kms_key_only_adds_encryption_spec(self, posted_bodies: Sequence[Mapping[str, object]]) -> None:
plain_body, cmek_body = posted_bodies
assert "encryptionSpec" not in plain_body
assert cmek_body["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME}
assert cmek_body["displayName"] != plain_body["displayName"]
untouched = ("encryptionSpec", "displayName")
untouched: Final = ("encryptionSpec", "displayName")
assert {k: v for k, v in cmek_body.items() if k not in untouched} == {
k: v for k, v in plain_body.items() if k not in untouched
}
assert plain_body["contents"][0]["parts"][0]["text"] == "Cached reference material"
assert plain_body["model"].endswith("models/gemini-2.5-pro")
assert plain_body["contents"] == [{"role": "user", "parts": [{"text": "Cached reference material"}]}]
assert str(plain_body["model"]).endswith("models/gemini-2.5-pro")
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
def test_check_and_create_cache_kms_key_adds_encryption_spec(self, custom_llm_provider):
def test_check_and_create_cache_kms_key_adds_encryption_spec(self, custom_llm_provider: str) -> None:
"""kms_key_name becomes Vertex's `encryptionSpec` on the cache-creation POST and changes nothing else.
It is forwarded for every provider on purpose: Google AI Studio has no CMEK, so it rejects the field
instead of silently caching the content without the customer's key.
"""
posted_bodies = []
client = HTTPHandler()
client.client = httpx.Client(transport=self._cached_contents_transport(posted_bodies))
kwargs = self._cmek_call_kwargs(custom_llm_provider)
fake: Final = _FakeCachedContents()
client: Final = HTTPHandler()
client.client = httpx.Client(transport=fake.transport)
kwargs: Final = self._cmek_call_kwargs(custom_llm_provider)
self.context_caching.check_and_create_cache(client=client, **kwargs)
_, _, returned_cache = self.context_caching.check_and_create_cache(
client=client, kms_key_name=self.KMS_KEY_NAME, **kwargs
)
assert returned_cache == "new_cache_name"
self._assert_kms_key_only_adds_encryption_spec(posted_bodies)
assert returned_cache == CREATED_CACHE_NAME
self._assert_kms_key_only_adds_encryption_spec(fake.created)
@pytest.mark.asyncio
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
async def test_async_check_and_create_cache_kms_key_adds_encryption_spec(self, custom_llm_provider):
async def test_async_check_and_create_cache_kms_key_adds_encryption_spec(self, custom_llm_provider: str) -> None:
"""Async variant of test_check_and_create_cache_kms_key_adds_encryption_spec."""
posted_bodies = []
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=self._cached_contents_transport(posted_bodies))
kwargs = self._cmek_call_kwargs(custom_llm_provider)
fake: Final = _FakeCachedContents()
client: Final = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=fake.transport)
kwargs: Final = self._cmek_call_kwargs(custom_llm_provider)
await self.context_caching.async_check_and_create_cache(client=client, **kwargs)
_, _, returned_cache = await self.context_caching.async_check_and_create_cache(
client=client, kms_key_name=self.KMS_KEY_NAME, **kwargs
)
assert returned_cache == "new_cache_name"
self._assert_kms_key_only_adds_encryption_spec(posted_bodies)
def _existing_cache_transport(self, listed_display_names, hits):
"""Fake Google holding caches already created under `listed_display_names`."""
assert returned_cache == CREATED_CACHE_NAME
self._assert_kms_key_only_adds_encryption_spec(fake.created)
def handle(request: httpx.Request) -> httpx.Response:
if request.method == "GET":
return httpx.Response(
200,
json={
"cachedContents": [
{"name": f"cache-for-{name}", "displayName": name} for name in listed_display_names
]
},
)
hits.append(json.loads(request.content))
return httpx.Response(200, json={"name": "freshly_created_cache", "model": "gemini-2.5-pro"})
return httpx.MockTransport(handle)
def _display_name_for(self, custom_llm_provider, kms_key_name):
recorded = []
client = HTTPHandler()
client.client = httpx.Client(transport=self._cached_contents_transport(recorded))
def _display_name_for(self, custom_llm_provider: str, kms_key_name: str | None) -> str:
fake: Final = _FakeCachedContents()
client: Final = HTTPHandler()
client.client = httpx.Client(transport=fake.transport)
self.context_caching.check_and_create_cache(
client=client, kms_key_name=kms_key_name, **self._cmek_call_kwargs(custom_llm_provider)
)
return recorded[0]["displayName"]
return str(fake.created[0]["displayName"])
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
def test_check_and_create_cache_does_not_reuse_cache_encrypted_under_another_policy(self, custom_llm_provider):
def test_check_and_create_cache_does_not_reuse_cache_encrypted_under_another_policy(
self, custom_llm_provider: str
) -> None:
"""A CMEK request must not reuse content cached with no key or a different key.
`encryptionSpec` is input-only, so a cache listing never reveals which key an existing entry
@ -1691,75 +1694,75 @@ class TestContextCachingEndpoints:
assert len({unencrypted_name, other_key_name, wanted_name}) == 3
for already_cached in ([], [unencrypted_name], [other_key_name], [unencrypted_name, other_key_name]):
created = []
for already_cached in ((), (unencrypted_name,), (other_key_name,), (unencrypted_name, other_key_name)):
fake = _FakeCachedContents(already_cached)
client = HTTPHandler()
client.client = httpx.Client(transport=self._existing_cache_transport(already_cached, created))
client.client = httpx.Client(transport=fake.transport)
_, _, returned_cache = self.context_caching.check_and_create_cache(
client=client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider)
)
assert returned_cache == "freshly_created_cache", f"reused a foreign-policy cache from {already_cached}"
assert created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME}
assert returned_cache == CREATED_CACHE_NAME, f"reused a foreign-policy cache from {already_cached}"
assert fake.created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME}
reused = []
client = HTTPHandler()
client.client = httpx.Client(transport=self._existing_cache_transport([wanted_name], reused))
reusable: Final = _FakeCachedContents((wanted_name,))
reuse_client: Final = HTTPHandler()
reuse_client.client = httpx.Client(transport=reusable.transport)
_, _, returned_cache = self.context_caching.check_and_create_cache(
client=client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider)
client=reuse_client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider)
)
assert returned_cache == f"cache-for-{wanted_name}"
assert reused == []
assert reusable.created == ()
@pytest.mark.asyncio
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
async def test_async_check_and_create_cache_does_not_reuse_cache_encrypted_under_another_policy(
self, custom_llm_provider
):
self, custom_llm_provider: str
) -> None:
"""Async variant: a cache created without CMEK is not reused for a CMEK request."""
unencrypted_name = self._display_name_for(custom_llm_provider, None)
wanted_name = self._display_name_for(custom_llm_provider, self.KMS_KEY_NAME)
created = []
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=self._existing_cache_transport([unencrypted_name], created))
fake: Final = _FakeCachedContents((unencrypted_name,))
client: Final = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=fake.transport)
_, _, returned_cache = await self.context_caching.async_check_and_create_cache(
client=client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider)
)
assert returned_cache == "freshly_created_cache"
assert created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME}
assert returned_cache == CREATED_CACHE_NAME
assert fake.created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME}
reused = []
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=self._existing_cache_transport([wanted_name], reused))
reusable: Final = _FakeCachedContents((wanted_name,))
reuse_client: Final = AsyncHTTPHandler()
reuse_client.client = httpx.AsyncClient(transport=reusable.transport)
_, _, returned_cache = await self.context_caching.async_check_and_create_cache(
client=client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider)
client=reuse_client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider)
)
assert returned_cache == f"cache-for-{wanted_name}"
assert reused == []
assert reusable.created == ()
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"])
def test_check_and_create_cache_without_kms_key_still_reuses_existing_cache(self, custom_llm_provider):
def test_check_and_create_cache_without_kms_key_still_reuses_existing_cache(self, custom_llm_provider: str) -> None:
"""Scoping must not break plain caching: a no-key request still hits a no-key cache."""
unencrypted_name = self._display_name_for(custom_llm_provider, None)
created = []
client = HTTPHandler()
client.client = httpx.Client(transport=self._existing_cache_transport([unencrypted_name], created))
fake: Final = _FakeCachedContents((unencrypted_name,))
client: Final = HTTPHandler()
client.client = httpx.Client(transport=fake.transport)
_, _, returned_cache = self.context_caching.check_and_create_cache(
client=client, **self._cmek_call_kwargs(custom_llm_provider)
)
assert returned_cache == f"cache-for-{unencrypted_name}"
assert created == []
assert fake.created == ()
def test_cached_messages_end_on_supported_turn():
from litellm.llms.vertex_ai.context_caching.transformation import (

View file

@ -1,6 +1,8 @@
import json
import uuid
from collections.abc import Mapping
from typing import Final
from unittest.mock import Mock
import httpx
@ -14,6 +16,7 @@ from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
)
from litellm.types.llms import openai
from litellm.types import completion
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai import RequestBody
@ -477,23 +480,28 @@ async def test_vertex_ai_async_transform_inlines_only_the_urls_gemini_cannot_fet
assert sorted(async_only_image_fetch.fetched) == sorted([plain_http_png, extensionless_https])
KMS_KEY_NAME = "projects/qa-project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache"
CACHE_NAME = "projects/qa-project/locations/us-central1/cachedContents/123"
KMS_KEY_NAME: Final = "projects/qa-project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache"
CACHE_NAME: Final = "projects/qa-project/locations/us-central1/cachedContents/123"
def _cmek_cache_transport(captured):
class _FakeCachedContents:
"""Fake Vertex: the cache list is empty, so the cachedContents create is exercised for real."""
def handle(request: httpx.Request) -> httpx.Response:
def __init__(self) -> None:
self.created: tuple[Mapping[str, object], ...] = ()
@property
def transport(self) -> httpx.MockTransport:
return httpx.MockTransport(self._handle)
def _handle(self, request: httpx.Request) -> httpx.Response:
if request.method == "GET":
return httpx.Response(200, json={})
captured["cache_create"] = json.loads(request.content)
self.created = (*self.created, json.loads(request.content))
return httpx.Response(200, json={"name": CACHE_NAME, "model": "gemini-2.5-flash"})
return httpx.MockTransport(handle)
def _cacheable_messages():
def _cacheable_messages() -> list[AllMessageValues]: # mutable-ok: the transform signature takes a list
return [
{
"role": "user",
@ -509,48 +517,61 @@ def _cacheable_messages():
]
def _transform_kwargs():
return {
"gemini_api_key": None,
"messages": _cacheable_messages(),
"api_base": None,
"model": "gemini-2.5-flash",
"timeout": None,
"extra_headers": None,
"optional_params": {"kms_key_name": KMS_KEY_NAME},
"logging_obj": Mock(),
"custom_llm_provider": "vertex_ai",
"litellm_params": {},
"vertex_project": "qa-project",
"vertex_location": "us-central1",
"vertex_auth_header": "qa-token",
}
def _assert_key_encrypts_cache_without_leaking(captured, body):
assert captured["cache_create"]["encryptionSpec"] == {"kmsKeyName": KMS_KEY_NAME}
def _assert_key_encrypts_cache_without_leaking(fake: _FakeCachedContents, body: Mapping[str, object]) -> None:
(cache_create,) = fake.created
assert cache_create["encryptionSpec"] == {"kmsKeyName": KMS_KEY_NAME}
assert body["cachedContent"] == CACHE_NAME
assert KMS_KEY_NAME not in json.dumps(body)
def test_sync_transform_request_body_forwards_kms_key_name_to_cache_creation():
def test_sync_transform_request_body_forwards_kms_key_name_to_cache_creation() -> None:
"""`kms_key_name` reaches the cachedContents POST as encryptionSpec and never the generateContent body."""
captured = {}
client = HTTPHandler()
client.client = httpx.Client(transport=_cmek_cache_transport(captured))
fake: Final = _FakeCachedContents()
client: Final = HTTPHandler()
client.client = httpx.Client(transport=fake.transport)
body = transformation.sync_transform_request_body(client=client, **_transform_kwargs())
body: Final = transformation.sync_transform_request_body(
client=client,
gemini_api_key=None,
messages=_cacheable_messages(),
api_base=None,
model="gemini-2.5-flash",
timeout=None,
extra_headers=None,
optional_params={"kms_key_name": KMS_KEY_NAME},
logging_obj=Mock(),
custom_llm_provider="vertex_ai",
litellm_params={},
vertex_project="qa-project",
vertex_location="us-central1",
vertex_auth_header="qa-token",
)
_assert_key_encrypts_cache_without_leaking(captured, body)
_assert_key_encrypts_cache_without_leaking(fake, body)
@pytest.mark.asyncio
async def test_async_transform_request_body_forwards_kms_key_name_to_cache_creation():
async def test_async_transform_request_body_forwards_kms_key_name_to_cache_creation() -> None:
"""The async transform path pops and forwards the key independently of the sync one."""
captured = {}
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=_cmek_cache_transport(captured))
fake: Final = _FakeCachedContents()
client: Final = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=fake.transport)
body = await transformation.async_transform_request_body(client=client, **_transform_kwargs())
body: Final = await transformation.async_transform_request_body(
client=client,
gemini_api_key=None,
messages=_cacheable_messages(),
api_base=None,
model="gemini-2.5-flash",
timeout=None,
extra_headers=None,
optional_params={"kms_key_name": KMS_KEY_NAME},
logging_obj=Mock(),
custom_llm_provider="vertex_ai",
litellm_params={},
vertex_project="qa-project",
vertex_location="us-central1",
vertex_auth_header="qa-token",
)
_assert_key_encrypts_cache_without_leaking(captured, body)
_assert_key_encrypts_cache_without_leaking(fake, body)