refactor(rust-bridge): unify native call inputs and Messages settings (#45126)

* refactor(rust): separate Messages settings from capability inputs

* refactor(rust-bridge): unify Messages OCR and Responses call inputs

* refactor(rust-bridge): share NativeCall across inference entrypoints

* fix(rust-bridge): preserve Responses URL aliases and public test inputs
This commit is contained in:
yujonglee 2026-10-07 15:05:43 -07:00 • committed by GitHub
parent 2803a16b36
commit 2aaa0b5d5c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
58 changed files with 1302 additions and 1192 deletions

View file

@ -1,8 +1,5 @@
use pyo3::{prelude::*, types::PyDict};
/// The caller's own object for a public argument: the keyword if given, even an explicit
/// `None`, else the bound request's attribute. Every reader of a public Python call uses
/// this rule, so the callbacks and the provider see one object per argument.
pub fn lookup<'py>(
kwargs: &Bound<'py, PyDict>,
request: &Bound<'py, PyAny>,
@ -11,6 +8,9 @@ pub fn lookup<'py>(
if let Some(value) = kwargs.get_item(name)? {
return Ok(Some(value));
}
if let Ok(bound) = request.cast::<PyDict>() {
return bound.get_item(name);
}
request.getattr_opt(name)
}
@ -48,4 +48,32 @@ kwargs = {'api_key': key, 'api_base': None}
assert!(find("model").is_none());
});
}
#[rstest::rstest]
#[case::prepared_value("{'api_key': 'replacement'}", Some("replacement"))]
#[case::explicit_none("{'api_key': None}", None)]
#[case::bound_fallback("{}", Some("original"))]
fn prepared_mapping_overrides_bound_values(
#[case] source: &str,
#[case] expected: Option<&str>,
) {
crate::initialize_python();
Python::attach(|py| {
let bound = PyDict::new(py);
bound.set_item("api_key", "original").unwrap();
let source = std::ffi::CString::new(source).unwrap();
let prepared = py
.eval(&source, None, None)
.unwrap()
.cast_into::<PyDict>()
.unwrap();
let value = lookup(&prepared, bound.as_any(), "api_key")
.unwrap()
.unwrap();
assert_eq!(
value.extract::<Option<String>>().unwrap().as_deref(),
expected
);
});
}
}

View file

@ -14,7 +14,9 @@ use litellm_secrets::source::SecretSource;
use std::sync::Arc;
pub use litellm_inference::RouteError as Error;
pub use types::{MessagesCall, MessagesCallResponse, MessagesShaping, messages_body};
pub use types::{
MessagesCall, MessagesCallResponse, MessagesSettings, MessagesShaping, messages_body,
};
#[derive(Clone)]
pub struct MessagesRoute {

View file

@ -80,12 +80,13 @@ fn prepare_provider_request(
let sanitized = config.shape_request(
MessagesRequest { model, ..body },
shaping.reasoning_auto_summary,
shaping.settings.reasoning_auto_summary,
)?;
let trimmed = without_additional_drop_params(sanitized, &shaping.additional_drop_params)?;
let trimmed =
without_additional_drop_params(sanitized, &shaping.settings.additional_drop_params)?;
let transformed = config.transform_anthropic_messages_request(
trimmed,
&MessagesTransformContext::new(shaping.capabilities, shaping.drop_params),
&MessagesTransformContext::new(shaping.capabilities, shaping.settings.drop_params),
)?;
let scoped =
@ -148,7 +149,7 @@ mod tests {
use serde_json::{Map, Value, json};
use super::*;
use crate::MessagesShaping;
use crate::{MessagesSettings, MessagesShaping};
#[fixture]
fn shaping() -> MessagesShaping {
@ -310,10 +311,13 @@ mod tests {
)
};
let shaping = MessagesShaping {
additional_drop_params: additional_drop_params
.iter()
.map(ToString::to_string)
.collect(),
settings: MessagesSettings {
additional_drop_params: additional_drop_params
.iter()
.map(ToString::to_string)
.collect(),
..shaping.settings
},
..shaping
};
assert_eq!(
@ -394,8 +398,11 @@ mod tests {
#[rstest]
fn dropped_thinking_display_is_not_restored_by_auto_summary(shaping: MessagesShaping) {
let shaping = MessagesShaping {
reasoning_auto_summary: true,
additional_drop_params: vec!["thinking.display".to_string()],
settings: MessagesSettings {
reasoning_auto_summary: true,
additional_drop_params: vec!["thinking.display".to_string()],
..shaping.settings
},
..shaping
};
assert_eq!(
@ -420,7 +427,10 @@ mod tests {
#[rstest]
fn dropping_an_invalid_metadata_user_id_does_not_skip_its_validation(shaping: MessagesShaping) {
let shaping = MessagesShaping {
additional_drop_params: vec!["metadata.user_id".to_string()],
settings: MessagesSettings {
additional_drop_params: vec!["metadata.user_id".to_string()],
..shaping.settings
},
..shaping
};
assert!(matches!(

View file

@ -38,6 +38,12 @@ pub type MessagesCallResponse =
pub struct MessagesShaping {
#[serde(default)]
pub capabilities: MessagesModelCapabilities,
#[serde(flatten)]
pub settings: MessagesSettings,
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct MessagesSettings {
#[serde(default)]
pub drop_params: bool,
#[serde(default)]
@ -58,16 +64,25 @@ mod tests {
#[case::nothing_projected(json!({}), MessagesShaping::default())]
#[case::only_drop_params(
json!({"drop_params": true}),
MessagesShaping { drop_params: true, ..MessagesShaping::default() },
MessagesShaping {
settings: MessagesSettings { drop_params: true, ..MessagesSettings::default() },
..MessagesShaping::default()
},
)]
#[case::only_reasoning_auto_summary(
json!({"reasoning_auto_summary": true}),
MessagesShaping { reasoning_auto_summary: true, ..MessagesShaping::default() },
MessagesShaping {
settings: MessagesSettings { reasoning_auto_summary: true, ..MessagesSettings::default() },
..MessagesShaping::default()
},
)]
#[case::only_additional_drop_params(
json!({"additional_drop_params": ["tools[*].input_examples"]}),
MessagesShaping {
additional_drop_params: vec!["tools[*].input_examples".to_string()],
settings: MessagesSettings {
additional_drop_params: vec!["tools[*].input_examples".to_string()],
..MessagesSettings::default()
},
..MessagesShaping::default()
},
)]
@ -98,6 +113,11 @@ mod tests {
"additional_drop_params": ["metadata.user_id", "thinking"]
}),
MessagesShaping {
settings: MessagesSettings {
drop_params: true,
reasoning_auto_summary: true,
additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()],
},
capabilities: MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
@ -115,9 +135,6 @@ mod tests {
max: false,
},
},
drop_params: true,
reasoning_auto_summary: true,
additional_drop_params: vec!["metadata.user_id".to_string(), "thinking".to_string()],
},
)]
fn shaping_deserializes_with_defaults_for_absent_fields(
@ -126,5 +143,22 @@ mod tests {
) {
let shaping: MessagesShaping = serde_json::from_value(projected).unwrap();
assert_eq!(shaping, expected);
let serialized = serde_json::to_value(&shaping).unwrap();
assert_eq!(
serialized["drop_params"],
json!(expected.settings.drop_params)
);
assert_eq!(
serialized["reasoning_auto_summary"],
json!(expected.settings.reasoning_auto_summary)
);
assert_eq!(
serialized["additional_drop_params"],
json!(expected.settings.additional_drop_params)
);
assert_eq!(
serialized["capabilities"],
serde_json::to_value(expected.capabilities).unwrap()
);
}
}

View file

@ -380,12 +380,14 @@ async fn the_request_context_carries_the_shaped_params_without_model_or_messages
let host = RecordingHost::passthrough(authenticated(
MessagesCall {
shaping: MessagesShaping {
settings: MessagesSettings {
drop_params: true,
..MessagesSettings::default()
},
capabilities: AnthropicModelCapabilities {
supports_sampling_params: false,
..AnthropicModelCapabilities::default()
},
drop_params: true,
..MessagesShaping::default()
},
..with_fields(call, json!({"temperature": 0.2}))
},

View file

@ -5,7 +5,7 @@ use std::{
use litellm_http::{HttpSettings, Resolution};
use litellm_inference_messages::{
Error, MessagesCall, MessagesShaping,
Error, MessagesCall, MessagesSettings, MessagesShaping,
route::{Messages, MessagesMachine, MessagesOutput},
};
use litellm_inference_testing::RecordingSecrets;

View file

@ -239,7 +239,10 @@ async fn additional_drop_params_remove_fields_before_sending(call: MessagesCall)
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
additional_drop_params: vec!["temperature".into()],
settings: MessagesSettings {
additional_drop_params: vec!["temperature".into()],
..MessagesSettings::default()
},
..MessagesShaping::default()
},
..with_fields(call, json!({"temperature": 0.5, "top_k": 3}))
@ -390,8 +393,10 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
capabilities,
drop_params,
..MessagesShaping::default()
settings: MessagesSettings {
drop_params,
..MessagesSettings::default()
},
},
body: call.body.clone(),
custom_llm_provider: call.custom_llm_provider.clone(),
@ -436,13 +441,15 @@ async fn reasoning_auto_summary_marks_active_thinking_on_the_wire(
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
settings: MessagesSettings {
reasoning_auto_summary: true,
..MessagesSettings::default()
},
capabilities: MessagesModelCapabilities {
supports_reasoning: true,
supports_adaptive_thinking: true,
..MessagesModelCapabilities::default()
},
reasoning_auto_summary: true,
..MessagesShaping::default()
},
..call
},
@ -634,7 +641,10 @@ async fn system_message_folding_is_selected_by_the_provider(
api_key: Some("sk-azure".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
additional_drop_params: drop_params.iter().map(ToString::to_string).collect(),
settings: MessagesSettings {
additional_drop_params: drop_params.iter().map(ToString::to_string).collect(),
..call.shaping.settings
},
..call.shaping
},
..call
@ -695,7 +705,10 @@ async fn provider_validation_runs_before_caller_parameter_removal(
api_key: Some("sk-test".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
additional_drop_params: vec!["metadata".into()],
settings: MessagesSettings {
additional_drop_params: vec!["metadata".into()],
..call.shaping.settings
},
..call.shaping
},
..call

View file

@ -6,6 +6,7 @@ use std::{
use litellm_auth::InputSource;
use litellm_host_python::{from_py, from_py_argument};
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
use serde::de::DeserializeOwned;
use serde_json::{Map, Value};
/// The keyword arguments every value route shares, validated at the Python boundary.
@ -25,18 +26,6 @@ pub(crate) fn messages_argument(value: &Bound<'_, PyAny>) -> PyResult<Vec<Value>
}
}
pub(crate) fn optional_params_argument(
value: &Bound<'_, PyAny>,
) -> PyResult<Option<Map<String, Value>>> {
optional_object("optional_params", value)
}
pub(crate) fn extra_headers_argument(
value: &Bound<'_, PyAny>,
) -> PyResult<Option<Map<String, Value>>> {
optional_object("extra_headers", value)
}
fn required_object(name: &'static str, value: Value) -> PyResult<Map<String, Value>> {
match value {
Value::Object(values) => Ok(values),
@ -71,6 +60,48 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py<PyAny>) -> PyRe
.extract()
}
pub(crate) fn required_field<'py>(
fields: &Bound<'py, PyDict>,
name: &str,
) -> PyResult<Bound<'py, PyAny>> {
fields
.get_item(name)?
.ok_or_else(|| PyValueError::new_err(format!("{name} is required")))
}
pub(crate) fn optional_field<T: DeserializeOwned>(
fields: &Bound<'_, PyDict>,
name: &str,
) -> PyResult<Option<T>> {
fields
.get_item(name)?
.map(|value| from_py_argument(&value))
.transpose()
.map(Option::flatten)
}
pub(crate) fn optional_object_field(
fields: &Bound<'_, PyDict>,
name: &'static str,
) -> PyResult<Option<Map<String, Value>>> {
fields
.get_item(name)?
.map(|value| optional_object(name, &value))
.transpose()
.map(Option::flatten)
}
pub(crate) fn value_route_options(fields: &Bound<'_, PyDict>) -> PyResult<RouteOptions> {
Ok(RouteOptions {
model: from_py_argument(&required_field(fields, "model")?)?,
api_key: optional_field(fields, "api_key")?,
api_base: optional_field(fields, "api_base")?,
custom_llm_provider: optional_field(fields, "custom_llm_provider")?,
extra_headers: optional_object_field(fields, "extra_headers")?,
timeout: optional_timeout(optional_field(fields, "timeout_seconds")?),
})
}
pub(crate) fn project_optional_fields(
kwargs: &Bound<'_, PyDict>,
names: &[&str],
@ -244,15 +275,15 @@ mod tests {
let params = py.eval(c"{'temperature': 0.2}", None, None).unwrap();
assert_eq!(
optional_params_argument(&params).unwrap(),
optional_object("optional_params", &params).unwrap(),
Some(required_object("optional_params", json!({"temperature": 0.2})).unwrap())
);
assert_eq!(
optional_params_argument(&py.None().into_bound(py)).unwrap(),
optional_object("optional_params", &py.None().into_bound(py)).unwrap(),
None
);
assert_eq!(
extra_headers_argument(&py.None().into_bound(py)).unwrap(),
optional_object("extra_headers", &py.None().into_bound(py)).unwrap(),
None
);
});

View file

@ -1,14 +1,13 @@
use crate::execution::{run_async, run_sync};
use litellm_host_python::from_py_argument;
use litellm_inference_transcription::{
AudioTranscriptionRoute, Error, types::AudioTranscriptionRequest,
};
use pyo3::{prelude::*, types::PyDict};
use pyo3::prelude::*;
use serde_json::{Map, Value};
use crate::{
errors::route_error_to_pyerr,
marshal::{RouteOptions, extra_headers_argument, optional_params_argument, optional_timeout},
marshal::{RouteOptions, optional_object_field, required_field, value_route_options},
};
async fn execute(
@ -41,81 +40,38 @@ async fn execute(
}
#[pyfunction]
#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[expect(
clippy::too_many_arguments,
reason = "one parameter per Python keyword"
)]
pub(crate) fn transcription(
py: Python<'_>,
model: String,
#[pyo3(from_py_with = from_py_argument)] audio: Value,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option<Map<String, Value>>,
#[pyo3(from_py_with = optional_params_argument)] optional_params: Option<Map<String, Value>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let options = RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let http = crate::http::provider_client(py, &PyDict::new(py), false)?;
pub(crate) fn transcription(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<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)?;
let optional_params =
optional_object_field(&call.bound, "optional_params")?.unwrap_or_default();
let http = crate::http::provider_client(py, &call.kwargs, false)?;
let secrets = crate::secrets::source(py)?;
run_sync(
py,
execute(
http,
secrets,
audio,
optional_params.unwrap_or_default(),
options,
),
execute(http, secrets, audio, optional_params, options),
route_error_to_pyerr,
)
}
#[pyfunction]
#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[expect(
clippy::too_many_arguments,
reason = "one parameter per Python keyword"
)]
pub(crate) fn atranscription<'py>(
py: Python<'py>,
model: String,
#[pyo3(from_py_with = from_py_argument)] audio: Value,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option<Map<String, Value>>,
#[pyo3(from_py_with = optional_params_argument)] optional_params: Option<Map<String, Value>>,
timeout_seconds: Option<f64>,
call: Bound<'py, PyAny>,
) -> PyResult<Bound<'py, PyAny>> {
let options = RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let http = crate::http::provider_client(py, &PyDict::new(py), true)?;
let call = super::NativeCall::extract(&call)?;
let audio: Value =
litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?;
let options = value_route_options(&call.bound)?;
let optional_params =
optional_object_field(&call.bound, "optional_params")?.unwrap_or_default();
let http = crate::http::provider_client(py, &call.kwargs, true)?;
let secrets = crate::secrets::source(py)?;
run_async(
py,
execute(
http,
secrets,
audio,
optional_params.unwrap_or_default(),
options,
),
execute(http, secrets, audio, optional_params, options),
route_error_to_pyerr,
)
}

View file

@ -11,8 +11,7 @@ use serde_json::{Map, Value};
use crate::{
errors::route_error_to_pyerr,
marshal::{
RouteOptions, extra_headers_argument, messages_argument, optional_params_argument,
optional_timeout,
RouteOptions, messages_argument, optional_object_field, required_field, value_route_options,
},
};
@ -50,81 +49,36 @@ async fn execute(
}
#[pyfunction]
#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))]
#[expect(
clippy::too_many_arguments,
reason = "one parameter per Python keyword"
)]
pub(crate) fn chat_completions(
py: Python<'_>,
model: String,
#[pyo3(from_py_with = messages_argument)] messages: Vec<Value>,
#[pyo3(from_py_with = optional_params_argument)] optional_params: Option<Map<String, Value>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option<Map<String, Value>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let options = RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let http = crate::http::provider_client(py, &PyDict::new(py), false)?;
pub(crate) fn chat_completions(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<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.unwrap_or_default(),
options,
),
execute(http, secrets, messages, optional_params, options),
route_error_to_pyerr,
)
}
#[pyfunction]
#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))]
#[expect(
clippy::too_many_arguments,
reason = "one parameter per Python keyword"
)]
pub(crate) fn achat_completions<'py>(
py: Python<'py>,
model: String,
#[pyo3(from_py_with = messages_argument)] messages: Vec<Value>,
#[pyo3(from_py_with = optional_params_argument)] optional_params: Option<Map<String, Value>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option<Map<String, Value>>,
timeout_seconds: Option<f64>,
call: Bound<'py, PyAny>,
) -> PyResult<Bound<'py, PyAny>> {
let options = RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
let http = crate::http::provider_client(py, &PyDict::new(py), true)?;
let call = super::NativeCall::extract(&call)?;
let messages: Vec<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.unwrap_or_default(),
options,
),
execute(http, secrets, messages, optional_params, options),
route_error_to_pyerr,
)
}
@ -184,21 +138,13 @@ fn run_public(
}
#[pyfunction]
pub(crate) fn completion(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_public(py, request, args, kwargs, false)
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<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_public(py, request, args, kwargs, true)
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,58 +1,49 @@
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use pyo3::prelude::*;
use crate::errors::RustBridgeDeclined;
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn embedding(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
pub(crate) fn embedding(call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
drop(super::NativeCall::extract(&call)?);
Err(RustBridgeDeclined::new_err(
"native embeddings route is not implemented",
))
}
#[pyfunction]
#[pyo3(signature = (request, args, kwargs))]
pub(crate) fn aembedding(
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
drop((request, args, kwargs));
Err(RustBridgeDeclined::new_err(
"native embeddings route is not implemented",
))
pub(crate) fn aembedding(call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
embedding(call)
}
#[cfg(test)]
mod tests {
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use pyo3::{prelude::*, types::PyDict};
use rstest::rstest;
use crate::errors::RustBridgeDeclined;
#[test]
fn both_entrypoints_decline_before_provider_execution() {
#[rstest]
#[case::sync(false)]
#[case::asynchronous(true)]
fn both_entrypoints_decline_before_provider_execution(#[case] asynchronous: bool) {
Python::initialize();
Python::attach(|py| {
let request = PyDict::new(py);
let args = PyTuple::empty(py);
let kwargs = PyDict::new(py);
for entrypoint in [super::embedding, super::aembedding] {
let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone())
.expect_err("native embeddings must decline until a route machine exists");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
let locals = PyDict::new(py);
py.run(
c"from types import SimpleNamespace
call = SimpleNamespace(args=(), kwargs={}, bound={'model':'test-model','input':'hello'})",
Some(&locals),
Some(&locals),
)
.unwrap();
let call = locals.get_item("call").unwrap().unwrap();
let error = if asynchronous {
super::aembedding(call)
} else {
super::embedding(call)
}
.expect_err("native embeddings must decline until a route machine exists");
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
});
}
}

View file

@ -86,6 +86,9 @@ impl InferenceHost {
if let Some(value) = lookup(arguments, request, name)? {
return Ok((!value.is_none()).then_some(value));
}
if request.is_instance_of::<PyDict>() {
return Ok(None);
}
let parameter = request
.getattr("parameters")?
.call_method1("get", (name,))?;

View file

@ -5,9 +5,10 @@ use bytes::Bytes;
use litellm_host_python::{InvokeError, PythonBinding, from_py, lookup, to_py};
use litellm_http::transport::Error as TransportError;
use litellm_inference_messages::{
Error, MessagesCall, MessagesShaping, messages_body,
Error, MessagesCall, MessagesSettings, MessagesShaping, messages_body,
route::{Messages, MessagesStreamHead},
};
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
use litellm_llms_types::headers::ProviderSpecificHeaders;
use pyo3::{
exceptions::{PyException, PyValueError},
@ -188,18 +189,24 @@ impl MessagesPythonHost {
custom_llm_provider: Option<&str>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<MessagesShaping> {
let projected = py.import(ROUTE_HOST_MODULE)?.getattr("shaping")?.call1((
model,
custom_llm_provider,
arguments,
))?;
from_py(&projected)
let module = py.import(ROUTE_HOST_MODULE)?;
let capabilities: MessagesModelCapabilities = from_py(
&py.import("litellm.rust_bridge.model_capabilities")?
.getattr("anthropic_model_capabilities")?
.call1((model, custom_llm_provider))?,
)?;
let settings: MessagesSettings =
from_py(&module.getattr("settings")?.call1((arguments,))?)?;
Ok(MessagesShaping {
capabilities,
settings,
})
}
fn provider(&self, py: Python<'_>) -> String {
self.request
.bind(py)
.getattr("custom_llm_provider")
.get_item("custom_llm_provider")
.and_then(|value| value.extract::<Option<String>>())
.ok()
.flatten()

View file

@ -64,21 +64,13 @@ fn run_messages(
}
#[pyfunction]
pub(crate) fn messages(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_messages(py, request, args, kwargs, false)
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)
}
#[pyfunction]
pub(crate) fn amessages(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_messages(py, request, args, kwargs, true)
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)
}

View file

@ -14,9 +14,35 @@ use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol}
use litellm_host_python::{HookChain, PythonBinding, PythonCallHooks, PythonHostCalls};
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
types::{PyDict, PyMapping, PyTuple},
};
struct NativeCall<'py> {
args: Bound<'py, PyTuple>,
kwargs: Bound<'py, PyDict>,
bound: Bound<'py, PyDict>,
}
impl<'py> NativeCall<'py> {
fn extract(call: &Bound<'py, PyAny>) -> PyResult<Self> {
Ok(Self {
args: call.getattr("args")?.cast_into()?,
kwargs: mapping_dict(&call.getattr("kwargs")?)?,
bound: mapping_dict(&call.getattr("bound")?)?,
})
}
}
fn mapping_dict<'py>(value: &Bound<'py, PyAny>) -> PyResult<Bound<'py, PyDict>> {
if let Ok(dict) = value.cast::<PyDict>() {
return Ok(dict.clone());
}
let mapping = value.cast::<PyMapping>()?;
let dict = PyDict::new(value.py());
dict.update(mapping)?;
Ok(dict)
}
fn call_hooks(
py: Python<'_>,
operation: LoggingOperation,
@ -74,40 +100,30 @@ mod tests {
types::{PyDict, PyList},
};
#[test]
fn sync_and_async_route_signatures_match_the_python_contract() {
Python::initialize();
Python::attach(|py| {
let module = crate::native_module(py);
let routes = [
(
"transcription",
"atranscription",
"(model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)",
),
(
"chat_completions",
"achat_completions",
"(model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)",
),
];
for (sync_name, async_name, expected) in routes {
let sync_signature: String = module
.getattr(sync_name)
.and_then(|function| function.getattr("__text_signature__"))
.and_then(|signature| signature.extract())
.expect("sync signature should be available");
let async_signature: String = module
.getattr(async_name)
.and_then(|function| function.getattr("__text_signature__"))
.and_then(|signature| signature.extract())
.expect("async signature should be available");
assert_eq!(sync_signature, expected);
assert_eq!(async_signature, expected);
}
});
fn value_call<'py>(
py: Python<'py>,
payload_name: &str,
payload: &Bound<'py, PyAny>,
kwargs: Option<&Bound<'py, PyDict>>,
) -> Bound<'py, PyAny> {
let fields = PyDict::new(py);
fields.set_item("model", "model").unwrap();
fields.set_item(payload_name, payload).unwrap();
if let Some(kwargs) = kwargs {
fields.update(kwargs.as_mapping()).unwrap();
}
let attributes = PyDict::new(py);
attributes
.set_item("args", pyo3::types::PyTuple::empty(py))
.unwrap();
attributes.set_item("kwargs", &fields).unwrap();
attributes.set_item("bound", &fields).unwrap();
py.import("types")
.unwrap()
.getattr("SimpleNamespace")
.unwrap()
.call((), Some(&attributes))
.unwrap()
}
#[test]
@ -138,7 +154,9 @@ value = Broken()
for name in ["chat_completions", "achat_completions"] {
let error = module
.getattr(name)
.and_then(|function| function.call1(("model", &broken)))
.and_then(|function| {
function.call1((value_call(py, "messages", &broken, None),))
})
.expect_err("route should reject a value it cannot convert");
assert!(
@ -158,11 +176,15 @@ value = Broken()
let invalid_messages = PyDict::new(py);
let sync_chat_error = module
.getattr("chat_completions")
.and_then(|function| function.call1(("model", &invalid_messages)))
.and_then(|function| {
function.call1((value_call(py, "messages", &invalid_messages, None),))
})
.expect_err("sync chat should reject a non-list messages value");
let async_chat_error = module
.getattr("achat_completions")
.and_then(|function| function.call1(("model", &invalid_messages)))
.and_then(|function| {
function.call1((value_call(py, "messages", &invalid_messages, None),))
})
.expect_err("async chat should reject a non-list messages value");
assert_eq!(
@ -180,11 +202,15 @@ value = Broken()
let sync_error = module
.getattr("transcription")
.and_then(|function| function.call(("model", &audio), Some(&kwargs)))
.and_then(|function| {
function.call1((value_call(py, "audio", &audio, Some(&kwargs)),))
})
.expect_err("sync route should reject non-dict extra_headers");
let async_error = module
.getattr("atranscription")
.and_then(|function| function.call(("model", &audio), Some(&kwargs)))
.and_then(|function| {
function.call1((value_call(py, "audio", &audio, Some(&kwargs)),))
})
.expect_err("async route should reject non-dict extra_headers");
assert_eq!(
@ -213,7 +239,12 @@ value = Broken()
let error = module
.getattr("chat_completions")
.and_then(|function| {
function.call(("model", &invalid_messages), Some(&chat_kwargs))
function.call1((value_call(
py,
"messages",
&invalid_messages,
Some(&chat_kwargs),
),))
})
.expect_err("messages should be validated first");
assert_eq!(error.to_string(), "ValueError: messages must be a list");
@ -221,7 +252,14 @@ value = Broken()
let valid_messages = PyList::empty(py);
let error = module
.getattr("chat_completions")
.and_then(|function| function.call(("model", &valid_messages), Some(&chat_kwargs)))
.and_then(|function| {
function.call1((value_call(
py,
"messages",
&valid_messages,
Some(&chat_kwargs),
),))
})
.expect_err("optional_params should be validated before headers");
assert_eq!(
error.to_string(),
@ -237,7 +275,12 @@ value = Broken()
let error = module
.getattr("transcription")
.and_then(|function| {
function.call(("model", &invalid_payload), Some(&headers_kwargs))
function.call1((value_call(
py,
"audio",
&invalid_payload,
Some(&headers_kwargs),
),))
})
.expect_err("payload should be validated before headers");
assert!(!error.to_string().contains("extra_headers"));
@ -265,11 +308,15 @@ value = Broken()
let omitted_error = module
.getattr("chat_completions")
.and_then(|function| function.call(("model", &messages), Some(&omitted)))
.and_then(|function| {
function.call1((value_call(py, "messages", &messages, Some(&omitted)),))
})
.expect_err("omitted optional_params should reach header validation");
let explicit_error = module
.getattr("chat_completions")
.and_then(|function| function.call(("model", &messages), Some(&explicit)))
.and_then(|function| {
function.call1((value_call(py, "messages", &messages, Some(&explicit)),))
})
.expect_err("None optional_params should reach header validation");
assert_eq!(
omitted_error.to_string(),

View file

@ -81,23 +81,15 @@ fn project_provider_defaults(snapshot: &Snapshot<'_>) -> PyResult<OcrSettings> {
}
#[pyfunction]
pub(crate) fn ocr(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_ocr(py, request, args, kwargs, false)
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)
}
#[pyfunction]
pub(crate) fn aocr(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_ocr(py, request, args, kwargs, true)
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)
}
#[pyfunction]

View file

@ -105,23 +105,15 @@ fn run_public(
}
#[pyfunction]
pub(crate) fn responses(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_public(py, request, args, kwargs, false)
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<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_public(py, request, args, kwargs, true)
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]

View file

@ -1,7 +1,7 @@
use std::time::Duration;
use litellm_config::Config;
use litellm_inference_messages::MessagesShaping;
use litellm_inference_messages::{MessagesSettings, MessagesShaping};
use litellm_router::{Deployment, Router};
use rstest::rstest;
@ -74,8 +74,11 @@ fn programmatic_deployments_preserve_overrides_and_last_entry_wins() {
custom_llm_provider: Some("test-provider".into()),
timeout: Some(Duration::from_secs(7)),
shaping: MessagesShaping {
drop_params: true,
additional_drop_params: vec!["metadata.test".into()],
settings: MessagesSettings {
drop_params: true,
additional_drop_params: vec!["metadata.test".into()],
..MessagesSettings::default()
},
..Default::default()
},
};

View file

@ -1,6 +1,5 @@
import inspect
from collections.abc import Awaitable, Callable, Coroutine, Mapping
from types import MappingProxyType
from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
from litellm import main
@ -8,13 +7,13 @@ from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.chat_completions.entrypoints import (
NATIVE_ACOMPLETION,
NATIVE_COMPLETION,
LiteLLMChatCompletionsRequest,
)
from litellm.rust_bridge.dispatch import PublicDispatch, call_hook
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.public_call import (
NativeCall,
bind,
optional_bool,
optional_mapping,
native_call,
native_call_hook,
optional_sequence,
optional_str,
signature,
@ -51,33 +50,22 @@ _ACOMPLETION: Final = signature(_PYTHON_ACOMPLETION)
def _public_request(
legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> LiteLLMChatCompletionsRequest | None:
) -> NativeCall | None:
fields: Final = bind(legacy, args, kwargs)
if fields is None:
return None
model: Final = fields.get("model")
messages: Final = optional_sequence(fields.get("messages"))
extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({})
if not isinstance(model, str) or messages is None:
return None
return LiteLLMChatCompletionsRequest(
model=model,
messages=messages,
stream=optional_bool(fields.get("stream")),
api_key=optional_str(fields.get("api_key")),
api_base=optional_str(extra.get("api_base")) or optional_str(fields.get("base_url")),
custom_llm_provider=optional_str(extra.get("custom_llm_provider")),
extra_headers=optional_mapping(fields.get("extra_headers")),
kwargs=extra,
parameters=MappingProxyType({name: value for name, value in fields.items() if name != "kwargs"}),
)
return native_call(args, kwargs, fields)
def _context(request: LiteLLMChatCompletionsRequest) -> RouteContext:
def _context(request: NativeCall) -> RouteContext:
return RouteContext(
Route.CHAT_COMPLETIONS,
provider=request.custom_llm_provider,
model=request.model,
provider=optional_str(request.bound.get("custom_llm_provider")),
model=str(request.bound["model"]),
)
@ -105,7 +93,7 @@ def completion(
kwargs,
python=python,
binding=NATIVE_COMPLETION,
native=call_hook,
native=native_call_hook,
)
@ -116,7 +104,7 @@ async def acompletion(*args: object, **kwargs: object) -> ChatResult: # kwargs-
kwargs,
python=python,
binding=NATIVE_ACOMPLETION,
native=call_hook,
native=native_call_hook,
)

View file

@ -1,17 +1,15 @@
import inspect
from collections.abc import Awaitable, Callable, Coroutine, Mapping
from types import MappingProxyType
from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
from litellm import main
from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.dispatch import PublicDispatch, call_hook
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.embeddings.entrypoints import (
NATIVE_AEMBEDDING,
NATIVE_EMBEDDING,
LiteLLMEmbeddingRequest,
)
from litellm.rust_bridge.public_call import bind, optional_mapping, optional_str, signature
from litellm.rust_bridge.public_call import NativeCall, bind, native_call, native_call_hook, optional_str, signature
from litellm.types.utils import EmbeddingResponse
__all__ = ("aembedding", "embedding")
@ -30,28 +28,24 @@ _EMBEDDING_SIGNATURE: Final = signature(_PYTHON_EMBEDDING)
def _public_request(
legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> LiteLLMEmbeddingRequest | None:
) -> NativeCall | None:
fields: Final = bind(legacy, args, kwargs)
if fields is None:
return None
model: Final = fields.get("model")
if not isinstance(model, str):
return None
extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({})
return LiteLLMEmbeddingRequest(
model=model,
input=fields.get("input"),
api_key=optional_str(fields.get("api_key")),
api_base=optional_str(fields.get("api_base")),
custom_llm_provider=optional_str(fields.get("custom_llm_provider")),
kwargs=extra,
return native_call(args, kwargs, fields)
def _context(request: NativeCall) -> RouteContext:
return RouteContext(
Route.EMBEDDINGS,
provider=optional_str(request.bound.get("custom_llm_provider")),
model=str(request.bound["model"]),
)
def _context(request: LiteLLMEmbeddingRequest) -> RouteContext:
return RouteContext(Route.EMBEDDINGS, provider=request.custom_llm_provider, model=request.model)
_DISPATCH: Final = PublicDispatch(
route=Route.EMBEDDINGS,
request=lambda args, kwargs: _public_request(_EMBEDDING_SIGNATURE, args, kwargs),
@ -75,7 +69,7 @@ def embedding(
kwargs,
python=_PYTHON_EMBEDDING,
binding=NATIVE_EMBEDDING,
native=call_hook,
native=native_call_hook,
)
@ -85,7 +79,7 @@ async def aembedding(*args: object, **kwargs: object) -> EmbeddingResponse: # k
kwargs,
python=_PYTHON_AEMBEDDING,
binding=NATIVE_AEMBEDDING,
native=call_hook,
native=native_call_hook,
)

View file

@ -6,6 +6,7 @@ import httpx
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
from litellm.rust_bridge import runtime
from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.public_call import NativeCall
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.rust_bridge.transcription.native import (
NATIVE_ATRANSCRIPTION,
@ -52,18 +53,18 @@ class BedrockAudioTranscriptionRustDispatch:
timeout: float | httpx.Timeout | None,
) -> TranscriptionResponse:
def native(rust: RustTranscription) -> TranscriptionResponse:
return TranscriptionResponse(
**rust(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_to_seconds(timeout),
)
)
fields: Final = {
"model": model,
"audio": self._audio_payload(audio_file),
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": timeout_to_seconds(timeout),
}
call: Final = NativeCall(args=(), kwargs=fields, bound=fields)
return TranscriptionResponse(**rust(call))
return runtime.run(
RouteContext(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model),
@ -85,18 +86,18 @@ class BedrockAudioTranscriptionRustDispatch:
timeout: float | httpx.Timeout | None,
) -> TranscriptionResponse:
async def native(rust: RustAtranscription) -> TranscriptionResponse:
return TranscriptionResponse(
**await rust(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_to_seconds(timeout),
)
)
fields: Final = {
"model": model,
"audio": self._audio_payload(audio_file),
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": timeout_to_seconds(timeout),
}
call: Final = NativeCall(args=(), kwargs=fields, bound=fields)
return TranscriptionResponse(**await rust(call))
return await runtime.arun(
RouteContext(Route.TRANSCRIPTION, provider=custom_llm_provider, model=model),

View file

@ -1,22 +1,21 @@
import inspect
from collections.abc import AsyncIterator, Awaitable, Callable, Coroutine, Iterator, Mapping
from types import MappingProxyType
from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
from litellm.exceptions import BadRequestError
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.anthropic.pass_through.messages import handler as main
from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.dispatch import PublicDispatch, call_hook
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.messages.entrypoints import (
NATIVE_AMESSAGES,
NATIVE_MESSAGES,
LiteLLMMessagesRequest,
)
from litellm.rust_bridge.public_call import (
NativeCall,
bind,
optional_bool,
optional_mapping,
native_call,
native_call_hook,
optional_sequence,
optional_str,
signature,
@ -52,7 +51,7 @@ _AMESSAGES: Final = signature(_PYTHON_AMESSAGES)
def _public_request(
legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> LiteLLMMessagesRequest | None:
) -> NativeCall | None:
fields: Final = bind(legacy, args, kwargs)
if fields is None:
return None
@ -61,30 +60,21 @@ def _public_request(
max_tokens: Final = fields.get("max_tokens")
if not isinstance(model, str) or messages is None or not isinstance(max_tokens, int):
return None
return LiteLLMMessagesRequest(
model=model,
messages=messages,
max_tokens=max_tokens,
stream=optional_bool(fields.get("stream")),
api_key=optional_str(fields.get("api_key")),
api_base=optional_str(fields.get("api_base")),
custom_llm_provider=optional_str(fields.get("custom_llm_provider")),
kwargs=optional_mapping(fields.get("kwargs")) or MappingProxyType({}),
)
return native_call(args, kwargs, fields)
def _resolved_provider(request: LiteLLMMessagesRequest) -> str | None:
def _resolved_provider(request: NativeCall) -> str | None:
try:
return get_llm_provider(request.model, request.custom_llm_provider)[1]
return get_llm_provider(str(request.bound["model"]), optional_str(request.bound.get("custom_llm_provider")))[1]
except BadRequestError:
return request.custom_llm_provider
return optional_str(request.bound.get("custom_llm_provider"))
def _context(request: LiteLLMMessagesRequest) -> RouteContext:
def _context(request: NativeCall) -> RouteContext:
return RouteContext(
Route.MESSAGES,
provider=_resolved_provider(request),
model=request.model,
model=str(request.bound["model"]),
)
@ -112,7 +102,7 @@ def anthropic_messages_handler(
kwargs,
python=python,
binding=NATIVE_MESSAGES,
native=call_hook,
native=native_call_hook,
)
@ -123,7 +113,7 @@ async def anthropic_messages(*args: object, **kwargs: object) -> MessagesResult:
kwargs,
python=python,
binding=NATIVE_AMESSAGES,
native=call_hook,
native=native_call_hook,
)

View file

@ -1,4 +1,5 @@
from collections.abc import Coroutine, Mapping
from types import MappingProxyType
from typing import Final
import httpx
@ -6,8 +7,9 @@ import httpx
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge import runtime
from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.dispatch import PublicDispatch, call_hook
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR
from litellm.rust_bridge.public_call import NativeCall, native_call, native_call_hook, optional_str
__all__ = ("aocr", "ocr")
@ -21,30 +23,33 @@ def _bind_request(
custom_llm_provider: str | None = None,
extra_headers: dict[str, object] | None = None,
**kwargs: object, # kwargs-ok: public OCR accepts provider-specific options
) -> LiteLLMOcrRequest:
return LiteLLMOcrRequest(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
kwargs=kwargs,
) -> Mapping[str, object]:
return MappingProxyType(
{
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"kwargs": kwargs,
}
)
def _public_request(name: str, args: tuple[object, ...], kwargs: Mapping[str, object]) -> LiteLLMOcrRequest:
def _public_request(name: str, args: tuple[object, ...], kwargs: Mapping[str, object]) -> NativeCall:
try:
return _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation
fields: Final = _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation
return native_call(args, kwargs, fields)
except TypeError as error:
raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None
def _context(request: LiteLLMOcrRequest) -> RouteContext:
prefix, separator, _ = request.model.partition("/")
provider: Final = request.custom_llm_provider or (prefix if separator else None)
return RouteContext(Route.OCR, provider=provider, model=request.model)
def _context(request: NativeCall) -> RouteContext:
prefix, separator, _ = str(request.bound["model"]).partition("/")
provider: Final = optional_str(request.bound.get("custom_llm_provider")) or (prefix if separator else None)
return RouteContext(Route.OCR, provider=provider, model=str(request.bound["model"]))
_DISPATCH: Final = PublicDispatch(
@ -70,7 +75,7 @@ def ocr(
kwargs,
python=runtime.NO_PYTHON,
binding=NATIVE_OCR,
native=call_hook,
native=native_call_hook,
)
@ -80,5 +85,5 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr
kwargs,
python=runtime.NO_PYTHON,
binding=NATIVE_AOCR,
native=call_hook,
native=native_call_hook,
)

View file

@ -1,17 +1,22 @@
import inspect
from collections.abc import Awaitable, Callable, Coroutine, Mapping
from types import MappingProxyType
from typing import Final, TypeAlias, cast # noqa: TID251 # native binding selects a sync result or an async awaitable
from litellm.responses import main
from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator
from litellm.rust_bridge.catalog import Route, RouteContext
from litellm.rust_bridge.dispatch import PublicDispatch, call_hook
from litellm.rust_bridge.public_call import bind, optional_bool, optional_mapping, optional_str, signature
from litellm.rust_bridge.dispatch import PublicDispatch
from litellm.rust_bridge.public_call import (
NativeCall,
bind,
native_call,
native_call_hook,
optional_str,
signature,
)
from litellm.rust_bridge.responses.entrypoints import (
NATIVE_ARESPONSES,
NATIVE_RESPONSES,
LiteLLMResponsesRequest,
)
from litellm.types.llms.openai import ResponsesAPIResponse
@ -44,32 +49,21 @@ _ARESPONSES: Final = signature(_PYTHON_ARESPONSES)
def _public_request(
legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> LiteLLMResponsesRequest | None:
) -> NativeCall | None:
fields: Final = bind(legacy, args, kwargs)
if fields is None:
return None
model: Final = fields.get("model")
extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({})
if not isinstance(model, str):
return None
return LiteLLMResponsesRequest(
model=model,
input=fields.get("input"),
stream=optional_bool(fields.get("stream")),
api_key=optional_str(extra.get("api_key")),
api_base=optional_str(extra.get("api_base")) or optional_str(extra.get("base_url")),
custom_llm_provider=optional_str(fields.get("custom_llm_provider")),
extra_headers=optional_mapping(fields.get("extra_headers")),
kwargs=extra,
parameters=MappingProxyType({name: value for name, value in fields.items() if name != "kwargs"}),
)
return native_call(args, kwargs, fields)
def _context(request: LiteLLMResponsesRequest) -> RouteContext:
def _context(request: NativeCall) -> RouteContext:
return RouteContext(
Route.RESPONSES,
provider=request.custom_llm_provider,
model=request.model,
provider=optional_str(request.bound.get("custom_llm_provider")),
model=str(request.bound["model"]),
)
@ -97,7 +91,7 @@ def responses(
kwargs,
python=python,
binding=NATIVE_RESPONSES,
native=call_hook,
native=native_call_hook,
)
@ -108,7 +102,7 @@ async def aresponses(*args: object, **kwargs: object) -> ResponsesResult: # kwa
kwargs,
python=python,
binding=NATIVE_ARESPONSES,
native=call_hook,
native=native_call_hook,
)

View file

@ -6,11 +6,7 @@ import httpx
from pydantic import JsonValue
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest
from litellm.rust_bridge.embeddings.entrypoints import LiteLLMEmbeddingRequest
from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest
from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest
from litellm.rust_bridge.public_call import NativeCall
from litellm.rust_bridge.trace.generated.types import QueryScope, ReadQueryName, TraceScope
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
from litellm.types.llms.openai import ResponsesAPIResponse
@ -41,7 +37,9 @@ class NativeTraceStorage:
def __new__(cls, config: NativeTraceConfig) -> NativeTraceStorage: ...
def ensure_schema(self) -> Future[None]: ...
def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Future[None]: ...
def ingest(self, payload: bytes, content_type: str | None, tenant: Mapping[str, str], logs: bool = False) -> Future[int]: ...
def ingest(
self, payload: bytes, content_type: str | None, tenant: Mapping[str, str], logs: bool = False
) -> Future[int]: ...
def list_traces(
self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None, limit: int
) -> Future[JsonValue]: ...
@ -54,7 +52,9 @@ class NativeTraceStorage:
) -> Future[JsonValue]: ...
def query_sql(self, sql: str, scope: QueryScope, secret: str) -> Future[str]: ...
def query_help(self, scope: QueryScope, secret: str) -> Future[JsonValue]: ...
def query(self, query: ReadQueryName, parameters: Mapping[str, str | int | float | Sequence[str]]) -> Future[str]: ...
def query(
self, query: ReadQueryName, parameters: Mapping[str, str | int | float | Sequence[str]]
) -> Future[str]: ...
@final
class NativeDiagnosticProcessor:
@ -73,96 +73,48 @@ class NativeDiagnosticProcessor:
def scrub_access_arguments(self, arguments: Sequence[str]) -> list[str]: ...
def ocr(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: dict[str, object],
call: NativeCall,
) -> OCRResponse: ...
def aocr(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: dict[str, object],
call: NativeCall,
) -> Coroutine[object, object, OCRResponse]: ...
def ocr_health_check_document(model: str, custom_llm_provider: str | None) -> dict[str, object]: ...
def ocr_passthrough_response(model: str, endpoint: str, body: bytes) -> dict[str, object] | None: ...
def embedding(
request: LiteLLMEmbeddingRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> EmbeddingResponse: ...
def aembedding(
request: LiteLLMEmbeddingRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> Coroutine[object, object, EmbeddingResponse]: ...
def transcription(
model: str,
audio: object,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
optional_params: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
call: NativeCall,
) -> dict[str, object]: ...
def atranscription(
model: str,
audio: object,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
optional_params: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
call: NativeCall,
) -> Future[dict[str, object]]: ...
def completion(
request: LiteLLMChatCompletionsRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> ModelResponse: ...
def acompletion(
request: LiteLLMChatCompletionsRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> Coroutine[object, object, ModelResponse]: ...
def responses(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> ResponsesAPIResponse: ...
def aresponses(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> Coroutine[object, object, ResponsesAPIResponse]: ...
def messages(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: dict[str, object],
call: NativeCall,
) -> AnthropicMessagesResponse | Iterator[bytes]: ...
def amessages(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: dict[str, object],
call: NativeCall,
) -> Coroutine[object, object, AnthropicMessagesResponse | AsyncIterator[bytes]]: ...
def chat_completions(
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None = None,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
call: NativeCall,
) -> dict[str, object]: ...
def achat_completions(
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None = None,
api_key: str | None = None,
api_base: str | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, object] | None = None,
timeout_seconds: float | None = None,
call: NativeCall,
) -> Future[dict[str, object]]: ...
@final
@ -338,27 +290,42 @@ class _SecretManagerRuntime:
def read_secret(self, name: str, settings: Mapping[str, object] | None = None) -> JsonValue: ...
def read_secret_async(self, name: str, settings: Mapping[str, object] | None = None) -> Future[JsonValue]: ...
def async_write_secret(
self, secret_name: str, secret_value: str, description: str | None = None,
self,
secret_name: str,
secret_value: str,
description: str | None = None,
optional_params: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None, tags: object = None,
timeout: float | httpx.Timeout | None = None,
tags: object = None,
) -> Future[dict[str, JsonValue]]: ...
def async_delete_secret(
self, secret_name: str, recovery_window_in_days: int | None = None,
self,
secret_name: str,
recovery_window_in_days: int | None = None,
optional_params: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
) -> Future[dict[str, JsonValue]]: ...
def async_rotate_secret(
self, current_secret_name: str, new_secret_name: str, new_secret_value: str,
self,
current_secret_name: str,
new_secret_name: str,
new_secret_value: str,
optional_params: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
) -> Future[dict[str, JsonValue]]: ...
def sync_read_secret(
self, secret_name: str, optional_params: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None, primary_secret_name: str | None = None,
self,
secret_name: str,
optional_params: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
primary_secret_name: str | None = None,
) -> JsonValue: ...
def async_read_secret(
self, secret_name: str, optional_params: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None, primary_secret_name: str | None = None,
self,
secret_name: str,
optional_params: Mapping[str, object] | None = None,
timeout: float | httpx.Timeout | None = None,
primary_secret_name: str | None = None,
) -> Future[JsonValue]: ...
@final
@ -366,11 +333,18 @@ class NativeCacheHandle:
def __new__(cls, _uninstantiable: Never, /) -> Never: ...
@staticmethod
def memory(
*, ttl: float = 600.0, capacity: int = 200, max_entry_bytes: int = 4194304,
*,
ttl: float = 600.0,
capacity: int = 200,
max_entry_bytes: int = 4194304,
) -> NativeCacheHandle: ...
@staticmethod
def redis(
url: str, *, namespace: str, ttl: float = 600.0, max_entry_bytes: int = 4194304,
url: str,
*,
namespace: str,
ttl: float = 600.0,
max_entry_bytes: int = 4194304,
) -> NativeCacheHandle: ...
def get(self, key: str) -> object: ...
def set(self, key: str, value: object, *, ttl: float | None = None) -> None: ...

View file

@ -1,42 +1,24 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from dataclasses import dataclass, field
from types import MappingProxyType
from collections.abc import Awaitable
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.utils import ModelResponse
@dataclass(frozen=True, slots=True)
class LiteLLMChatCompletionsRequest:
model: str
messages: Sequence[object]
stream: bool | None
api_key: str | None
api_base: str | None
custom_llm_provider: str | None
extra_headers: Mapping[str, object] | None
kwargs: Mapping[str, object]
parameters: Mapping[str, object] = field(default_factory=lambda: MappingProxyType({}))
class NativeCompletion(Protocol):
def __call__(
self,
request: LiteLLMChatCompletionsRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> ModelResponse: ...
class NativeAcompletion(Protocol):
def __call__(
self,
request: LiteLLMChatCompletionsRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> Awaitable[ModelResponse]: ...

View file

@ -6,7 +6,7 @@ from typing import Final
import litellm
from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS
from litellm.rust_bridge import failures
from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest
from litellm.rust_bridge.public_call import optional_str
from litellm.types.utils import ModelResponse
_TRANSPORT_PARAMETERS: Final = frozenset(
@ -37,10 +37,16 @@ def response(value: Mapping[str, object]) -> ModelResponse:
return ModelResponse(**value)
def arguments(request: LiteLLMChatCompletionsRequest) -> Mapping[str, object]:
return request.kwargs
def arguments(request: Mapping[str, object]) -> Mapping[str, object]:
return request
def map_failure(error: Exception, request: LiteLLMChatCompletionsRequest) -> Exception:
provider: Final = request.custom_llm_provider or request.model.partition("/")[0]
return failures.map_native_failure(error, request.model, provider, arguments(request), request.api_base)
def map_failure(error: Exception, request: Mapping[str, object]) -> Exception:
provider: Final = optional_str(request.get("custom_llm_provider")) or str(request["model"]).partition("/")[0]
return failures.map_native_failure(
error,
str(request["model"]),
provider,
arguments(request),
optional_str(request.get("api_base")) or optional_str(request.get("base_url")),
)

View file

@ -1,38 +1,24 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping
from dataclasses import dataclass
from collections.abc import Awaitable
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.utils import EmbeddingResponse
@dataclass(frozen=True, slots=True)
class LiteLLMEmbeddingRequest:
model: str
input: object
api_key: str | None
api_base: str | None
custom_llm_provider: str | None
kwargs: Mapping[str, object]
class NativeEmbedding(Protocol):
def __call__(
self,
request: LiteLLMEmbeddingRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> EmbeddingResponse: ...
class NativeAembedding(Protocol):
def __call__(
self,
request: LiteLLMEmbeddingRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
call: NativeCall,
) -> Awaitable[EmbeddingResponse]: ...

View file

@ -1,40 +1,24 @@
from __future__ import annotations
from collections.abc import AsyncIterator, Awaitable, Iterator, Mapping, Sequence
from dataclasses import dataclass
from collections.abc import AsyncIterator, Awaitable, Iterator
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
@dataclass(frozen=True, slots=True)
class LiteLLMMessagesRequest:
model: str
messages: Sequence[object]
max_tokens: int
stream: bool | None
api_key: str | None
api_base: str | None
custom_llm_provider: str | None
kwargs: Mapping[str, object]
class NativeMessages(Protocol):
def __call__(
self,
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> AnthropicMessagesResponse | Iterator[bytes]: ...
class NativeAmessages(Protocol):
def __call__(
self,
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> Awaitable[AnthropicMessagesResponse | AsyncIterator[bytes]]: ...

View file

@ -11,37 +11,14 @@ import litellm
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.llms.anthropic.pass_through.utils import is_reasoning_auto_summary_enabled
from litellm.rust_bridge import failures
from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
from litellm.rust_bridge.public_call import optional_str
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
_DROP_PATHS: Final = TypeAdapter(list[object])
@dataclass(frozen=True, slots=True)
class EffortTiers:
minimal: bool
low: bool
medium: bool
high: bool
xhigh: bool
max: bool
@dataclass(frozen=True, slots=True)
class ModelCapabilities:
supports_reasoning: bool
supports_adaptive_thinking: bool
thinking_always_on: bool
supports_legacy_thinking: bool
supports_output_config: bool
supports_sampling_params: bool
supports_speed: bool
effort_tiers: EffortTiers
@dataclass(frozen=True, slots=True)
class MessagesShaping:
capabilities: ModelCapabilities
class MessagesSettings:
drop_params: bool
reasoning_auto_summary: bool
additional_drop_params: Sequence[str]
@ -62,56 +39,19 @@ def stream_hidden_params(headers: Sequence[tuple[str, str]]) -> Mapping[str, obj
return anthropic_messages_stream_hidden_params(httpx.Headers(list(headers)))
def arguments(request: LiteLLMMessagesRequest) -> Mapping[str, object]:
return request.kwargs
def arguments(request: Mapping[str, object]) -> Mapping[str, object]:
return request
def map_failure(error: Exception, request: LiteLLMMessagesRequest, request_provider: str) -> Exception:
def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception:
if getattr(error, "messages_request_error", False):
return litellm.BadRequestError(
message=str(error),
model=request.model.removeprefix(f"{request_provider}/"),
model=str(request["model"]).removeprefix(f"{request_provider}/"),
llm_provider=request_provider,
)
return failures.map_native_failure(error, request.model, request_provider, arguments(request), request.api_base)
def _resolved_provider(model: str, custom_llm_provider: str | None) -> tuple[str, str]:
try:
resolved_model, provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider)
except Exception: # noqa: BLE001 # an unroutable model still shapes as a bare Anthropic id
return model, custom_llm_provider or "anthropic"
return resolved_model, provider
def model_capabilities(model: str, custom_llm_provider: str | None) -> ModelCapabilities:
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
resolved_model, provider = _resolved_provider(model, custom_llm_provider)
def supports(flag: str) -> bool:
return AnthropicModelInfo._supports_model_capability(model, flag, provider) # pyright: ignore[reportPrivateUsage] # same probes the Python transform runs; forking them would drift
def tier(level: str) -> bool:
return AnthropicConfig._supports_effort_level(model, level, provider) # pyright: ignore[reportPrivateUsage] # same probe the Python transform runs
return ModelCapabilities(
supports_reasoning=supports("supports_reasoning"),
supports_adaptive_thinking=supports("supports_adaptive_thinking"),
thinking_always_on=supports("thinking_always_on"),
supports_legacy_thinking=supports("supports_legacy_thinking"),
supports_output_config=supports("supports_output_config"),
supports_sampling_params=AnthropicModelInfo._supports_sampling_params(resolved_model), # pyright: ignore[reportPrivateUsage] # same gate the handler applies
supports_speed=AnthropicConfig._model_supports_speed_param(resolved_model, provider), # pyright: ignore[reportPrivateUsage] # same gate the handler applies
effort_tiers=EffortTiers(
minimal=tier("minimal"),
low=tier("low"),
medium=tier("medium"),
high=tier("high"),
xhigh=tier("xhigh"),
max=tier("max"),
),
return failures.map_native_failure(
error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base"))
)
@ -127,10 +67,9 @@ def _additional_drop_params(kwargs: Mapping[str, object]) -> tuple[str, ...]:
return tuple(path for path in configured if isinstance(path, str))
def shaping(model: str, custom_llm_provider: str | None, kwargs: Mapping[str, object]) -> dict[str, object]:
def settings(kwargs: Mapping[str, object]) -> dict[str, object]:
return asdict(
MessagesShaping(
capabilities=model_capabilities(model, custom_llm_provider),
MessagesSettings(
drop_params=_drop_params(kwargs),
reasoning_auto_summary=is_reasoning_auto_summary_enabled(),
additional_drop_params=_additional_drop_params(kwargs),

View file

@ -0,0 +1,35 @@
from __future__ import annotations
import litellm
def _resolved_provider(model: str, custom_llm_provider: str | None) -> tuple[str, str]:
try:
resolved_model, provider, _, _ = litellm.get_llm_provider(model=model, custom_llm_provider=custom_llm_provider)
except Exception: # noqa: BLE001 # an unroutable model still shapes as a bare Anthropic id
return model, custom_llm_provider or "anthropic"
return resolved_model, provider
def anthropic_model_capabilities(model: str, custom_llm_provider: str | None) -> dict[str, object]:
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
resolved_model, provider = _resolved_provider(model, custom_llm_provider)
def supports(flag: str) -> bool:
return AnthropicModelInfo._supports_model_capability(model, flag, provider) # pyright: ignore[reportPrivateUsage] # same probes the Python transform runs; forking them would drift
def tier(level: str) -> bool:
return AnthropicConfig._supports_effort_level(model, level, provider) # pyright: ignore[reportPrivateUsage] # same probe the Python transform runs
return {
"supports_reasoning": supports("supports_reasoning"),
"supports_adaptive_thinking": supports("supports_adaptive_thinking"),
"thinking_always_on": supports("thinking_always_on"),
"supports_legacy_thinking": supports("supports_legacy_thinking"),
"supports_output_config": supports("supports_output_config"),
"supports_sampling_params": AnthropicModelInfo._supports_sampling_params(resolved_model), # pyright: ignore[reportPrivateUsage] # same gate the handler applies
"supports_speed": AnthropicConfig._model_supports_speed_param(resolved_model, provider), # pyright: ignore[reportPrivateUsage] # same gate the handler applies
"effort_tiers": {level: tier(level) for level in ("minimal", "low", "medium", "high", "xhigh", "max")},
}

View file

@ -1,43 +1,24 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping
from dataclasses import dataclass
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
import httpx
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRResponse
from litellm.rust_bridge.bindings import NativeBinding
@dataclass(frozen=True, slots=True)
class LiteLLMOcrRequest:
model: str
document: Mapping[str, object]
api_key: str | None
api_base: str | None
timeout: float | httpx.Timeout | None
custom_llm_provider: str | None
extra_headers: dict[str, object] | None
kwargs: Mapping[str, object]
input_sources: Mapping[str, str] | None = None
from litellm.rust_bridge.public_call import NativeCall
class NativeOcr(Protocol):
def __call__(
self,
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse: ...
class NativeAocr(Protocol):
def __call__(
self,
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> Awaitable[OCRResponse]: ...

View file

@ -10,7 +10,7 @@ import litellm
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
from litellm.rust_bridge import failures
from litellm.rust_bridge.failures import UpstreamFailure
from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest
from litellm.rust_bridge.public_call import optional_str
__all__ = ("UpstreamFailure", "arguments", "map_failure", "response")
@ -27,15 +27,17 @@ def response(value: Mapping[str, object]) -> OCRResponse:
return normalized
def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]:
return request.kwargs
def arguments(request: Mapping[str, object]) -> Mapping[str, object]:
return request
def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception:
def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception:
if getattr(error, "ocr_request_format_error", False):
return litellm.UnsupportedParamsError(
message=f"Invalid `req_format`: {request.kwargs.get('req_format')!r}. Expected 'native' or 'litellm'.",
model=request.model.removeprefix(f"{request_provider}/"),
message=f"Invalid `req_format`: {request.get('req_format')!r}. Expected 'native' or 'litellm'.",
model=str(request["model"]).removeprefix(f"{request_provider}/"),
llm_provider=request_provider,
)
return failures.map_native_failure(error, request.model, request_provider, arguments(request), request.api_base)
return failures.map_native_failure(
error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base"))
)

View file

@ -4,7 +4,9 @@ from __future__ import annotations
import inspect
from collections.abc import Callable, Mapping, Sequence
from typing import Final, cast # noqa: TID251 # narrows caller-owned containers without copying them
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, TypeVar, cast # noqa: TID251 # narrows caller-owned containers without copying them
import litellm
@ -80,3 +82,28 @@ def inference_decline_reason(parameters: tuple[str, ...], kwargs: Mapping[str, o
if name not in parameters and name not in _INFERENCE_CONTEXT:
return f"native inference does not implement {name}"
return None
@dataclass(frozen=True, slots=True)
class NativeCall:
args: tuple[object, ...]
kwargs: Mapping[str, object]
bound: Mapping[str, object]
def native_call(args: tuple[object, ...], kwargs: Mapping[str, object], fields: Mapping[str, object]) -> NativeCall:
extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({})
named: Final = {name: value for name, value in fields.items() if name != "kwargs"}
return NativeCall(args=args, kwargs=kwargs, bound=MappingProxyType({**named, **extra}))
NativeResultT: Final = TypeVar("NativeResultT")
def native_call_hook(
hook: Callable[[NativeCall], NativeResultT],
call: NativeCall,
_args: tuple[object, ...],
_kwargs: Mapping[str, object],
) -> NativeResultT:
return hook(call)

View file

@ -1,42 +1,24 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping
from dataclasses import dataclass, field
from types import MappingProxyType
from collections.abc import Awaitable
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.llms.openai import ResponsesAPIResponse
@dataclass(frozen=True, slots=True)
class LiteLLMResponsesRequest:
model: str
input: object
stream: bool | None
api_key: str | None
api_base: str | None
custom_llm_provider: str | None
extra_headers: Mapping[str, object] | None
kwargs: Mapping[str, object]
parameters: Mapping[str, object] = field(default_factory=lambda: MappingProxyType({}))
class NativeResponses(Protocol):
def __call__(
self,
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> ResponsesAPIResponse: ...
class NativeAresponses(Protocol):
def __call__(
self,
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> Awaitable[ResponsesAPIResponse]: ...

View file

@ -6,8 +6,7 @@ from typing import Final
import litellm
from litellm import get_llm_provider
from litellm.rust_bridge import failures
from litellm.rust_bridge.public_call import inference_decline_reason
from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest
from litellm.rust_bridge.public_call import inference_decline_reason, optional_str
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams, ResponsesAPIResponse
PARAMETERS: Final = tuple(ResponsesAPIOptionalRequestParams.__annotations__)
@ -21,21 +20,27 @@ def response(value: Mapping[str, object]) -> ResponsesAPIResponse:
return ResponsesAPIResponse.model_validate(value)
def arguments(request: LiteLLMResponsesRequest) -> Mapping[str, object]:
return request.kwargs
def arguments(request: Mapping[str, object]) -> Mapping[str, object]:
return request
def map_failure(error: Exception, request: LiteLLMResponsesRequest) -> Exception:
provider: Final = request.custom_llm_provider or "openai"
return failures.map_native_failure(error, request.model, provider, arguments(request), request.api_base)
def map_failure(error: Exception, request: Mapping[str, object]) -> Exception:
provider: Final = optional_str(request.get("custom_llm_provider")) or "openai"
return failures.map_native_failure(
error,
str(request["model"]),
provider,
arguments(request),
optional_str(request.get("api_base")) or optional_str(request.get("base_url")),
)
def decline_reason(request: LiteLLMResponsesRequest) -> str | None:
if request.custom_llm_provider is None and "/" not in request.model:
def decline_reason(request: Mapping[str, object]) -> str | None:
if optional_str(request.get("custom_llm_provider")) is None and "/" not in str(request["model"]):
try:
_, provider, _, _ = get_llm_provider(model=request.model)
_, provider, _, _ = get_llm_provider(model=str(request["model"]))
except litellm.exceptions.BadRequestError:
return "native Responses could not resolve the provider"
if provider != "openai":
return "native HTTP responses provider"
return inference_decline_reason(PARAMETERS, {**request.parameters, **request.kwargs})
return inference_decline_reason(PARAMETERS, request)

View file

@ -4,19 +4,13 @@ from collections.abc import Awaitable
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.public_call import NativeCall
class RustTranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
call: NativeCall,
) -> dict[str, object]:
raise NotImplementedError
@ -24,14 +18,7 @@ class RustTranscription(Protocol):
class RustAtranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
call: NativeCall,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError

View file

@ -22,7 +22,7 @@ from litellm.rust_bridge import runtime
from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.dispatch import call_hook
from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.caching import CachingSupportedCallTypes
from tests.test_litellm_rust.support.cache import cache_key, collect, invoke, payload
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging
@ -443,15 +443,26 @@ def test_sync_rust_messages_calls_python_cache(recording_server: RecordingServer
"api_key": "test-key",
"api_base": recording_server.base_url,
}
request: Final = LiteLLMMessagesRequest(
MESSAGES_MODEL, list(MESSAGES), 32, None, "test-key", recording_server.base_url, "anthropic", arguments
request: Final = NativeCall(
args=(),
kwargs=arguments,
bound={
"model": MESSAGES_MODEL,
"messages": list(MESSAGES),
"max_tokens": 32,
"stream": None,
"api_key": "test-key",
"api_base": recording_server.base_url,
"custom_llm_provider": "anthropic",
**arguments,
},
)
def call() -> object:
return runtime.run(
RouteContext(Route.MESSAGES),
binding=NATIVE_MESSAGES,
native=lambda hook: call_hook(hook, request, (), arguments),
native=lambda hook: hook(request),
python=runtime.NO_PYTHON,
rules=(RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),),
)

View file

@ -235,3 +235,83 @@ async def test_non_string_metadata_user_id_is_rejected_before_the_provider_call(
await litellm.anthropic.messages.acreate(**arguments(messages_server, metadata={"user_id": 123}))
assert messages_server.requests == []
@pytest.mark.asyncio
async def test_native_messages_observes_runtime_capabilities_and_separate_caller_settings(
messages_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES
from litellm.rust_bridge.public_call import NativeCall
native: Final = NATIVE_AMESSAGES.load()
assert native is not None
model: Final = "claude-test-runtime-capabilities"
messages_server.expected_requests = 2
request: Final = NativeCall(
args=(),
kwargs={"temperature": 0.2, "drop_params": True},
bound={
"model": model,
"messages": MESSAGES,
"max_tokens": 16,
"stream": None,
"api_key": "test-key",
"api_base": messages_server.base_url,
"custom_llm_provider": "anthropic",
"temperature": 0.2,
"drop_params": True,
},
)
monkeypatch.setitem(
litellm.model_cost,
model,
{
"litellm_provider": "anthropic",
"mode": "chat",
"supports_sampling_params": True,
},
)
first: Final = await native(request)
monkeypatch.setitem(
litellm.model_cost,
model,
{
"litellm_provider": "anthropic",
"mode": "chat",
"supports_sampling_params": False,
},
)
second: Final = await native(request)
assert isinstance(first, dict)
assert isinstance(second, dict)
assert first["id"] == second["id"] == MESSAGES_RESPONSE["id"]
assert len(messages_server.requests) == 2
first_body: Final = messages_server.requests[0].body
second_body: Final = messages_server.requests[1].body
assert isinstance(first_body, dict)
assert isinstance(second_body, dict)
assert first_body["temperature"] == request.kwargs["temperature"]
assert second_body == {name: value for name, value in first_body.items() if name != "temperature"}
@pytest.mark.asyncio
async def test_native_messages_reads_optional_positional_body_parameters(messages_server: RecordingServer) -> None:
from litellm.messages.dispatch import _MESSAGES, _public_request
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES
native: Final = NATIVE_AMESSAGES.load()
assert native is not None
metadata: Final = {"user_id": "caller"}
args: Final = (16, MESSAGES, "anthropic/claude-test", metadata, None, False, "Be brief", 0.25)
kwargs: Final = {"api_key": "test-key", "api_base": messages_server.base_url}
call: Final = _public_request(_MESSAGES, args, kwargs)
assert call is not None
await native(call)
body, _ = sent(messages_server)
assert body["temperature"] == args[7]
assert body["system"] == args[6]
assert body["metadata"] == metadata

View file

@ -296,7 +296,7 @@ def test_unstarted_native_coroutine_releases_input_without_reading_file(ocr_serv
def create():
file: Final = File()
kwargs: Final = {"model": "mistral/mistral-ocr-latest", "document": {"type": "file", "file": file}}
coroutine: Final = _native.aocr(_public_request("aocr", (), kwargs), (), kwargs)
coroutine: Final = _native.aocr(_public_request("aocr", (), kwargs))
file.owner = coroutine
coroutine.close()
return weakref.ref(file)

View file

@ -610,7 +610,8 @@ def test_native_projection_errors_never_select_python(
from litellm.rust_bridge import runtime, settings
from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR, LiteLLMOcrRequest
from litellm.rust_bridge.ocr.entrypoints import NATIVE_OCR
from litellm.rust_bridge.public_call import NativeCall
ocr_server.expected_requests = 0
snapshot: Final = dataclasses.replace(settings.http_settings(), user_agent=1)
@ -620,15 +621,18 @@ def test_native_projection_errors_never_select_python(
monkeypatch.setattr(
litellm, "ssl_verify", ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) if failure == "live" else object()
)
request: Final = LiteLLMOcrRequest(
model="mistral/mistral-ocr-latest",
document=OCR_DOCUMENT,
api_key="test-key",
api_base=ocr_server.base_url,
timeout=None,
custom_llm_provider="mistral",
extra_headers=None,
request: Final = NativeCall(
args=(),
kwargs={},
bound={
"model": "mistral/mistral-ocr-latest",
"document": OCR_DOCUMENT,
"api_key": "test-key",
"api_base": ocr_server.base_url,
"timeout": None,
"custom_llm_provider": "mistral",
"extra_headers": None,
},
)
def python_fallback() -> NoReturn:
@ -638,7 +642,7 @@ def test_native_projection_errors_never_select_python(
runtime.run(
RouteContext(Route.OCR, provider="mistral"),
binding=NATIVE_OCR,
native=lambda native: native(request, (), {}),
native=lambda native: native(request),
python=python_fallback,
rules=(RouteRule(Route.OCR, Rollout.RUST_REQUIRED if required else Rollout.RUST_OPT_OUT),),
)

View file

@ -10,11 +10,11 @@ import litellm
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
from litellm.rust_bridge import runtime
from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule
from litellm.rust_bridge.chat_completions.entrypoints import NATIVE_ACOMPLETION, LiteLLMChatCompletionsRequest
from litellm.rust_bridge.chat_completions.entrypoints import NATIVE_ACOMPLETION
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.dispatch import call_hook
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, LiteLLMMessagesRequest
from litellm.rust_bridge.responses.entrypoints import NATIVE_ARESPONSES, LiteLLMResponsesRequest
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES
from litellm.rust_bridge.public_call import NativeCall
from litellm.rust_bridge.responses.entrypoints import NATIVE_ARESPONSES
from litellm.types.utils import ModelResponse
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_EVENTS, MESSAGES_MODEL, MESSAGES_RESPONSE
@ -48,13 +48,24 @@ async def invoke(
arguments: Final = {"model": RESPONSES_MODEL, "input": "hello", **common}
if not native:
return await litellm.aresponses(**arguments)
request: Final = LiteLLMResponsesRequest(
RESPONSES_MODEL, "hello", None, "test-key", server.base_url, "openai", None, arguments
request: Final = NativeCall(
args=(),
kwargs=arguments,
bound={
"model": RESPONSES_MODEL,
"input": "hello",
"stream": None,
"api_key": "test-key",
"api_base": server.base_url,
"custom_llm_provider": "openai",
"extra_headers": None,
**arguments,
},
)
return await runtime.arun(
RouteContext(Route.RESPONSES),
binding=NATIVE_ARESPONSES,
native=lambda hook: call_hook(hook, request, (), arguments),
native=lambda hook: hook(request),
python=runtime.NO_PYTHON,
rules=(RouteRule(Route.RESPONSES, Rollout.RUST_REQUIRED),),
)
@ -67,25 +78,34 @@ async def invoke(
if route == "chat":
if not native:
return await litellm.acompletion(**parameters)
chat: Final = LiteLLMChatCompletionsRequest(
MESSAGES_MODEL, list(MESSAGES), None, "test-key", server.base_url, None, None, parameters
)
chat: Final = NativeCall(args=(), kwargs=parameters, bound=parameters)
return await runtime.arun(
RouteContext(Route.CHAT_COMPLETIONS),
binding=NATIVE_ACOMPLETION,
native=lambda hook: call_hook(hook, chat, (), parameters),
native=lambda hook: hook(chat),
python=runtime.NO_PYTHON,
rules=(RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),),
)
if not native:
return await litellm.anthropic_messages(**parameters)
messages: Final = LiteLLMMessagesRequest(
MESSAGES_MODEL, list(MESSAGES), 32, None, "test-key", server.base_url, "anthropic", parameters
messages: Final = NativeCall(
args=(),
kwargs=parameters,
bound={
"model": MESSAGES_MODEL,
"messages": list(MESSAGES),
"max_tokens": 32,
"stream": None,
"api_key": "test-key",
"api_base": server.base_url,
"custom_llm_provider": "anthropic",
**parameters,
},
)
return await runtime.arun(
RouteContext(Route.MESSAGES),
binding=NATIVE_AMESSAGES,
native=lambda hook: call_hook(hook, messages, (), parameters),
native=lambda hook: hook(messages),
python=runtime.NO_PYTHON,
rules=(RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),),
)

View file

@ -7,12 +7,12 @@ from pydantic import JsonValue, TypeAdapter
import litellm
from litellm import RateLimitError
from litellm.chat_completions import dispatch as chat_dispatch
from litellm.integrations.custom_logger import CustomLogger
from litellm.models.credentials import CredentialItem
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.rust_bridge import _native
from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest
from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import CallTypes, ModelResponse
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger
@ -66,10 +66,8 @@ def native_call(
"max_tokens": 32,
**options,
}
request: Final = LiteLLMChatCompletionsRequest(
MESSAGES_MODEL, list(MESSAGES), None, "test-key", server.base_url, None, None, kwargs
)
return (_native.acompletion if asynchronous else _native.completion)(request, (), kwargs)
request: Final = NativeCall(args=(), kwargs=kwargs, bound=kwargs)
return (_native.acompletion if asynchronous else _native.completion)(request)
response_kwargs: Final = {
"model": RESPONSES_MODEL,
"input": "hello",
@ -78,10 +76,21 @@ def native_call(
"max_output_tokens": 32,
**options,
}
response_request: Final = LiteLLMResponsesRequest(
RESPONSES_MODEL, "hello", None, "test-key", server.base_url, "openai", None, response_kwargs
response_request: Final = NativeCall(
args=(),
kwargs=response_kwargs,
bound={
"model": RESPONSES_MODEL,
"input": "hello",
"stream": None,
"api_key": "test-key",
"api_base": server.base_url,
"custom_llm_provider": "openai",
"extra_headers": None,
**response_kwargs,
},
)
return (_native.aresponses if asynchronous else _native.responses)(response_request, (), response_kwargs)
return (_native.aresponses if asynchronous else _native.responses)(response_request)
async def execute(route: Route, asynchronous: bool, server: RecordingServer, options: Mapping[str, object]) -> object:
@ -284,13 +293,13 @@ async def test_native_projection_reads_positional_parameters(route: Route, recor
args: Final = (MESSAGES_MODEL, list(MESSAGES), 12.0, 0.25)
request: Final = chat_dispatch.request(args, kwargs)
assert request is not None
await asyncio.to_thread(_native.completion, request, args, kwargs)
await asyncio.to_thread(_native.completion, request)
assert _OBJECT.validate_python(recording_server.requests[0].body)["temperature"] == 0.25
else:
response_args: Final = ("hello", RESPONSES_MODEL, None, "Be brief", 16)
response_request: Final = responses_dispatch.request(response_args, kwargs)
assert response_request is not None
await asyncio.to_thread(_native.responses, response_request, response_args, kwargs)
await asyncio.to_thread(_native.responses, response_request)
body: Final = _OBJECT.validate_python(recording_server.requests[0].body)
assert body["instructions"] == "Be brief"
assert body["max_output_tokens"] == 16
@ -351,3 +360,24 @@ async def test_native_responses_decode_continuation_ids(
)
await execute("responses", asynchronous, recording_server, {"previous_response_id": previous})
assert _OBJECT.validate_python(recording_server.requests[0].body)["previous_response_id"] == original
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_native_chat_uses_bound_positional_parameters(
asynchronous: bool, recording_server: RecordingServer
) -> None:
recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE)
arguments: Final = (MESSAGES_MODEL, list(MESSAGES), 12.0, 0.35)
supplied: Final = {"base_url": recording_server.base_url, "api_key": "test-key", "max_tokens": 32}
call: Final = chat_dispatch._DISPATCH.request(arguments, supplied) # pyright: ignore[reportPrivateUsage] # exercise the native request produced by public binding
assert call is not None
result: Final = (
await _native.acompletion(call) if asynchronous else await asyncio.to_thread(_native.completion, call)
)
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "Hello from native Messages"
assert len(recording_server.requests) == 1
body: Final = _OBJECT.validate_python(recording_server.requests[0].body)
assert body["temperature"] == arguments[3]
assert body["max_tokens"] == supplied["max_tokens"]

View file

@ -4,24 +4,23 @@ from typing import Final, cast # noqa: TID251 # narrows legacy callable signat
import pytest
import litellm
from litellm.chat_completions import dispatch
from litellm.chat_completions.dispatch import (
_ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
_DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
)
from litellm.rust_bridge import catalog
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, RouteRule
from litellm.rust_bridge.catalog import Route, RouteRule, Rules
from litellm.rust_bridge.chat_completions.entrypoints import (
NATIVE_ACOMPLETION,
NATIVE_COMPLETION,
LiteLLMChatCompletionsRequest,
NativeAcompletion,
NativeCompletion,
)
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.public_call import NativeCall, native_call_hook
from litellm.types.utils import ModelResponse
from litellm.chat_completions import dispatch
from litellm.rust_bridge.catalog import Rules
MESSAGES: Final = [{"role": "user", "content": "hi"}]
PYTHON_RULES: Final = ()
@ -51,9 +50,7 @@ def test_python_route_forwards_original_call_shape() -> None:
captured.append((call_args, call_kwargs))
return response
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
def native(request: NativeCall) -> ModelResponse:
pytest.fail("Python-only dispatch must not call native")
assert (
@ -62,7 +59,7 @@ def test_python_route_forwards_original_call_shape() -> None:
kwargs,
python=python,
binding=completion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=native_call_hook,
rules=PYTHON_RULES,
)
is response
@ -87,9 +84,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None:
captured.append((call_args, call_kwargs))
return response
async def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
async def native(request: NativeCall) -> ModelResponse:
pytest.fail("Python-only dispatch must not call native")
result: Final = await _ADISPATCH.arun(
@ -97,7 +92,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None:
kwargs,
python=python,
binding=acompletion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=native_call_hook,
rules=PYTHON_RULES,
)
assert result is response
@ -119,15 +114,13 @@ def test_native_receives_bound_request_and_original_call_shape() -> None:
"custom_llm_provider": "anthropic",
"metadata": metadata,
}
captured: Final[list[tuple[LiteLLMChatCompletionsRequest, tuple[object, ...], Mapping[str, object]]]] = []
captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = []
def python(*call_args: object, **call_kwargs: object) -> ModelResponse: # kwargs-ok: rejected Rust fallback
pytest.fail("Required Rust dispatch must not call Python")
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
captured.append((request, args, kwargs))
def native(request: NativeCall) -> ModelResponse:
captured.append((request, request.args, request.kwargs))
return ModelResponse()
args: Final[tuple[object, ...]] = ("anthropic/claude-sonnet-4-5", MESSAGES)
@ -136,19 +129,19 @@ def test_native_receives_bound_request_and_original_call_shape() -> None:
kwargs,
python=python,
binding=completion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=native_call_hook,
rules=RUST_RULES,
)
request, call_args, call_kwargs = captured[0]
assert request.model == "anthropic/claude-sonnet-4-5"
assert request.messages is MESSAGES
assert request.stream is True
assert request.api_key == "sk-test"
assert request.api_base == "https://example.invalid"
assert request.custom_llm_provider == "anthropic"
assert request.extra_headers == {"x-test": "1"}
assert request.kwargs == {"custom_llm_provider": "anthropic", "metadata": metadata}
assert request.bound["model"] == "anthropic/claude-sonnet-4-5"
assert request.bound["messages"] is MESSAGES
assert request.bound["stream"] is True
assert request.bound["api_key"] == "sk-test"
assert request.bound["base_url"] == "https://example.invalid"
assert request.bound["custom_llm_provider"] == "anthropic"
assert request.bound["extra_headers"] == {"x-test": "1"}
assert request.kwargs is kwargs
assert call_args == args
assert call_kwargs == kwargs
assert call_kwargs["metadata"] is metadata
@ -162,9 +155,7 @@ def test_internal_async_marker_bypasses_native() -> None:
called.append(True)
return response
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
def native(request: NativeCall) -> ModelResponse:
pytest.fail("acompletion's inner completion call must stay on Python")
result: Final = _DISPATCH.run(
@ -172,7 +163,7 @@ def test_internal_async_marker_bypasses_native() -> None:
{"acompletion": True},
python=python,
binding=completion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=native_call_hook,
rules=RUST_RULES,
)
assert result is response
@ -194,9 +185,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map
captured.append((call_args, call_kwargs))
return response
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
def native(request: NativeCall) -> ModelResponse:
pytest.fail("Binding failures must be delegated to Python")
assert (
@ -205,7 +194,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map
kwargs,
python=python,
binding=completion_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=native_call_hook,
rules=RUST_RULES,
)
is response
@ -214,14 +203,10 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map
def test_public_completion_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None:
captured: Final[list[LiteLLMChatCompletionsRequest]] = []
captured: Final[list[NativeCall]] = []
expected: Final = ModelResponse()
def native(
request: LiteLLMChatCompletionsRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> ModelResponse:
def native(request: NativeCall) -> ModelResponse:
captured.append(request)
return expected
@ -233,19 +218,15 @@ def test_public_completion_routes_through_dispatch(monkeypatch: pytest.MonkeyPat
finally:
NATIVE_COMPLETION.reset()
assert result is expected
assert [request.model for request in captured] == ["gpt-4o"]
assert [request.bound["model"] for request in captured] == ["gpt-4o"]
@pytest.mark.asyncio
async def test_public_acompletion_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None:
captured: Final[list[LiteLLMChatCompletionsRequest]] = []
captured: Final[list[NativeCall]] = []
expected: Final = ModelResponse()
async def native(
request: LiteLLMChatCompletionsRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
) -> ModelResponse:
async def native(request: NativeCall) -> ModelResponse:
captured.append(request)
return expected
@ -257,7 +238,7 @@ async def test_public_acompletion_routes_through_dispatch(monkeypatch: pytest.Mo
finally:
NATIVE_ACOMPLETION.reset()
assert result is expected
assert [request.model for request in captured] == ["gpt-4o"]
assert [request.bound["model"] for request in captured] == ["gpt-4o"]
@pytest.mark.asyncio
@ -275,13 +256,11 @@ def test_sync_completion_request_projects_public_arguments() -> None:
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_REQUIRED),)
expected: Final = ModelResponse()
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
assert request.model == "test-model"
assert request.messages == MESSAGES
assert request.custom_llm_provider == "openai"
assert request.stream is True
def native(request: NativeCall) -> ModelResponse:
assert request.bound["model"] == "test-model"
assert request.bound["messages"] == MESSAGES
assert request.bound["custom_llm_provider"] == "openai"
assert request.bound["stream"] is True
return expected
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
@ -291,7 +270,7 @@ def test_sync_completion_request_projects_public_arguments() -> None:
{"custom_llm_provider": "openai", "stream": True},
python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"),
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=native_call_hook,
rules=rules,
)
@ -309,9 +288,7 @@ async def test_async_completion_falls_back_after_native_declines() -> None:
expected: Final = ModelResponse()
rules: Final[Rules] = (RouteRule(Route.CHAT_COMPLETIONS, Rollout.RUST_OPT_OUT),)
async def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
async def native(request: NativeCall) -> ModelResponse:
raise declined("unsupported")
async def python(*args: object, **kwargs: object) -> ModelResponse:
@ -324,7 +301,7 @@ async def test_async_completion_falls_back_after_native_declines() -> None:
{},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=native_call_hook,
rules=rules,
)
@ -338,9 +315,7 @@ def test_internal_acompletion_marker_bypasses_native() -> None:
def python(*args: object, **kwargs: object) -> ModelResponse:
return expected
def native(
request: LiteLLMChatCompletionsRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> ModelResponse:
def native(request: NativeCall) -> ModelResponse:
pytest.fail("acompletion's inner completion call must stay on Python")
binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None)
@ -350,7 +325,7 @@ def test_internal_acompletion_marker_bypasses_native() -> None:
{"custom_llm_provider": "openai", "acompletion": True},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=native_call_hook,
rules=rules,
)
@ -360,6 +335,6 @@ def test_internal_acompletion_marker_bypasses_native() -> None:
def test_positional_parameters_remain_available_to_native_projection() -> None:
request: Final = _DISPATCH.request(("anthropic/test-model", MESSAGES, 12.0, 0.25), {})
assert request is not None
assert request.parameters["timeout"] == 12.0
assert request.parameters["temperature"] == 0.25
assert request.messages is MESSAGES
assert request.bound["timeout"] == 12.0
assert request.bound["temperature"] == 0.25
assert request.bound["messages"] is MESSAGES

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping
from collections.abc import Awaitable, Callable
from typing import Final
import pytest
@ -11,7 +11,7 @@ from litellm.embeddings import dispatch
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, RouteRule, Rules
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.embeddings.entrypoints import LiteLLMEmbeddingRequest
from litellm.rust_bridge.public_call import NativeCall, native_call_hook
from litellm.types.utils import EmbeddingResponse
@ -32,24 +32,22 @@ def test_sync_embedding_request_projects_public_arguments() -> None:
rules: Final[Rules] = (RouteRule(Route.EMBEDDINGS, Rollout.RUST_REQUIRED),)
expected: Final = EmbeddingResponse(model="test-model", data=[])
def native(
request: LiteLLMEmbeddingRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> EmbeddingResponse:
assert request.model == "test-model"
assert request.input == "hello"
assert request.custom_llm_provider == "openai"
def native(request: NativeCall) -> EmbeddingResponse:
assert request.bound["model"] == "test-model"
assert request.bound["input"] == "hello"
assert request.bound["custom_llm_provider"] == "openai"
return expected
binding: Final[
NativeBinding[Callable[[LiteLLMEmbeddingRequest, tuple[object, ...], Mapping[str, object]], EmbeddingResponse]]
] = NativeBinding("embedding", validate=lambda _: None)
binding: Final[NativeBinding[Callable[[NativeCall], EmbeddingResponse]]] = NativeBinding(
"embedding", validate=lambda _: None
)
binding.override(native)
response: Final = dispatch._DISPATCH.run( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
("test-model", "hello"),
{"custom_llm_provider": "openai", "dimensions": 8},
python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"),
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=native_call_hook,
rules=rules,
)
@ -67,26 +65,22 @@ async def test_async_embedding_falls_back_after_native_declines() -> None:
expected: Final = EmbeddingResponse(model="test-model", data=[])
rules: Final[Rules] = (RouteRule(Route.EMBEDDINGS, Rollout.RUST_OPT_OUT),)
async def native(
request: LiteLLMEmbeddingRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> EmbeddingResponse:
async def native(request: NativeCall) -> EmbeddingResponse:
raise declined("unsupported")
async def python(*args: object, **kwargs: object) -> EmbeddingResponse:
return expected
binding: Final[
NativeBinding[
Callable[[LiteLLMEmbeddingRequest, tuple[object, ...], Mapping[str, object]], Awaitable[EmbeddingResponse]]
]
] = NativeBinding("aembedding", validate=lambda _: None)
binding: Final[NativeBinding[Callable[[NativeCall], Awaitable[EmbeddingResponse]]]] = NativeBinding(
"aembedding", validate=lambda _: None
)
binding.override(native)
response: Final = await dispatch._ADISPATCH.arun( # pyright: ignore[reportPrivateUsage] # test an explicit route decision
("test-model", "hello"),
{},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=native_call_hook,
rules=rules,
)

View file

@ -3,9 +3,11 @@ from collections.abc import Awaitable, Callable, Mapping
from typing import Final, cast # noqa: TID251 # narrows legacy callable signatures for inspect
import pytest
from pydantic import TypeAdapter
import litellm
from litellm.llms.anthropic.pass_through.messages import handler as python_messages
from litellm.messages import dispatch
from litellm.messages.dispatch import (
_ADISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
_DISPATCH, # pyright: ignore[reportPrivateUsage] # tests configured dispatch
@ -14,16 +16,9 @@ from litellm.rust_bridge import catalog
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, RouteRule, Rules
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.messages.entrypoints import (
NATIVE_AMESSAGES,
NATIVE_MESSAGES,
LiteLLMMessagesRequest,
NativeAmessages,
NativeMessages,
)
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, NATIVE_MESSAGES, NativeAmessages, NativeMessages
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
from pydantic import TypeAdapter
from litellm.messages import dispatch
MESSAGES: Final = [{"role": "user", "content": "hi"}]
PYTHON_RULES: Final[Rules] = ()
@ -67,9 +62,7 @@ def test_python_route_forwards_original_call_shape() -> None:
return expected
def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> AnthropicMessagesResponse:
pytest.fail("Python-only dispatch must not call native")
@ -78,7 +71,7 @@ def test_python_route_forwards_original_call_shape() -> None:
kwargs,
python=python,
binding=messages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=PYTHON_RULES,
)
assert result is expected
@ -106,9 +99,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None:
return expected
async def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> AnthropicMessagesResponse:
pytest.fail("Python-only dispatch must not call native")
@ -117,7 +108,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None:
kwargs,
python=python,
binding=amessages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=PYTHON_RULES,
)
assert result is expected
@ -139,17 +130,17 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None:
"custom_llm_provider": "anthropic",
"litellm_metadata": metadata,
}
captured: Final[list[tuple[LiteLLMMessagesRequest, tuple[object, ...], Mapping[str, object]]]] = []
captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response("anthropic/claude-sonnet-4-5")
def python(*call_args: object, **call_kwargs: object) -> AnthropicMessagesResponse: # kwargs-ok: rejected fallback
pytest.fail("Required Rust dispatch must not call Python")
def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> AnthropicMessagesResponse:
args: Final = request.args
kwargs: Final = request.kwargs
captured.append((request, args, kwargs))
return expected
@ -158,19 +149,19 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None:
kwargs,
python=python,
binding=messages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
assert result is expected
request, call_args, call_kwargs = captured[0]
assert request.model == "anthropic/claude-sonnet-4-5"
assert request.messages is MESSAGES
assert request.max_tokens == 16
assert request.stream is True
assert request.api_key == "sk-test"
assert request.api_base == "https://example.invalid"
assert request.custom_llm_provider == "anthropic"
assert request.kwargs == {"litellm_metadata": metadata}
assert request.bound["model"] == "anthropic/claude-sonnet-4-5"
assert request.bound["messages"] is MESSAGES
assert request.bound["max_tokens"] == 16
assert request.bound["stream"] is True
assert request.bound["api_key"] == "sk-test"
assert request.bound["api_base"] == "https://example.invalid"
assert request.bound["custom_llm_provider"] == "anthropic"
assert request.kwargs == kwargs
assert request.kwargs["litellm_metadata"] is metadata
assert call_args == args
assert call_args[1] is MESSAGES
@ -189,9 +180,7 @@ def test_internal_async_marker_bypasses_native() -> None:
return expected
def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> AnthropicMessagesResponse:
pytest.fail("The async handler's inner sync call must stay on Python")
@ -200,7 +189,7 @@ def test_internal_async_marker_bypasses_native() -> None:
kwargs,
python=python,
binding=messages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
assert result is expected
@ -223,9 +212,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map
return expected
def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> AnthropicMessagesResponse:
pytest.fail("Binding failures must be delegated to Python")
@ -234,7 +221,7 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map
kwargs,
python=python,
binding=messages_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
assert result is expected
@ -242,13 +229,11 @@ def test_binding_errors_delegate_to_python(args: tuple[object, ...], kwargs: Map
def test_anthropic_create_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None:
captured: Final[list[LiteLLMMessagesRequest]] = []
captured: Final[list[NativeCall]] = []
expected: Final = response()
def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> AnthropicMessagesResponse:
captured.append(request)
return expected
@ -261,18 +246,16 @@ def test_anthropic_create_routes_through_dispatch(monkeypatch: pytest.MonkeyPatc
finally:
NATIVE_MESSAGES.reset()
assert result is expected
assert [request.model for request in captured] == ["claude-sonnet-4-5"]
assert [request.bound["model"] for request in captured] == ["claude-sonnet-4-5"]
@pytest.mark.asyncio
async def test_anthropic_acreate_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None:
captured: Final[list[LiteLLMMessagesRequest]] = []
captured: Final[list[NativeCall]] = []
expected: Final = response()
async def native(
request: LiteLLMMessagesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> AnthropicMessagesResponse:
captured.append(request)
return expected
@ -285,7 +268,7 @@ async def test_anthropic_acreate_routes_through_dispatch(monkeypatch: pytest.Mon
finally:
NATIVE_AMESSAGES.reset()
assert result is expected
assert [request.model for request in captured] == ["claude-sonnet-4-5"]
assert [request.bound["model"] for request in captured] == ["claude-sonnet-4-5"]
@pytest.mark.asyncio
@ -303,13 +286,11 @@ def test_sync_messages_request_projects_public_arguments() -> None:
rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_REQUIRED),)
expected: Final = AnthropicMessagesResponse(model="claude-test")
def native(
request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> AnthropicMessagesResponse:
assert request.model == "claude-test"
assert request.messages == MESSAGES
assert request.max_tokens == 10
assert request.custom_llm_provider == "anthropic"
def native(request: NativeCall) -> AnthropicMessagesResponse:
assert request.bound["model"] == "claude-test"
assert request.bound["messages"] == MESSAGES
assert request.bound["max_tokens"] == 10
assert request.bound["custom_llm_provider"] == "anthropic"
return expected
binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None)
@ -324,7 +305,7 @@ def test_sync_messages_request_projects_public_arguments() -> None:
},
python=lambda *args, **kwargs: pytest.fail("required native route must handle this call"),
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=lambda hook, request, args, kwargs: hook(request),
rules=rules,
)
@ -338,9 +319,7 @@ def test_messages_binding_error_delegates_unchanged_to_python() -> None:
def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse:
return expected
def native(
request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> AnthropicMessagesResponse:
def native(request: NativeCall) -> AnthropicMessagesResponse:
pytest.fail("a call without max_tokens cannot project a request and must stay on Python")
binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None)
@ -350,7 +329,7 @@ def test_messages_binding_error_delegates_unchanged_to_python() -> None:
{"model": "claude-test", "messages": MESSAGES, "custom_llm_provider": "anthropic"},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=lambda hook, request, args, kwargs: hook(request),
rules=rules,
)
@ -368,9 +347,7 @@ async def test_async_messages_falls_back_after_native_declines() -> None:
expected: Final = AnthropicMessagesResponse(model="claude-test")
rules: Final[Rules] = (RouteRule(Route.MESSAGES, Rollout.RUST_OPT_OUT),)
async def native(
request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> AnthropicMessagesResponse:
async def native(request: NativeCall) -> AnthropicMessagesResponse:
raise declined("unsupported")
async def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse:
@ -383,7 +360,7 @@ async def test_async_messages_falls_back_after_native_declines() -> None:
{"model": "claude-test", "messages": MESSAGES, "max_tokens": 10},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=lambda hook, request, args, kwargs: hook(request),
rules=rules,
)
@ -397,9 +374,7 @@ def test_internal_is_async_marker_bypasses_native() -> None:
def python(*args: object, **kwargs: object) -> AnthropicMessagesResponse:
return expected
def native(
request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object]
) -> AnthropicMessagesResponse:
def native(request: NativeCall) -> AnthropicMessagesResponse:
pytest.fail("anthropic_messages' inner handler call must stay on Python")
binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None)
@ -415,7 +390,7 @@ def test_internal_is_async_marker_bypasses_native() -> None:
},
python=python,
binding=binding,
native=lambda hook, request, args, kwargs: hook(request, args, kwargs),
native=lambda hook, request, args, kwargs: hook(request),
rules=rules,
)

View file

@ -14,13 +14,8 @@ from litellm.rust_bridge import catalog, runtime
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, RouteRule, Rules
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.ocr.entrypoints import (
NATIVE_AOCR,
NATIVE_OCR,
LiteLLMOcrRequest,
NativeAocr,
NativeOcr,
)
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, NativeAocr, NativeOcr
from litellm.rust_bridge.public_call import NativeCall
RUST_RULES: Final[Rules] = (RouteRule(Route.OCR, Rollout.RUST_REQUIRED),)
@ -55,14 +50,14 @@ def test_native_receives_normalized_positional_request_and_original_call_shape()
"extra_headers": extra_headers,
"pages": pages,
}
captured: Final[list[tuple[LiteLLMOcrRequest, tuple[object, ...], Mapping[str, object]]]] = []
captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse:
args: Final = request.args
kwargs: Final = request.kwargs
captured.append((request, args, kwargs))
return expected
@ -71,20 +66,20 @@ def test_native_receives_normalized_positional_request_and_original_call_shape()
kwargs,
python=runtime.NO_PYTHON,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
request, call_args, call_kwargs = captured[0]
assert result is expected
assert request.model == "mistral/mistral-ocr-latest"
assert request.document is document
assert request.api_key == "test-key"
assert request.api_base == "https://example.invalid"
assert request.timeout is timeout
assert request.custom_llm_provider == "mistral"
assert request.extra_headers is extra_headers
assert request.kwargs == {"pages": pages}
assert request.bound["model"] == "mistral/mistral-ocr-latest"
assert request.bound["document"] is document
assert request.bound["api_key"] == "test-key"
assert request.bound["api_base"] == "https://example.invalid"
assert request.bound["timeout"] is timeout
assert request.bound["custom_llm_provider"] == "mistral"
assert request.bound["extra_headers"] is extra_headers
assert request.kwargs == kwargs
assert request.kwargs["pages"] is pages
assert call_args is args
assert call_kwargs is kwargs
@ -102,14 +97,14 @@ def test_native_preserves_keyword_model_and_document_in_original_call_shape() ->
"document": document,
"pages": pages,
}
captured: Final[list[tuple[LiteLLMOcrRequest, tuple[object, ...], Mapping[str, object]]]] = []
captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = []
expected: Final = response()
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse:
args: Final = request.args
kwargs: Final = request.kwargs
captured.append((request, args, kwargs))
return expected
@ -118,15 +113,15 @@ def test_native_preserves_keyword_model_and_document_in_original_call_shape() ->
kwargs,
python=runtime.NO_PYTHON,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
request, call_args, call_kwargs = captured[0]
assert result is expected
assert request.model == "mistral/mistral-ocr-latest"
assert request.document is document
assert request.kwargs == {"pages": pages}
assert request.bound["model"] == "mistral/mistral-ocr-latest"
assert request.bound["document"] is document
assert request.kwargs == kwargs
assert call_args is args
assert call_kwargs is kwargs
assert call_kwargs["model"] == "mistral/mistral-ocr-latest"
@ -139,9 +134,7 @@ def test_aocr_marker_cannot_be_served_without_python() -> None:
kwargs: Final[Mapping[str, object]] = {"aocr": True}
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse:
pytest.fail("the aocr bypass marker must not reach native")
@ -151,7 +144,7 @@ def test_aocr_marker_cannot_be_served_without_python() -> None:
kwargs,
python=runtime.NO_PYTHON,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
@ -165,7 +158,7 @@ def test_missing_native_binding_is_a_required_rust_error() -> None:
{},
python=runtime.NO_PYTHON,
binding=ocr_binding(None),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
@ -174,9 +167,7 @@ def test_non_required_rule_cannot_be_served_without_python() -> None:
args: Final[tuple[object, ...]] = ("mistral/mistral-ocr-latest", {"type": "file", "file": b"pdf"})
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse:
return response()
@ -186,7 +177,7 @@ def test_non_required_rule_cannot_be_served_without_python() -> None:
{},
python=runtime.NO_PYTHON,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=(RouteRule(Route.OCR, Rollout.PYTHON_ONLY),),
)
@ -206,13 +197,9 @@ def test_non_required_rule_cannot_be_served_without_python() -> None:
),
),
)
def test_ocr_parser_errors_before_native(
args: tuple[object, ...], kwargs: Mapping[str, object], message: str
) -> None:
def test_ocr_parser_errors_before_native(args: tuple[object, ...], kwargs: Mapping[str, object], message: str) -> None:
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse:
pytest.fail("OCR parser failures must not call native")
@ -222,7 +209,7 @@ def test_ocr_parser_errors_before_native(
kwargs,
python=runtime.NO_PYTHON,
binding=ocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
@ -247,9 +234,7 @@ async def test_aocr_parser_errors_before_native(
args: tuple[object, ...], kwargs: Mapping[str, object], message: str
) -> None:
async def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse:
pytest.fail("OCR parser failures must not call native")
@ -259,7 +244,7 @@ async def test_aocr_parser_errors_before_native(
kwargs,
python=runtime.NO_PYTHON,
binding=aocr_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
@ -269,13 +254,11 @@ def test_public_ocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) ->
"type": "document_url",
"document_url": "https://example.invalid/document.pdf",
}
captured: Final[list[LiteLLMOcrRequest]] = []
captured: Final[list[NativeCall]] = []
expected: Final = response()
def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse:
captured.append(request)
return expected
@ -288,7 +271,7 @@ def test_public_ocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) ->
finally:
NATIVE_OCR.reset()
assert result is expected
assert [request.model for request in captured] == ["mistral/mistral-ocr-latest"]
assert [request.bound["model"] for request in captured] == ["mistral/mistral-ocr-latest"]
@pytest.mark.asyncio
@ -297,13 +280,11 @@ async def test_public_aocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPat
"type": "document_url",
"document_url": "https://example.invalid/document.pdf",
}
captured: Final[list[LiteLLMOcrRequest]] = []
captured: Final[list[NativeCall]] = []
expected: Final = response()
async def native(
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> OCRResponse:
captured.append(request)
return expected
@ -316,4 +297,4 @@ async def test_public_aocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPat
finally:
NATIVE_AOCR.reset()
assert result is expected
assert [request.model for request in captured] == ["mistral/mistral-ocr-latest"]
assert [request.bound["model"] for request in captured] == ["mistral/mistral-ocr-latest"]

View file

@ -15,10 +15,10 @@ from litellm.rust_bridge import catalog
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.catalog import Route, RouteRule
from litellm.rust_bridge.configuration import Rollout
from litellm.rust_bridge.public_call import NativeCall
from litellm.rust_bridge.responses.entrypoints import (
NATIVE_ARESPONSES,
NATIVE_RESPONSES,
LiteLLMResponsesRequest,
NativeAresponses,
NativeResponses,
)
@ -68,9 +68,7 @@ def test_python_route_forwards_original_call_shape() -> None:
return response
def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> ResponsesAPIResponse:
pytest.fail("Python-only dispatch must not call native")
@ -80,7 +78,7 @@ def test_python_route_forwards_original_call_shape() -> None:
kwargs,
python=python,
binding=responses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=PYTHON_RULES,
)
is response
@ -109,9 +107,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None:
return response
async def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> ResponsesAPIResponse:
pytest.fail("Python-only dispatch must not call native")
@ -120,7 +116,7 @@ async def test_async_python_route_forwards_original_call_shape() -> None:
kwargs,
python=python,
binding=aresponses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=PYTHON_RULES,
)
assert result is response
@ -144,17 +140,17 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None:
"custom_llm_provider": "anthropic",
"litellm_metadata": metadata,
}
captured: Final[list[tuple[LiteLLMResponsesRequest, tuple[object, ...], Mapping[str, object]]]] = []
captured: Final[list[tuple[NativeCall, tuple[object, ...], Mapping[str, object]]]] = []
response: Final = _response("anthropic/claude-sonnet-4-5")
def python(*call_args: object, **call_kwargs: object) -> ResponsesAPIResponse: # kwargs-ok: rejected fallback
pytest.fail("Required Rust dispatch must not call Python")
def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> ResponsesAPIResponse:
args: Final = request.args
kwargs: Final = request.kwargs
captured.append((request, args, kwargs))
return response
@ -163,24 +159,20 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None:
kwargs,
python=python,
binding=responses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
request, call_args, call_kwargs = captured[0]
assert result is response
assert request.model == "anthropic/claude-sonnet-4-5"
assert request.input is INPUT
assert request.stream is True
assert request.api_key == "sk-test"
assert request.api_base == "https://example.invalid"
assert request.custom_llm_provider == "anthropic"
assert request.extra_headers is extra_headers
assert request.kwargs == {
"api_key": "sk-test",
"base_url": "https://example.invalid",
"litellm_metadata": metadata,
}
assert request.bound["model"] == "anthropic/claude-sonnet-4-5"
assert request.bound["input"] is INPUT
assert request.bound["stream"] is True
assert request.bound["api_key"] == "sk-test"
assert request.bound["base_url"] == "https://example.invalid"
assert request.bound["custom_llm_provider"] == "anthropic"
assert request.bound["extra_headers"] is extra_headers
assert request.kwargs == kwargs
assert request.kwargs["litellm_metadata"] is metadata
assert call_args == args
assert call_args[0] is INPUT
@ -200,9 +192,7 @@ def test_internal_async_marker_bypasses_native() -> None:
return response
def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> ResponsesAPIResponse:
pytest.fail("aresponses' inner responses call must stay on Python")
@ -212,7 +202,7 @@ def test_internal_async_marker_bypasses_native() -> None:
kwargs,
python=python,
binding=responses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
is response
@ -236,9 +226,7 @@ def test_binding_errors_delegate_unchanged_to_python(args: tuple[object, ...], k
return response
def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> ResponsesAPIResponse:
pytest.fail("Binding failures must be delegated to Python")
@ -248,7 +236,7 @@ def test_binding_errors_delegate_unchanged_to_python(args: tuple[object, ...], k
kwargs,
python=python,
binding=responses_binding(native),
native=lambda hook, request, call_args, call_kwargs: hook(request, call_args, call_kwargs),
native=lambda hook, request, call_args, call_kwargs: hook(request),
rules=RUST_RULES,
)
is response
@ -257,13 +245,11 @@ def test_binding_errors_delegate_unchanged_to_python(args: tuple[object, ...], k
def test_public_responses_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None:
captured: Final[list[LiteLLMResponsesRequest]] = []
captured: Final[list[NativeCall]] = []
expected: Final = _response()
def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> ResponsesAPIResponse:
captured.append(request)
return expected
@ -276,18 +262,16 @@ def test_public_responses_routes_through_dispatch(monkeypatch: pytest.MonkeyPatc
finally:
NATIVE_RESPONSES.reset()
assert result is expected
assert [request.model for request in captured] == ["gpt-4o"]
assert [request.bound["model"] for request in captured] == ["gpt-4o"]
@pytest.mark.asyncio
async def test_public_aresponses_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> None:
captured: Final[list[LiteLLMResponsesRequest]] = []
captured: Final[list[NativeCall]] = []
expected: Final = _response()
async def native(
request: LiteLLMResponsesRequest,
args: tuple[object, ...],
kwargs: Mapping[str, object],
request: NativeCall,
) -> ResponsesAPIResponse:
captured.append(request)
return expected
@ -300,7 +284,7 @@ async def test_public_aresponses_routes_through_dispatch(monkeypatch: pytest.Mon
finally:
NATIVE_ARESPONSES.reset()
assert result is expected
assert [request.model for request in captured] == ["gpt-4o"]
assert [request.bound["model"] for request in captured] == ["gpt-4o"]
def test_responses_with_retries_uses_the_dispatch_entrypoint(monkeypatch: pytest.MonkeyPatch) -> None:
@ -323,6 +307,6 @@ def test_positional_parameters_remain_available_to_native_projection() -> None:
include: Final = ["reasoning.encrypted_content"]
request: Final = _DISPATCH.request((INPUT, "openai/test-model", include, "Be brief", 16), {})
assert request is not None
assert request.parameters["include"] is include
assert request.parameters["instructions"] == "Be brief"
assert request.parameters["max_output_tokens"] == 16
assert request.bound["include"] is include
assert request.bound["instructions"] == "Be brief"
assert request.bound["max_output_tokens"] == 16

View file

@ -2,7 +2,7 @@
Test what each side of the bridge does, not the rollout policy that picks a side. `LITELLM_RUST` and `catalog.RULES` change every time a route or backend rolls forward, so a test that sets the env var or patches the catalog to reach a path goes red on a policy change even when the code under test is fine
Call each path directly with an explicit decision instead. The Python path is the implementation the dispatcher falls back to, e.g. `litellm.responses.main.responses`. The Rust path is the native binding, e.g. `NATIVE_OCR.load()` from `litellm/rust_bridge/ocr/entrypoints.py`, called with the request, args and kwargs that dispatch would hand it. When the native side reads a policy-derived setting such as `settings.secret_manager().native`, pin that field in the test instead of deriving it from the catalog. `ocr/test_secrets.py` shows the pattern
Call each path directly with an explicit decision instead. The Python path is the implementation the dispatcher falls back to, e.g. `litellm.responses.main.responses`. The Rust path is the native binding, e.g. `NATIVE_OCR.load()` from `litellm/rust_bridge/ocr/entrypoints.py`, called with the `NativeCall` envelope that dispatch would hand it for every public inference route. When the native side reads a policy-derived setting such as `settings.secret_manager().native`, pin that field in the test instead of deriving it from the catalog. `ocr/test_secrets.py` shows the pattern
Rollout policy itself, meaning which rule matches and what `LITELLM_RUST` changes, belongs in `test_catalog.py`, `test_configuration.py` and `test_dispatch.py`, tested against rules the test builds rather than the shipped `catalog.RULES`

View file

@ -5,7 +5,6 @@ import pytest
import litellm
from litellm.rust_bridge.chat_completions.route_host import arguments, connection_defaults, response
from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest
from litellm.types.utils import ModelResponse
@ -38,18 +37,8 @@ def test_response_builds_the_public_model_response() -> None:
def test_arguments_are_the_public_kwargs_view() -> None:
kwargs: Final = MappingProxyType({"metadata": {"user_id": "u"}})
request: Final = LiteLLMChatCompletionsRequest(
model="anthropic/claude-sonnet-4-5",
messages=[{"role": "user", "content": "hi"}],
stream=None,
api_key=None,
api_base=None,
custom_llm_provider="anthropic",
extra_headers=None,
kwargs=kwargs,
)
assert arguments(request) is kwargs
assert arguments(kwargs) is kwargs
@pytest.mark.parametrize(

View file

@ -2,8 +2,7 @@ from types import MappingProxyType
from typing import Final
from litellm.rust_bridge.messages.route_host import arguments, response
from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
from dataclasses import astuple
from litellm.rust_bridge.public_call import NativeCall
import pytest
import litellm
from litellm.rust_bridge.messages import route_host
@ -30,67 +29,37 @@ def test_response_is_a_detached_public_messages_dict() -> None:
assert "_hidden_params" not in native
def test_arguments_are_the_public_kwargs_view() -> None:
def test_arguments_preserve_the_bound_view() -> None:
kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}})
request: Final = LiteLLMMessagesRequest(
model="claude-sonnet-4-5",
messages=[{"role": "user", "content": "hi"}],
max_tokens=16,
stream=None,
api_key=None,
api_base=None,
custom_llm_provider="anthropic",
request: Final = NativeCall(
args=(),
kwargs=kwargs,
)
assert arguments(request) is kwargs
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
def _flag_model(monkeypatch: pytest.MonkeyPatch, name: str, **flags: bool) -> None:
monkeypatch.setitem(
litellm.model_cost,
name,
{
"litellm_provider": "anthropic",
"mode": "chat",
"input_cost_per_token": 0,
"output_cost_per_token": 0,
**flags,
bound={
"model": "claude-sonnet-4-5",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 16,
"stream": None,
"api_key": None,
"api_base": None,
"custom_llm_provider": "anthropic",
**kwargs,
},
)
def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeypatch: pytest.MonkeyPatch) -> None:
_flag_model(
monkeypatch,
"claude-test-adaptive",
supports_reasoning=True,
supports_adaptive_thinking=True,
supports_output_config=True,
supports_xhigh_reasoning_effort=True,
supports_sampling_params=False,
)
capabilities: Final = route_host.model_capabilities("anthropic/claude-test-adaptive", None)
assert capabilities.supports_adaptive_thinking
assert capabilities.supports_output_config
assert not capabilities.supports_legacy_thinking
assert not capabilities.supports_sampling_params
assert capabilities.effort_tiers.xhigh
assert not capabilities.effort_tiers.max
assert arguments(request.bound) is request.bound
def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> None:
capabilities: Final = route_host.model_capabilities("anthropic/not-a-real-model", None)
def test_settings_project_caller_configuration_without_resolving_a_model(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "drop_params", False)
monkeypatch.setattr(litellm, "reasoning_auto_summary", True)
assert capabilities.supports_sampling_params
assert not capabilities.supports_reasoning
assert not capabilities.supports_adaptive_thinking
assert not any(astuple(capabilities.effort_tiers))
projected: Final = route_host.settings({"drop_params": "true", "additional_drop_params": ["metadata.user_id"]})
assert projected == {
"drop_params": True,
"reasoning_auto_summary": True,
"additional_drop_params": ("metadata.user_id",),
}
@pytest.mark.parametrize(
@ -108,7 +77,7 @@ def test_drop_params_merges_the_global_flag_with_the_request(
) -> None:
monkeypatch.setattr(litellm, "drop_params", global_flag)
assert route_host.shaping("anthropic/not-a-real-model", None, kwargs)["drop_params"] is expected
assert route_host.settings(kwargs)["drop_params"] is expected
@pytest.mark.parametrize(
@ -120,36 +89,42 @@ def test_drop_params_merges_the_global_flag_with_the_request(
],
)
def test_additional_drop_params_keep_only_string_paths(configured: object, expected: tuple[str, ...]) -> None:
shaping: Final = route_host.shaping("anthropic/not-a-real-model", None, {"additional_drop_params": configured})
settings: Final = route_host.settings({"additional_drop_params": configured})
assert shaping["additional_drop_params"] == expected
assert settings["additional_drop_params"] == expected
def test_native_request_rejections_map_to_the_public_400() -> None:
from types import MappingProxyType
from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest
from litellm.rust_bridge.public_call import NativeCall
request: Final = LiteLLMMessagesRequest(
model="anthropic/claude-sonnet-5",
messages=(),
max_tokens=8,
stream=None,
api_key=None,
api_base=None,
custom_llm_provider=None,
request: Final = NativeCall(
args=(),
kwargs=MappingProxyType({}),
bound={
"model": "anthropic/claude-sonnet-5",
"messages": (),
"max_tokens": 8,
"stream": None,
"api_key": None,
"api_base": None,
"custom_llm_provider": None,
**MappingProxyType({}),
},
)
rejected: Final = ValueError("claude-sonnet-5 does not support top_k=5")
rejected.messages_request_error = True # pyright: ignore[reportAttributeAccessIssue] # marker the native host sets
mapped: Final = route_host.map_failure(rejected, request, "anthropic")
mapped: Final = route_host.map_failure(rejected, request.bound, "anthropic")
assert isinstance(mapped, litellm.BadRequestError)
assert mapped.status_code == 400
assert "does not support top_k=5" in mapped.message
assert mapped.model == "claude-sonnet-5"
assert not isinstance(route_host.map_failure(ValueError("plain"), request, "anthropic"), litellm.BadRequestError)
assert not isinstance(
route_host.map_failure(ValueError("plain"), request.bound, "anthropic"), litellm.BadRequestError
)
def test_stream_hidden_params_projects_upstream_headers_the_way_the_python_handler_does() -> None:

View file

@ -2,7 +2,6 @@ from __future__ import annotations
from collections.abc import Awaitable, Mapping
from dataclasses import replace
from types import MappingProxyType
from typing import Final, Protocol, cast # noqa: TID251 # narrows the parametrized path to its protocol
import httpx
@ -12,7 +11,8 @@ import litellm
from litellm.integrations.custom_secret_manager import CustomSecretManager
from litellm.llms.anthropic.pass_through.messages.handler import anthropic_messages
from litellm.rust_bridge import settings
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, NATIVE_MESSAGES, LiteLLMMessagesRequest
from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, NATIVE_MESSAGES
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem
from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service
from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_MODEL, MESSAGES_RESPONSE
@ -48,16 +48,21 @@ class _ManagedSecrets(CustomSecretManager):
return self.values.get(secret_name)
def _native_request() -> LiteLLMMessagesRequest:
return LiteLLMMessagesRequest(
model=MESSAGES_MODEL,
messages=MESSAGES,
max_tokens=8,
stream=None,
api_key=None,
api_base=None,
custom_llm_provider=None,
kwargs=MappingProxyType({}),
def _native_request() -> NativeCall:
supplied: Final = _public_kwargs()
return NativeCall(
args=(),
kwargs=supplied,
bound={
"model": MESSAGES_MODEL,
"messages": MESSAGES,
"max_tokens": 8,
"stream": None,
"api_key": None,
"api_base": None,
"custom_llm_provider": None,
**supplied,
},
)
@ -72,13 +77,13 @@ async def _python_messages() -> object:
async def _rust_messages() -> object:
route: Final = NATIVE_MESSAGES.load()
assert route is not None
return route(_native_request(), (), _public_kwargs())
return route(_native_request())
async def _rust_amessages() -> object:
route: Final = NATIVE_AMESSAGES.load()
assert route is not None
return await route(_native_request(), (), _public_kwargs())
return await route(_native_request())
@pytest.fixture(

View file

@ -14,6 +14,7 @@ from http.client import HTTPMessage
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from socket import socket as Socket
from types import SimpleNamespace
from typing import Final
REQUEST_STARTED: Final = threading.Event()
@ -27,6 +28,10 @@ ANTHROPIC_RESPONSE: Final = (
)
class NativeRouteServer(ThreadingHTTPServer):
request_queue_size = 64
class NativeRouteHandler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
@ -110,6 +115,11 @@ def load_native(native_path: Path) -> object:
return native_module
def route_call(route: str, api_base: str, outcome: str) -> SimpleNamespace:
fields: Final = route_kwargs(route, api_base, outcome)
return SimpleNamespace(args=(), kwargs=fields, bound=fields)
def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
common: Final = {
"api_base": api_base,
@ -161,9 +171,9 @@ def assert_rate_limit(route: str, error: BaseException) -> None:
def exercise_sync(native: object, api_base: str) -> None:
for route in ("transcription", "chat_completions"):
function: Final = getattr(native, route)
assert_success(route, function(**route_kwargs(route, api_base, "success")))
assert_success(route, function(route_call(route, api_base, "success")))
try:
function(**route_kwargs(route, api_base, "429"))
function(route_call(route, api_base, "429"))
except native.RustUpstreamError as error:
assert_rate_limit(route, error)
else:
@ -173,18 +183,18 @@ def exercise_sync(native: object, api_base: str) -> None:
async def exercise_async(native: object, api_base: str) -> None:
for route in ("transcription", "chat_completions"):
function: Final = getattr(native, f"a{route}")
assert_success(route, await function(**route_kwargs(route, api_base, "success")))
assert_success(route, await function(route_call(route, api_base, "success")))
try:
await function(**route_kwargs(route, api_base, "429"))
await function(route_call(route, api_base, "429"))
except native.RustUpstreamError as error:
assert_rate_limit(route, error)
else:
raise AssertionError(f"a{route} accepted a 429 response")
async def exercise_async_concurrency(native: object, api_base: str) -> None:
responses: Final = await asyncio.wait_for(
asyncio.gather(*(native.achat_completions(**route_kwargs("chat_completions", api_base, "success")) for _ in range(32))),
asyncio.gather(
*(native.achat_completions(route_call("chat_completions", api_base, "success")) for _ in range(32))
),
timeout=15,
)
for response in responses:
@ -195,14 +205,13 @@ def exercise_routes(native_path: Path, api_base: str) -> object:
native: Final = load_native(native_path)
exercise_sync(native, api_base)
asyncio.run(exercise_async(native, api_base))
asyncio.run(exercise_async_concurrency(native, api_base))
return native
def exercise_signal(native: object, api_base: str) -> int:
try:
native.chat_completions(
**route_kwargs("chat_completions", api_base, "hang"),
route_call("chat_completions", api_base, "hang"),
)
except KeyboardInterrupt:
sys.stdout.write("KeyboardInterrupt\n")
@ -264,7 +273,7 @@ def verify_wheel(wheel: Path) -> int:
raise AssertionError(f"expected one native extension, found {len(native_members)}")
native_path: Final = wheel_root / native_members[0].filename
server: Final = ThreadingHTTPServer(("127.0.0.1", 0), NativeRouteHandler)
server: Final = NativeRouteServer(("127.0.0.1", 0), NativeRouteHandler)
server_thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
server_thread.start()
api_base: Final = f"http://127.0.0.1:{server.server_address[1]}"

View file

@ -5,17 +5,21 @@ import pytest
import litellm
from litellm.rust_bridge.ocr.route_host import UpstreamFailure, map_failure
from litellm.rust_bridge.ocr.route_host import response as build_ocr_response
from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest
from litellm.rust_bridge.public_call import NativeCall
REQUEST: Final = LiteLLMOcrRequest(
model="mistral/mistral-ocr-latest",
document={"type": "document_url", "document_url": "https://example.com/file.pdf"},
api_key="test-key",
api_base=None,
timeout=None,
custom_llm_provider=None,
extra_headers=None,
REQUEST: Final = NativeCall(
args=(),
kwargs={"req_format": "markdown"},
bound={
"model": "mistral/mistral-ocr-latest",
"document": {"type": "document_url", "document_url": "https://example.com/file.pdf"},
"api_key": "test-key",
"api_base": None,
"timeout": None,
"custom_llm_provider": None,
"extra_headers": None,
**{"req_format": "markdown"},
},
)
@ -49,7 +53,7 @@ def test_rust_ocr_response_retains_provider_native_response():
def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> None:
error: Final = RustUpstreamError(429, '{"message": "slow down"}', (("retry-after", "7"),))
public_error: Final = map_failure(error, REQUEST, "mistral")
public_error: Final = map_failure(error, REQUEST.bound, "mistral")
assert isinstance(public_error, litellm.RateLimitError)
assert public_error.status_code == 429
@ -62,7 +66,7 @@ def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> N
def test_map_failure_maps_upstream_401_to_authentication_error() -> None:
error: Final = RustUpstreamError(401, '{"message": "Unauthorized"}', ())
public_error: Final = map_failure(error, REQUEST, "mistral")
public_error: Final = map_failure(error, REQUEST.bound, "mistral")
assert isinstance(public_error, litellm.AuthenticationError)
assert public_error.status_code == 401
@ -73,7 +77,7 @@ def test_map_failure_maps_upstream_401_to_authentication_error() -> None:
def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None:
error: Final = RuntimeError("bridge exploded")
public_error: Final = map_failure(error, REQUEST, "mistral")
public_error: Final = map_failure(error, REQUEST.bound, "mistral")
assert not isinstance(public_error, UpstreamFailure)
assert isinstance(public_error, litellm.APIConnectionError)
@ -82,4 +86,4 @@ def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None:
def test_map_failure_reports_invalid_request_format_as_unsupported_params() -> None:
with pytest.raises(litellm.UnsupportedParamsError, match="Invalid `req_format`: 'markdown'"):
raise map_failure(RustFormatError(), REQUEST, "mistral")
raise map_failure(RustFormatError(), REQUEST.bound, "mistral")

View file

@ -4,7 +4,6 @@ import asyncio
from collections.abc import Awaitable, Generator, Mapping
from contextlib import contextmanager
from dataclasses import replace
from types import MappingProxyType
from typing import Final, Literal, Protocol, TypeAlias, cast
import httpx
@ -14,7 +13,8 @@ import litellm
from litellm.integrations.custom_secret_manager import CustomSecretManager
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge import settings
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest
from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec, recording_service
from tests.test_litellm_rust.support.requests import OCR_DOCUMENT, OCR_MODEL, OCR_RESPONSE
@ -59,16 +59,21 @@ class _VaultSecrets(CustomSecretManager):
return tuple(params for name, params in self.reads if name == "MISTRAL_API_KEY")
def _native_request(api_base: str) -> LiteLLMOcrRequest:
return LiteLLMOcrRequest(
model=OCR_MODEL,
document=OCR_DOCUMENT,
api_key=None,
api_base=api_base,
timeout=None,
custom_llm_provider=None,
extra_headers=None,
kwargs=MappingProxyType({}),
def _native_request(api_base: str) -> NativeCall:
supplied: Final = _public_kwargs(api_base)
return NativeCall(
args=(),
kwargs=supplied,
bound={
"model": OCR_MODEL,
"document": OCR_DOCUMENT,
"api_key": None,
"api_base": api_base,
"timeout": None,
"custom_llm_provider": None,
"extra_headers": None,
**supplied,
},
)
@ -79,13 +84,13 @@ def _public_kwargs(api_base: str) -> dict[str, object]:
async def _rust_ocr(api_base: str) -> OCRResponse:
route: Final = NATIVE_OCR.load()
assert route is not None
return route(_native_request(api_base), (), _public_kwargs(api_base))
return route(_native_request(api_base))
async def _rust_aocr(api_base: str) -> OCRResponse:
route: Final = NATIVE_AOCR.load()
assert route is not None
return await route(_native_request(api_base), (), _public_kwargs(api_base))
return await route(_native_request(api_base))
_RUST_PATHS: Final = (_rust_ocr, _rust_aocr)

View file

@ -5,8 +5,8 @@ import pytest
from pydantic import ValidationError
import litellm
from litellm.rust_bridge.responses.route_host import arguments, connection_defaults, response
from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest
from litellm.rust_bridge.responses.route_host import arguments, connection_defaults, map_failure, response
from litellm.rust_bridge.public_call import NativeCall
from litellm.types.llms.openai import ResponsesAPIResponse
@ -42,20 +42,24 @@ def test_response_rejects_a_payload_missing_required_fields() -> None:
response(MappingProxyType({"object": "response"}))
def test_arguments_are_the_public_kwargs_view() -> None:
def test_arguments_preserve_the_bound_view() -> None:
kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}})
request: Final = LiteLLMResponsesRequest(
model="gpt-4o",
input="hi",
stream=None,
api_key=None,
api_base=None,
custom_llm_provider="openai",
extra_headers=None,
request: Final = NativeCall(
args=(),
kwargs=kwargs,
bound={
"model": "gpt-4o",
"input": "hi",
"stream": None,
"api_key": None,
"api_base": None,
"custom_llm_provider": "openai",
"extra_headers": None,
**kwargs,
},
)
assert arguments(request) is kwargs
assert arguments(request.bound) is request.bound
@pytest.mark.parametrize(
@ -74,3 +78,23 @@ def test_connection_defaults_preserve_openai_precedence(
monkeypatch.setattr(litellm, "openai_key", provider_key)
monkeypatch.setattr(litellm, "api_base", "https://configured.invalid/v1")
assert connection_defaults("openai") == (expected, litellm.api_base)
class _UpstreamFailure(Exception):
headers: Final = ()
@pytest.mark.parametrize(
("api_base", "base_url", "expected"),
(
(None, "https://alias.invalid/v1", "https://alias.invalid/v1"),
("", "https://alias.invalid/v1", "https://alias.invalid/v1"),
("https://base.invalid/v1", "https://alias.invalid/v1", "https://base.invalid/v1"),
),
)
def test_failure_preserves_the_explicit_endpoint(api_base: str | None, base_url: str, expected: str) -> None:
upstream: Final = _UpstreamFailure(429, '{"error":{"message":"rate limited"}}')
mapped: Final = map_failure(upstream, {"model": "openai/test-model", "api_base": api_base, "base_url": base_url})
assert isinstance(mapped, litellm.RateLimitError)
assert str(mapped.response.request.url) == expected
assert mapped.__context__ is upstream

View file

@ -0,0 +1,73 @@
from typing import Final
import pytest
import litellm
from litellm.rust_bridge.model_capabilities import anthropic_model_capabilities
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
def _flag_model(monkeypatch: pytest.MonkeyPatch, name: str, **flags: bool) -> None:
monkeypatch.setitem(
litellm.model_cost,
name,
{
"litellm_provider": "anthropic",
"mode": "chat",
"input_cost_per_token": 0,
"output_cost_per_token": 0,
**flags,
},
)
def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeypatch: pytest.MonkeyPatch) -> None:
_flag_model(
monkeypatch,
"claude-test-adaptive",
supports_reasoning=True,
supports_adaptive_thinking=True,
supports_output_config=True,
supports_xhigh_reasoning_effort=True,
supports_sampling_params=False,
)
capabilities: Final = anthropic_model_capabilities("anthropic/claude-test-adaptive", None)
assert capabilities["supports_adaptive_thinking"]
assert capabilities["supports_output_config"]
assert not capabilities["supports_legacy_thinking"]
assert not capabilities["supports_sampling_params"]
assert capabilities["effort_tiers"] == {
"minimal": False,
"low": False,
"medium": False,
"high": False,
"xhigh": True,
"max": False,
}
def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> None:
capabilities: Final = anthropic_model_capabilities("anthropic/not-a-real-model", None)
assert capabilities["supports_sampling_params"]
assert not capabilities["supports_reasoning"]
assert not capabilities["supports_adaptive_thinking"]
assert capabilities["effort_tiers"] == dict.fromkeys(("minimal", "low", "medium", "high", "xhigh", "max"), False)
def test_capability_source_observes_runtime_registration_changes(monkeypatch: pytest.MonkeyPatch) -> None:
model: Final = "claude-test-runtime-registration"
_flag_model(monkeypatch, model, supports_reasoning=True, supports_output_config=True)
before: Final = anthropic_model_capabilities(model, "anthropic")
_flag_model(monkeypatch, model, supports_reasoning=False, supports_output_config=False)
after: Final = anthropic_model_capabilities(model, "anthropic")
assert before["supports_reasoning"] is True
assert before["supports_output_config"] is True
assert after["supports_reasoning"] is False
assert after["supports_output_config"] is False

View file

@ -0,0 +1,61 @@
from collections.abc import Mapping, Sequence
from typing import Final
import pytest
from litellm.rust_bridge.public_call import bind, native_call, signature
def _messages(
max_tokens: int,
messages: Sequence[object],
model: str,
temperature: float | None = None,
api_key: str | None = None,
**kwargs: object, # kwargs-ok: exercise the public signature binding contract
) -> None:
return None
@pytest.mark.parametrize("supplied", ({}, {"api_key": None}, {"api_key": "explicit"}))
def test_native_call_preserves_omission_separately_from_bound_defaults(supplied: Mapping[str, object]) -> None:
messages: Final[Sequence[object]] = [{"role": "user", "content": "hello"}]
args: Final = (128, messages, "model", 0.25)
fields: Final = bind(signature(_messages), args, supplied)
assert fields is not None
call: Final = native_call(args, supplied, fields)
assert call.args is args
assert call.kwargs is supplied
assert call.bound == {
"max_tokens": 128,
"messages": messages,
"model": "model",
"temperature": 0.25,
"api_key": supplied.get("api_key"),
}
assert call.bound["messages"] is messages
assert ("api_key" in call.kwargs) == ("api_key" in supplied)
def test_native_call_keeps_extra_option_objects_without_nested_kwargs() -> None:
messages: Final[Sequence[object]] = []
metadata: Final = {"trace": "caller"}
supplied: Final = {"metadata": metadata}
args: Final = (128, messages, "model")
fields: Final = bind(signature(_messages), args, supplied)
assert fields is not None
call: Final = native_call(args, supplied, fields)
assert call.bound == {
"max_tokens": 128,
"messages": messages,
"model": "model",
"temperature": None,
"api_key": None,
"metadata": metadata,
}
assert call.bound["metadata"] is metadata
assert supplied == {"metadata": metadata}