refactor(python-bridge): take NativeCall directly and fold routes into per-route folders (#45413)

* refactor(python-bridge): take NativeCall directly and fold routes into per-route folders

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* test(python-bridge): use rstest for updated tests and keep embedding's stub parameter name

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

---------

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
yujonglee 2026-10-08 13:48:28 -07:00 • committed by GitHub
parent 85a3869dfd
commit 0721cffab2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
32 changed files with 559 additions and 908 deletions

View file

@ -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(

View file

@ -13,21 +13,21 @@ use pyo3::{
pub struct PublicCall {
args: Py<PyTuple>,
kwargs: Py<PyDict>,
request: Py<PyAny>,
bound: Py<PyDict>,
}
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<Self> {
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<Option<Bound<'py, PyAny>>> {
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::<PyDict>()
.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::<PyDict>()
.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"))
);
});
}
}

View file

@ -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::<PyDict>().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::<PyDict>().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)
}

View file

@ -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<Option<Bound<'py, PyAny>>> {
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::<PyDict>() {
return bound.get_item(name);
}
request.getattr_opt(name)
}
pub fn present<'py>(
kwargs: &Bound<'py, PyDict>,
bound: &Bound<'py, PyDict>,
name: &str,
) -> PyResult<Option<Bound<'py, PyAny>>> {
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::<PyDict>()
.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<Option<&str>>,
) {
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::<PyDict>().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::<Option<String>>().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::<String>().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::<PyDict>()
.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::<Option<String>>().unwrap().as_deref(),
expected
);
assert!(found.is(&document));
});
}
}

View file

@ -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};

View file

@ -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",

View file

@ -19,13 +19,6 @@ pub(crate) struct RouteOptions {
pub(crate) timeout: Option<Duration>,
}
pub(crate) fn messages_argument(value: &Bound<'_, PyAny>) -> PyResult<Vec<Value>> {
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<Map<String, Value>> {
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", &params).unwrap(),

View file

@ -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_<route>` 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 `<route>.rs`

View file

@ -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<Py<PyAny>> {
let call = super::NativeCall::extract(&call)?;
pub(crate) fn transcription(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
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<Bound<'py, PyAny>> {
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)?;

View file

@ -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<litellm_http::Client, litellm_http::Error>,
secrets: std::sync::Arc<dyn litellm_secrets::source::SecretSource>,
messages: Vec<Value>,
optional_params: Map<String, Value>,
options: RouteOptions,
) -> Result<ChatCompletionsResponse, Error> {
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<Py<PyAny>> {
let call = super::NativeCall::extract(&call)?;
let messages: Vec<Value> = 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<Bound<'py, PyAny>> {
let call = super::NativeCall::extract(&call)?;
let messages: Vec<Value> = 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<Py<PyAny>> {
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<Py<PyAny>> {
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<Py<PyAny>> {
let call = super::NativeCall::extract(&call)?;
run_public(py, call.bound.into_any(), call.args, call.kwargs, true)
}

View file

@ -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::{

View file

@ -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<Py<PyAny>> {
let host = InferenceHost::new(
call.bound.clone().unbind(),
"litellm.rust_bridge.chat_completions.route_host",
);
run_inference::<ChatCompletionsRoute, _>(
py,
call,
asynchronous,
ChatCompletionsPythonHost(host),
)
}
#[pyfunction]
pub(crate) fn completion(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
run_chat_completions(py, call, false)
}
#[pyfunction]
pub(crate) fn acompletion(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
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<AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Self {
Self::new(http, auth, secrets)
}
fn with_cache(self, cache: ScopedCache) -> Self {
self.with_cache(cache)
}
fn machine(
self,
call: <ChatCompletions as Protocol>::Request,
policy: CachePolicy,
) -> HostedMachine<ChatCompletions> {
self.machine(call, policy)
}
}

View file

@ -1,17 +1,18 @@
use pyo3::prelude::*;
use super::NativeCall;
use crate::errors::RustBridgeDeclined;
#[pyfunction]
pub(crate) fn embedding(call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
drop(super::NativeCall::extract(&call)?);
pub(crate) fn embedding(call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
drop(call);
Err(RustBridgeDeclined::new_err(
"native embeddings route is not implemented",
))
}
#[pyfunction]
pub(crate) fn aembedding(call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
pub(crate) fn aembedding(call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
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 {

View file

@ -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<PyAny>,
pub request: Py<PyDict>,
module: &'static str,
}
@ -26,7 +34,7 @@ pub(super) struct ProjectedCall {
}
impl InferenceHost {
pub fn new(request: Py<PyAny>, module: &'static str) -> Self {
pub fn new(request: Py<PyDict>, module: &'static str) -> Self {
Self { request, module }
}
@ -85,21 +93,7 @@ impl InferenceHost {
arguments: &Bound<'py, PyDict>,
name: &str,
) -> PyResult<Option<Bound<'py, PyAny>>> {
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::<PyDict>() {
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<Error = RouteError>;
const OPERATION: LoggingOperation;
const SYNC_CALL_TYPE: &'static str;
const ASYNC_CALL_TYPE: &'static str;
fn new(
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Self;
fn with_cache(self, cache: ScopedCache) -> Self;
fn machine(
self,
call: <Self::Protocol as Protocol>::Request,
policy: CachePolicy,
) -> HostedMachine<Self::Protocol>;
}
pub(super) fn run_inference<R, H>(
py: Python<'_>,
call: NativeCall<'_>,
asynchronous: bool,
host: H,
) -> PyResult<Py<PyAny>>
where
R: InferenceRoute,
H: PythonBinding<Protocol = R::Protocol> + PythonHostCalls<R::Protocol> + '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,
)
}

View file

@ -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<PyErr> {
/// 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<PyAny>,
request: Py<PyDict>,
cache: PythonCache,
}
impl MessagesPythonHost {
pub(super) fn new(request: Py<PyAny>, asynchronous: bool) -> Self {
pub(super) fn new(request: Py<PyDict>, asynchronous: bool) -> Self {
Self {
request,
cache: PythonCache::new(asynchronous),
@ -107,9 +107,7 @@ impl MessagesPythonHost {
arguments: &Bound<'_, PyDict>,
) -> PyResult<Result<MessagesCall, Error>> {
let request = self.request.bind(py);
let argument = |name: &str| -> PyResult<Option<Bound<'_, PyAny>>> {
Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none()))
};
let argument = |name: &str| present(arguments, request, name);
let string = |name: &str| -> PyResult<Option<String>> {
argument(name)?.map(|value| value.extract()).transpose()
};
@ -153,8 +151,7 @@ impl MessagesPythonHost {
) -> PyResult<Option<Map<String, Value>>> {
let request = self.request.bind(py);
let mapping = |name: &str| -> PyResult<Option<Map<String, Value>>> {
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<Option<ProviderSpecificHeaders>> {
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::<Option<String>>())
.ok()
.flatten()
.and_then(|value| value.extract::<Option<String>>().ok().flatten())
.unwrap_or_else(|| "anthropic".into())
}

View file

@ -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<Py<PyAny>> {
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<Py<PyAny>> {
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<Py<PyAny>> {
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<Py<PyAny>> {
run_messages(py, call, false)
}
#[pyfunction]
pub(crate) fn amessages(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
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<Py<PyAny>> {
run_messages(py, call, true)
}

View file

@ -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<Self> {
impl<'py> FromPyObject<'_, 'py> for NativeCall<'py> {
type Error = PyErr;
fn extract(call: Borrowed<'_, 'py, PyAny>) -> PyResult<Self> {
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<Bound<'py, PyDict>>
fn call_hooks(
py: Python<'_>,
operation: LoggingOperation,
request: &Bound<'_, PyAny>,
args: &Bound<'_, PyTuple>,
kwargs: &Bound<'_, PyDict>,
call: &NativeCall<'_>,
asynchronous: bool,
) -> PyResult<(Py<PyDict>, 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::<pyo3::exceptions::PyValueError>(py),
"{name} surfaced {error} instead of ValueError"
);
}
assert!(
error.is_instance_of::<pyo3::exceptions::PyValueError>(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<Bound<'_, PyAny>>| {
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());
});
}
}

View file

@ -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<PyAny>,
request: Py<PyDict>,
data: OcrHostData,
}
impl OcrPythonHost {
pub(super) fn new(request: Py<PyAny>) -> Self {
pub(super) fn new(request: Py<PyDict>) -> Self {
Self {
request,
data: OcrHostData::Unprojected,
@ -209,7 +209,7 @@ del provider
.unwrap()
.cast_into::<PyDict>()
.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);

View file

@ -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<bool> =
Ok(field.exact_true())
});
fn run_ocr(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let (arguments, hooks) = crate::routes::call_hooks(
py,
LoggingOperation::Ocr,
&request,
&args,
&kwargs,
asynchronous,
)?;
fn run_ocr(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult<Py<PyAny>> {
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<OcrSettings> {
}
#[pyfunction]
pub(crate) fn ocr(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
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<Py<PyAny>> {
run_ocr(py, call, false)
}
#[pyfunction]
pub(crate) fn aocr(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
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<Py<PyAny>> {
run_ocr(py, call, true)
}
#[pyfunction]

View file

@ -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<Bound<'py, PyAny>> {
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<Option<Map<String, Value>>> {
self.lookup("extra_headers")?
.extract::<Option<Py<PyAny>>>()?
.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::<Option<Py<PyAny>>>()?
.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<OcrDocumentInput>, 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::<Vec<_>>();
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::<PyDict>()
.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<OcrDocumentInput> {
@ -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::<PyDict>()
.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<String> = request.getattr("accesses").unwrap().extract().unwrap();
assert_eq!(accesses, Vec::<String>::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::<PyDict>()
.unwrap();
assert_eq!(
arguments(&request, &kwargs).model().unwrap(),
"mistral-ocr-latest"
);
assert_eq!(
request.getattr("reads").unwrap().extract::<i32>().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::<PyDict>()
.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::<PyDict>()
.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::<PyDict>()
.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::<PyDict>()
.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::<Vec<_>>(),
["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<Value>,
) {
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!(

View file

@ -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::{

View file

@ -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<Py<PyAny>> {
if let Some(reason) = py
.import(ROUTE_HOST_MODULE)?
.getattr("decline_reason")?
.call1((&call.bound,))?
.extract::<Option<String>>()?
{
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::<String>()?;
let provider = argument("custom_llm_provider")?
.map(|value| value.extract::<String>())
.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>(&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::<ResponsesRoute, _>(py, call, asynchronous, ResponsesPythonHost(host))
}
#[pyfunction]
pub(crate) fn responses(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
run_responses(py, call, false)
}
#[pyfunction]
pub(crate) fn aresponses(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
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<AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Self {
Self::new(http, auth, secrets)
}
fn with_cache(self, cache: ScopedCache) -> Self {
self.with_cache(cache)
}
fn machine(
self,
call: <Responses as Protocol>::Request,
policy: CachePolicy,
) -> HostedMachine<Responses> {
self.machine(call, policy)
}
}

View file

@ -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<Py<PyAny>> {
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::<Option<String>>()?
{
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::<String>()?;
let provider = host
.argument(py, &kwargs, "custom_llm_provider")?
.map(|value| value.extract::<String>())
.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>(&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<Py<PyAny>> {
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<Py<PyAny>> {
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,

View file

@ -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",

View file

@ -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")),
)

View file

@ -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"))
)

View file

@ -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"))
)

View file

@ -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")),
)

View file

@ -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"),
(

View file

@ -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)

View file

@ -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")

View file

@ -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"),
(