refactor(native): run acceptance checks inside shared dispatch

This commit is contained in:
Yujong Lee 2026-09-05 13:38:56 -07:00 committed by yujonglee
parent cfcaefd09b
commit 6b3a9af0e3
3 changed files with 113 additions and 2 deletions

View file

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

View file

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

View file

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