add PublicDispatch

This commit is contained in:
Yujong Lee 2026-09-16 15:44:46 -07:00
parent a84f68b6e3
commit 9617312ab2
11 changed files with 1198 additions and 676 deletions

View file

@ -9,8 +9,8 @@ from litellm.rust_bridge.chat_completions.entrypoints import (
NATIVE_ACOMPLETION,
NATIVE_COMPLETION,
LiteLLMChatCompletionsRequest,
NativeAcompletion,
)
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.public_call import (
bind,
optional_bool,
@ -19,7 +19,6 @@ from litellm.rust_bridge.public_call import (
optional_str,
signature,
)
from litellm.rust_bridge.runtime import arun, run
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
@ -69,33 +68,42 @@ def _public_request(
)
_DISPATCH: Final = PublicDispatch(
route=Route.CHAT_COMPLETIONS,
request=lambda args, kwargs: _public_request(_COMPLETION, args, kwargs),
context=lambda request: _context(request),
bypass=lambda request: request.kwargs.get("acompletion") is True,
)
_ADISPATCH: Final = PublicDispatch(
route=Route.CHAT_COMPLETIONS,
request=lambda args, kwargs: _public_request(_ACOMPLETION, args, kwargs),
context=lambda request: _context(request),
)
def completion(
*args: object,
**kwargs: object, # kwargs-ok: preserve the public chat completions call shape
) -> ChatResult | Coroutine[object, object, ChatResult]:
python: Final = _python_completion()
request: Final = _public_request(_COMPLETION, args, kwargs)
if request is None or request.kwargs.get("acompletion") is True:
return python(*args, **kwargs)
return run(
_context(request),
return _DISPATCH.run(
args,
kwargs,
python=python,
binding=NATIVE_COMPLETION,
native=lambda hook: hook(request, args, kwargs),
python=lambda: python(*args, **kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
)
async def acompletion(*args: object, **kwargs: object) -> ChatResult: # kwargs-ok: preserve the public call shape
python: Final = _python_acompletion()
request: Final = _public_request(_ACOMPLETION, args, kwargs)
if request is None:
return await python(*args, **kwargs)
async def native(hook: NativeAcompletion) -> ChatResult:
return await hook(request, args, kwargs)
return await arun(
_context(request), binding=NATIVE_ACOMPLETION, native=native, python=lambda: python(*args, **kwargs)
return await _ADISPATCH.arun(
args,
kwargs,
python=python,
binding=NATIVE_ACOMPLETION,
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
)

View file

@ -5,11 +5,11 @@ from typing import Final, TypeAlias, cast # noqa: TID251 # native binding sele
from litellm.llms.anthropic.experimental_pass_through.messages import handler as main
from litellm.rust_bridge.catalog import Context, Delivery, Route
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.messages.entrypoints import (
NATIVE_AMESSAGES,
NATIVE_MESSAGES,
LiteLLMMessagesRequest,
NativeAmessages,
)
from litellm.rust_bridge.public_call import (
bind,
@ -19,7 +19,6 @@ from litellm.rust_bridge.public_call import (
optional_str,
signature,
)
from litellm.rust_bridge.runtime import arun, run
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
__all__ = ("anthropic_messages", "anthropic_messages_handler")
@ -68,33 +67,42 @@ def _public_request(
)
_DISPATCH: Final = PublicDispatch(
route=Route.MESSAGES,
request=lambda args, kwargs: _public_request(_MESSAGES, args, kwargs),
context=lambda request: _context(request),
bypass=lambda request: request.kwargs.get("is_async") is True,
)
_ADISPATCH: Final = PublicDispatch(
route=Route.MESSAGES,
request=lambda args, kwargs: _public_request(_AMESSAGES, args, kwargs),
context=lambda request: _context(request),
)
def anthropic_messages_handler(
*args: object,
**kwargs: object, # kwargs-ok: preserve the public Anthropic Messages call shape
) -> MessagesResult | Coroutine[object, object, MessagesResult]:
python: Final = _python_messages()
request: Final = _public_request(_MESSAGES, args, kwargs)
if request is None or request.kwargs.get("is_async") is True:
return python(*args, **kwargs)
return run(
_context(request),
return _DISPATCH.run(
args,
kwargs,
python=python,
binding=NATIVE_MESSAGES,
native=lambda hook: hook(request, args, kwargs),
python=lambda: python(*args, **kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
)
async def anthropic_messages(*args: object, **kwargs: object) -> MessagesResult: # kwargs-ok: public call shape
python: Final = _python_amessages()
request: Final = _public_request(_AMESSAGES, args, kwargs)
if request is None:
return await python(*args, **kwargs)
async def native(hook: NativeAmessages) -> MessagesResult:
return await hook(request, args, kwargs)
return await arun(
_context(request), binding=NATIVE_AMESSAGES, native=native, python=lambda: python(*args, **kwargs)
return await _ADISPATCH.arun(
args,
kwargs,
python=python,
binding=NATIVE_AMESSAGES,
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
)

View file

@ -7,8 +7,8 @@ from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import main
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
from litellm.rust_bridge.catalog import Context, Route
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest, NativeAocr
from litellm.rust_bridge.runtime import arun, run
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest
__all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr")
@ -35,41 +35,54 @@ def _bind_request(
)
def _public_request(name: str, args: tuple[object, ...], kwargs: dict[str, object]) -> LiteLLMOcrRequest:
def _public_request(name: str, args: tuple[object, ...], kwargs: Mapping[str, object]) -> LiteLLMOcrRequest:
try:
return _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation
except TypeError as error:
raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None
_DISPATCH: Final = PublicDispatch(
route=Route.OCR,
request=lambda args, kwargs: _public_request("ocr", args, kwargs),
context=lambda request: _context(request),
bypass=lambda request: request.kwargs.get("aocr") is True,
)
_ADISPATCH: Final = PublicDispatch(
route=Route.OCR,
request=lambda args, kwargs: _public_request("aocr", args, kwargs),
context=lambda request: _context(request),
)
def ocr(
*args: object,
**kwargs: object, # kwargs-ok: preserve the public OCR call shape
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
request: Final = _public_request("ocr", args, kwargs)
python_ocr: Final = cast( # cast-ok: forward the original call shape through the Python @client decorator
Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]], main.ocr
)
if request.kwargs.get("aocr") is True:
return python_ocr(*args, **kwargs)
return run(
_context(request),
return _DISPATCH.run(
args,
kwargs,
python=python_ocr,
binding=NATIVE_OCR,
native=lambda hook: hook(request, args, kwargs),
python=lambda: python_ocr(*args, **kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
)
async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: preserve the public OCR call shape
request: Final = _public_request("aocr", args, kwargs)
fallback: Final = cast( # cast-ok: forward the original call shape through the Python @client decorator
Callable[..., Awaitable[OCRResponse]], main.aocr
)
async def native(hook: NativeAocr) -> OCRResponse:
return await hook(request, args, kwargs)
return await arun(_context(request), binding=NATIVE_AOCR, native=native, python=lambda: fallback(*args, **kwargs))
return await _ADISPATCH.arun(
args,
kwargs,
python=fallback,
binding=NATIVE_AOCR,
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
)
def _context(request: LiteLLMOcrRequest) -> Context:

View file

@ -6,14 +6,13 @@ from typing import Final, TypeAlias, cast # noqa: TID251 # native binding sele
from litellm.responses import main
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
from litellm.rust_bridge.catalog import Context, Delivery, Route
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.public_call import bind, optional_bool, optional_mapping, optional_str, signature
from litellm.rust_bridge.responses.entrypoints import (
NATIVE_ARESPONSES,
NATIVE_RESPONSES,
LiteLLMResponsesRequest,
NativeAresponses,
)
from litellm.rust_bridge.runtime import arun, run
from litellm.types.llms.openai import ResponsesAPIResponse
__all__ = ("aresponses", "responses")
@ -61,33 +60,42 @@ def _public_request(
)
_DISPATCH: Final = PublicDispatch(
route=Route.RESPONSES,
request=lambda args, kwargs: _public_request(_RESPONSES, args, kwargs),
context=lambda request: _context(request),
bypass=lambda request: request.kwargs.get("aresponses") is True,
)
_ADISPATCH: Final = PublicDispatch(
route=Route.RESPONSES,
request=lambda args, kwargs: _public_request(_ARESPONSES, args, kwargs),
context=lambda request: _context(request),
)
def responses(
*args: object,
**kwargs: object, # kwargs-ok: preserve the public Responses call shape
) -> ResponsesResult | Coroutine[object, object, ResponsesResult]:
python: Final = _python_responses()
request: Final = _public_request(_RESPONSES, args, kwargs)
if request is None or request.kwargs.get("aresponses") is True:
return python(*args, **kwargs)
return run(
_context(request),
return _DISPATCH.run(
args,
kwargs,
python=python,
binding=NATIVE_RESPONSES,
native=lambda hook: hook(request, args, kwargs),
python=lambda: python(*args, **kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
)
async def aresponses(*args: object, **kwargs: object) -> ResponsesResult: # kwargs-ok: preserve the public call shape
python: Final = _python_aresponses()
request: Final = _public_request(_ARESPONSES, args, kwargs)
if request is None:
return await python(*args, **kwargs)
async def native(hook: NativeAresponses) -> ResponsesResult:
return await hook(request, args, kwargs)
return await arun(
_context(request), binding=NATIVE_ARESPONSES, native=native, python=lambda: python(*args, **kwargs)
return await _ADISPATCH.arun(
args,
kwargs,
python=python,
binding=NATIVE_ARESPONSES,
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
)

View file

@ -0,0 +1,83 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from typing import Final, Generic, TypeVar
from litellm.rust_bridge import catalog
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Context, Route, Rules
from litellm.rust_bridge.configuration import Decision
from litellm.rust_bridge.configuration import decision as rollout_decision
from litellm.rust_bridge.runtime import arun, run
RequestT = TypeVar("RequestT")
NativeT = TypeVar("NativeT")
ResultT = TypeVar("ResultT")
@dataclass(frozen=True, slots=True)
class PublicDispatch(Generic[RequestT]):
route: Route
request: Callable[[tuple[object, ...], Mapping[str, object]], RequestT | None]
context: Callable[[RequestT], Context]
bypass: Callable[[RequestT], bool] | None = None
def _requires_projection(self, rules: Rules) -> bool:
for rule in rules:
if rule.route is not self.route:
continue
if rule.providers is not None or rule.models is not None or rule.deliveries is not None:
if rollout_decision(rule.rollout) is not Decision.PYTHON:
return True
continue
return rollout_decision(rule.rollout) is not Decision.PYTHON
return False
def run(
self,
args: tuple[object, ...],
kwargs: Mapping[str, object],
*,
python: Callable[..., ResultT],
binding: NativeBinding[NativeT],
native: Callable[[NativeT, RequestT, tuple[object, ...], Mapping[str, object]], ResultT],
rules: Rules | None = None,
) -> ResultT:
selected_rules: Final = catalog.RULES if rules is None else rules
if not self._requires_projection(selected_rules):
return python(*args, **kwargs)
request: Final = self.request(args, kwargs)
if request is None or (self.bypass is not None and self.bypass(request)):
return python(*args, **kwargs)
return run(
self.context(request),
binding=binding,
native=lambda hook: native(hook, request, args, kwargs),
python=lambda: python(*args, **kwargs),
rules=selected_rules,
)
async def arun(
self,
args: tuple[object, ...],
kwargs: Mapping[str, object],
*,
python: Callable[..., Awaitable[ResultT]],
binding: NativeBinding[NativeT],
native: Callable[[NativeT, RequestT, tuple[object, ...], Mapping[str, object]], Awaitable[ResultT]],
rules: Rules | None = None,
) -> ResultT:
selected_rules: Final = catalog.RULES if rules is None else rules
if not self._requires_projection(selected_rules):
return await python(*args, **kwargs)
request: Final = self.request(args, kwargs)
if request is None or (self.bypass is not None and self.bypass(request)):
return await python(*args, **kwargs)
return await arun(
self.context(request),
binding=binding,
native=lambda hook: native(hook, request, args, kwargs),
python=lambda: python(*args, **kwargs),
rules=selected_rules,
)

View file

@ -46,9 +46,9 @@ def run(
binding: NativeBinding[NativeT],
native: Callable[[NativeT], ResultT],
python: Callable[[], ResultT],
rules: Rules = RULES,
rules: Rules | None = None,
) -> ResultT:
selected: Final = decision(context, rules)
selected: Final = decision(context, RULES if rules is None else rules)
match selected:
case Decision.PYTHON:
return python()
@ -74,9 +74,9 @@ async def arun(
binding: NativeBinding[NativeT],
native: Callable[[NativeT], Awaitable[ResultT]],
python: Callable[[], Awaitable[ResultT]],
rules: Rules = RULES,
rules: Rules | None = None,
) -> ResultT:
selected: Final = decision(context, rules)
selected: Final = decision(context, RULES if rules is None else rules)
match selected:
case Decision.PYTHON:
return await python()

View file

@ -1,116 +1,154 @@
import inspect
from collections.abc import Generator, Mapping
from typing import Final
from unittest.mock import AsyncMock, Mock
from collections.abc import Callable, Mapping
from typing import Final, cast # noqa: TID251 # narrows legacy callable signatures for inspect
import pytest
import litellm
from litellm import main as python_chat
from litellm.rust_bridge import configuration, runtime
from litellm.rust_bridge.catalog import Route, Rule, decision
from litellm.chat_completions.dispatch import (
_ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
_DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
)
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, Rule
from litellm.rust_bridge.chat_completions.entrypoints import (
NATIVE_ACOMPLETION,
NATIVE_COMPLETION,
LiteLLMChatCompletionsRequest,
NativeAcompletion,
NativeCompletion,
)
from litellm.rust_bridge.configuration import Rollout
from litellm.types.utils import ModelResponse
MESSAGES: Final = [{"role": "user", "content": "hi"}]
RUST_RULES: Final = (Rule(Route.CHAT_COMPLETIONS, Rollout.RUST_OPT_OUT),)
PYTHON_RULES: Final = ()
RUST_RULES: Final = (Rule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),)
@pytest.fixture(autouse=True)
def isolated_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
yield
NATIVE_COMPLETION.reset()
NATIVE_ACOMPLETION.reset()
configuration.reset_rust_configuration()
def completion_binding(native: NativeCompletion | None) -> NativeBinding[NativeCompletion]:
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
binding.override(native)
return binding
@pytest.fixture
def rust_route(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(runtime, "decision", lambda context, rules=RUST_RULES: decision(context, RUST_RULES))
def acompletion_binding(native: NativeAcompletion | None) -> NativeBinding[NativeAcompletion]:
binding: Final[NativeBinding[NativeAcompletion]] = NativeBinding("acompletion", validate=lambda _: None)
binding.override(native)
return binding
def test_public_signature_is_the_legacy_signature() -> None:
assert inspect.signature(litellm.completion) == inspect.signature(python_chat.completion)
assert inspect.signature(litellm.acompletion) == inspect.signature(python_chat.acompletion)
public_completion: Final = cast(Callable[..., object], litellm.completion)
legacy_completion: Final = cast(Callable[..., object], python_chat.completion)
public_acompletion: Final = cast(Callable[..., object], litellm.acompletion)
legacy_acompletion: Final = cast(Callable[..., object], python_chat.acompletion)
assert inspect.signature(public_completion) == inspect.signature(legacy_completion)
assert inspect.signature(public_acompletion) == inspect.signature(legacy_acompletion)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_python_only_route_never_loads_native(monkeypatch: pytest.MonkeyPatch, asynchronous: bool) -> None:
def test_python_route_forwards_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
args: Final[tuple[object, ...]] = ("gpt-4o", MESSAGES)
kwargs: Final[Mapping[str, object]] = {"temperature": 0.1, "metadata": metadata}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
response: Final = ModelResponse()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_chat, "acompletion" if asynchronous else "completion", fallback)
monkeypatch.setattr(
NATIVE_ACOMPLETION if asynchronous else NATIVE_COMPLETION,
"load",
Mock(side_effect=AssertionError("native must not be loaded")),
)
litellm.rust(True)
result: Final = (
await litellm.acompletion("gpt-4o", MESSAGES, temperature=0.1)
if asynchronous
else litellm.completion("gpt-4o", MESSAGES, temperature=0.1)
)
assert result is response
fallback.assert_called_once_with("gpt-4o", MESSAGES, temperature=0.1)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_unavailable_native_uses_python(
monkeypatch: pytest.MonkeyPatch, rust_route: None, asynchronous: bool
) -> None:
response: Final = ModelResponse()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_chat, "acompletion" if asynchronous else "completion", fallback)
(NATIVE_ACOMPLETION if asynchronous else NATIVE_COMPLETION).override(None)
result: Final = (
await litellm.acompletion("gpt-4o", MESSAGES, temperature=0.1)
if asynchronous
else litellm.completion("gpt-4o", MESSAGES, temperature=0.1)
)
assert result is response
fallback.assert_called_once_with("gpt-4o", MESSAGES, temperature=0.1)
def test_native_receives_the_bound_request_and_original_call_shape(rust_route: None) -> None:
captured: Final[list[tuple[LiteLLMChatCompletionsRequest, tuple[object, ...], Mapping[str, object]]]] = []
def python(*call_args: object, **call_kwargs: object) -> ModelResponse: # kwargs-ok: records public call shape
captured.append((call_args, call_kwargs))
return response
def native(
request: LiteLLMChatCompletionsRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
pytest.fail("Python-only dispatch must not call native")
assert (
_DISPATCH.run(
args,
kwargs,
python=python,
binding=completion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=PYTHON_RULES,
)
is response
)
call_args, call_kwargs = captured[0]
assert call_args == args
assert call_args[1] is MESSAGES
assert call_kwargs == kwargs
assert call_kwargs["metadata"] is metadata
assert kwargs == {"temperature": 0.1, "metadata": metadata}
@pytest.mark.asyncio
async def test_async_python_route_forwards_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
args: Final[tuple[object, ...]] = ("gpt-4o", MESSAGES)
kwargs: Final[Mapping[str, object]] = {"temperature": 0.1, "metadata": metadata}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
response: Final = ModelResponse()
async def python(*call_args: object, **call_kwargs: object) -> ModelResponse: # kwargs-ok: records call shape
captured.append((call_args, call_kwargs))
return response
async def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
pytest.fail("Python-only dispatch must not call native")
result: Final = await _ADISPATCH.arun(
args,
kwargs,
python=python,
binding=acompletion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=PYTHON_RULES,
)
assert result is response
call_args, call_kwargs = captured[0]
assert call_args == args
assert call_args[1] is MESSAGES
assert call_kwargs == kwargs
assert call_kwargs["metadata"] is metadata
assert kwargs == {"temperature": 0.1, "metadata": metadata}
def test_native_receives_bound_request_and_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
kwargs: Final[Mapping[str, object]] = {
"stream": True,
"api_key": "sk-test",
"base_url": "https://example.invalid",
"extra_headers": {"x-test": "1"},
"custom_llm_provider": "anthropic",
"metadata": metadata,
}
captured: Final[
list[tuple[LiteLLMChatCompletionsRequest, tuple[object, ...], Mapping[str, object]]]
] = []
def python(*call_args: object, **call_kwargs: object) -> ModelResponse: # kwargs-ok: rejected Rust fallback
pytest.fail("Required Rust dispatch must not call Python")
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
captured.append((request, args, kwargs))
return ModelResponse(model=request.model)
return ModelResponse()
NATIVE_COMPLETION.override(native)
response: Final = litellm.completion(
"anthropic/claude-sonnet-4-5",
MESSAGES,
stream=True,
api_key="sk-test",
base_url="https://example.invalid",
extra_headers={"x-test": "1"},
custom_llm_provider="anthropic",
metadata={"user_id": "u"},
args: Final[tuple[object, ...]] = ("anthropic/claude-sonnet-4-5", MESSAGES)
_DISPATCH.run(
args,
kwargs,
python=python,
binding=completion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
request, call_args, hook_kwargs = captured[0]
assert isinstance(response, ModelResponse)
assert response.model == "anthropic/claude-sonnet-4-5"
request, call_args, call_kwargs = captured[0]
assert request.model == "anthropic/claude-sonnet-4-5"
assert request.messages is MESSAGES
assert request.stream is True
@ -118,69 +156,66 @@ def test_native_receives_the_bound_request_and_original_call_shape(rust_route: N
assert request.api_base == "https://example.invalid"
assert request.custom_llm_provider == "anthropic"
assert request.extra_headers == {"x-test": "1"}
assert request.kwargs == {"custom_llm_provider": "anthropic", "metadata": {"user_id": "u"}}
assert call_args == ("anthropic/claude-sonnet-4-5", MESSAGES)
assert hook_kwargs["metadata"] == {"user_id": "u"}
assert "temperature" not in hook_kwargs
assert request.kwargs == {"custom_llm_provider": "anthropic", "metadata": metadata}
assert call_args == args
assert call_kwargs == kwargs
assert call_kwargs["metadata"] is metadata
def test_internal_async_dispatch_marker_stays_on_python(monkeypatch: pytest.MonkeyPatch, rust_route: None) -> None:
native: Final = Mock(side_effect=AssertionError("acompletion's inner completion() call must stay on Python"))
NATIVE_COMPLETION.override(native)
def test_internal_async_marker_bypasses_native() -> None:
response: Final = ModelResponse()
fallback: Final = Mock(return_value=response)
monkeypatch.setattr(python_chat, "completion", fallback)
called: Final[list[bool]] = []
assert litellm.completion("gpt-4o", MESSAGES, acompletion=True) is response
native.assert_not_called()
def python(*call_args: object, **call_kwargs: object) -> ModelResponse: # kwargs-ok: records public call shape
called.append(True)
return response
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
pytest.fail("acompletion's inner completion call must stay on Python")
result: Final = _DISPATCH.run(
("gpt-4o", MESSAGES),
{"acompletion": True},
python=python,
binding=completion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
assert result is response
assert called == [True]
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
def test_public_binding_errors_do_not_depend_on_native_selection(rust_route: None, enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
litellm.rust(enabled)
NATIVE_COMPLETION.override(native)
with pytest.raises(TypeError, match=r"completion\(\) got multiple values for argument 'model'"):
litellm.completion("gpt-4o", MESSAGES, model="duplicate")
with pytest.raises(TypeError, match=r"completion\(\) missing 1 required positional argument: 'model'"):
litellm.completion()
native.assert_not_called()
class Declined(Exception):
pass
class Upstream(Exception):
pass
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("declined", [False, True])
async def test_only_native_declines_replay_on_python(
monkeypatch: pytest.MonkeyPatch, rust_route: None, asynchronous: bool, declined: bool
) -> None:
failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called")
native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure)
(NATIVE_ACOMPLETION if asynchronous else NATIVE_COMPLETION).override(native)
monkeypatch.setattr(runtime, "native_exception_types", lambda: (Declined, Upstream))
@pytest.mark.parametrize(
("args", "kwargs"),
(
(("gpt-4o", MESSAGES), {"model": "duplicate"}),
((), {}),
),
)
def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Mapping[str, object]) -> None:
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
response: Final = ModelResponse()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_chat, "acompletion" if asynchronous else "completion", fallback)
async def call() -> object:
if asynchronous:
return await litellm.acompletion("gpt-4o", MESSAGES)
return litellm.completion("gpt-4o", MESSAGES)
def python(*call_args: object, **call_kwargs: object) -> ModelResponse: # kwargs-ok: records invalid call shape
captured.append((call_args, call_kwargs))
return response
if declined:
assert await call() is response
fallback.assert_called_once_with("gpt-4o", MESSAGES)
else:
with pytest.raises(RuntimeError) as caught:
await call()
assert caught.value is failure
fallback.assert_not_called()
assert native.call_count == 1
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
pytest.fail("Binding failures must be delegated to Python")
assert (
_DISPATCH.run(
args,
kwargs,
python=python,
binding=completion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
is response
)
assert captured == [(args, kwargs)]

View file

@ -1,101 +1,145 @@
import inspect
from collections.abc import Generator, Mapping
from typing import Final
from unittest.mock import AsyncMock, Mock
from collections.abc import Callable, Mapping
from typing import Final, cast # noqa: TID251 # narrows legacy callable signatures for inspect
import pytest
import litellm
from litellm.llms.anthropic.experimental_pass_through.messages import handler as python_messages
from litellm.rust_bridge import configuration, runtime
from litellm.rust_bridge.catalog import Route, Rule, decision
from litellm.messages.dispatch import (
_ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
_DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
)
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, Rule, Rules
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.messages.entrypoints import (
NATIVE_AMESSAGES,
NATIVE_MESSAGES,
LiteLLMMessagesRequest,
NativeAmessages,
NativeMessages,
)
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
MESSAGES: Final = [{"role": "user", "content": "hi"}]
RUST_RULES: Final = (Rule(Route.MESSAGES, Rollout.RUST_OPT_OUT),)
PYTHON_RULES: Final[Rules] = ()
RUST_RULES: Final[Rules] = (Rule(Route.MESSAGES, Rollout.RUST_REQUIRED),)
def _response(model: str = "claude-sonnet-4-5") -> AnthropicMessagesResponse:
def messages_binding(native: NativeMessages | None) -> NativeBinding[NativeMessages]:
binding: Final[NativeBinding[NativeMessages]] = NativeBinding(
"anthropic_messages_handler", validate=lambda _: None
)
binding.override(native)
return binding
def amessages_binding(native: NativeAmessages | None) -> NativeBinding[NativeAmessages]:
binding: Final[NativeBinding[NativeAmessages]] = NativeBinding("anthropic_messages", validate=lambda _: None)
binding.override(native)
return binding
def response(model: str = "claude-sonnet-4-5") -> AnthropicMessagesResponse:
return AnthropicMessagesResponse(id="msg_test", type="message", role="assistant", model=model, content=[])
@pytest.fixture(autouse=True)
def isolated_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
yield
NATIVE_MESSAGES.reset()
NATIVE_AMESSAGES.reset()
configuration.reset_rust_configuration()
@pytest.fixture
def rust_route(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(runtime, "decision", lambda context, rules=RUST_RULES: decision(context, RUST_RULES))
def test_public_signature_is_the_legacy_signature() -> None:
assert inspect.signature(litellm.anthropic_messages_handler) == inspect.signature(
python_messages.anthropic_messages_handler
public_messages: Final = cast(Callable[..., object], litellm.anthropic_messages_handler)
legacy_messages: Final = cast(Callable[..., object], python_messages.anthropic_messages_handler)
public_amessages: Final = cast(Callable[..., object], litellm.anthropic_messages)
legacy_amessages: Final = cast(Callable[..., object], python_messages.anthropic_messages)
assert inspect.signature(public_messages) == inspect.signature(legacy_messages)
assert inspect.signature(public_amessages) == inspect.signature(legacy_amessages)
def test_python_route_forwards_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
args: Final[tuple[object, ...]] = (16, MESSAGES, "claude-sonnet-4-5")
kwargs: Final[Mapping[str, object]] = {"temperature": 0.1, "litellm_metadata": metadata}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
def python(*call_args: object, **call_kwargs: object) -> AnthropicMessagesResponse: # kwargs-ok: records call shape
captured.append((call_args, call_kwargs))
return expected
def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> AnthropicMessagesResponse:
pytest.fail("Python-only dispatch must not call native")
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=messages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=PYTHON_RULES,
)
assert inspect.signature(litellm.anthropic_messages) == inspect.signature(python_messages.anthropic_messages)
assert result is expected
call_args, call_kwargs = captured[0]
assert call_args == args
assert call_args[1] is MESSAGES
assert call_kwargs == kwargs
assert call_kwargs["litellm_metadata"] is metadata
assert kwargs == {"temperature": 0.1, "litellm_metadata": metadata}
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_python_only_route_never_loads_native(monkeypatch: pytest.MonkeyPatch, asynchronous: bool) -> None:
response: Final = _response()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(
python_messages, "anthropic_messages" if asynchronous else "anthropic_messages_handler", fallback
async def test_async_python_route_forwards_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
args: Final[tuple[object, ...]] = (16, MESSAGES, "claude-sonnet-4-5")
kwargs: Final[Mapping[str, object]] = {"temperature": 0.1, "litellm_metadata": metadata}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
async def python(
*call_args: object, **call_kwargs: object # kwargs-ok: records call shape
) -> AnthropicMessagesResponse:
captured.append((call_args, call_kwargs))
return expected
async def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> AnthropicMessagesResponse:
pytest.fail("Python-only dispatch must not call native")
result: Final = await _ADISPATCH.arun(
args,
kwargs,
python=python,
binding=amessages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=PYTHON_RULES,
)
monkeypatch.setattr(
NATIVE_AMESSAGES if asynchronous else NATIVE_MESSAGES,
"load",
Mock(side_effect=AssertionError("native must not be loaded")),
)
litellm.rust(True)
result: Final = (
await litellm.anthropic_messages(16, MESSAGES, "claude-sonnet-4-5", temperature=0.1)
if asynchronous
else litellm.anthropic_messages_handler(16, MESSAGES, "claude-sonnet-4-5", temperature=0.1)
)
assert result is response
fallback.assert_called_once_with(16, MESSAGES, "claude-sonnet-4-5", temperature=0.1)
assert result is expected
call_args, call_kwargs = captured[0]
assert call_args == args
assert call_args[1] is MESSAGES
assert call_kwargs == kwargs
assert call_kwargs["litellm_metadata"] is metadata
assert kwargs == {"temperature": 0.1, "litellm_metadata": metadata}
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_unavailable_native_uses_python(
monkeypatch: pytest.MonkeyPatch, rust_route: None, asynchronous: bool
) -> None:
response: Final = _response()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(
python_messages, "anthropic_messages" if asynchronous else "anthropic_messages_handler", fallback
)
(NATIVE_AMESSAGES if asynchronous else NATIVE_MESSAGES).override(None)
result: Final = (
await litellm.anthropic_messages(16, MESSAGES, "claude-sonnet-4-5", temperature=0.1)
if asynchronous
else litellm.anthropic_messages_handler(16, MESSAGES, "claude-sonnet-4-5", temperature=0.1)
)
assert result is response
fallback.assert_called_once_with(16, MESSAGES, "claude-sonnet-4-5", temperature=0.1)
def test_native_receives_the_bound_request_and_original_call_shape(rust_route: None) -> None:
def test_native_receives_normalized_request_and_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
args: Final[tuple[object, ...]] = (16, MESSAGES, "anthropic/claude-sonnet-4-5")
kwargs: Final[Mapping[str, object]] = {
"stream": True,
"api_key": "sk-test",
"api_base": "https://example.invalid",
"custom_llm_provider": "anthropic",
"litellm_metadata": metadata,
}
captured: Final[list[tuple[LiteLLMMessagesRequest, tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response("anthropic/claude-sonnet-4-5")
def python(*call_args: object, **call_kwargs: object) -> AnthropicMessagesResponse: # kwargs-ok: rejected fallback
pytest.fail("Required Rust dispatch must not call Python")
def native(
request: LiteLLMMessagesRequest,
@ -103,24 +147,18 @@ def test_native_receives_the_bound_request_and_original_call_shape(rust_route: N
kwargs: Mapping[str, object],
) -> AnthropicMessagesResponse:
captured.append((request, args, kwargs))
return _response(request.model)
return expected
NATIVE_MESSAGES.override(native)
response: Final = litellm.anthropic_messages_handler(
16,
MESSAGES,
"anthropic/claude-sonnet-4-5",
stream=True,
api_key="sk-test",
api_base="https://example.invalid",
custom_llm_provider="anthropic",
litellm_metadata={"user_id": "u"},
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=messages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
request, call_args, hook_kwargs = captured[0]
assert isinstance(response, dict)
assert response["model"] == "anthropic/claude-sonnet-4-5"
assert result is expected
request, call_args, call_kwargs = captured[0]
assert request.model == "anthropic/claude-sonnet-4-5"
assert request.messages is MESSAGES
assert request.max_tokens == 16
@ -128,71 +166,72 @@ def test_native_receives_the_bound_request_and_original_call_shape(rust_route: N
assert request.api_key == "sk-test"
assert request.api_base == "https://example.invalid"
assert request.custom_llm_provider == "anthropic"
assert request.kwargs == {"litellm_metadata": {"user_id": "u"}}
assert call_args == (16, MESSAGES, "anthropic/claude-sonnet-4-5")
assert hook_kwargs["litellm_metadata"] == {"user_id": "u"}
assert "temperature" not in hook_kwargs
assert request.kwargs == {"litellm_metadata": metadata}
assert request.kwargs["litellm_metadata"] is metadata
assert call_args == args
assert call_args[1] is MESSAGES
assert call_kwargs == kwargs
assert call_kwargs["litellm_metadata"] is metadata
def test_internal_async_dispatch_marker_stays_on_python(monkeypatch: pytest.MonkeyPatch, rust_route: None) -> None:
native: Final = Mock(side_effect=AssertionError("the async handler's inner sync call must stay on Python"))
NATIVE_MESSAGES.override(native)
response: Final = _response()
fallback: Final = Mock(return_value=response)
monkeypatch.setattr(python_messages, "anthropic_messages_handler", fallback)
def test_internal_async_marker_bypasses_native() -> None:
args: Final[tuple[object, ...]] = (16, MESSAGES, "claude-sonnet-4-5")
kwargs: Final[Mapping[str, object]] = {"is_async": True}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
assert litellm.anthropic_messages_handler(16, MESSAGES, "claude-sonnet-4-5", is_async=True) is response
native.assert_not_called()
def python(*call_args: object, **call_kwargs: object) -> AnthropicMessagesResponse: # kwargs-ok: records call shape
captured.append((call_args, call_kwargs))
return expected
def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> AnthropicMessagesResponse:
pytest.fail("The async handler's inner sync call must stay on Python")
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
def test_public_binding_errors_do_not_depend_on_native_selection(rust_route: None, enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
litellm.rust(enabled)
NATIVE_MESSAGES.override(native)
with pytest.raises(TypeError, match=r"anthropic_messages_handler\(\) got multiple values for argument 'model'"):
litellm.anthropic_messages_handler(16, MESSAGES, "claude-sonnet-4-5", model="duplicate")
with pytest.raises(TypeError, match=r"anthropic_messages_handler\(\) missing 3 required positional arguments"):
litellm.anthropic_messages_handler()
native.assert_not_called()
class Declined(Exception):
pass
class Upstream(Exception):
pass
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("declined", [False, True])
async def test_only_native_declines_replay_on_python(
monkeypatch: pytest.MonkeyPatch, rust_route: None, asynchronous: bool, declined: bool
) -> None:
failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called")
native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure)
(NATIVE_AMESSAGES if asynchronous else NATIVE_MESSAGES).override(native)
monkeypatch.setattr(runtime, "native_exception_types", lambda: (Declined, Upstream))
response: Final = _response()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(
python_messages, "anthropic_messages" if asynchronous else "anthropic_messages_handler", fallback
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=messages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
assert result is expected
assert captured == [(args, kwargs)]
async def call() -> object:
if asynchronous:
return await litellm.anthropic_messages(16, MESSAGES, "claude-sonnet-4-5")
return litellm.anthropic_messages_handler(16, MESSAGES, "claude-sonnet-4-5")
if declined:
assert await call() is response
fallback.assert_called_once_with(16, MESSAGES, "claude-sonnet-4-5")
else:
with pytest.raises(RuntimeError) as caught:
await call()
assert caught.value is failure
fallback.assert_not_called()
assert native.call_count == 1
@pytest.mark.parametrize(
("args", "kwargs"),
(
((16, MESSAGES, "claude-sonnet-4-5"), {"model": "duplicate"}),
((), {}),
),
)
def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Mapping[str, object]) -> None:
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
def python(*call_args: object, **call_kwargs: object) -> AnthropicMessagesResponse: # kwargs-ok: records invalid call
captured.append((call_args, call_kwargs))
return expected
def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> AnthropicMessagesResponse:
pytest.fail("Binding failures must be delegated to Python")
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=messages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
assert result is expected
assert captured == [(args, kwargs)]

View file

@ -1,66 +1,139 @@
from collections.abc import Generator, Mapping
from collections.abc import Mapping
from typing import Final
from unittest.mock import AsyncMock, Mock
import httpx
import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import main as python_ocr
from litellm.rust_bridge import bindings, configuration, runtime
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest
from litellm.ocr.dispatch import (
_ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
_DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
)
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, Rule, Rules
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest, NativeAocr, NativeOcr
PYTHON_RULES: Final[Rules] = (Rule(Route.OCR, Rollout.PYTHON_ONLY),)
RUST_RULES: Final[Rules] = (Rule(Route.OCR, Rollout.RUST_REQUIRED),)
@pytest.fixture(autouse=True)
def isolated_ocr_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
yield
NATIVE_OCR.reset()
NATIVE_AOCR.reset()
configuration.reset_rust_configuration()
def ocr_binding(native: NativeOcr | None) -> NativeBinding[NativeOcr]:
binding: Final[NativeBinding[NativeOcr]] = NativeBinding("ocr", validate=lambda _: None)
binding.override(native)
return binding
def aocr_binding(native: NativeAocr | None) -> NativeBinding[NativeAocr]:
binding: Final[NativeBinding[NativeAocr]] = NativeBinding("aocr", validate=lambda _: None)
binding.override(native)
return binding
def response(model: str = "mistral/mistral-ocr-latest") -> OCRResponse:
return OCRResponse(pages=[], model=model)
def test_python_route_forwards_original_call_shape() -> None:
document: Final[Mapping[str, object]] = {
"type": "document_url",
"document_url": "https://example.invalid/document.pdf",
}
pages: Final = [0]
args: Final[tuple[object, ...]] = ("mistral/mistral-ocr-latest", document)
kwargs: Final[Mapping[str, object]] = {"pages": pages}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
def python(*call_args: object, **call_kwargs: object) -> OCRResponse: # kwargs-ok: records public call shape
captured.append((call_args, call_kwargs))
return expected
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> OCRResponse:
pytest.fail("Python-only dispatch must not call native")
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=PYTHON_RULES,
)
assert result is expected
call_args, call_kwargs = captured[0]
assert call_args == args
assert call_args[1] is document
assert call_kwargs == kwargs
assert call_kwargs["pages"] is pages
assert kwargs == {"pages": pages}
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_unavailable_native_uses_python(monkeypatch: pytest.MonkeyPatch, asynchronous: bool) -> None:
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_ocr, "aocr" if asynchronous else "ocr", fallback)
if asynchronous:
NATIVE_AOCR.override(None)
else:
NATIVE_OCR.override(None)
document: Final = {"type": "document_url", "document_url": "https://example.com"}
async def test_async_python_route_forwards_original_call_shape() -> None:
document: Final[Mapping[str, object]] = {"type": "file", "file": b"pdf"}
pages: Final = [1]
args: Final[tuple[object, ...]] = ("mistral/mistral-ocr-latest", document)
kwargs: Final[Mapping[str, object]] = {"pages": pages}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
result: Final = (
await litellm.aocr("mistral/mistral-ocr-latest", document, pages=[0])
if asynchronous
else litellm.ocr("mistral/mistral-ocr-latest", document, pages=[0])
async def python(
*call_args: object,
**call_kwargs: object, # kwargs-ok: records public call shape
) -> OCRResponse:
captured.append((call_args, call_kwargs))
return expected
async def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> OCRResponse:
pytest.fail("Python-only dispatch must not call native")
result: Final = await _ADISPATCH.arun(
args,
kwargs,
python=python,
binding=aocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=PYTHON_RULES,
)
assert result is response
fallback.assert_called_once_with("mistral/mistral-ocr-latest", document, pages=[0])
assert result is expected
call_args, call_kwargs = captured[0]
assert call_args == args
assert call_args[1] is document
assert call_kwargs == kwargs
assert call_kwargs["pages"] is pages
assert kwargs == {"pages": pages}
def test_admitted_failure_is_returned_without_replay() -> None:
failure: Final = RuntimeError("admitted")
native: Final = Mock(side_effect=failure)
litellm.rust(True)
NATIVE_OCR.override(native)
try:
with pytest.raises(RuntimeError) as caught:
litellm.ocr("mistral/mistral-ocr-latest", {"type": "document_url", "document_url": "https://example.com"})
assert caught.value is failure
finally:
NATIVE_OCR.reset()
litellm.rust(None)
assert native.call_count == 1
def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_kwargs() -> None:
document: Final = {"type": "document_url", "document_url": "https://example.com"}
def test_native_receives_normalized_positional_request_and_original_call_shape() -> None:
document: Final[Mapping[str, object]] = {"type": "file", "file": b"pdf"}
timeout: Final = httpx.Timeout(30)
extra_headers: Final[dict[str, object]] = {"x-test": "1"}
pages: Final = [0, 2]
args: Final[tuple[object, ...]] = ("mistral/mistral-ocr-latest", document)
kwargs: Final[Mapping[str, object]] = {
"api_key": "test-key",
"api_base": "https://example.invalid",
"timeout": timeout,
"custom_llm_provider": "mistral",
"extra_headers": extra_headers,
"pages": pages,
}
captured: Final[list[tuple[LiteLLMOcrRequest, tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
def python(*call_args: object, **call_kwargs: object) -> OCRResponse: # kwargs-ok: rejected fallback
pytest.fail("Required Rust dispatch must not call Python")
def native(
request: LiteLLMOcrRequest,
@ -68,170 +141,186 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_
kwargs: Mapping[str, object],
) -> OCRResponse:
captured.append((request, args, kwargs))
return OCRResponse(pages=[], model=request.model)
return expected
litellm.rust(True)
NATIVE_OCR.override(native)
try:
response: Final = litellm.ocr("mistral/mistral-ocr-latest", document)
finally:
NATIVE_OCR.reset()
litellm.rust(None)
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
request, call_args, hook_kwargs = captured[0]
assert response.model == "mistral/mistral-ocr-latest"
request, call_args, call_kwargs = captured[0]
assert result is expected
assert request.model == "mistral/mistral-ocr-latest"
assert request.document is document
assert call_args == ("mistral/mistral-ocr-latest", document)
assert hook_kwargs == {}
assert request.api_key == "test-key"
assert request.api_base == "https://example.invalid"
assert request.timeout is timeout
assert request.custom_llm_provider == "mistral"
assert request.extra_headers is extra_headers
assert request.kwargs == {"pages": pages}
assert request.kwargs["pages"] is pages
assert call_args is args
assert call_kwargs is kwargs
def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs() -> None:
document: Final = {"type": "document_url", "document_url": "https://example.com"}
captured: Final[list[Mapping[str, object]]] = []
def test_native_preserves_keyword_model_and_document_in_original_call_shape() -> None:
document: Final[Mapping[str, object]] = {
"type": "document_url",
"document_url": "https://example.invalid/document.pdf",
}
pages: Final = [1]
args: Final[tuple[object, ...]] = ()
kwargs: Final[Mapping[str, object]] = {
"model": "mistral/mistral-ocr-latest",
"document": document,
"pages": pages,
}
captured: Final[list[tuple[LiteLLMOcrRequest, tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
def python(*call_args: object, **call_kwargs: object) -> OCRResponse: # kwargs-ok: rejected fallback
pytest.fail("Required Rust dispatch must not call Python")
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> OCRResponse:
assert args == ()
captured.append(kwargs)
return OCRResponse(pages=[], model=request.model)
captured.append((request, args, kwargs))
return expected
litellm.rust(True)
NATIVE_OCR.override(native)
try:
litellm.ocr(model="mistral/mistral-ocr-latest", document=document)
finally:
NATIVE_OCR.reset()
litellm.rust(None)
assert captured[0]["model"] == "mistral/mistral-ocr-latest"
assert captured[0]["document"] is document
assert "timeout" not in captured[0]
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
def test_public_duplicate_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
document: Final = {"type": "document_url", "document_url": "https://example.com"}
litellm.rust(enabled)
NATIVE_OCR.override(native)
try:
with pytest.raises(TypeError, match=r"ocr\(\) got multiple values for argument 'model'"):
litellm.ocr("mistral/mistral-ocr-latest", document, model="duplicate")
finally:
NATIVE_OCR.reset()
litellm.rust(None)
assert native.call_count == 0
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
def test_public_missing_required_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
litellm.rust(enabled)
NATIVE_OCR.override(native)
try:
with pytest.raises(TypeError, match=r"ocr\(\) missing 1 required positional argument: 'document'"):
litellm.ocr("mistral/mistral-ocr-latest")
finally:
NATIVE_OCR.reset()
litellm.rust(None)
assert native.call_count == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("enabled", [False, True, None])
async def test_environment_opt_out_never_loads_native(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, enabled: bool | None
) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_ocr, "aocr" if asynchronous else "ocr", fallback)
load: Final = Mock(side_effect=AssertionError("native must not be loaded"))
monkeypatch.setattr(bindings, "get_native_bridge", load)
litellm.rust(enabled)
document: Final = {"type": "file", "file": b"pdf"}
result: Final = (
await litellm.aocr("mistral/mistral-ocr-latest", document, pages=[1])
if asynchronous
else litellm.ocr("mistral/mistral-ocr-latest", document, pages=[1])
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
assert result is response
fallback.assert_called_once_with("mistral/mistral-ocr-latest", document, pages=[1])
load.assert_not_called()
request, call_args, call_kwargs = captured[0]
assert result is expected
assert request.model == "mistral/mistral-ocr-latest"
assert request.document is document
assert request.kwargs == {"pages": pages}
assert call_args is args
assert call_kwargs is kwargs
assert call_kwargs["model"] == "mistral/mistral-ocr-latest"
assert call_kwargs["document"] is document
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("environment", [None, "1"])
async def test_native_is_enabled_by_default(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, environment: str | None
) -> None:
if environment is not None:
monkeypatch.setenv("LITELLM_RUST", environment)
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
native: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
if asynchronous:
NATIVE_AOCR.override(native)
else:
NATIVE_OCR.override(native)
fallback: Final = Mock(side_effect=AssertionError("Python must not run"))
monkeypatch.setattr(python_ocr, "aocr" if asynchronous else "ocr", fallback)
def test_aocr_marker_bypasses_native() -> None:
document: Final[Mapping[str, object]] = {"type": "file", "file": b"pdf"}
args: Final[tuple[object, ...]] = ("mistral/mistral-ocr-latest", document)
kwargs: Final[Mapping[str, object]] = {"aocr": True}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
result: Final = (
await litellm.aocr("mistral/mistral-ocr-latest", {})
if asynchronous
else litellm.ocr("mistral/mistral-ocr-latest", {})
def python(*call_args: object, **call_kwargs: object) -> OCRResponse: # kwargs-ok: records public call shape
captured.append((call_args, call_kwargs))
return expected
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> OCRResponse:
pytest.fail("aocr's inner ocr call must stay on Python")
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
assert result is response
assert native.call_count == 1
fallback.assert_not_called()
assert result is expected
assert captured == [(args, kwargs)]
class Declined(Exception):
pass
@pytest.mark.parametrize(
("args", "kwargs", "message"),
(
(
("mistral/mistral-ocr-latest", {"type": "file", "file": b"pdf"}),
{"model": "duplicate"},
r"ocr\(\) got multiple values for argument 'model'",
),
(
("mistral/mistral-ocr-latest",),
{},
r"ocr\(\) missing 1 required positional argument: 'document'",
),
),
)
def test_ocr_parser_errors_before_python_or_native(
args: tuple[object, ...], kwargs: Mapping[str, object], message: str
) -> None:
def python(*call_args: object, **call_kwargs: object) -> OCRResponse: # kwargs-ok: rejects parser failures
pytest.fail("OCR parser failures must not call Python")
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> OCRResponse:
pytest.fail("OCR parser failures must not call native")
class Upstream(Exception):
pass
with pytest.raises(TypeError, match=message):
_DISPATCH.run(
args,
kwargs,
python=python,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("declined", [False, True])
async def test_only_native_declines_replay_on_python(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, declined: bool
@pytest.mark.parametrize(
("args", "kwargs", "message"),
(
(
("mistral/mistral-ocr-latest", {"type": "file", "file": b"pdf"}),
{"model": "duplicate"},
r"aocr\(\) got multiple values for argument 'model'",
),
(
("mistral/mistral-ocr-latest",),
{},
r"aocr\(\) missing 1 required positional argument: 'document'",
),
),
)
async def test_aocr_parser_errors_before_python_or_native(
args: tuple[object, ...], kwargs: Mapping[str, object], message: str
) -> None:
failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called")
native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure)
if asynchronous:
NATIVE_AOCR.override(native)
else:
NATIVE_OCR.override(native)
monkeypatch.setattr(runtime, "native_exception_types", lambda: (Declined, Upstream))
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_ocr, "aocr" if asynchronous else "ocr", fallback)
document: Final = {"type": "file", "file": b"pdf"}
async def python(
*call_args: object,
**call_kwargs: object, # kwargs-ok: rejects parser failures
) -> OCRResponse:
pytest.fail("OCR parser failures must not call Python")
async def call() -> object:
if asynchronous:
return await litellm.aocr("mistral/mistral-ocr-latest", document, pages=[0])
return litellm.ocr("mistral/mistral-ocr-latest", document, pages=[0])
async def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> OCRResponse:
pytest.fail("OCR parser failures must not call native")
if declined:
assert await call() is response
fallback.assert_called_once_with("mistral/mistral-ocr-latest", document, pages=[0])
else:
with pytest.raises(RuntimeError) as caught:
await call()
assert caught.value is failure
fallback.assert_not_called()
assert native.call_count == 1
with pytest.raises(TypeError, match=message):
await _ADISPATCH.arun(
args,
kwargs,
python=python,
binding=aocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)

View file

@ -1,23 +1,28 @@
import inspect
from collections.abc import Generator, Mapping
from typing import Final
from unittest.mock import AsyncMock, Mock
from collections.abc import Callable, Mapping
from typing import Final, cast # noqa: TID251 # narrows legacy callable signatures for inspect
import pytest
import litellm
from litellm.responses import main as python_responses
from litellm.rust_bridge import configuration, runtime
from litellm.rust_bridge.catalog import Route, Rule, decision
from litellm.responses.dispatch import (
_ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
_DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
)
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, Rule
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.responses.entrypoints import (
NATIVE_ARESPONSES,
NATIVE_RESPONSES,
LiteLLMResponsesRequest,
NativeAresponses,
NativeResponses,
)
from litellm.types.llms.openai import ResponsesAPIResponse
RUST_RULES: Final = (Rule(Route.RESPONSES, Rollout.RUST_OPT_OUT),)
INPUT: Final = [{"role": "user", "content": "hi"}]
PYTHON_RULES: Final = ()
RUST_RULES: Final = (Rule(Route.RESPONSES, Rollout.RUST_REQUIRED),)
def _response(model: str = "gpt-4o") -> ResponsesAPIResponse:
@ -26,71 +31,121 @@ def _response(model: str = "gpt-4o") -> ResponsesAPIResponse:
)
@pytest.fixture(autouse=True)
def isolated_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
yield
NATIVE_RESPONSES.reset()
NATIVE_ARESPONSES.reset()
configuration.reset_rust_configuration()
def responses_binding(native: NativeResponses | None) -> NativeBinding[NativeResponses]:
binding: Final[NativeBinding[NativeResponses]] = NativeBinding("responses", validate=lambda _: None)
binding.override(native)
return binding
@pytest.fixture
def rust_route(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(runtime, "decision", lambda context, rules=RUST_RULES: decision(context, RUST_RULES))
def aresponses_binding(native: NativeAresponses | None) -> NativeBinding[NativeAresponses]:
binding: Final[NativeBinding[NativeAresponses]] = NativeBinding("aresponses", validate=lambda _: None)
binding.override(native)
return binding
def test_public_signature_is_the_legacy_signature() -> None:
assert inspect.signature(litellm.responses) == inspect.signature(python_responses.responses)
assert inspect.signature(litellm.aresponses) == inspect.signature(python_responses.aresponses)
public_responses: Final = cast(Callable[..., object], litellm.responses)
legacy_responses: Final = cast(Callable[..., object], python_responses.responses)
public_aresponses: Final = cast(Callable[..., object], litellm.aresponses)
legacy_aresponses: Final = cast(Callable[..., object], python_responses.aresponses)
assert inspect.signature(public_responses) == inspect.signature(legacy_responses)
assert inspect.signature(public_aresponses) == inspect.signature(legacy_aresponses)
def test_python_route_forwards_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
args: Final[tuple[object, ...]] = (INPUT, "gpt-4o")
kwargs: Final[Mapping[str, object]] = {"temperature": 0.1, "litellm_metadata": metadata}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
response: Final = _response()
def python(*call_args: object, **call_kwargs: object) -> ResponsesAPIResponse: # kwargs-ok: records call shape
captured.append((call_args, call_kwargs))
return response
def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> ResponsesAPIResponse:
pytest.fail("Python-only dispatch must not call native")
assert (
_DISPATCH.run(
args,
kwargs,
python=python,
binding=responses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=PYTHON_RULES,
)
is response
)
call_args, call_kwargs = captured[0]
assert call_args == args
assert call_args[0] is INPUT
assert call_kwargs == kwargs
assert call_kwargs["litellm_metadata"] is metadata
assert kwargs == {"temperature": 0.1, "litellm_metadata": metadata}
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_python_only_route_never_loads_native(monkeypatch: pytest.MonkeyPatch, asynchronous: bool) -> None:
async def test_async_python_route_forwards_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
args: Final[tuple[object, ...]] = (INPUT, "gpt-4o")
kwargs: Final[Mapping[str, object]] = {"temperature": 0.1, "litellm_metadata": metadata}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
response: Final = _response()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_responses, "aresponses" if asynchronous else "responses", fallback)
monkeypatch.setattr(
NATIVE_ARESPONSES if asynchronous else NATIVE_RESPONSES,
"load",
Mock(side_effect=AssertionError("native must not be loaded")),
)
litellm.rust(True)
result: Final = (
await litellm.aresponses("hi", "gpt-4o", temperature=0.1)
if asynchronous
else litellm.responses("hi", "gpt-4o", temperature=0.1)
)
async def python(
*call_args: object, **call_kwargs: object # kwargs-ok: records call shape
) -> ResponsesAPIResponse:
captured.append((call_args, call_kwargs))
return response
async def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> ResponsesAPIResponse:
pytest.fail("Python-only dispatch must not call native")
result: Final = await _ADISPATCH.arun(
args,
kwargs,
python=python,
binding=aresponses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=PYTHON_RULES,
)
assert result is response
fallback.assert_called_once_with("hi", "gpt-4o", temperature=0.1)
call_args, call_kwargs = captured[0]
assert call_args == args
assert call_args[0] is INPUT
assert call_kwargs == kwargs
assert call_kwargs["litellm_metadata"] is metadata
assert kwargs == {"temperature": 0.1, "litellm_metadata": metadata}
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_unavailable_native_uses_python(
monkeypatch: pytest.MonkeyPatch, rust_route: None, asynchronous: bool
) -> None:
response: Final = _response()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_responses, "aresponses" if asynchronous else "responses", fallback)
(NATIVE_ARESPONSES if asynchronous else NATIVE_RESPONSES).override(None)
def test_native_receives_normalized_request_and_original_call_shape() -> None:
metadata: Final = {"user_id": "u"}
extra_headers: Final = {"x-test": "1"}
args: Final[tuple[object, ...]] = (INPUT, "anthropic/claude-sonnet-4-5")
kwargs: Final[Mapping[str, object]] = {
"stream": True,
"api_key": "sk-test",
"base_url": "https://example.invalid",
"extra_headers": extra_headers,
"custom_llm_provider": "anthropic",
"litellm_metadata": metadata,
}
captured: Final[
list[tuple[LiteLLMResponsesRequest, tuple[object, ...], Mapping[str, object]]]
] = []
response: Final = _response("anthropic/claude-sonnet-4-5")
result: Final = (
await litellm.aresponses("hi", "gpt-4o", temperature=0.1)
if asynchronous
else litellm.responses("hi", "gpt-4o", temperature=0.1)
)
assert result is response
fallback.assert_called_once_with("hi", "gpt-4o", temperature=0.1)
def test_native_receives_the_bound_request_and_original_call_shape(rust_route: None) -> None:
captured: Final[list[tuple[LiteLLMResponsesRequest, tuple[object, ...], Mapping[str, object]]]] = []
def python(*call_args: object, **call_kwargs: object) -> ResponsesAPIResponse: # kwargs-ok: rejected fallback
pytest.fail("Required Rust dispatch must not call Python")
def native(
request: LiteLLMResponsesRequest,
@ -98,98 +153,103 @@ def test_native_receives_the_bound_request_and_original_call_shape(rust_route: N
kwargs: Mapping[str, object],
) -> ResponsesAPIResponse:
captured.append((request, args, kwargs))
return _response(request.model)
return response
NATIVE_RESPONSES.override(native)
response: Final = litellm.responses(
"hi",
"anthropic/claude-sonnet-4-5",
stream=True,
api_key="sk-test",
api_base="https://example.invalid",
extra_headers={"x-test": "1"},
custom_llm_provider="anthropic",
litellm_metadata={"user_id": "u"},
result: Final = _DISPATCH.run(
args,
kwargs,
python=python,
binding=responses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
request, call_args, hook_kwargs = captured[0]
assert isinstance(response, ResponsesAPIResponse)
assert response.model == "anthropic/claude-sonnet-4-5"
request, call_args, call_kwargs = captured[0]
assert result is response
assert request.model == "anthropic/claude-sonnet-4-5"
assert request.input == "hi"
assert request.input is INPUT
assert request.stream is True
assert request.api_key == "sk-test"
assert request.api_base == "https://example.invalid"
assert request.custom_llm_provider == "anthropic"
assert request.extra_headers == {"x-test": "1"}
assert request.extra_headers is extra_headers
assert request.kwargs == {
"api_key": "sk-test",
"api_base": "https://example.invalid",
"litellm_metadata": {"user_id": "u"},
"base_url": "https://example.invalid",
"litellm_metadata": metadata,
}
assert call_args == ("hi", "anthropic/claude-sonnet-4-5")
assert hook_kwargs["litellm_metadata"] == {"user_id": "u"}
assert "temperature" not in hook_kwargs
assert request.kwargs["litellm_metadata"] is metadata
assert call_args == args
assert call_args[0] is INPUT
assert call_kwargs == kwargs
assert call_kwargs["extra_headers"] is extra_headers
assert call_kwargs["litellm_metadata"] is metadata
def test_internal_async_dispatch_marker_stays_on_python(monkeypatch: pytest.MonkeyPatch, rust_route: None) -> None:
native: Final = Mock(side_effect=AssertionError("aresponses's inner responses() call must stay on Python"))
NATIVE_RESPONSES.override(native)
def test_internal_async_marker_bypasses_native() -> None:
args: Final[tuple[object, ...]] = (INPUT, "gpt-4o")
kwargs: Final[Mapping[str, object]] = {"aresponses": True}
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
response: Final = _response()
fallback: Final = Mock(return_value=response)
monkeypatch.setattr(python_responses, "responses", fallback)
assert litellm.responses("hi", "gpt-4o", aresponses=True) is response
native.assert_not_called()
def python(*call_args: object, **call_kwargs: object) -> ResponsesAPIResponse: # kwargs-ok: records call shape
captured.append((call_args, call_kwargs))
return response
def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> ResponsesAPIResponse:
pytest.fail("aresponses' inner responses call must stay on Python")
assert (
_DISPATCH.run(
args,
kwargs,
python=python,
binding=responses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
is response
)
assert captured == [(args, kwargs)]
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
def test_public_binding_errors_do_not_depend_on_native_selection(rust_route: None, enabled: bool) -> None:
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
litellm.rust(enabled)
NATIVE_RESPONSES.override(native)
with pytest.raises(TypeError, match=r"responses\(\) got multiple values for argument 'model'"):
litellm.responses("hi", "gpt-4o", model="duplicate")
with pytest.raises(TypeError, match=r"responses\(\) missing 2 required positional arguments: 'input' and 'model'"):
litellm.responses()
native.assert_not_called()
class Declined(Exception):
pass
class Upstream(Exception):
pass
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("declined", [False, True])
async def test_only_native_declines_replay_on_python(
monkeypatch: pytest.MonkeyPatch, rust_route: None, asynchronous: bool, declined: bool
@pytest.mark.parametrize(
("args", "kwargs"),
(
((INPUT, "gpt-4o"), {"model": "duplicate"}),
((), {}),
),
)
def test_binding_errors_delegate_unchanged_to_python(
args: tuple[object, ...], kwargs: Mapping[str, object]
) -> None:
failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called")
native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure)
(NATIVE_ARESPONSES if asynchronous else NATIVE_RESPONSES).override(native)
monkeypatch.setattr(runtime, "native_exception_types", lambda: (Declined, Upstream))
captured: Final[list[tuple[tuple[object, ...], Mapping[str, object]]]] = []
response: Final = _response()
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(python_responses, "aresponses" if asynchronous else "responses", fallback)
async def call() -> object:
if asynchronous:
return await litellm.aresponses("hi", "gpt-4o")
return litellm.responses("hi", "gpt-4o")
def python(*call_args: object, **call_kwargs: object) -> ResponsesAPIResponse: # kwargs-ok: records invalid call
captured.append((call_args, call_kwargs))
return response
if declined:
assert await call() is response
fallback.assert_called_once_with("hi", "gpt-4o")
else:
with pytest.raises(RuntimeError) as caught:
await call()
assert caught.value is failure
fallback.assert_not_called()
assert native.call_count == 1
def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> ResponsesAPIResponse:
pytest.fail("Binding failures must be delegated to Python")
assert (
_DISPATCH.run(
args,
kwargs,
python=python,
binding=responses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
rules=RUST_RULES,
)
is response
)
assert captured == [(args, kwargs)]

View file

@ -0,0 +1,179 @@
from collections.abc import AsyncGenerator, Awaitable, Callable, Iterator, Mapping
from dataclasses import dataclass
from typing import Final
import pytest
from litellm.rust_bridge import configuration
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Context, Delivery, Route, Rule, Rules
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.dispatch import PublicDispatch
@dataclass(frozen=True, slots=True)
class Request:
model: str
def binding() -> NativeBinding[object]:
bound: Final[NativeBinding[object]] = NativeBinding("unused", validate=lambda value: value)
bound.override(None)
return bound
def test_route_without_rules_forwards_before_request_projection() -> None:
stream: Final[Iterator[int]] = iter((1, 2))
def reject_request(args: tuple[object, ...], kwargs: Mapping[str, object]) -> Request:
pytest.fail("Python-only routes must not project the request")
dispatch: Final = PublicDispatch(route=Route.CHAT_COMPLETIONS, request=reject_request, context=lambda _: Context(Route.CHAT_COMPLETIONS))
result: Final = dispatch.run(
("model",),
{"stream": True},
python=lambda *args, **kwargs: stream,
binding=binding(),
native=lambda hook, request, args, kwargs: pytest.fail("Python-only routes must not call native"),
rules=(),
)
assert result is stream
def test_unconditional_python_rule_prevents_later_rust_rule_projection() -> None:
rules: Final[Rules] = (
Rule(Route.CHAT_COMPLETIONS, Rollout.PYTHON_ONLY),
Rule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),
)
def reject_request(args: tuple[object, ...], kwargs: Mapping[str, object]) -> Request:
pytest.fail("First-match Python rule must prevent request projection")
dispatch: Final = PublicDispatch(
route=Route.CHAT_COMPLETIONS,
request=reject_request,
context=lambda _: Context(Route.CHAT_COMPLETIONS),
)
expected: Final = object()
result: Final = dispatch.run(
("model",),
{},
python=lambda *args, **kwargs: expected,
binding=binding(),
native=lambda hook, request, args, kwargs: pytest.fail("First-match Python rule must prevent native"),
rules=rules,
)
assert result is expected
def test_disabled_optional_rust_rule_forwards_before_projection() -> None:
rules: Final[Rules] = (Rule(Route.OCR, Rollout.RUST_OPT_OUT),)
def reject_request(args: tuple[object, ...], kwargs: Mapping[str, object]) -> Request:
pytest.fail("Disabled optional Rust must not project the request")
dispatch: Final = PublicDispatch(route=Route.OCR, request=reject_request, context=lambda _: Context(Route.OCR))
expected: Final = object()
configuration.rust(False)
try:
result: Final = dispatch.run(
("model",),
{},
python=lambda *args, **kwargs: expected,
binding=binding(),
native=lambda hook, request, args, kwargs: pytest.fail("Disabled optional Rust must not call native"),
rules=rules,
)
finally:
configuration.rust(None)
assert result is expected
def test_native_stream_result_is_not_consumed_or_wrapped() -> None:
request: Final = Request(model="streaming-model")
stream: Final[Iterator[int]] = iter((1, 2))
rules: Final[Rules] = (
Rule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED, deliveries=frozenset({Delivery.STREAMING})),
)
dispatch: Final = PublicDispatch(
route=Route.CHAT_COMPLETIONS,
request=lambda args, kwargs: request,
context=lambda value: Context(Route.CHAT_COMPLETIONS, model=value.model, delivery=Delivery.STREAMING),
)
def native(request: Request, args: tuple[object, ...], kwargs: Mapping[str, object]) -> Iterator[int]:
return stream
native_binding: Final[
NativeBinding[Callable[[Request, tuple[object, ...], Mapping[str, object]], Iterator[int]]]
] = NativeBinding("stream", validate=lambda _: None)
native_binding.override(native)
result: Final = dispatch.run(
("streaming-model",),
{"stream": True},
python=lambda *args, **kwargs: pytest.fail("Required native stream dispatch must not call Python"),
binding=native_binding,
native=lambda hook, value, args, kwargs: hook(value, args, kwargs),
rules=rules,
)
assert result is stream
@pytest.mark.asyncio
async def test_async_route_without_rules_preserves_async_iterator_result() -> None:
async def chunks() -> AsyncGenerator[int, None]:
yield 1
stream: Final = chunks()
def reject_request(args: tuple[object, ...], kwargs: Mapping[str, object]) -> Request:
pytest.fail("Python-only routes must not project the request")
async def python(*args: object, **kwargs: object) -> AsyncGenerator[int, None]: # kwargs-ok: pass-through shape
return stream
dispatch: Final = PublicDispatch(route=Route.RESPONSES, request=reject_request, context=lambda _: Context(Route.RESPONSES))
result: Final = await dispatch.arun(
("model",),
{"stream": True},
python=python,
binding=binding(),
native=lambda hook, request, args, kwargs: pytest.fail("Python-only routes must not call native"),
rules=(),
)
assert result is stream
await stream.aclose()
@pytest.mark.asyncio
async def test_async_dispatch_accepts_websocket_style_none_result() -> None:
request: Final = Request(model="realtime-model")
rules: Final[Rules] = (
Rule(Route.RESPONSES, Rollout.RUST_REQUIRED, deliveries=frozenset({Delivery.WEBSOCKET})),
)
dispatch: Final = PublicDispatch(
route=Route.RESPONSES,
request=lambda args, kwargs: request,
context=lambda value: Context(Route.RESPONSES, model=value.model, delivery=Delivery.WEBSOCKET),
)
async def python(*args: object, **kwargs: object) -> None: # kwargs-ok: public pass-through shape
pytest.fail("Required native WebSocket dispatch must not call Python")
async def native(request: Request, args: tuple[object, ...], kwargs: Mapping[str, object]) -> None:
return None
native_binding: Final[NativeBinding[Callable[[Request, tuple[object, ...], Mapping[str, object]], Awaitable[None]]]] = NativeBinding(
"websocket", validate=lambda _: None
)
native_binding.override(native)
result: Final = await dispatch.arun(
("realtime-model",),
{},
python=python,
binding=native_binding,
native=lambda hook, value, args, kwargs: hook(value, args, kwargs),
rules=rules,
)
assert result is None