litellm/tests/test_litellm_rust/cache/test_v2.py
devin-ai-integration[bot] 5e38a08741
feat(cache): select Rust caching through explicit cache objects (#43601)
* refactor(cache): organize v2 cache as a package

* docs: clarify experimental v2 guidance

* fix(cache): verify cache-hit accounting and preserve logging metadata

* refactor(cache): separate execution facts from host accounting

* refactor(rust): build messages routes with named dependencies

* wip

* fix(cache): preserve facade policy and preflight fallback

* refactor(cache): defer shared Python logging changes

* test(gateway-inference): allow dead code in shared test helpers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(cache): key prepared requests and honor facade controls

* feat(cache): use Python caches from Rust Messages inference

* refactor(cache): separate native and Python cache adapters

* refactor(cache): enforce shared composition and adapter boundaries

* fix(cache): let Python key delegated Rust Messages entries

---------

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-09-29 00:01:44 +00:00

794 lines
35 KiB
Python

import asyncio
from collections.abc import AsyncIterator, Mapping
from types import MappingProxyType
from typing import Final, Literal
import pytest
from pydantic import BaseModel, TypeAdapter
import litellm
from litellm import _v2
from litellm._v2.cache import NativeBackend
from litellm.caching.caching import Cache, CacheMode
from litellm.caching.caching_handler import (
_PENDING_CACHE_WRITES, # pyright: ignore[reportPrivateUsage] # await the existing background cache writer before the next request
)
from litellm.proxy._types import Litellm_EntityType, UserAPIKeyAuth
from litellm.proxy.hooks.model_max_budget_limiter import (
_PROXY_VirtualKeyModelMaxBudgetLimiter,
model_budget_spend_cache_key,
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
from litellm.proxy.utils import InternalUsageCache
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
from litellm.rust_bridge import runtime
from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule
from litellm.rust_bridge.chat_completions.entrypoints import NATIVE_ACOMPLETION, LiteLLMChatCompletionsRequest
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.dispatch import call_hook
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, LiteLLMMessagesRequest
from litellm.rust_bridge.responses.entrypoints import NATIVE_ARESPONSES, LiteLLMResponsesRequest
from litellm.types.caching import CachingSupportedCallTypes
from litellm.types.utils import ModelResponse
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_EVENTS, MESSAGES_MODEL, MESSAGES_RESPONSE
from tests.test_litellm_rust.test_inference import RESPONSES_MODEL, RESPONSES_RESPONSE
pytestmark = pytest.mark.requires_rust_extension
def payload(value: object) -> object:
if isinstance(value, ModelResponse):
return value.model_dump_json(exclude=MappingProxyType({"id": True, "created": True}))
if isinstance(value, dict):
fields: Final = TypeAdapter(dict[str, object]).validate_python(value)
return {name: field for name, field in fields.items() if name != "_hidden_params"}
return value.model_dump_json() if isinstance(value, BaseModel) else value
def cache_key(response: object) -> object:
hidden: Final = get_hidden_params_dict(response)
headers: Final = TypeAdapter(dict[str, object]).validate_python(hidden.get("additional_headers", {}))
return headers.get("x-litellm-cache-key")
async def invoke(
route: Literal["chat", "messages", "responses"],
server: RecordingServer,
options: Mapping[str, object],
native: bool = True,
) -> object:
common: Final = {"api_key": "test-key", "api_base": server.base_url, **options}
if route == "responses":
server.default_response = ResponseSpec(body=RESPONSES_RESPONSE)
arguments: Final = {"model": RESPONSES_MODEL, "input": "hello", **common}
if not native:
return await litellm.aresponses(**arguments)
request: Final = LiteLLMResponsesRequest(
RESPONSES_MODEL, "hello", None, "test-key", server.base_url, "openai", None, arguments
)
return await runtime.arun(
RouteContext(Route.RESPONSES),
binding=NATIVE_ARESPONSES,
native=lambda hook: call_hook(hook, request, (), arguments),
python=runtime.NO_PYTHON,
rules=(RouteRule(Route.RESPONSES, Rollout.RUST_REQUIRED),),
)
server.default_response = (
ResponseSpec(body=None, events=MESSAGES_EVENTS)
if options.get("stream")
else ResponseSpec(body=MESSAGES_RESPONSE)
)
parameters: Final = {"model": MESSAGES_MODEL, "messages": list(MESSAGES), "max_tokens": 32, **common}
if route == "chat":
if not native:
return await litellm.acompletion(**parameters)
chat: Final = LiteLLMChatCompletionsRequest(
MESSAGES_MODEL, list(MESSAGES), None, "test-key", server.base_url, None, None, parameters
)
return await runtime.arun(
RouteContext(Route.CHAT_COMPLETIONS),
binding=NATIVE_ACOMPLETION,
native=lambda hook: call_hook(hook, chat, (), parameters),
python=runtime.NO_PYTHON,
rules=(RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),),
)
if not native:
return await litellm.anthropic_messages(**parameters)
messages: Final = LiteLLMMessagesRequest(
MESSAGES_MODEL, list(MESSAGES), 32, None, "test-key", server.base_url, "anthropic", parameters
)
return await runtime.arun(
RouteContext(Route.MESSAGES),
binding=NATIVE_AMESSAGES,
native=lambda hook: call_hook(hook, messages, (), parameters),
python=runtime.NO_PYTHON,
rules=(RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ("chat", "messages", "responses"))
@pytest.mark.parametrize("backend", ("memory", "redis"))
async def test_v2_cache_skips_provider_and_reports_one_success_per_call(
recording_server: RecordingServer,
route: Literal["chat", "messages", "responses"],
backend: Literal["memory", "redis"],
redis_url: str,
) -> None:
recording_server.expected_requests = 2
litellm.cache = _v2.Cache.memory() if backend == "memory" else _v2.Cache.redis(redis_url, namespace="headers")
recorder: Final = RecordingLogger()
first: Final = await invoke(route, recording_server, {"callbacks": [recorder]})
await recorder.wait_for_async("async_log_success_event")
second: Final = await invoke(route, recording_server, {"callbacks": [recorder]})
assert payload(first) == payload(second)
assert cache_key(first) is None
key: Final = cache_key(second)
assert isinstance(key, str)
assert key == get_hidden_params_dict(second)["cache_key"]
assert len(recording_server.requests) == 1
await drain_logging()
successes: Final = await recorder.wait_for_async("async_log_success_event", count=2)
assert len(successes) == 2
cached_log: Final = TypeAdapter(dict[str, object]).validate_python(successes[-1].kwargs)
assert cached_log["cache_hit"] is True
assert cached_log["response_cost"] == 0
await litellm.cache.delete_cache_keys([key])
refreshed: Final = await invoke(route, recording_server, {"callbacks": [recorder]})
assert cache_key(refreshed) is None
assert len(recording_server.requests) == 2
assert len(await recorder.wait_for_async("async_log_success_event", count=3)) == 3
await litellm.cache.disconnect()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("route", "stream", "legacy"),
(
("chat", False, False),
("messages", False, False),
("responses", False, False),
("messages", True, False),
("messages", False, True),
("messages", True, True),
),
)
@pytest.mark.parametrize("native", (False, True), ids=("python", "rust"))
async def test_cache_hit_keeps_model_budget_spend_but_accounts_for_usage(
recording_server: RecordingServer,
route: Literal["chat", "messages", "responses"],
stream: bool,
native: bool,
monkeypatch: pytest.MonkeyPatch,
legacy: bool,
) -> None:
monkeypatch.setenv("LITELLM_RUST", "1" if native else "0")
litellm.cache = Cache() if legacy else _v2.Cache.memory()
counters: Final = litellm.DualCache()
budget: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(counters)
limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(
InternalUsageCache(counters), model_group_resolver=lambda model: model
)
recorder: Final = RecordingLogger()
key_hash: Final = "a" * 64
metadata: Final = {
"user_api_key": key_hash,
"model_group": "cached-model",
"user_api_key_model_max_budget": {"cached-model": {"max_budget": 1, "budget_duration": "1h"}},
}
options: Final = {
"callbacks": [budget, limiter, recorder],
"metadata": metadata,
"stream": stream,
}
spend_key: Final = model_budget_spend_cache_key(Litellm_EntityType.KEY, key_hash, "cached-model", "1h")
token_key: Final = limiter.create_rate_limit_keys("api_key", key_hash, "tokens")
first: Final = await invoke(route, recording_server, options, native=native)
if stream:
await collect(first)
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
await drain_logging()
first_events: Final = await recorder.wait_for_async("async_log_success_event")
first_log: Final = TypeAdapter(dict[str, object]).validate_python(first_events[0].kwargs)
first_payload: Final = TypeAdapter(dict[str, object]).validate_python(first_log["standard_logging_object"])
expected_cost: Final = TypeAdapter(float).validate_python(first_log["response_cost"])
usage: Final = RESPONSES_RESPONSE["usage"] if route == "responses" else MESSAGES_RESPONSE["usage"]
expected_tokens: Final = usage["input_tokens"] + usage["output_tokens"]
assert expected_cost > 0
assert counters.get_cache(spend_key) == pytest.approx(expected_cost)
assert counters.get_cache(token_key) == first_payload["total_tokens"] == expected_tokens
second: Final = await invoke(route, recording_server, options, native=native)
if stream:
await collect(second)
await drain_logging()
successes: Final = await recorder.wait_for_async("async_log_success_event", count=2)
cached_log: Final = TypeAdapter(dict[str, object]).validate_python(successes[-1].kwargs)
cached_payload: Final = TypeAdapter(dict[str, object]).validate_python(cached_log["standard_logging_object"])
assert len(recording_server.requests) == 1
assert len(successes) == 2
assert cached_log["cache_hit"] is True
assert cached_log["response_cost"] == cached_payload["response_cost"] == 0
assert cached_payload["cache_hit"] is True
assert cached_payload["id"] != first_payload["id"]
assert cached_payload["custom_llm_provider"] == first_payload["custom_llm_provider"]
assert cached_payload["custom_llm_provider"] == ("openai" if route == "responses" else "anthropic"), {
"miss_provider": first_log.get("custom_llm_provider"),
"hit_provider": cached_log.get("custom_llm_provider"),
}
assert cached_payload["total_tokens"] == first_payload["total_tokens"]
assert counters.get_cache(spend_key) == pytest.approx(expected_cost)
assert counters.get_cache(token_key) == 2 * expected_tokens
@pytest.mark.asyncio
@pytest.mark.parametrize("native", (False, True), ids=("python", "rust"))
@pytest.mark.parametrize("backend", ("disabled", "memory", "redis"))
async def test_response_cache_backend_does_not_control_coordination(
recording_server: RecordingServer,
native: bool,
monkeypatch: pytest.MonkeyPatch,
backend: Literal["disabled", "memory", "redis"],
redis_url: str,
) -> None:
monkeypatch.setenv("LITELLM_RUST", "1" if native else "0")
litellm.cache = (
None
if backend == "disabled"
else _v2.Cache.memory()
if backend == "memory"
else _v2.Cache.redis(redis_url, namespace="independent-coordination")
)
recording_server.expected_requests = 2 if backend == "disabled" else 1
counters: Final = litellm.DualCache()
budget: Final = _PROXY_VirtualKeyModelMaxBudgetLimiter(counters)
limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(
InternalUsageCache(counters), model_group_resolver=lambda model: model
)
key_hash: Final = "b" * 64
identity: Final = UserAPIKeyAuth(api_key=key_hash, rpm_limit=2, tpm_limit=1000, max_parallel_requests=1)
spend_key: Final = model_budget_spend_cache_key(Litellm_EntityType.KEY, key_hash, "cached-model", "1h")
request_key: Final = limiter.create_rate_limit_keys("api_key", key_hash, "requests")
token_key: Final = limiter.create_rate_limit_keys("api_key", key_hash, "tokens")
parallel_key: Final = limiter.create_rate_limit_keys("api_key", key_hash, "max_parallel_requests")
expected_tokens: Final = MESSAGES_RESPONSE["usage"]["input_tokens"] + MESSAGES_RESPONSE["usage"]["output_tokens"]
recorder: Final = RecordingLogger()
async def request(call_id: str, successes: int) -> object:
data: Final = {
"model": MESSAGES_MODEL,
"messages": list(MESSAGES),
"litellm_call_id": call_id,
"max_tokens": 32,
"metadata": {
"user_api_key": key_hash,
"model_group": "cached-model",
"user_api_key_model_max_budget": {"cached-model": {"max_budget": 1, "budget_duration": "1h"}},
},
}
await limiter.async_pre_call_hook(identity, counters, data, "acompletion")
assert len(TypeAdapter(dict[str, float]).validate_python(counters.get_cache(parallel_key))) == 1
response: Final = await invoke(
"chat", recording_server, {**data, "callbacks": [budget, limiter, recorder]}, native=native
)
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
await recorder.wait_for_async("async_log_success_event", count=successes)
return response
await asyncio.create_task(request("cache-miss", 1))
assert counters.get_cache(parallel_key) == {}
assert counters.get_cache(request_key) == 1
assert counters.get_cache(token_key) == expected_tokens
first_events: Final = await recorder.wait_for_async("async_log_success_event")
first_cost: Final = TypeAdapter(float).validate_python(first_events[0].kwargs["response_cost"])
assert first_cost > 0
assert counters.get_cache(spend_key) == pytest.approx(first_cost)
await asyncio.create_task(request("cache-hit", 2))
expected_spend: Final = first_cost * recording_server.expected_requests
assert counters.get_cache(spend_key) == pytest.approx(expected_spend)
assert len(recording_server.requests) == recording_server.expected_requests
assert counters.get_cache(parallel_key) == {}
assert counters.get_cache(request_key) == 2
assert counters.get_cache(token_key) == 2 * expected_tokens
with pytest.raises(litellm.RateLimitError):
await asyncio.create_task(request("over-rpm-limit", 3))
assert len(recording_server.requests) == recording_server.expected_requests
assert counters.get_cache(parallel_key) == {}
assert counters.get_cache(token_key) == 2 * expected_tokens
assert counters.get_cache(spend_key) == pytest.approx(expected_spend)
if litellm.cache is not None:
await litellm.cache.disconnect()
@pytest.mark.asyncio
async def test_v2_global_cache_leaves_legacy_only_calls_usable() -> None:
litellm.cache = _v2.Cache.memory()
response: Final = await litellm.aembedding(
model="openai/cache-test-embedding",
input=["hello"],
api_key="test-key",
mock_response=[0.25, 0.75],
)
assert response.model_dump(include={"data"}) == {
"data": [{"embedding": [0.25, 0.75], "index": 0, "object": "embedding"}]
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("route", "legacy"), (("chat", False), ("messages", False), ("responses", False), ("messages", True))
)
async def test_cache_controls_and_backend_credential_key_semantics(
recording_server: RecordingServer,
route: Literal["chat", "messages", "responses"],
legacy: bool,
) -> None:
recording_server.expected_requests = 3 if legacy else 4
litellm.cache = Cache() if legacy else _v2.Cache.memory()
await invoke(route, recording_server, {"cache": {"no-store": True}})
await invoke(route, recording_server, {})
await invoke(route, recording_server, {})
assert len(recording_server.requests) == 2
await invoke(route, recording_server, {"cache": {"no-cache": True}})
await invoke(route, recording_server, {"api_key": "another-key"})
assert len(recording_server.requests) == recording_server.expected_requests
async def collect(stream: object) -> bytes:
assert isinstance(stream, AsyncIterator)
return b"".join([chunk_bytes(chunk) async for chunk in stream])
def chunk_bytes(value: object) -> bytes:
assert isinstance(value, bytes)
return value
@pytest.mark.asyncio
@pytest.mark.parametrize("legacy", (False, True))
async def test_v2_messages_replays_a_completed_stream(recording_server: RecordingServer, legacy: bool) -> None:
recording_server.default_response = ResponseSpec(body=None, events=MESSAGES_EVENTS)
litellm.cache = Cache() if legacy else _v2.Cache.memory()
recorder: Final = RecordingLogger()
parameters: Final = {
"model": MESSAGES_MODEL,
"messages": list(MESSAGES),
"max_tokens": 32,
"api_key": "test-key",
"api_base": recording_server.base_url,
"stream": True,
"callbacks": [recorder],
}
first_stream: Final = await litellm.anthropic_messages(**parameters)
assert cache_key(first_stream) is None
first: Final = await collect(first_stream)
await recorder.wait_for_async("async_log_success_event")
second_stream: Final = await litellm.anthropic_messages(**parameters)
assert isinstance(cache_key(second_stream), str)
assert cache_key(second_stream) == get_hidden_params_dict(second_stream)["cache_key"]
second: Final = await collect(second_stream)
assert payload(first) == payload(second)
assert first == b"".join(recording_server.default_response.payloads())
assert len(recording_server.requests) == 1
await drain_logging()
successes: Final = await recorder.wait_for_async("async_log_success_event", count=2)
cached_log: Final = TypeAdapter(dict[str, object]).validate_python(successes[-1].kwargs)
assert cached_log["cache_hit"] is True
assert cached_log["response_cost"] == 0
@pytest.mark.parametrize("route", ("chat", "responses"))
def test_v2_cache_works_through_python_inference(
recording_server: RecordingServer, route: Literal["chat", "messages", "responses"], monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
litellm.cache = _v2.Cache.memory()
common: Final = {"api_key": "test-key", "api_base": recording_server.base_url}
if route == "responses":
recording_server.default_response = ResponseSpec(body=RESPONSES_RESPONSE)
parameters: Final = {"model": RESPONSES_MODEL, "input": "hello", **common}
first: Final = litellm.responses(**parameters)
second: Final = litellm.responses(**parameters)
assert payload(first) == payload(second)
else:
recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE)
arguments: Final = {"model": MESSAGES_MODEL, "messages": list(MESSAGES), "max_tokens": 32, **common}
initial: Final = litellm.completion(**arguments)
cached: Final = litellm.completion(**arguments)
assert isinstance(initial, ModelResponse) and isinstance(cached, ModelResponse)
assert (
initial.choices[0].message.content
== cached.choices[0].message.content
== MESSAGES_RESPONSE["content"][0]["text"]
)
assert len(recording_server.requests) == 1
@pytest.mark.asyncio
async def test_v2_facade_and_backend_share_storage_and_management() -> None:
cache: Final = _v2.Cache.memory()
await cache.async_add_cache({"answer": 7}, cache_key="shared")
assert cache.get_cache(cache_key="shared") == {"answer": 7}
assert await cache.ping() is True
await cache.delete_cache_keys(["shared"])
assert await cache.async_get_cache(cache_key="shared") is None
cache.add_cache({"answer": 8}, cache_key="flush")
backend: Final = cache.cache
assert isinstance(backend, NativeBackend)
backend.flush_cache()
assert cache.get_cache(cache_key="flush") is None
await cache.disconnect()
@pytest.mark.asyncio
@pytest.mark.parametrize("control", ("s-maxage", "s-max-age"))
async def test_v2_native_cache_accepts_existing_freshness_aliases(
recording_server: RecordingServer, control: str
) -> None:
litellm.cache = _v2.Cache.memory()
first: Final = await invoke("responses", recording_server, {})
second: Final = await invoke("responses", recording_server, {"cache": {control: 600}})
assert payload(first) == payload(second)
assert len(recording_server.requests) == 1
@pytest.mark.asyncio
async def test_v2_cache_does_not_force_native_responses_streaming(
recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
litellm.cache = _v2.Cache.memory()
recording_server.default_response = ResponseSpec(
body=None,
events=(
("response.created", {"type": "response.created", "sequence_number": 0, "response": RESPONSES_RESPONSE}),
(
"response.completed",
{"type": "response.completed", "sequence_number": 1, "response": RESPONSES_RESPONSE},
),
),
)
response: Final = await litellm.aresponses(
model=RESPONSES_MODEL,
input="hello",
stream=True,
caching=False,
api_key="test-key",
api_base=recording_server.base_url,
)
assert isinstance(response, AsyncIterator)
chunks: Final = [chunk async for chunk in response]
assert chunks[-1].type == "response.completed"
assert chunks[-1].response.output[0].content[0].text == "native response"
@pytest.mark.asyncio
async def test_v2_cache_works_through_python_messages(
recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
litellm.cache = _v2.Cache.memory()
recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE)
parameters: Final = {
"model": MESSAGES_MODEL,
"messages": list(MESSAGES),
"max_tokens": 32,
"api_key": "test-key",
"api_base": recording_server.base_url,
}
first: Final = await litellm.anthropic_messages(**parameters)
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
second: Final = await litellm.anthropic_messages(**parameters)
assert payload(first) == payload(second)
assert len(recording_server.requests) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("backend", ("memory", "redis"))
async def test_rust_messages_uses_a_legacy_cache_without_python_inference(
recording_server: RecordingServer,
monkeypatch: pytest.MonkeyPatch,
backend: Literal["memory", "redis"],
redis_url: str,
) -> None:
from litellm.caching.caching import Cache
monkeypatch.setenv("LITELLM_RUST", "1")
litellm.cache = Cache() if backend == "memory" else Cache(type="redis", url=redis_url, namespace="rust-host")
logger: Final = RecordingLogger()
litellm.callbacks = [logger]
recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE)
parameters: Final = {
"model": MESSAGES_MODEL,
"messages": list(MESSAGES),
"max_tokens": 32,
"api_key": "test-key",
"api_base": recording_server.base_url,
}
first: Final = await invoke("messages", recording_server, parameters)
second: Final = await invoke("messages", recording_server, parameters)
assert cache_key(second)
assert cache_key(first) is None
assert payload(first) == payload(second)
assert len(recording_server.requests) == 1
await logger.wait_for_async("async_log_success_event", count=2)
assert logger.names.count("async_log_success_event") == 2
assert "log_failure_event" not in logger.names
assert "async_log_failure_event" not in logger.names
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ("chat", "messages", "responses"))
@pytest.mark.parametrize("native", (False, True))
@pytest.mark.parametrize("excluded", (None, [], ["embedding"]))
async def test_v2_cache_honors_supported_call_types_for_reads_and_writes(
recording_server: RecordingServer,
monkeypatch: pytest.MonkeyPatch,
route: Literal["chat", "messages", "responses"],
native: bool,
excluded: list[CachingSupportedCallTypes] | None,
) -> None:
monkeypatch.setenv("LITELLM_RUST", "1" if native else "0")
litellm.cache = _v2.Cache.memory()
call_type: Final[CachingSupportedCallTypes] = (
"acompletion" if route == "chat" else "anthropic_messages" if route == "messages" else "aresponses"
)
recording_server.expected_requests = 4
litellm.cache.supported_call_types = excluded
await invoke(route, recording_server, {}, native=native)
await invoke(route, recording_server, {}, native=native)
assert len(recording_server.requests) == 2
litellm.cache.supported_call_types = [call_type]
await invoke(route, recording_server, {}, native=native)
assert len(recording_server.requests) == 3
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
await invoke(route, recording_server, {}, native=native)
assert len(recording_server.requests) == 3
litellm.cache.supported_call_types = excluded
await invoke(route, recording_server, {}, native=native)
assert len(recording_server.requests) == 4
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_v2_redis_flush_only_removes_its_namespace(
redis_url: str, recording_server: RecordingServer, asynchronous: bool
) -> None:
own: Final = _v2.Cache.redis(redis_url, namespace="flush-own")
other: Final = _v2.Cache.redis(redis_url, namespace="flush-other")
litellm.cache = own
recording_server.expected_requests = 2
await own.async_add_cache({"answer": "own"}, cache_key="shared")
await other.async_add_cache({"answer": "other"}, cache_key="shared")
await invoke("responses", recording_server, {})
hit: Final = await invoke("responses", recording_server, {})
assert isinstance(cache_key(hit), str)
assert await own.async_get_cache(cache_key="shared") == {"answer": "own"}
backend: Final = own.cache
assert isinstance(backend, NativeBackend)
if asynchronous:
await backend.async_flush_cache()
else:
backend.flush_cache()
assert await own.async_get_cache(cache_key="shared") is None
assert await other.async_get_cache(cache_key="shared") == {"answer": "other"}
refreshed: Final = await invoke("responses", recording_server, {})
assert cache_key(refreshed) is None
assert len(recording_server.requests) == 2
await own.disconnect()
await other.disconnect()
@pytest.mark.asyncio
@pytest.mark.parametrize("native", (False, True))
@pytest.mark.parametrize(
("route", "legacy"), (("chat", False), ("messages", False), ("responses", False), ("messages", True))
)
async def test_v2_default_off_requires_opt_in_even_for_existing_entries(
recording_server: RecordingServer,
monkeypatch: pytest.MonkeyPatch,
route: Literal["chat", "messages", "responses"],
native: bool,
legacy: bool,
) -> None:
monkeypatch.setenv("LITELLM_RUST", "1" if native else "0")
litellm.cache = Cache() if legacy else _v2.Cache.memory()
litellm.cache.mode = CacheMode.default_off
recording_server.expected_requests = 4
await invoke(route, recording_server, {}, native=native)
await invoke(route, recording_server, {}, native=native)
assert len(recording_server.requests) == 2
await invoke(route, recording_server, {"cache": {"use-cache": True}}, native=native)
assert len(recording_server.requests) == 3
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
await invoke(route, recording_server, {"cache": {"use-cache": True}}, native=native)
assert len(recording_server.requests) == 3
await invoke(route, recording_server, {}, native=native)
assert len(recording_server.requests) == 4
@pytest.mark.asyncio
@pytest.mark.parametrize(
("route", "legacy"), (("chat", False), ("messages", False), ("responses", False), ("messages", True))
)
async def test_cache_lookup_uses_backend_request_callback_semantics(
recording_server: RecordingServer,
route: Literal["chat", "messages", "responses"],
legacy: bool,
) -> None:
from tests.test_litellm_rust.support.requests import request_body
class Rewrite(RecordingLogger):
temperature = 0.1
def log_pre_api_call(self, model: str, messages: object, kwargs: dict[str, object]) -> None:
request_body(kwargs)["temperature"] = self.temperature
super().log_pre_api_call(model, messages, kwargs)
logger: Final = Rewrite()
litellm.cache = Cache() if legacy else _v2.Cache.memory()
recording_server.expected_requests = 1 if legacy else 2
await invoke(route, recording_server, {"callbacks": [logger]})
first_hit: Final = await invoke(route, recording_server, {"callbacks": [logger]})
logger.temperature = 0.8
await invoke(route, recording_server, {"callbacks": [logger]})
second_hit: Final = await invoke(route, recording_server, {"callbacks": [logger]})
assert logger.names.count("log_pre_api_call") == 4
assert len(recording_server.requests) == recording_server.expected_requests
assert recording_server.requests[0].body["temperature"] == 0.1
if not legacy:
assert recording_server.requests[1].body["temperature"] == 0.8
assert isinstance(cache_key(first_hit), str)
assert isinstance(cache_key(second_hit), str)
if not legacy:
assert cache_key(first_hit) != cache_key(second_hit)
@pytest.mark.asyncio
@pytest.mark.parametrize("cancel_lookup", (False, True))
async def test_python_cache_operations_stay_in_the_rust_callers_task(
recording_server: RecordingServer,
cancel_lookup: bool,
) -> None:
from litellm.caching.base_cache import BaseCache
from litellm.caching.in_memory_cache import InMemoryCache
caller: Final = asyncio.current_task()
entered: Final = asyncio.Event()
release: Final = asyncio.Event()
storage: Final = InMemoryCache()
class CallerCache(BaseCache):
async def async_set_cache_pipeline(
self, cache_list: list[tuple[str, object]], ttl: float | None = None
) -> None:
await storage.async_set_cache_pipeline(cache_list, ttl=ttl)
async def async_get_cache(self, key: str, **kwargs: object) -> object:
if cancel_lookup:
entered.set()
await release.wait()
else:
assert asyncio.current_task() is caller
return storage.get_cache(key, **kwargs)
async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None:
assert asyncio.current_task() is caller
await asyncio.sleep(0)
storage.set_cache(key, value, **kwargs)
litellm.cache = Cache(_backend=CallerCache())
if cancel_lookup:
recording_server.expected_requests = 0
task: Final = asyncio.create_task(invoke("messages", recording_server, {}))
await asyncio.wait_for(entered.wait(), timeout=5)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
release.set()
await asyncio.sleep(0)
assert len(recording_server.requests) == 0
assert storage.cache_dict == {}
return
first: Final = await invoke("messages", recording_server, {})
second: Final = await invoke("messages", recording_server, {})
assert payload(first) == payload(second)
assert cache_key(second)
assert len(recording_server.requests) == 1
def test_sync_rust_messages_calls_python_cache(recording_server: RecordingServer) -> None:
from litellm.rust_bridge.messages.entrypoints import NATIVE_MESSAGES
litellm.cache = Cache()
recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE)
arguments: Final = {
"model": MESSAGES_MODEL,
"messages": list(MESSAGES),
"max_tokens": 32,
"api_key": "test-key",
"api_base": recording_server.base_url,
}
request: Final = LiteLLMMessagesRequest(
MESSAGES_MODEL, list(MESSAGES), 32, None, "test-key", recording_server.base_url, "anthropic", arguments
)
def call() -> object:
return runtime.run(
RouteContext(Route.MESSAGES),
binding=NATIVE_MESSAGES,
native=lambda hook: call_hook(hook, request, (), arguments),
python=runtime.NO_PYTHON,
rules=(RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),),
)
first: Final = call()
second: Final = call()
assert payload(first) == payload(second)
assert cache_key(second)
assert len(recording_server.requests) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("namespace_source", ("cache", "metadata"))
async def test_rust_messages_legacy_cache_honors_request_namespaces(
recording_server: RecordingServer, namespace_source: str
) -> None:
litellm.cache = Cache()
recording_server.expected_requests = 2
first_options: Final = (
{"cache": {"namespace": "first"}} if namespace_source == "cache" else {"metadata": {"redis_namespace": "first"}}
)
second_options: Final = (
{"cache": {"namespace": "second"}}
if namespace_source == "cache"
else {"metadata": {"redis_namespace": "second"}}
)
await invoke("messages", recording_server, first_options)
second: Final = await invoke("messages", recording_server, second_options)
first_hit: Final = await invoke("messages", recording_server, first_options)
second_hit: Final = await invoke("messages", recording_server, second_options)
assert cache_key(second) is None
assert cache_key(first_hit) is not None
assert cache_key(second_hit) is not None
assert len(recording_server.requests) == 2
@pytest.mark.asyncio
async def test_rust_messages_legacy_semantic_cache_preserves_python_scope(
recording_server: RecordingServer,
) -> None:
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.types.caching import LiteLLMCacheType
litellm.cache = Cache(type=LiteLLMCacheType.REDIS_SEMANTIC, _backend=InMemoryCache())
first_options: Final = {"messages": [{"role": "user", "content": "hello"}]}
second_options: Final = {"messages": [{"role": "user", "content": "hi"}]}
first: Final = await invoke("messages", recording_server, first_options)
second: Final = await invoke("messages", recording_server, second_options)
assert cache_key(first) is None
assert cache_key(second) is not None
assert payload(first) == payload(second)
assert len(recording_server.requests) == 1
@pytest.mark.asyncio
@pytest.mark.parametrize("rust_first", (False, True), ids=("python_to_rust", "rust_to_python"))
@pytest.mark.parametrize("stream", (False, True), ids=("response", "stream"))
async def test_legacy_cache_keeps_public_messages_responses_compatible(
recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, rust_first: bool, stream: bool
) -> None:
litellm.cache = Cache()
monkeypatch.setenv("LITELLM_RUST", "0")
options: Final = {"litellm_params": {"preset_cache_key": "shared-messages"}, "stream": stream}
first: Final = await invoke("messages", recording_server, options, native=rust_first)
first_payload: Final = await collect(first) if stream else payload(first)
await asyncio.gather(*tuple(_PENDING_CACHE_WRITES))
second: Final = await invoke("messages", recording_server, options, native=not rust_first)
second_payload: Final = await collect(second) if stream else payload(second)
assert second_payload == first_payload
assert len(recording_server.requests) == 1