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.
This commit is contained in:
Ondra Zahradnik 2026-09-09 14:43:19 +02:00
parent 943b029ef1
commit 5e2fa50724
5 changed files with 197 additions and 34 deletions

View file

@ -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]]:

View file

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

View file

@ -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):

View file

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

View file

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