diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index d5478920de0..d630728bf1d 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -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 datetime import datetime, timezone @@ -11,7 +12,11 @@ from types import MappingProxyType 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 @@ -106,6 +111,19 @@ def _normalize_ttl_to_seconds(ttl: object) -> str | None: return f"{seconds:.9f}".rstrip("0").rstrip(".") + "s" +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]]: @@ -165,6 +183,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) @@ -199,4 +218,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..2e85bc7e815 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 @@ -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, ) @@ -290,6 +291,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 @@ -366,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, @@ -393,6 +396,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 +453,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 @@ -520,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, @@ -548,6 +554,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 e3cc3bbb2dc..f1ad4db2a04 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 ce51e46ef15..9114c8e8973 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] + + 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/gemini/test_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py index f135acd094f..8c277d66afa 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_transformation.py @@ -1,19 +1,22 @@ import json import uuid +from collections.abc import Mapping +from typing import Final from unittest.mock import Mock 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, ) from litellm.types.llms import openai from litellm.types import completion +from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.vertex_ai import RequestBody @@ -475,3 +478,100 @@ 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]) + + +KMS_KEY_NAME: Final = "projects/qa-project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache" +CACHE_NAME: Final = "projects/qa-project/locations/us-central1/cachedContents/123" + + +class _FakeCachedContents: + """Fake Vertex: the cache list is empty, so the cachedContents create is exercised for real.""" + + def __init__(self) -> None: + self.created: tuple[Mapping[str, object], ...] = () + + @property + def transport(self) -> httpx.MockTransport: + return httpx.MockTransport(self._handle) + + def _handle(self, request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response(200, json={}) + self.created = (*self.created, json.loads(request.content)) + return httpx.Response(200, json={"name": CACHE_NAME, "model": "gemini-2.5-flash"}) + + +def _cacheable_messages() -> list[AllMessageValues]: # mutable-ok: the transform signature takes a list + return [ + { + "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?"}, + ] + + +def _assert_key_encrypts_cache_without_leaking(fake: _FakeCachedContents, body: Mapping[str, object]) -> None: + (cache_create,) = fake.created + assert 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() -> None: + """`kms_key_name` reaches the cachedContents POST as encryptionSpec and never the generateContent body.""" + fake: Final = _FakeCachedContents() + client: Final = HTTPHandler() + client.client = httpx.Client(transport=fake.transport) + + body: Final = transformation.sync_transform_request_body( + client=client, + 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", + ) + + _assert_key_encrypts_cache_without_leaking(fake, body) + + +@pytest.mark.asyncio +async def test_async_transform_request_body_forwards_kms_key_name_to_cache_creation() -> None: + """The async transform path pops and forwards the key independently of the sync one.""" + fake: Final = _FakeCachedContents() + client: Final = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=fake.transport) + + body: Final = await transformation.async_transform_request_body( + client=client, + 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", + ) + + _assert_key_encrypts_cache_without_leaking(fake, body) diff --git a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 7913700c8a7..5c4cbcdfcaa 100644 --- a/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/unit/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1,4 +1,6 @@ -from typing import List +import json +from collections.abc import Mapping, Sequence +from typing import Final, List from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -14,6 +16,34 @@ from litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching import ( ) +CREATED_CACHE_NAME: Final = "freshly_created_cache" + + +class _FakeCachedContents: + """Fake Google: GET lists the caches in `already_cached`, POST records the create body and succeeds.""" + + def __init__(self, already_cached: Sequence[str] = ()) -> None: + self._already_cached: Final = tuple(already_cached) + self.created: tuple[Mapping[str, object], ...] = () + + @property + def transport(self) -> httpx.MockTransport: + return httpx.MockTransport(self._handle) + + def _handle(self, 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 self._already_cached + ] + }, + ) + self.created = (*self.created, json.loads(request.content)) + return httpx.Response(200, json={"name": CREATED_CACHE_NAME, "model": "gemini-2.5-pro"}) + + @pytest.fixture def local_model_cost_map(monkeypatch): """Force the bundled in-repo cost map so capability and pricing assertions do not @@ -1614,6 +1644,181 @@ 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: str) -> Mapping[str, object]: + 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 _assert_kms_key_only_adds_encryption_spec(self, posted_bodies: Sequence[Mapping[str, object]]) -> None: + plain_body, cmek_body = posted_bodies + assert "encryptionSpec" not in plain_body + assert cmek_body["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME} + assert cmek_body["displayName"] != plain_body["displayName"] + + untouched: Final = ("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"] == [{"role": "user", "parts": [{"text": "Cached reference material"}]}] + assert str(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: str) -> None: + """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. + """ + fake: Final = _FakeCachedContents() + client: Final = HTTPHandler() + client.client = httpx.Client(transport=fake.transport) + kwargs: Final = 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 == CREATED_CACHE_NAME + self._assert_kms_key_only_adds_encryption_spec(fake.created) + + @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: str) -> None: + """Async variant of test_check_and_create_cache_kms_key_adds_encryption_spec.""" + fake: Final = _FakeCachedContents() + client: Final = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=fake.transport) + kwargs: Final = 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 == CREATED_CACHE_NAME + self._assert_kms_key_only_adds_encryption_spec(fake.created) + + def _display_name_for(self, custom_llm_provider: str, kms_key_name: str | None) -> str: + fake: Final = _FakeCachedContents() + client: Final = HTTPHandler() + client.client = httpx.Client(transport=fake.transport) + self.context_caching.check_and_create_cache( + client=client, kms_key_name=kms_key_name, **self._cmek_call_kwargs(custom_llm_provider) + ) + return str(fake.created[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: str + ) -> None: + """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)): + fake = _FakeCachedContents(already_cached) + client = HTTPHandler() + client.client = httpx.Client(transport=fake.transport) + + _, _, 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 == CREATED_CACHE_NAME, f"reused a foreign-policy cache from {already_cached}" + assert fake.created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME} + + reusable: Final = _FakeCachedContents((wanted_name,)) + reuse_client: Final = HTTPHandler() + reuse_client.client = httpx.Client(transport=reusable.transport) + + _, _, returned_cache = self.context_caching.check_and_create_cache( + client=reuse_client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider) + ) + + assert returned_cache == f"cache-for-{wanted_name}" + assert reusable.created == () + + @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: str + ) -> None: + """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) + + fake: Final = _FakeCachedContents((unencrypted_name,)) + client: Final = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=fake.transport) + + _, _, 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 == CREATED_CACHE_NAME + assert fake.created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME} + + reusable: Final = _FakeCachedContents((wanted_name,)) + reuse_client: Final = AsyncHTTPHandler() + reuse_client.client = httpx.AsyncClient(transport=reusable.transport) + + _, _, returned_cache = await self.context_caching.async_check_and_create_cache( + client=reuse_client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider) + ) + + assert returned_cache == f"cache-for-{wanted_name}" + assert reusable.created == () + + @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: str) -> None: + """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) + + fake: Final = _FakeCachedContents((unencrypted_name,)) + client: Final = HTTPHandler() + client.client = httpx.Client(transport=fake.transport) + + _, _, 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 fake.created == () def test_cached_messages_end_on_supported_turn(): from litellm.llms.vertex_ai.context_caching.transformation import (