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
This commit is contained in:
Ondra Zahradnik 2026-09-09 14:02:31 +02:00
parent 47b15ffb67
commit 943b029ef1
6 changed files with 162 additions and 3 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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