diff --git a/litellm-rust/crates/host-python/src/argument.rs b/litellm-rust/crates/host-python/src/argument.rs index 34e07cdfbd5..13214e0e9ed 100644 --- a/litellm-rust/crates/host-python/src/argument.rs +++ b/litellm-rust/crates/host-python/src/argument.rs @@ -1,8 +1,5 @@ use pyo3::{prelude::*, types::PyDict}; -/// The caller's own object for a public argument: the keyword if given, even an explicit -/// `None`, else the bound request's attribute. Every reader of a public Python call uses -/// this rule, so the callbacks and the provider see one object per argument. pub fn lookup<'py>( kwargs: &Bound<'py, PyDict>, request: &Bound<'py, PyAny>, @@ -11,6 +8,9 @@ pub fn lookup<'py>( if let Some(value) = kwargs.get_item(name)? { return Ok(Some(value)); } + if let Ok(bound) = request.cast::() { + return bound.get_item(name); + } request.getattr_opt(name) } @@ -48,4 +48,32 @@ kwargs = {'api_key': key, 'api_base': None} assert!(find("model").is_none()); }); } + + #[rstest::rstest] + #[case::prepared_value("{'api_key': 'replacement'}", Some("replacement"))] + #[case::explicit_none("{'api_key': None}", None)] + #[case::bound_fallback("{}", Some("original"))] + fn prepared_mapping_overrides_bound_values( + #[case] source: &str, + #[case] expected: Option<&str>, + ) { + crate::initialize_python(); + Python::attach(|py| { + let bound = PyDict::new(py); + bound.set_item("api_key", "original").unwrap(); + let source = std::ffi::CString::new(source).unwrap(); + let prepared = py + .eval(&source, None, None) + .unwrap() + .cast_into::() + .unwrap(); + let value = lookup(&prepared, bound.as_any(), "api_key") + .unwrap() + .unwrap(); + assert_eq!( + value.extract::>().unwrap().as_deref(), + expected + ); + }); + } } diff --git a/litellm-rust/crates/inference-messages/src/lib.rs b/litellm-rust/crates/inference-messages/src/lib.rs index cc27ba5e652..cc48a2a39b1 100644 --- a/litellm-rust/crates/inference-messages/src/lib.rs +++ b/litellm-rust/crates/inference-messages/src/lib.rs @@ -14,7 +14,9 @@ use litellm_secrets::source::SecretSource; use std::sync::Arc; pub use litellm_inference::RouteError as Error; -pub use types::{MessagesCall, MessagesCallResponse, MessagesShaping, messages_body}; +pub use types::{ + MessagesCall, MessagesCallResponse, MessagesSettings, MessagesShaping, messages_body, +}; #[derive(Clone)] pub struct MessagesRoute { diff --git a/litellm-rust/crates/inference-messages/src/prepare.rs b/litellm-rust/crates/inference-messages/src/prepare.rs index 312bf144d21..b2a09f3baf9 100644 --- a/litellm-rust/crates/inference-messages/src/prepare.rs +++ b/litellm-rust/crates/inference-messages/src/prepare.rs @@ -80,12 +80,13 @@ fn prepare_provider_request( let sanitized = config.shape_request( MessagesRequest { model, ..body }, - shaping.reasoning_auto_summary, + shaping.settings.reasoning_auto_summary, )?; - let trimmed = without_additional_drop_params(sanitized, &shaping.additional_drop_params)?; + let trimmed = + without_additional_drop_params(sanitized, &shaping.settings.additional_drop_params)?; let transformed = config.transform_anthropic_messages_request( trimmed, - &MessagesTransformContext::new(shaping.capabilities, shaping.drop_params), + &MessagesTransformContext::new(shaping.capabilities, shaping.settings.drop_params), )?; let scoped = @@ -148,7 +149,7 @@ mod tests { use serde_json::{Map, Value, json}; use super::*; - use crate::MessagesShaping; + use crate::{MessagesSettings, MessagesShaping}; #[fixture] fn shaping() -> MessagesShaping { @@ -310,10 +311,13 @@ mod tests { ) }; let shaping = MessagesShaping { - additional_drop_params: additional_drop_params - .iter() - .map(ToString::to_string) - .collect(), + settings: MessagesSettings { + additional_drop_params: additional_drop_params + .iter() + .map(ToString::to_string) + .collect(), + ..shaping.settings + }, ..shaping }; assert_eq!( @@ -394,8 +398,11 @@ mod tests { #[rstest] fn dropped_thinking_display_is_not_restored_by_auto_summary(shaping: MessagesShaping) { let shaping = MessagesShaping { - reasoning_auto_summary: true, - additional_drop_params: vec!["thinking.display".to_string()], + settings: MessagesSettings { + reasoning_auto_summary: true, + additional_drop_params: vec!["thinking.display".to_string()], + ..shaping.settings + }, ..shaping }; assert_eq!( @@ -420,7 +427,10 @@ mod tests { #[rstest] fn dropping_an_invalid_metadata_user_id_does_not_skip_its_validation(shaping: MessagesShaping) { let shaping = MessagesShaping { - additional_drop_params: vec!["metadata.user_id".to_string()], + settings: MessagesSettings { + additional_drop_params: vec!["metadata.user_id".to_string()], + ..shaping.settings + }, ..shaping }; assert!(matches!( diff --git a/litellm-rust/crates/inference-messages/src/types.rs b/litellm-rust/crates/inference-messages/src/types.rs index 6736e9178ba..8521345b00e 100644 --- a/litellm-rust/crates/inference-messages/src/types.rs +++ b/litellm-rust/crates/inference-messages/src/types.rs @@ -38,6 +38,12 @@ pub type MessagesCallResponse = pub struct MessagesShaping { #[serde(default)] pub capabilities: MessagesModelCapabilities, + #[serde(flatten)] + pub settings: MessagesSettings, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct MessagesSettings { #[serde(default)] pub drop_params: bool, #[serde(default)] @@ -58,16 +64,25 @@ mod tests { #[case::nothing_projected(json!({}), MessagesShaping::default())] #[case::only_drop_params( json!({"drop_params": true}), - MessagesShaping { drop_params: true, ..MessagesShaping::default() }, + MessagesShaping { + settings: MessagesSettings { drop_params: true, ..MessagesSettings::default() }, + ..MessagesShaping::default() + }, )] #[case::only_reasoning_auto_summary( json!({"reasoning_auto_summary": true}), - MessagesShaping { reasoning_auto_summary: true, ..MessagesShaping::default() }, + MessagesShaping { + settings: MessagesSettings { reasoning_auto_summary: true, ..MessagesSettings::default() }, + ..MessagesShaping::default() + }, )] #[case::only_additional_drop_params( json!({"additional_drop_params": ["tools[*].input_examples"]}), MessagesShaping { - additional_drop_params: vec!["tools[*].input_examples".to_string()], + settings: MessagesSettings { + additional_drop_params: vec!["tools[*].input_examples".to_string()], + ..MessagesSettings::default() + }, ..MessagesShaping::default() }, )] @@ -98,6 +113,11 @@ mod tests { "additional_drop_params": ["metadata.user_id", "thinking"] }), MessagesShaping { + settings: MessagesSettings { + drop_params: true, + reasoning_auto_summary: true, + additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()], + }, capabilities: MessagesModelCapabilities { supports_reasoning: true, supports_adaptive_thinking: true, @@ -115,9 +135,6 @@ mod tests { max: false, }, }, - drop_params: true, - reasoning_auto_summary: true, - additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()], }, )] fn shaping_deserializes_with_defaults_for_absent_fields( @@ -126,5 +143,22 @@ mod tests { ) { let shaping: MessagesShaping = serde_json::from_value(projected).unwrap(); assert_eq!(shaping, expected); + let serialized = serde_json::to_value(&shaping).unwrap(); + assert_eq!( + serialized["drop_params"], + json!(expected.settings.drop_params) + ); + assert_eq!( + serialized["reasoning_auto_summary"], + json!(expected.settings.reasoning_auto_summary) + ); + assert_eq!( + serialized["additional_drop_params"], + json!(expected.settings.additional_drop_params) + ); + assert_eq!( + serialized["capabilities"], + serde_json::to_value(expected.capabilities).unwrap() + ); } } diff --git a/litellm-rust/crates/inference-messages/tests/messages/host.rs b/litellm-rust/crates/inference-messages/tests/messages/host.rs index e2adf04a6c0..ba20fb9eec4 100644 --- a/litellm-rust/crates/inference-messages/tests/messages/host.rs +++ b/litellm-rust/crates/inference-messages/tests/messages/host.rs @@ -380,12 +380,14 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages let host = RecordingHost::passthrough(authenticated( MessagesCall { shaping: MessagesShaping { + settings: MessagesSettings { + drop_params: true, + ..MessagesSettings::default() + }, capabilities: AnthropicModelCapabilities { supports_sampling_params: false, ..AnthropicModelCapabilities::default() }, - drop_params: true, - ..MessagesShaping::default() }, ..with_fields(call, json!({"temperature": 0.2})) }, diff --git a/litellm-rust/crates/inference-messages/tests/messages/main.rs b/litellm-rust/crates/inference-messages/tests/messages/main.rs index a9e4a881c0a..509a8b19847 100644 --- a/litellm-rust/crates/inference-messages/tests/messages/main.rs +++ b/litellm-rust/crates/inference-messages/tests/messages/main.rs @@ -5,7 +5,7 @@ use std::{ use litellm_http::{HttpSettings, Resolution}; use litellm_inference_messages::{ - Error, MessagesCall, MessagesShaping, + Error, MessagesCall, MessagesSettings, MessagesShaping, route::{Messages, MessagesMachine, MessagesOutput}, }; use litellm_inference_testing::RecordingSecrets; diff --git a/litellm-rust/crates/inference-messages/tests/messages/request.rs b/litellm-rust/crates/inference-messages/tests/messages/request.rs index 6a01be2b4f4..707e91c2d0d 100644 --- a/litellm-rust/crates/inference-messages/tests/messages/request.rs +++ b/litellm-rust/crates/inference-messages/tests/messages/request.rs @@ -239,7 +239,10 @@ async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall) api_key: Some("sk".into()), api_base: Some(upstream.uri()), shaping: MessagesShaping { - additional_drop_params: vec!["temperature".into()], + settings: MessagesSettings { + additional_drop_params: vec!["temperature".into()], + ..MessagesSettings::default() + }, ..MessagesShaping::default() }, ..with_fields(call, json!({"temperature": 0.5, "top_k": 3})) @@ -390,8 +393,10 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i api_base: Some(upstream.uri()), shaping: MessagesShaping { capabilities, - drop_params, - ..MessagesShaping::default() + settings: MessagesSettings { + drop_params, + ..MessagesSettings::default() + }, }, body: call.body.clone(), custom_llm_provider: call.custom_llm_provider.clone(), @@ -436,13 +441,15 @@ async fn reasoning_auto_summary_marks_active_thinking_on_the_wire( api_key: Some("sk".into()), api_base: Some(upstream.uri()), shaping: MessagesShaping { + settings: MessagesSettings { + reasoning_auto_summary: true, + ..MessagesSettings::default() + }, capabilities: MessagesModelCapabilities { supports_reasoning: true, supports_adaptive_thinking: true, ..MessagesModelCapabilities::default() }, - reasoning_auto_summary: true, - ..MessagesShaping::default() }, ..call }, @@ -634,7 +641,10 @@ async fn system_message_folding_is_selected_by_the_provider( api_key: Some("sk-azure".into()), api_base: Some(upstream.uri()), shaping: MessagesShaping { - additional_drop_params: drop_params.iter().map(ToString::to_string).collect(), + settings: MessagesSettings { + additional_drop_params: drop_params.iter().map(ToString::to_string).collect(), + ..call.shaping.settings + }, ..call.shaping }, ..call @@ -695,7 +705,10 @@ async fn provider_validation_runs_before_caller_parameter_removal( api_key: Some("sk-test".into()), api_base: Some(upstream.uri()), shaping: MessagesShaping { - additional_drop_params: vec!["metadata".into()], + settings: MessagesSettings { + additional_drop_params: vec!["metadata".into()], + ..call.shaping.settings + }, ..call.shaping }, ..call diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 7858b695edf..d1d0e5ff270 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -6,6 +6,7 @@ use std::{ use litellm_auth::InputSource; use litellm_host_python::{from_py, from_py_argument}; use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; +use serde::de::DeserializeOwned; use serde_json::{Map, Value}; /// The keyword arguments every value route shares, validated at the Python boundary. @@ -25,18 +26,6 @@ pub(crate) fn messages_argument(value: &Bound<'_, PyAny>) -> PyResult } } -pub(crate) fn optional_params_argument( - value: &Bound<'_, PyAny>, -) -> PyResult>> { - optional_object("optional_params", value) -} - -pub(crate) fn extra_headers_argument( - value: &Bound<'_, PyAny>, -) -> PyResult>> { - optional_object("extra_headers", value) -} - fn required_object(name: &'static str, value: Value) -> PyResult> { match value { Value::Object(values) => Ok(values), @@ -71,6 +60,48 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py) -> PyRe .extract() } +pub(crate) fn required_field<'py>( + fields: &Bound<'py, PyDict>, + name: &str, +) -> PyResult> { + fields + .get_item(name)? + .ok_or_else(|| PyValueError::new_err(format!("{name} is required"))) +} + +pub(crate) fn optional_field( + fields: &Bound<'_, PyDict>, + name: &str, +) -> PyResult> { + fields + .get_item(name)? + .map(|value| from_py_argument(&value)) + .transpose() + .map(Option::flatten) +} + +pub(crate) fn optional_object_field( + fields: &Bound<'_, PyDict>, + name: &'static str, +) -> PyResult>> { + fields + .get_item(name)? + .map(|value| optional_object(name, &value)) + .transpose() + .map(Option::flatten) +} + +pub(crate) fn value_route_options(fields: &Bound<'_, PyDict>) -> PyResult { + Ok(RouteOptions { + model: from_py_argument(&required_field(fields, "model")?)?, + api_key: optional_field(fields, "api_key")?, + api_base: optional_field(fields, "api_base")?, + custom_llm_provider: optional_field(fields, "custom_llm_provider")?, + extra_headers: optional_object_field(fields, "extra_headers")?, + timeout: optional_timeout(optional_field(fields, "timeout_seconds")?), + }) +} + pub(crate) fn project_optional_fields( kwargs: &Bound<'_, PyDict>, names: &[&str], @@ -244,15 +275,15 @@ mod tests { let params = py.eval(c"{'temperature': 0.2}", None, None).unwrap(); assert_eq!( - optional_params_argument(¶ms).unwrap(), + optional_object("optional_params", ¶ms).unwrap(), Some(required_object("optional_params", json!({"temperature": 0.2})).unwrap()) ); assert_eq!( - optional_params_argument(&py.None().into_bound(py)).unwrap(), + optional_object("optional_params", &py.None().into_bound(py)).unwrap(), None ); assert_eq!( - extra_headers_argument(&py.None().into_bound(py)).unwrap(), + optional_object("extra_headers", &py.None().into_bound(py)).unwrap(), None ); }); diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index c26b9c75734..fc7bfb08b75 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -1,14 +1,13 @@ use crate::execution::{run_async, run_sync}; -use litellm_host_python::from_py_argument; use litellm_inference_transcription::{ AudioTranscriptionRoute, Error, types::AudioTranscriptionRequest, }; -use pyo3::{prelude::*, types::PyDict}; +use pyo3::prelude::*; use serde_json::{Map, Value}; use crate::{ errors::route_error_to_pyerr, - marshal::{RouteOptions, extra_headers_argument, optional_params_argument, optional_timeout}, + marshal::{RouteOptions, optional_object_field, required_field, value_route_options}, }; async fn execute( @@ -41,81 +40,38 @@ async fn execute( } #[pyfunction] -#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))] -#[expect( - clippy::too_many_arguments, - reason = "one parameter per Python keyword" -)] -pub(crate) fn transcription( - py: Python<'_>, - model: String, - #[pyo3(from_py_with = from_py_argument)] audio: Value, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option>, - #[pyo3(from_py_with = optional_params_argument)] optional_params: Option>, - timeout_seconds: Option, -) -> PyResult> { - let options = RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout: optional_timeout(timeout_seconds), - }; - let http = crate::http::provider_client(py, &PyDict::new(py), false)?; +pub(crate) fn transcription(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + let audio: Value = + litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?; + let options = value_route_options(&call.bound)?; + let optional_params = + optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); + let http = crate::http::provider_client(py, &call.kwargs, false)?; let secrets = crate::secrets::source(py)?; run_sync( py, - execute( - http, - secrets, - audio, - optional_params.unwrap_or_default(), - options, - ), + execute(http, secrets, audio, optional_params, options), route_error_to_pyerr, ) } #[pyfunction] -#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))] -#[expect( - clippy::too_many_arguments, - reason = "one parameter per Python keyword" -)] pub(crate) fn atranscription<'py>( py: Python<'py>, - model: String, - #[pyo3(from_py_with = from_py_argument)] audio: Value, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option>, - #[pyo3(from_py_with = optional_params_argument)] optional_params: Option>, - timeout_seconds: Option, + call: Bound<'py, PyAny>, ) -> PyResult> { - let options = RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout: optional_timeout(timeout_seconds), - }; - let http = crate::http::provider_client(py, &PyDict::new(py), true)?; + let call = super::NativeCall::extract(&call)?; + let audio: Value = + litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?; + let options = value_route_options(&call.bound)?; + let optional_params = + optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); + let http = crate::http::provider_client(py, &call.kwargs, true)?; let secrets = crate::secrets::source(py)?; run_async( py, - execute( - http, - secrets, - audio, - optional_params.unwrap_or_default(), - options, - ), + execute(http, secrets, audio, optional_params, options), route_error_to_pyerr, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 14f6cbc83b4..27f076b7cce 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -11,8 +11,7 @@ use serde_json::{Map, Value}; use crate::{ errors::route_error_to_pyerr, marshal::{ - RouteOptions, extra_headers_argument, messages_argument, optional_params_argument, - optional_timeout, + RouteOptions, messages_argument, optional_object_field, required_field, value_route_options, }, }; @@ -50,81 +49,36 @@ async fn execute( } #[pyfunction] -#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] -#[expect( - clippy::too_many_arguments, - reason = "one parameter per Python keyword" -)] -pub(crate) fn chat_completions( - py: Python<'_>, - model: String, - #[pyo3(from_py_with = messages_argument)] messages: Vec, - #[pyo3(from_py_with = optional_params_argument)] optional_params: Option>, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult> { - let options = RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout: optional_timeout(timeout_seconds), - }; - let http = crate::http::provider_client(py, &PyDict::new(py), false)?; +pub(crate) fn chat_completions(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + let messages: Vec = messages_argument(&required_field(&call.bound, "messages")?)?; + let optional_params = + optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); + let options = value_route_options(&call.bound)?; + let http = crate::http::provider_client(py, &call.kwargs, false)?; let secrets = crate::secrets::source(py)?; run_sync( py, - execute( - http, - secrets, - messages, - optional_params.unwrap_or_default(), - options, - ), + execute(http, secrets, messages, optional_params, options), route_error_to_pyerr, ) } #[pyfunction] -#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] -#[expect( - clippy::too_many_arguments, - reason = "one parameter per Python keyword" -)] pub(crate) fn achat_completions<'py>( py: Python<'py>, - model: String, - #[pyo3(from_py_with = messages_argument)] messages: Vec, - #[pyo3(from_py_with = optional_params_argument)] optional_params: Option>, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option>, - timeout_seconds: Option, + call: Bound<'py, PyAny>, ) -> PyResult> { - let options = RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout: optional_timeout(timeout_seconds), - }; - let http = crate::http::provider_client(py, &PyDict::new(py), true)?; + let call = super::NativeCall::extract(&call)?; + let messages: Vec = messages_argument(&required_field(&call.bound, "messages")?)?; + let optional_params = + optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); + let options = value_route_options(&call.bound)?; + let http = crate::http::provider_client(py, &call.kwargs, true)?; let secrets = crate::secrets::source(py)?; run_async( py, - execute( - http, - secrets, - messages, - optional_params.unwrap_or_default(), - options, - ), + execute(http, secrets, messages, optional_params, options), route_error_to_pyerr, ) } @@ -184,21 +138,13 @@ fn run_public( } #[pyfunction] -pub(crate) fn completion( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - run_public(py, request, args, kwargs, false) +pub(crate) fn completion(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + run_public(py, call.bound.into_any(), call.args, call.kwargs, false) } #[pyfunction] -pub(crate) fn acompletion( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - run_public(py, request, args, kwargs, true) +pub(crate) fn acompletion(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + run_public(py, call.bound.into_any(), call.args, call.kwargs, true) } diff --git a/litellm-rust/crates/python-bridge/src/routes/embeddings.rs b/litellm-rust/crates/python-bridge/src/routes/embeddings.rs index b1681a2e652..1d34e3c21ff 100644 --- a/litellm-rust/crates/python-bridge/src/routes/embeddings.rs +++ b/litellm-rust/crates/python-bridge/src/routes/embeddings.rs @@ -1,58 +1,49 @@ -use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, -}; +use pyo3::prelude::*; use crate::errors::RustBridgeDeclined; #[pyfunction] -#[pyo3(signature = (request, args, kwargs))] -pub(crate) fn embedding( - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - drop((request, args, kwargs)); +pub(crate) fn embedding(call: Bound<'_, PyAny>) -> PyResult> { + drop(super::NativeCall::extract(&call)?); Err(RustBridgeDeclined::new_err( "native embeddings route is not implemented", )) } #[pyfunction] -#[pyo3(signature = (request, args, kwargs))] -pub(crate) fn aembedding( - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - drop((request, args, kwargs)); - Err(RustBridgeDeclined::new_err( - "native embeddings route is not implemented", - )) +pub(crate) fn aembedding(call: Bound<'_, PyAny>) -> PyResult> { + embedding(call) } #[cfg(test)] mod tests { - use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, - }; + use pyo3::{prelude::*, types::PyDict}; + use rstest::rstest; use crate::errors::RustBridgeDeclined; - #[test] - fn both_entrypoints_decline_before_provider_execution() { + #[rstest] + #[case::sync(false)] + #[case::asynchronous(true)] + fn both_entrypoints_decline_before_provider_execution(#[case] asynchronous: bool) { Python::initialize(); Python::attach(|py| { - let request = PyDict::new(py); - let args = PyTuple::empty(py); - let kwargs = PyDict::new(py); - - for entrypoint in [super::embedding, super::aembedding] { - let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone()) - .expect_err("native embeddings must decline until a route machine exists"); - assert!(error.is_instance_of::(py)); + let locals = PyDict::new(py); + py.run( + c"from types import SimpleNamespace +call = SimpleNamespace(args=(), kwargs={}, bound={'model':'test-model','input':'hello'})", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let call = locals.get_item("call").unwrap().unwrap(); + let error = if asynchronous { + super::aembedding(call) + } else { + super::embedding(call) } + .expect_err("native embeddings must decline until a route machine exists"); + assert!(error.is_instance_of::(py)); }); } } diff --git a/litellm-rust/crates/python-bridge/src/routes/inference.rs b/litellm-rust/crates/python-bridge/src/routes/inference.rs index cb1e31c7c85..ed18ed659b9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/inference.rs +++ b/litellm-rust/crates/python-bridge/src/routes/inference.rs @@ -86,6 +86,9 @@ impl InferenceHost { if let Some(value) = lookup(arguments, request, name)? { return Ok((!value.is_none()).then_some(value)); } + if request.is_instance_of::() { + return Ok(None); + } let parameter = request .getattr("parameters")? .call_method1("get", (name,))?; diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 7b80b34c50b..0271e744b5c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -5,9 +5,10 @@ use bytes::Bytes; use litellm_host_python::{InvokeError, PythonBinding, from_py, lookup, to_py}; use litellm_http::transport::Error as TransportError; use litellm_inference_messages::{ - Error, MessagesCall, MessagesShaping, messages_body, + Error, MessagesCall, MessagesSettings, MessagesShaping, messages_body, route::{Messages, MessagesStreamHead}, }; +use litellm_llms::base_llm::messages::context::MessagesModelCapabilities; use litellm_llms_types::headers::ProviderSpecificHeaders; use pyo3::{ exceptions::{PyException, PyValueError}, @@ -188,18 +189,24 @@ impl MessagesPythonHost { custom_llm_provider: Option<&str>, arguments: &Bound<'_, PyDict>, ) -> PyResult { - let projected = py.import(ROUTE_HOST_MODULE)?.getattr("shaping")?.call1(( - model, - custom_llm_provider, - arguments, - ))?; - from_py(&projected) + let module = py.import(ROUTE_HOST_MODULE)?; + let capabilities: MessagesModelCapabilities = from_py( + &py.import("litellm.rust_bridge.model_capabilities")? + .getattr("anthropic_model_capabilities")? + .call1((model, custom_llm_provider))?, + )?; + let settings: MessagesSettings = + from_py(&module.getattr("settings")?.call1((arguments,))?)?; + Ok(MessagesShaping { + capabilities, + settings, + }) } fn provider(&self, py: Python<'_>) -> String { self.request .bind(py) - .getattr("custom_llm_provider") + .get_item("custom_llm_provider") .and_then(|value| value.extract::>()) .ok() .flatten() diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index c234e88b842..318aae27121 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -64,21 +64,13 @@ fn run_messages( } #[pyfunction] -pub(crate) fn messages( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - run_messages(py, request, args, kwargs, false) +pub(crate) fn messages(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + run_messages(py, call.bound.into_any(), call.args, call.kwargs, false) } #[pyfunction] -pub(crate) fn amessages( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - run_messages(py, request, args, kwargs, true) +pub(crate) fn amessages(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + run_messages(py, call.bound.into_any(), call.args, call.kwargs, true) } diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 2380274001e..bf52fe3bfe8 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -14,9 +14,35 @@ use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol} use litellm_host_python::{HookChain, PythonBinding, PythonCallHooks, PythonHostCalls}; use pyo3::{ prelude::*, - types::{PyDict, PyTuple}, + types::{PyDict, PyMapping, PyTuple}, }; +struct NativeCall<'py> { + args: Bound<'py, PyTuple>, + kwargs: Bound<'py, PyDict>, + bound: Bound<'py, PyDict>, +} + +impl<'py> NativeCall<'py> { + fn extract(call: &Bound<'py, PyAny>) -> PyResult { + Ok(Self { + args: call.getattr("args")?.cast_into()?, + kwargs: mapping_dict(&call.getattr("kwargs")?)?, + bound: mapping_dict(&call.getattr("bound")?)?, + }) + } +} + +fn mapping_dict<'py>(value: &Bound<'py, PyAny>) -> PyResult> { + if let Ok(dict) = value.cast::() { + return Ok(dict.clone()); + } + let mapping = value.cast::()?; + let dict = PyDict::new(value.py()); + dict.update(mapping)?; + Ok(dict) +} + fn call_hooks( py: Python<'_>, operation: LoggingOperation, @@ -74,40 +100,30 @@ mod tests { types::{PyDict, PyList}, }; - #[test] - fn sync_and_async_route_signatures_match_the_python_contract() { - Python::initialize(); - Python::attach(|py| { - let module = crate::native_module(py); - let routes = [ - ( - "transcription", - "atranscription", - "(model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)", - ), - ( - "chat_completions", - "achat_completions", - "(model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)", - ), - ]; - - for (sync_name, async_name, expected) in routes { - let sync_signature: String = module - .getattr(sync_name) - .and_then(|function| function.getattr("__text_signature__")) - .and_then(|signature| signature.extract()) - .expect("sync signature should be available"); - let async_signature: String = module - .getattr(async_name) - .and_then(|function| function.getattr("__text_signature__")) - .and_then(|signature| signature.extract()) - .expect("async signature should be available"); - - assert_eq!(sync_signature, expected); - assert_eq!(async_signature, expected); - } - }); + fn value_call<'py>( + py: Python<'py>, + payload_name: &str, + payload: &Bound<'py, PyAny>, + kwargs: Option<&Bound<'py, PyDict>>, + ) -> Bound<'py, PyAny> { + let fields = PyDict::new(py); + fields.set_item("model", "model").unwrap(); + fields.set_item(payload_name, payload).unwrap(); + if let Some(kwargs) = kwargs { + fields.update(kwargs.as_mapping()).unwrap(); + } + let attributes = PyDict::new(py); + attributes + .set_item("args", pyo3::types::PyTuple::empty(py)) + .unwrap(); + attributes.set_item("kwargs", &fields).unwrap(); + attributes.set_item("bound", &fields).unwrap(); + py.import("types") + .unwrap() + .getattr("SimpleNamespace") + .unwrap() + .call((), Some(&attributes)) + .unwrap() } #[test] @@ -138,7 +154,9 @@ value = Broken() for name in ["chat_completions", "achat_completions"] { let error = module .getattr(name) - .and_then(|function| function.call1(("model", &broken))) + .and_then(|function| { + function.call1((value_call(py, "messages", &broken, None),)) + }) .expect_err("route should reject a value it cannot convert"); assert!( @@ -158,11 +176,15 @@ value = Broken() let invalid_messages = PyDict::new(py); let sync_chat_error = module .getattr("chat_completions") - .and_then(|function| function.call1(("model", &invalid_messages))) + .and_then(|function| { + function.call1((value_call(py, "messages", &invalid_messages, None),)) + }) .expect_err("sync chat should reject a non-list messages value"); let async_chat_error = module .getattr("achat_completions") - .and_then(|function| function.call1(("model", &invalid_messages))) + .and_then(|function| { + function.call1((value_call(py, "messages", &invalid_messages, None),)) + }) .expect_err("async chat should reject a non-list messages value"); assert_eq!( @@ -180,11 +202,15 @@ value = Broken() let sync_error = module .getattr("transcription") - .and_then(|function| function.call(("model", &audio), Some(&kwargs))) + .and_then(|function| { + function.call1((value_call(py, "audio", &audio, Some(&kwargs)),)) + }) .expect_err("sync route should reject non-dict extra_headers"); let async_error = module .getattr("atranscription") - .and_then(|function| function.call(("model", &audio), Some(&kwargs))) + .and_then(|function| { + function.call1((value_call(py, "audio", &audio, Some(&kwargs)),)) + }) .expect_err("async route should reject non-dict extra_headers"); assert_eq!( @@ -213,7 +239,12 @@ value = Broken() let error = module .getattr("chat_completions") .and_then(|function| { - function.call(("model", &invalid_messages), Some(&chat_kwargs)) + function.call1((value_call( + py, + "messages", + &invalid_messages, + Some(&chat_kwargs), + ),)) }) .expect_err("messages should be validated first"); assert_eq!(error.to_string(), "ValueError: messages must be a list"); @@ -221,7 +252,14 @@ value = Broken() let valid_messages = PyList::empty(py); let error = module .getattr("chat_completions") - .and_then(|function| function.call(("model", &valid_messages), Some(&chat_kwargs))) + .and_then(|function| { + function.call1((value_call( + py, + "messages", + &valid_messages, + Some(&chat_kwargs), + ),)) + }) .expect_err("optional_params should be validated before headers"); assert_eq!( error.to_string(), @@ -237,7 +275,12 @@ value = Broken() let error = module .getattr("transcription") .and_then(|function| { - function.call(("model", &invalid_payload), Some(&headers_kwargs)) + function.call1((value_call( + py, + "audio", + &invalid_payload, + Some(&headers_kwargs), + ),)) }) .expect_err("payload should be validated before headers"); assert!(!error.to_string().contains("extra_headers")); @@ -265,11 +308,15 @@ value = Broken() let omitted_error = module .getattr("chat_completions") - .and_then(|function| function.call(("model", &messages), Some(&omitted))) + .and_then(|function| { + function.call1((value_call(py, "messages", &messages, Some(&omitted)),)) + }) .expect_err("omitted optional_params should reach header validation"); let explicit_error = module .getattr("chat_completions") - .and_then(|function| function.call(("model", &messages), Some(&explicit))) + .and_then(|function| { + function.call1((value_call(py, "messages", &messages, Some(&explicit)),)) + }) .expect_err("None optional_params should reach header validation"); assert_eq!( omitted_error.to_string(), diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index 3b914c70c19..71b1054fc3c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -81,23 +81,15 @@ fn project_provider_defaults(snapshot: &Snapshot<'_>) -> PyResult { } #[pyfunction] -pub(crate) fn ocr( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - run_ocr(py, request, args, kwargs, false) +pub(crate) fn ocr(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + run_ocr(py, call.bound.into_any(), call.args, call.kwargs, false) } #[pyfunction] -pub(crate) fn aocr( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - run_ocr(py, request, args, kwargs, true) +pub(crate) fn aocr(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + run_ocr(py, call.bound.into_any(), call.args, call.kwargs, true) } #[pyfunction] diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index 840bca89b51..780f2e6929b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -105,23 +105,15 @@ fn run_public( } #[pyfunction] -pub(crate) fn responses( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - run_public(py, request, args, kwargs, false) +pub(crate) fn responses(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + run_public(py, call.bound.into_any(), call.args, call.kwargs, false) } #[pyfunction] -pub(crate) fn aresponses( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, -) -> PyResult> { - run_public(py, request, args, kwargs, true) +pub(crate) fn aresponses(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { + let call = super::NativeCall::extract(&call)?; + run_public(py, call.bound.into_any(), call.args, call.kwargs, true) } #[pyclass] diff --git a/litellm-rust/crates/router/tests/router.rs b/litellm-rust/crates/router/tests/router.rs index a2dc21a0708..97ee544f120 100644 --- a/litellm-rust/crates/router/tests/router.rs +++ b/litellm-rust/crates/router/tests/router.rs @@ -1,7 +1,7 @@ use std::time::Duration; use litellm_config::Config; -use litellm_inference_messages::MessagesShaping; +use litellm_inference_messages::{MessagesSettings, MessagesShaping}; use litellm_router::{Deployment, Router}; use rstest::rstest; @@ -74,8 +74,11 @@ fn programmatic_deployments_preserve_overrides_and_last_entry_wins() { custom_llm_provider: Some("test-provider".into()), timeout: Some(Duration::from_secs(7)), shaping: MessagesShaping { - drop_params: true, - additional_drop_params: vec!["metadata.test".into()], + settings: MessagesSettings { + drop_params: true, + additional_drop_params: vec!["metadata.test".into()], + ..MessagesSettings::default() + }, ..Default::default() }, }; diff --git a/litellm/chat_completions/dispatch.py b/litellm/chat_completions/dispatch.py index b2b9662dd78..2eee37d41ed 100644 --- a/litellm/chat_completions/dispatch.py +++ b/litellm/chat_completions/dispatch.py @@ -1,6 +1,5 @@ import inspect from collections.abc import Awaitable, Callable, Coroutine, Mapping -from types import MappingProxyType from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable from litellm import main @@ -8,13 +7,13 @@ from litellm.rust_bridge.catalog import Route, RouteContext from litellm.rust_bridge.chat_completions.entrypoints import ( NATIVE_ACOMPLETION, NATIVE_COMPLETION, - LiteLLMChatCompletionsRequest, ) -from litellm.rust_bridge.dispatch import PublicDispatch, call_hook +from litellm.rust_bridge.dispatch import PublicDispatch from litellm.rust_bridge.public_call import ( + NativeCall, bind, - optional_bool, - optional_mapping, + native_call, + native_call_hook, optional_sequence, optional_str, signature, @@ -51,33 +50,22 @@ _ACOMPLETION: Final = signature(_PYTHON_ACOMPLETION) def _public_request( legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object] -) -> LiteLLMChatCompletionsRequest | None: +) -> NativeCall | None: fields: Final = bind(legacy, args, kwargs) if fields is None: return None model: Final = fields.get("model") messages: Final = optional_sequence(fields.get("messages")) - extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({}) if not isinstance(model, str) or messages is None: return None - return LiteLLMChatCompletionsRequest( - model=model, - messages=messages, - stream=optional_bool(fields.get("stream")), - api_key=optional_str(fields.get("api_key")), - api_base=optional_str(extra.get("api_base")) or optional_str(fields.get("base_url")), - custom_llm_provider=optional_str(extra.get("custom_llm_provider")), - extra_headers=optional_mapping(fields.get("extra_headers")), - kwargs=extra, - parameters=MappingProxyType({name: value for name, value in fields.items() if name != "kwargs"}), - ) + return native_call(args, kwargs, fields) -def _context(request: LiteLLMChatCompletionsRequest) -> RouteContext: +def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.CHAT_COMPLETIONS, - provider=request.custom_llm_provider, - model=request.model, + provider=optional_str(request.bound.get("custom_llm_provider")), + model=str(request.bound["model"]), ) @@ -105,7 +93,7 @@ def completion( kwargs, python=python, binding=NATIVE_COMPLETION, - native=call_hook, + native=native_call_hook, ) @@ -116,7 +104,7 @@ async def acompletion(*args: object, **kwargs: object) -> ChatResult: # kwargs- kwargs, python=python, binding=NATIVE_ACOMPLETION, - native=call_hook, + native=native_call_hook, ) diff --git a/litellm/embeddings/dispatch.py b/litellm/embeddings/dispatch.py index bba68d2c0f1..54408691910 100644 --- a/litellm/embeddings/dispatch.py +++ b/litellm/embeddings/dispatch.py @@ -1,17 +1,15 @@ import inspect from collections.abc import Awaitable, Callable, Coroutine, Mapping -from types import MappingProxyType from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable from litellm import main from litellm.rust_bridge.catalog import Route, RouteContext -from litellm.rust_bridge.dispatch import PublicDispatch, call_hook +from litellm.rust_bridge.dispatch import PublicDispatch from litellm.rust_bridge.embeddings.entrypoints import ( NATIVE_AEMBEDDING, NATIVE_EMBEDDING, - LiteLLMEmbeddingRequest, ) -from litellm.rust_bridge.public_call import bind, optional_mapping, optional_str, signature +from litellm.rust_bridge.public_call import NativeCall, bind, native_call, native_call_hook, optional_str, signature from litellm.types.utils import EmbeddingResponse __all__ = ("aembedding", "embedding") @@ -30,28 +28,24 @@ _EMBEDDING_SIGNATURE: Final = signature(_PYTHON_EMBEDDING) def _public_request( legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object] -) -> LiteLLMEmbeddingRequest | None: +) -> NativeCall | None: fields: Final = bind(legacy, args, kwargs) if fields is None: return None model: Final = fields.get("model") if not isinstance(model, str): return None - extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({}) - return LiteLLMEmbeddingRequest( - model=model, - input=fields.get("input"), - api_key=optional_str(fields.get("api_key")), - api_base=optional_str(fields.get("api_base")), - custom_llm_provider=optional_str(fields.get("custom_llm_provider")), - kwargs=extra, + return native_call(args, kwargs, fields) + + +def _context(request: NativeCall) -> RouteContext: + return RouteContext( + Route.EMBEDDINGS, + provider=optional_str(request.bound.get("custom_llm_provider")), + model=str(request.bound["model"]), ) -def _context(request: LiteLLMEmbeddingRequest) -> RouteContext: - return RouteContext(Route.EMBEDDINGS, provider=request.custom_llm_provider, model=request.model) - - _DISPATCH: Final = PublicDispatch( route=Route.EMBEDDINGS, request=lambda args, kwargs: _public_request(_EMBEDDING_SIGNATURE, args, kwargs), @@ -75,7 +69,7 @@ def embedding( kwargs, python=_PYTHON_EMBEDDING, binding=NATIVE_EMBEDDING, - native=call_hook, + native=native_call_hook, ) @@ -85,7 +79,7 @@ async def aembedding(*args: object, **kwargs: object) -> EmbeddingResponse: # k kwargs, python=_PYTHON_AEMBEDDING, binding=NATIVE_AEMBEDDING, - native=call_hook, + native=native_call_hook, ) diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index c62587566c0..0afa5efc29d 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -6,6 +6,7 @@ import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.rust_bridge import runtime from litellm.rust_bridge.catalog import Route, RouteContext +from litellm.rust_bridge.public_call import NativeCall from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.rust_bridge.transcription.native import ( NATIVE_ATRANSCRIPTION, @@ -52,18 +53,18 @@ class BedrockAudioTranscriptionRustDispatch: timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: def native(rust: RustTranscription) -> TranscriptionResponse: - return TranscriptionResponse( - **rust( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_to_seconds(timeout), - ) - ) + fields: Final = { + "model": model, + "audio": self._audio_payload(audio_file), + "api_key": api_key, + "api_base": api_base, + "custom_llm_provider": custom_llm_provider, + "extra_headers": extra_headers, + "optional_params": optional_params, + "timeout_seconds": timeout_to_seconds(timeout), + } + call: Final = NativeCall(args=(), kwargs=fields, bound=fields) + return TranscriptionResponse(**rust(call)) return runtime.run( RouteContext(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model), @@ -85,18 +86,18 @@ class BedrockAudioTranscriptionRustDispatch: timeout: float | httpx.Timeout | None, ) -> TranscriptionResponse: async def native(rust: RustAtranscription) -> TranscriptionResponse: - return TranscriptionResponse( - **await rust( - model=model, - audio=self._audio_payload(audio_file), - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_to_seconds(timeout), - ) - ) + fields: Final = { + "model": model, + "audio": self._audio_payload(audio_file), + "api_key": api_key, + "api_base": api_base, + "custom_llm_provider": custom_llm_provider, + "extra_headers": extra_headers, + "optional_params": optional_params, + "timeout_seconds": timeout_to_seconds(timeout), + } + call: Final = NativeCall(args=(), kwargs=fields, bound=fields) + return TranscriptionResponse(**await rust(call)) return await runtime.arun( RouteContext(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model), diff --git a/litellm/messages/dispatch.py b/litellm/messages/dispatch.py index 54aad495d1b..74736011f20 100644 --- a/litellm/messages/dispatch.py +++ b/litellm/messages/dispatch.py @@ -1,22 +1,21 @@ import inspect from collections.abc import AsyncIterator, Awaitable, Callable, Coroutine, Iterator, Mapping -from types import MappingProxyType from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable from litellm.exceptions import BadRequestError from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.llms.anthropic.pass_through.messages import handler as main from litellm.rust_bridge.catalog import Route, RouteContext -from litellm.rust_bridge.dispatch import PublicDispatch, call_hook +from litellm.rust_bridge.dispatch import PublicDispatch from litellm.rust_bridge.messages.entrypoints import ( NATIVE_AMESSAGES, NATIVE_MESSAGES, - LiteLLMMessagesRequest, ) from litellm.rust_bridge.public_call import ( + NativeCall, bind, - optional_bool, - optional_mapping, + native_call, + native_call_hook, optional_sequence, optional_str, signature, @@ -52,7 +51,7 @@ _AMESSAGES: Final = signature(_PYTHON_AMESSAGES) def _public_request( legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object] -) -> LiteLLMMessagesRequest | None: +) -> NativeCall | None: fields: Final = bind(legacy, args, kwargs) if fields is None: return None @@ -61,30 +60,21 @@ def _public_request( max_tokens: Final = fields.get("max_tokens") if not isinstance(model, str) or messages is None or not isinstance(max_tokens, int): return None - return LiteLLMMessagesRequest( - model=model, - messages=messages, - max_tokens=max_tokens, - stream=optional_bool(fields.get("stream")), - api_key=optional_str(fields.get("api_key")), - api_base=optional_str(fields.get("api_base")), - custom_llm_provider=optional_str(fields.get("custom_llm_provider")), - kwargs=optional_mapping(fields.get("kwargs")) or MappingProxyType({}), - ) + return native_call(args, kwargs, fields) -def _resolved_provider(request: LiteLLMMessagesRequest) -> str | None: +def _resolved_provider(request: NativeCall) -> str | None: try: - return get_llm_provider(request.model, request.custom_llm_provider)[1] + return get_llm_provider(str(request.bound["model"]), optional_str(request.bound.get("custom_llm_provider")))[1] except BadRequestError: - return request.custom_llm_provider + return optional_str(request.bound.get("custom_llm_provider")) -def _context(request: LiteLLMMessagesRequest) -> RouteContext: +def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.MESSAGES, provider=_resolved_provider(request), - model=request.model, + model=str(request.bound["model"]), ) @@ -112,7 +102,7 @@ def anthropic_messages_handler( kwargs, python=python, binding=NATIVE_MESSAGES, - native=call_hook, + native=native_call_hook, ) @@ -123,7 +113,7 @@ async def anthropic_messages(*args: object, **kwargs: object) -> MessagesResult: kwargs, python=python, binding=NATIVE_AMESSAGES, - native=call_hook, + native=native_call_hook, ) diff --git a/litellm/ocr/dispatch.py b/litellm/ocr/dispatch.py index a94e9122ecc..a6a6c5d0c50 100644 --- a/litellm/ocr/dispatch.py +++ b/litellm/ocr/dispatch.py @@ -1,4 +1,5 @@ from collections.abc import Coroutine, Mapping +from types import MappingProxyType from typing import Final import httpx @@ -6,8 +7,9 @@ import httpx from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge import runtime from litellm.rust_bridge.catalog import Route, RouteContext -from litellm.rust_bridge.dispatch import PublicDispatch, call_hook -from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest +from litellm.rust_bridge.dispatch import PublicDispatch +from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR +from litellm.rust_bridge.public_call import NativeCall, native_call, native_call_hook, optional_str __all__ = ("aocr", "ocr") @@ -21,30 +23,33 @@ def _bind_request( custom_llm_provider: str | None = None, extra_headers: dict[str, object] | None = None, **kwargs: object, # kwargs-ok: public OCR accepts provider-specific options -) -> LiteLLMOcrRequest: - return LiteLLMOcrRequest( - model=model, - document=document, - api_key=api_key, - api_base=api_base, - timeout=timeout, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - kwargs=kwargs, +) -> Mapping[str, object]: + return MappingProxyType( + { + "model": model, + "document": document, + "api_key": api_key, + "api_base": api_base, + "timeout": timeout, + "custom_llm_provider": custom_llm_provider, + "extra_headers": extra_headers, + "kwargs": kwargs, + } ) -def _public_request(name: str, args: tuple[object, ...], kwargs: Mapping[str, object]) -> LiteLLMOcrRequest: +def _public_request(name: str, args: tuple[object, ...], kwargs: Mapping[str, object]) -> NativeCall: try: - return _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation + fields: Final = _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation + return native_call(args, kwargs, fields) except TypeError as error: raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None -def _context(request: LiteLLMOcrRequest) -> RouteContext: - prefix, separator, _ = request.model.partition("/") - provider: Final = request.custom_llm_provider or (prefix if separator else None) - return RouteContext(Route.OCR, provider=provider, model=request.model) +def _context(request: NativeCall) -> RouteContext: + prefix, separator, _ = str(request.bound["model"]).partition("/") + provider: Final = optional_str(request.bound.get("custom_llm_provider")) or (prefix if separator else None) + return RouteContext(Route.OCR, provider=provider, model=str(request.bound["model"])) _DISPATCH: Final = PublicDispatch( @@ -70,7 +75,7 @@ def ocr( kwargs, python=runtime.NO_PYTHON, binding=NATIVE_OCR, - native=call_hook, + native=native_call_hook, ) @@ -80,5 +85,5 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr kwargs, python=runtime.NO_PYTHON, binding=NATIVE_AOCR, - native=call_hook, + native=native_call_hook, ) diff --git a/litellm/responses/dispatch.py b/litellm/responses/dispatch.py index 6fe9451cc9a..c728dd0ea13 100644 --- a/litellm/responses/dispatch.py +++ b/litellm/responses/dispatch.py @@ -1,17 +1,22 @@ import inspect from collections.abc import Awaitable, Callable, Coroutine, Mapping -from types import MappingProxyType from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable from litellm.responses import main from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator from litellm.rust_bridge.catalog import Route, RouteContext -from litellm.rust_bridge.dispatch import PublicDispatch, call_hook -from litellm.rust_bridge.public_call import bind, optional_bool, optional_mapping, optional_str, signature +from litellm.rust_bridge.dispatch import PublicDispatch +from litellm.rust_bridge.public_call import ( + NativeCall, + bind, + native_call, + native_call_hook, + optional_str, + signature, +) from litellm.rust_bridge.responses.entrypoints import ( NATIVE_ARESPONSES, NATIVE_RESPONSES, - LiteLLMResponsesRequest, ) from litellm.types.llms.openai import ResponsesAPIResponse @@ -44,32 +49,21 @@ _ARESPONSES: Final = signature(_PYTHON_ARESPONSES) def _public_request( legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object] -) -> LiteLLMResponsesRequest | None: +) -> NativeCall | None: fields: Final = bind(legacy, args, kwargs) if fields is None: return None model: Final = fields.get("model") - extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({}) if not isinstance(model, str): return None - return LiteLLMResponsesRequest( - model=model, - input=fields.get("input"), - stream=optional_bool(fields.get("stream")), - api_key=optional_str(extra.get("api_key")), - api_base=optional_str(extra.get("api_base")) or optional_str(extra.get("base_url")), - custom_llm_provider=optional_str(fields.get("custom_llm_provider")), - extra_headers=optional_mapping(fields.get("extra_headers")), - kwargs=extra, - parameters=MappingProxyType({name: value for name, value in fields.items() if name != "kwargs"}), - ) + return native_call(args, kwargs, fields) -def _context(request: LiteLLMResponsesRequest) -> RouteContext: +def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.RESPONSES, - provider=request.custom_llm_provider, - model=request.model, + provider=optional_str(request.bound.get("custom_llm_provider")), + model=str(request.bound["model"]), ) @@ -97,7 +91,7 @@ def responses( kwargs, python=python, binding=NATIVE_RESPONSES, - native=call_hook, + native=native_call_hook, ) @@ -108,7 +102,7 @@ async def aresponses(*args: object, **kwargs: object) -> ResponsesResult: # kwa kwargs, python=python, binding=NATIVE_ARESPONSES, - native=call_hook, + native=native_call_hook, ) diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 4f9a0b2c492..5ca7b79d127 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -6,11 +6,7 @@ import httpx from pydantic import JsonValue from litellm.llms.base_llm.ocr.transformation import OCRResponse -from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest -from litellm.rust_bridge.embeddings.entrypoints import LiteLLMEmbeddingRequest -from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest -from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest -from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest +from litellm.rust_bridge.public_call import NativeCall from litellm.rust_bridge.trace.generated.types import QueryScope, ReadQueryName, TraceScope from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse from litellm.types.llms.openai import ResponsesAPIResponse @@ -41,7 +37,9 @@ class NativeTraceStorage: def __new__(cls, config: NativeTraceConfig) -> NativeTraceStorage: ... def ensure_schema(self) -> Future[None]: ... def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Future[None]: ... - def ingest(self, payload: bytes, content_type: str | None, tenant: Mapping[str, str], logs: bool = False) -> Future[int]: ... + def ingest( + self, payload: bytes, content_type: str | None, tenant: Mapping[str, str], logs: bool = False + ) -> Future[int]: ... def list_traces( self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None, limit: int ) -> Future[JsonValue]: ... @@ -54,7 +52,9 @@ class NativeTraceStorage: ) -> Future[JsonValue]: ... def query_sql(self, sql: str, scope: QueryScope, secret: str) -> Future[str]: ... def query_help(self, scope: QueryScope, secret: str) -> Future[JsonValue]: ... - def query(self, query: ReadQueryName, parameters: Mapping[str, str | int | float | Sequence[str]]) -> Future[str]: ... + def query( + self, query: ReadQueryName, parameters: Mapping[str, str | int | float | Sequence[str]] + ) -> Future[str]: ... @final class NativeDiagnosticProcessor: @@ -73,96 +73,48 @@ class NativeDiagnosticProcessor: def scrub_access_arguments(self, arguments: Sequence[str]) -> list[str]: ... def ocr( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: dict[str, object], + call: NativeCall, ) -> OCRResponse: ... def aocr( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: dict[str, object], + call: NativeCall, ) -> Coroutine[object, object, OCRResponse]: ... def ocr_health_check_document(model: str, custom_llm_provider: str | None) -> dict[str, object]: ... def ocr_passthrough_response(model: str, endpoint: str, body: bytes) -> dict[str, object] | None: ... def embedding( - request: LiteLLMEmbeddingRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> EmbeddingResponse: ... def aembedding( - request: LiteLLMEmbeddingRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> Coroutine[object, object, EmbeddingResponse]: ... def transcription( - model: str, - audio: object, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: Mapping[str, object] | None = None, - optional_params: Mapping[str, object] | None = None, - timeout_seconds: float | None = None, + call: NativeCall, ) -> dict[str, object]: ... def atranscription( - model: str, - audio: object, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: Mapping[str, object] | None = None, - optional_params: Mapping[str, object] | None = None, - timeout_seconds: float | None = None, + call: NativeCall, ) -> Future[dict[str, object]]: ... def completion( - request: LiteLLMChatCompletionsRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> ModelResponse: ... def acompletion( - request: LiteLLMChatCompletionsRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> Coroutine[object, object, ModelResponse]: ... def responses( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> ResponsesAPIResponse: ... def aresponses( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> Coroutine[object, object, ResponsesAPIResponse]: ... def messages( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: dict[str, object], + call: NativeCall, ) -> AnthropicMessagesResponse | Iterator[bytes]: ... def amessages( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: dict[str, object], + call: NativeCall, ) -> Coroutine[object, object, AnthropicMessagesResponse | AsyncIterator[bytes]]: ... def chat_completions( - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None = None, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: Mapping[str, object] | None = None, - timeout_seconds: float | None = None, + call: NativeCall, ) -> dict[str, object]: ... def achat_completions( - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object] | None = None, - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: Mapping[str, object] | None = None, - timeout_seconds: float | None = None, + call: NativeCall, ) -> Future[dict[str, object]]: ... @final @@ -338,27 +290,42 @@ class _SecretManagerRuntime: def read_secret(self, name: str, settings: Mapping[str, object] | None = None) -> JsonValue: ... def read_secret_async(self, name: str, settings: Mapping[str, object] | None = None) -> Future[JsonValue]: ... def async_write_secret( - self, secret_name: str, secret_value: str, description: str | None = None, + self, + secret_name: str, + secret_value: str, + description: str | None = None, optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, tags: object = None, + timeout: float | httpx.Timeout | None = None, + tags: object = None, ) -> Future[dict[str, JsonValue]]: ... def async_delete_secret( - self, secret_name: str, recovery_window_in_days: int | None = None, + self, + secret_name: str, + recovery_window_in_days: int | None = None, optional_params: Mapping[str, object] | None = None, timeout: float | httpx.Timeout | None = None, ) -> Future[dict[str, JsonValue]]: ... def async_rotate_secret( - self, current_secret_name: str, new_secret_name: str, new_secret_value: str, + self, + current_secret_name: str, + new_secret_name: str, + new_secret_value: str, optional_params: Mapping[str, object] | None = None, timeout: float | httpx.Timeout | None = None, ) -> Future[dict[str, JsonValue]]: ... def sync_read_secret( - self, secret_name: str, optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, primary_secret_name: str | None = None, + self, + secret_name: str, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + primary_secret_name: str | None = None, ) -> JsonValue: ... def async_read_secret( - self, secret_name: str, optional_params: Mapping[str, object] | None = None, - timeout: float | httpx.Timeout | None = None, primary_secret_name: str | None = None, + self, + secret_name: str, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + primary_secret_name: str | None = None, ) -> Future[JsonValue]: ... @final @@ -366,11 +333,18 @@ class NativeCacheHandle: def __new__(cls, _uninstantiable: Never, /) -> Never: ... @staticmethod def memory( - *, ttl: float = 600.0, capacity: int = 200, max_entry_bytes: int = 4194304, + *, + ttl: float = 600.0, + capacity: int = 200, + max_entry_bytes: int = 4194304, ) -> NativeCacheHandle: ... @staticmethod def redis( - url: str, *, namespace: str, ttl: float = 600.0, max_entry_bytes: int = 4194304, + url: str, + *, + namespace: str, + ttl: float = 600.0, + max_entry_bytes: int = 4194304, ) -> NativeCacheHandle: ... def get(self, key: str) -> object: ... def set(self, key: str, value: object, *, ttl: float | None = None) -> None: ... diff --git a/litellm/rust_bridge/chat_completions/entrypoints.py b/litellm/rust_bridge/chat_completions/entrypoints.py index d8cde8d0c66..6bfcffa2c21 100644 --- a/litellm/rust_bridge/chat_completions/entrypoints.py +++ b/litellm/rust_bridge/chat_completions/entrypoints.py @@ -1,42 +1,24 @@ from __future__ import annotations -from collections.abc import Awaitable, Mapping, Sequence -from dataclasses import dataclass, field -from types import MappingProxyType +from collections.abc import Awaitable from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.public_call import NativeCall from litellm.types.utils import ModelResponse -@dataclass(frozen=True, slots=True) -class LiteLLMChatCompletionsRequest: - model: str - messages: Sequence[object] - stream: bool | None - api_key: str | None - api_base: str | None - custom_llm_provider: str | None - extra_headers: Mapping[str, object] | None - kwargs: Mapping[str, object] - parameters: Mapping[str, object] = field(default_factory=lambda: MappingProxyType({})) - - class NativeCompletion(Protocol): def __call__( self, - request: LiteLLMChatCompletionsRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> ModelResponse: ... class NativeAcompletion(Protocol): def __call__( self, - request: LiteLLMChatCompletionsRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> Awaitable[ModelResponse]: ... diff --git a/litellm/rust_bridge/chat_completions/route_host.py b/litellm/rust_bridge/chat_completions/route_host.py index d6b222dcf41..cb06ea8e213 100644 --- a/litellm/rust_bridge/chat_completions/route_host.py +++ b/litellm/rust_bridge/chat_completions/route_host.py @@ -6,7 +6,7 @@ from typing import Final import litellm from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS from litellm.rust_bridge import failures -from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest +from litellm.rust_bridge.public_call import optional_str from litellm.types.utils import ModelResponse _TRANSPORT_PARAMETERS: Final = frozenset( @@ -37,10 +37,16 @@ def response(value: Mapping[str, object]) -> ModelResponse: return ModelResponse(**value) -def arguments(request: LiteLLMChatCompletionsRequest) -> Mapping[str, object]: - return request.kwargs +def arguments(request: Mapping[str, object]) -> Mapping[str, object]: + return request -def map_failure(error: Exception, request: LiteLLMChatCompletionsRequest) -> Exception: - provider: Final = request.custom_llm_provider or request.model.partition("/")[0] - return failures.map_native_failure(error, request.model, provider, arguments(request), request.api_base) +def map_failure(error: Exception, request: Mapping[str, object]) -> Exception: + provider: Final = optional_str(request.get("custom_llm_provider")) or str(request["model"]).partition("/")[0] + return failures.map_native_failure( + error, + str(request["model"]), + provider, + arguments(request), + optional_str(request.get("api_base")) or optional_str(request.get("base_url")), + ) diff --git a/litellm/rust_bridge/embeddings/entrypoints.py b/litellm/rust_bridge/embeddings/entrypoints.py index da17434df02..2fed4eb512b 100644 --- a/litellm/rust_bridge/embeddings/entrypoints.py +++ b/litellm/rust_bridge/embeddings/entrypoints.py @@ -1,38 +1,24 @@ from __future__ import annotations -from collections.abc import Awaitable, Mapping -from dataclasses import dataclass +from collections.abc import Awaitable from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.public_call import NativeCall from litellm.types.utils import EmbeddingResponse -@dataclass(frozen=True, slots=True) -class LiteLLMEmbeddingRequest: - model: str - input: object - api_key: str | None - api_base: str | None - custom_llm_provider: str | None - kwargs: Mapping[str, object] - - class NativeEmbedding(Protocol): def __call__( self, - request: LiteLLMEmbeddingRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> EmbeddingResponse: ... class NativeAembedding(Protocol): def __call__( self, - request: LiteLLMEmbeddingRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + call: NativeCall, ) -> Awaitable[EmbeddingResponse]: ... diff --git a/litellm/rust_bridge/messages/entrypoints.py b/litellm/rust_bridge/messages/entrypoints.py index d25c906c4c1..c49877bf174 100644 --- a/litellm/rust_bridge/messages/entrypoints.py +++ b/litellm/rust_bridge/messages/entrypoints.py @@ -1,40 +1,24 @@ from __future__ import annotations -from collections.abc import AsyncIterator, Awaitable, Iterator, Mapping, Sequence -from dataclasses import dataclass +from collections.abc import AsyncIterator, Awaitable, Iterator from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.public_call import NativeCall from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse -@dataclass(frozen=True, slots=True) -class LiteLLMMessagesRequest: - model: str - messages: Sequence[object] - max_tokens: int - stream: bool | None - api_key: str | None - api_base: str | None - custom_llm_provider: str | None - kwargs: Mapping[str, object] - - class NativeMessages(Protocol): def __call__( self, - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> AnthropicMessagesResponse | Iterator[bytes]: ... class NativeAmessages(Protocol): def __call__( self, - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> Awaitable[AnthropicMessagesResponse | AsyncIterator[bytes]]: ... diff --git a/litellm/rust_bridge/messages/route_host.py b/litellm/rust_bridge/messages/route_host.py index 19e3126ad82..88a81fdae4a 100644 --- a/litellm/rust_bridge/messages/route_host.py +++ b/litellm/rust_bridge/messages/route_host.py @@ -11,37 +11,14 @@ import litellm from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.llms.anthropic.pass_through.utils import is_reasoning_auto_summary_enabled from litellm.rust_bridge import failures -from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest +from litellm.rust_bridge.public_call import optional_str from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse _DROP_PATHS: Final = TypeAdapter(list[object]) @dataclass(frozen=True, slots=True) -class EffortTiers: - minimal: bool - low: bool - medium: bool - high: bool - xhigh: bool - max: bool - - -@dataclass(frozen=True, slots=True) -class ModelCapabilities: - supports_reasoning: bool - supports_adaptive_thinking: bool - thinking_always_on: bool - supports_legacy_thinking: bool - supports_output_config: bool - supports_sampling_params: bool - supports_speed: bool - effort_tiers: EffortTiers - - -@dataclass(frozen=True, slots=True) -class MessagesShaping: - capabilities: ModelCapabilities +class MessagesSettings: drop_params: bool reasoning_auto_summary: bool additional_drop_params: Sequence[str] @@ -62,56 +39,19 @@ def stream_hidden_params(headers: Sequence[tuple[str, str]]) -> Mapping[str, obj return anthropic_messages_stream_hidden_params(httpx.Headers(list(headers))) -def arguments(request: LiteLLMMessagesRequest) -> Mapping[str, object]: - return request.kwargs +def arguments(request: Mapping[str, object]) -> Mapping[str, object]: + return request -def map_failure(error: Exception, request: LiteLLMMessagesRequest, request_provider: str) -> Exception: +def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception: if getattr(error, "messages_request_error", False): return litellm.BadRequestError( message=str(error), - model=request.model.removeprefix(f"{request_provider}/"), + model=str(request["model"]).removeprefix(f"{request_provider}/"), llm_provider=request_provider, ) - return failures.map_native_failure(error, request.model, request_provider, arguments(request), request.api_base) - - -def _resolved_provider(model: str, custom_llm_provider: str | None) -> tuple[str, str]: - try: - resolved_model, provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider) - except Exception: # noqa: BLE001 # an unroutable model still shapes as a bare Anthropic id - return model, custom_llm_provider or "anthropic" - return resolved_model, provider - - -def model_capabilities(model: str, custom_llm_provider: str | None) -> ModelCapabilities: - from litellm.llms.anthropic.chat.transformation import AnthropicConfig - from litellm.llms.anthropic.common_utils import AnthropicModelInfo - - resolved_model, provider = _resolved_provider(model, custom_llm_provider) - - def supports(flag: str) -> bool: - return AnthropicModelInfo._supports_model_capability(model, flag, provider) # pyright: ignore[reportPrivateUsage] # same probes the Python transform runs; forking them would drift - - def tier(level: str) -> bool: - return AnthropicConfig._supports_effort_level(model, level, provider) # pyright: ignore[reportPrivateUsage] # same probe the Python transform runs - - return ModelCapabilities( - supports_reasoning=supports("supports_reasoning"), - supports_adaptive_thinking=supports("supports_adaptive_thinking"), - thinking_always_on=supports("thinking_always_on"), - supports_legacy_thinking=supports("supports_legacy_thinking"), - supports_output_config=supports("supports_output_config"), - supports_sampling_params=AnthropicModelInfo._supports_sampling_params(resolved_model), # pyright: ignore[reportPrivateUsage] # same gate the handler applies - supports_speed=AnthropicConfig._model_supports_speed_param(resolved_model, provider), # pyright: ignore[reportPrivateUsage] # same gate the handler applies - effort_tiers=EffortTiers( - minimal=tier("minimal"), - low=tier("low"), - medium=tier("medium"), - high=tier("high"), - xhigh=tier("xhigh"), - max=tier("max"), - ), + return failures.map_native_failure( + error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base")) ) @@ -127,10 +67,9 @@ def _additional_drop_params(kwargs: Mapping[str, object]) -> tuple[str, ...]: return tuple(path for path in configured if isinstance(path, str)) -def shaping(model: str, custom_llm_provider: str | None, kwargs: Mapping[str, object]) -> dict[str, object]: +def settings(kwargs: Mapping[str, object]) -> dict[str, object]: return asdict( - MessagesShaping( - capabilities=model_capabilities(model, custom_llm_provider), + MessagesSettings( drop_params=_drop_params(kwargs), reasoning_auto_summary=is_reasoning_auto_summary_enabled(), additional_drop_params=_additional_drop_params(kwargs), diff --git a/litellm/rust_bridge/model_capabilities.py b/litellm/rust_bridge/model_capabilities.py new file mode 100644 index 00000000000..eb8496c722f --- /dev/null +++ b/litellm/rust_bridge/model_capabilities.py @@ -0,0 +1,35 @@ +from __future__ import annotations + +import litellm + + +def _resolved_provider(model: str, custom_llm_provider: str | None) -> tuple[str, str]: + try: + resolved_model, provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider) + except Exception: # noqa: BLE001 # an unroutable model still shapes as a bare Anthropic id + return model, custom_llm_provider or "anthropic" + return resolved_model, provider + + +def anthropic_model_capabilities(model: str, custom_llm_provider: str | None) -> dict[str, object]: + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + resolved_model, provider = _resolved_provider(model, custom_llm_provider) + + def supports(flag: str) -> bool: + return AnthropicModelInfo._supports_model_capability(model, flag, provider) # pyright: ignore[reportPrivateUsage] # same probes the Python transform runs; forking them would drift + + def tier(level: str) -> bool: + return AnthropicConfig._supports_effort_level(model, level, provider) # pyright: ignore[reportPrivateUsage] # same probe the Python transform runs + + return { + "supports_reasoning": supports("supports_reasoning"), + "supports_adaptive_thinking": supports("supports_adaptive_thinking"), + "thinking_always_on": supports("thinking_always_on"), + "supports_legacy_thinking": supports("supports_legacy_thinking"), + "supports_output_config": supports("supports_output_config"), + "supports_sampling_params": AnthropicModelInfo._supports_sampling_params(resolved_model), # pyright: ignore[reportPrivateUsage] # same gate the handler applies + "supports_speed": AnthropicConfig._model_supports_speed_param(resolved_model, provider), # pyright: ignore[reportPrivateUsage] # same gate the handler applies + "effort_tiers": {level: tier(level) for level in ("minimal", "low", "medium", "high", "xhigh", "max")}, + } diff --git a/litellm/rust_bridge/ocr/entrypoints.py b/litellm/rust_bridge/ocr/entrypoints.py index 0c67700de6b..2796f0148e3 100644 --- a/litellm/rust_bridge/ocr/entrypoints.py +++ b/litellm/rust_bridge/ocr/entrypoints.py @@ -1,43 +1,24 @@ from __future__ import annotations from collections.abc import Awaitable, Mapping -from dataclasses import dataclass from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables -import httpx - from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRResponse from litellm.rust_bridge.bindings import NativeBinding - - -@dataclass(frozen=True, slots=True) -class LiteLLMOcrRequest: - model: str - document: Mapping[str, object] - api_key: str | None - api_base: str | None - timeout: float | httpx.Timeout | None - custom_llm_provider: str | None - extra_headers: dict[str, object] | None - kwargs: Mapping[str, object] - input_sources: Mapping[str, str] | None = None +from litellm.rust_bridge.public_call import NativeCall class NativeOcr(Protocol): def __call__( self, - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: ... class NativeAocr(Protocol): def __call__( self, - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> Awaitable[OCRResponse]: ... diff --git a/litellm/rust_bridge/ocr/route_host.py b/litellm/rust_bridge/ocr/route_host.py index bfbd5c11d4e..a7eb0829f5c 100644 --- a/litellm/rust_bridge/ocr/route_host.py +++ b/litellm/rust_bridge/ocr/route_host.py @@ -10,7 +10,7 @@ import litellm from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse from litellm.rust_bridge import failures from litellm.rust_bridge.failures import UpstreamFailure -from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest +from litellm.rust_bridge.public_call import optional_str __all__ = ("UpstreamFailure", "arguments", "map_failure", "response") @@ -27,15 +27,17 @@ def response(value: Mapping[str, object]) -> OCRResponse: return normalized -def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]: - return request.kwargs +def arguments(request: Mapping[str, object]) -> Mapping[str, object]: + return request -def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception: +def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception: if getattr(error, "ocr_request_format_error", False): return litellm.UnsupportedParamsError( - message=f"Invalid `req_format`: {request.kwargs.get('req_format')!r}. Expected 'native' or 'litellm'.", - model=request.model.removeprefix(f"{request_provider}/"), + message=f"Invalid `req_format`: {request.get('req_format')!r}. Expected 'native' or 'litellm'.", + model=str(request["model"]).removeprefix(f"{request_provider}/"), llm_provider=request_provider, ) - return failures.map_native_failure(error, request.model, request_provider, arguments(request), request.api_base) + return failures.map_native_failure( + error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base")) + ) diff --git a/litellm/rust_bridge/public_call.py b/litellm/rust_bridge/public_call.py index 3cf19026de1..a25e9802593 100644 --- a/litellm/rust_bridge/public_call.py +++ b/litellm/rust_bridge/public_call.py @@ -4,7 +4,9 @@ from __future__ import annotations import inspect from collections.abc import Callable, Mapping, Sequence -from typing import Final, cast # noqa: TID251 # narrows caller-owned containers without copying them +from dataclasses import dataclass +from types import MappingProxyType +from typing import Final, TypeVar, cast # noqa: TID251 # narrows caller-owned containers without copying them import litellm @@ -80,3 +82,28 @@ def inference_decline_reason(parameters: tuple[str, ...], kwargs: Mapping[str, o if name not in parameters and name not in _INFERENCE_CONTEXT: return f"native inference does not implement {name}" return None + + +@dataclass(frozen=True, slots=True) +class NativeCall: + args: tuple[object, ...] + kwargs: Mapping[str, object] + bound: Mapping[str, object] + + +def native_call(args: tuple[object, ...], kwargs: Mapping[str, object], fields: Mapping[str, object]) -> NativeCall: + extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({}) + named: Final = {name: value for name, value in fields.items() if name != "kwargs"} + return NativeCall(args=args, kwargs=kwargs, bound=MappingProxyType({**named, **extra})) + + +NativeResultT: Final = TypeVar("NativeResultT") + + +def native_call_hook( + hook: Callable[[NativeCall], NativeResultT], + call: NativeCall, + _args: tuple[object, ...], + _kwargs: Mapping[str, object], +) -> NativeResultT: + return hook(call) diff --git a/litellm/rust_bridge/responses/entrypoints.py b/litellm/rust_bridge/responses/entrypoints.py index d9f9489ac22..e0bb973075c 100644 --- a/litellm/rust_bridge/responses/entrypoints.py +++ b/litellm/rust_bridge/responses/entrypoints.py @@ -1,42 +1,24 @@ from __future__ import annotations -from collections.abc import Awaitable, Mapping -from dataclasses import dataclass, field -from types import MappingProxyType +from collections.abc import Awaitable from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.public_call import NativeCall from litellm.types.llms.openai import ResponsesAPIResponse -@dataclass(frozen=True, slots=True) -class LiteLLMResponsesRequest: - model: str - input: object - stream: bool | None - api_key: str | None - api_base: str | None - custom_llm_provider: str | None - extra_headers: Mapping[str, object] | None - kwargs: Mapping[str, object] - parameters: Mapping[str, object] = field(default_factory=lambda: MappingProxyType({})) - - class NativeResponses(Protocol): def __call__( self, - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> ResponsesAPIResponse: ... class NativeAresponses(Protocol): def __call__( self, - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> Awaitable[ResponsesAPIResponse]: ... diff --git a/litellm/rust_bridge/responses/route_host.py b/litellm/rust_bridge/responses/route_host.py index 4f491064185..00f24244f49 100644 --- a/litellm/rust_bridge/responses/route_host.py +++ b/litellm/rust_bridge/responses/route_host.py @@ -6,8 +6,7 @@ from typing import Final import litellm from litellm import get_llm_provider from litellm.rust_bridge import failures -from litellm.rust_bridge.public_call import inference_decline_reason -from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest +from litellm.rust_bridge.public_call import inference_decline_reason, optional_str from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams, ResponsesAPIResponse PARAMETERS: Final = tuple(ResponsesAPIOptionalRequestParams.__annotations__) @@ -21,21 +20,27 @@ def response(value: Mapping[str, object]) -> ResponsesAPIResponse: return ResponsesAPIResponse.model_validate(value) -def arguments(request: LiteLLMResponsesRequest) -> Mapping[str, object]: - return request.kwargs +def arguments(request: Mapping[str, object]) -> Mapping[str, object]: + return request -def map_failure(error: Exception, request: LiteLLMResponsesRequest) -> Exception: - provider: Final = request.custom_llm_provider or "openai" - return failures.map_native_failure(error, request.model, provider, arguments(request), request.api_base) +def map_failure(error: Exception, request: Mapping[str, object]) -> Exception: + provider: Final = optional_str(request.get("custom_llm_provider")) or "openai" + return failures.map_native_failure( + error, + str(request["model"]), + provider, + arguments(request), + optional_str(request.get("api_base")) or optional_str(request.get("base_url")), + ) -def decline_reason(request: LiteLLMResponsesRequest) -> str | None: - if request.custom_llm_provider is None and "/" not in request.model: +def decline_reason(request: Mapping[str, object]) -> str | None: + if optional_str(request.get("custom_llm_provider")) is None and "/" not in str(request["model"]): try: - _, provider, _, _ = get_llm_provider(model=request.model) + _, provider, _, _ = get_llm_provider(model=str(request["model"])) except litellm.exceptions.BadRequestError: return "native Responses could not resolve the provider" if provider != "openai": return "native HTTP responses provider" - return inference_decline_reason(PARAMETERS, {**request.parameters, **request.kwargs}) + return inference_decline_reason(PARAMETERS, request) diff --git a/litellm/rust_bridge/transcription/native.py b/litellm/rust_bridge/transcription/native.py index 25ee8d362df..550746ba7c3 100644 --- a/litellm/rust_bridge/transcription/native.py +++ b/litellm/rust_bridge/transcription/native.py @@ -4,19 +4,13 @@ from collections.abc import Awaitable from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.public_call import NativeCall class RustTranscription(Protocol): def __call__( self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, + call: NativeCall, ) -> dict[str, object]: raise NotImplementedError @@ -24,14 +18,7 @@ class RustTranscription(Protocol): class RustAtranscription(Protocol): def __call__( self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, + call: NativeCall, ) -> Awaitable[dict[str, object]]: raise NotImplementedError diff --git a/tests/test_litellm_rust/cache/test_python_cache.py b/tests/test_litellm_rust/cache/test_python_cache.py index 4aa59d8b3d1..3a2e3b8c143 100644 --- a/tests/test_litellm_rust/cache/test_python_cache.py +++ b/tests/test_litellm_rust/cache/test_python_cache.py @@ -22,7 +22,7 @@ from litellm.rust_bridge import runtime from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule from litellm.rust_bridge.configuration import Rollout from litellm.rust_bridge.dispatch import call_hook -from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest +from litellm.rust_bridge.public_call import NativeCall from litellm.types.caching import CachingSupportedCallTypes from tests.test_litellm_rust.support.cache import cache_key, collect, invoke, payload from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging @@ -443,15 +443,26 @@ def test_sync_rust_messages_calls_python_cache(recording_server: RecordingServer "api_key": "test-key", "api_base": recording_server.base_url, } - request: Final = LiteLLMMessagesRequest( - MESSAGES_MODEL, list(MESSAGES), 32, None, "test-key", recording_server.base_url, "anthropic", arguments + request: Final = NativeCall( + args=(), + kwargs=arguments, + bound={ + "model": MESSAGES_MODEL, + "messages": list(MESSAGES), + "max_tokens": 32, + "stream": None, + "api_key": "test-key", + "api_base": recording_server.base_url, + "custom_llm_provider": "anthropic", + **arguments, + }, ) def call() -> object: return runtime.run( RouteContext(Route.MESSAGES), binding=NATIVE_MESSAGES, - native=lambda hook: call_hook(hook, request, (), arguments), + native=lambda hook: hook(request), python=runtime.NO_PYTHON, rules=(RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),), ) diff --git a/tests/test_litellm_rust/messages/test_request_shaping.py b/tests/test_litellm_rust/messages/test_request_shaping.py index f885fda5f42..a1c99f245b4 100644 --- a/tests/test_litellm_rust/messages/test_request_shaping.py +++ b/tests/test_litellm_rust/messages/test_request_shaping.py @@ -235,3 +235,83 @@ async def test_non_string_metadata_user_id_is_rejected_before_the_provider_call( await litellm.anthropic.messages.acreate(**arguments(messages_server, metadata={"user_id": 123})) assert messages_server.requests == [] + + +@pytest.mark.asyncio +async def test_native_messages_observes_runtime_capabilities_and_separate_caller_settings( + messages_server: RecordingServer, monkeypatch: pytest.MonkeyPatch +) -> None: + from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES + from litellm.rust_bridge.public_call import NativeCall + + native: Final = NATIVE_AMESSAGES.load() + assert native is not None + model: Final = "claude-test-runtime-capabilities" + messages_server.expected_requests = 2 + request: Final = NativeCall( + args=(), + kwargs={"temperature": 0.2, "drop_params": True}, + bound={ + "model": model, + "messages": MESSAGES, + "max_tokens": 16, + "stream": None, + "api_key": "test-key", + "api_base": messages_server.base_url, + "custom_llm_provider": "anthropic", + "temperature": 0.2, + "drop_params": True, + }, + ) + monkeypatch.setitem( + litellm.model_cost, + model, + { + "litellm_provider": "anthropic", + "mode": "chat", + "supports_sampling_params": True, + }, + ) + first: Final = await native(request) + monkeypatch.setitem( + litellm.model_cost, + model, + { + "litellm_provider": "anthropic", + "mode": "chat", + "supports_sampling_params": False, + }, + ) + second: Final = await native(request) + + assert isinstance(first, dict) + assert isinstance(second, dict) + assert first["id"] == second["id"] == MESSAGES_RESPONSE["id"] + assert len(messages_server.requests) == 2 + first_body: Final = messages_server.requests[0].body + second_body: Final = messages_server.requests[1].body + assert isinstance(first_body, dict) + assert isinstance(second_body, dict) + assert first_body["temperature"] == request.kwargs["temperature"] + assert second_body == {name: value for name, value in first_body.items() if name != "temperature"} + + +@pytest.mark.asyncio +async def test_native_messages_reads_optional_positional_body_parameters(messages_server: RecordingServer) -> None: + from litellm.messages.dispatch import _MESSAGES, _public_request + from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES + + native: Final = NATIVE_AMESSAGES.load() + assert native is not None + metadata: Final = {"user_id": "caller"} + args: Final = (16, MESSAGES, "anthropic/claude-test", metadata, None, False, "Be brief", 0.25) + kwargs: Final = {"api_key": "test-key", "api_base": messages_server.base_url} + call: Final = _public_request(_MESSAGES, args, kwargs) + assert call is not None + + await native(call) + + body, _ = sent(messages_server) + assert body["temperature"] == args[7] + assert body["system"] == args[6] + assert body["metadata"] == metadata diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 18c8861ed6d..40bcb992b22 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -296,7 +296,7 @@ def test_unstarted_native_coroutine_releases_input_without_reading_file(ocr_serv def create(): file: Final = File() kwargs: Final = {"model": "mistral/mistral-ocr-latest", "document": {"type": "file", "file": file}} - coroutine: Final = _native.aocr(_public_request("aocr", (), kwargs), (), kwargs) + coroutine: Final = _native.aocr(_public_request("aocr", (), kwargs)) file.owner = coroutine coroutine.close() return weakref.ref(file) diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index d09e60784fa..1eae999ff70 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -610,7 +610,8 @@ def test_native_projection_errors_never_select_python( from litellm.rust_bridge import runtime, settings from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule from litellm.rust_bridge.configuration import Rollout - from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR, LiteLLMOcrRequest + from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR + from litellm.rust_bridge.public_call import NativeCall ocr_server.expected_requests = 0 snapshot: Final = dataclasses.replace(settings.http_settings(), user_agent=1) @@ -620,15 +621,18 @@ def test_native_projection_errors_never_select_python( monkeypatch.setattr( litellm, "ssl_verify", ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) if failure == "live" else object() ) - request: Final = LiteLLMOcrRequest( - model="mistral/mistral-ocr-latest", - document=OCR_DOCUMENT, - api_key="test-key", - api_base=ocr_server.base_url, - timeout=None, - custom_llm_provider="mistral", - extra_headers=None, + request: Final = NativeCall( + args=(), kwargs={}, + bound={ + "model": "mistral/mistral-ocr-latest", + "document": OCR_DOCUMENT, + "api_key": "test-key", + "api_base": ocr_server.base_url, + "timeout": None, + "custom_llm_provider": "mistral", + "extra_headers": None, + }, ) def python_fallback() -> NoReturn: @@ -638,7 +642,7 @@ def test_native_projection_errors_never_select_python( runtime.run( RouteContext(Route.OCR, provider="mistral"), binding=NATIVE_OCR, - native=lambda native: native(request, (), {}), + native=lambda native: native(request), python=python_fallback, rules=(RouteRule(Route.OCR, Rollout.RUST_REQUIRED if required else Rollout.RUST_OPT_OUT),), ) diff --git a/tests/test_litellm_rust/support/cache.py b/tests/test_litellm_rust/support/cache.py index 564578478cd..352722f3269 100644 --- a/tests/test_litellm_rust/support/cache.py +++ b/tests/test_litellm_rust/support/cache.py @@ -10,11 +10,11 @@ import litellm from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.rust_bridge import runtime from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule -from litellm.rust_bridge.chat_completions.entrypoints import NATIVE_ACOMPLETION, LiteLLMChatCompletionsRequest +from litellm.rust_bridge.chat_completions.entrypoints import NATIVE_ACOMPLETION from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.dispatch import call_hook -from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, LiteLLMMessagesRequest -from litellm.rust_bridge.responses.entrypoints import NATIVE_ARESPONSES, LiteLLMResponsesRequest +from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES +from litellm.rust_bridge.public_call import NativeCall +from litellm.rust_bridge.responses.entrypoints import NATIVE_ARESPONSES from litellm.types.utils import ModelResponse from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_EVENTS, MESSAGES_MODEL, MESSAGES_RESPONSE @@ -48,13 +48,24 @@ async def invoke( arguments: Final = {"model": RESPONSES_MODEL, "input": "hello", **common} if not native: return await litellm.aresponses(**arguments) - request: Final = LiteLLMResponsesRequest( - RESPONSES_MODEL, "hello", None, "test-key", server.base_url, "openai", None, arguments + request: Final = NativeCall( + args=(), + kwargs=arguments, + bound={ + "model": RESPONSES_MODEL, + "input": "hello", + "stream": None, + "api_key": "test-key", + "api_base": server.base_url, + "custom_llm_provider": "openai", + "extra_headers": None, + **arguments, + }, ) return await runtime.arun( RouteContext(Route.RESPONSES), binding=NATIVE_ARESPONSES, - native=lambda hook: call_hook(hook, request, (), arguments), + native=lambda hook: hook(request), python=runtime.NO_PYTHON, rules=(RouteRule(Route.RESPONSES, Rollout.RUST_REQUIRED),), ) @@ -67,25 +78,34 @@ async def invoke( if route == "chat": if not native: return await litellm.acompletion(**parameters) - chat: Final = LiteLLMChatCompletionsRequest( - MESSAGES_MODEL, list(MESSAGES), None, "test-key", server.base_url, None, None, parameters - ) + chat: Final = NativeCall(args=(), kwargs=parameters, bound=parameters) return await runtime.arun( RouteContext(Route.CHAT_COMPLETIONS), binding=NATIVE_ACOMPLETION, - native=lambda hook: call_hook(hook, chat, (), parameters), + native=lambda hook: hook(chat), python=runtime.NO_PYTHON, rules=(RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),), ) if not native: return await litellm.anthropic_messages(**parameters) - messages: Final = LiteLLMMessagesRequest( - MESSAGES_MODEL, list(MESSAGES), 32, None, "test-key", server.base_url, "anthropic", parameters + messages: Final = NativeCall( + args=(), + kwargs=parameters, + bound={ + "model": MESSAGES_MODEL, + "messages": list(MESSAGES), + "max_tokens": 32, + "stream": None, + "api_key": "test-key", + "api_base": server.base_url, + "custom_llm_provider": "anthropic", + **parameters, + }, ) return await runtime.arun( RouteContext(Route.MESSAGES), binding=NATIVE_AMESSAGES, - native=lambda hook: call_hook(hook, messages, (), parameters), + native=lambda hook: hook(messages), python=runtime.NO_PYTHON, rules=(RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),), ) diff --git a/tests/test_litellm_rust/test_inference.py b/tests/test_litellm_rust/test_inference.py index f1b6071a844..5e4b7b284bb 100644 --- a/tests/test_litellm_rust/test_inference.py +++ b/tests/test_litellm_rust/test_inference.py @@ -7,12 +7,12 @@ from pydantic import JsonValue, TypeAdapter import litellm from litellm import RateLimitError +from litellm.chat_completions import dispatch as chat_dispatch from litellm.integrations.custom_logger import CustomLogger from litellm.models.credentials import CredentialItem from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.rust_bridge import _native -from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest -from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest +from litellm.rust_bridge.public_call import NativeCall from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import CallTypes, ModelResponse from tests.test_litellm_rust.support.callback_recorder import RecordingLogger @@ -66,10 +66,8 @@ def native_call( "max_tokens": 32, **options, } - request: Final = LiteLLMChatCompletionsRequest( - MESSAGES_MODEL, list(MESSAGES), None, "test-key", server.base_url, None, None, kwargs - ) - return (_native.acompletion if asynchronous else _native.completion)(request, (), kwargs) + request: Final = NativeCall(args=(), kwargs=kwargs, bound=kwargs) + return (_native.acompletion if asynchronous else _native.completion)(request) response_kwargs: Final = { "model": RESPONSES_MODEL, "input": "hello", @@ -78,10 +76,21 @@ def native_call( "max_output_tokens": 32, **options, } - response_request: Final = LiteLLMResponsesRequest( - RESPONSES_MODEL, "hello", None, "test-key", server.base_url, "openai", None, response_kwargs + response_request: Final = NativeCall( + args=(), + kwargs=response_kwargs, + bound={ + "model": RESPONSES_MODEL, + "input": "hello", + "stream": None, + "api_key": "test-key", + "api_base": server.base_url, + "custom_llm_provider": "openai", + "extra_headers": None, + **response_kwargs, + }, ) - return (_native.aresponses if asynchronous else _native.responses)(response_request, (), response_kwargs) + return (_native.aresponses if asynchronous else _native.responses)(response_request) async def execute(route: Route, asynchronous: bool, server: RecordingServer, options: Mapping[str, object]) -> object: @@ -284,13 +293,13 @@ async def test_native_projection_reads_positional_parameters(route: Route, recor args: Final = (MESSAGES_MODEL, list(MESSAGES), 12.0, 0.25) request: Final = chat_dispatch.request(args, kwargs) assert request is not None - await asyncio.to_thread(_native.completion, request, args, kwargs) + await asyncio.to_thread(_native.completion, request) assert _OBJECT.validate_python(recording_server.requests[0].body)["temperature"] == 0.25 else: response_args: Final = ("hello", RESPONSES_MODEL, None, "Be brief", 16) response_request: Final = responses_dispatch.request(response_args, kwargs) assert response_request is not None - await asyncio.to_thread(_native.responses, response_request, response_args, kwargs) + await asyncio.to_thread(_native.responses, response_request) body: Final = _OBJECT.validate_python(recording_server.requests[0].body) assert body["instructions"] == "Be brief" assert body["max_output_tokens"] == 16 @@ -351,3 +360,24 @@ async def test_native_responses_decode_continuation_ids( ) await execute("responses", asynchronous, recording_server, {"previous_response_id": previous}) assert _OBJECT.validate_python(recording_server.requests[0].body)["previous_response_id"] == original + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_native_chat_uses_bound_positional_parameters( + asynchronous: bool, recording_server: RecordingServer +) -> None: + recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE) + arguments: Final = (MESSAGES_MODEL, list(MESSAGES), 12.0, 0.35) + supplied: Final = {"base_url": recording_server.base_url, "api_key": "test-key", "max_tokens": 32} + call: Final = chat_dispatch._DISPATCH.request(arguments, supplied) # pyright: ignore[reportPrivateUsage] # exercise the native request produced by public binding + assert call is not None + result: Final = ( + await _native.acompletion(call) if asynchronous else await asyncio.to_thread(_native.completion, call) + ) + assert isinstance(result, ModelResponse) + assert result.choices[0].message.content == "Hello from native Messages" + assert len(recording_server.requests) == 1 + body: Final = _OBJECT.validate_python(recording_server.requests[0].body) + assert body["temperature"] == arguments[3] + assert body["max_tokens"] == supplied["max_tokens"] diff --git a/tests/unit/chat_completions/test_dispatch.py b/tests/unit/chat_completions/test_dispatch.py index c9274321a0f..fbfc4875aa1 100644 --- a/tests/unit/chat_completions/test_dispatch.py +++ b/tests/unit/chat_completions/test_dispatch.py @@ -4,24 +4,23 @@ from typing import Final, cast # noqa: TID251 # narrows legacy callable signat import pytest import litellm +from litellm.chat_completions import dispatch from litellm.chat_completions.dispatch import ( _ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch _DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch ) from litellm.rust_bridge import catalog from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.catalog import Route, RouteRule +from litellm.rust_bridge.catalog import Route, RouteRule, Rules from litellm.rust_bridge.chat_completions.entrypoints import ( NATIVE_ACOMPLETION, NATIVE_COMPLETION, - LiteLLMChatCompletionsRequest, NativeAcompletion, NativeCompletion, ) from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge.public_call import NativeCall, native_call_hook from litellm.types.utils import ModelResponse -from litellm.chat_completions import dispatch -from litellm.rust_bridge.catalog import Rules MESSAGES: Final = [{"role": "user", "content": "hi"}] PYTHON_RULES: Final = () @@ -51,9 +50,7 @@ def test_python_route_forwards_original_call_shape() -> None: captured.append((call_args, call_kwargs)) return response - def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: + def native(request: NativeCall) -> ModelResponse: pytest.fail("Python-only dispatch must not call native") assert ( @@ -62,7 +59,7 @@ def test_python_route_forwards_original_call_shape() -> None: kwargs, python=python, binding=completion_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=native_call_hook, rules=PYTHON_RULES, ) is response @@ -87,9 +84,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None: captured.append((call_args, call_kwargs)) return response - async def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: + async def native(request: NativeCall) -> ModelResponse: pytest.fail("Python-only dispatch must not call native") result: Final = await _ADISPATCH.arun( @@ -97,7 +92,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None: kwargs, python=python, binding=acompletion_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=native_call_hook, rules=PYTHON_RULES, ) assert result is response @@ -119,15 +114,13 @@ def test_native_receives_bound_request_and_original_call_shape() -> None: "custom_llm_provider": "anthropic", "metadata": metadata, } - captured: Final[list[tuple[LiteLLMChatCompletionsRequest, tuple[object, ...], Mapping[str, object]]]] = [] + captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = [] def python(*call_args: object, **call_kwargs: object) -> ModelResponse: # kwargs-ok: rejected Rust fallback pytest.fail("Required Rust dispatch must not call Python") - def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: - captured.append((request, args, kwargs)) + def native(request: NativeCall) -> ModelResponse: + captured.append((request, request.args, request.kwargs)) return ModelResponse() args: Final[tuple[object, ...]] = ("anthropic/claude-sonnet-4-5", MESSAGES) @@ -136,19 +129,19 @@ def test_native_receives_bound_request_and_original_call_shape() -> None: kwargs, python=python, binding=completion_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=native_call_hook, rules=RUST_RULES, ) request, call_args, call_kwargs = captured[0] - assert request.model == "anthropic/claude-sonnet-4-5" - assert request.messages is MESSAGES - assert request.stream is True - assert request.api_key == "sk-test" - assert request.api_base == "https://example.invalid" - assert request.custom_llm_provider == "anthropic" - assert request.extra_headers == {"x-test": "1"} - assert request.kwargs == {"custom_llm_provider": "anthropic", "metadata": metadata} + assert request.bound["model"] == "anthropic/claude-sonnet-4-5" + assert request.bound["messages"] is MESSAGES + assert request.bound["stream"] is True + assert request.bound["api_key"] == "sk-test" + assert request.bound["base_url"] == "https://example.invalid" + assert request.bound["custom_llm_provider"] == "anthropic" + assert request.bound["extra_headers"] == {"x-test": "1"} + assert request.kwargs is kwargs assert call_args == args assert call_kwargs == kwargs assert call_kwargs["metadata"] is metadata @@ -162,9 +155,7 @@ def test_internal_async_marker_bypasses_native() -> None: called.append(True) return response - def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: + def native(request: NativeCall) -> ModelResponse: pytest.fail("acompletion's inner completion call must stay on Python") result: Final = _DISPATCH.run( @@ -172,7 +163,7 @@ def test_internal_async_marker_bypasses_native() -> None: {"acompletion": True}, python=python, binding=completion_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=native_call_hook, rules=RUST_RULES, ) assert result is response @@ -194,9 +185,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map captured.append((call_args, call_kwargs)) return response - def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: + def native(request: NativeCall) -> ModelResponse: pytest.fail("Binding failures must be delegated to Python") assert ( @@ -205,7 +194,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map kwargs, python=python, binding=completion_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=native_call_hook, rules=RUST_RULES, ) is response @@ -214,14 +203,10 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map def test_public_completion_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None: - captured: Final[list[LiteLLMChatCompletionsRequest]] = [] + captured: Final[list[NativeCall]] = [] expected: Final = ModelResponse() - def native( - request: LiteLLMChatCompletionsRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], - ) -> ModelResponse: + def native(request: NativeCall) -> ModelResponse: captured.append(request) return expected @@ -233,19 +218,15 @@ def test_public_completion_routes_through_dispatch(monkeypatch: pytest.MonkeyPat finally: NATIVE_COMPLETION.reset() assert result is expected - assert [request.model for request in captured] == ["gpt-4o"] + assert [request.bound["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio async def test_public_acompletion_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None: - captured: Final[list[LiteLLMChatCompletionsRequest]] = [] + captured: Final[list[NativeCall]] = [] expected: Final = ModelResponse() - async def native( - request: LiteLLMChatCompletionsRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], - ) -> ModelResponse: + async def native(request: NativeCall) -> ModelResponse: captured.append(request) return expected @@ -257,7 +238,7 @@ async def test_public_acompletion_routes_through_dispatch(monkeypatch: pytest.Mo finally: NATIVE_ACOMPLETION.reset() assert result is expected - assert [request.model for request in captured] == ["gpt-4o"] + assert [request.bound["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio @@ -275,13 +256,11 @@ def test_sync_completion_request_projects_public_arguments() -> None: rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),) expected: Final = ModelResponse() - def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: - assert request.model == "test-model" - assert request.messages == MESSAGES - assert request.custom_llm_provider == "openai" - assert request.stream is True + def native(request: NativeCall) -> ModelResponse: + assert request.bound["model"] == "test-model" + assert request.bound["messages"] == MESSAGES + assert request.bound["custom_llm_provider"] == "openai" + assert request.bound["stream"] is True return expected binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None) @@ -291,7 +270,7 @@ def test_sync_completion_request_projects_public_arguments() -> None: {"custom_llm_provider": "openai", "stream": True}, python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"), binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=native_call_hook, rules=rules, ) @@ -309,9 +288,7 @@ async def test_async_completion_falls_back_after_native_declines() -> None: expected: Final = ModelResponse() rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_OPT_OUT),) - async def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: + async def native(request: NativeCall) -> ModelResponse: raise declined("unsupported") async def python(*args: object, **kwargs: object) -> ModelResponse: @@ -324,7 +301,7 @@ async def test_async_completion_falls_back_after_native_declines() -> None: {}, python=python, binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=native_call_hook, rules=rules, ) @@ -338,9 +315,7 @@ def test_internal_acompletion_marker_bypasses_native() -> None: def python(*args: object, **kwargs: object) -> ModelResponse: return expected - def native( - request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> ModelResponse: + def native(request: NativeCall) -> ModelResponse: pytest.fail("acompletion's inner completion call must stay on Python") binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None) @@ -350,7 +325,7 @@ def test_internal_acompletion_marker_bypasses_native() -> None: {"custom_llm_provider": "openai", "acompletion": True}, python=python, binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=native_call_hook, rules=rules, ) @@ -360,6 +335,6 @@ def test_internal_acompletion_marker_bypasses_native() -> None: def test_positional_parameters_remain_available_to_native_projection() -> None: request: Final = _DISPATCH.request(("anthropic/test-model", MESSAGES, 12.0, 0.25), {}) assert request is not None - assert request.parameters["timeout"] == 12.0 - assert request.parameters["temperature"] == 0.25 - assert request.messages is MESSAGES + assert request.bound["timeout"] == 12.0 + assert request.bound["temperature"] == 0.25 + assert request.bound["messages"] is MESSAGES diff --git a/tests/unit/embeddings/test_dispatch.py b/tests/unit/embeddings/test_dispatch.py index 1062c320cbb..88a2e7532c2 100644 --- a/tests/unit/embeddings/test_dispatch.py +++ b/tests/unit/embeddings/test_dispatch.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable from typing import Final import pytest @@ -11,7 +11,7 @@ from litellm.embeddings import dispatch from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.catalog import Route, RouteRule, Rules from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.embeddings.entrypoints import LiteLLMEmbeddingRequest +from litellm.rust_bridge.public_call import NativeCall, native_call_hook from litellm.types.utils import EmbeddingResponse @@ -32,24 +32,22 @@ def test_sync_embedding_request_projects_public_arguments() -> None: rules: Final[Rules] = (RouteRule(Route.EMBEDDINGS, Rollout.RUST_REQUIRED),) expected: Final = EmbeddingResponse(model="test-model", data=[]) - def native( - request: LiteLLMEmbeddingRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> EmbeddingResponse: - assert request.model == "test-model" - assert request.input == "hello" - assert request.custom_llm_provider == "openai" + def native(request: NativeCall) -> EmbeddingResponse: + assert request.bound["model"] == "test-model" + assert request.bound["input"] == "hello" + assert request.bound["custom_llm_provider"] == "openai" return expected - binding: Final[ - NativeBinding[Callable[[LiteLLMEmbeddingRequest, tuple[object, ...], Mapping[str, object]], EmbeddingResponse]] - ] = NativeBinding("embedding", validate=lambda _: None) + binding: Final[NativeBinding[Callable[[NativeCall], EmbeddingResponse]]] = NativeBinding( + "embedding", validate=lambda _: None + ) binding.override(native) response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision ("test-model", "hello"), {"custom_llm_provider": "openai", "dimensions": 8}, python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"), binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=native_call_hook, rules=rules, ) @@ -67,26 +65,22 @@ async def test_async_embedding_falls_back_after_native_declines() -> None: expected: Final = EmbeddingResponse(model="test-model", data=[]) rules: Final[Rules] = (RouteRule(Route.EMBEDDINGS, Rollout.RUST_OPT_OUT),) - async def native( - request: LiteLLMEmbeddingRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> EmbeddingResponse: + async def native(request: NativeCall) -> EmbeddingResponse: raise declined("unsupported") async def python(*args: object, **kwargs: object) -> EmbeddingResponse: return expected - binding: Final[ - NativeBinding[ - Callable[[LiteLLMEmbeddingRequest, tuple[object, ...], Mapping[str, object]], Awaitable[EmbeddingResponse]] - ] - ] = NativeBinding("aembedding", validate=lambda _: None) + binding: Final[NativeBinding[Callable[[NativeCall], Awaitable[EmbeddingResponse]]]] = NativeBinding( + "aembedding", validate=lambda _: None + ) binding.override(native) response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision ("test-model", "hello"), {}, python=python, binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=native_call_hook, rules=rules, ) diff --git a/tests/unit/messages/test_dispatch.py b/tests/unit/messages/test_dispatch.py index bf2f373d35b..8668e11ee65 100644 --- a/tests/unit/messages/test_dispatch.py +++ b/tests/unit/messages/test_dispatch.py @@ -3,9 +3,11 @@ from collections.abc import Awaitable, Callable, Mapping from typing import Final, cast # noqa: TID251 # narrows legacy callable signatures for inspect import pytest +from pydantic import TypeAdapter import litellm from litellm.llms.anthropic.pass_through.messages import handler as python_messages +from litellm.messages import dispatch from litellm.messages.dispatch import ( _ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch _DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch @@ -14,16 +16,9 @@ from litellm.rust_bridge import catalog from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.catalog import Route, RouteRule, Rules from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.messages.entrypoints import ( - NATIVE_AMESSAGES, - NATIVE_MESSAGES, - LiteLLMMessagesRequest, - NativeAmessages, - NativeMessages, -) +from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, NATIVE_MESSAGES, NativeAmessages, NativeMessages +from litellm.rust_bridge.public_call import NativeCall from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse -from pydantic import TypeAdapter -from litellm.messages import dispatch MESSAGES: Final = [{"role": "user", "content": "hi"}] PYTHON_RULES: Final[Rules] = () @@ -67,9 +62,7 @@ def test_python_route_forwards_original_call_shape() -> None: return expected def native( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> AnthropicMessagesResponse: pytest.fail("Python-only dispatch must not call native") @@ -78,7 +71,7 @@ def test_python_route_forwards_original_call_shape() -> None: kwargs, python=python, binding=messages_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=PYTHON_RULES, ) assert result is expected @@ -106,9 +99,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None: return expected async def native( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> AnthropicMessagesResponse: pytest.fail("Python-only dispatch must not call native") @@ -117,7 +108,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None: kwargs, python=python, binding=amessages_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=PYTHON_RULES, ) assert result is expected @@ -139,17 +130,17 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None: "custom_llm_provider": "anthropic", "litellm_metadata": metadata, } - captured: Final[list[tuple[LiteLLMMessagesRequest, tuple[object, ...], Mapping[str, object]]]] = [] + captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = [] expected: Final = response("anthropic/claude-sonnet-4-5") def python(*call_args: object, **call_kwargs: object) -> AnthropicMessagesResponse: # kwargs-ok: rejected fallback pytest.fail("Required Rust dispatch must not call Python") def native( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> AnthropicMessagesResponse: + args: Final = request.args + kwargs: Final = request.kwargs captured.append((request, args, kwargs)) return expected @@ -158,19 +149,19 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None: kwargs, python=python, binding=messages_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) assert result is expected request, call_args, call_kwargs = captured[0] - assert request.model == "anthropic/claude-sonnet-4-5" - assert request.messages is MESSAGES - assert request.max_tokens == 16 - assert request.stream is True - assert request.api_key == "sk-test" - assert request.api_base == "https://example.invalid" - assert request.custom_llm_provider == "anthropic" - assert request.kwargs == {"litellm_metadata": metadata} + assert request.bound["model"] == "anthropic/claude-sonnet-4-5" + assert request.bound["messages"] is MESSAGES + assert request.bound["max_tokens"] == 16 + assert request.bound["stream"] is True + assert request.bound["api_key"] == "sk-test" + assert request.bound["api_base"] == "https://example.invalid" + assert request.bound["custom_llm_provider"] == "anthropic" + assert request.kwargs == kwargs assert request.kwargs["litellm_metadata"] is metadata assert call_args == args assert call_args[1] is MESSAGES @@ -189,9 +180,7 @@ def test_internal_async_marker_bypasses_native() -> None: return expected def native( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> AnthropicMessagesResponse: pytest.fail("The async handler's inner sync call must stay on Python") @@ -200,7 +189,7 @@ def test_internal_async_marker_bypasses_native() -> None: kwargs, python=python, binding=messages_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) assert result is expected @@ -223,9 +212,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map return expected def native( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> AnthropicMessagesResponse: pytest.fail("Binding failures must be delegated to Python") @@ -234,7 +221,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map kwargs, python=python, binding=messages_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) assert result is expected @@ -242,13 +229,11 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map def test_anthropic_create_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None: - captured: Final[list[LiteLLMMessagesRequest]] = [] + captured: Final[list[NativeCall]] = [] expected: Final = response() def native( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> AnthropicMessagesResponse: captured.append(request) return expected @@ -261,18 +246,16 @@ def test_anthropic_create_routes_through_dispatch(monkeypatch: pytest.MonkeyPatc finally: NATIVE_MESSAGES.reset() assert result is expected - assert [request.model for request in captured] == ["claude-sonnet-4-5"] + assert [request.bound["model"] for request in captured] == ["claude-sonnet-4-5"] @pytest.mark.asyncio async def test_anthropic_acreate_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None: - captured: Final[list[LiteLLMMessagesRequest]] = [] + captured: Final[list[NativeCall]] = [] expected: Final = response() async def native( - request: LiteLLMMessagesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> AnthropicMessagesResponse: captured.append(request) return expected @@ -285,7 +268,7 @@ async def test_anthropic_acreate_routes_through_dispatch(monkeypatch: pytest.Mon finally: NATIVE_AMESSAGES.reset() assert result is expected - assert [request.model for request in captured] == ["claude-sonnet-4-5"] + assert [request.bound["model"] for request in captured] == ["claude-sonnet-4-5"] @pytest.mark.asyncio @@ -303,13 +286,11 @@ def test_sync_messages_request_projects_public_arguments() -> None: rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),) expected: Final = AnthropicMessagesResponse(model="claude-test") - def native( - request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> AnthropicMessagesResponse: - assert request.model == "claude-test" - assert request.messages == MESSAGES - assert request.max_tokens == 10 - assert request.custom_llm_provider == "anthropic" + def native(request: NativeCall) -> AnthropicMessagesResponse: + assert request.bound["model"] == "claude-test" + assert request.bound["messages"] == MESSAGES + assert request.bound["max_tokens"] == 10 + assert request.bound["custom_llm_provider"] == "anthropic" return expected binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) @@ -324,7 +305,7 @@ def test_sync_messages_request_projects_public_arguments() -> None: }, python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"), binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=lambda hook, request, args, kwargs: hook(request), rules=rules, ) @@ -338,9 +319,7 @@ def test_messages_binding_error_delegates_unchanged_to_python() -> None: def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: return expected - def native( - request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> AnthropicMessagesResponse: + def native(request: NativeCall) -> AnthropicMessagesResponse: pytest.fail("a call without max_tokens cannot project a request and must stay on Python") binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) @@ -350,7 +329,7 @@ def test_messages_binding_error_delegates_unchanged_to_python() -> None: {"model": "claude-test", "messages": MESSAGES, "custom_llm_provider": "anthropic"}, python=python, binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=lambda hook, request, args, kwargs: hook(request), rules=rules, ) @@ -368,9 +347,7 @@ async def test_async_messages_falls_back_after_native_declines() -> None: expected: Final = AnthropicMessagesResponse(model="claude-test") rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_OPT_OUT),) - async def native( - request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> AnthropicMessagesResponse: + async def native(request: NativeCall) -> AnthropicMessagesResponse: raise declined("unsupported") async def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: @@ -383,7 +360,7 @@ async def test_async_messages_falls_back_after_native_declines() -> None: {"model": "claude-test", "messages": MESSAGES, "max_tokens": 10}, python=python, binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=lambda hook, request, args, kwargs: hook(request), rules=rules, ) @@ -397,9 +374,7 @@ def test_internal_is_async_marker_bypasses_native() -> None: def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse: return expected - def native( - request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object] - ) -> AnthropicMessagesResponse: + def native(request: NativeCall) -> AnthropicMessagesResponse: pytest.fail("anthropic_messages' inner handler call must stay on Python") binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) @@ -415,7 +390,7 @@ def test_internal_is_async_marker_bypasses_native() -> None: }, python=python, binding=binding, - native=lambda hook, request, args, kwargs: hook(request, args, kwargs), + native=lambda hook, request, args, kwargs: hook(request), rules=rules, ) diff --git a/tests/unit/ocr/test_dispatch.py b/tests/unit/ocr/test_dispatch.py index 531f392b17a..3b0d75ae07a 100644 --- a/tests/unit/ocr/test_dispatch.py +++ b/tests/unit/ocr/test_dispatch.py @@ -14,13 +14,8 @@ from litellm.rust_bridge import catalog, runtime from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.catalog import Route, RouteRule, Rules from litellm.rust_bridge.configuration import Rollout -from litellm.rust_bridge.ocr.entrypoints import ( - NATIVE_AOCR, - NATIVE_OCR, - LiteLLMOcrRequest, - NativeAocr, - NativeOcr, -) +from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, NativeAocr, NativeOcr +from litellm.rust_bridge.public_call import NativeCall RUST_RULES: Final[Rules] = (RouteRule(Route.OCR, Rollout.RUST_REQUIRED),) @@ -55,14 +50,14 @@ def test_native_receives_normalized_positional_request_and_original_call_shape() "extra_headers": extra_headers, "pages": pages, } - captured: Final[list[tuple[LiteLLMOcrRequest, tuple[object, ...], Mapping[str, object]]]] = [] + captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = [] expected: Final = response() def native( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: + args: Final = request.args + kwargs: Final = request.kwargs captured.append((request, args, kwargs)) return expected @@ -71,20 +66,20 @@ def test_native_receives_normalized_positional_request_and_original_call_shape() kwargs, python=runtime.NO_PYTHON, binding=ocr_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) request, call_args, call_kwargs = captured[0] assert result is expected - assert request.model == "mistral/mistral-ocr-latest" - assert request.document is document - assert request.api_key == "test-key" - assert request.api_base == "https://example.invalid" - assert request.timeout is timeout - assert request.custom_llm_provider == "mistral" - assert request.extra_headers is extra_headers - assert request.kwargs == {"pages": pages} + assert request.bound["model"] == "mistral/mistral-ocr-latest" + assert request.bound["document"] is document + assert request.bound["api_key"] == "test-key" + assert request.bound["api_base"] == "https://example.invalid" + assert request.bound["timeout"] is timeout + assert request.bound["custom_llm_provider"] == "mistral" + assert request.bound["extra_headers"] is extra_headers + assert request.kwargs == kwargs assert request.kwargs["pages"] is pages assert call_args is args assert call_kwargs is kwargs @@ -102,14 +97,14 @@ def test_native_preserves_keyword_model_and_document_in_original_call_shape() -> "document": document, "pages": pages, } - captured: Final[list[tuple[LiteLLMOcrRequest, tuple[object, ...], Mapping[str, object]]]] = [] + captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = [] expected: Final = response() def native( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: + args: Final = request.args + kwargs: Final = request.kwargs captured.append((request, args, kwargs)) return expected @@ -118,15 +113,15 @@ def test_native_preserves_keyword_model_and_document_in_original_call_shape() -> kwargs, python=runtime.NO_PYTHON, binding=ocr_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) request, call_args, call_kwargs = captured[0] assert result is expected - assert request.model == "mistral/mistral-ocr-latest" - assert request.document is document - assert request.kwargs == {"pages": pages} + assert request.bound["model"] == "mistral/mistral-ocr-latest" + assert request.bound["document"] is document + assert request.kwargs == kwargs assert call_args is args assert call_kwargs is kwargs assert call_kwargs["model"] == "mistral/mistral-ocr-latest" @@ -139,9 +134,7 @@ def test_aocr_marker_cannot_be_served_without_python() -> None: kwargs: Final[Mapping[str, object]] = {"aocr": True} def native( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: pytest.fail("the aocr bypass marker must not reach native") @@ -151,7 +144,7 @@ def test_aocr_marker_cannot_be_served_without_python() -> None: kwargs, python=runtime.NO_PYTHON, binding=ocr_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) @@ -165,7 +158,7 @@ def test_missing_native_binding_is_a_required_rust_error() -> None: {}, python=runtime.NO_PYTHON, binding=ocr_binding(None), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) @@ -174,9 +167,7 @@ def test_non_required_rule_cannot_be_served_without_python() -> None: args: Final[tuple[object, ...]] = ("mistral/mistral-ocr-latest", {"type": "file", "file": b"pdf"}) def native( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: return response() @@ -186,7 +177,7 @@ def test_non_required_rule_cannot_be_served_without_python() -> None: {}, python=runtime.NO_PYTHON, binding=ocr_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=(RouteRule(Route.OCR, Rollout.PYTHON_ONLY),), ) @@ -206,13 +197,9 @@ def test_non_required_rule_cannot_be_served_without_python() -> None: ), ), ) -def test_ocr_parser_errors_before_native( - args: tuple[object, ...], kwargs: Mapping[str, object], message: str -) -> None: +def test_ocr_parser_errors_before_native(args: tuple[object, ...], kwargs: Mapping[str, object], message: str) -> None: def native( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: pytest.fail("OCR parser failures must not call native") @@ -222,7 +209,7 @@ def test_ocr_parser_errors_before_native( kwargs, python=runtime.NO_PYTHON, binding=ocr_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) @@ -247,9 +234,7 @@ async def test_aocr_parser_errors_before_native( args: tuple[object, ...], kwargs: Mapping[str, object], message: str ) -> None: async def native( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: pytest.fail("OCR parser failures must not call native") @@ -259,7 +244,7 @@ async def test_aocr_parser_errors_before_native( kwargs, python=runtime.NO_PYTHON, binding=aocr_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) @@ -269,13 +254,11 @@ def test_public_ocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> "type": "document_url", "document_url": "https://example.invalid/document.pdf", } - captured: Final[list[LiteLLMOcrRequest]] = [] + captured: Final[list[NativeCall]] = [] expected: Final = response() def native( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: captured.append(request) return expected @@ -288,7 +271,7 @@ def test_public_ocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> finally: NATIVE_OCR.reset() assert result is expected - assert [request.model for request in captured] == ["mistral/mistral-ocr-latest"] + assert [request.bound["model"] for request in captured] == ["mistral/mistral-ocr-latest"] @pytest.mark.asyncio @@ -297,13 +280,11 @@ async def test_public_aocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPat "type": "document_url", "document_url": "https://example.invalid/document.pdf", } - captured: Final[list[LiteLLMOcrRequest]] = [] + captured: Final[list[NativeCall]] = [] expected: Final = response() async def native( - request: LiteLLMOcrRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> OCRResponse: captured.append(request) return expected @@ -316,4 +297,4 @@ async def test_public_aocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPat finally: NATIVE_AOCR.reset() assert result is expected - assert [request.model for request in captured] == ["mistral/mistral-ocr-latest"] + assert [request.bound["model"] for request in captured] == ["mistral/mistral-ocr-latest"] diff --git a/tests/unit/responses/test_dispatch.py b/tests/unit/responses/test_dispatch.py index 637d4bc0a1e..befe2d0000e 100644 --- a/tests/unit/responses/test_dispatch.py +++ b/tests/unit/responses/test_dispatch.py @@ -15,10 +15,10 @@ from litellm.rust_bridge import catalog from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.catalog import Route, RouteRule from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge.public_call import NativeCall from litellm.rust_bridge.responses.entrypoints import ( NATIVE_ARESPONSES, NATIVE_RESPONSES, - LiteLLMResponsesRequest, NativeAresponses, NativeResponses, ) @@ -68,9 +68,7 @@ def test_python_route_forwards_original_call_shape() -> None: return response def native( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> ResponsesAPIResponse: pytest.fail("Python-only dispatch must not call native") @@ -80,7 +78,7 @@ def test_python_route_forwards_original_call_shape() -> None: kwargs, python=python, binding=responses_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=PYTHON_RULES, ) is response @@ -109,9 +107,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None: return response async def native( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> ResponsesAPIResponse: pytest.fail("Python-only dispatch must not call native") @@ -120,7 +116,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None: kwargs, python=python, binding=aresponses_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=PYTHON_RULES, ) assert result is response @@ -144,17 +140,17 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None: "custom_llm_provider": "anthropic", "litellm_metadata": metadata, } - captured: Final[list[tuple[LiteLLMResponsesRequest, tuple[object, ...], Mapping[str, object]]]] = [] + captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = [] response: Final = _response("anthropic/claude-sonnet-4-5") def python(*call_args: object, **call_kwargs: object) -> ResponsesAPIResponse: # kwargs-ok: rejected fallback pytest.fail("Required Rust dispatch must not call Python") def native( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> ResponsesAPIResponse: + args: Final = request.args + kwargs: Final = request.kwargs captured.append((request, args, kwargs)) return response @@ -163,24 +159,20 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None: kwargs, python=python, binding=responses_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) request, call_args, call_kwargs = captured[0] assert result is response - assert request.model == "anthropic/claude-sonnet-4-5" - assert request.input is INPUT - assert request.stream is True - assert request.api_key == "sk-test" - assert request.api_base == "https://example.invalid" - assert request.custom_llm_provider == "anthropic" - assert request.extra_headers is extra_headers - assert request.kwargs == { - "api_key": "sk-test", - "base_url": "https://example.invalid", - "litellm_metadata": metadata, - } + assert request.bound["model"] == "anthropic/claude-sonnet-4-5" + assert request.bound["input"] is INPUT + assert request.bound["stream"] is True + assert request.bound["api_key"] == "sk-test" + assert request.bound["base_url"] == "https://example.invalid" + assert request.bound["custom_llm_provider"] == "anthropic" + assert request.bound["extra_headers"] is extra_headers + assert request.kwargs == kwargs assert request.kwargs["litellm_metadata"] is metadata assert call_args == args assert call_args[0] is INPUT @@ -200,9 +192,7 @@ def test_internal_async_marker_bypasses_native() -> None: return response def native( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> ResponsesAPIResponse: pytest.fail("aresponses' inner responses call must stay on Python") @@ -212,7 +202,7 @@ def test_internal_async_marker_bypasses_native() -> None: kwargs, python=python, binding=responses_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) is response @@ -236,9 +226,7 @@ def test_binding_errors_delegate_unchanged_to_python(args: tuple[object, ...], k return response def native( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> ResponsesAPIResponse: pytest.fail("Binding failures must be delegated to Python") @@ -248,7 +236,7 @@ def test_binding_errors_delegate_unchanged_to_python(args: tuple[object, ...], k kwargs, python=python, binding=responses_binding(native), - native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs), + native=lambda hook, request, call_args, call_kwargs: hook(request), rules=RUST_RULES, ) is response @@ -257,13 +245,11 @@ def test_binding_errors_delegate_unchanged_to_python(args: tuple[object, ...], k def test_public_responses_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None: - captured: Final[list[LiteLLMResponsesRequest]] = [] + captured: Final[list[NativeCall]] = [] expected: Final = _response() def native( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> ResponsesAPIResponse: captured.append(request) return expected @@ -276,18 +262,16 @@ def test_public_responses_routes_through_dispatch(monkeypatch: pytest.MonkeyPatc finally: NATIVE_RESPONSES.reset() assert result is expected - assert [request.model for request in captured] == ["gpt-4o"] + assert [request.bound["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio async def test_public_aresponses_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None: - captured: Final[list[LiteLLMResponsesRequest]] = [] + captured: Final[list[NativeCall]] = [] expected: Final = _response() async def native( - request: LiteLLMResponsesRequest, - args: tuple[object, ...], - kwargs: Mapping[str, object], + request: NativeCall, ) -> ResponsesAPIResponse: captured.append(request) return expected @@ -300,7 +284,7 @@ async def test_public_aresponses_routes_through_dispatch(monkeypatch: pytest.Mon finally: NATIVE_ARESPONSES.reset() assert result is expected - assert [request.model for request in captured] == ["gpt-4o"] + assert [request.bound["model"] for request in captured] == ["gpt-4o"] def test_responses_with_retries_uses_the_dispatch_entrypoint(monkeypatch: pytest.MonkeyPatch) -> None: @@ -323,6 +307,6 @@ def test_positional_parameters_remain_available_to_native_projection() -> None: include: Final = ["reasoning.encrypted_content"] request: Final = _DISPATCH.request((INPUT, "openai/test-model", include, "Be brief", 16), {}) assert request is not None - assert request.parameters["include"] is include - assert request.parameters["instructions"] == "Be brief" - assert request.parameters["max_output_tokens"] == 16 + assert request.bound["include"] is include + assert request.bound["instructions"] == "Be brief" + assert request.bound["max_output_tokens"] == 16 diff --git a/tests/unit/rust_bridge/AGENTS.md b/tests/unit/rust_bridge/AGENTS.md index 2e46d704c0f..ab88d6b49ce 100644 --- a/tests/unit/rust_bridge/AGENTS.md +++ b/tests/unit/rust_bridge/AGENTS.md @@ -2,7 +2,7 @@ Test what each side of the bridge does, not the rollout policy that picks a side. `LITELLM_RUST` and `catalog.RULES` change every time a route or backend rolls forward, so a test that sets the env var or patches the catalog to reach a path goes red on a policy change even when the code under test is fine -Call each path directly with an explicit decision instead. The Python path is the implementation the dispatcher falls back to, e.g. `litellm.responses.main.responses`. The Rust path is the native binding, e.g. `NATIVE_OCR.load()` from `litellm/rust_bridge/ocr/entrypoints.py`, called with the request, args and kwargs that dispatch would hand it. When the native side reads a policy-derived setting such as `settings.secret_manager().native`, pin that field in the test instead of deriving it from the catalog. `ocr/test_secrets.py` shows the pattern +Call each path directly with an explicit decision instead. The Python path is the implementation the dispatcher falls back to, e.g. `litellm.responses.main.responses`. The Rust path is the native binding, e.g. `NATIVE_OCR.load()` from `litellm/rust_bridge/ocr/entrypoints.py`, called with the `NativeCall` envelope that dispatch would hand it for every public inference route. When the native side reads a policy-derived setting such as `settings.secret_manager().native`, pin that field in the test instead of deriving it from the catalog. `ocr/test_secrets.py` shows the pattern Rollout policy itself, meaning which rule matches and what `LITELLM_RUST` changes, belongs in `test_catalog.py`, `test_configuration.py` and `test_dispatch.py`, tested against rules the test builds rather than the shipped `catalog.RULES` diff --git a/tests/unit/rust_bridge/chat_completions/test_route_host.py b/tests/unit/rust_bridge/chat_completions/test_route_host.py index 7f9295e93b3..92eed17de0f 100644 --- a/tests/unit/rust_bridge/chat_completions/test_route_host.py +++ b/tests/unit/rust_bridge/chat_completions/test_route_host.py @@ -5,7 +5,6 @@ import pytest import litellm from litellm.rust_bridge.chat_completions.route_host import arguments, connection_defaults, response -from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest from litellm.types.utils import ModelResponse @@ -38,18 +37,8 @@ def test_response_builds_the_public_model_response() -> None: def test_arguments_are_the_public_kwargs_view() -> None: kwargs: Final = MappingProxyType({"metadata": {"user_id": "u"}}) - request: Final = LiteLLMChatCompletionsRequest( - model="anthropic/claude-sonnet-4-5", - messages=[{"role": "user", "content": "hi"}], - stream=None, - api_key=None, - api_base=None, - custom_llm_provider="anthropic", - extra_headers=None, - kwargs=kwargs, - ) - assert arguments(request) is kwargs + assert arguments(kwargs) is kwargs @pytest.mark.parametrize( diff --git a/tests/unit/rust_bridge/messages/test_route_host.py b/tests/unit/rust_bridge/messages/test_route_host.py index 7b15553a055..dde76ee5e82 100644 --- a/tests/unit/rust_bridge/messages/test_route_host.py +++ b/tests/unit/rust_bridge/messages/test_route_host.py @@ -2,8 +2,7 @@ from types import MappingProxyType from typing import Final from litellm.rust_bridge.messages.route_host import arguments, response -from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest -from dataclasses import astuple +from litellm.rust_bridge.public_call import NativeCall import pytest import litellm from litellm.rust_bridge.messages import route_host @@ -30,67 +29,37 @@ def test_response_is_a_detached_public_messages_dict() -> None: assert "_hidden_params" not in native -def test_arguments_are_the_public_kwargs_view() -> None: +def test_arguments_preserve_the_bound_view() -> None: kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}}) - request: Final = LiteLLMMessagesRequest( - model="claude-sonnet-4-5", - messages=[{"role": "user", "content": "hi"}], - max_tokens=16, - stream=None, - api_key=None, - api_base=None, - custom_llm_provider="anthropic", + request: Final = NativeCall( + args=(), kwargs=kwargs, - ) - - assert arguments(request) is kwargs - - -pytestmark = pytest.mark.usefixtures("local_model_cost_map") - - -def _flag_model(monkeypatch: pytest.MonkeyPatch, name: str, **flags: bool) -> None: - monkeypatch.setitem( - litellm.model_cost, - name, - { - "litellm_provider": "anthropic", - "mode": "chat", - "input_cost_per_token": 0, - "output_cost_per_token": 0, - **flags, + bound={ + "model": "claude-sonnet-4-5", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 16, + "stream": None, + "api_key": None, + "api_base": None, + "custom_llm_provider": "anthropic", + **kwargs, }, ) - -def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeypatch: pytest.MonkeyPatch) -> None: - _flag_model( - monkeypatch, - "claude-test-adaptive", - supports_reasoning=True, - supports_adaptive_thinking=True, - supports_output_config=True, - supports_xhigh_reasoning_effort=True, - supports_sampling_params=False, - ) - - capabilities: Final = route_host.model_capabilities("anthropic/claude-test-adaptive", None) - - assert capabilities.supports_adaptive_thinking - assert capabilities.supports_output_config - assert not capabilities.supports_legacy_thinking - assert not capabilities.supports_sampling_params - assert capabilities.effort_tiers.xhigh - assert not capabilities.effort_tiers.max + assert arguments(request.bound) is request.bound -def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> None: - capabilities: Final = route_host.model_capabilities("anthropic/not-a-real-model", None) +def test_settings_project_caller_configuration_without_resolving_a_model(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "drop_params", False) + monkeypatch.setattr(litellm, "reasoning_auto_summary", True) - assert capabilities.supports_sampling_params - assert not capabilities.supports_reasoning - assert not capabilities.supports_adaptive_thinking - assert not any(astuple(capabilities.effort_tiers)) + projected: Final = route_host.settings({"drop_params": "true", "additional_drop_params": ["metadata.user_id"]}) + + assert projected == { + "drop_params": True, + "reasoning_auto_summary": True, + "additional_drop_params": ("metadata.user_id",), + } @pytest.mark.parametrize( @@ -108,7 +77,7 @@ def test_drop_params_merges_the_global_flag_with_the_request( ) -> None: monkeypatch.setattr(litellm, "drop_params", global_flag) - assert route_host.shaping("anthropic/not-a-real-model", None, kwargs)["drop_params"] is expected + assert route_host.settings(kwargs)["drop_params"] is expected @pytest.mark.parametrize( @@ -120,36 +89,42 @@ def test_drop_params_merges_the_global_flag_with_the_request( ], ) def test_additional_drop_params_keep_only_string_paths(configured: object, expected: tuple[str, ...]) -> None: - shaping: Final = route_host.shaping("anthropic/not-a-real-model", None, {"additional_drop_params": configured}) + settings: Final = route_host.settings({"additional_drop_params": configured}) - assert shaping["additional_drop_params"] == expected + assert settings["additional_drop_params"] == expected def test_native_request_rejections_map_to_the_public_400() -> None: from types import MappingProxyType - from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest + from litellm.rust_bridge.public_call import NativeCall - request: Final = LiteLLMMessagesRequest( - model="anthropic/claude-sonnet-5", - messages=(), - max_tokens=8, - stream=None, - api_key=None, - api_base=None, - custom_llm_provider=None, + request: Final = NativeCall( + args=(), kwargs=MappingProxyType({}), + bound={ + "model": "anthropic/claude-sonnet-5", + "messages": (), + "max_tokens": 8, + "stream": None, + "api_key": None, + "api_base": None, + "custom_llm_provider": None, + **MappingProxyType({}), + }, ) rejected: Final = ValueError("claude-sonnet-5 does not support top_k=5") rejected.messages_request_error = True # pyright: ignore[reportAttributeAccessIssue] # marker the native host sets - mapped: Final = route_host.map_failure(rejected, request, "anthropic") + mapped: Final = route_host.map_failure(rejected, request.bound, "anthropic") assert isinstance(mapped, litellm.BadRequestError) assert mapped.status_code == 400 assert "does not support top_k=5" in mapped.message assert mapped.model == "claude-sonnet-5" - assert not isinstance(route_host.map_failure(ValueError("plain"), request, "anthropic"), litellm.BadRequestError) + assert not isinstance( + route_host.map_failure(ValueError("plain"), request.bound, "anthropic"), litellm.BadRequestError + ) def test_stream_hidden_params_projects_upstream_headers_the_way_the_python_handler_does() -> None: diff --git a/tests/unit/rust_bridge/messages/test_secrets.py b/tests/unit/rust_bridge/messages/test_secrets.py index bd5dc97cedd..53e1376b978 100644 --- a/tests/unit/rust_bridge/messages/test_secrets.py +++ b/tests/unit/rust_bridge/messages/test_secrets.py @@ -2,7 +2,6 @@ from __future__ import annotations from collections.abc import Awaitable, Mapping from dataclasses import replace -from types import MappingProxyType from typing import Final, Protocol, cast # noqa: TID251 # narrows the parametrized path to its protocol import httpx @@ -12,7 +11,8 @@ import litellm from litellm.integrations.custom_secret_manager import CustomSecretManager from litellm.llms.anthropic.pass_through.messages.handler import anthropic_messages from litellm.rust_bridge import settings -from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, NATIVE_MESSAGES, LiteLLMMessagesRequest +from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, NATIVE_MESSAGES +from litellm.rust_bridge.public_call import NativeCall from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_MODEL, MESSAGES_RESPONSE @@ -48,16 +48,21 @@ class _ManagedSecrets(CustomSecretManager): return self.values.get(secret_name) -def _native_request() -> LiteLLMMessagesRequest: - return LiteLLMMessagesRequest( - model=MESSAGES_MODEL, - messages=MESSAGES, - max_tokens=8, - stream=None, - api_key=None, - api_base=None, - custom_llm_provider=None, - kwargs=MappingProxyType({}), +def _native_request() -> NativeCall: + supplied: Final = _public_kwargs() + return NativeCall( + args=(), + kwargs=supplied, + bound={ + "model": MESSAGES_MODEL, + "messages": MESSAGES, + "max_tokens": 8, + "stream": None, + "api_key": None, + "api_base": None, + "custom_llm_provider": None, + **supplied, + }, ) @@ -72,13 +77,13 @@ async def _python_messages() -> object: async def _rust_messages() -> object: route: Final = NATIVE_MESSAGES.load() assert route is not None - return route(_native_request(), (), _public_kwargs()) + return route(_native_request()) async def _rust_amessages() -> object: route: Final = NATIVE_AMESSAGES.load() assert route is not None - return await route(_native_request(), (), _public_kwargs()) + return await route(_native_request()) @pytest.fixture( diff --git a/tests/unit/rust_bridge/native_route_wheel_test.py b/tests/unit/rust_bridge/native_route_wheel_test.py index a665418e511..9e7523aa29e 100644 --- a/tests/unit/rust_bridge/native_route_wheel_test.py +++ b/tests/unit/rust_bridge/native_route_wheel_test.py @@ -14,6 +14,7 @@ from http.client import HTTPMessage from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from socket import socket as Socket +from types import SimpleNamespace from typing import Final REQUEST_STARTED: Final = threading.Event() @@ -27,6 +28,10 @@ ANTHROPIC_RESPONSE: Final = ( ) +class NativeRouteServer(ThreadingHTTPServer): + request_queue_size = 64 + + class NativeRouteHandler(BaseHTTPRequestHandler): protocol_version = "HTTP/1.1" @@ -110,6 +115,11 @@ def load_native(native_path: Path) -> object: return native_module +def route_call(route: str, api_base: str, outcome: str) -> SimpleNamespace: + fields: Final = route_kwargs(route, api_base, outcome) + return SimpleNamespace(args=(), kwargs=fields, bound=fields) + + def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]: common: Final = { "api_base": api_base, @@ -161,9 +171,9 @@ def assert_rate_limit(route: str, error: BaseException) -> None: def exercise_sync(native: object, api_base: str) -> None: for route in ("transcription", "chat_completions"): function: Final = getattr(native, route) - assert_success(route, function(**route_kwargs(route, api_base, "success"))) + assert_success(route, function(route_call(route, api_base, "success"))) try: - function(**route_kwargs(route, api_base, "429")) + function(route_call(route, api_base, "429")) except native.RustUpstreamError as error: assert_rate_limit(route, error) else: @@ -173,18 +183,18 @@ def exercise_sync(native: object, api_base: str) -> None: async def exercise_async(native: object, api_base: str) -> None: for route in ("transcription", "chat_completions"): function: Final = getattr(native, f"a{route}") - assert_success(route, await function(**route_kwargs(route, api_base, "success"))) + assert_success(route, await function(route_call(route, api_base, "success"))) try: - await function(**route_kwargs(route, api_base, "429")) + await function(route_call(route, api_base, "429")) except native.RustUpstreamError as error: assert_rate_limit(route, error) else: raise AssertionError(f"a{route} accepted a 429 response") - -async def exercise_async_concurrency(native: object, api_base: str) -> None: responses: Final = await asyncio.wait_for( - asyncio.gather(*(native.achat_completions(**route_kwargs("chat_completions", api_base, "success")) for _ in range(32))), + asyncio.gather( + *(native.achat_completions(route_call("chat_completions", api_base, "success")) for _ in range(32)) + ), timeout=15, ) for response in responses: @@ -195,14 +205,13 @@ def exercise_routes(native_path: Path, api_base: str) -> object: native: Final = load_native(native_path) exercise_sync(native, api_base) asyncio.run(exercise_async(native, api_base)) - asyncio.run(exercise_async_concurrency(native, api_base)) return native def exercise_signal(native: object, api_base: str) -> int: try: native.chat_completions( - **route_kwargs("chat_completions", api_base, "hang"), + route_call("chat_completions", api_base, "hang"), ) except KeyboardInterrupt: sys.stdout.write("KeyboardInterrupt\n") @@ -264,7 +273,7 @@ def verify_wheel(wheel: Path) -> int: raise AssertionError(f"expected one native extension, found {len(native_members)}") native_path: Final = wheel_root / native_members[0].filename - server: Final = ThreadingHTTPServer(("127.0.0.1", 0), NativeRouteHandler) + server: Final = NativeRouteServer(("127.0.0.1", 0), NativeRouteHandler) server_thread: Final = threading.Thread(target=server.serve_forever, daemon=True) server_thread.start() api_base: Final = f"http://127.0.0.1:{server.server_address[1]}" diff --git a/tests/unit/rust_bridge/ocr/test_route_host.py b/tests/unit/rust_bridge/ocr/test_route_host.py index 699492e4424..0b8af515a64 100644 --- a/tests/unit/rust_bridge/ocr/test_route_host.py +++ b/tests/unit/rust_bridge/ocr/test_route_host.py @@ -5,17 +5,21 @@ import pytest import litellm from litellm.rust_bridge.ocr.route_host import UpstreamFailure, map_failure from litellm.rust_bridge.ocr.route_host import response as build_ocr_response -from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest +from litellm.rust_bridge.public_call import NativeCall -REQUEST: Final = LiteLLMOcrRequest( - model="mistral/mistral-ocr-latest", - document={"type": "document_url", "document_url": "https://example.com/file.pdf"}, - api_key="test-key", - api_base=None, - timeout=None, - custom_llm_provider=None, - extra_headers=None, +REQUEST: Final = NativeCall( + args=(), kwargs={"req_format": "markdown"}, + bound={ + "model": "mistral/mistral-ocr-latest", + "document": {"type": "document_url", "document_url": "https://example.com/file.pdf"}, + "api_key": "test-key", + "api_base": None, + "timeout": None, + "custom_llm_provider": None, + "extra_headers": None, + **{"req_format": "markdown"}, + }, ) @@ -49,7 +53,7 @@ def test_rust_ocr_response_retains_provider_native_response(): def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> None: error: Final = RustUpstreamError(429, '{"message": "slow down"}', (("retry-after", "7"),)) - public_error: Final = map_failure(error, REQUEST, "mistral") + public_error: Final = map_failure(error, REQUEST.bound, "mistral") assert isinstance(public_error, litellm.RateLimitError) assert public_error.status_code == 429 @@ -62,7 +66,7 @@ def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> N def test_map_failure_maps_upstream_401_to_authentication_error() -> None: error: Final = RustUpstreamError(401, '{"message": "Unauthorized"}', ()) - public_error: Final = map_failure(error, REQUEST, "mistral") + public_error: Final = map_failure(error, REQUEST.bound, "mistral") assert isinstance(public_error, litellm.AuthenticationError) assert public_error.status_code == 401 @@ -73,7 +77,7 @@ def test_map_failure_maps_upstream_401_to_authentication_error() -> None: def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None: error: Final = RuntimeError("bridge exploded") - public_error: Final = map_failure(error, REQUEST, "mistral") + public_error: Final = map_failure(error, REQUEST.bound, "mistral") assert not isinstance(public_error, UpstreamFailure) assert isinstance(public_error, litellm.APIConnectionError) @@ -82,4 +86,4 @@ def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None: def test_map_failure_reports_invalid_request_format_as_unsupported_params() -> None: with pytest.raises(litellm.UnsupportedParamsError, match="Invalid `req_format`: 'markdown'"): - raise map_failure(RustFormatError(), REQUEST, "mistral") + raise map_failure(RustFormatError(), REQUEST.bound, "mistral") diff --git a/tests/unit/rust_bridge/ocr/test_secrets.py b/tests/unit/rust_bridge/ocr/test_secrets.py index b14681d9ad7..5d051fbffe0 100644 --- a/tests/unit/rust_bridge/ocr/test_secrets.py +++ b/tests/unit/rust_bridge/ocr/test_secrets.py @@ -4,7 +4,6 @@ import asyncio from collections.abc import Awaitable, Generator, Mapping from contextlib import contextmanager from dataclasses import replace -from types import MappingProxyType from typing import Final, Literal, Protocol, TypeAlias, cast import httpx @@ -14,7 +13,8 @@ import litellm from litellm.integrations.custom_secret_manager import CustomSecretManager from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge import settings -from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest +from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR +from litellm.rust_bridge.public_call import NativeCall from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec, recording_service from tests.test_litellm_rust.support.requests import OCR_DOCUMENT, OCR_MODEL, OCR_RESPONSE @@ -59,16 +59,21 @@ class _VaultSecrets(CustomSecretManager): return tuple(params for name, params in self.reads if name == "MISTRAL_API_KEY") -def _native_request(api_base: str) -> LiteLLMOcrRequest: - return LiteLLMOcrRequest( - model=OCR_MODEL, - document=OCR_DOCUMENT, - api_key=None, - api_base=api_base, - timeout=None, - custom_llm_provider=None, - extra_headers=None, - kwargs=MappingProxyType({}), +def _native_request(api_base: str) -> NativeCall: + supplied: Final = _public_kwargs(api_base) + return NativeCall( + args=(), + kwargs=supplied, + bound={ + "model": OCR_MODEL, + "document": OCR_DOCUMENT, + "api_key": None, + "api_base": api_base, + "timeout": None, + "custom_llm_provider": None, + "extra_headers": None, + **supplied, + }, ) @@ -79,13 +84,13 @@ def _public_kwargs(api_base: str) -> dict[str, object]: async def _rust_ocr(api_base: str) -> OCRResponse: route: Final = NATIVE_OCR.load() assert route is not None - return route(_native_request(api_base), (), _public_kwargs(api_base)) + return route(_native_request(api_base)) async def _rust_aocr(api_base: str) -> OCRResponse: route: Final = NATIVE_AOCR.load() assert route is not None - return await route(_native_request(api_base), (), _public_kwargs(api_base)) + return await route(_native_request(api_base)) _RUST_PATHS: Final = (_rust_ocr, _rust_aocr) diff --git a/tests/unit/rust_bridge/responses/test_route_host.py b/tests/unit/rust_bridge/responses/test_route_host.py index d04e02b0dda..1d67f2c368c 100644 --- a/tests/unit/rust_bridge/responses/test_route_host.py +++ b/tests/unit/rust_bridge/responses/test_route_host.py @@ -5,8 +5,8 @@ import pytest from pydantic import ValidationError import litellm -from litellm.rust_bridge.responses.route_host import arguments, connection_defaults, response -from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest +from litellm.rust_bridge.responses.route_host import arguments, connection_defaults, map_failure, response +from litellm.rust_bridge.public_call import NativeCall from litellm.types.llms.openai import ResponsesAPIResponse @@ -42,20 +42,24 @@ def test_response_rejects_a_payload_missing_required_fields() -> None: response(MappingProxyType({"object": "response"})) -def test_arguments_are_the_public_kwargs_view() -> None: +def test_arguments_preserve_the_bound_view() -> None: kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}}) - request: Final = LiteLLMResponsesRequest( - model="gpt-4o", - input="hi", - stream=None, - api_key=None, - api_base=None, - custom_llm_provider="openai", - extra_headers=None, + request: Final = NativeCall( + args=(), kwargs=kwargs, + bound={ + "model": "gpt-4o", + "input": "hi", + "stream": None, + "api_key": None, + "api_base": None, + "custom_llm_provider": "openai", + "extra_headers": None, + **kwargs, + }, ) - assert arguments(request) is kwargs + assert arguments(request.bound) is request.bound @pytest.mark.parametrize( @@ -74,3 +78,23 @@ def test_connection_defaults_preserve_openai_precedence( monkeypatch.setattr(litellm, "openai_key", provider_key) monkeypatch.setattr(litellm, "api_base", "https://configured.invalid/v1") assert connection_defaults("openai") == (expected, litellm.api_base) + + +class _UpstreamFailure(Exception): + headers: Final = () + + +@pytest.mark.parametrize( + ("api_base", "base_url", "expected"), + ( + (None, "https://alias.invalid/v1", "https://alias.invalid/v1"), + ("", "https://alias.invalid/v1", "https://alias.invalid/v1"), + ("https://base.invalid/v1", "https://alias.invalid/v1", "https://base.invalid/v1"), + ), +) +def test_failure_preserves_the_explicit_endpoint(api_base: str | None, base_url: str, expected: str) -> None: + upstream: Final = _UpstreamFailure(429, '{"error":{"message":"rate limited"}}') + mapped: Final = map_failure(upstream, {"model": "openai/test-model", "api_base": api_base, "base_url": base_url}) + assert isinstance(mapped, litellm.RateLimitError) + assert str(mapped.response.request.url) == expected + assert mapped.__context__ is upstream diff --git a/tests/unit/rust_bridge/test_model_capabilities.py b/tests/unit/rust_bridge/test_model_capabilities.py new file mode 100644 index 00000000000..9f02660a98f --- /dev/null +++ b/tests/unit/rust_bridge/test_model_capabilities.py @@ -0,0 +1,73 @@ +from typing import Final + +import pytest + +import litellm +from litellm.rust_bridge.model_capabilities import anthropic_model_capabilities + + +pytestmark = pytest.mark.usefixtures("local_model_cost_map") + + +def _flag_model(monkeypatch: pytest.MonkeyPatch, name: str, **flags: bool) -> None: + monkeypatch.setitem( + litellm.model_cost, + name, + { + "litellm_provider": "anthropic", + "mode": "chat", + "input_cost_per_token": 0, + "output_cost_per_token": 0, + **flags, + }, + ) + + +def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeypatch: pytest.MonkeyPatch) -> None: + _flag_model( + monkeypatch, + "claude-test-adaptive", + supports_reasoning=True, + supports_adaptive_thinking=True, + supports_output_config=True, + supports_xhigh_reasoning_effort=True, + supports_sampling_params=False, + ) + + capabilities: Final = anthropic_model_capabilities("anthropic/claude-test-adaptive", None) + + assert capabilities["supports_adaptive_thinking"] + assert capabilities["supports_output_config"] + assert not capabilities["supports_legacy_thinking"] + assert not capabilities["supports_sampling_params"] + assert capabilities["effort_tiers"] == { + "minimal": False, + "low": False, + "medium": False, + "high": False, + "xhigh": True, + "max": False, + } + + +def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> None: + capabilities: Final = anthropic_model_capabilities("anthropic/not-a-real-model", None) + + assert capabilities["supports_sampling_params"] + assert not capabilities["supports_reasoning"] + assert not capabilities["supports_adaptive_thinking"] + assert capabilities["effort_tiers"] == dict.fromkeys(("minimal", "low", "medium", "high", "xhigh", "max"), False) + + +def test_capability_source_observes_runtime_registration_changes(monkeypatch: pytest.MonkeyPatch) -> None: + model: Final = "claude-test-runtime-registration" + _flag_model(monkeypatch, model, supports_reasoning=True, supports_output_config=True) + before: Final = anthropic_model_capabilities(model, "anthropic") + + _flag_model(monkeypatch, model, supports_reasoning=False, supports_output_config=False) + after: Final = anthropic_model_capabilities(model, "anthropic") + + assert before["supports_reasoning"] is True + assert before["supports_output_config"] is True + assert after["supports_reasoning"] is False + assert after["supports_output_config"] is False diff --git a/tests/unit/rust_bridge/test_public_call.py b/tests/unit/rust_bridge/test_public_call.py new file mode 100644 index 00000000000..24d5db057d6 --- /dev/null +++ b/tests/unit/rust_bridge/test_public_call.py @@ -0,0 +1,61 @@ +from collections.abc import Mapping, Sequence +from typing import Final + +import pytest + +from litellm.rust_bridge.public_call import bind, native_call, signature + + +def _messages( + max_tokens: int, + messages: Sequence[object], + model: str, + temperature: float | None = None, + api_key: str | None = None, + **kwargs: object, # kwargs-ok: exercise the public signature binding contract +) -> None: + return None + + +@pytest.mark.parametrize("supplied", ({}, {"api_key": None}, {"api_key": "explicit"})) +def test_native_call_preserves_omission_separately_from_bound_defaults(supplied: Mapping[str, object]) -> None: + messages: Final[Sequence[object]] = [{"role": "user", "content": "hello"}] + args: Final = (128, messages, "model", 0.25) + fields: Final = bind(signature(_messages), args, supplied) + assert fields is not None + + call: Final = native_call(args, supplied, fields) + + assert call.args is args + assert call.kwargs is supplied + assert call.bound == { + "max_tokens": 128, + "messages": messages, + "model": "model", + "temperature": 0.25, + "api_key": supplied.get("api_key"), + } + assert call.bound["messages"] is messages + assert ("api_key" in call.kwargs) == ("api_key" in supplied) + + +def test_native_call_keeps_extra_option_objects_without_nested_kwargs() -> None: + messages: Final[Sequence[object]] = [] + metadata: Final = {"trace": "caller"} + supplied: Final = {"metadata": metadata} + args: Final = (128, messages, "model") + fields: Final = bind(signature(_messages), args, supplied) + assert fields is not None + + call: Final = native_call(args, supplied, fields) + + assert call.bound == { + "max_tokens": 128, + "messages": messages, + "model": "model", + "temperature": None, + "api_key": None, + "metadata": metadata, + } + assert call.bound["metadata"] is metadata + assert supplied == {"metadata": metadata}