mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
* 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>
794 lines
35 KiB
Python
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
|