diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index e23374d57a1..834f5319402 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -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)) 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 75d4ffbed86..14b06cc1ad1 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 @@ -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 diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 13e2238fdf6..0f9b3f1e100 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -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, diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 3b95b786631..d620de55c69 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -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 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 f666829d2e8..26034f4a53d 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 @@ -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 ( 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 f135acd094f..62686b69026 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py @@ -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)