mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
2803a16b36
commit
2aaa0b5d5c
58 changed files with 1302 additions and 1192 deletions
|
|
@ -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
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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}))
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(¶ms).unwrap(),
|
||||
optional_object("optional_params", ¶ms).unwrap(),
|
||||
Some(required_object("optional_params", json!({"temperature": 0.2})).unwrap())
|
||||
);
|
||||
assert_eq!(
|
||||
optional_params_argument(&py.None().into_bound(py)).unwrap(),
|
||||
optional_object("optional_params", &py.None().into_bound(py)).unwrap(),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
extra_headers_argument(&py.None().into_bound(py)).unwrap(),
|
||||
optional_object("extra_headers", &py.None().into_bound(py)).unwrap(),
|
||||
None
|
||||
);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,))?;
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
},
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
|
|
|||
|
|
@ -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]: ...
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]: ...
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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]]: ...
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
35
litellm/rust_bridge/model_capabilities.py
Normal file
35
litellm/rust_bridge/model_capabilities.py
Normal 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")},
|
||||
}
|
||||
|
|
@ -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]: ...
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]: ...
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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),),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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),),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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`
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]}"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
73
tests/unit/rust_bridge/test_model_capabilities.py
Normal file
73
tests/unit/rust_bridge/test_model_capabilities.py
Normal 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
|
||||
61
tests/unit/rust_bridge/test_public_call.py
Normal file
61
tests/unit/rust_bridge/test_public_call.py
Normal 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}
|
||||
Loading…
Add table
Reference in a new issue