mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor(native): replace preflight APIs with typed execution declines
This commit is contained in:
parent
8e20a8644f
commit
2e9a24f67d
15 changed files with 31 additions and 683 deletions
|
|
@ -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",
|
||||
];
|
||||
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue