mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
test(vertex_ai): type the CMEK test helpers and stop mutating caller collections
The cachedContents fakes took a caller-owned list and appended to it, which mutates a function parameter, and none of the new helpers carried annotations. Replace both transports with a `_FakeCachedContents` class that owns its recording and grows an immutable tuple of create bodies, and annotate every new helper and test. The kwargs bag in the gemini test is gone: its only precise annotation would be a mutable collection or Any, so the two transform calls now pass their arguments explicitly instead. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
5e2fa50724
commit
1ead3ae91f
2 changed files with 145 additions and 121 deletions
|
|
@ -1,5 +1,6 @@
|
|||
import json
|
||||
from typing import List
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final, List
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -15,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"})
|
||||
|
||||
|
||||
class TestContextCachingEndpoints:
|
||||
"""Test class for ContextCachingEndpoints methods"""
|
||||
|
||||
|
|
@ -1561,7 +1590,7 @@ class TestContextCachingEndpoints:
|
|||
|
||||
KMS_KEY_NAME = "projects/test_project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache"
|
||||
|
||||
def _cmek_call_kwargs(self, custom_llm_provider):
|
||||
def _cmek_call_kwargs(self, custom_llm_provider: str) -> Mapping[str, object]:
|
||||
return {
|
||||
"messages": [
|
||||
{
|
||||
|
|
@ -1588,95 +1617,69 @@ class TestContextCachingEndpoints:
|
|||
"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):
|
||||
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 = ("encryptionSpec", "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"][0]["parts"][0]["text"] == "Cached reference material"
|
||||
assert plain_body["model"].endswith("models/gemini-2.5-pro")
|
||||
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):
|
||||
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.
|
||||
"""
|
||||
posted_bodies = []
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=self._cached_contents_transport(posted_bodies))
|
||||
kwargs = self._cmek_call_kwargs(custom_llm_provider)
|
||||
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 == "new_cache_name"
|
||||
self._assert_kms_key_only_adds_encryption_spec(posted_bodies)
|
||||
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):
|
||||
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."""
|
||||
posted_bodies = []
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=self._cached_contents_transport(posted_bodies))
|
||||
kwargs = self._cmek_call_kwargs(custom_llm_provider)
|
||||
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 == "new_cache_name"
|
||||
self._assert_kms_key_only_adds_encryption_spec(posted_bodies)
|
||||
def _existing_cache_transport(self, listed_display_names, hits):
|
||||
"""Fake Google holding caches already created under `listed_display_names`."""
|
||||
assert returned_cache == CREATED_CACHE_NAME
|
||||
self._assert_kms_key_only_adds_encryption_spec(fake.created)
|
||||
|
||||
def handle(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 listed_display_names
|
||||
]
|
||||
},
|
||||
)
|
||||
hits.append(json.loads(request.content))
|
||||
return httpx.Response(200, json={"name": "freshly_created_cache", "model": "gemini-2.5-pro"})
|
||||
|
||||
return httpx.MockTransport(handle)
|
||||
|
||||
def _display_name_for(self, custom_llm_provider, kms_key_name):
|
||||
recorded = []
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=self._cached_contents_transport(recorded))
|
||||
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 recorded[0]["displayName"]
|
||||
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):
|
||||
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
|
||||
|
|
@ -1691,75 +1694,75 @@ class TestContextCachingEndpoints:
|
|||
|
||||
assert len({unencrypted_name, other_key_name, wanted_name}) == 3
|
||||
|
||||
for already_cached in ([], [unencrypted_name], [other_key_name], [unencrypted_name, other_key_name]):
|
||||
created = []
|
||||
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=self._existing_cache_transport(already_cached, created))
|
||||
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 == "freshly_created_cache", f"reused a foreign-policy cache from {already_cached}"
|
||||
assert created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME}
|
||||
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}
|
||||
|
||||
reused = []
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=self._existing_cache_transport([wanted_name], reused))
|
||||
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=client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider)
|
||||
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 reused == []
|
||||
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
|
||||
):
|
||||
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)
|
||||
|
||||
created = []
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=self._existing_cache_transport([unencrypted_name], created))
|
||||
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 == "freshly_created_cache"
|
||||
assert created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME}
|
||||
assert returned_cache == CREATED_CACHE_NAME
|
||||
assert fake.created[0]["encryptionSpec"] == {"kmsKeyName": self.KMS_KEY_NAME}
|
||||
|
||||
reused = []
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=self._existing_cache_transport([wanted_name], reused))
|
||||
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=client, kms_key_name=self.KMS_KEY_NAME, **self._cmek_call_kwargs(custom_llm_provider)
|
||||
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 reused == []
|
||||
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):
|
||||
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)
|
||||
|
||||
created = []
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=self._existing_cache_transport([unencrypted_name], created))
|
||||
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 created == []
|
||||
assert fake.created == ()
|
||||
|
||||
def test_cached_messages_end_on_supported_turn():
|
||||
from litellm.llms.vertex_ai.context_caching.transformation import (
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
from unittest.mock import Mock
|
||||
|
||||
import httpx
|
||||
|
|
@ -14,6 +16,7 @@ from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
|||
)
|
||||
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
|
||||
|
||||
|
||||
|
|
@ -477,23 +480,28 @@ async def test_vertex_ai_async_transform_inlines_only_the_urls_gemini_cannot_fet
|
|||
assert sorted(async_only_image_fetch.fetched) == sorted([plain_http_png, extensionless_https])
|
||||
|
||||
|
||||
KMS_KEY_NAME = "projects/qa-project/locations/us-central1/keyRings/litellm/cryptoKeys/context-cache"
|
||||
CACHE_NAME = "projects/qa-project/locations/us-central1/cachedContents/123"
|
||||
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"
|
||||
|
||||
|
||||
def _cmek_cache_transport(captured):
|
||||
class _FakeCachedContents:
|
||||
"""Fake Vertex: the cache list is empty, so the cachedContents create is exercised for real."""
|
||||
|
||||
def handle(request: httpx.Request) -> httpx.Response:
|
||||
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={})
|
||||
captured["cache_create"] = json.loads(request.content)
|
||||
self.created = (*self.created, json.loads(request.content))
|
||||
return httpx.Response(200, json={"name": CACHE_NAME, "model": "gemini-2.5-flash"})
|
||||
|
||||
return httpx.MockTransport(handle)
|
||||
|
||||
|
||||
def _cacheable_messages():
|
||||
def _cacheable_messages() -> list[AllMessageValues]: # mutable-ok: the transform signature takes a list
|
||||
return [
|
||||
{
|
||||
"role": "user",
|
||||
|
|
@ -509,48 +517,61 @@ def _cacheable_messages():
|
|||
]
|
||||
|
||||
|
||||
def _transform_kwargs():
|
||||
return {
|
||||
"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",
|
||||
}
|
||||
|
||||
|
||||
def _assert_key_encrypts_cache_without_leaking(captured, body):
|
||||
assert captured["cache_create"]["encryptionSpec"] == {"kmsKeyName": KMS_KEY_NAME}
|
||||
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():
|
||||
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."""
|
||||
captured = {}
|
||||
client = HTTPHandler()
|
||||
client.client = httpx.Client(transport=_cmek_cache_transport(captured))
|
||||
fake: Final = _FakeCachedContents()
|
||||
client: Final = HTTPHandler()
|
||||
client.client = httpx.Client(transport=fake.transport)
|
||||
|
||||
body = transformation.sync_transform_request_body(client=client, **_transform_kwargs())
|
||||
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(captured, body)
|
||||
_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():
|
||||
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."""
|
||||
captured = {}
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=_cmek_cache_transport(captured))
|
||||
fake: Final = _FakeCachedContents()
|
||||
client: Final = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=fake.transport)
|
||||
|
||||
body = await transformation.async_transform_request_body(client=client, **_transform_kwargs())
|
||||
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(captured, body)
|
||||
_assert_key_encrypts_cache_without_leaking(fake, body)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue