This commit is contained in:
Ondra Zahradnik 2026-09-23 21:12:54 +02:00 • committed by GitHub
commit cd1a545d12
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 354 additions and 8 deletions

View file

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

View file

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

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]
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,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)

View file

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