From 6b3a9af0e31bc757add5b3b956ddc568a80d7d10 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 13:38:56 -0700 Subject: [PATCH] refactor(native): run acceptance checks inside shared dispatch --- litellm/rust_bridge/bindings.py | 11 ++- litellm/rust_bridge/runtime.py | 35 ++++++++++ .../test_litellm/rust_bridge/test_runtime.py | 69 +++++++++++++++++++ 3 files changed, 113 insertions(+), 2 deletions(-) diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index e0ecdba5b58..a170ec4edeb 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections.abc import Callable +from types import ModuleType from typing import Final, Generic, TypeVar, cast # noqa: TID251 # PyO3 module boundary from litellm.rust_bridge.loader import get_native_bridge @@ -26,14 +27,20 @@ UNCHANGED: Final = Unchanged() class NativeBinding(Generic[BindingT]): """Resolve one native attribute with an explicit, resettable test override.""" - def __init__(self, select: Callable[[NativeModule], BindingT]) -> None: + def __init__( + self, + select: Callable[[NativeModule], BindingT], + *, + module_loader: Callable[[], ModuleType | None] | None = None, + ) -> None: self._select: Final = select + self._module_loader: Final = module_loader self._override: BindingT | None | _Unset = _UNSET def load(self) -> BindingT | None: if not isinstance(self._override, _Unset): return self._override - native: Final = get_native_bridge() + native: Final = self._module_loader() if self._module_loader is not None else get_native_bridge() if native is None: return None module: Final = cast(NativeModule, native) # cast-ok: PyO3 exports are validated individually below diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index 355fa1b17eb..4cd74ddf24f 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -103,12 +103,16 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> DispatchResult[ResultT]: binding_or_fallback: Final = self._binding_or_python_fallback( eligible=eligible, ) if isinstance(binding_or_fallback, PythonFallback): return binding_or_fallback + preflight_result: Final = preflight() if preflight is not None else None + if preflight_result is not None: + return preflight_result return self._attempt_call( call=lambda: call(binding_or_fallback, prepare()), adapt=adapt, @@ -123,12 +127,16 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> DispatchResult[ResultT]: binding_or_fallback: Final = self._binding_or_python_fallback( eligible=eligible, ) if isinstance(binding_or_fallback, PythonFallback): return binding_or_fallback + preflight_result: Final = preflight() if preflight is not None else None + if preflight_result is not None: + return preflight_result return await self._attempt_acall( call=lambda: call(binding_or_fallback, prepare()), adapt=adapt, @@ -144,6 +152,7 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: result: Final = self._attempt( prepare=prepare, @@ -151,6 +160,7 @@ class EndpointBinding(Generic[BindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) match result: case Handled(value=value): @@ -167,6 +177,7 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: result: Final = await self._aattempt( prepare=prepare, @@ -174,6 +185,7 @@ class EndpointBinding(Generic[BindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) match result: case Handled(value=value): @@ -181,6 +193,17 @@ class EndpointBinding(Generic[BindingT]): case PythonFallback(): return await fallback() + def assess( + self, + *, + check: Callable[[BindingT], str | None], + ) -> PythonFallback | None: + binding: Final = self._binding_or_python_fallback(eligible=True) + if isinstance(binding, PythonFallback): + return binding + reason: Final = check(binding) + return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, reason) if reason is not None else None + def accepts( self, *, @@ -206,6 +229,7 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: result: Final = self._attempt( prepare=prepare, @@ -213,6 +237,7 @@ class EndpointBinding(Generic[BindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) match result: case Handled(value=value): @@ -228,6 +253,7 @@ class EndpointBinding(Generic[BindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: result: Final = await self._aattempt( prepare=prepare, @@ -235,6 +261,7 @@ class EndpointBinding(Generic[BindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) match result: case Handled(value=value): @@ -385,6 +412,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: return self.sync.invoke( prepare=prepare, @@ -393,6 +421,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) async def ainvoke( @@ -404,6 +433,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: return await self.asynchronous.ainvoke( prepare=prepare, @@ -412,6 +442,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) def require( @@ -422,6 +453,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: return self.sync.require( prepare=prepare, @@ -429,6 +461,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) async def arequire( @@ -439,6 +472,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt: Callable[[NativeT], ResultT], error_context: BridgeErrorContext, eligible: bool = True, + preflight: Callable[[], PythonFallback | None] | None = None, ) -> ResultT: return await self.asynchronous.arequire( prepare=prepare, @@ -446,6 +480,7 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]): adapt=adapt, error_context=error_context, eligible=eligible, + preflight=preflight, ) diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index fff0f18c312..c00d3fc1950 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -429,3 +429,72 @@ async def test_async_propagate_policy_preserves_native_errors(error: Exception) await bridge.ainvoke(prepare=lambda: None, call=fail, fallback=fallback, adapt=str, error_context=context()) assert caught.value is error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("available, accepted", ((False, False), (True, False), (True, True))) +async def test_preflight_runs_after_binding_selection_before_preparation( + asynchronous: bool, available: bool, accepted: bool +) -> None: + events: list[str] = [] + + def load() -> object | None: + events.append("load") + return object() if available else None + + def preflight() -> runtime.PythonFallback | None: + events.append("preflight") + return None if accepted else runtime.PythonFallback(runtime.PythonFallbackReason.NATIVE_DECLINED) + + def prepare() -> int: + events.append("prepare") + return 7 + + def call(binding: object, request: int) -> int: + events.append("native") + return request + + async def acall(binding: object, request: int) -> int: + return call(binding, request) + + def fallback() -> str: + events.append("python") + return "3" + + async def afallback() -> str: + return fallback() + + endpoint: Final = runtime.EndpointBinding(route="ocr", load=load, enabled=enabled) + result: Final = ( + await endpoint.ainvoke( + prepare=prepare, call=acall, fallback=afallback, adapt=str, error_context=context(), preflight=preflight + ) + if asynchronous + else endpoint.invoke( + prepare=prepare, call=call, fallback=fallback, adapt=str, error_context=context(), preflight=preflight + ) + ) + assert result == ("7" if available and accepted else "3") + assert events == ( + ["load", "preflight", "prepare", "native"] + if available and accepted + else ["load", "preflight", "python"] if available else ["load", "python"] + ) + + +def test_preflight_failure_is_not_a_native_decline() -> None: + endpoint: Final = runtime.EndpointBinding(route="ocr", load=object, enabled=enabled) + + def preflight() -> runtime.PythonFallback | None: + raise ValueError("invalid acceptance contract") + + with pytest.raises(ValueError, match="invalid acceptance contract"): + endpoint.invoke( + prepare=lambda: pytest.fail("must not prepare"), + call=lambda binding, request: pytest.fail("must not invoke"), + fallback=lambda: pytest.fail("must not fall back"), + adapt=str, + error_context=context(), + preflight=preflight, + )