From 943b029ef1b09acdc0b1c254bf9f1c4ebfbf2494 Mon Sep 17 00:00:00 2001 From: Ondra Zahradnik Date: Wed, 9 Sep 2026 14:02:31 +0200 Subject: [PATCH 1/3] feat(vertex_ai): support CMEK for explicit context caching Forward a `kms_key_name` param to the Vertex AI cachedContents create call as `encryptionSpec.kmsKeyName`, so caches created by LiteLLM are protected by the customer-managed Cloud KMS key. The param is popped from optional_params next to `cached_content`, so it never reaches the generateContent body. Google AI Studio has no CMEK, so the field is forwarded there too and rejected upstream rather than silently caching the content unencrypted --- .../context_caching/transformation.py | 11 ++- .../vertex_ai_context_caching.py | 4 + .../llms/vertex_ai/gemini/transformation.py | 2 + litellm/types/llms/vertex_ai.py | 11 +++ .../test_vertex_ai_context_caching.py | 85 +++++++++++++++++++ .../vertex_ai/gemini/test_transformation.py | 52 +++++++++++- 6 files changed, 162 insertions(+), 3 deletions(-) diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index e23374d57a1..834f5319402 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -9,7 +9,11 @@ from collections.abc import Sequence from typing import Final, Literal from litellm.types.llms.openai import AllMessageValues -from litellm.types.llms.vertex_ai import CachedContentRequestBody +from litellm.types.llms.vertex_ai import ( + CachedContentRequestBody, + EncryptedCachedContentRequestBody, + EncryptionSpec, +) from litellm.utils import is_cached_message from ..common_utils import get_supports_system_message @@ -174,6 +178,7 @@ def transform_openai_messages_to_gemini_context_caching( cache_key: str, vertex_project: str | None, vertex_location: str | None, + kms_key_name: str | None = None, ) -> CachedContentRequestBody: # Extract TTL from cached messages BEFORE system message transformation ttl: Final = extract_ttl_from_cached_messages(messages) @@ -208,4 +213,6 @@ def transform_openai_messages_to_gemini_context_caching( if transformed_system_messages is not None: data["system_instruction"] = transformed_system_messages - return data + if kms_key_name is None: + return data + return EncryptedCachedContentRequestBody(**data, encryptionSpec=EncryptionSpec(kmsKeyName=kms_key_name)) diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 75d4ffbed86..14b06cc1ad1 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -290,6 +290,7 @@ class ContextCachingEndpoints(VertexBase): vertex_auth_header: str | None, extra_headers: dict | None = None, cached_content: str | None = None, + kms_key_name: str | None = None, ) -> tuple[list[AllMessageValues], dict, str | None]: """ Receives @@ -393,6 +394,7 @@ class ContextCachingEndpoints(VertexBase): custom_llm_provider=custom_llm_provider, vertex_project=vertex_project, vertex_location=vertex_location, + kms_key_name=kms_key_name, ) cached_content_request_body["tools"] = tools @@ -449,6 +451,7 @@ class ContextCachingEndpoints(VertexBase): vertex_auth_header: str | None, extra_headers: dict | None = None, cached_content: str | None = None, + kms_key_name: str | None = None, ) -> tuple[list[AllMessageValues], dict, str | None]: """ Receives @@ -548,6 +551,7 @@ class ContextCachingEndpoints(VertexBase): custom_llm_provider=custom_llm_provider, vertex_project=vertex_project, vertex_location=vertex_location, + kms_key_name=kms_key_name, ) cached_content_request_body["tools"] = tools diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 13e2238fdf6..0f9b3f1e100 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -1293,6 +1293,7 @@ def sync_transform_request_body( timeout=timeout, extra_headers=extra_headers, cached_content=optional_params.pop("cached_content", None), + kms_key_name=optional_params.pop("kms_key_name", None), logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, vertex_project=vertex_project, @@ -1361,6 +1362,7 @@ async def async_transform_request_body( timeout=timeout, extra_headers=extra_headers, cached_content=optional_params.pop("cached_content", None), + kms_key_name=optional_params.pop("kms_key_name", None), logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, vertex_project=vertex_project, diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 3b95b786631..d620de55c69 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -2,6 +2,7 @@ from enum import Enum from typing import Any, Final, Literal, Protocol from typing_extensions import ( + ReadOnly, Required, TypedDict, ) @@ -348,6 +349,10 @@ class RequestBody(TypedDict, total=False): serviceTier: str +class EncryptionSpec(TypedDict): + kmsKeyName: ReadOnly[str] # Format: projects/{project}/locations/{location}/keyRings/{ring}/cryptoKeys/{key} + + class CachedContentRequestBody(TypedDict, total=False): contents: Required[list[ContentType]] system_instruction: SystemInstructions @@ -358,6 +363,12 @@ class CachedContentRequestBody(TypedDict, total=False): displayName: str +class EncryptedCachedContentRequestBody(CachedContentRequestBody, total=False): + """Vertex AI only: Google AI Studio's cachedContents has no encryptionSpec.""" + + encryptionSpec: ReadOnly[EncryptionSpec] + + class CachedContentListAllResponseBody(TypedDict, total=False): cachedContents: list[CachedContent] nextPageToken: str 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 f666829d2e8..26034f4a53d 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,3 +1,4 @@ +import json from typing import List from unittest.mock import AsyncMock, MagicMock, patch @@ -1558,6 +1559,90 @@ class TestContextCachingEndpoints: self.mock_async_client.get.assert_not_called() self.mock_async_client.post.assert_not_called() + KMS_KEY_NAME = "projects/test_project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache" + + def _cmek_call_kwargs(self, custom_llm_provider): + return { + "messages": [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "Cached reference material", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + {"role": "user", "content": "Question about the material"}, + ], + "optional_params": {}, + "api_key": "test_key", + "api_base": None, + "model": "gemini-2.5-pro", + "timeout": 30.0, + "logging_obj": self.mock_logging, + "custom_llm_provider": custom_llm_provider, + "vertex_project": "test_project", + "vertex_location": "us-central1", + "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): + plain_body, cmek_body = posted_bodies + assert "encryptionSpec" not in plain_body + assert cmek_body["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME} + assert {k: v for k, v in cmek_body.items() if k != "encryptionSpec"} == plain_body + assert plain_body["contents"][0]["parts"][0]["text"] == "Cached reference material" + assert 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): + """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) + + 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) + + @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 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) + + 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 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 f135acd094f..62686b69026 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py @@ -7,7 +7,7 @@ import httpx import pytest import litellm -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.vertex_ai.gemini import transformation from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, @@ -475,3 +475,53 @@ async def test_vertex_ai_async_transform_inlines_only_the_urls_gemini_cannot_fet {"file_data": {"mime_type": "application/pdf", "file_uri": files_api_pdf}}, ] assert sorted(async_only_image_fetch.fetched) == sorted([plain_http_png, extensionless_https]) + + +def test_sync_transform_request_body_forwards_kms_key_name_to_cache_creation(): + """`kms_key_name` reaches the cachedContents POST as encryptionSpec and never the generateContent body.""" + kms_key_name = "projects/qa-project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache" + cache_name = "projects/qa-project/locations/us-central1/cachedContents/123" + captured = {} + + def handle(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response(200, json={}) + captured["cache_create"] = json.loads(request.content) + return httpx.Response(200, json={"name": cache_name, "model": "gemini-2.5-flash"}) + + client = HTTPHandler() + client.client = httpx.Client(transport=httpx.MockTransport(handle)) + messages = [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": " ".join(f"clause {i}" for i in range(2000)), + "cache_control": {"type": "ephemeral"}, + } + ], + }, + {"role": "user", "content": "Which clause covers termination?"}, + ] + + body = transformation.sync_transform_request_body( + gemini_api_key=None, + messages=messages, + api_base=None, + model="gemini-2.5-flash", + client=client, + 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 captured["cache_create"]["encryptionSpec"] == {"kmsKeyName": kms_key_name} + assert body["cachedContent"] == cache_name + assert kms_key_name not in json.dumps(body) From 5e2fa507240504d73fa881515b38e59bd27e4bcd Mon Sep 17 00:00:00 2001 From: Ondra Zahradnik Date: Wed, 9 Sep 2026 14:43:19 +0200 Subject: [PATCH 2/3] fix(vertex_ai): scope context cache identity by the CMEK key The cachedContents displayName is the only thing the cache lookup can match on, and `encryptionSpec` is input-only, so a listing never reveals which key an existing cache uses. Content already cached under a Google-managed key, or under a different customer key, was therefore reused for a request that asked for a specific key, leaving that content outside the caller's encryption policy. Fold the key into the displayName so a cache is only reused when it was created under the very same key. Requests without a key keep the unscoped name, so plain caching still hits as before. --- .../context_caching/transformation.py | 14 +++ .../vertex_ai_context_caching.py | 11 +- litellm/types/llms/vertex_ai.py | 2 +- .../test_vertex_ai_context_caching.py | 119 +++++++++++++++++- .../vertex_ai/gemini/test_transformation.py | 85 ++++++++----- 5 files changed, 197 insertions(+), 34 deletions(-) diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index 834f5319402..77b74b7d0ec 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -4,6 +4,7 @@ Transformation logic for context caching. Why separate file? Make it easy to see how transformation works """ +import hashlib import re from collections.abc import Sequence from typing import Final, Literal @@ -119,6 +120,19 @@ def _is_valid_ttl_format(ttl: str) -> bool: return False +def scope_cache_key_to_encryption_key(cache_key: str, kms_key_name: str | None) -> str: + """ + Namespace the cache's displayName by the CMEK key, since displayName is the only thing + check_cache can match on and `encryptionSpec` is input-only, so Google never tells us + which key an existing cache uses. Without this, content already cached under a + Google-managed key (or a different CMEK key) would be reused for a request that asked + for a specific key, silently escaping the caller's encryption policy. + """ + if kms_key_name is None: + return cache_key + return f"{cache_key}-cmek-{hashlib.sha256(kms_key_name.encode()).hexdigest()[:16]}" + + def separate_cached_messages( messages: list[AllMessageValues], ) -> tuple[list[AllMessageValues], list[AllMessageValues]]: diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 14b06cc1ad1..2e85bc7e815 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -23,6 +23,7 @@ from ..common_utils import VertexAIError, get_vertex_base_url from ..vertex_llm_base import VertexBase from .transformation import ( cached_messages_end_on_supported_turn, + scope_cache_key_to_encryption_key, separate_cached_messages, transform_openai_messages_to_gemini_context_caching, ) @@ -367,8 +368,9 @@ class ContextCachingEndpoints(VertexBase): client = client ## CHECK IF CACHED ALREADY - generated_cache_key: Final = local_cache_obj.get_cache_key( - messages=cached_messages, tools=tools, tool_choice=tool_choice, model=model + generated_cache_key: Final = scope_cache_key_to_encryption_key( + local_cache_obj.get_cache_key(messages=cached_messages, tools=tools, tool_choice=tool_choice, model=model), + kms_key_name, ) google_cache_name: Final = self.check_cache( cache_key=generated_cache_key, @@ -523,8 +525,9 @@ class ContextCachingEndpoints(VertexBase): client = client ## CHECK IF CACHED ALREADY - generated_cache_key: Final = local_cache_obj.get_cache_key( - messages=cached_messages, tools=tools, tool_choice=tool_choice, model=model + generated_cache_key: Final = scope_cache_key_to_encryption_key( + local_cache_obj.get_cache_key(messages=cached_messages, tools=tools, tool_choice=tool_choice, model=model), + kms_key_name, ) google_cache_name: Final = await self.async_check_cache( cache_key=generated_cache_key, diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index d620de55c69..46aa39bff54 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -350,7 +350,7 @@ class RequestBody(TypedDict, total=False): class EncryptionSpec(TypedDict): - kmsKeyName: ReadOnly[str] # Format: projects/{project}/locations/{location}/keyRings/{ring}/cryptoKeys/{key} + kmsKeyName: ReadOnly[str] class CachedContentRequestBody(TypedDict, total=False): 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 26034f4a53d..f410a797186 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 @@ -1603,7 +1603,12 @@ class TestContextCachingEndpoints: plain_body, cmek_body = posted_bodies assert "encryptionSpec" not in plain_body assert cmek_body["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME} - assert {k: v for k, v in cmek_body.items() if k != "encryptionSpec"} == plain_body + assert cmek_body["displayName"] != plain_body["displayName"] + + untouched = ("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") @@ -1643,6 +1648,118 @@ class TestContextCachingEndpoints: 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`.""" + + 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)) + 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"] + + @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): + """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 + uses. The displayName is therefore scoped by the key, and the only reusable cache is one + created under the very same key. + """ + unencrypted_name = self._display_name_for(custom_llm_provider, None) + other_key_name = self._display_name_for( + custom_llm_provider, "projects/test_project/locations/us-central1/keyRings/litellm/cryptoKeys/rotated" + ) + wanted_name = self._display_name_for(custom_llm_provider, self.KMS_KEY_NAME) + + 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 = [] + client = HTTPHandler() + client.client = httpx.Client(transport=self._existing_cache_transport(already_cached, created)) + + _, _, 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} + + reused = [] + client = HTTPHandler() + client.client = httpx.Client(transport=self._existing_cache_transport([wanted_name], reused)) + + _, _, 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 == f"cache-for-{wanted_name}" + assert reused == [] + + @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 + ): + """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)) + + _, _, 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} + + reused = [] + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=self._existing_cache_transport([wanted_name], reused)) + + _, _, 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 == f"cache-for-{wanted_name}" + assert reused == [] + + @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): + """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)) + + _, _, 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 == [] 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 62686b69026..38e0a51c811 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py @@ -477,21 +477,24 @@ 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]) -def test_sync_transform_request_body_forwards_kms_key_name_to_cache_creation(): - """`kms_key_name` reaches the cachedContents POST as encryptionSpec and never the generateContent body.""" - kms_key_name = "projects/qa-project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache" - cache_name = "projects/qa-project/locations/us-central1/cachedContents/123" - captured = {} +KMS_KEY_NAME = "projects/qa-project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache" +CACHE_NAME = "projects/qa-project/locations/us-central1/cachedContents/123" + + +def _cmek_cache_transport(captured): + """Fake Vertex: the cache list is empty, so the cachedContents create is exercised for real.""" def handle(request: httpx.Request) -> httpx.Response: if request.method == "GET": return httpx.Response(200, json={}) captured["cache_create"] = json.loads(request.content) - return httpx.Response(200, json={"name": cache_name, "model": "gemini-2.5-flash"}) + return httpx.Response(200, json={"name": CACHE_NAME, "model": "gemini-2.5-flash"}) - client = HTTPHandler() - client.client = httpx.Client(transport=httpx.MockTransport(handle)) - messages = [ + return httpx.MockTransport(handle) + + +def _cacheable_messages(): + return [ { "role": "user", "content": [ @@ -505,23 +508,49 @@ def test_sync_transform_request_body_forwards_kms_key_name_to_cache_creation(): {"role": "user", "content": "Which clause covers termination?"}, ] - body = transformation.sync_transform_request_body( - gemini_api_key=None, - messages=messages, - api_base=None, - model="gemini-2.5-flash", - client=client, - 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 captured["cache_create"]["encryptionSpec"] == {"kmsKeyName": kms_key_name} - assert body["cachedContent"] == cache_name - assert kms_key_name not in json.dumps(body) +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} + 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(): + """`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)) + + body = transformation.sync_transform_request_body(client=client, **_transform_kwargs()) + + _assert_key_encrypts_cache_without_leaking(captured, body) + + +@pytest.mark.asyncio +async def test_async_transform_request_body_forwards_kms_key_name_to_cache_creation(): + """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)) + + body = await transformation.async_transform_request_body(client=client, **_transform_kwargs()) + + _assert_key_encrypts_cache_without_leaking(captured, body) From 1ead3ae91fc1a701badd8462233ce4965a755ff0 Mon Sep 17 00:00:00 2001 From: Ondra Zahradnik Date: Fri, 18 Sep 2026 10:56:13 +0200 Subject: [PATCH 3/3] 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)