mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
add PublicDispatch
This commit is contained in:
parent
a84f68b6e3
commit
9617312ab2
11 changed files with 1198 additions and 676 deletions
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
83
litellm/rust_bridge/dispatch.py
Normal file
83
litellm/rust_bridge/dispatch.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
179
tests/test_litellm/rust_bridge/test_dispatch.py
Normal file
179
tests/test_litellm/rust_bridge/test_dispatch.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue