mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge d9ad7ae2e6 into d2fc8b2438
This commit is contained in:
commit
0f28962e00
41 changed files with 283 additions and 1581 deletions
|
|
@ -37,7 +37,7 @@ Python builds the frozen request dataclasses in `litellm/rust_bridge/request.py`
|
|||
PyO3 extracts their fields before execution. AWS credentials and metadata policy
|
||||
belong in `options.bedrock`; Vertex project/location belongs in `options.vertex`.
|
||||
|
||||
This boundary preserves existing Python provider preparation, preflight decisions,
|
||||
This boundary preserves existing Python provider preparation and admission decisions,
|
||||
fallback, and callbacks
|
||||
|
||||
## Crates
|
||||
|
|
|
|||
|
|
@ -275,7 +275,6 @@ fn guardrail_error_to_core_error(error: GuardrailError) -> Error {
|
|||
|
||||
fn core_error_kind(error: &Error) -> &'static str {
|
||||
match error {
|
||||
Error::Declined(_) => "UnsupportedRequest",
|
||||
Error::Auth(_) => "AuthError",
|
||||
Error::InvalidProvider(_) => "InvalidProvider",
|
||||
Error::InvalidRequest(_) => "InvalidRequest",
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
pub use crate::ocr::{OcrRequest, ocr, ocr_provider_supported, ocr_with_observer};
|
||||
pub use crate::ocr::{OcrRequest, ocr, ocr_admitted, ocr_with_observer};
|
||||
|
|
|
|||
|
|
@ -380,7 +380,6 @@ fn guardrail_error_to_core_error(error: GuardrailError) -> Error {
|
|||
|
||||
fn core_error_kind(error: &Error) -> &'static str {
|
||||
match error {
|
||||
Error::Declined(_) => "UnsupportedRequest",
|
||||
Error::Auth(_) => "AuthError",
|
||||
Error::InvalidProvider(_) => "InvalidProvider",
|
||||
Error::InvalidRequest(_) => "InvalidRequest",
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ where
|
|||
.await
|
||||
}
|
||||
|
||||
pub fn ocr_provider_supported(model: &str, provider: &str, request_format: Option<&str>) -> bool {
|
||||
pub fn ocr_admitted(model: &str, provider: &str, request_format: Option<&str>) -> bool {
|
||||
common_utils::ocr_provider_config(provider, model).is_some_and(|config| {
|
||||
request_format != Some("native") || config.supported_ocr_params().contains(&"req_format")
|
||||
})
|
||||
|
|
|
|||
|
|
@ -121,7 +121,7 @@ impl IntoResponse for MessagesRouteError {
|
|||
// The gateway has no Python implementation to decline to, so a
|
||||
// request the core cannot serve is reported to the caller. The
|
||||
// reason is a fixed internal string, never provider content.
|
||||
Error::Declined(reason) | Error::Unsupported(reason) => (
|
||||
Error::Unsupported(reason) => (
|
||||
StatusCode::BAD_REQUEST,
|
||||
format!("messages request is not supported: {reason}"),
|
||||
),
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ pub async fn audio_transcription(
|
|||
.await
|
||||
}
|
||||
|
||||
pub fn transcription_provider_supported(provider: &str) -> bool {
|
||||
pub fn transcription_admitted(provider: &str) -> bool {
|
||||
prepare::provider_config(provider).is_some()
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@
|
|||
|
||||
use crate::Error;
|
||||
use crate::eligibility::native_route_decline;
|
||||
use crate::native_outcome::{Decline, NativeOutcome};
|
||||
use crate::request_context::LiteLlmRequestContext;
|
||||
use crate::request_options::RequestOptions;
|
||||
mod client;
|
||||
|
|
@ -31,18 +32,24 @@ pub async fn chat_completions(
|
|||
request: ChatCompletionsRequest<'_>,
|
||||
options: &RequestOptions,
|
||||
context: &LiteLlmRequestContext,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
) -> Result<NativeOutcome<ChatCompletionsResponse>, Error> {
|
||||
if let Some(reason) = chat_completions_admission(
|
||||
request.model,
|
||||
options.custom_llm_provider.as_deref(),
|
||||
request.messages.clone(),
|
||||
&request.optional_params,
|
||||
options,
|
||||
context,
|
||||
) {
|
||||
return Ok(NativeOutcome::Declined(Decline::new(reason)));
|
||||
}
|
||||
execute_chat_completions_provider_call(resolve_request(request, options.clone(), context)?)
|
||||
.await
|
||||
.map(NativeOutcome::Completed)
|
||||
}
|
||||
|
||||
/// Whether the core would accept this request, without resolving credentials or
|
||||
/// touching the network.
|
||||
///
|
||||
/// A host that keeps the Python implementation asks this first so it can emit
|
||||
/// its pre-call logging exactly once, on whichever path is about to run.
|
||||
/// Returns the decline reason, or `None` when the request is accepted.
|
||||
pub fn chat_completions_decline_reason(
|
||||
/// Pure admission for the normal route entrypoint.
|
||||
fn chat_completions_admission(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
messages: Value,
|
||||
|
|
|
|||
|
|
@ -50,11 +50,11 @@ pub(super) fn resolve_request(
|
|||
) -> Result<ResolvedChatCompletionsRequest, Error> {
|
||||
let (model, provider, config) =
|
||||
resolve_provider_config(request.model, options.custom_llm_provider.as_deref())
|
||||
.map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?;
|
||||
let messages =
|
||||
parse_messages(request.messages).map_err(|_| Error::Declined("unreadable message list"))?;
|
||||
.map_err(|_| Error::Unsupported("provider is not on the rust chat completions path"))?;
|
||||
let messages = parse_messages(request.messages)
|
||||
.map_err(|_| Error::Unsupported("unreadable message list"))?;
|
||||
if messages.is_empty() {
|
||||
return Err(Error::Declined("empty message list"));
|
||||
return Err(Error::Unsupported("empty message list"));
|
||||
}
|
||||
if let Some(reason) = super::unsupported_reason(
|
||||
provider,
|
||||
|
|
@ -64,7 +64,7 @@ pub(super) fn resolve_request(
|
|||
&options,
|
||||
context,
|
||||
) {
|
||||
return Err(Error::Declined(reason.0));
|
||||
return Err(Error::Unsupported(reason.0));
|
||||
}
|
||||
Ok(ResolvedChatCompletionsRequest {
|
||||
model,
|
||||
|
|
|
|||
|
|
@ -220,7 +220,7 @@ fn declines_an_unsupported_request_before_resolving_credentials() {
|
|||
call.options.api_key = None;
|
||||
// No api_key is set and no env is consulted: the gate must run first, so the
|
||||
// error is the decline rather than a missing-credential error.
|
||||
assert_eq!(decline(call), Error::Declined("streaming"));
|
||||
assert_eq!(decline(call), Error::Unsupported("streaming"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -232,7 +232,7 @@ fn rejects_an_unknown_provider() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
)),
|
||||
Error::Declined("provider is not on the rust chat completions path")
|
||||
Error::Unsupported("provider is not on the rust chat completions path")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -245,7 +245,7 @@ fn rejects_a_model_with_no_resolvable_provider() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
)),
|
||||
Error::Declined("provider is not on the rust chat completions path")
|
||||
Error::Unsupported("provider is not on the rust chat completions path")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -258,7 +258,7 @@ fn rejects_an_empty_or_malformed_message_list() {
|
|||
json!([]),
|
||||
json!({}),
|
||||
)),
|
||||
Error::Declined("empty message list")
|
||||
Error::Unsupported("empty message list")
|
||||
);
|
||||
assert_eq!(
|
||||
decline(request(
|
||||
|
|
@ -267,7 +267,7 @@ fn rejects_an_empty_or_malformed_message_list() {
|
|||
json!("not a list"),
|
||||
json!({}),
|
||||
)),
|
||||
Error::Declined("unreadable message list")
|
||||
Error::Unsupported("unreadable message list")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
@ -513,7 +513,7 @@ fn decline_reason(
|
|||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
};
|
||||
super::chat_completions_decline_reason(
|
||||
super::chat_completions_admission(
|
||||
model,
|
||||
provider,
|
||||
messages,
|
||||
|
|
@ -524,7 +524,7 @@ fn decline_reason(
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_accepts_what_prepare_accepts() {
|
||||
fn admission_accepts_a_supported_call() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
|
|
@ -537,7 +537,7 @@ fn the_gate_accepts_what_prepare_accepts() {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
|
||||
fn admission_declines_without_resolving_credentials_or_calling_out() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
|
|
@ -581,9 +581,7 @@ fn the_gate_declines_without_resolving_credentials_or_calling_out() {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
|
||||
// A gate that accepts what prepare then declines would make the host emit
|
||||
// its pre-call logging on a path that falls back, so pin the agreement.
|
||||
fn admission_agrees_with_preparation_on_supported_cases() {
|
||||
for (messages, params) in [
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
|
|
@ -606,7 +604,7 @@ fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
|
|||
params.clone()
|
||||
),
|
||||
None,
|
||||
"gate declined {messages}"
|
||||
"admission declined {messages}"
|
||||
);
|
||||
prepare_chat_completions_call(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
|
|
@ -711,7 +709,12 @@ mod round_trip {
|
|||
call: TestChatCompletionsCall<'_>,
|
||||
context: &LiteLlmRequestContext,
|
||||
) -> Result<super::super::types::ChatCompletionsResponse, Error> {
|
||||
run_chat_completions(call.request, &call.options, context).await
|
||||
match run_chat_completions(call.request, &call.options, context).await? {
|
||||
crate::native_outcome::NativeOutcome::Completed(response) => Ok(response),
|
||||
crate::native_outcome::NativeOutcome::Declined(decline) => {
|
||||
panic!("round-trip fixture was declined: {}", decline.reason())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
|
||||
|
|
@ -886,7 +889,7 @@ mod round_trip {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn preflight_and_execution_share_provider_metadata_eligibility() {
|
||||
fn admission_and_preparation_share_provider_metadata_eligibility() {
|
||||
let messages = json!([{"role": "user", "content": "hi"}]);
|
||||
let cases = [
|
||||
(
|
||||
|
|
@ -930,7 +933,7 @@ fn preflight_and_execution_share_provider_metadata_eligibility() {
|
|||
for (provider, options, expected_decline) in cases {
|
||||
let context = LiteLlmRequestContext::default();
|
||||
let params = Map::new();
|
||||
let preflight = super::chat_completions_decline_reason(
|
||||
let admission = super::chat_completions_admission(
|
||||
"claude-sonnet-4-5",
|
||||
Some(provider),
|
||||
messages.clone(),
|
||||
|
|
@ -948,9 +951,9 @@ fn preflight_and_execution_share_provider_metadata_eligibility() {
|
|||
&context,
|
||||
);
|
||||
assert_eq!(
|
||||
preflight.is_some(),
|
||||
admission.is_some(),
|
||||
expected_decline,
|
||||
"{provider} preflight"
|
||||
"{provider} admission"
|
||||
);
|
||||
assert_eq!(execution.is_err(), expected_decline, "{provider} execution");
|
||||
}
|
||||
|
|
|
|||
|
|
@ -20,10 +20,10 @@ impl NativeRouteDecline {
|
|||
}
|
||||
|
||||
pub fn native_route_decline(
|
||||
provider_supported: bool,
|
||||
provider_admitted: bool,
|
||||
capabilities: &RequestCapabilities,
|
||||
) -> Option<NativeRouteDecline> {
|
||||
if !provider_supported {
|
||||
if !provider_admitted {
|
||||
return Some(NativeRouteDecline::UnsupportedProvider);
|
||||
}
|
||||
if capabilities.stream {
|
||||
|
|
|
|||
|
|
@ -2,8 +2,6 @@ use thiserror::Error as ThisError;
|
|||
|
||||
#[derive(Debug, ThisError, PartialEq, Eq)]
|
||||
pub enum Error {
|
||||
#[error("native execution declined: {0}")]
|
||||
Declined(&'static str),
|
||||
#[error("expected {expected}, got {actual}")]
|
||||
InvalidType {
|
||||
expected: &'static str,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ pub mod error;
|
|||
pub mod hook_contracts;
|
||||
pub mod http_utils;
|
||||
pub mod messages;
|
||||
pub mod native_outcome;
|
||||
#[cfg(any(feature = "observability", test))]
|
||||
pub mod observability;
|
||||
pub mod ocr;
|
||||
|
|
|
|||
|
|
@ -37,7 +37,7 @@ pub async fn messages_stream(
|
|||
execute_messages_provider_stream(request, options.clone()).await
|
||||
}
|
||||
|
||||
pub fn messages_provider_supported(provider: &str) -> bool {
|
||||
pub fn messages_admitted(provider: &str) -> bool {
|
||||
common_utils::messages_provider_config(provider).is_some()
|
||||
}
|
||||
|
||||
|
|
|
|||
29
litellm-rust/crates/core/src/native_outcome.rs
Normal file
29
litellm-rust/crates/core/src/native_outcome.rs
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct Decline {
|
||||
reason: &'static str,
|
||||
}
|
||||
|
||||
impl Decline {
|
||||
pub const fn new(reason: &'static str) -> Self {
|
||||
Self { reason }
|
||||
}
|
||||
|
||||
pub const fn reason(self) -> &'static str {
|
||||
self.reason
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum NativeOutcome<T> {
|
||||
Completed(T),
|
||||
Declined(Decline),
|
||||
}
|
||||
|
||||
impl<T> NativeOutcome<T> {
|
||||
pub fn map<U>(self, map: impl FnOnce(T) -> U) -> NativeOutcome<U> {
|
||||
match self {
|
||||
Self::Completed(value) => NativeOutcome::Completed(map(value)),
|
||||
Self::Declined(decline) => NativeOutcome::Declined(decline),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -18,7 +18,6 @@ pyo3::create_exception!(
|
|||
|
||||
pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr {
|
||||
match err {
|
||||
Error::Declined(message) => RustBridgeDeclined::new_err(message),
|
||||
Error::Auth(message) => PyValueError::new_err(message),
|
||||
Error::InvalidProvider(_)
|
||||
| Error::InvalidRequest(_)
|
||||
|
|
@ -30,7 +29,6 @@ pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr {
|
|||
|
||||
pub(crate) fn chat_completions_error_to_pyerr(err: Error) -> PyErr {
|
||||
match err {
|
||||
Error::Declined(message) => RustBridgeDeclined::new_err(message),
|
||||
Error::Auth(message) => RustUpstreamError::new_err((401u16, message)),
|
||||
Error::Http { status, body } => {
|
||||
RustUpstreamError::new_err((status, format!("{status}: {body}")))
|
||||
|
|
@ -91,15 +89,6 @@ mod tests {
|
|||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn only_explicit_decline_authorizes_python_fallback(#[from(initialized_python)] (): ()) {
|
||||
Python::attach(|py| {
|
||||
let mapped = chat_completions_error_to_pyerr(Error::Declined("unsupported request"));
|
||||
assert!(mapped.is_instance_of::<RustBridgeDeclined>(py));
|
||||
assert_eq!(mapped.value(py).to_string(), "unsupported request");
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn request_failures_do_not_authorize_python_fallback(#[from(initialized_python)] (): ()) {
|
||||
Python::attach(|py| {
|
||||
|
|
|
|||
|
|
@ -47,12 +47,14 @@ impl ResponsesWebSocketConnection {
|
|||
context: NativeRequestContext,
|
||||
callback_adapter: Option<Py<PyAny>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let provider_supported = litellm_core::responses::websocket::native_websocket_supported(
|
||||
let provider_admitted = litellm_core::responses::websocket::native_websocket_supported(
|
||||
options.provider("openai"),
|
||||
);
|
||||
let context: litellm_core::request_context::LiteLlmRequestContext = context.into();
|
||||
if let Some(reason) = routes::definition::request_decline(provider_supported, &context) {
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
if let litellm_core::native_outcome::NativeOutcome::Declined(decline) =
|
||||
routes::definition::admission(provider_admitted, &context)
|
||||
{
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(decline.reason()));
|
||||
}
|
||||
let options: litellm_core::request_options::RequestOptions = options.into();
|
||||
let call_id = context.litellm_call_id.clone().unwrap_or_default();
|
||||
|
|
@ -141,35 +143,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 +182,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",
|
||||
];
|
||||
|
||||
|
|
|
|||
|
|
@ -11,8 +11,7 @@ use std::future::Future;
|
|||
#[derive(FromPyObject)]
|
||||
struct AudioTranscriptionInputs {
|
||||
model: String,
|
||||
#[pyo3(from_py_with = litellm_python_interop::from_py)]
|
||||
audio: Value,
|
||||
audio: Py<PyAny>,
|
||||
#[pyo3(from_py_with = litellm_python_interop::from_py)]
|
||||
optional_params: Map<String, Value>,
|
||||
}
|
||||
|
|
@ -22,16 +21,22 @@ fn prepare_transcription(
|
|||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
_callback_adapter: Option<Py<PyAny>>,
|
||||
_python_context: crate::execution::PythonCallContext<'_>,
|
||||
python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
|
||||
let provider_supported = litellm_core::audio_transcription::transcription_provider_supported(
|
||||
options.provider("bedrock"),
|
||||
);
|
||||
let provider_admitted =
|
||||
litellm_core::audio_transcription::transcription_admitted(options.provider("bedrock"));
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
if let Some(reason) = super::definition::request_decline(provider_supported, &context) {
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
if let litellm_core::native_outcome::NativeOutcome::Declined(decline) =
|
||||
super::definition::admission(provider_admitted, &context)
|
||||
{
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(decline.reason()));
|
||||
}
|
||||
let audio = input.audio;
|
||||
let py = python_context.py;
|
||||
let audio = py
|
||||
.import("litellm.rust_bridge.transcription")?
|
||||
.getattr("_consume_audio_for_native")?
|
||||
.call1((input.audio.bind(py),))?;
|
||||
let audio: Value = litellm_python_interop::from_py(&audio)?;
|
||||
Ok(async move {
|
||||
run_route(
|
||||
AudioTranscriptionRequest {
|
||||
|
|
@ -46,25 +51,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],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,14 +2,19 @@ use crate::errors::chat_completions_error_to_pyerr;
|
|||
use crate::marshal::{NativeRequestContext, NativeRequestOptions, required_value};
|
||||
use litellm_core::Error;
|
||||
use litellm_core::chat_completions::chat_completions as run_route;
|
||||
use litellm_core::chat_completions::chat_completions_decline_reason;
|
||||
use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse};
|
||||
use litellm_core::native_outcome::NativeOutcome;
|
||||
use litellm_core::request_context::LiteLlmRequestContext;
|
||||
use litellm_core::request_options::RequestOptions;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
use std::future::Future;
|
||||
|
||||
enum ChatCompletionsRouteError {
|
||||
Declined(String),
|
||||
Terminal(Error),
|
||||
}
|
||||
|
||||
#[derive(FromPyObject)]
|
||||
struct ChatCompletionsInputs {
|
||||
model: String,
|
||||
|
|
@ -25,57 +30,40 @@ fn prepare_chat_completions(
|
|||
context: NativeRequestContext,
|
||||
_callback_adapter: Option<Py<PyAny>>,
|
||||
_python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
|
||||
) -> PyResult<
|
||||
impl Future<Output = Result<ChatCompletionsResponse, ChatCompletionsRouteError>> + Send + 'static,
|
||||
> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
let messages = required_value("messages", input.messages, Value::is_array, "list")?;
|
||||
let options: RequestOptions = options.into();
|
||||
Ok(async move {
|
||||
run_route(
|
||||
match run_route(
|
||||
ChatCompletionsRequest {
|
||||
model: &input.model,
|
||||
messages,
|
||||
optional_params: input.optional_params,
|
||||
},
|
||||
&options.into(),
|
||||
&options,
|
||||
&context,
|
||||
)
|
||||
.await
|
||||
.map_err(ChatCompletionsRouteError::Terminal)?
|
||||
{
|
||||
NativeOutcome::Completed(response) => Ok(response),
|
||||
NativeOutcome::Declined(decline) => Err(ChatCompletionsRouteError::Declined(
|
||||
decline.reason().to_string(),
|
||||
)),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
#[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",
|
||||
));
|
||||
fn chat_completions_route_error_to_pyerr(error: ChatCompletionsRouteError) -> PyErr {
|
||||
match error {
|
||||
ChatCompletionsRouteError::Declined(reason) => {
|
||||
crate::errors::RustBridgeDeclined::new_err(reason)
|
||||
}
|
||||
};
|
||||
Ok(chat_completions_decline_reason(
|
||||
&model,
|
||||
custom_llm_provider.as_deref(),
|
||||
messages,
|
||||
&optional_params,
|
||||
&options,
|
||||
&context,
|
||||
)
|
||||
.map(str::to_string))
|
||||
ChatCompletionsRouteError::Terminal(error) => chat_completions_error_to_pyerr(error),
|
||||
}
|
||||
}
|
||||
|
||||
bridge_route! {
|
||||
|
|
@ -83,6 +71,5 @@ bridge_route! {
|
|||
asynchronous = achat_completions,
|
||||
request = ChatCompletionsInputs,
|
||||
prepare = prepare_chat_completions,
|
||||
errors = chat_completions_error_to_pyerr,
|
||||
extra = [chat_completions_decline],
|
||||
errors = chat_completions_route_error_to_pyerr,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -102,12 +102,17 @@ pub(super) fn add_function(
|
|||
module.add_function(function)
|
||||
}
|
||||
|
||||
pub(crate) fn request_decline(
|
||||
provider_supported: bool,
|
||||
pub(crate) fn admission(
|
||||
provider_admitted: bool,
|
||||
context: &litellm_core::request_context::LiteLlmRequestContext,
|
||||
) -> Option<String> {
|
||||
litellm_core::eligibility::native_route_decline(provider_supported, &context.capabilities)
|
||||
.map(|reason| reason.reason().to_string())
|
||||
) -> litellm_core::native_outcome::NativeOutcome<()> {
|
||||
match litellm_core::eligibility::native_route_decline(provider_admitted, &context.capabilities)
|
||||
{
|
||||
Some(reason) => litellm_core::native_outcome::NativeOutcome::Declined(
|
||||
litellm_core::native_outcome::Decline::new(reason.reason()),
|
||||
),
|
||||
None => litellm_core::native_outcome::NativeOutcome::Completed(()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
@ -285,7 +290,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,69 +298,48 @@ 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"
|
||||
import asyncio
|
||||
|
||||
async def invoke_async(execute, request, options):
|
||||
return await execute(request, options=options, context=context)
|
||||
|
||||
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',
|
||||
)
|
||||
unsupported_options = Options(custom_llm_provider='unsupported-native-provider')
|
||||
functions = (
|
||||
(routes.ResponsesWebSocketConnection.connect,)
|
||||
((routes.ResponsesWebSocketConnection.connect, True),)
|
||||
if route == 'responses_websocket'
|
||||
else (getattr(routes, route), getattr(routes, 'a' + route))
|
||||
else ((getattr(routes, route), False), (getattr(routes, 'a' + route), True))
|
||||
)
|
||||
for execute in functions:
|
||||
for execute, is_async in functions:
|
||||
try:
|
||||
execute(request, options=unsupported_options, context=context)
|
||||
if is_async:
|
||||
asyncio.run(invoke_async(execute, request, unsupported_options))
|
||||
else:
|
||||
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");
|
||||
});
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -22,11 +22,13 @@ fn prepare_messages(
|
|||
_callback_adapter: Option<Py<PyAny>>,
|
||||
_python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<AnthropicMessagesResponse, Error>> + Send + 'static> {
|
||||
let provider_supported =
|
||||
litellm_core::messages::messages_provider_supported(options.provider("anthropic"));
|
||||
let provider_admitted =
|
||||
litellm_core::messages::messages_admitted(options.provider("anthropic"));
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
if let Some(reason) = super::definition::request_decline(provider_supported, &context) {
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
if let litellm_core::native_outcome::NativeOutcome::Declined(decline) =
|
||||
super::definition::admission(provider_admitted, &context)
|
||||
{
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(decline.reason()));
|
||||
}
|
||||
let body = required_value("body", input.body, Value::is_object, "dict")?;
|
||||
Ok(async move {
|
||||
|
|
@ -42,25 +44,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],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -13,8 +13,7 @@ use std::future::Future;
|
|||
#[derive(FromPyObject)]
|
||||
struct OcrInputs {
|
||||
model: String,
|
||||
#[pyo3(from_py_with = litellm_python_interop::from_py)]
|
||||
document: Value,
|
||||
document: Py<PyAny>,
|
||||
#[pyo3(from_py_with = litellm_python_interop::from_py)]
|
||||
optional_params: Map<String, Value>,
|
||||
}
|
||||
|
|
@ -27,15 +26,32 @@ fn prepare_ocr(
|
|||
python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
let provider_supported = litellm_ai_gateway::io::ocr::ocr_provider_supported(
|
||||
let provider_admitted = litellm_ai_gateway::io::ocr::ocr_admitted(
|
||||
&input.model,
|
||||
options.provider("mistral"),
|
||||
context.capabilities.request_format.as_deref(),
|
||||
);
|
||||
if let Some(reason) = super::definition::request_decline(provider_supported, &context) {
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
if let litellm_core::native_outcome::NativeOutcome::Declined(decline) =
|
||||
super::definition::admission(provider_admitted, &context)
|
||||
{
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(decline.reason()));
|
||||
}
|
||||
let document = input.document;
|
||||
let py = python_context.py;
|
||||
let document = if input
|
||||
.document
|
||||
.bind(py)
|
||||
.get_item("type")
|
||||
.and_then(|value| value.extract::<String>())
|
||||
.is_ok_and(|kind| kind == "file")
|
||||
{
|
||||
py.import("litellm.ocr.main")?
|
||||
.getattr("convert_file_document_to_url_document")?
|
||||
.call1((input.document.bind(py),))?
|
||||
.unbind()
|
||||
} else {
|
||||
input.document
|
||||
};
|
||||
let document: Value = litellm_python_interop::from_py(document.bind(py))?;
|
||||
let mut observer = PythonProviderObserver::new(callback_adapter, python_context)?;
|
||||
Ok(async move {
|
||||
run_route(
|
||||
|
|
@ -56,27 +72,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],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@ import mimetypes
|
|||
import os
|
||||
import re
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from io import IOBase
|
||||
from typing import Any, Final, cast
|
||||
|
|
@ -82,11 +81,7 @@ def _prepare_ocr_request(
|
|||
|
||||
doc_type = document.get("type")
|
||||
|
||||
if doc_type == "file":
|
||||
document = convert_file_document_to_url_document(document)
|
||||
doc_type = document.get("type")
|
||||
|
||||
if doc_type not in ["document_url", "image_url"]:
|
||||
if doc_type not in ["document_url", "image_url", "file"]:
|
||||
raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url', 'image_url', or 'file'")
|
||||
|
||||
caller_supplied_api_base: Final = api_base is not None
|
||||
|
|
@ -188,7 +183,6 @@ def _prepare_ocr_request(
|
|||
|
||||
def _rust_bridge_optional_params(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> dict[str, object]:
|
||||
optional_params: Final = dict(prepared_request.optional_params)
|
||||
if prepared_request.custom_llm_provider == "vertex_ai":
|
||||
|
|
@ -196,14 +190,11 @@ def _rust_bridge_optional_params(
|
|||
prepared_request.litellm_params.get("vertex_project")
|
||||
or prepared_request.litellm_params.get("vertex_ai_project")
|
||||
or litellm.vertex_project
|
||||
or resolve_secret("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_location: Final = (
|
||||
prepared_request.litellm_params.get("vertex_location")
|
||||
or prepared_request.litellm_params.get("vertex_ai_location")
|
||||
or litellm.vertex_location
|
||||
or resolve_secret("VERTEXAI_LOCATION")
|
||||
or resolve_secret("VERTEX_LOCATION")
|
||||
)
|
||||
if vertex_project is not None:
|
||||
optional_params["vertex_project"] = vertex_project
|
||||
|
|
@ -212,37 +203,16 @@ def _rust_bridge_optional_params(
|
|||
return optional_params
|
||||
|
||||
|
||||
def _rust_bridge_api_base(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> str | None:
|
||||
if prepared_request.api_base is not None:
|
||||
return prepared_request.api_base
|
||||
if prepared_request.custom_llm_provider == "azure_ai":
|
||||
if is_azure_document_intelligence_model(prepared_request.model):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
return resolve_secret("AZURE_AI_API_BASE")
|
||||
return None
|
||||
|
||||
|
||||
def _prepare_rust_ocr_call(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
_resolve_api_key: Callable[[str], str | None],
|
||||
) -> PreparedNativeCall[rust_ocr_bridge.NativeOCRRequest]:
|
||||
provider_config: Final = prepared_request.provider_config
|
||||
api_key_env_var: Final = provider_config.get_api_key_env_var()
|
||||
resolved_api_key: Final = prepared_request.api_key or (
|
||||
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
|
||||
_resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
|
||||
)
|
||||
resolved_headers: Final = provider_config.validate_environment(
|
||||
headers=prepared_request.extra_headers or {},
|
||||
model=prepared_request.model,
|
||||
api_key=resolved_api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
|
||||
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
|
||||
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request)
|
||||
return PreparedNativeCall(
|
||||
request=rust_ocr_bridge.NativeOCRRequest(
|
||||
model=prepared_request.model,
|
||||
|
|
@ -252,11 +222,9 @@ def _prepare_rust_ocr_call(
|
|||
options=NativeRequestOptions(
|
||||
vertex=vertex_options(rust_optional_params),
|
||||
api_key=resolved_api_key,
|
||||
api_base=rust_api_base,
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=cast( # cast-ok: provider header normalization returns string-object pairs
|
||||
dict[str, object], resolved_headers
|
||||
),
|
||||
extra_headers=prepared_request.extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
|
||||
),
|
||||
context=request_context(
|
||||
|
|
@ -287,21 +255,16 @@ class _OCROperation:
|
|||
request: _PreparedOCRRequest
|
||||
resolve_api_key: Callable[[str], str | None]
|
||||
python: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]]
|
||||
logged: bool = False
|
||||
|
||||
def prepare(self) -> PreparedNativeCall[rust_ocr_bridge.NativeOCRRequest]:
|
||||
prepared: Final = _prepare_rust_ocr_call(self.request, self.resolve_api_key)
|
||||
self.logged = True
|
||||
return prepared
|
||||
return _prepare_rust_ocr_call(self.request, self.resolve_api_key)
|
||||
|
||||
def fallback(self) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
with self.request.litellm_logging_obj.suppress_next_pre_call() if self.logged else nullcontext():
|
||||
return self.python()
|
||||
return self.python()
|
||||
|
||||
async def afallback(self) -> OCRResponse:
|
||||
with self.request.litellm_logging_obj.suppress_next_pre_call() if self.logged else nullcontext():
|
||||
result: Final = self.python()
|
||||
return await result if isinstance(result, Coroutine) else result
|
||||
result: Final = self.python()
|
||||
return await result if isinstance(result, Coroutine) else result
|
||||
|
||||
|
||||
def _run_rust_ocr(
|
||||
|
|
@ -316,9 +279,6 @@ def _run_rust_ocr(
|
|||
adapt=OCRResponse.model_validate,
|
||||
model=prepared_request.model,
|
||||
provider=prepared_request.custom_llm_provider,
|
||||
request_format=(
|
||||
"native" if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -334,9 +294,6 @@ async def _run_rust_aocr(
|
|||
adapt=OCRResponse.model_validate,
|
||||
model=prepared_request.model,
|
||||
provider=prepared_request.custom_llm_provider,
|
||||
request_format=(
|
||||
"native" if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from .callbacks import CallbackDecision, CallbackUnchanged, SessionCallbackHandl
|
|||
|
||||
|
||||
class PreCallArguments(TypedDict):
|
||||
complete_input_dict: ReadOnly[Mapping[str, JsonValue]]
|
||||
complete_input_dict: Mapping[str, JsonValue] # writable-ok: provider hooks may replace request fields
|
||||
api_base: ReadOnly[str]
|
||||
headers: ReadOnly[Mapping[str, str]]
|
||||
|
||||
|
|
@ -90,12 +90,16 @@ class ProviderLoggingAdapter:
|
|||
|
||||
def pre_call(self, payload: object, /) -> CallbackDecision:
|
||||
event: Final = ProviderPreCall.model_validate(payload)
|
||||
request: Final = dict(event.request)
|
||||
additional_args: Final[PreCallArguments] = {
|
||||
"complete_input_dict": event.request,
|
||||
"complete_input_dict": request,
|
||||
"api_base": event.api_base,
|
||||
"headers": event.headers,
|
||||
}
|
||||
self.logging_obj.pre_call(input=self.input, api_key=self.api_key, additional_args=additional_args)
|
||||
mutated: Final = additional_args["complete_input_dict"]
|
||||
if dict(mutated) != dict(event.request):
|
||||
return {"action": "replace", "payload": dict(mutated)}
|
||||
return _unchanged()
|
||||
|
||||
def post_call(self, payload: object, /) -> CallbackDecision:
|
||||
|
|
|
|||
|
|
@ -33,7 +33,6 @@ 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 +51,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 +114,12 @@ _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 +133,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(
|
||||
|
|
@ -166,56 +152,11 @@ def _provider_eligibility_options(
|
|||
return NativeRequestOptions(custom_llm_provider=provider, bedrock=bedrock, anthropic=anthropic)
|
||||
|
||||
|
||||
def _eligibility_context(
|
||||
*,
|
||||
execution_mode: str | None = None,
|
||||
stream: bool,
|
||||
has_custom_client: bool = False,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> NativeRequestContext:
|
||||
return NativeRequestContext(
|
||||
capabilities=NativeRequestCapabilities(
|
||||
execution_mode=execution_mode,
|
||||
stream=stream,
|
||||
has_custom_client=has_custom_client,
|
||||
has_agentic_hook=has_agentic_hook,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
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 +326,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 +430,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 +437,4 @@ def dispatch_completion(
|
|||
fallback=operation.fallback,
|
||||
adapt=operation.adapt,
|
||||
error_context=error_context,
|
||||
preflight=operation.assess,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,193 +1 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from functools import wraps
|
||||
from typing import Final, ParamSpec, TypeAlias, TypeVar
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.exceptions import APIError, AuthenticationError, InternalServerError, RateLimitError
|
||||
from litellm.rust_bridge.bindings import native_declined_types, native_upstream_types
|
||||
from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
|
||||
NativeT = TypeVar("NativeT")
|
||||
PythonT = TypeVar("PythonT")
|
||||
P = ParamSpec("P")
|
||||
|
||||
|
||||
class ErrorAction(Enum):
|
||||
RAISE = "raise"
|
||||
SKIP = "skip"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class APIErrorMapping:
|
||||
provider: str
|
||||
model: str
|
||||
|
||||
|
||||
FailureAction: TypeAlias = ErrorAction | APIErrorMapping
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ErrorHandling:
|
||||
declined: FailureAction = ErrorAction.RAISE
|
||||
upstream: FailureAction = ErrorAction.RAISE
|
||||
unknown: FailureAction = ErrorAction.RAISE
|
||||
missing_metadata: FailureAction = ErrorAction.RAISE
|
||||
unexpected: FailureAction = ErrorAction.RAISE
|
||||
|
||||
|
||||
PROPAGATE: Final = ErrorHandling()
|
||||
|
||||
|
||||
def provider_errors(provider: str, model: str) -> ErrorHandling:
|
||||
return ErrorHandling(
|
||||
declined=ErrorAction.SKIP,
|
||||
upstream=APIErrorMapping(provider=provider, model=model),
|
||||
)
|
||||
|
||||
|
||||
def _handle_error(error: Exception, action: FailureAction, route: str, reason: NativeSkipReason) -> NativeSkipped:
|
||||
match action:
|
||||
case ErrorAction.SKIP:
|
||||
return NativeSkipped(reason, str(error))
|
||||
case ErrorAction.RAISE:
|
||||
raise error
|
||||
case APIErrorMapping(provider, model):
|
||||
args: Final[tuple[object, ...]] = error.args
|
||||
attribute_status: Final = getattr(error, "status_code", None)
|
||||
attribute_message: Final = getattr(error, "message", None)
|
||||
status_value: Final = attribute_status if isinstance(attribute_status, int) else (args[0] if args else 0)
|
||||
message_value: Final = (
|
||||
attribute_message if isinstance(attribute_message, str) else (args[1] if len(args) > 1 else str(error))
|
||||
)
|
||||
status: Final = status_value if isinstance(status_value, int) else 0
|
||||
message: Final = message_value if isinstance(message_value, str) else str(message_value)
|
||||
error_message: Final = f"litellm rust {route}: {message}"
|
||||
if status == 401:
|
||||
raise AuthenticationError(message=error_message, llm_provider=provider, model=model) from error
|
||||
if status == 429:
|
||||
raise RateLimitError(message=error_message, llm_provider=provider, model=model) from error
|
||||
if status == 500:
|
||||
raise InternalServerError(message=error_message, llm_provider=provider, model=model) from error
|
||||
raise APIError(
|
||||
status_code=status or 500,
|
||||
message=error_message,
|
||||
llm_provider=provider,
|
||||
model=model,
|
||||
) from error
|
||||
|
||||
|
||||
def _resolve(result: DispatchResult[NativeT], errors: ErrorHandling, route: str) -> Handled[NativeT] | NativeSkipped:
|
||||
if not isinstance(result, NativeFailed):
|
||||
return result
|
||||
declined: Final = native_declined_types()
|
||||
upstream: Final = native_upstream_types()
|
||||
if not declined or not upstream:
|
||||
return _handle_error(result.error, errors.missing_metadata, route, NativeSkipReason.FAILED)
|
||||
if isinstance(result.error, declined):
|
||||
return _handle_error(result.error, errors.declined, route, NativeSkipReason.DECLINED)
|
||||
if isinstance(result.error, upstream):
|
||||
return _handle_error(result.error, errors.upstream, route, NativeSkipReason.FAILED)
|
||||
return _handle_error(result.error, errors.unknown, route, NativeSkipReason.FAILED)
|
||||
|
||||
|
||||
def _log_skip(route: str, skipped: NativeSkipped) -> None:
|
||||
verbose_logger.debug("Native %s skipped (%s): %s", route, skipped.reason.value, skipped.detail or "")
|
||||
|
||||
|
||||
def native_first(
|
||||
*,
|
||||
native: Callable[P, DispatchResult[NativeT]],
|
||||
route: str,
|
||||
errors: Callable[P, ErrorHandling],
|
||||
) -> Callable[[Callable[P, PythonT]], Callable[P, NativeT | PythonT]]:
|
||||
def wrap(implementation: Callable[P, PythonT]) -> Callable[P, NativeT | PythonT]:
|
||||
@wraps(implementation)
|
||||
def run(
|
||||
*args: P.args,
|
||||
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
|
||||
) -> NativeT | PythonT:
|
||||
rules: Final = errors(*args, **kwargs)
|
||||
try:
|
||||
attempted: Final = native(*args, **kwargs)
|
||||
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
|
||||
skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED)
|
||||
_log_skip(route, skipped)
|
||||
else:
|
||||
result: Final = _resolve(attempted, rules, route)
|
||||
if isinstance(result, Handled):
|
||||
return result.value
|
||||
_log_skip(route, result)
|
||||
return implementation(*args, **kwargs)
|
||||
|
||||
return run
|
||||
|
||||
return wrap
|
||||
|
||||
|
||||
def anative_first(
|
||||
*,
|
||||
native: Callable[P, Awaitable[DispatchResult[NativeT]]],
|
||||
route: str,
|
||||
errors: Callable[P, ErrorHandling],
|
||||
) -> Callable[[Callable[P, Awaitable[PythonT]]], Callable[P, Awaitable[NativeT | PythonT]]]:
|
||||
def wrap(implementation: Callable[P, Awaitable[PythonT]]) -> Callable[P, Awaitable[NativeT | PythonT]]:
|
||||
@wraps(implementation)
|
||||
async def run(
|
||||
*args: P.args,
|
||||
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
|
||||
) -> NativeT | PythonT:
|
||||
rules: Final = errors(*args, **kwargs)
|
||||
try:
|
||||
attempted: Final = await native(*args, **kwargs)
|
||||
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
|
||||
skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED)
|
||||
_log_skip(route, skipped)
|
||||
else:
|
||||
result: Final = _resolve(attempted, rules, route)
|
||||
if isinstance(result, Handled):
|
||||
return result.value
|
||||
_log_skip(route, result)
|
||||
return await implementation(*args, **kwargs)
|
||||
|
||||
return run
|
||||
|
||||
return wrap
|
||||
|
||||
|
||||
def anative_context(
|
||||
*,
|
||||
native: Callable[P, Awaitable[DispatchResult[AbstractAsyncContextManager[NativeT]]]],
|
||||
route: str,
|
||||
errors: Callable[P, ErrorHandling],
|
||||
) -> Callable[
|
||||
[Callable[P, AbstractAsyncContextManager[PythonT]]],
|
||||
Callable[P, AbstractAsyncContextManager[NativeT | PythonT]],
|
||||
]:
|
||||
def wrap(
|
||||
implementation: Callable[P, AbstractAsyncContextManager[PythonT]],
|
||||
) -> Callable[P, AbstractAsyncContextManager[NativeT | PythonT]]:
|
||||
@anative_first(native=native, route=route, errors=errors)
|
||||
async def acquire(
|
||||
*args: P.args,
|
||||
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
|
||||
) -> AbstractAsyncContextManager[PythonT]:
|
||||
return implementation(*args, **kwargs)
|
||||
|
||||
@wraps(implementation)
|
||||
@asynccontextmanager
|
||||
async def run(
|
||||
*args: P.args,
|
||||
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
|
||||
) -> AsyncGenerator[NativeT | PythonT, None]:
|
||||
manager: Final = await acquire(*args, **kwargs)
|
||||
async with manager as connection:
|
||||
yield connection
|
||||
|
||||
return run
|
||||
|
||||
return wrap
|
||||
"""Compatibility module retained after native dispatch moved into route harnesses."""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
@ -63,8 +54,6 @@ def dispatch_ocr(
|
|||
adapt: Callable[[Mapping[str, object]], ResultT],
|
||||
model: str,
|
||||
provider: str,
|
||||
eligible: bool = True,
|
||||
request_format: str | None = None,
|
||||
) -> ResultT:
|
||||
return _OCR.invoke(
|
||||
prepare=prepare,
|
||||
|
|
@ -72,8 +61,6 @@ def dispatch_ocr(
|
|||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=provider, model=model),
|
||||
eligible=eligible,
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, provider, request_format=request_format),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -84,8 +71,6 @@ async def adispatch_ocr(
|
|||
adapt: Callable[[Mapping[str, object]], ResultT],
|
||||
model: str,
|
||||
provider: str,
|
||||
eligible: bool = True,
|
||||
request_format: str | None = None,
|
||||
) -> ResultT:
|
||||
return await _OCR.ainvoke(
|
||||
prepare=prepare,
|
||||
|
|
@ -93,6 +78,4 @@ async def adispatch_ocr(
|
|||
fallback=fallback,
|
||||
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: ...
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from __future__ import annotations
|
|||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
from types import MappingProxyType
|
||||
from typing import Generic, Protocol, TypeVar
|
||||
from typing import Generic, Protocol
|
||||
|
||||
from .callbacks import OneShotCallbackHandle
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from __future__ import annotations
|
|||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Final, Generic, NoReturn, Protocol, TypeAlias, TypeVar, assert_never
|
||||
from typing import Final, Generic, NoReturn, Protocol, TypeAlias, TypeVar
|
||||
|
||||
from litellm.exceptions import APIError, AuthenticationError, InternalServerError, RateLimitError
|
||||
from litellm.rust_bridge.bindings import (
|
||||
|
|
@ -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")
|
||||
|
|
@ -33,14 +32,6 @@ class PythonFallbackReason(Enum):
|
|||
NATIVE_DECLINED = "native_declined"
|
||||
|
||||
|
||||
class NativeSkipReason(Enum):
|
||||
DISABLED = "disabled"
|
||||
INELIGIBLE = "ineligible"
|
||||
UNAVAILABLE = "unavailable"
|
||||
DECLINED = "declined"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Handled(Generic[ResultT]):
|
||||
value: ResultT
|
||||
|
|
@ -52,18 +43,7 @@ class PythonFallback:
|
|||
detail: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeSkipped:
|
||||
reason: NativeSkipReason
|
||||
detail: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeFailed:
|
||||
error: Exception
|
||||
|
||||
|
||||
DispatchResult: TypeAlias = Handled[ResultT] | PythonFallback | NativeSkipped | NativeFailed
|
||||
DispatchResult: TypeAlias = Handled[ResultT] | PythonFallback
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -118,17 +98,10 @@ class EndpointBinding(Generic[BindingT]):
|
|||
call: Callable[[BindingT, RequestT], NativeT],
|
||||
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,
|
||||
)
|
||||
binding_or_fallback: Final = self._binding_or_python_fallback()
|
||||
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,
|
||||
|
|
@ -142,17 +115,10 @@ class EndpointBinding(Generic[BindingT]):
|
|||
call: Callable[[BindingT, RequestT], Awaitable[NativeT]],
|
||||
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,
|
||||
)
|
||||
binding_or_fallback: Final = self._binding_or_python_fallback()
|
||||
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,
|
||||
|
|
@ -167,24 +133,18 @@ class EndpointBinding(Generic[BindingT]):
|
|||
fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
result: Final = self._attempt(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
match result:
|
||||
case Handled(value=value):
|
||||
return value
|
||||
case PythonFallback():
|
||||
return fallback()
|
||||
case _ as unreachable:
|
||||
assert_never(unreachable)
|
||||
|
||||
async def ainvoke(
|
||||
self,
|
||||
|
|
@ -194,52 +154,18 @@ class EndpointBinding(Generic[BindingT]):
|
|||
fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
result: Final = await self._aattempt(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
match result:
|
||||
case Handled(value=value):
|
||||
return value
|
||||
case PythonFallback():
|
||||
return await fallback()
|
||||
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,
|
||||
|
|
@ -248,24 +174,18 @@ class EndpointBinding(Generic[BindingT]):
|
|||
call: Callable[[BindingT, RequestT], NativeT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
result: Final = self._attempt(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
match result:
|
||||
case Handled(value=value):
|
||||
return value
|
||||
case PythonFallback():
|
||||
self._raise_required(result)
|
||||
case _ as unreachable:
|
||||
assert_never(unreachable)
|
||||
|
||||
async def arequire(
|
||||
self,
|
||||
|
|
@ -274,46 +194,26 @@ class EndpointBinding(Generic[BindingT]):
|
|||
call: Callable[[BindingT, RequestT], Awaitable[NativeT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
result: Final = await self._aattempt(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
match result:
|
||||
case Handled(value=value):
|
||||
return value
|
||||
case PythonFallback():
|
||||
self._raise_required(result)
|
||||
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)
|
||||
raise RuntimeError(f"native {self.route} endpoint {reason}{detail}")
|
||||
|
||||
def _binding_or_python_fallback(
|
||||
self,
|
||||
*,
|
||||
eligible: bool,
|
||||
) -> BindingT | PythonFallback:
|
||||
if not eligible or not self.enabled():
|
||||
def _binding_or_python_fallback(self) -> BindingT | PythonFallback:
|
||||
if not self.enabled():
|
||||
return PythonFallback(PythonFallbackReason.NATIVE_DISABLED)
|
||||
binding: Final = self.load()
|
||||
if binding is None:
|
||||
|
|
@ -440,8 +340,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
|
|||
fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
return self.sync.invoke(
|
||||
prepare=prepare,
|
||||
|
|
@ -449,8 +347,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
|
|||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
async def ainvoke(
|
||||
|
|
@ -461,8 +357,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
|
|||
fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
return await self.asynchronous.ainvoke(
|
||||
prepare=prepare,
|
||||
|
|
@ -470,8 +364,6 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
|
|||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
def require(
|
||||
|
|
@ -481,16 +373,12 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
|
|||
call: Callable[[SyncBindingT, RequestT], NativeT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
return self.sync.require(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
async def arequire(
|
||||
|
|
@ -500,16 +388,12 @@ class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
|
|||
call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
return await self.asynchronous.arequire(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -544,30 +428,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,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,23 +3,19 @@ from __future__ import annotations
|
|||
import base64
|
||||
import json
|
||||
from collections.abc import Callable, Coroutine
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from io import IOBase
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from typing import Final, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
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.callback_adapters import ProviderLoggingAdapter
|
||||
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
|
||||
from litellm.rust_bridge.request import (
|
||||
NativePreCallDetails,
|
||||
NativeRequestCapabilities,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
|
|
@ -32,16 +28,12 @@ from litellm.rust_bridge.request import (
|
|||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
EndpointDispatch,
|
||||
PythonFallback,
|
||||
always_enabled,
|
||||
assess_route,
|
||||
async_none,
|
||||
identity,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native(
|
||||
|
|
@ -52,24 +44,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 +110,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 +157,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),
|
||||
|
|
@ -200,6 +177,18 @@ def _input_source_kind(file: FileTypes) -> str:
|
|||
return "opaque"
|
||||
|
||||
|
||||
def _consume_audio_for_native(file: object) -> dict[str, object]:
|
||||
"""Read audio only after the native route has admitted the request."""
|
||||
if isinstance(file, dict):
|
||||
return TypeAdapter(dict[str, object]).validate_python(file)
|
||||
processed: Final = process_audio_file(cast(FileTypes, file))
|
||||
return {
|
||||
"data": base64.b64encode(processed.file_content).decode("ascii"),
|
||||
"format": processed.filename.rsplit(".", 1)[-1].lower() if "." in processed.filename else "wav",
|
||||
"filename": processed.filename,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class _TranscriptionOperation:
|
||||
model: str
|
||||
|
|
@ -214,55 +203,17 @@ class _TranscriptionOperation:
|
|||
python: Callable[[FileTypes], TranscriptionResult]
|
||||
asynchronous: bool = False
|
||||
has_custom_client: bool = False
|
||||
fallback_file: FileTypes | None = None
|
||||
logged: bool = False
|
||||
|
||||
def prepare(self) -> PreparedNativeCall[NativeTranscriptionRequest]:
|
||||
key: Final = (
|
||||
self.api_key
|
||||
or litellm.api_key
|
||||
or TypeAdapter(str | None).validate_python(getattr(litellm, f"{self.provider}_key", None))
|
||||
or get_secret_str(f"{self.provider.upper()}_API_KEY")
|
||||
)
|
||||
base: Final = (
|
||||
self.api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str(f"{self.provider.upper()}_BASE_URL")
|
||||
or get_secret_str(f"{self.provider.upper()}_API_BASE")
|
||||
)
|
||||
content: Final = self.file[1] if isinstance(self.file, tuple) else self.file
|
||||
position: Final = content.tell() if isinstance(content, IOBase) and content.seekable() else None
|
||||
try:
|
||||
processed: Final = process_audio_file(self.file)
|
||||
finally:
|
||||
if position is not None and isinstance(content, IOBase):
|
||||
content.seek(position)
|
||||
self.fallback_file = (processed.filename, processed.file_content, processed.content_type)
|
||||
audio: Final = TypeAdapter(dict[str, object]).validate_python(
|
||||
MappingProxyType(
|
||||
{
|
||||
"data": base64.b64encode(processed.file_content).decode("ascii"),
|
||||
"format": processed.filename.rsplit(".", 1)[-1].lower() if "." in processed.filename else "wav",
|
||||
"filename": processed.filename,
|
||||
}
|
||||
)
|
||||
)
|
||||
log_details: Final[NativePreCallDetails] = {
|
||||
"api_base": base or "",
|
||||
"headers": self.headers,
|
||||
"complete_input_dict": {"model": self.model, **self.optional_params},
|
||||
}
|
||||
self.logging.pre_call(input="audio transcription", api_key=key, additional_args=log_details)
|
||||
self.logged = True
|
||||
return PreparedNativeCall(
|
||||
NativeTranscriptionRequest(
|
||||
model=self.model,
|
||||
audio=audio,
|
||||
audio=self.file,
|
||||
optional_params=self.optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=key,
|
||||
api_base=base,
|
||||
api_key=self.api_key,
|
||||
api_base=self.api_base,
|
||||
custom_llm_provider=self.provider,
|
||||
extra_headers=self.headers,
|
||||
timeout_seconds=timeout_to_seconds(self.timeout),
|
||||
|
|
@ -279,16 +230,15 @@ class _TranscriptionOperation:
|
|||
input_source_kind=_input_source_kind(self.file),
|
||||
),
|
||||
),
|
||||
callback_adapter=ProviderLoggingAdapter(self.logging, "audio transcription", self.api_key),
|
||||
)
|
||||
|
||||
def fallback(self) -> TranscriptionResult:
|
||||
with self.logging.suppress_next_pre_call() if self.logged else nullcontext():
|
||||
return self.python(self.fallback_file if self.fallback_file is not None else self.file)
|
||||
return self.python(self.file)
|
||||
|
||||
async def afallback(self) -> TranscriptionResponse:
|
||||
with self.logging.suppress_next_pre_call() if self.logged else nullcontext():
|
||||
result: Final = self.python(self.fallback_file if self.fallback_file is not None else self.file)
|
||||
return await result if isinstance(result, Coroutine) else result
|
||||
result: Final = self.python(self.file)
|
||||
return await result if isinstance(result, Coroutine) else result
|
||||
|
||||
def adapt(self, response: dict[str, object]) -> TranscriptionResponse:
|
||||
text: Final = TypeAdapter(str).validate_python(response["text"])
|
||||
|
|
@ -329,15 +279,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 +287,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(
|
||||
|
|
@ -362,8 +301,6 @@ def dispatch_transcription(
|
|||
adapt=operation.adapt,
|
||||
fallback=operation.afallback,
|
||||
error_context=error_context,
|
||||
eligible=rust_enabled(),
|
||||
preflight=preflight,
|
||||
)
|
||||
return _TRANSCRIPTION.invoke(
|
||||
prepare=operation.prepare,
|
||||
|
|
@ -371,6 +308,4 @@ def dispatch_transcription(
|
|||
adapt=operation.adapt,
|
||||
fallback=operation.fallback,
|
||||
error_context=error_context,
|
||||
eligible=rust_enabled(),
|
||||
preflight=preflight,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ class RecordingMessages:
|
|||
*,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -65,6 +66,7 @@ class RecordingAsyncMessages:
|
|||
*,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -114,19 +116,11 @@ class RaisingAsyncMessages:
|
|||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_rust_flag():
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None, decline=None)
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None)
|
||||
configuration.reset_rust_configuration()
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
rust_messages.set_rust_messages(
|
||||
decline=lambda model, custom_llm_provider, *, context: (
|
||||
"unsupported feature"
|
||||
if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client"))
|
||||
or context.capabilities.request_format == "native"
|
||||
else None
|
||||
)
|
||||
)
|
||||
yield
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None, decline=None)
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None)
|
||||
configuration.reset_rust_configuration()
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
|
||||
|
|
@ -285,7 +279,7 @@ def test_public_messages_strips_provider_specific_fields_before_native_dispatch(
|
|||
assert "provider_specific_fields" in messages[0]["content"][0]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("condition", ["disabled", "declined", "missing_binding", "missing_preflight", "stream"])
|
||||
@pytest.mark.parametrize("condition", ["disabled", "missing_binding"])
|
||||
def test_public_messages_fallback_once(monkeypatch, condition):
|
||||
module = importlib.import_module("litellm.llms.anthropic.experimental_pass_through.messages.handler")
|
||||
python = PythonMessages()
|
||||
|
|
@ -293,18 +287,13 @@ def test_public_messages_fallback_once(monkeypatch, condition):
|
|||
bridge = RecordingMessages()
|
||||
litellm.rust(condition != "disabled")
|
||||
rust_messages.set_rust_messages(messages=bridge)
|
||||
if condition == "declined":
|
||||
rust_messages.set_rust_messages(decline=lambda model, custom_llm_provider, **features: "unsupported provider")
|
||||
elif condition == "missing_binding":
|
||||
if condition == "missing_binding":
|
||||
rust_messages._MESSAGES.sync.override(None)
|
||||
elif condition == "missing_preflight":
|
||||
rust_messages._PREFLIGHT.override(None)
|
||||
litellm.anthropic.messages.create(
|
||||
model="anthropic/test-model",
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="key",
|
||||
stream=condition == "stream",
|
||||
)
|
||||
assert python.calls == 1
|
||||
assert bridge.calls == []
|
||||
|
|
@ -316,7 +305,9 @@ def test_public_messages_invalid_response_does_not_fallback(monkeypatch, respons
|
|||
python = PythonMessages()
|
||||
monkeypatch.setattr(module, "base_llm_http_handler", python)
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(messages=lambda request, *, options, context: response)
|
||||
rust_messages.set_rust_messages(
|
||||
messages=lambda request, *, options, context, callback_adapter=None: response
|
||||
)
|
||||
with pytest.raises(ValidationError):
|
||||
litellm.anthropic.messages.create(
|
||||
model="anthropic/test-model",
|
||||
|
|
|
|||
|
|
@ -4,15 +4,18 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
from litellm.rust_bridge.request import NativeChatCompletionsRequest, NativeRequestContext, NativeRequestOptions
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeChatCompletionsRequest,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def native_bridge(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
bridge.set_rust_chat_completions(decline=lambda **features: None)
|
||||
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)
|
||||
|
||||
|
||||
def test_native_bedrock_receives_explicit_auth_and_endpoint():
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import litellm
|
|||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge.callback_adapters import ProviderLoggingAdapter
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.callback_adapters import ProviderLoggingAdapter
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeOCRRequest,
|
||||
NativeRequestContext,
|
||||
|
|
@ -224,24 +225,11 @@ def _reset_rust_flag():
|
|||
"""Keep the global toggle isolated between tests."""
|
||||
rust_bridge._OCR.sync.reset()
|
||||
rust_bridge._OCR.asynchronous.reset()
|
||||
rust_bridge._PREFLIGHT.reset()
|
||||
configuration.reset_rust_configuration()
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
rust_bridge._PREFLIGHT.override(
|
||||
lambda model, custom_llm_provider, *, context: (
|
||||
"unsupported feature"
|
||||
if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client"))
|
||||
or (
|
||||
context.capabilities.request_format == "native"
|
||||
and not (custom_llm_provider == "azure_ai" and "doc-intelligence" in model)
|
||||
)
|
||||
else None
|
||||
)
|
||||
)
|
||||
yield
|
||||
rust_bridge._OCR.sync.reset()
|
||||
rust_bridge._OCR.asynchronous.reset()
|
||||
rust_bridge._PREFLIGHT.reset()
|
||||
configuration.reset_rust_configuration()
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
|
||||
|
|
@ -378,7 +366,6 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
|
|||
|
||||
rust_bridge._OCR.sync.reset()
|
||||
rust_bridge._OCR.asynchronous.reset()
|
||||
rust_bridge._PREFLIGHT.reset()
|
||||
assert rust_bridge.load_rust_ocr() is None
|
||||
assert rust_bridge.load_rust_aocr() is None
|
||||
|
||||
|
|
@ -446,7 +433,6 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
|
|||
adapt=dict,
|
||||
model="mistral-ocr-latest",
|
||||
provider="mistral",
|
||||
eligible=True,
|
||||
)
|
||||
|
||||
assert response == FAKE_OCR_RESPONSE
|
||||
|
|
@ -495,7 +481,6 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
|
|||
adapt=dict,
|
||||
model="mistral-ocr-maas",
|
||||
provider="vertex_ai",
|
||||
eligible=True,
|
||||
)
|
||||
|
||||
assert response == FAKE_OCR_RESPONSE
|
||||
|
|
@ -538,17 +523,14 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
|
|||
"api_key": "sk-test",
|
||||
"api_base": "https://proxy.internal",
|
||||
"custom_llm_provider": "mistral",
|
||||
"extra_headers": {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
},
|
||||
"extra_headers": {"x-trace-id": "trace-1"},
|
||||
"optional_params": {"include_image_base64": True},
|
||||
"vertex": NativeVertexOptions(),
|
||||
"timeout_seconds": 12.5,
|
||||
}
|
||||
|
||||
|
||||
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
||||
def test_run_rust_ocr_does_not_resolve_credentials_before_admission():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
|
@ -556,10 +538,10 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
|||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(api_key=None, timeout=None),
|
||||
resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None,
|
||||
resolve_api_key=lambda name: pytest.fail(f"unexpected pre-admission lookup: {name}"),
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["api_key"] == "sk-from-vault"
|
||||
assert bridge.calls[0]["api_key"] is None
|
||||
|
||||
|
||||
def test_run_rust_ocr_prefers_explicit_key_over_resolver():
|
||||
|
|
@ -582,7 +564,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver():
|
|||
assert bridge.calls[0]["api_key"] == "sk-explicit"
|
||||
|
||||
|
||||
def test_run_rust_ocr_uses_provider_api_key_env_var():
|
||||
def test_run_rust_ocr_leaves_provider_discovery_to_native_admission():
|
||||
bridge = RecordingBridge()
|
||||
resolver_calls = []
|
||||
litellm.rust(True)
|
||||
|
|
@ -603,8 +585,8 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
|
|||
resolve_api_key=_resolver,
|
||||
)
|
||||
|
||||
assert resolver_calls == ["PROVIDER_OCR_API_KEY"]
|
||||
assert bridge.calls[0]["api_key"] == "sk-provider-env"
|
||||
assert resolver_calls == []
|
||||
assert bridge.calls[0]["api_key"] is None
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
||||
|
|
@ -632,7 +614,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
|||
assert bridge.calls[0]["vertex"] == NativeVertexOptions(project="project-1", location="us-central1")
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager():
|
||||
def test_prepare_rust_ocr_call_does_not_resolve_vertex_metadata_before_admission():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
|
@ -653,10 +635,10 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana
|
|||
resolve_api_key=_resolver,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["vertex"] == NativeVertexOptions(project="project-from-secret", location="us-east5")
|
||||
assert bridge.calls[0]["vertex"] == NativeVertexOptions()
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
||||
def test_prepare_rust_ocr_call_leaves_azure_base_discovery_to_native():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
|
@ -672,10 +654,10 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
|||
resolve_api_key=lambda name: "https://azure.example.com" if name == "AZURE_AI_API_BASE" else None,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["api_base"] == "https://azure.example.com"
|
||||
assert bridge.calls[0]["api_base"] is None
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
||||
def test_prepare_rust_ocr_call_leaves_document_intelligence_endpoint_to_native():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
|
@ -693,7 +675,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
|||
),
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["api_base"] == "https://document-intelligence.example.com"
|
||||
assert bridge.calls[0]["api_base"] is None
|
||||
|
||||
|
||||
def test_run_rust_ocr_passes_provider_logging_adapter():
|
||||
|
|
@ -735,7 +717,7 @@ def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge):
|
|||
assert fake_bridge.calls[0]["custom_llm_provider"] == "azure_ai"
|
||||
|
||||
|
||||
def test_ocr_rust_path_converts_file_document_before_bridge(fake_bridge):
|
||||
def test_ocr_rust_path_keeps_file_document_opaque_until_native_admission(fake_bridge):
|
||||
response = litellm.ocr(
|
||||
model=MODEL,
|
||||
document={"type": "file", "file": b"%PDF-1.4", "mime_type": "application/pdf"},
|
||||
|
|
@ -744,8 +726,7 @@ def test_ocr_rust_path_converts_file_document_before_bridge(fake_bridge):
|
|||
|
||||
assert isinstance(response, OCRResponse)
|
||||
document = fake_bridge.calls[0]["document"]
|
||||
assert document["type"] == "document_url"
|
||||
assert document["document_url"].startswith("data:application/pdf;base64,")
|
||||
assert document == {"type": "file", "file": b"%PDF-1.4", "mime_type": "application/pdf"}
|
||||
|
||||
|
||||
def test_ocr_exception_type_uses_resolved_provider_context(
|
||||
|
|
@ -786,10 +767,7 @@ async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge):
|
|||
assert call["document"] == DOCUMENT
|
||||
assert call["api_key"] == "sk-test"
|
||||
assert call["custom_llm_provider"] == "mistral"
|
||||
assert call["extra_headers"] == {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
}
|
||||
assert call["extra_headers"] == {"x-trace-id": "trace-1"}
|
||||
assert call["optional_params"].get("include_image_base64") is True
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -44,18 +44,10 @@ class _FakeNativeBridge:
|
|||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_responses_websocket():
|
||||
responses_websocket.set_rust_responses_websocket(connection=None, decline=None)
|
||||
responses_websocket.set_rust_responses_websocket(connection=None)
|
||||
configuration.reset_rust_configuration()
|
||||
responses_websocket.set_rust_responses_websocket(
|
||||
decline=lambda model, custom_llm_provider, *, context: (
|
||||
"unsupported feature"
|
||||
if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client"))
|
||||
or context.capabilities.request_format == "native"
|
||||
else None
|
||||
)
|
||||
)
|
||||
yield
|
||||
responses_websocket.set_rust_responses_websocket(connection=None, decline=None)
|
||||
responses_websocket.set_rust_responses_websocket(connection=None)
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
|
|
@ -181,11 +173,7 @@ async def test_connection_dispatch_cleans_up_without_reconnecting(native, sessio
|
|||
finally:
|
||||
await python_socket.close()
|
||||
|
||||
responses_websocket.set_rust_responses_websocket(connection=Native)
|
||||
if not native:
|
||||
responses_websocket.set_rust_responses_websocket(
|
||||
decline=lambda model, custom_llm_provider, **features: "declined"
|
||||
)
|
||||
responses_websocket.set_rust_responses_websocket(connection=Native if native else None)
|
||||
|
||||
async def run():
|
||||
async with responses_websocket.open_connection(
|
||||
|
|
@ -210,7 +198,7 @@ async def test_connection_dispatch_cleans_up_without_reconnecting(native, sessio
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_acceptance_export_uses_python_connection_once():
|
||||
async def test_missing_acceptance_export_keeps_native_failure_terminal():
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
calls = []
|
||||
|
|
@ -226,10 +214,10 @@ async def test_missing_acceptance_export_uses_python_connection_once():
|
|||
|
||||
configuration.rust(True)
|
||||
responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge)
|
||||
responses_websocket._PREFLIGHT.override(None)
|
||||
async with responses_websocket.open_connection(
|
||||
url="wss://example.test", headers={}, timeout=1, model="model", provider="openai", fallback=python
|
||||
) as connection:
|
||||
assert connection is socket
|
||||
assert calls == ["python"]
|
||||
assert socket.closed
|
||||
with pytest.raises(RuntimeError, match="connection failed"):
|
||||
async with responses_websocket.open_connection(
|
||||
url="wss://example.test", headers={}, timeout=1, model="model", provider="openai", fallback=python
|
||||
):
|
||||
pass
|
||||
assert calls == []
|
||||
assert not socket.closed
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
@ -488,14 +351,14 @@ def test_typed_capability_and_provider_metadata_facts_are_isolated():
|
|||
assert anthropic_options({"metadata": {"user_id": None}}).has_user_id is False
|
||||
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["anthropic", "bedrock", "openai"])
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
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,14 +372,11 @@ 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"]
|
||||
)
|
||||
@pytest.mark.parametrize("failure", ["decline", "unavailable", "error", "malformed", "cancelled"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_completion_fallback_contract(monkeypatch, asynchronous, failure):
|
||||
import asyncio
|
||||
|
|
@ -564,10 +424,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)
|
||||
|
|
|
|||
|
|
@ -1,337 +0,0 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.exceptions import APIError, AuthenticationError, InternalServerError, RateLimitError
|
||||
from litellm.rust_bridge import bindings
|
||||
from litellm.rust_bridge.dispatch import PROPAGATE, anative_first, native_first, provider_errors
|
||||
from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
|
||||
|
||||
class Declined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class Upstream(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def native_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(
|
||||
bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("reason", tuple(NativeSkipReason))
|
||||
async def test_shared_dispatch_calls_python_once_and_logs_skip(
|
||||
asynchronous: bool, reason: NativeSkipReason, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
caplog.set_level(logging.DEBUG, logger="LiteLLM")
|
||||
calls: Final[list[str]] = []
|
||||
|
||||
def native() -> DispatchResult[str]:
|
||||
calls.append("native")
|
||||
return NativeSkipped(reason, "diagnostic detail")
|
||||
|
||||
async def anative() -> DispatchResult[str]:
|
||||
return native()
|
||||
|
||||
def python() -> str:
|
||||
calls.append("python")
|
||||
return "python response"
|
||||
|
||||
async def apython() -> str:
|
||||
return python()
|
||||
|
||||
result: Final = (
|
||||
await anative_first(native=anative, route="test", errors=lambda: PROPAGATE)(apython)()
|
||||
if asynchronous
|
||||
else native_first(native=native, route="test", errors=lambda: PROPAGATE)(python)()
|
||||
)
|
||||
assert result == "python response"
|
||||
assert calls == ["native", "python"]
|
||||
assert f"Native test skipped ({reason.value}): diagnostic detail" in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_native_success_does_not_run_python_even_when_value_is_none(asynchronous: bool) -> None:
|
||||
|
||||
async def native() -> DispatchResult[None]:
|
||||
return Handled(None)
|
||||
|
||||
def python() -> str:
|
||||
pytest.fail("handled results must not run Python")
|
||||
|
||||
async def apython() -> str:
|
||||
return python()
|
||||
|
||||
result: Final = (
|
||||
await anative_first(native=native, route="test", errors=lambda: PROPAGATE)(apython)()
|
||||
if asynchronous
|
||||
else native_first(native=lambda: Handled(None), route="test", errors=lambda: PROPAGATE)(python)()
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("policy", ("chat", "propagate"))
|
||||
@pytest.mark.parametrize("kind", ("declined", "upstream", "unknown", "unexpected", "missing"))
|
||||
async def test_declarations_control_endpoint_error_behavior(
|
||||
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, policy: str, kind: str
|
||||
) -> None:
|
||||
if kind == "missing":
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: None)
|
||||
error: Final = (
|
||||
Declined("unsupported")
|
||||
if kind == "declined"
|
||||
else Upstream(429, "rate limited")
|
||||
if kind == "upstream"
|
||||
else RuntimeError("failed")
|
||||
)
|
||||
rules: Final = (
|
||||
provider_errors("anthropic", "model")
|
||||
if policy == "chat"
|
||||
else PROPAGATE
|
||||
)
|
||||
calls: Final[list[str]] = []
|
||||
|
||||
def native() -> DispatchResult[str]:
|
||||
if kind == "unexpected":
|
||||
raise error
|
||||
return NativeFailed(error)
|
||||
|
||||
async def anative() -> DispatchResult[str]:
|
||||
return native()
|
||||
|
||||
def python() -> str:
|
||||
calls.append("python")
|
||||
return "python response"
|
||||
|
||||
async def apython() -> str:
|
||||
return python()
|
||||
|
||||
async def run() -> str:
|
||||
if asynchronous:
|
||||
return await anative_first(native=anative, route="chat_completions", errors=lambda: rules)(apython)()
|
||||
return native_first(native=native, route="chat_completions", errors=lambda: rules)(python)()
|
||||
|
||||
if policy == "chat" and kind == "declined":
|
||||
assert await run() == "python response"
|
||||
assert calls == ["python"]
|
||||
elif policy == "chat" and kind == "upstream":
|
||||
with pytest.raises(RateLimitError) as caught:
|
||||
await run()
|
||||
assert caught.value.status_code == 429
|
||||
assert caught.value.model == "model"
|
||||
assert caught.value.llm_provider == "anthropic"
|
||||
assert caught.value.__cause__ is error
|
||||
assert calls == []
|
||||
else:
|
||||
with pytest.raises(type(error)) as caught_original:
|
||||
await run()
|
||||
assert caught_original.value is error
|
||||
assert calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_python_failure_is_never_reclassified_as_native_failure(asynchronous: bool) -> None:
|
||||
error: Final = RuntimeError("Python failed")
|
||||
calls: Final[list[str]] = []
|
||||
|
||||
async def native() -> DispatchResult[str]:
|
||||
return NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
|
||||
def python() -> str:
|
||||
calls.append("python")
|
||||
raise error
|
||||
|
||||
async def apython() -> str:
|
||||
return python()
|
||||
|
||||
async def run() -> str:
|
||||
if asynchronous:
|
||||
return await anative_first(native=native, route="test", errors=lambda: PROPAGATE)(apython)()
|
||||
return native_first(
|
||||
native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE), route="test", errors=lambda: PROPAGATE
|
||||
)(python)()
|
||||
|
||||
with pytest.raises(RuntimeError) as caught:
|
||||
await run()
|
||||
assert caught.value is error
|
||||
assert calls == ["python"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancellation_does_not_run_python() -> None:
|
||||
|
||||
async def native() -> DispatchResult[str]:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
async def python() -> str:
|
||||
pytest.fail("cancellation must not dispatch Python")
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await anative_first(native=native, route="test", errors=lambda: PROPAGATE)(python)()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize(
|
||||
"status,exception_type",
|
||||
(
|
||||
(0, APIError),
|
||||
(401, AuthenticationError),
|
||||
(403, APIError),
|
||||
(429, RateLimitError),
|
||||
(500, InternalServerError),
|
||||
(503, APIError),
|
||||
),
|
||||
)
|
||||
async def test_upstream_mapping_preserves_status_message_and_context(
|
||||
asynchronous: bool, status: int, exception_type: type[Exception]
|
||||
) -> None:
|
||||
error: Final = Upstream(status, "upstream failed")
|
||||
|
||||
async def native() -> DispatchResult[str]:
|
||||
return NativeFailed(error)
|
||||
|
||||
async def python() -> str:
|
||||
pytest.fail("upstream errors must not run Python")
|
||||
|
||||
async def run() -> str:
|
||||
if asynchronous:
|
||||
return await anative_first(
|
||||
native=native, route="chat_completions", errors=lambda: provider_errors("anthropic", "model")
|
||||
)(python)()
|
||||
return native_first(
|
||||
native=lambda: NativeFailed(error),
|
||||
route="chat_completions",
|
||||
errors=lambda: provider_errors("anthropic", "model"),
|
||||
)(lambda: pytest.fail("upstream errors must not run Python"))()
|
||||
|
||||
with pytest.raises(exception_type, match="upstream failed") as caught:
|
||||
await run()
|
||||
assert caught.value.status_code == (status or 500)
|
||||
assert caught.value.model == "model"
|
||||
assert caught.value.llm_provider == "anthropic"
|
||||
assert caught.value.__cause__ is error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_registered_wrapper_preserves_arguments_and_request_error_context(asynchronous: bool) -> None:
|
||||
calls: Final[list[tuple[str, str, str]]] = []
|
||||
|
||||
def native(provider: str, *, model: str) -> DispatchResult[str]:
|
||||
calls.append(("native", provider, model))
|
||||
return (
|
||||
NativeFailed(Upstream(429, "limited"))
|
||||
if model == "limited"
|
||||
else NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
)
|
||||
|
||||
async def anative(provider: str, *, model: str) -> DispatchResult[str]:
|
||||
return native(provider, model=model)
|
||||
|
||||
def rules(provider: str, *, model: str):
|
||||
return provider_errors(provider, model)
|
||||
|
||||
@native_first(native=native, route="chat_completions", errors=rules)
|
||||
def execute(provider: str, *, model: str) -> str:
|
||||
calls.append(("python", provider, model))
|
||||
return model
|
||||
|
||||
@anative_first(native=anative, route="chat_completions", errors=rules)
|
||||
async def aexecute(provider: str, *, model: str) -> str:
|
||||
calls.append(("python", provider, model))
|
||||
return model
|
||||
|
||||
assert (await aexecute("first", model="ok") if asynchronous else execute("first", model="ok")) == "ok"
|
||||
|
||||
async def fail() -> None:
|
||||
if asynchronous:
|
||||
await aexecute("second", model="limited")
|
||||
else:
|
||||
execute("second", model="limited")
|
||||
|
||||
with pytest.raises(RateLimitError) as caught:
|
||||
await fail()
|
||||
assert caught.value.llm_provider == "second"
|
||||
assert caught.value.model == "limited"
|
||||
assert calls == [("native", "first", "ok"), ("python", "first", "ok"), ("native", "second", "limited")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("selection", ("native", "unavailable", "failed"))
|
||||
@pytest.mark.parametrize("failure", ("none", "body", "cleanup", "cancel"))
|
||||
async def test_context_selection_and_lifetime_are_separate(selection: str, failure: str) -> None:
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
|
||||
from litellm.rust_bridge.dispatch import anative_context
|
||||
|
||||
events: Final[list[str]] = []
|
||||
error: Final = RuntimeError("connection use failed")
|
||||
|
||||
@asynccontextmanager
|
||||
async def connection(name: str) -> AsyncGenerator[str, None]:
|
||||
events.append(f"{name}:enter")
|
||||
try:
|
||||
yield name
|
||||
finally:
|
||||
events.append(f"{name}:exit")
|
||||
if failure == "cleanup":
|
||||
raise error
|
||||
|
||||
async def native() -> DispatchResult[AbstractAsyncContextManager[str]]:
|
||||
events.append("attempt")
|
||||
if selection == "failed":
|
||||
raise RuntimeError("connect failed")
|
||||
if selection == "unavailable":
|
||||
return NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
return Handled(connection("native"))
|
||||
|
||||
@anative_context(native=native, route="websocket", errors=lambda: PROPAGATE)
|
||||
def execute() -> AbstractAsyncContextManager[str]:
|
||||
events.append("python")
|
||||
return connection("python")
|
||||
|
||||
async def run() -> None:
|
||||
async with execute() as name:
|
||||
assert name == ("native" if selection == "native" else "python")
|
||||
if failure == "body":
|
||||
raise error
|
||||
if failure == "cancel":
|
||||
raise asyncio.CancelledError
|
||||
|
||||
if selection == "failed":
|
||||
with pytest.raises(RuntimeError, match="connect failed"):
|
||||
await run()
|
||||
elif failure == "none":
|
||||
await run()
|
||||
elif failure == "cancel":
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await run()
|
||||
else:
|
||||
with pytest.raises(RuntimeError) as caught:
|
||||
await run()
|
||||
assert caught.value is error
|
||||
expected: Final = (
|
||||
["attempt", "native:enter", "native:exit"]
|
||||
if selection == "native"
|
||||
else ["attempt", "python", "python:enter", "python:exit"]
|
||||
if selection == "unavailable"
|
||||
else ["attempt"]
|
||||
)
|
||||
assert events == expected
|
||||
|
|
@ -41,7 +41,6 @@ def enabled() -> bool:
|
|||
@dataclass(frozen=True, slots=True)
|
||||
class FallbackCase:
|
||||
process_enabled: bool | None = None
|
||||
eligible: bool = True
|
||||
binding_available: bool = True
|
||||
declined: bool = False
|
||||
expected_events: tuple[str, ...] = ()
|
||||
|
|
@ -52,10 +51,6 @@ FALLBACK_CASES: Final = (
|
|||
FallbackCase(process_enabled=False, expected_events=("python",)),
|
||||
id="process-disabled",
|
||||
),
|
||||
pytest.param(
|
||||
FallbackCase(eligible=False, expected_events=("python",)),
|
||||
id="request-ineligible",
|
||||
),
|
||||
pytest.param(
|
||||
FallbackCase(binding_available=False, expected_events=("load", "python")),
|
||||
id="bridge-unavailable",
|
||||
|
|
@ -90,7 +85,6 @@ def test_invoke_falls_back_only_before_provider_success(case: FallbackCase) -> N
|
|||
fallback=lambda: events.append("python") or "fallback",
|
||||
adapt=str,
|
||||
error_context=context(),
|
||||
eligible=case.eligible,
|
||||
)
|
||||
|
||||
assert result == "fallback"
|
||||
|
|
@ -125,7 +119,6 @@ async def test_ainvoke_matches_sync_fallback_contract(case: FallbackCase) -> Non
|
|||
fallback=fallback,
|
||||
adapt=str,
|
||||
error_context=context(),
|
||||
eligible=case.eligible,
|
||||
)
|
||||
|
||||
assert result == "fallback"
|
||||
|
|
@ -292,37 +285,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 +357,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 +386,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)
|
||||
|
|
|
|||
|
|
@ -15,13 +15,9 @@ rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
|
|||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_rust_transcription() -> None:
|
||||
rust_bridge.configure_rust_transcription(
|
||||
transcription=None,
|
||||
atranscription=None,
|
||||
decline=lambda model, custom_llm_provider, *, context: None,
|
||||
)
|
||||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None)
|
||||
yield
|
||||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None, decline=None)
|
||||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None)
|
||||
|
||||
|
||||
class SyncBridge:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue