mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
refactor(native): run acceptance checks inside shared dispatch
This commit is contained in:
parent
cfcaefd09b
commit
6b3a9af0e3
3 changed files with 113 additions and 2 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue