mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
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:
parent
943b029ef1
commit
5e2fa50724
5 changed files with 197 additions and 34 deletions
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue