From 1ead3ae91fc1a701badd8462233ce4965a755ff0 Mon Sep 17 00:00:00 2001 From: Ondra Zahradnik Date: Fri, 18 Sep 2026 10:56:13 +0200 Subject: [PATCH] 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) --- .../test_vertex_ai_context_caching.py | 165 +++++++++--------- .../vertex_ai/gemini/test_transformation.py | 101 ++++++----- 2 files changed, 145 insertions(+), 121 deletions(-) diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index f410a797186..f8afb7bdeaf 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -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 ( diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py index 38e0a51c811..8c277d66afa 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py @@ -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)