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