mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge 125c2f1b1e into 9fd25b2228
This commit is contained in:
commit
cd1a545d12
6 changed files with 354 additions and 8 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
||||
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,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)
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue