diff --git a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs index 70b7ecb45fd..dd3da8991cc 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs @@ -936,9 +936,6 @@ mod payload_tests { /// The payload phases of `Logging` on top of `StubLogger`, with `pre_call` handing the /// payload to the case's `on_pre_call`. const PAYLOAD_LOGGER: &CStr = c" -class Request: - pass - class PayloadLogger(StubLogger): def update_from_kwargs(self, **update): self.update = update @@ -954,7 +951,7 @@ class PayloadLogger(StubLogger): self.record('post_call', None) self.post = (original_response, api_key, additional_args) -request = Request() +bound = {} kwargs = {} logger = PayloadLogger() on_pre_call = lambda additional_args: None @@ -1225,10 +1222,10 @@ on_pre_call = lambda args: observed.append( def check(): assert observed == [(True, True)], observed ")] - #[case::request_attribute_behind_an_omitted_keyword(c" + #[case::bound_value_behind_an_omitted_keyword(c" document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'} pages = [0] -request.document = document +bound['document'] = document kwargs = {'pages': pages} observed = [] on_pre_call = lambda args: observed.append( diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index 8e7645cae7b..8ea67292ae7 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -13,21 +13,21 @@ use pyo3::{ pub struct PublicCall { args: Py, kwargs: Py, - request: Py, + bound: Py, } impl PublicCall { /// Copies the keyword arguments once, so the legacy path's rewrites never reach the /// caller's own dict while every value keeps its identity. pub fn capture( - request: &Bound<'_, PyAny>, + bound: &Bound<'_, PyDict>, args: &Bound<'_, PyTuple>, kwargs: &Bound<'_, PyDict>, ) -> PyResult { Ok(Self { args: args.clone().unbind(), kwargs: kwargs.copy()?.unbind(), - request: request.clone().unbind(), + bound: bound.clone().unbind(), }) } @@ -55,31 +55,30 @@ impl PublicCall { py: Python<'py>, name: &str, ) -> PyResult>> { - lookup(self.kwargs.bind(py), self.request.bind(py), name) + lookup(self.kwargs.bind(py), self.bound.bind(py), name) } pub(crate) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.args)?; visit.call(&self.kwargs)?; - visit.call(&self.request) + visit.call(&self.bound) } } #[cfg(test)] mod tests { use super::*; + use crate::test_support::{local, local_dict}; fn capture<'py>(py: Python<'py>, source: &std::ffi::CStr) -> (PublicCall, Bound<'py, PyDict>) { let locals = PyDict::new(py); py.run(source, Some(&locals), Some(&locals)).unwrap(); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap(); + let call = PublicCall::capture( + &local_dict(&locals, "bound"), + &PyTuple::empty(py), + &local_dict(&locals, "kwargs"), + ) + .unwrap(); (call, locals) } @@ -91,25 +90,22 @@ mod tests { py, c" pages = [0] -class Request: - pass -request = Request() +bound = {'pages': [1]} kwargs = {'pages': pages} ", ); - let caller = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); + let caller = local_dict(&locals, "kwargs"); call.kwargs() .bind(py) .set_item("litellm_call_id", "call") .unwrap(); assert!(!caller.contains("litellm_call_id").unwrap()); - let pages = locals.get_item("pages").unwrap().unwrap(); - assert!(call.lookup(py, "pages").unwrap().unwrap().is(&pages)); + assert!( + call.lookup(py, "pages") + .unwrap() + .unwrap() + .is(local(&locals, "pages")) + ); }); } } diff --git a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs index f252a45b562..a7b28f1b0ca 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/test_support.rs @@ -173,21 +173,23 @@ pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, locals.get_item(name).unwrap().unwrap() } -/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`). +pub(crate) fn local_dict<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyDict> { + local(locals, name).cast_into().unwrap() +} + +/// A legacy call over the namespace's `kwargs` and `bound` dicts, each empty when absent. pub(crate) fn legacy_call( py: Python<'_>, locals: &Bound<'_, PyDict>, asynchronous: bool, ) -> LegacyLogging { - let request = locals - .get_item("request") - .unwrap() - .unwrap_or_else(|| py.None().into_bound(py)); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .map(|kwargs| kwargs.cast_into::().unwrap()) - .unwrap_or_else(|| PyDict::new(py)); - let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap(); + let dict = |name: &str| { + locals + .get_item(name) + .unwrap() + .map(|value| value.cast_into::().unwrap()) + .unwrap_or_else(|| PyDict::new(py)) + }; + let call = PublicCall::capture(&dict("bound"), &PyTuple::empty(py), &dict("kwargs")).unwrap(); LegacyLogging::new(py, crate::LoggingOperation::Ocr, call, asynchronous) } diff --git a/litellm-rust/crates/host-python/src/argument.rs b/litellm-rust/crates/host-python/src/argument.rs index 13214e0e9ed..0b89f94cd74 100644 --- a/litellm-rust/crates/host-python/src/argument.rs +++ b/litellm-rust/crates/host-python/src/argument.rs @@ -2,78 +2,97 @@ use pyo3::{prelude::*, types::PyDict}; pub fn lookup<'py>( kwargs: &Bound<'py, PyDict>, - request: &Bound<'py, PyAny>, + bound: &Bound<'py, PyDict>, name: &str, ) -> PyResult>> { - if let Some(value) = kwargs.get_item(name)? { - return Ok(Some(value)); + match kwargs.get_item(name)? { + Some(value) => Ok(Some(value)), + None => bound.get_item(name), } - if let Ok(bound) = request.cast::() { - return bound.get_item(name); - } - request.getattr_opt(name) +} + +pub fn present<'py>( + kwargs: &Bound<'py, PyDict>, + bound: &Bound<'py, PyDict>, + name: &str, +) -> PyResult>> { + Ok(lookup(kwargs, bound, name)?.filter(|value| !value.is_none())) } #[cfg(test)] mod tests { + use rstest::rstest; + use super::*; - #[test] - fn lookup_prefers_the_keyword_even_when_none_and_falls_back_to_the_request() { + fn dicts<'py>( + py: Python<'py>, + kwargs: &str, + bound: &str, + ) -> (Bound<'py, PyDict>, Bound<'py, PyDict>) { + let eval = |source: &str| { + py.eval(&std::ffi::CString::new(source).unwrap(), None, None) + .unwrap() + .cast_into::() + .unwrap() + }; + (eval(kwargs), eval(bound)) + } + + #[rstest] + #[case::keyword_wins( + "{'api_key': 'keyword'}", + "{'api_key': 'bound'}", + Some(Some("keyword")) + )] + #[case::explicit_none_wins("{'api_key': None}", "{'api_key': 'bound'}", Some(None))] + #[case::bound_fallback("{}", "{'api_key': 'bound'}", Some(Some("bound")))] + #[case::missing("{}", "{}", None)] + fn lookup_prefers_the_keyword_and_falls_back_to_bound( + #[case] kwargs: &str, + #[case] bound: &str, + #[case] expected: Option>, + ) { crate::initialize_python(); Python::attach(|py| { - let locals = PyDict::new(py); - py.run( - c" -key = object() -document = {'type': 'document_url'} -class Request: - api_key = 'from-request' - api_base = 'from-request' - document = document -request = Request() -kwargs = {'api_key': key, 'api_base': None} -", - Some(&locals), - Some(&locals), - ) - .unwrap(); - let item = |name: &str| locals.get_item(name).unwrap().unwrap(); - let kwargs = item("kwargs").cast_into::().unwrap(); - let request = item("request"); - let find = |name: &str| lookup(&kwargs, &request, name).unwrap(); - assert!(find("api_key").unwrap().is(item("key"))); - assert!(find("api_base").unwrap().is_none()); - assert!(find("document").unwrap().is(item("document"))); - assert!(find("model").is_none()); + let (kwargs, bound) = dicts(py, kwargs, bound); + let value = lookup(&kwargs, &bound, "api_key") + .unwrap() + .map(|value| value.extract::>().unwrap()); + assert_eq!(value, expected.map(|value| value.map(str::to_owned))); }); } - #[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, + #[rstest] + #[case::explicit_none_hides_bound("{'api_key': None}", "{'api_key': 'bound'}", None)] + #[case::bound_none("{}", "{'api_key': None}", None)] + #[case::bound_value("{}", "{'api_key': 'bound'}", Some("bound"))] + fn present_treats_none_as_unset( + #[case] kwargs: &str, + #[case] bound: &str, #[case] expected: Option<&str>, ) { crate::initialize_python(); Python::attach(|py| { + let (kwargs, bound) = dicts(py, kwargs, bound); + let value = present(&kwargs, &bound, "api_key") + .unwrap() + .map(|value| value.extract::().unwrap()); + assert_eq!(value.as_deref(), expected); + }); + } + + #[rstest] + fn lookup_returns_the_callers_object() { + crate::initialize_python(); + Python::attach(|py| { + let document = PyDict::new(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") + bound.set_item("document", &document).unwrap(); + let found = lookup(&PyDict::new(py), &bound, "document") .unwrap() .unwrap(); - assert_eq!( - value.extract::>().unwrap().as_deref(), - expected - ); + assert!(found.is(&document)); }); } } diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 00543f64085..323f5b7a79b 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -19,7 +19,7 @@ mod owned; mod runtime; mod services; -pub use argument::lookup; +pub use argument::{lookup, present}; pub use binding::PythonBinding; pub use conversion_cache::{FromPythonCache, ToPythonCache}; pub use driver::{CallOptions, run_call}; diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index cf1080a3881..92acf796e6b 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -17,6 +17,10 @@ mod tokenizer; #[pymodule(gil_used = true)] mod _native { + #[pymodule_export] + use litellm_host_python::{ForkedAfterNativeRuntimeStarted, ProcessReservedForForking}; + use pyo3::{prelude::*, types::PyModule}; + #[cfg(feature = "panic-test")] #[pymodule_export] use crate::diagnostics::_panic_for_test; @@ -29,9 +33,7 @@ mod _native { #[pymodule_export] use crate::routes::audio_transcription::{atranscription, transcription}; #[pymodule_export] - use crate::routes::chat_completions::{ - achat_completions, acompletion, chat_completions, completion, - }; + use crate::routes::chat_completions::{acompletion, completion}; #[pymodule_export] use crate::routes::embeddings::{aembedding, embedding}; #[pymodule_export] @@ -51,9 +53,6 @@ mod _native { use crate::tokenizer::HuggingFaceEncoding; #[pymodule_export] use crate::tokenizer::Tokenizer; - #[pymodule_export] - use litellm_host_python::{ForkedAfterNativeRuntimeStarted, ProcessReservedForForking}; - use pyo3::{prelude::*, types::PyModule}; #[pymodule_init] fn init(module: &Bound<'_, PyModule>) -> PyResult<()> { @@ -101,8 +100,6 @@ mod tests { "atranscription", "messages", "amessages", - "chat_completions", - "achat_completions", "completion", "acompletion", "responses", diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index df15a9591a1..4a7efeffdd2 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -19,13 +19,6 @@ pub(crate) struct RouteOptions { pub(crate) timeout: Option, } -pub(crate) fn messages_argument(value: &Bound<'_, PyAny>) -> PyResult> { - match from_py_argument(value)? { - Value::Array(values) => Ok(values), - _ => Err(PyValueError::new_err("messages must be a list")), - } -} - fn required_object(name: &'static str, value: Value) -> PyResult> { match value { Value::Object(values) => Ok(values), @@ -449,18 +442,6 @@ module.__getattr__ = fail fn argument_converters_keep_nested_values_and_accept_explicit_none() { Python::initialize(); Python::attach(|py| { - let messages = py - .eval( - c"[{'role': 'user', 'content': [{'type': 'text', 'text': 'hi'}]}]", - None, - None, - ) - .unwrap(); - assert_eq!( - Value::Array(messages_argument(&messages).unwrap()), - json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]) - ); - let params = py.eval(c"{'temperature': 0.2}", None, None).unwrap(); assert_eq!( optional_object("optional_params", ¶ms).unwrap(), diff --git a/litellm-rust/crates/python-bridge/src/routes/AGENTS.md b/litellm-rust/crates/python-bridge/src/routes/AGENTS.md index c2afed49b45..948c67eb647 100644 --- a/litellm-rust/crates/python-bridge/src/routes/AGENTS.md +++ b/litellm-rust/crates/python-bridge/src/routes/AGENTS.md @@ -11,3 +11,7 @@ The host driver owns sequencing and terminal events; the bridge supplies fallibl Use the shared `run_public_call` boundary with hooks supplied by bridge composition. `callbacks-legacy-python` owns legacy argument sharing and `Logging` dispatch behind `PublicCall` and `LegacyLogging`. Route bindings supply `callbacks-legacy-python::LoggingOperation` when composing legacy logging and may retain the request needed for projection, but must not duplicate the legacy callback contract Regression tests must observe that an unstarted call does no setup, hook and preflight rewrites affect resource configuration, setup failures reach the selected failure handler once, and provider work is not replayed. Retain existing read-point and object-identity guarantees while changing setup timing + +## Layout + +`messages/` is the reference shape. Each lifecycle route is a folder: `mod.rs` holds `run_` and the thin sync and async `#[pyfunction]` entrypoints that take `NativeCall` and call it, `host.rs` holds the `PythonBinding` host, and anything else route-specific (projection, error mapping, extra pyclasses) gets its own sibling file. Scaffolded routes without a lifecycle (embeddings, audio transcription) stay a single `.rs` 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 fc7bfb08b75..4c0c2760875 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -5,6 +5,8 @@ use litellm_inference_transcription::{ use pyo3::prelude::*; use serde_json::{Map, Value}; +use super::NativeCall; + use crate::{ errors::route_error_to_pyerr, marshal::{RouteOptions, optional_object_field, required_field, value_route_options}, @@ -40,8 +42,7 @@ async fn execute( } #[pyfunction] -pub(crate) fn transcription(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; +pub(crate) fn transcription(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { let audio: Value = litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?; let options = value_route_options(&call.bound)?; @@ -59,9 +60,8 @@ pub(crate) fn transcription(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult< #[pyfunction] pub(crate) fn atranscription<'py>( py: Python<'py>, - call: Bound<'py, PyAny>, + call: NativeCall<'py>, ) -> 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)?; diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs deleted file mode 100644 index 27f076b7cce..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ /dev/null @@ -1,150 +0,0 @@ -mod host; - -use pyo3::types::{PyDict, PyTuple}; - -use crate::execution::{run_async, run_sync}; -use litellm_inference_chat::{ChatCompletionsRoute, Error, types::ChatCompletionsRequest}; -use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse; -use pyo3::prelude::*; -use serde_json::{Map, Value}; - -use crate::{ - errors::route_error_to_pyerr, - marshal::{ - RouteOptions, messages_argument, optional_object_field, required_field, value_route_options, - }, -}; - -async fn execute( - http: Result, - secrets: std::sync::Arc, - messages: Vec, - optional_params: Map, - options: RouteOptions, -) -> Result { - let RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout, - } = options; - ChatCompletionsRoute::new(http?, crate::http::resources().auth.clone(), secrets) - .execute( - ChatCompletionsRequest { - model: &model, - messages: Value::Array(messages), - optional_params, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }, - &(), - None, - ) - .await -} - -#[pyfunction] -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, options), - route_error_to_pyerr, - ) -} - -#[pyfunction] -pub(crate) fn achat_completions<'py>( - py: Python<'py>, - call: Bound<'py, 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, true)?; - let secrets = crate::secrets::source(py)?; - run_async( - py, - execute(http, secrets, messages, optional_params, options), - route_error_to_pyerr, - ) -} - -fn run_public( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - asynchronous: bool, -) -> PyResult> { - use super::inference::InferenceHost; - use litellm_callbacks_legacy_python::LoggingOperation; - let host = InferenceHost::new( - request.clone().unbind(), - "litellm.rust_bridge.chat_completions.route_host", - ); - let cache_call_type = if asynchronous { - "acompletion" - } else { - "completion" - }; - crate::cache::admit_native(py, &kwargs, cache_call_type)?; - let (arguments, hooks) = crate::routes::call_hooks( - py, - LoggingOperation::Completion, - &request, - &args, - &kwargs, - asynchronous, - )?; - crate::routes::run_public_call( - py, - arguments, - move |py, arguments, request| { - let route = ChatCompletionsRoute::new( - crate::http::provider_client(py, arguments, asynchronous)? - .map_err(crate::http::client_error)?, - crate::http::resources().auth.clone(), - crate::secrets::source(py)?, - ); - let (cache, cache_options) = - crate::cache::configured_native(py, arguments, cache_call_type)?; - let route = match cache { - Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( - cache, - litellm_cache_response::CacheScope::Shared, - )), - None => route, - }; - Ok(route.machine(request, cache_options.policy)) - }, - host::ChatCompletionsPythonHost(host), - hooks, - asynchronous, - ) -} - -#[pyfunction] -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<'_>, 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/chat_completions/host.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs index e7525cd3a91..1547a4daa93 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs @@ -1,6 +1,6 @@ use std::convert::Infallible; -use super::super::inference::InferenceHost; +use crate::routes::inference::InferenceHost; use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned}; use litellm_inference_chat::{Error, route::ChatCompletions, types::ChatCompletionsCall}; use pyo3::{ diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs new file mode 100644 index 00000000000..1146436bfe6 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs @@ -0,0 +1,71 @@ +mod host; + +use std::sync::Arc; + +use host::ChatCompletionsPythonHost; +use litellm_auth::AuthServices; +use litellm_cache_response::{CachePolicy, ScopedCache}; +use litellm_callbacks_legacy_python::LoggingOperation; +use litellm_host::{call::HostedMachine, protocol::Protocol}; +use litellm_inference_chat::{ChatCompletionsRoute, route::ChatCompletions}; +use litellm_secrets::source::SecretSource; +use pyo3::prelude::*; + +use super::{ + NativeCall, + inference::{InferenceHost, InferenceRoute, run_inference}, +}; + +fn run_chat_completions( + py: Python<'_>, + call: NativeCall<'_>, + asynchronous: bool, +) -> PyResult> { + let host = InferenceHost::new( + call.bound.clone().unbind(), + "litellm.rust_bridge.chat_completions.route_host", + ); + run_inference::( + py, + call, + asynchronous, + ChatCompletionsPythonHost(host), + ) +} + +#[pyfunction] +pub(crate) fn completion(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_chat_completions(py, call, false) +} + +#[pyfunction] +pub(crate) fn acompletion(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_chat_completions(py, call, true) +} + +impl InferenceRoute for ChatCompletionsRoute { + type Protocol = ChatCompletions; + const OPERATION: LoggingOperation = LoggingOperation::Completion; + const SYNC_CALL_TYPE: &'static str = "completion"; + const ASYNC_CALL_TYPE: &'static str = "acompletion"; + + fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self::new(http, auth, secrets) + } + + fn with_cache(self, cache: ScopedCache) -> Self { + self.with_cache(cache) + } + + fn machine( + self, + call: ::Request, + policy: CachePolicy, + ) -> HostedMachine { + self.machine(call, policy) + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/embeddings.rs b/litellm-rust/crates/python-bridge/src/routes/embeddings.rs index 1d34e3c21ff..2b184274041 100644 --- a/litellm-rust/crates/python-bridge/src/routes/embeddings.rs +++ b/litellm-rust/crates/python-bridge/src/routes/embeddings.rs @@ -1,17 +1,18 @@ use pyo3::prelude::*; +use super::NativeCall; use crate::errors::RustBridgeDeclined; #[pyfunction] -pub(crate) fn embedding(call: Bound<'_, PyAny>) -> PyResult> { - drop(super::NativeCall::extract(&call)?); +pub(crate) fn embedding(call: NativeCall<'_>) -> PyResult> { + drop(call); Err(RustBridgeDeclined::new_err( "native embeddings route is not implemented", )) } #[pyfunction] -pub(crate) fn aembedding(call: Bound<'_, PyAny>) -> PyResult> { +pub(crate) fn aembedding(call: NativeCall<'_>) -> PyResult> { embedding(call) } @@ -36,7 +37,8 @@ call = SimpleNamespace(args=(), kwargs={}, bound={'model':'test-model','input':' Some(&locals), ) .unwrap(); - let call = locals.get_item("call").unwrap().unwrap(); + let call: super::NativeCall<'_> = + locals.get_item("call").unwrap().unwrap().extract().unwrap(); let error = if asynchronous { super::aembedding(call) } else { diff --git a/litellm-rust/crates/python-bridge/src/routes/inference.rs b/litellm-rust/crates/python-bridge/src/routes/inference.rs index 28025b80b0f..0333aaeb04a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/inference.rs +++ b/litellm-rust/crates/python-bridge/src/routes/inference.rs @@ -1,11 +1,19 @@ +use std::sync::Arc; + +use litellm_auth::AuthServices; +use litellm_cache_response::{CachePolicy, CacheScope, ScopedCache}; +use litellm_callbacks_legacy_python::LoggingOperation; use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; -use litellm_host_python::{from_py, lookup}; +use litellm_host::{call::HostedMachine, protocol::Protocol}; +use litellm_host_python::{PythonBinding, PythonHostCalls, from_py, present}; use litellm_http::transport::Error as TransportError; use litellm_inference::RouteError; +use litellm_secrets::source::SecretSource; use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; use serde::Serialize; use serde_json::{Map, Value}; +use super::NativeCall; use crate::{ errors::{RustUpstreamError, route_error_to_pyerr}, marshal::{ @@ -15,7 +23,7 @@ use crate::{ }; pub(super) struct InferenceHost { - pub request: Py, + pub request: Py, module: &'static str, } @@ -26,7 +34,7 @@ pub(super) struct ProjectedCall { } impl InferenceHost { - pub fn new(request: Py, module: &'static str) -> Self { + pub fn new(request: Py, module: &'static str) -> Self { Self { request, module } } @@ -85,21 +93,7 @@ impl InferenceHost { arguments: &Bound<'py, PyDict>, name: &str, ) -> PyResult>> { - let request = self.request.bind(py); - 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,))?; - if !parameter.is_none() { - return Ok(Some(parameter)); - } - let extra = request.getattr("kwargs")?.call_method1("get", (name,))?; - Ok((!extra.is_none()).then_some(extra)) + present(arguments, self.request.bind(py), name) } pub fn parameters( @@ -140,3 +134,62 @@ impl InferenceHost { Ok(PyErr::from_value(mapped)) } } + +pub(super) trait InferenceRoute: Sized + 'static { + type Protocol: Protocol; + const OPERATION: LoggingOperation; + const SYNC_CALL_TYPE: &'static str; + const ASYNC_CALL_TYPE: &'static str; + + fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self; + fn with_cache(self, cache: ScopedCache) -> Self; + fn machine( + self, + call: ::Request, + policy: CachePolicy, + ) -> HostedMachine; +} + +pub(super) fn run_inference( + py: Python<'_>, + call: NativeCall<'_>, + asynchronous: bool, + host: H, +) -> PyResult> +where + R: InferenceRoute, + H: PythonBinding + PythonHostCalls + 'static, +{ + let call_type = if asynchronous { + R::ASYNC_CALL_TYPE + } else { + R::SYNC_CALL_TYPE + }; + crate::cache::admit_native(py, &call.kwargs, call_type)?; + let (arguments, hooks) = super::call_hooks(py, R::OPERATION, &call, asynchronous)?; + super::run_public_call( + py, + arguments, + move |py, arguments, request| { + let route = R::new( + crate::http::provider_client(py, arguments, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); + let (cache, cache_options) = crate::cache::configured_native(py, arguments, call_type)?; + let route = match cache { + Some(cache) => route.with_cache(ScopedCache::new(cache, CacheScope::Shared)), + None => route, + }; + Ok(route.machine(request, cache_options.policy)) + }, + host, + hooks, + asynchronous, + ) +} 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 bde1cc7269c..841ef3081e3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -2,7 +2,7 @@ use crate::cache::{CacheCall, Cached, PythonCache, Selection}; use litellm_host_python::{PythonHostCalls, PythonOwned}; use bytes::Bytes; -use litellm_host_python::{InvokeError, PythonBinding, from_py, lookup, to_py}; +use litellm_host_python::{InvokeError, PythonBinding, from_py, present, to_py}; use litellm_http::transport::Error as TransportError; use litellm_inference_messages::{ Error, MessagesCall, MessagesSettings, MessagesShaping, messages_body, @@ -89,12 +89,12 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { /// The Python side of the Messages route: projects the prepared arguments and builds the /// public response, chunks and exceptions. pub(super) struct MessagesPythonHost { - request: Py, + request: Py, cache: PythonCache, } impl MessagesPythonHost { - pub(super) fn new(request: Py, asynchronous: bool) -> Self { + pub(super) fn new(request: Py, asynchronous: bool) -> Self { Self { request, cache: PythonCache::new(asynchronous), @@ -107,9 +107,7 @@ impl MessagesPythonHost { arguments: &Bound<'_, PyDict>, ) -> PyResult> { let request = self.request.bind(py); - let argument = |name: &str| -> PyResult>> { - Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none())) - }; + let argument = |name: &str| present(arguments, request, name); let string = |name: &str| -> PyResult> { argument(name)?.map(|value| value.extract()).transpose() }; @@ -153,8 +151,7 @@ impl MessagesPythonHost { ) -> PyResult>> { let request = self.request.bind(py); let mapping = |name: &str| -> PyResult>> { - lookup(arguments, request, name)? - .filter(|value| !value.is_none()) + present(arguments, request, name)? .map(|value| from_py(&value)) .transpose() }; @@ -169,8 +166,7 @@ impl MessagesPythonHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> PyResult> { - lookup(arguments, self.request.bind(py), "provider_specific_header")? - .filter(|value| !value.is_none()) + present(arguments, self.request.bind(py), "provider_specific_header")? .map(|value| from_py(&value)) .transpose() } @@ -200,9 +196,9 @@ impl MessagesPythonHost { self.request .bind(py) .get_item("custom_llm_provider") - .and_then(|value| value.extract::>()) .ok() .flatten() + .and_then(|value| value.extract::>().ok().flatten()) .unwrap_or_else(|| "anthropic".into()) } 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 318aae27121..6ed6656cd21 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -2,26 +2,13 @@ mod host; use host::MessagesPythonHost; use litellm_callbacks_legacy_python::LoggingOperation; -use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, -}; +use pyo3::prelude::*; -fn run_messages( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - asynchronous: bool, -) -> PyResult> { - let (arguments, hooks) = crate::routes::call_hooks( - py, - LoggingOperation::Messages, - &request, - &args, - &kwargs, - asynchronous, - )?; +use super::NativeCall; + +fn run_messages(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult> { + let (arguments, hooks) = + crate::routes::call_hooks(py, LoggingOperation::Messages, &call, asynchronous)?; crate::routes::run_public_call( py, arguments, @@ -57,20 +44,18 @@ fn run_messages( }, )) }, - MessagesPythonHost::new(request.unbind(), asynchronous), + MessagesPythonHost::new(call.bound.unbind(), asynchronous), hooks, asynchronous, ) } #[pyfunction] -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) +pub(crate) fn messages(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_messages(py, call, false) } #[pyfunction] -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) +pub(crate) fn amessages(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_messages(py, call, true) } diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index bf52fe3bfe8..646b7799730 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -8,8 +8,7 @@ pub(crate) mod responses; pub(crate) mod token_counter; pub(crate) mod traces; -use litellm_callbacks_legacy_python::LoggingOperation; -use litellm_callbacks_legacy_python::{LegacyLogging, PublicCall}; +use litellm_callbacks_legacy_python::{LegacyLogging, LoggingOperation, PublicCall}; use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol}; use litellm_host_python::{HookChain, PythonBinding, PythonCallHooks, PythonHostCalls}; use pyo3::{ @@ -17,14 +16,16 @@ use pyo3::{ types::{PyDict, PyMapping, PyTuple}, }; -struct NativeCall<'py> { +pub(crate) 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 { +impl<'py> FromPyObject<'_, 'py> for NativeCall<'py> { + type Error = PyErr; + + fn extract(call: Borrowed<'_, 'py, PyAny>) -> PyResult { Ok(Self { args: call.getattr("args")?.cast_into()?, kwargs: mapping_dict(&call.getattr("kwargs")?)?, @@ -46,12 +47,10 @@ fn mapping_dict<'py>(value: &Bound<'py, PyAny>) -> PyResult> fn call_hooks( py: Python<'_>, operation: LoggingOperation, - request: &Bound<'_, PyAny>, - args: &Bound<'_, PyTuple>, - kwargs: &Bound<'_, PyDict>, + call: &NativeCall<'_>, asynchronous: bool, ) -> PyResult<(Py, impl PythonCallHooks + use<>)> { - let call = PublicCall::capture(request, args, kwargs)?; + let call = PublicCall::capture(&call.bound, &call.args, &call.kwargs)?; let arguments = call.arguments(py); Ok(( arguments, @@ -99,6 +98,7 @@ mod tests { prelude::*, types::{PyDict, PyList}, }; + use rstest::rstest; fn value_call<'py>( py: Python<'py>, @@ -126,8 +126,10 @@ mod tests { .unwrap() } - #[test] - fn route_arguments_that_fail_to_convert_raise_value_error() { + #[rstest] + #[case::sync("transcription")] + #[case::asynchronous("atranscription")] + fn route_arguments_that_fail_to_convert_raise_value_error(#[case] name: &str) { Python::initialize(); Python::attach(|py| { let module = crate::native_module(py); @@ -151,48 +153,24 @@ value = Broken() .expect("locals should be readable") .expect("helper value should exist"); - for name in ["chat_completions", "achat_completions"] { - let error = module - .getattr(name) - .and_then(|function| { - function.call1((value_call(py, "messages", &broken, None),)) - }) - .expect_err("route should reject a value it cannot convert"); + let error = module + .getattr(name) + .and_then(|function| function.call1((value_call(py, "audio", &broken, None),))) + .expect_err("route should reject a value it cannot convert"); - assert!( - error.is_instance_of::(py), - "{name} surfaced {error} instead of ValueError" - ); - } + assert!( + error.is_instance_of::(py), + "{name} surfaced {error} instead of ValueError" + ); }); } - #[test] + #[rstest] fn sync_and_async_routes_apply_the_same_input_validation() { Python::initialize(); Python::attach(|py| { let module = crate::native_module(py); - let invalid_messages = PyDict::new(py); - let sync_chat_error = module - .getattr("chat_completions") - .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((value_call(py, "messages", &invalid_messages, None),)) - }) - .expect_err("async chat should reject a non-list messages value"); - - assert_eq!( - sync_chat_error.to_string(), - "ValueError: messages must be a list" - ); - assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string()); - let invalid_headers = PyList::empty(py); let kwargs = PyDict::new(py); kwargs @@ -221,51 +199,43 @@ value = Broken() }); } - #[test] + #[rstest] + fn missing_and_explicit_none_optional_params_share_the_next_error() { + Python::initialize(); + Python::attach(|py| { + let module = crate::native_module(py); + let audio = PyDict::new(py); + let transcribe = |optional_params: Option>| { + let kwargs = PyDict::new(py); + if let Some(optional_params) = optional_params { + kwargs.set_item("optional_params", optional_params).unwrap(); + } + module + .getattr("transcription") + .and_then(|function| { + function.call1((value_call(py, "audio", &audio, Some(&kwargs)),)) + }) + .expect_err("model 'model' has no provider") + .to_string() + }; + + let omitted = transcribe(None); + assert_eq!(transcribe(Some(py.None().into_bound(py))), omitted); + assert_ne!(omitted, "ValueError: optional_params must be a dict"); + assert_eq!( + transcribe(Some(PyList::empty(py).into_any())), + "ValueError: optional_params must be a dict" + ); + }); + } + + #[rstest] fn route_input_validation_preserves_left_to_right_order() { Python::initialize(); Python::attach(|py| { let module = crate::native_module(py); let invalid = PyList::empty(py); - let chat_kwargs = PyDict::new(py); - chat_kwargs - .set_item("optional_params", &invalid) - .expect("kwargs should accept optional_params"); - chat_kwargs - .set_item("extra_headers", &invalid) - .expect("kwargs should accept extra_headers"); - let invalid_messages = PyDict::new(py); - let error = module - .getattr("chat_completions") - .and_then(|function| { - 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"); - - let valid_messages = PyList::empty(py); - let error = module - .getattr("chat_completions") - .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(), - "ValueError: optional_params must be a dict" - ); - let headers_kwargs = PyDict::new(py); headers_kwargs .set_item("extra_headers", &invalid) @@ -286,43 +256,4 @@ value = Broken() assert!(!error.to_string().contains("extra_headers")); }); } - - #[test] - fn missing_and_explicit_none_optional_params_share_the_next_error() { - Python::initialize(); - Python::attach(|py| { - let module = crate::native_module(py); - let messages = PyList::empty(py); - let headers = PyList::empty(py); - let omitted = PyDict::new(py); - omitted - .set_item("extra_headers", &headers) - .expect("kwargs should accept extra_headers"); - let explicit = PyDict::new(py); - explicit - .set_item("optional_params", py.None()) - .expect("kwargs should accept optional_params"); - explicit - .set_item("extra_headers", &headers) - .expect("kwargs should accept extra_headers"); - - let omitted_error = module - .getattr("chat_completions") - .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.call1((value_call(py, "messages", &messages, Some(&explicit)),)) - }) - .expect_err("None optional_params should reach header validation"); - assert_eq!( - omitted_error.to_string(), - "ValueError: extra_headers must be a dict" - ); - assert_eq!(explicit_error.to_string(), omitted_error.to_string()); - }); - } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 7ddaba70937..ca683600d71 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -27,12 +27,12 @@ enum OcrHostData { /// document as it goes), acquires Azure AD tokens, and builds the public response and /// exception. pub(super) struct OcrPythonHost { - request: Py, + request: Py, data: OcrHostData, } impl OcrPythonHost { - pub(super) fn new(request: Py) -> Self { + pub(super) fn new(request: Py) -> Self { Self { request, data: OcrHostData::Unprojected, @@ -209,7 +209,7 @@ del provider .unwrap() .cast_into::() .unwrap(); - let mut host = OcrPythonHost::new(py.None()); + let mut host = OcrPythonHost::new(PyDict::new(py).unbind()); assert!(host.decode_request(py, &kwargs).unwrap().caller_token); locals.del_item("kwargs").unwrap(); drop(kwargs); 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 71b1054fc3c..4dcb2f4c53e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -9,10 +9,9 @@ use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::to_py; use litellm_inference_ocr::provider_config; use litellm_llms::base_llm::ocr::settings::OcrSettings; -use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, -}; +use pyo3::prelude::*; + +use super::NativeCall; use crate::{ coercion::FieldSpec, @@ -29,21 +28,9 @@ const ENABLE_AZURE_AD_TOKEN_REFRESH: FieldSpec = Ok(field.exact_true()) }); -fn run_ocr( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - asynchronous: bool, -) -> PyResult> { - let (arguments, hooks) = crate::routes::call_hooks( - py, - LoggingOperation::Ocr, - &request, - &args, - &kwargs, - asynchronous, - )?; +fn run_ocr(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult> { + let (arguments, hooks) = + crate::routes::call_hooks(py, LoggingOperation::Ocr, &call, asynchronous)?; crate::routes::run_public_call( py, arguments, @@ -61,7 +48,7 @@ fn run_ocr( let route = litellm_inference_ocr::OcrRoute::new(client); Ok(route.machine(request, None)) }, - OcrPythonHost::new(request.unbind()), + OcrPythonHost::new(call.bound.unbind()), hooks, asynchronous, ) @@ -81,15 +68,13 @@ fn project_provider_defaults(snapshot: &Snapshot<'_>) -> PyResult { } #[pyfunction] -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) +pub(crate) fn ocr(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_ocr(py, call, false) } #[pyfunction] -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) +pub(crate) fn aocr(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_ocr(py, call, true) } #[pyfunction] diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index ee7f6238510..7ae5a1f3ce2 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -1,5 +1,5 @@ use litellm_auth::SecretValue; -use litellm_host_python::from_py; +use litellm_host_python::{from_py, present}; use litellm_inference_ocr::{ types::{LiteLLMOcrRequest, OcrDocumentInput}, wire::{OcrWireRequest, consumed_optional_params, decode_document, decode_request_input}, @@ -22,13 +22,13 @@ pub(super) struct OcrHostHandles { } struct OcrArguments<'a, 'py> { - request: &'a Bound<'py, PyAny>, + bound: &'a Bound<'py, PyDict>, kwargs: &'a Bound<'py, PyDict>, } impl<'py> OcrArguments<'_, 'py> { fn lookup(&self, name: &str) -> PyResult> { - litellm_host_python::lookup(self.kwargs, self.request, name)? + litellm_host_python::lookup(self.kwargs, self.bound, name)? .ok_or_else(|| PyValueError::new_err(format!("missing argument: {name}"))) } @@ -58,7 +58,7 @@ impl<'py> OcrArguments<'_, 'py> { fn extra_headers(&self) -> PyResult>> { self.lookup("extra_headers")? .extract::>>()? - .map(|value| from_py(value.bind(self.request.py()))) + .map(|value| from_py(value.bind(self.bound.py()))) .transpose() } @@ -66,7 +66,7 @@ impl<'py> OcrArguments<'_, 'py> { Ok(self .lookup("timeout")? .extract::>>()? - .map(|value| python_timeout_seconds(self.request.py(), value)) + .map(|value| python_timeout_seconds(self.bound.py(), value)) .transpose()? .flatten()) } @@ -110,10 +110,10 @@ impl ProjectedDocument { } pub(super) fn project_request( - request: &Bound<'_, PyAny>, + bound: &Bound<'_, PyDict>, kwargs: &Bound<'_, PyDict>, ) -> PyResult<(LiteLLMOcrRequest, OcrHostHandles)> { - let arguments = OcrArguments { request, kwargs }; + let arguments = OcrArguments { bound, kwargs }; let model = arguments.model()?; let custom_llm_provider = arguments.custom_llm_provider()?; let document = ProjectedDocument::project(&arguments.document()?)?; @@ -122,7 +122,7 @@ pub(super) fn project_request( .map_err(ocr_error_to_pyerr)?; let names = specs.iter().map(|spec| spec.name).collect::>(); let optional_params = - project_optional_fields(names.iter().copied(), |name| kwargs.get_item(name))?; + project_optional_fields(names.iter().copied(), |name| present(kwargs, bound, name))?; let input_sources = request_input_sources( kwargs, names @@ -136,7 +136,7 @@ pub(super) fn project_request( let timeout_seconds = arguments.timeout_seconds()?; let wire = OcrWireRequest { model, - document: document.resolve(request.py())?, + document: document.resolve(bound.py())?, api_key, api_base, custom_llm_provider, @@ -169,11 +169,20 @@ mod tests { locals } + fn dict<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyDict> { + locals + .get_item(name) + .unwrap() + .unwrap() + .cast_into::() + .unwrap() + } + fn arguments<'a, 'py>( - request: &'a Bound<'py, PyAny>, + bound: &'a Bound<'py, PyDict>, kwargs: &'a Bound<'py, PyDict>, ) -> OcrArguments<'a, 'py> { - OcrArguments { request, kwargs } + OcrArguments { bound, kwargs } } fn project_document(document: &Bound<'_, PyAny>) -> PyResult { @@ -204,141 +213,26 @@ sys.modules['litellm.rust_bridge.timeouts'] = timeouts } #[test] - fn kwargs_override_request_attributes_including_explicit_none() { + fn kwargs_override_bound_values_including_explicit_none() { Python::initialize(); Python::attach(|py| { let locals = eval( py, c" -class Request: - def __init__(self): - self.accesses = [] - def __getattribute__(self, name): - if name != 'accesses': - object.__getattribute__(self, 'accesses').append(name) - return object.__getattribute__(self, name) -request = Request() -request.model = 'from-request' -request.custom_llm_provider = 'mistral' +bound = {'model': 'from-bound', 'custom_llm_provider': 'mistral'} kwargs = {'model': 'from-kwargs', 'custom_llm_provider': None} ", ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let arguments = arguments(&request, &kwargs); + let (bound, kwargs) = (dict(&locals, "bound"), dict(&locals, "kwargs")); + let arguments = arguments(&bound, &kwargs); assert_eq!(arguments.model().unwrap(), "from-kwargs"); assert_eq!(arguments.custom_llm_provider().unwrap(), None); - let accesses: Vec = request.getattr("accesses").unwrap().extract().unwrap(); - assert_eq!(accesses, Vec::::new()); }); } - #[test] - fn missing_kwargs_read_the_request_property_once() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Request: - def __init__(self): - self.reads = 0 - @property - def model(self): - self.reads += 1 - return 'mistral-ocr-latest' -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - assert_eq!( - arguments(&request, &kwargs).model().unwrap(), - "mistral-ocr-latest" - ); - assert_eq!( - request.getattr("reads").unwrap().extract::().unwrap(), - 1 - ); - }); - } - - #[test] - fn request_property_exceptions_keep_their_identity() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -failure = LookupError('model failed') -class Request: - @property - def model(self): - raise failure -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let error = arguments(&request, &kwargs).model().unwrap_err(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - }); - } - - #[test] - fn unused_raising_property_is_never_inspected() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Request: - @property - def unused(self): - raise RuntimeError('unused') - model = 'mistral-ocr-latest' - custom_llm_provider = None -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let arguments = arguments(&request, &kwargs); - assert_eq!(arguments.model().unwrap(), "mistral-ocr-latest"); - assert_eq!(arguments.custom_llm_provider().unwrap(), None); - }); - } - - /// A reader that rewrites the request while it runs shows which arguments projection - /// read before it and which after: every other argument is read first, and the read - /// happens exactly once. + /// A reader that rewrites the bound arguments while it runs shows which arguments + /// projection read before it and which after: every other argument is read first, and + /// the read happens exactly once. #[test] fn document_readers_are_read_once_after_every_other_argument() { Python::initialize(); @@ -347,37 +241,28 @@ kwargs = {} let locals = eval( py, c" -class Request: - model = 'mistral/mistral-ocr-latest' - custom_llm_provider = None - api_key = None - api_base = 'https://original.example.com' - extra_headers = {'x-source': 'original'} - timeout = 1 - @property - def document(self): - return document class Reader: reads = 0 def read(self): Reader.reads += 1 - Request.api_base = 'https://mutated.example.com' - Request.extra_headers = {'x-source': 'mutated'} - Request.timeout = 9 + bound['api_base'] = 'https://mutated.example.com' + bound['extra_headers'] = {'x-source': 'mutated'} + bound['timeout'] = 9 return b'abc' -document = {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'} -request = Request() +bound = { + 'model': 'mistral/mistral-ocr-latest', + 'custom_llm_provider': None, + 'api_key': None, + 'api_base': 'https://original.example.com', + 'extra_headers': {'x-source': 'original'}, + 'timeout': 1, + 'document': {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'}, +} kwargs = {} ", ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let (projected, _) = project_request(&request, &kwargs).unwrap(); + let (projected, _) = + project_request(&dict(&locals, "bound"), &dict(&locals, "kwargs")).unwrap(); assert_eq!( py.eval(c"Reader.reads", Some(&locals), Some(&locals)) .unwrap() @@ -524,34 +409,26 @@ document = Document() }); } - fn request_and_kwargs<'py>( + fn bound_and_kwargs<'py>( py: Python<'py>, kwargs: &std::ffi::CStr, - ) -> (Bound<'py, PyAny>, Bound<'py, PyDict>) { + ) -> (Bound<'py, PyDict>, Bound<'py, PyDict>) { let locals = eval( py, c" -class Request: - model = 'mistral/mistral-ocr-latest' - custom_llm_provider = 'mistral' - document = {'type': 'document_url', 'document_url': 'https://example.com/request.pdf'} - api_key = None - api_base = 'https://request.example.com' - extra_headers = {'x-source': 'request'} - timeout = 1 -request = Request() +bound = { + 'model': 'mistral/mistral-ocr-latest', + 'custom_llm_provider': 'mistral', + 'document': {'type': 'document_url', 'document_url': 'https://example.com/bound.pdf'}, + 'api_key': None, + 'api_base': 'https://bound.example.com', + 'extra_headers': {'x-source': 'bound'}, + 'timeout': 1, +} ", ); py.run(kwargs, Some(&locals), Some(&locals)).unwrap(); - ( - locals.get_item("request").unwrap().unwrap(), - locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(), - ) + (dict(&locals, "bound"), dict(&locals, "kwargs")) } #[test] @@ -559,7 +436,7 @@ request = Request() Python::initialize(); Python::attach(|py| { stub_timeout_conversion(py); - let (request, kwargs) = request_and_kwargs( + let (bound, kwargs) = bound_and_kwargs( py, c" kwargs = { @@ -575,7 +452,7 @@ kwargs = { } ", ); - let (projected, _) = project_request(&request, &kwargs).unwrap(); + let (projected, _) = project_request(&bound, &kwargs).unwrap(); assert_eq!( projected.optional_params.keys().collect::>(), ["pages"] @@ -584,12 +461,29 @@ kwargs = { }); } + #[rstest::rstest] + #[case::explicit_none_is_unset(c"bound['pages'] = [1]\nkwargs = {'pages': None}", None)] + #[case::bound_fallback(c"bound['pages'] = [1]\nkwargs = {}", Some(serde_json::json!([1])))] + #[case::keyword_wins(c"bound['pages'] = [1]\nkwargs = {'pages': [0]}", Some(serde_json::json!([0])))] + fn optional_params_read_through_bound_and_drop_none( + #[case] script: &std::ffi::CStr, + #[case] expected: Option, + ) { + Python::initialize(); + Python::attach(|py| { + stub_timeout_conversion(py); + let (bound, kwargs) = bound_and_kwargs(py, script); + let (projected, _) = project_request(&bound, &kwargs).unwrap(); + assert_eq!(projected.optional_params.get("pages").cloned(), expected); + }); + } + #[test] fn replacement_kwargs_project_provider_connection_and_timeout() { Python::initialize(); Python::attach(|py| { stub_timeout_conversion(py); - let (request, kwargs) = request_and_kwargs( + let (bound, kwargs) = bound_and_kwargs( py, c" kwargs = { @@ -602,7 +496,7 @@ kwargs = { } ", ); - let (projected, handles) = project_request(&request, &kwargs).unwrap(); + let (projected, handles) = project_request(&bound, &kwargs).unwrap(); assert_eq!(handles.provider, "azure_ai"); assert_eq!(projected.model, "mistral-ocr-latest"); assert_eq!( diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs index 25d3fc07332..b22bbc6192b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs @@ -1,6 +1,6 @@ use std::convert::Infallible; -use super::super::inference::InferenceHost; +use crate::routes::inference::InferenceHost; use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned}; use litellm_inference_responses::{Error, route::Responses, types::ResponsesCall}; use pyo3::{ diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs new file mode 100644 index 00000000000..a5b81b91752 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs @@ -0,0 +1,102 @@ +mod host; +mod websocket; + +use std::sync::Arc; + +use host::ResponsesPythonHost; +use litellm_auth::AuthServices; +use litellm_cache_response::{CachePolicy, ScopedCache}; +use litellm_callbacks_legacy_python::LoggingOperation; +use litellm_host::{call::HostedMachine, protocol::Protocol}; +use litellm_host_python::present; +use litellm_inference_responses::{ResponsesRoute, route::Responses}; +use litellm_secrets::source::SecretSource; +use pyo3::prelude::*; +use serde_json::Value; +pub(crate) use websocket::ResponsesWebSocketConnection; + +use super::{ + NativeCall, + inference::{InferenceHost, InferenceRoute, run_inference}, +}; +use crate::errors::RustBridgeDeclined; + +const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.responses.route_host"; + +fn run_responses(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult> { + if let Some(reason) = py + .import(ROUTE_HOST_MODULE)? + .getattr("decline_reason")? + .call1((&call.bound,))? + .extract::>()? + { + return Err(RustBridgeDeclined::new_err(reason)); + } + let argument = |name: &str| present(&call.kwargs, &call.bound, name); + let model = argument("model")? + .ok_or_else(|| pyo3::exceptions::PyValueError::new_err("model is required"))? + .extract::()?; + let provider = argument("custom_llm_provider")? + .map(|value| value.extract::()) + .transpose()?; + if provider + .as_deref() + .is_some_and(|provider| provider != "openai") + || model + .strip_prefix("openai/") + .unwrap_or(&model) + .contains('/') + { + return Err(RustBridgeDeclined::new_err( + "native HTTP responses provider", + )); + } + if argument("stream")? + .map(|value| litellm_host_python::from_py::(&value)) + .transpose()? + .is_some_and(|value| value == Value::Bool(true)) + { + return Err(RustBridgeDeclined::new_err( + "native Python responses streaming", + )); + } + let host = InferenceHost::new(call.bound.clone().unbind(), ROUTE_HOST_MODULE); + run_inference::(py, call, asynchronous, ResponsesPythonHost(host)) +} + +#[pyfunction] +pub(crate) fn responses(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_responses(py, call, false) +} + +#[pyfunction] +pub(crate) fn aresponses(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + run_responses(py, call, true) +} + +impl InferenceRoute for ResponsesRoute { + type Protocol = Responses; + const OPERATION: LoggingOperation = LoggingOperation::Responses; + const SYNC_CALL_TYPE: &'static str = "responses"; + const ASYNC_CALL_TYPE: &'static str = "aresponses"; + + fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self::new(http, auth, secrets) + } + + fn with_cache(self, cache: ScopedCache) -> Self { + self.with_cache(cache) + } + + fn machine( + self, + call: ::Request, + policy: CachePolicy, + ) -> HostedMachine { + self.machine(call, policy) + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs similarity index 57% rename from litellm-rust/crates/python-bridge/src/routes/responses.rs rename to litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs index 780f2e6929b..615a838a597 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs @@ -1,121 +1,12 @@ -mod host; - use litellm_inference_responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; -use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, -}; +use pyo3::prelude::*; use serde_json::Value; use crate::{ - errors::{RustBridgeDeclined, route_error_to_pyerr}, + errors::route_error_to_pyerr, marshal::{marshal_headers, optional_timeout}, }; -fn run_public( - py: Python<'_>, - request: Bound<'_, PyAny>, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - asynchronous: bool, -) -> PyResult> { - use super::inference::InferenceHost; - use litellm_callbacks_legacy_python::LoggingOperation; - let host = InferenceHost::new( - request.clone().unbind(), - "litellm.rust_bridge.responses.route_host", - ); - if let Some(reason) = py - .import("litellm.rust_bridge.responses.route_host")? - .getattr("decline_reason")? - .call1((&request,))? - .extract::>()? - { - return Err(RustBridgeDeclined::new_err(reason)); - } - let model = host - .argument(py, &kwargs, "model")? - .ok_or_else(|| pyo3::exceptions::PyValueError::new_err("model is required"))? - .extract::()?; - let provider = host - .argument(py, &kwargs, "custom_llm_provider")? - .map(|value| value.extract::()) - .transpose()?; - if provider - .as_deref() - .is_some_and(|provider| provider != "openai") - || model - .strip_prefix("openai/") - .unwrap_or(&model) - .contains('/') - { - return Err(RustBridgeDeclined::new_err( - "native HTTP responses provider", - )); - } - if host - .argument(py, &kwargs, "stream")? - .map(|value| litellm_host_python::from_py::(&value)) - .transpose()? - .is_some_and(|value| value == Value::Bool(true)) - { - return Err(RustBridgeDeclined::new_err( - "native Python responses streaming", - )); - } - let cache_call_type = if asynchronous { - "aresponses" - } else { - "responses" - }; - crate::cache::admit_native(py, &kwargs, cache_call_type)?; - let (arguments, hooks) = crate::routes::call_hooks( - py, - LoggingOperation::Responses, - &request, - &args, - &kwargs, - asynchronous, - )?; - crate::routes::run_public_call( - py, - arguments, - move |py, arguments, request| { - let route = litellm_inference_responses::ResponsesRoute::new( - crate::http::provider_client(py, arguments, asynchronous)? - .map_err(crate::http::client_error)?, - crate::http::resources().auth.clone(), - crate::secrets::source(py)?, - ); - let (cache, cache_options) = - crate::cache::configured_native(py, arguments, cache_call_type)?; - let route = match cache { - Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new( - cache, - litellm_cache_response::CacheScope::Shared, - )), - None => route, - }; - Ok(route.machine(request, cache_options.policy)) - }, - host::ResponsesPythonHost(host), - hooks, - asynchronous, - ) -} - -#[pyfunction] -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<'_>, call: Bound<'_, PyAny>) -> PyResult> { - let call = super::NativeCall::extract(&call)?; - run_public(py, call.bound.into_any(), call.args, call.kwargs, true) -} - #[pyclass] pub(crate) struct ResponsesWebSocketConnection { inner: RustResponsesWebSocketConnection, diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 5ca7b79d127..c83e4e5318c 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -110,12 +110,6 @@ def messages( def amessages( call: NativeCall, ) -> Coroutine[object, object, AnthropicMessagesResponse | AsyncIterator[bytes]]: ... -def chat_completions( - call: NativeCall, -) -> dict[str, object]: ... -def achat_completions( - call: NativeCall, -) -> Future[dict[str, object]]: ... @final class ResponsesWebSocketConnection: @@ -251,14 +245,12 @@ __all__ = [ "RustUpstreamError", "TokenCounter", "Tokenizer", - "achat_completions", "acompletion", "aembedding", "amessages", "aocr", "aresponses", "atranscription", - "chat_completions", "completion", "embedding", "gil_stats", diff --git a/litellm/rust_bridge/chat_completions/route_host.py b/litellm/rust_bridge/chat_completions/route_host.py index cb06ea8e213..d600495c919 100644 --- a/litellm/rust_bridge/chat_completions/route_host.py +++ b/litellm/rust_bridge/chat_completions/route_host.py @@ -37,16 +37,12 @@ def response(value: Mapping[str, object]) -> ModelResponse: return ModelResponse(**value) -def arguments(request: Mapping[str, object]) -> Mapping[str, object]: - return request - - 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), + request, optional_str(request.get("api_base")) or optional_str(request.get("base_url")), ) diff --git a/litellm/rust_bridge/messages/route_host.py b/litellm/rust_bridge/messages/route_host.py index 88a81fdae4a..69e27e78223 100644 --- a/litellm/rust_bridge/messages/route_host.py +++ b/litellm/rust_bridge/messages/route_host.py @@ -39,10 +39,6 @@ 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: Mapping[str, object]) -> Mapping[str, object]: - return request - - def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception: if getattr(error, "messages_request_error", False): return litellm.BadRequestError( @@ -51,7 +47,7 @@ def map_failure(error: Exception, request: Mapping[str, object], request_provide llm_provider=request_provider, ) return failures.map_native_failure( - error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base")) + error, str(request["model"]), request_provider, request, optional_str(request.get("api_base")) ) diff --git a/litellm/rust_bridge/ocr/route_host.py b/litellm/rust_bridge/ocr/route_host.py index a7eb0829f5c..8f563a381ef 100644 --- a/litellm/rust_bridge/ocr/route_host.py +++ b/litellm/rust_bridge/ocr/route_host.py @@ -12,7 +12,7 @@ from litellm.rust_bridge import failures from litellm.rust_bridge.failures import UpstreamFailure from litellm.rust_bridge.public_call import optional_str -__all__ = ("UpstreamFailure", "arguments", "map_failure", "response") +__all__ = ("UpstreamFailure", "map_failure", "response") _RESPONSE_ADAPTER: Final = TypeAdapter(dict[str, object]) @@ -27,10 +27,6 @@ def response(value: Mapping[str, object]) -> OCRResponse: return normalized -def arguments(request: Mapping[str, object]) -> Mapping[str, object]: - return request - - def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception: if getattr(error, "ocr_request_format_error", False): return litellm.UnsupportedParamsError( @@ -39,5 +35,5 @@ def map_failure(error: Exception, request: Mapping[str, object], request_provide llm_provider=request_provider, ) return failures.map_native_failure( - error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base")) + error, str(request["model"]), request_provider, request, optional_str(request.get("api_base")) ) diff --git a/litellm/rust_bridge/responses/route_host.py b/litellm/rust_bridge/responses/route_host.py index 00f24244f49..d8ad2a94b75 100644 --- a/litellm/rust_bridge/responses/route_host.py +++ b/litellm/rust_bridge/responses/route_host.py @@ -20,17 +20,13 @@ def response(value: Mapping[str, object]) -> ResponsesAPIResponse: return ResponsesAPIResponse.model_validate(value) -def arguments(request: Mapping[str, object]) -> Mapping[str, object]: - return request - - 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), + request, optional_str(request.get("api_base")) or optional_str(request.get("base_url")), ) 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 92eed17de0f..402c032f340 100644 --- a/tests/unit/rust_bridge/chat_completions/test_route_host.py +++ b/tests/unit/rust_bridge/chat_completions/test_route_host.py @@ -4,7 +4,7 @@ from typing import Final import pytest import litellm -from litellm.rust_bridge.chat_completions.route_host import arguments, connection_defaults, response +from litellm.rust_bridge.chat_completions.route_host import connection_defaults, response from litellm.types.utils import ModelResponse @@ -35,12 +35,6 @@ def test_response_builds_the_public_model_response() -> None: assert built.usage.total_tokens == 5 -def test_arguments_are_the_public_kwargs_view() -> None: - kwargs: Final = MappingProxyType({"metadata": {"user_id": "u"}}) - - assert arguments(kwargs) is kwargs - - @pytest.mark.parametrize( ("provider", "global_key", "provider_key", "expected_key", "expected_base"), ( diff --git a/tests/unit/rust_bridge/messages/test_route_host.py b/tests/unit/rust_bridge/messages/test_route_host.py index dde76ee5e82..16b9634beb9 100644 --- a/tests/unit/rust_bridge/messages/test_route_host.py +++ b/tests/unit/rust_bridge/messages/test_route_host.py @@ -1,8 +1,7 @@ from types import MappingProxyType from typing import Final -from litellm.rust_bridge.messages.route_host import arguments, response -from litellm.rust_bridge.public_call import NativeCall +from litellm.rust_bridge.messages.route_host import response import pytest import litellm from litellm.rust_bridge.messages import route_host @@ -29,26 +28,6 @@ def test_response_is_a_detached_public_messages_dict() -> None: assert "_hidden_params" not in native -def test_arguments_preserve_the_bound_view() -> None: - kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}}) - request: Final = NativeCall( - args=(), - kwargs=kwargs, - 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, - }, - ) - - assert arguments(request.bound) is request.bound - - 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) diff --git a/tests/unit/rust_bridge/native_route_wheel_test.py b/tests/unit/rust_bridge/native_route_wheel_test.py index 9e7523aa29e..268eca0f663 100644 --- a/tests/unit/rust_bridge/native_route_wheel_test.py +++ b/tests/unit/rust_bridge/native_route_wheel_test.py @@ -20,13 +20,6 @@ from typing import Final REQUEST_STARTED: Final = threading.Event() REQUEST_CANCELLED: Final = threading.Event() -ANTHROPIC_RESPONSE: Final = ( - b'{"id":"msg_native","type":"message","role":"assistant",' - b'"model":"claude-sonnet-4-5","content":[{"type":"text","text":"native-message"}],' - b'"stop_reason":"end_turn","stop_sequence":null,' - b'"usage":{"input_tokens":2,"output_tokens":3}}' -) - class NativeRouteServer(ThreadingHTTPServer): request_queue_size = 64 @@ -49,7 +42,7 @@ class NativeRouteHandler(BaseHTTPRequestHandler): return status: Final = 429 if outcome == "429" else 200 - response_body: Final = native_response(status, route) + response_body: Final = native_response(status) self.send_response(status) self.send_header("content-type", "application/json") @@ -78,32 +71,23 @@ def assert_native_request( headers: HTTPMessage, body: object, ) -> None: - if route not in {"transcription", "chat_completions"}: + if route != "transcription": raise AssertionError(f"unexpected route marker: {route!r}") if outcome not in {"success", "429", "hang"}: raise AssertionError(f"unexpected outcome marker: {outcome!r}") if not isinstance(body, dict): raise TypeError(f"{route} sent {type(body).__name__}, expected a JSON object") - if route == "transcription": - assert path == "/model/mistral.voxtral-mini-3b-2507/converse" - assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ") - assert headers.get("x-amz-date") - assert body["messages"][0]["content"][0]["audio"]["source"]["bytes"] == "AQI=" - assert "The audio language is en" in body["messages"][0]["content"][1]["text"] - return - assert path == "/v1/messages" - assert headers.get("x-api-key") == "sk-native" - assert body["model"] == "claude-sonnet-4-5" - assert body["max_tokens"] == 17 - assert body["messages"][0]["content"] == [{"type": "text", "text": "hello-from-chat"}] + assert path == "/model/mistral.voxtral-mini-3b-2507/converse" + assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ") + assert headers.get("x-amz-date") + assert body["messages"][0]["content"][0]["audio"]["source"]["bytes"] == "AQI=" + assert "The audio language is en" in body["messages"][0]["content"][1]["text"] -def native_response(status: int, route: str | None) -> bytes: +def native_response(status: int) -> bytes: if status == 429: return b'{"error":"native-rate-limit"}' - if route == "transcription": - return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}' - return ANTHROPIC_RESPONSE + return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}' def load_native(native_path: Path) -> object: @@ -138,38 +122,23 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]: "language": "en", }, } - if route == "chat_completions": - return common | { - "model": "anthropic/claude-sonnet-4-5", - "messages": [{"role": "user", "content": "hello-from-chat"}], - "optional_params": {"max_tokens": 17}, - "api_key": "sk-native", - } raise AssertionError(f"unknown route: {route}") def assert_success(route: str, response: object) -> None: if not isinstance(response, dict): raise TypeError(f"{route} returned {type(response).__name__}, expected dict") - actual: Final = success_value(route, response) - expected: Final = "native-transcription" if route == "transcription" else "native-message" - if actual != expected: - raise AssertionError(f"{route} returned {actual!r}, expected {expected!r}") - - -def success_value(route: str, response: dict[object, object]) -> object: - if route == "transcription": - return response["text"] - return response["choices"][0]["message"]["content"] + if response["text"] != "native-transcription": + raise AssertionError(f"{route} returned {response['text']!r}, expected 'native-transcription'") def assert_rate_limit(route: str, error: BaseException) -> None: - if error.args != (429, native_response(429, route).decode()): + if error.args != (429, native_response(429).decode()): raise AssertionError(f"{route} returned the wrong 429 error: {error!r}") def exercise_sync(native: object, api_base: str) -> None: - for route in ("transcription", "chat_completions"): + for route in ("transcription",): function: Final = getattr(native, route) assert_success(route, function(route_call(route, api_base, "success"))) try: @@ -181,7 +150,7 @@ 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"): + for route in ("transcription",): function: Final = getattr(native, f"a{route}") assert_success(route, await function(route_call(route, api_base, "success"))) try: @@ -192,13 +161,11 @@ async def exercise_async(native: object, api_base: str) -> None: raise AssertionError(f"a{route} accepted a 429 response") responses: Final = await asyncio.wait_for( - asyncio.gather( - *(native.achat_completions(route_call("chat_completions", api_base, "success")) for _ in range(32)) - ), + asyncio.gather(*(native.atranscription(route_call("transcription", api_base, "success")) for _ in range(32))), timeout=15, ) for response in responses: - assert_success("chat_completions", response) + assert_success("transcription", response) def exercise_routes(native_path: Path, api_base: str) -> object: @@ -210,8 +177,8 @@ def exercise_routes(native_path: Path, api_base: str) -> object: def exercise_signal(native: object, api_base: str) -> int: try: - native.chat_completions( - route_call("chat_completions", api_base, "hang"), + native.transcription( + route_call("transcription", api_base, "hang"), ) except KeyboardInterrupt: sys.stdout.write("KeyboardInterrupt\n") diff --git a/tests/unit/rust_bridge/responses/test_route_host.py b/tests/unit/rust_bridge/responses/test_route_host.py index 1d67f2c368c..e110f16350b 100644 --- a/tests/unit/rust_bridge/responses/test_route_host.py +++ b/tests/unit/rust_bridge/responses/test_route_host.py @@ -5,8 +5,7 @@ import pytest from pydantic import ValidationError import litellm -from litellm.rust_bridge.responses.route_host import arguments, connection_defaults, map_failure, response -from litellm.rust_bridge.public_call import NativeCall +from litellm.rust_bridge.responses.route_host import connection_defaults, map_failure, response from litellm.types.llms.openai import ResponsesAPIResponse @@ -42,26 +41,6 @@ def test_response_rejects_a_payload_missing_required_fields() -> None: response(MappingProxyType({"object": "response"})) -def test_arguments_preserve_the_bound_view() -> None: - kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}}) - 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.bound) is request.bound - - @pytest.mark.parametrize( ("global_key", "provider_key", "expected"), (