From 5e2fa507240504d73fa881515b38e59bd27e4bcd Mon Sep 17 00:00:00 2001 From: Ondra Zahradnik Date: Wed, 9 Sep 2026 14:43:19 +0200 Subject: [PATCH] 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)