refactor(native): replace preflight APIs with typed execution declines

This commit is contained in:
Yujong Lee 2026-09-05 16:51:15 -07:00
parent 8e20a8644f
commit 2e9a24f67d
15 changed files with 31 additions and 683 deletions

View file

@ -141,35 +141,21 @@ fn session_event(session_id: &str, call_id: &str, message: Option<String>) -> Se
}
}
#[pyfunction]
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
fn responses_websocket_decline(
_model: &str,
custom_llm_provider: &str,
context: NativeRequestContext,
) -> Option<String> {
let context: litellm_core::request_context::LiteLlmRequestContext = context.into();
routes::definition::request_decline(
litellm_core::responses::websocket::native_websocket_supported(custom_llm_provider),
&context,
)
}
#[pymodule(gil_used = false)]
mod _native {
use pyo3::prelude::*;
#[pymodule_init]
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::errors::register(module)?;
use pyo3::types::PyDict;
litellm_python_interop::callback_runtime::register(module)?;
super::callback_bindings::register(module)?;
super::errors::register(module)?;
let ready_endpoints = PyDict::new(module.py());
module.add("ready_endpoints", ready_endpoints)?;
super::routes::register(module)?;
module.add_class::<super::ResponsesWebSocketConnection>()?;
module.add_function(wrap_pyfunction!(
super::responses_websocket_decline,
module
)?)?;
super::diagnostics::register(module)
}
}
@ -194,20 +180,16 @@ mod tests {
let expected = [
"RustBridgeDeclined",
"RustUpstreamError",
"ocr_decline",
"ready_endpoints",
"ocr",
"aocr",
"transcription_decline",
"transcription",
"atranscription",
"messages_decline",
"messages",
"amessages",
"chat_completions_decline",
"chat_completions",
"achat_completions",
"ResponsesWebSocketConnection",
"responses_websocket_decline",
"gil_stats",
];

View file

@ -46,25 +46,10 @@ fn prepare_transcription(
})
}
#[pyfunction]
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
fn transcription_decline(
_model: &str,
custom_llm_provider: &str,
context: NativeRequestContext,
) -> Option<String> {
let context: LiteLlmRequestContext = context.into();
super::definition::request_decline(
litellm_core::audio_transcription::transcription_provider_supported(custom_llm_provider),
&context,
)
}
bridge_route! {
sync = transcription,
asynchronous = atranscription,
request = AudioTranscriptionInputs,
prepare = prepare_transcription,
errors = core_error_to_pyerr,
extra = [transcription_decline],
}

View file

@ -28,6 +28,17 @@ fn prepare_chat_completions(
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
let context: LiteLlmRequestContext = context.into();
let messages = required_value("messages", input.messages, Value::is_array, "list")?;
let options: RequestOptions = options.into();
if let Some(reason) = chat_completions_decline_reason(
&input.model,
options.custom_llm_provider.as_deref(),
messages.clone(),
&input.optional_params,
&options,
&context,
) {
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
}
Ok(async move {
run_route(
ChatCompletionsRequest {
@ -35,54 +46,17 @@ fn prepare_chat_completions(
messages,
optional_params: input.optional_params,
},
&options.into(),
&options,
&context,
)
.await
})
}
#[pyfunction]
#[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None, *, options, context))]
#[allow(
clippy::too_many_arguments,
reason = "PyO3 preserves chat preflight inputs alongside separated options and context"
)]
fn chat_completions_decline(
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] messages: Value,
#[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option<Value>,
custom_llm_provider: Option<String>,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> PyResult<Option<String>> {
let context: LiteLlmRequestContext = context.into();
let options: RequestOptions = options.into();
let optional_params = match optional_params {
None | Some(Value::Null) => Map::new(),
Some(Value::Object(params)) => params,
Some(_) => {
return Err(pyo3::exceptions::PyValueError::new_err(
"optional_params must be a dict",
));
}
};
Ok(chat_completions_decline_reason(
&model,
custom_llm_provider.as_deref(),
messages,
&optional_params,
&options,
&context,
)
.map(str::to_string))
}
bridge_route! {
sync = chat_completions,
asynchronous = achat_completions,
request = ChatCompletionsInputs,
prepare = prepare_chat_completions,
errors = chat_completions_error_to_pyerr,
extra = [chat_completions_decline],
}

View file

@ -285,7 +285,7 @@ for field in ('litellm_call_id', 'trace_id', 'request_model'):
}
#[test]
fn acceptance_and_execution_decline_unsupported_requests_before_io() {
fn normal_execution_declines_unsupported_requests_without_acceptance_exports() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "routes").expect("module should be created");
@ -293,31 +293,18 @@ for field in ('litellm_call_id', 'trace_id', 'request_model'):
module
.add_class::<crate::ResponsesWebSocketConnection>()
.unwrap();
module
.add_function(
wrap_pyfunction!(crate::responses_websocket_decline, &module).unwrap(),
)
.unwrap();
let locals = crate::marshal::request_fixtures(py);
locals.set_item("routes", module).unwrap();
py.run(
c"
for route, provider in (
('chat_completions', 'anthropic'),
('messages', 'anthropic'),
('transcription', 'bedrock'),
('ocr', 'mistral'),
('responses_websocket', 'openai'),
):
decline = getattr(routes, route + '_decline')
assert decline('model', provider, context=context) is None, route
for flag in ('stream', 'has_agentic_hook', 'has_custom_client'):
flagged_context = replace(
context,
capabilities=replace(context.capabilities, **{flag: True}),
)
assert decline('model', provider, context=flagged_context) is not None, (route, flag)
reason = decline('model', 'unsupported-native-provider', context=context)
assert reason is not None, route
assert not hasattr(routes, route + '_decline'), route
request = Request(
messages=[], body={}, audio={}, document={}, optional_params={},
url='invalid-url-must-not-be-used',
@ -333,29 +320,13 @@ for route, provider in (
execute(request, options=unsupported_options, context=context)
except Exception as error:
assert type(error).__name__ == 'RustBridgeDeclined', (route, error)
assert str(error) == reason, (route, reason, error)
else:
raise AssertionError('unsupported request reached provider execution')
native_context = replace(
context,
capabilities=replace(context.capabilities, request_format='native'),
)
litellm_context = replace(
context,
capabilities=replace(context.capabilities, request_format='litellm'),
)
assert routes.ocr_decline('model', 'mistral', context=native_context) is not None
assert routes.ocr_decline('model', 'mistral', context=litellm_context) is None
assert routes.ocr_decline(
'doc-intelligence/prebuilt-layout',
'azure_ai',
context=native_context,
) is None
",
Some(&locals),
Some(&locals),
)
.expect("acceptance must match execution eligibility without I/O");
.expect("normal execution must decline unsupported requests without I/O");
});
}

View file

@ -42,25 +42,10 @@ fn prepare_messages(
})
}
#[pyfunction]
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
fn messages_decline(
_model: &str,
custom_llm_provider: &str,
context: NativeRequestContext,
) -> Option<String> {
let context: LiteLlmRequestContext = context.into();
super::definition::request_decline(
litellm_core::messages::messages_provider_supported(custom_llm_provider),
&context,
)
}
bridge_route! {
sync = messages,
asynchronous = amessages,
request = MessagesInputs,
prepare = prepare_messages,
errors = core_error_to_pyerr,
extra = [messages_decline],
}

View file

@ -56,27 +56,10 @@ fn prepare_ocr(
})
}
#[pyfunction]
#[pyo3(signature = (model, custom_llm_provider, *, context))]
fn ocr_decline(
model: &str,
custom_llm_provider: &str,
context: NativeRequestContext,
) -> Option<String> {
let context: LiteLlmRequestContext = context.into();
let provider_supported = litellm_ai_gateway::io::ocr::ocr_provider_supported(
model,
custom_llm_provider,
context.capabilities.request_format.as_deref(),
);
super::definition::request_decline(provider_supported, &context)
}
bridge_route! {
sync = ocr,
asynchronous = aocr,
request = OcrInputs,
prepare = prepare_ocr,
errors = ocr_error_to_pyerr,
extra = [ocr_decline],
}

View file

@ -27,13 +27,11 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
convert_to_model_response_object,
)
from litellm.llms.bedrock.request_metadata import get_bedrock_request_metadata_fields
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.protocols import (
RustAchatCompletions,
RustChatCompletions,
RustChatCompletionsDecline,
)
from litellm.rust_bridge.request import (
NativeAnthropicOptions,
@ -52,9 +50,7 @@ from litellm.rust_bridge.request import (
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointBinding,
EndpointDispatch,
PythonFallback,
async_none,
)
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -117,18 +113,10 @@ _CHAT: Final[EndpointDispatch[RustChatCompletions, RustAchatCompletions]] = Endp
asynchronous=lambda native: native.achat_completions,
enabled=rust_enabled,
)
_CHAT_PREFLIGHT: Final[EndpointBinding[RustChatCompletionsDecline]] = EndpointBinding.native(
route="chat_completions",
select=lambda native: native.chat_completions_decline,
enabled=rust_enabled,
)
def set_rust_chat_completions(
*,
chat_completions: RustChatCompletions | None | Unchanged = UNCHANGED,
achat_completions: RustAchatCompletions | None | Unchanged = UNCHANGED,
decline: RustChatCompletionsDecline | None | Unchanged = UNCHANGED,
) -> None:
"""Inject the native callables, so tests can supply a double instead of
patching module attributes."""
@ -142,11 +130,6 @@ def set_rust_chat_completions(
_CHAT.asynchronous.reset()
else:
_CHAT.asynchronous.override(achat_completions)
if not isinstance(decline, Unchanged):
if decline is None:
_CHAT_PREFLIGHT.reset()
else:
_CHAT_PREFLIGHT.override(decline)
def _provider_eligibility_options(
@ -183,39 +166,6 @@ def _eligibility_context(
)
def _execution_context(context: NativeRequestContext | None, mode: str) -> NativeRequestContext:
current = context or NativeRequestContext()
return with_capabilities(current, replace(current.capabilities, execution_mode=mode))
def rust_chat_completions_accepts(
*,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object],
custom_llm_provider: str | None,
litellm_params: Mapping[str, object] | None,
stream: object,
) -> bool:
"""Whether the Rust path will serve this request.
Asked before the caller commits to either path, so pre-call logging is
emitted exactly once, on whichever path actually runs. The core's own
capability gate answers the second half; it resolves no credentials and
performs no I/O.
"""
return _CHAT_PREFLIGHT.accepts(
check=lambda decline: decline(
model=model,
messages=messages,
optional_params=optional_params,
custom_llm_provider=custom_llm_provider,
options=_provider_eligibility_options(custom_llm_provider, litellm_params, optional_params),
context=_eligibility_context(stream=bool(stream)),
),
)
def _build_model_response(
rust_response: Mapping[str, object],
model_response: ModelResponse,
@ -385,24 +335,6 @@ class _ChatOperation:
python: Callable[[], _CompletionDispatchResult]
pre_call_logged: bool = False
def assess(self) -> PythonFallback | None:
ctx: Final = self.context
return _CHAT_PREFLIGHT.assess(
check=lambda decline: decline(
model=ctx.model,
messages=ctx.messages,
optional_params=ctx.optional_params,
custom_llm_provider=ctx.custom_llm_provider,
options=_provider_eligibility_options(ctx.custom_llm_provider, ctx.litellm_params, ctx.optional_params),
context=_eligibility_context(
execution_mode="async" if ctx.acompletion else "sync",
stream=bool(ctx.stream),
has_custom_client=ctx.client is not None or ctx.shared_session is not None,
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(ctx.logging),
),
),
)
def prepare(self) -> PreparedNativeCall[NativeChatCompletionsRequest]:
ctx: Final = self.context
config: Final = ctx.provider_config
@ -507,7 +439,6 @@ def dispatch_completion(
fallback=operation.afallback,
adapt=operation.adapt,
error_context=error_context,
preflight=operation.assess,
)
return _CHAT.invoke(
prepare=operation.prepare,
@ -515,5 +446,4 @@ def dispatch_completion(
fallback=operation.fallback,
adapt=operation.adapt,
error_context=error_context,
preflight=operation.assess,
)

View file

@ -21,7 +21,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.protocols import RustAmessages, RustMessages, RustRouteDecline
from litellm.rust_bridge.protocols import RustAmessages, RustMessages
from litellm.rust_bridge.request import (
NativeMessagesRequest,
NativePreCallDetails,
@ -35,10 +35,7 @@ from litellm.rust_bridge.request import (
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointBinding,
EndpointDispatch,
PythonFallback,
assess_route,
async_none,
identity,
)
@ -57,24 +54,11 @@ _MESSAGES: Final[EndpointDispatch[RustMessages, RustAmessages]] = EndpointDispat
)
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
route="messages",
select=lambda native: native.messages_decline,
enabled=rust_enabled,
)
def set_rust_messages(
*,
messages: RustMessages | None | Unchanged = UNCHANGED,
amessages: RustAmessages | None | Unchanged = UNCHANGED,
decline: RustRouteDecline | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(decline, Unchanged):
if decline is None:
_PREFLIGHT.reset()
else:
_PREFLIGHT.override(decline)
if not isinstance(messages, Unchanged):
if messages is None:
_MESSAGES.sync.reset()
@ -133,7 +117,6 @@ def messages(
),
),
call=call_native,
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
fallback=lambda: None,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -178,7 +161,6 @@ async def amessages(
),
),
call=call_native,
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -368,16 +350,6 @@ def dispatch_messages(
has_custom_client,
)
def preflight() -> PythonFallback | None:
return assess_route(
_PREFLIGHT,
model,
provider,
stream=stream,
has_custom_client=has_custom_client,
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(logging),
)
error_context: Final = BridgeErrorContext(provider=provider, model=model)
if asynchronous:
return _MESSAGES.ainvoke(
@ -385,7 +357,6 @@ def dispatch_messages(
call=call_native,
adapt=operation.adapt,
fallback=operation.afallback,
preflight=preflight,
error_context=error_context,
)
return _MESSAGES.invoke(
@ -393,6 +364,5 @@ def dispatch_messages(
call=call_native,
adapt=operation.adapt,
fallback=operation.fallback,
preflight=preflight,
error_context=error_context,
)

View file

@ -6,13 +6,11 @@ from collections.abc import Awaitable, Callable, Mapping
from typing import Final, TypeVar
from . import configuration as _configuration
from .protocols import RustAocr, RustOcr, RustRouteDecline
from .protocols import RustAocr, RustOcr
from .request import NativeOCRRequest, PreparedNativeCall, call_native
from .runtime import (
BridgeErrorContext,
EndpointBinding,
EndpointDispatch,
assess_route,
)
ResultT = TypeVar("ResultT")
@ -26,13 +24,6 @@ _OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native(
)
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
route="ocr",
select=lambda native: native.ocr_decline,
enabled=_configuration.rust_enabled,
)
def load_rust_ocr() -> RustOcr | None:
return _OCR.sync.load()
@ -73,7 +64,6 @@ def dispatch_ocr(
adapt=adapt,
error_context=BridgeErrorContext(provider=provider, model=model),
eligible=eligible,
preflight=lambda: assess_route(_PREFLIGHT, model, provider, request_format=request_format),
)
@ -94,5 +84,4 @@ async def adispatch_ocr(
adapt=adapt,
error_context=BridgeErrorContext(provider=provider, model=model),
eligible=eligible,
preflight=lambda: assess_route(_PREFLIGHT, model, provider, request_format=request_format),
)

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from collections.abc import Awaitable, Mapping
from typing import Protocol
from .callbacks import SessionCallbackHandle
@ -25,19 +25,6 @@ RustTranscription = NativeFunction[NativeTranscriptionRequest, dict[str, object]
RustAtranscription = NativeFunction[NativeTranscriptionRequest, Awaitable[dict[str, object]]]
class RustChatCompletionsDecline(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
custom_llm_provider: str | None,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> str | None: ...
class RustResponsesWebSocket(Protocol):
async def send_text(self, text: str) -> None: ...
@ -58,16 +45,6 @@ class RustResponsesWebSocketConnection(Protocol):
) -> RustResponsesWebSocket: ...
class RustRouteDecline(Protocol):
def __call__(
self,
model: str,
custom_llm_provider: str,
*,
context: NativeRequestContext,
) -> str | None: ...
class NativeModule(Protocol):
@property
def chat_completions(self) -> RustChatCompletions: ...
@ -75,9 +52,6 @@ class NativeModule(Protocol):
@property
def achat_completions(self) -> RustAchatCompletions: ...
@property
def chat_completions_decline(self) -> RustChatCompletionsDecline: ...
@property
def ResponsesWebSocketConnection(self) -> type[RustResponsesWebSocketConnection]: ...
@ -104,15 +78,3 @@ class NativeModule(Protocol):
@property
def atranscription(self) -> RustAtranscription: ...
@property
def ocr_decline(self) -> RustRouteDecline: ...
@property
def messages_decline(self) -> RustRouteDecline: ...
@property
def transcription_decline(self) -> RustRouteDecline: ...
@property
def responses_websocket_decline(self) -> RustRouteDecline: ...

View file

@ -15,7 +15,6 @@ from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.protocols import (
RustResponsesWebSocket,
RustResponsesWebSocketConnection,
RustRouteDecline,
)
from litellm.rust_bridge.request import (
NativeRequestCapabilities,
@ -29,7 +28,6 @@ from litellm.rust_bridge.request import (
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointBinding,
assess_route,
async_none,
)
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -41,23 +39,10 @@ _RESPONSES_WEBSOCKET: Final[EndpointBinding[RustResponsesWebSocketConnection]] =
)
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
route="responses_websocket",
select=lambda native: native.responses_websocket_decline,
enabled=rust_enabled,
)
def set_rust_responses_websocket(
*,
connection: RustResponsesWebSocketConnection | None | Unchanged = UNCHANGED,
decline: RustRouteDecline | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(decline, Unchanged):
if decline is None:
_PREFLIGHT.reset()
else:
_PREFLIGHT.override(decline)
if not isinstance(connection, Unchanged):
if connection is None:
_RESPONSES_WEBSOCKET.reset()
@ -120,7 +105,6 @@ async def connect(
callback_adapter=callback_adapter,
),
call=lambda connection_type, request: call_native(connection_type.connect, request),
preflight=lambda: assess_route(_PREFLIGHT, model, provider),
fallback=fallback,
adapt=_ConnectionAdapter,
error_context=BridgeErrorContext(provider=provider, model=model),

View file

@ -13,8 +13,7 @@ from litellm.rust_bridge.bindings import (
native_declined_types,
native_upstream_types,
)
from litellm.rust_bridge.protocols import NativeModule, RustRouteDecline
from litellm.rust_bridge.request import NativeRequestCapabilities, NativeRequestContext
from litellm.rust_bridge.protocols import NativeModule
BindingT = TypeVar("BindingT")
SelectedT = TypeVar("SelectedT")
@ -119,16 +118,12 @@ 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,
@ -143,16 +138,12 @@ 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,
@ -168,7 +159,6 @@ 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,
@ -176,7 +166,6 @@ class EndpointBinding(Generic[BindingT]):
adapt=adapt,
error_context=error_context,
eligible=eligible,
preflight=preflight,
)
match result:
case Handled(value=value):
@ -195,7 +184,6 @@ 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,
@ -203,7 +191,6 @@ class EndpointBinding(Generic[BindingT]):
adapt=adapt,
error_context=error_context,
eligible=eligible,
preflight=preflight,
)
match result:
case Handled(value=value):
@ -213,34 +200,6 @@ class EndpointBinding(Generic[BindingT]):
case _ as unreachable:
assert_never(unreachable)
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,
*,
check: Callable[[BindingT], str | None],
eligible: bool = True,
) -> bool:
binding_or_fallback: Final = self._binding_or_python_fallback(
eligible=eligible,
)
if isinstance(binding_or_fallback, PythonFallback):
return False
try:
reason: Final = check(binding_or_fallback)
except Exception: # noqa: BLE001 # preflight performs no provider I/O, so Python handoff is safe
return False
return reason is None
def require(
self,
*,
@ -249,7 +208,6 @@ 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,
@ -257,7 +215,6 @@ class EndpointBinding(Generic[BindingT]):
adapt=adapt,
error_context=error_context,
eligible=eligible,
preflight=preflight,
)
match result:
case Handled(value=value):
@ -275,7 +232,6 @@ 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,
@ -283,7 +239,6 @@ class EndpointBinding(Generic[BindingT]):
adapt=adapt,
error_context=error_context,
eligible=eligible,
preflight=preflight,
)
match result:
case Handled(value=value):
@ -293,16 +248,6 @@ class EndpointBinding(Generic[BindingT]):
case _ as unreachable:
assert_never(unreachable)
def can_attempt(
self,
*,
eligible: bool = True,
) -> bool:
return not isinstance(
self._binding_or_python_fallback(eligible=eligible),
PythonFallback,
)
def _raise_required(self, fallback: PythonFallback) -> NoReturn:
detail: Final = f": {fallback.detail}" if fallback.detail else ""
reason: Final = _required_reason(fallback.reason)
@ -441,7 +386,6 @@ 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,
@ -450,7 +394,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
adapt=adapt,
error_context=error_context,
eligible=eligible,
preflight=preflight,
)
async def ainvoke(
@ -462,7 +405,6 @@ 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,
@ -471,7 +413,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
adapt=adapt,
error_context=error_context,
eligible=eligible,
preflight=preflight,
)
def require(
@ -482,7 +423,6 @@ 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,
@ -490,7 +430,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
adapt=adapt,
error_context=error_context,
eligible=eligible,
preflight=preflight,
)
async def arequire(
@ -501,7 +440,6 @@ 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,
@ -509,7 +447,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
adapt=adapt,
error_context=error_context,
eligible=eligible,
preflight=preflight,
)
@ -544,30 +481,3 @@ def adapt_result(result: DispatchResult[NativeT], adapt: Callable[[NativeT], Res
async def async_none() -> None:
return None
def assess_route(
binding: EndpointBinding[RustRouteDecline],
model: str,
provider: str,
*,
stream: bool = False,
has_agentic_hook: bool = False,
has_custom_client: bool = False,
request_format: str | None = None,
) -> PythonFallback | None:
context: Final = NativeRequestContext(
capabilities=NativeRequestCapabilities(
stream=stream,
has_agentic_hook=has_agentic_hook,
has_custom_client=has_custom_client,
request_format=request_format,
)
)
return binding.assess(
check=lambda decline: decline(
model,
provider,
context=context,
),
)

View file

@ -17,7 +17,7 @@ from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.protocols import RustAtranscription, RustRouteDecline, RustTranscription
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
from litellm.rust_bridge.request import (
NativePreCallDetails,
NativeRequestCapabilities,
@ -32,11 +32,8 @@ from litellm.rust_bridge.request import (
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointBinding,
EndpointDispatch,
PythonFallback,
always_enabled,
assess_route,
async_none,
identity,
)
@ -52,24 +49,11 @@ _TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] =
)
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
route="transcription",
select=lambda native: native.transcription_decline,
enabled=always_enabled,
)
def configure_rust_transcription(
*,
transcription: RustTranscription | None | Unchanged = UNCHANGED,
atranscription: RustAtranscription | None | Unchanged = UNCHANGED,
decline: RustRouteDecline | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(decline, Unchanged):
if decline is None:
_PREFLIGHT.reset()
else:
_PREFLIGHT.override(decline)
if not isinstance(transcription, Unchanged):
if transcription is None:
_TRANSCRIPTION.sync.reset()
@ -131,7 +115,6 @@ def transcription(
),
),
call=call_native,
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
fallback=lambda: None,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -179,7 +162,6 @@ async def atranscription(
),
),
call=call_native,
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -329,15 +311,6 @@ def dispatch_transcription(
has_custom_client,
)
def preflight() -> PythonFallback | None:
return assess_route(
_PREFLIGHT,
model,
provider,
stream=optional_params.get("stream") is True,
has_custom_client=has_custom_client,
)
error_context: Final = BridgeErrorContext(provider=provider, model=model)
if provider == "bedrock":
if asynchronous:
@ -346,14 +319,12 @@ def dispatch_transcription(
call=call_native,
adapt=operation.adapt,
error_context=error_context,
preflight=preflight,
)
return _TRANSCRIPTION.require(
prepare=operation.prepare,
call=call_native,
adapt=operation.adapt,
error_context=error_context,
preflight=preflight,
)
if asynchronous:
return _TRANSCRIPTION.ainvoke(
@ -363,7 +334,6 @@ def dispatch_transcription(
fallback=operation.afallback,
error_context=error_context,
eligible=rust_enabled(),
preflight=preflight,
)
return _TRANSCRIPTION.invoke(
prepare=operation.prepare,
@ -372,5 +342,4 @@ def dispatch_transcription(
fallback=operation.fallback,
error_context=error_context,
eligible=rust_enabled(),
preflight=preflight,
)

View file

@ -75,26 +75,14 @@ def _hide_native_bridge(monkeypatch):
@pytest.fixture(autouse=True)
def reset_bridge(monkeypatch):
"""Every test starts with no injected callables, and leaves none behind."""
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
configuration.reset_rust_configuration()
monkeypatch.setenv("LITELLM_RUST", "1")
yield
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
configuration.reset_rust_configuration()
class _RecordingDecline:
"""A stand-in for the native gate that records what it was asked."""
def __init__(self, reason: str | None = None):
self.reason = reason
self.calls: list[dict] = []
def __call__(self, **kwargs):
self.calls.append(kwargs)
return self.reason
class _RecordingCall:
def __init__(self, result=None, error: Exception | None = None):
self.result = result if result is not None else dict(RUST_RESPONSE)
@ -119,131 +107,6 @@ class _RecordingAsyncCall(_RecordingCall):
)
def _accepts(**overrides) -> bool:
kwargs = {
"model": "claude-sonnet-4-5",
"messages": MESSAGES,
"optional_params": {"max_tokens": 16},
"custom_llm_provider": "anthropic",
"litellm_params": {},
"stream": None,
}
kwargs.update(overrides)
return bridge.rust_chat_completions_accepts(**kwargs)
class TestGate:
def test_declines_when_the_deployment_did_not_opt_in(self, monkeypatch):
monkeypatch.delenv("LITELLM_RUST", raising=False)
gate = _RecordingDecline()
bridge.set_rust_chat_completions(decline=gate)
assert _accepts(litellm_params={}) is False
assert _accepts(litellm_params=None) is False
assert gate.calls == [], "the gate must not be consulted before opt-in"
def test_accepts_when_the_deployment_opted_in_and_the_core_agrees(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
gate = _RecordingDecline()
bridge.set_rust_chat_completions(decline=gate)
assert _accepts() is True
assert gate.calls[0]["model"] == "claude-sonnet-4-5"
assert gate.calls[0]["custom_llm_provider"] == "anthropic"
def test_process_enable_applies_without_request_override(self):
bridge.set_rust_chat_completions(decline=_RecordingDecline())
configuration.rust(True)
assert _accepts(litellm_params={}) is True
def test_the_env_var_opts_in_without_a_per_model_flag(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "true")
bridge.set_rust_chat_completions(decline=_RecordingDecline())
assert _accepts(litellm_params={}) is True
def test_declines_streaming_and_providers_off_the_path(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
gate = pytest.importorskip("litellm.rust_bridge._native").chat_completions_decline
bridge.set_rust_chat_completions(decline=gate)
assert _accepts(stream=True) is False
assert _accepts(custom_llm_provider="openai") is False
assert _accepts(custom_llm_provider=None) is False
def test_declines_an_anthropic_request_carrying_a_litellm_metadata_user_id(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
gate = pytest.importorskip("litellm.rust_bridge._native").chat_completions_decline
bridge.set_rust_chat_completions(decline=gate)
assert _accepts(litellm_params={"metadata": {"user_id": "u-123"}}) is False
# Bedrock's Converse transform reads no `user_id`, and an Anthropic request
# whose metadata carries none is one Python would not attribute either.
assert (
_accepts(
custom_llm_provider="bedrock",
model="bedrock/us-east-1/anthropic.claude-v2",
optional_params={"maxTokens": 16},
litellm_params={"metadata": {"user_id": "u-123"}},
)
is True
)
assert _accepts(litellm_params={"metadata": {"trace_id": "t-1"}}) is True
assert _accepts(litellm_params={"metadata": {"user_id": None}}) is True
assert _accepts(litellm_params={"metadata": None}) is True
assert _accepts(litellm_params={"litellm_metadata": {"user_id": "u-123"}}) is True
assert _accepts(litellm_params={"metadata": "invalid"}) is True
assert _accepts(litellm_params={"metadata": {"trace": object()}}) is True
assert _accepts(litellm_params={"metadata": {"user_id": object()}}) is False
assert (
_accepts(
custom_llm_provider="bedrock",
model="bedrock/us-east-1/anthropic.claude-v2",
optional_params={"maxTokens": 16},
litellm_params={"metadata": {"user_id": object()}},
)
is True
)
def test_declines_a_bedrock_request_while_the_proxy_owns_request_metadata(self, monkeypatch):
"""`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the
Converse body from `litellm_params`, and owning that field also means
evicting a caller-supplied one. The core can do neither, so an operator
who armed `bedrock_request_metadata_fields` keeps the Python path.
"""
monkeypatch.setenv("LITELLM_RUST", "1")
gate = pytest.importorskip("litellm.rust_bridge._native").chat_completions_decline
bridge.set_rust_chat_completions(decline=gate)
bedrock = {
"custom_llm_provider": "bedrock",
"model": "bedrock/us-east-1/anthropic.claude-v2",
"optional_params": {"maxTokens": 16},
}
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["user_api_key_team_id"])
assert _accepts(**bedrock) is False
assert _accepts() is True, "arming Bedrock attribution must not decline Anthropic"
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None)
assert _accepts(**bedrock) is True, "the decline follows the operator's opt-in alone"
def test_declines_when_the_core_declines(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
bridge.set_rust_chat_completions(decline=_RecordingDecline("streaming"))
assert _accepts() is False
def test_declines_when_the_bridge_is_unavailable(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
_hide_native_bridge(monkeypatch)
assert _accepts() is False
def test_declines_when_the_gate_itself_raises(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
def exploding(**_kwargs):
raise RuntimeError("boom")
bridge.set_rust_chat_completions(decline=exploding)
assert _accepts() is False
def _call_kwargs(model_response: ModelResponse) -> dict:
return {
"model": "claude-sonnet-4-5",
@ -494,8 +357,7 @@ def test_typed_capability_and_provider_metadata_facts_are_isolated():
async def test_public_completion_discovers_any_provider(provider, asynchronous):
native = _RecordingCall()
anative = _RecordingAsyncCall()
gate = _RecordingDecline()
bridge.set_rust_chat_completions(chat_completions=native, achat_completions=anative, decline=gate)
bridge.set_rust_chat_completions(chat_completions=native, achat_completions=anative)
kwargs = {
"model": f"{provider}/test-model",
"messages": MESSAGES,
@ -509,13 +371,12 @@ async def test_public_completion_discovers_any_provider(provider, asynchronous):
assert len(calls) == 1
assert calls[0]["options"].custom_llm_provider == provider
assert calls[0]["request"].messages == MESSAGES
assert gate.calls[0]["custom_llm_provider"] == provider
assert len(native.calls) + len(anative.calls) == 1
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize(
"failure", ["preflight", "missing_preflight", "decline", "unavailable", "error", "malformed", "cancelled"]
"failure", ["decline", "unavailable", "error", "malformed", "cancelled"]
)
@pytest.mark.asyncio
async def test_public_completion_fallback_contract(monkeypatch, asynchronous, failure):
@ -564,10 +425,7 @@ async def test_public_completion_fallback_contract(monkeypatch, asynchronous, fa
bridge.set_rust_chat_completions(
chat_completions=native,
achat_completions=anative,
decline=_RecordingDecline("unsupported" if failure == "preflight" else None),
)
if failure == "missing_preflight":
bridge._CHAT_PREFLIGHT.override(None)
if failure == "unavailable":
bridge._CHAT.sync.override(None)
bridge._CHAT.asynchronous.override(None)

View file

@ -292,37 +292,6 @@ def test_require_explains_why_rust_did_not_handle_request(
)
@pytest.mark.parametrize(
("state", "expected", "expected_events"),
(
pytest.param("disabled", False, (), id="disabled"),
pytest.param("ineligible", False, (), id="ineligible"),
pytest.param("unavailable", False, ("load",), id="unavailable"),
pytest.param("available", True, ("load",), id="available"),
),
)
def test_can_attempt_only_enabled_available_requests(
state: str,
expected: bool,
expected_events: tuple[str, ...],
) -> None:
events: list[str] = []
def load() -> object | None:
events.append("load")
return None if state == "unavailable" else object()
bridge: Final = runtime.EndpointBinding(route="messages", load=load, enabled=lambda: state != "disabled")
assert (
bridge.can_attempt(
eligible=state != "ineligible",
)
is expected
)
assert tuple(events) == expected_events
def test_native_endpoint_applies_partial_overrides_and_reset(monkeypatch: pytest.MonkeyPatch) -> None:
def native_sync() -> str:
return "native"
@ -395,77 +364,6 @@ async def test_response_adaptation_failure_never_authorizes_fallback(asynchronou
await invoke()
@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,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"route",
@ -495,12 +393,10 @@ async def test_unready_routes_never_prepare_or_call_native(
)
arguments: Final = {
"prepare": unexpected,
"preflight": unexpected,
"call": unexpected,
"adapt": unexpected,
"error_context": runtime.BridgeErrorContext(provider="test", model="test-model"),
}
assert not endpoint.can_attempt()
assert endpoint.invoke(**arguments, fallback=lambda: "python") == "python"
with pytest.raises(RuntimeError, match=f"native {route} endpoint is unavailable"):
endpoint.require(**arguments)