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:
devin-ai-integration[bot] 2026-10-05 20:03:00 -07:00 • committed by GitHub
parent 922ea71fe5
commit 12dcce8db2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 310 additions and 278 deletions

View file

@ -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(

View file

View 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()

View file

@ -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