mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(tests): extract Rust cache test split (#44780)
Co-authored-by: Yujong Lee <yujong@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
922ea71fe5
commit
12dcce8db2
4 changed files with 310 additions and 278 deletions
|
|
@ -1,14 +1,11 @@
|
|||
import asyncio
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
from pydantic import 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
|
||||
|
|
@ -23,126 +20,19 @@ 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.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
|
||||
from litellm.types.caching import CachingSupportedCallTypes
|
||||
from litellm.types.utils import ModelResponse
|
||||
from tests.test_litellm_rust.support.cache import cache_key, collect, invoke, payload
|
||||
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
|
||||
from tests.test_litellm_rust.test_inference import 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"),
|
||||
|
|
@ -303,20 +193,6 @@ async def test_response_cache_backend_does_not_control_coordination(
|
|||
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))
|
||||
|
|
@ -337,16 +213,6 @@ async def test_cache_controls_and_backend_credential_key_semantics(
|
|||
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:
|
||||
|
|
@ -380,112 +246,6 @@ async def test_v2_messages_replays_a_completed_stream(recording_server: Recordin
|
|||
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(
|
||||
|
|
@ -553,36 +313,6 @@ async def test_v2_cache_honors_supported_call_types_for_reads_and_writes(
|
|||
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(
|
||||
0
tests/test_litellm_rust/cache/v2/__init__.py
vendored
Normal file
0
tests/test_litellm_rust/cache/v2/__init__.py
vendored
Normal file
207
tests/test_litellm_rust/cache/v2/test_cache.py
vendored
Normal file
207
tests/test_litellm_rust/cache/v2/test_cache.py
vendored
Normal file
|
|
@ -0,0 +1,207 @@
|
|||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm import _v2
|
||||
from litellm._v2.cache import NativeBackend
|
||||
from litellm.caching.caching_handler import (
|
||||
_PENDING_CACHE_WRITES, # pyright: ignore[reportPrivateUsage] # await the existing background cache writer before the next request
|
||||
)
|
||||
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
|
||||
from litellm.types.utils import ModelResponse
|
||||
from tests.test_litellm_rust.support.cache import cache_key, invoke, payload
|
||||
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_MODEL, MESSAGES_RESPONSE
|
||||
from tests.test_litellm_rust.test_inference import RESPONSES_MODEL, RESPONSES_RESPONSE
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
|
||||
@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
|
||||
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.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("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()
|
||||
|
|
@ -1,13 +1,28 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Protocol, TypeAlias
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, Protocol, TypeAlias
|
||||
from uuid import uuid4
|
||||
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.rust_bridge import _native
|
||||
from litellm.rust_bridge.response_cache import ResponseCacheRuntime
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
|
||||
from litellm.rust_bridge import _native, 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.response_cache import ResponseCacheRuntime
|
||||
from litellm.rust_bridge.responses.entrypoints import NATIVE_ARESPONSES, LiteLLMResponsesRequest
|
||||
from litellm.types.utils import ModelResponse
|
||||
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
|
||||
|
||||
CacheRuntime: TypeAlias = _native._ResponseCacheRuntime # pyright: ignore[reportPrivateUsage] # private runtime under test
|
||||
|
||||
|
|
@ -52,3 +67,83 @@ def assert_native_runtime(facade: Cache) -> ResponseCacheRuntime:
|
|||
|
||||
def completion_kwargs(label: str) -> dict[str, object]:
|
||||
return {"model": "gpt-4o", "messages": [{"role": "user", "content": f"{label} {uuid4().hex}"}]}
|
||||
|
||||
|
||||
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),),
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue