diff --git a/tests/test_litellm_rust/cache/test_v2.py b/tests/test_litellm_rust/cache/test_python_cache.py similarity index 66% rename from tests/test_litellm_rust/cache/test_v2.py rename to tests/test_litellm_rust/cache/test_python_cache.py index d086a02c600..4aa59d8b3d1 100644 --- a/tests/test_litellm_rust/cache/test_v2.py +++ b/tests/test_litellm_rust/cache/test_python_cache.py @@ -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( diff --git a/tests/test_litellm_rust/cache/v2/__init__.py b/tests/test_litellm_rust/cache/v2/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm_rust/cache/v2/test_cache.py b/tests/test_litellm_rust/cache/v2/test_cache.py new file mode 100644 index 00000000000..5f930bd4ba7 --- /dev/null +++ b/tests/test_litellm_rust/cache/v2/test_cache.py @@ -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() diff --git a/tests/test_litellm_rust/support/cache.py b/tests/test_litellm_rust/support/cache.py index c5f2844170a..0498eb1d10d 100644 --- a/tests/test_litellm_rust/support/cache.py +++ b/tests/test_litellm_rust/support/cache.py @@ -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