This commit is contained in:
Yujong Lee 2026-09-14 16:52:42 -07:00
parent c8c916e549
commit 6276022f34
105 changed files with 2955 additions and 2672 deletions

View file

@ -299,6 +299,9 @@ test-rust-extension:
[ "$$#" -eq 1 ] && \
UV_PROJECT_ENVIRONMENT="$$temporary/venv" $(UV) sync --python 3.12 --frozen --no-install-project --all-groups --all-extras && \
$(UV) pip install --python "$$temporary/venv/bin/python" --no-deps "$$1" && \
"$$temporary/venv/bin/python" -I -m mypy.stubtest \
--mypy-config-file tests/test_litellm/rust_bridge/stubtest.ini \
litellm.rust_bridge._native && \
LITELLM_RUST=1 LITELLM_LOCAL_MODEL_COST_MAP=True \
"$$temporary/venv/bin/python" -I -m pytest --import-mode=importlib -m requires_rust_extension tests/test_litellm_rust

View file

@ -17,5 +17,25 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu
.await
}
pub fn admit(
model: &str,
provider: Option<&str>,
audio: &Value,
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
use crate::call_lifecycle::admission::AdmissionDecline;
let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider);
let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider));
if provider.and_then(prepare::provider_config).is_none() {
return Err(AdmissionDecline::Provider);
}
if let Some(format) = audio.get("format").and_then(Value::as_str)
&& !matches!(format, "wav" | "mp3" | "flac" | "ogg")
{
return Err(AdmissionDecline::Feature("unsupported audio format"));
}
Ok(())
}
#[cfg(test)]
mod tests;

View file

@ -8,7 +8,9 @@ use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderCo
use super::types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
pub(super) fn provider_config(
provider: &str,
) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
#[cfg(feature = "bedrock-auth")]
if provider == "bedrock" {
return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG);

View file

@ -18,3 +18,13 @@ pub enum UnimplementedRoute {
pub fn admit_unimplemented(route: UnimplementedRoute) -> Result<Infallible, UnimplementedRoute> {
Err(route)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, strum::Display)]
pub enum AdmissionDecline {
#[strum(to_string = "provider is not supported by this native route")]
Provider,
#[strum(to_string = "required host operations are not supported")]
HostOperations,
#[strum(to_string = "{0}")]
Feature(&'static str),
}

View file

@ -55,5 +55,55 @@ pub fn chat_completions_decline_reason(
.map(|reason| reason.0)
}
#[derive(Clone, Copy, Debug, Default, serde::Deserialize)]
#[serde(deny_unknown_fields)]
pub struct AdmissionContext {
#[serde(default)]
pub stream: bool,
#[serde(default)]
pub anthropic_user_id: bool,
#[serde(default)]
pub bedrock_metadata_owned: bool,
}
pub fn admit(
model: &str,
provider: Option<&str>,
messages: Value,
params: &Map<String, Value>,
headers: Option<&Map<String, Value>>,
context: AdmissionContext,
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
use crate::call_lifecycle::admission::AdmissionDecline;
let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider);
let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider));
if context.stream {
return Err(AdmissionDecline::Feature("streaming"));
}
if (provider == Some("anthropic") && context.anthropic_user_id)
|| (provider == Some("bedrock") && context.bedrock_metadata_owned)
{
return Err(AdmissionDecline::HostOperations);
}
#[cfg(feature = "bedrock-auth")]
if provider == Some("bedrock")
&& headers.is_some_and(|headers| {
headers
.keys()
.any(|name| crate::providers::bedrock::aws_base::is_sigv4_computed_header(name))
})
{
return Err(AdmissionDecline::Feature(
"request forwards a header AWS SigV4 computes",
));
}
let _ = headers;
match chat_completions_decline_reason(model, provider, messages, params) {
Some(reason) => Err(AdmissionDecline::Feature(reason)),
None => Ok(()),
}
}
#[cfg(test)]
mod tests;

View file

@ -27,5 +27,26 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Re
execute_messages_provider_stream(request).await
}
pub fn admit(
model: &str,
provider: Option<&str>,
has_agentic_hook: bool,
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
use crate::call_lifecycle::admission::AdmissionDecline;
let resolved = crate::routing_utils::provider::get_custom_llm_provider(model, provider);
let provider = provider.or_else(|| resolved.as_ref().map(|value| value.custom_llm_provider));
if provider
.and_then(common_utils::messages_provider_config)
.is_none()
{
return Err(AdmissionDecline::Provider);
}
if has_agentic_hook {
return Err(AdmissionDecline::HostOperations);
}
Ok(())
}
#[cfg(test)]
mod tests;

View file

@ -19,6 +19,15 @@ pub use lifecycle::{
};
pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
pub fn admit_value(
model: &str,
provider: Option<&str>,
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
registry::resolve_wire_adapter(model, provider)
.map(|_| ())
.map_err(|_| crate::call_lifecycle::admission::AdmissionDecline::Provider)
}
#[cfg(test)]
#[path = "../../tests/azure_ai_ocr.rs"]
mod azure_ai_tests;

View file

@ -335,3 +335,12 @@ mod tests {
assert!(!nested_without_flat.data.contains_key("model"));
}
}
pub fn admit(
provider: Option<&str>,
) -> Result<(), crate::call_lifecycle::admission::AdmissionDecline> {
match provider {
Some("openai") => Ok(()),
_ => Err(crate::call_lifecycle::admission::AdmissionDecline::Provider),
}
}

View file

@ -9,6 +9,18 @@ pyo3::create_exception!(
"Core admission declined without effects, so the host may select its legacy path once."
);
pyo3::create_exception!(
_native,
RustHostCallbackError,
pyo3::exceptions::PyException
);
pyo3::create_exception!(
_native,
RustBridgeUnavailable,
pyo3::exceptions::PyException
);
pyo3::create_exception!(
_native,
RustUpstreamError,
@ -31,7 +43,7 @@ pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr {
pub(crate) fn execution_error_to_pyerr(error: Error) -> PyErr {
match error {
Error::Http { status, body } => RustUpstreamError::new_err((status, body)),
Error::Network(message) | Error::InvalidResponse(message) => {
Error::Connect(message) | Error::Network(message) | Error::InvalidResponse(message) => {
RustUpstreamError::new_err((0u16, message))
}
other => core_error_to_pyerr(other),
@ -40,10 +52,25 @@ pub(crate) fn execution_error_to_pyerr(error: Error) -> PyErr {
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
let py = module.py();
module.add(
"RustBridgeUnavailable",
py.get_type::<RustBridgeUnavailable>(),
)?;
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
module.add(
"RustHostCallbackError",
py.get_type::<RustHostCallbackError>(),
)?;
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
}
pub(crate) fn host_callback_error(py: Python<'_>, error: PyErr) -> PyErr {
let wrapped = RustHostCallbackError::new_err(error.to_string());
wrapped.set_context(py, Some(error.clone_ref(py)));
wrapped.set_cause(py, Some(error));
wrapped
}
#[cfg(test)]
mod tests {
use super::*;
@ -70,7 +97,6 @@ mod tests {
Error::MissingAzureDocumentIntelligenceCredentials,
Error::MissingReductoApiKey,
Error::Routing("routing failed".into()),
Error::Connect("connection refused".into()),
] {
let expected = core_error_to_pyerr(error.clone());
let actual = execution_error_to_pyerr(error);
@ -86,6 +112,7 @@ mod tests {
Python::initialize();
Python::attach(|py| {
for (error, status, message) in [
(Error::Connect("connection refused".into()), 0, "connection refused"),
(
Error::Http {
status: 429,
@ -114,3 +141,9 @@ mod tests {
});
}
}
pub(crate) fn admit(
result: Result<(), litellm_core::call_lifecycle::admission::AdmissionDecline>,
) -> PyResult<()> {
result.map_err(|reason| RustBridgeDeclined::new_err(reason.to_string()))
}

View file

@ -8,7 +8,6 @@ mod function_trace;
mod lifecycle;
mod marshal;
mod routes;
mod token_counter;
use pyo3::prelude::*;
@ -20,7 +19,6 @@ mod _native {
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::errors::register(module)?;
super::routes::register(module)?;
super::token_counter::register(module)?;
super::diagnostics::register(module)
}
}
@ -44,7 +42,9 @@ mod tests {
let module = pyo3::wrap_pymodule!(_native)(py).into_bound(py);
let expected = [
"RustBridgeUnavailable",
"RustBridgeDeclined",
"RustHostCallbackError",
"RustUpstreamError",
"ocr",
"aocr",
@ -52,11 +52,10 @@ mod tests {
"atranscription",
"messages",
"amessages",
"chat_completions_decline",
"chat_completions",
"achat_completions",
"ResponsesWebSocketConnection",
"TokenCounter",
"count_input_tokens",
"gil_stats",
];
@ -148,7 +147,7 @@ mod tests {
import asyncio
async def exercise():
connection = await native.ResponsesWebSocketConnection.connect(url)
connection = await native.ResponsesWebSocketConnection.connect(url, custom_llm_provider="openai")
assert type(connection) is native.ResponsesWebSocketConnection
await connection.send_text("from-python")
assert await connection.recv_text() == "from-server"

View file

@ -2,9 +2,7 @@ use litellm_core::Error;
use std::future::Future;
use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse};
use litellm_core::chat_completions::{
chat_completions as run_chat_completions, chat_completions_decline_reason,
};
use litellm_core::chat_completions::{AdmissionContext, chat_completions as run_chat_completions};
use pyo3::prelude::*;
use serde_json::Value;
@ -16,6 +14,12 @@ fn prepare_chat_completions(
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
let messages = required_array("messages", inputs.messages)?;
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
let context: AdmissionContext = inputs
.host_facts
.map(serde_json::from_value)
.transpose()
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?
.unwrap_or_default();
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
@ -25,6 +29,23 @@ fn prepare_chat_completions(
timeout_seconds: inputs.timeout_seconds,
})?;
crate::errors::admit(litellm_core::chat_completions::admit(
&options.model,
options.custom_llm_provider.as_deref(),
Value::Array(messages.clone()),
&optional_params,
options.extra_headers.as_ref(),
context,
))?;
if let Some(on_request) = inputs.on_request {
Python::attach(|py| {
on_request
.call0(py)
.map(|_| ())
.map_err(|error| crate::errors::host_callback_error(py, error))
})?;
}
Ok(async move {
let RouteOptions {
model,
@ -48,24 +69,6 @@ fn prepare_chat_completions(
})
}
#[pyfunction]
#[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))]
fn chat_completions_decline(
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] messages: Value,
#[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option<Value>,
custom_llm_provider: Option<String>,
) -> PyResult<Option<String>> {
let optional_params = object_or_empty("optional_params", optional_params)?;
Ok(chat_completions_decline_reason(
&model,
custom_llm_provider.as_deref(),
messages,
&optional_params,
)
.map(str::to_string))
}
bridge_route! {
sync = chat_completions,
asynchronous = achat_completions,
@ -84,8 +87,10 @@ bridge_route! {
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<serde_json::Value>,
timeout_seconds: Option<f64>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
host_facts: Option<Value>,
on_request: Option<Py<PyAny>>,
},
prepare = prepare_chat_completions,
errors = execution_error_to_pyerr,
extra = [chat_completions_decline],
}

View file

@ -5,17 +5,18 @@ use pyo3::types::PyCFunction;
macro_rules! unimplemented_lifecycle_route {
($route:ident, $entrypoint:ident) => {
#[pyo3::pyfunction]
#[pyo3(signature = (request, args, kwargs, asynchronous))]
#[pyo3(signature = (request, args, kwargs, asynchronous, host))]
fn $entrypoint(
request: pyo3::Bound<'_, pyo3::PyAny>,
args: pyo3::Bound<'_, pyo3::types::PyTuple>,
kwargs: pyo3::Bound<'_, pyo3::types::PyDict>,
asynchronous: bool,
host: pyo3::Bound<'_, pyo3::PyAny>,
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
use litellm_core::call_lifecycle::admission::{
UnimplementedRoute, admit_unimplemented,
};
let _ = (request, args, kwargs, asynchronous);
let _ = (request, args, kwargs, asynchronous, host);
match admit_unimplemented(UnimplementedRoute::$route) {
Ok(never) => match never {},
Err(route) => Err($crate::errors::RustBridgeDeclined::new_err(format!(
@ -268,12 +269,12 @@ mod tests {
(
"messages",
"amessages",
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)",
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, has_agentic_hook=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)",
"(model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, host_facts=None, on_request=None)",
),
];
@ -458,46 +459,6 @@ mod tests {
});
}
#[test]
fn chat_completions_decline_keeps_existing_reasons() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "routes").expect("module should be created");
crate::routes::register(&module).expect("routes should register");
let decline = module
.getattr("chat_completions_decline")
.expect("decline helper should be registered");
let empty = PyList::empty(py);
let unreadable = py
.eval(c"'nope'", None, None)
.expect("string messages should convert");
let unknown: Option<String> = decline
.call1(("unknown-model", &empty))
.and_then(|value| value.extract())
.expect("unknown providers should decline");
assert_eq!(
unknown.as_deref(),
Some("provider is not on the rust chat completions path")
);
let empty_reason: Option<String> = decline
.call1(("anthropic/claude-sonnet-4-5", &empty))
.and_then(|value| value.extract())
.expect("empty lists should decline");
assert_eq!(empty_reason.as_deref(), Some("empty message list"));
let unreadable_reason: Option<String> = decline
.call1(("anthropic/claude-sonnet-4-5", unreadable))
.and_then(|value| value.extract())
.expect("non-list messages should decline");
assert_eq!(
unreadable_reason.as_deref(),
Some("unreadable message list")
);
});
}
#[test]
fn generated_routes_execute_sync_and_async_contracts() {
Python::initialize();

View file

@ -5,7 +5,7 @@ use pyo3::prelude::*;
use serde_json::Value;
use std::future::Future;
use crate::errors::core_error_to_pyerr;
use crate::errors::{admit, execution_error_to_pyerr};
use crate::marshal::{RouteOptions, RouteOptionsInputs, required_object};
fn prepare_messages(
@ -21,6 +21,12 @@ fn prepare_messages(
timeout_seconds: inputs.timeout_seconds,
})?;
admit(litellm_core::messages::admit(
&options.model,
options.custom_llm_provider.as_deref(),
inputs.has_agentic_hook.unwrap_or(false),
))?;
Ok(async move {
let RouteOptions {
model,
@ -59,7 +65,8 @@ bridge_route! {
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<serde_json::Value>,
timeout_seconds: Option<f64>,
has_agentic_hook: Option<bool>,
},
prepare = prepare_messages,
errors = core_error_to_pyerr,
errors = execution_error_to_pyerr,
}

View file

@ -13,6 +13,7 @@ mod ocr;
mod rerank;
mod responses;
mod speech;
mod token_counter;
mod transcription;
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
@ -27,6 +28,7 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
rerank::register(module)?;
responses::register(module)?;
speech::register(module)?;
token_counter::register(module)?;
#[cfg(feature = "trace-parity")]
{

View file

@ -29,6 +29,7 @@ impl PythonLogger {
pub(super) fn update_ocr(
&self,
py: Python<'_>,
host: &Py<PyAny>,
kwargs: &Py<PyDict>,
pre_call: &OcrLoggingFields,
secret_fields: &[&str],
@ -58,7 +59,7 @@ impl PythonLogger {
params.set_item(name, value)?;
}
}
for name in custom_pricing_fields(py)? {
for name in custom_pricing_fields(py, host)? {
if let Some(value) = kwargs.bind(py).get_item(&name)?
&& !value.is_none()
{
@ -127,15 +128,10 @@ impl PythonLogger {
}
}
fn custom_pricing_fields(py: Python<'_>) -> PyResult<Vec<String>> {
py.import("litellm.types.utils")?
.getattr("CustomPricingLiteLLMParams")?
.getattr("model_fields")?
.cast_into::<PyDict>()?
.keys()
.iter()
.map(|name| name.extract::<String>())
.collect()
fn custom_pricing_fields(py: Python<'_>, host: &Py<PyAny>) -> PyResult<Vec<String>> {
host.bind(py)
.call_method0("custom_pricing_fields")?
.extract()
}
fn redact(
@ -158,21 +154,26 @@ fn redact(
Ok(redacted.unbind())
}
pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
py.import("litellm.rust_bridge.ocr.value")?
.getattr("_response")?
pub(super) fn response(
py: Python<'_>,
host: &Py<PyAny>,
response: &LiteLLMOcrResponse,
) -> PyResult<Py<PyAny>> {
host.bind(py)
.getattr("response")?
.call1((to_py(py, response)?,))
.map(Bound::unbind)
}
pub(super) fn map_failure(
py: Python<'_>,
host: &Py<PyAny>,
error: &Py<PyBaseException>,
request: &Bound<'_, PyAny>,
provider: &str,
) -> PyResult<Py<PyBaseException>> {
Ok(py
.import("litellm.rust_bridge.ocr.lifecycle")?
Ok(host
.bind(py)
.getattr("map_failure")?
.call1((error, request, provider))?
.extract()?)

View file

@ -17,6 +17,7 @@ use crate::lifecycle::{
struct PythonOcrHost {
state: PythonCallState,
adapter: Py<PyAny>,
data: OcrHostData,
}
@ -92,6 +93,7 @@ impl PythonOcrHost {
let pre_call = projected.pre_call.as_ref().ok_or_else(missing_state)?;
self.state.logger()?.update_ocr(
py,
&self.adapter,
&self.state.kwargs,
pre_call,
&projected.fields.secret_fields,
@ -215,7 +217,8 @@ impl PythonRoute for PythonOcrHost {
}
OcrHostOperation::ConstructResponse(response) => {
self.state.end = Some(now(py)?);
self.state.response = Some(callbacks::response(py, response.as_ref())?);
self.state.response =
Some(callbacks::response(py, &self.adapter, response.as_ref())?);
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::MapFailure(error) => {
@ -234,7 +237,7 @@ impl PythonRoute for PythonOcrHost {
),
OcrHostData::Released => return Err(missing_state()),
};
let mapped = callbacks::map_failure(py, error, request, provider)?;
let mapped = callbacks::map_failure(py, &self.adapter, error, request, provider)?;
self.state
.retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any()));
OcrHostResult::Lifecycle(Ok(()))
@ -249,6 +252,7 @@ impl PythonRoute for PythonOcrHost {
self.data = OcrHostData::Released;
}
fn traverse(&self, visit: &pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> {
visit.call(&self.adapter)?;
match &self.data {
OcrHostData::Unprojected { request } => visit.call(request),
OcrHostData::Projected(projected) => {
@ -282,6 +286,7 @@ fn _ocr_lifecycle(
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
host: Bound<'_, PyAny>,
) -> PyResult<Py<PyAny>> {
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
let call = admitted_call(OcrCall::admit(
@ -299,6 +304,7 @@ fn _ocr_lifecycle(
asynchronous,
if asynchronous { "aocr" } else { "ocr" },
)?,
adapter: host.unbind(),
data: OcrHostData::Unprojected {
request: request.unbind(),
},

View file

@ -28,6 +28,11 @@ fn prepare_ocr(
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?
.unwrap_or_default();
crate::errors::admit(litellm_core::ocr::admit_value(
&options.model,
options.custom_llm_provider.as_deref(),
))?;
Ok(async move {
let RouteOptions {
model,

View file

@ -3,7 +3,7 @@ use pyo3::prelude::*;
use pyo3::types::PyAny;
use serde_json::Value;
use crate::errors::core_error_to_pyerr;
use crate::errors::{admit, execution_error_to_pyerr};
use crate::marshal::{marshal_headers, optional_timeout};
#[pyclass]
@ -14,20 +14,24 @@ struct ResponsesWebSocketConnection {
#[pymethods]
impl ResponsesWebSocketConnection {
#[classmethod]
#[pyo3(signature = (url, headers=None, timeout_seconds=None))]
#[pyo3(signature = (url, headers=None, timeout_seconds=None, custom_llm_provider=None))]
fn connect<'py>(
_cls: &Bound<'py, pyo3::types::PyType>,
py: Python<'py>,
url: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] headers: Option<Value>,
timeout_seconds: Option<f64>,
custom_llm_provider: Option<String>,
) -> PyResult<Bound<'py, PyAny>> {
admit(litellm_core::responses::websocket::admit(
custom_llm_provider.as_deref(),
))?;
let headers = marshal_headers(headers)?;
let timeout = optional_timeout(timeout_seconds);
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout)
.await
.map_err(core_error_to_pyerr)?;
.map_err(execution_error_to_pyerr)?;
Ok(ResponsesWebSocketConnection { inner })
})
}
@ -35,21 +39,24 @@ impl ResponsesWebSocketConnection {
fn send_text<'py>(&self, py: Python<'py>, text: String) -> PyResult<Bound<'py, PyAny>> {
let inner = self.inner.clone();
pyo3_async_runtimes::tokio::future_into_py(py, async move {
inner.send_text(text).await.map_err(core_error_to_pyerr)
inner
.send_text(text)
.await
.map_err(execution_error_to_pyerr)
})
}
fn recv_text<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let inner = self.inner.clone();
pyo3_async_runtimes::tokio::future_into_py(py, async move {
inner.recv_text().await.map_err(core_error_to_pyerr)
inner.recv_text().await.map_err(execution_error_to_pyerr)
})
}
fn close<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
let inner = self.inner.clone();
pyo3_async_runtimes::tokio::future_into_py(py, async move {
inner.close().await.map_err(core_error_to_pyerr)
inner.close().await.map_err(execution_error_to_pyerr)
})
}
}

View file

@ -0,0 +1,127 @@
use std::collections::HashMap;
use std::num::NonZero;
use std::sync::{Arc, Mutex, OnceLock};
use std::thread::available_parallelism;
use litellm_python_interop::release_gil;
use litellm_token_counter::{
CountableRequest, Error, InputTokenCount, TokenCounter as CoreTokenCounter,
};
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
use pyo3::types::PyAny;
use tokio::sync::Semaphore;
use crate::constants::TOKEN_COUNT_FALLBACK_PARALLELISM;
use crate::errors::{RustBridgeDeclined, RustBridgeUnavailable};
use crate::execution::run_async;
struct CachedCounter {
counter: Arc<CoreTokenCounter>,
encode_slots: Arc<Semaphore>,
}
static COUNTERS: OnceLock<Mutex<HashMap<&'static str, Arc<CachedCounter>>>> = OnceLock::new();
#[pyfunction]
#[pyo3(signature = (body, kind, encoding, disabled, legacy_accounting, resource_loader))]
fn count_input_tokens<'py>(
py: Python<'py>,
body: &[u8],
kind: Option<&str>,
encoding: &str,
disabled: bool,
legacy_accounting: bool,
resource_loader: Py<PyAny>,
) -> PyResult<Bound<'py, PyAny>> {
let tokenizer = litellm_token_counter::admit_tokenizer(kind, encoding, disabled, legacy_accounting)
.map_err(admission_error_to_pyerr)?;
CoreTokenCounter::admit_request(body).map_err(admission_error_to_pyerr)?;
let cached = cached_counter(py, tokenizer, resource_loader)?;
let body = body.to_vec();
run_async(
py,
async move {
let _slot = Arc::clone(&cached.encode_slots)
.acquire_owned()
.await
.map_err(|error| Error::Task(error.to_string()))?;
tokio::task::spawn_blocking(move || count_body(&cached.counter, &body))
.await
.map_err(|error| Error::Task(error.to_string()))?
},
token_count_error_to_pyerr,
)
}
fn cached_counter(
py: Python<'_>,
tokenizer: &'static str,
resource_loader: Py<PyAny>,
) -> PyResult<Arc<CachedCounter>> {
let counters = COUNTERS.get_or_init(|| Mutex::new(HashMap::new()));
if let Some(counter) = counters
.lock()
.map_err(|error| PyRuntimeError::new_err(error.to_string()))?
.get(tokenizer)
.cloned()
{
return Ok(counter);
}
let resource: String = resource_loader
.call1(py, (tokenizer,))
.and_then(|value| value.extract(py))
.map_err(|error| RustBridgeUnavailable::new_err(error.to_string()))?;
let counter = release_gil(py, move || load_counter(tokenizer, &resource)).map_err(token_count_error_to_pyerr)?;
let cached = Arc::new(CachedCounter {
counter: Arc::new(counter),
encode_slots: Arc::new(Semaphore::new(encode_parallelism())),
});
counters
.lock()
.map_err(|error| PyRuntimeError::new_err(error.to_string()))?
.insert(tokenizer, Arc::clone(&cached));
Ok(cached)
}
fn load_counter(tokenizer: &str, resource: &str) -> Result<CoreTokenCounter, Error> {
match tokenizer {
"anthropic" => CoreTokenCounter::from_json(resource),
"cl100k_base" => CoreTokenCounter::from_cl100k_ranks(resource),
"o200k_base" => CoreTokenCounter::from_o200k_ranks(resource),
_ => Err(Error::UnsupportedTokenizer),
}
}
fn encode_parallelism() -> usize {
available_parallelism().map_or(TOKEN_COUNT_FALLBACK_PARALLELISM, NonZero::get)
}
fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result<InputTokenCount, Error> {
let request = CountableRequest::parse(body)?;
counter.count_request(&request)
}
fn admission_error_to_pyerr(error: Error) -> PyErr {
RustBridgeDeclined::new_err(error.to_string())
}
fn token_count_error_to_pyerr(error: Error) -> PyErr {
let message = error.to_string();
match error {
Error::Load(_) | Error::Ranks(_) | Error::UnicodeClasses => RustBridgeUnavailable::new_err(message),
Error::UnsupportedTokenizer
| Error::RequestParse(_)
| Error::MissingInput
| Error::FloatText
| Error::ContentBlock
| Error::ArrayItems
| Error::JsonSerialization(_)
| Error::JsonUtf8(_) => RustBridgeDeclined::new_err(message),
Error::Encode(_) | Error::Task(_) => PyRuntimeError::new_err(message),
}
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(count_input_tokens, module)?)
}

View file

@ -7,7 +7,7 @@ use litellm_core::audio_transcription::{
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::core_error_to_pyerr;
use crate::errors::{admit, execution_error_to_pyerr};
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
fn prepare_transcription(
@ -24,6 +24,12 @@ fn prepare_transcription(
})?;
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
admit(litellm_core::audio_transcription::admit(
&options.model,
options.custom_llm_provider.as_deref(),
&audio,
))?;
Ok(async move {
let RouteOptions {
model,
@ -67,5 +73,5 @@ bridge_route! {
timeout_seconds: Option<f64>,
},
prepare = prepare_transcription,
errors = core_error_to_pyerr,
errors = execution_error_to_pyerr,
}

View file

@ -1,105 +0,0 @@
use std::num::NonZero;
use std::sync::Arc;
use std::thread::available_parallelism;
use litellm_python_interop::release_gil;
use litellm_token_counter::{
CountableRequest, Error, InputTokenCount, TokenCounter as CoreTokenCounter,
};
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::PyAny;
use tokio::sync::Semaphore;
use crate::constants::TOKEN_COUNT_FALLBACK_PARALLELISM;
use crate::errors::RustBridgeDeclined;
use crate::execution::run_async;
/// Counts the input tokens of a raw request body off the Python event loop with
/// the GIL released. Python owns which requests get here and what to do with
/// the count. At most one encode per core runs at a time; the rest wait in the
/// async task, where a cancelled Python awaiter drops them before any blocking
/// work is scheduled.
#[pyclass(frozen)]
struct TokenCounter {
inner: Arc<CoreTokenCounter>,
encode_slots: Arc<Semaphore>,
}
#[pymethods]
impl TokenCounter {
#[new]
fn new(py: Python<'_>, tokenizer_json: &str) -> PyResult<Self> {
Self::load(py, || CoreTokenCounter::from_json(tokenizer_json))
}
#[staticmethod]
fn from_cl100k_ranks(py: Python<'_>, rank_file: &str) -> PyResult<Self> {
Self::load(py, || CoreTokenCounter::from_cl100k_ranks(rank_file))
}
#[staticmethod]
fn from_o200k_ranks(py: Python<'_>, rank_file: &str) -> PyResult<Self> {
Self::load(py, || CoreTokenCounter::from_o200k_ranks(rank_file))
}
fn acount_request<'py>(&self, py: Python<'py>, body: &[u8]) -> PyResult<Bound<'py, PyAny>> {
let counter = Arc::clone(&self.inner);
let encode_slots = Arc::clone(&self.encode_slots);
let body = body.to_vec();
run_async(
py,
async move {
let _slot = encode_slots
.acquire_owned()
.await
.map_err(|error| Error::Task(error.to_string()))?;
tokio::task::spawn_blocking(move || count_body(&counter, &body))
.await
.map_err(|error| Error::Task(error.to_string()))?
},
token_count_error_to_pyerr,
)
}
}
impl TokenCounter {
fn load(
py: Python<'_>,
load: impl FnOnce() -> Result<CoreTokenCounter, Error> + Send,
) -> PyResult<Self> {
let inner = release_gil(py, load).map_err(token_count_error_to_pyerr)?;
Ok(Self {
inner: Arc::new(inner),
encode_slots: Arc::new(Semaphore::new(encode_parallelism())),
})
}
}
fn encode_parallelism() -> usize {
available_parallelism().map_or(TOKEN_COUNT_FALLBACK_PARALLELISM, NonZero::get)
}
fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result<InputTokenCount, Error> {
let request = CountableRequest::parse(body)?;
counter.count_request(&request)
}
fn token_count_error_to_pyerr(error: Error) -> PyErr {
let message = error.to_string();
match error {
Error::Load(_) | Error::Ranks(_) | Error::UnicodeClasses => PyValueError::new_err(message),
Error::RequestParse(_)
| Error::MissingInput
| Error::FloatText
| Error::ContentBlock
| Error::ArrayItems
| Error::JsonSerialization(_)
| Error::JsonUtf8(_) => RustBridgeDeclined::new_err(message),
Error::Encode(_) | Error::Task(_) => PyRuntimeError::new_err(message),
}
}
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_class::<TokenCounter>()
}

View file

@ -25,6 +25,7 @@ pub struct InputTokenCount {
}
enum Encoder {
Admission,
HuggingFace {
tokenizer: Box<tokenizers::Tokenizer>,
byte_level: Option<ByteLevelCounter>,
@ -40,6 +41,15 @@ pub struct TokenCounter {
}
impl TokenCounter {
pub fn admit_request(body: &[u8]) -> Result<(), Error> {
let request = CountableRequest::parse(body)?;
Self {
encoder: Encoder::Admission,
}
.count_request(&request)
.map(|_| ())
}
/// Load a HuggingFace `tokenizer.json` document. The host reads the file.
pub fn from_json(tokenizer_json: &str) -> Result<Self, Error> {
let tokenizer = tokenizer_json
@ -74,6 +84,7 @@ impl TokenCounter {
pub fn count_text(&self, text: &str) -> Result<usize, Error> {
match &self.encoder {
Encoder::Admission => Ok(0),
Encoder::Tiktoken(counter) => Ok(counter.count(text)),
Encoder::HuggingFace {
tokenizer,

View file

@ -4,6 +4,8 @@ use thiserror::Error as ThisError;
#[derive(Debug, ThisError)]
pub enum Error {
#[error("tokenizer or accounting configuration is not supported")]
UnsupportedTokenizer,
#[error("failed to load tokenizer: {0}")]
Load(#[source] tokenizers::Error),
#[error("failed to load tokenizer: tiktoken rank file: {0}")]

View file

@ -19,3 +19,20 @@ mod unicode_classes;
pub use counter::{InputTokenCount, TokenCounter};
pub use error::Error;
pub use types::CountableRequest;
pub fn admit_tokenizer(
kind: Option<&str>,
encoding: &str,
disabled: bool,
legacy_accounting: bool,
) -> Result<&'static str, Error> {
if disabled {
return Err(Error::UnsupportedTokenizer);
}
match (kind, encoding, legacy_accounting) {
(Some("anthropic"), _, _) => Ok("anthropic"),
(None, "cl100k_base", false) => Ok("cl100k_base"),
(None, "o200k_base", false) => Ok("o200k_base"),
_ => Err(Error::UnsupportedTokenizer),
}
}

View file

@ -911,7 +911,7 @@ openai_compatible_endpoints: Final[list] = [
]
openai_compatible_providers: Final[list] = [
openai_compatible_providers: Final[list[str]] = [
"anyscale",
"groq",
"nvidia_nim",

View file

@ -26,7 +26,6 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
)
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.types.llms.anthropic import (
ContentBlockDelta,
ContentBlockStart,
@ -372,27 +371,21 @@ class AnthropicChatCompletion(BaseLLM):
transform_params: Final = {**optional_params, "is_vertex_request": is_vertex_request}
def finish_request(request_data: dict) -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream
"""Filter beta headers and emit pre_call, returning `(headers, data)`.
The pair stays mutable because the streaming path rewrites it in
place (`data["stream"] = True`) before sending. A Rust attempt that
declined already emitted pre_call for this request, so skip it there.
"""
"""Filter beta headers and emit pre_call, returning `(headers, data)`."""
request_headers, data = update_request_with_filtered_beta(
headers=headers,
request_data=request_data,
provider=custom_llm_provider,
)
if not serves_via_rust:
logging_obj.pre_call(
input=messages,
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": request_headers,
},
)
logging_obj.pre_call(
input=messages,
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": request_headers,
},
)
print_verbose(f"_is_function_call: {_is_function_call}")
return request_headers, data
@ -456,54 +449,30 @@ class AnthropicChatCompletion(BaseLLM):
timeout=timeout,
)
# The Rust core owns the whole call for the subset it accepts, so ask
# before transforming: whichever path runs emits pre_call exactly once.
# `get_config` merges the class-level defaults (Anthropic's required
# `max_tokens` among them) that `transform_request` would have applied.
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
**AnthropicConfig.get_config(model=model),
**optional_params,
}
serves_via_rust: Final = rust_chat_completions_accepts(
model=model,
messages=messages,
optional_params=rust_optional_params,
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
stream=stream,
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
"model": model,
"messages": messages,
**rust_optional_params,
},
"api_base": api_base,
"headers": headers,
}
log_rust_pre_call: Final = lambda: logging_obj.pre_call(
input=messages, api_key=api_key, additional_args=rust_logging_args
)
if serves_via_rust:
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
"model": model,
"messages": messages,
**rust_optional_params,
},
"api_base": api_base,
"headers": headers,
}
logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args)
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key=api_key,
additional_args=rust_logging_args,
)
if acompletion is True:
return rust_chat_completions_bridge.achat_completions_or_fallback(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
python_fallback=acompletion_dispatch,
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key=api_key,
additional_args=rust_logging_args,
)
if acompletion is True:
return rust_chat_completions_bridge.achat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
@ -513,10 +482,30 @@ class AnthropicChatCompletion(BaseLLM):
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
stream=stream,
litellm_params=litellm_params,
on_request=log_rust_pre_call,
on_response=log_rust_post_call,
python_fallback=acompletion_dispatch,
)
if rust_response is not None:
return rust_response
rust_response: Final = rust_chat_completions_bridge.chat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
stream=stream,
litellm_params=litellm_params,
on_request=log_rust_pre_call,
on_response=log_rust_post_call,
python_fallback=lambda: None,
)
if rust_response is not None:
return rust_response
if acompletion is True:
return acompletion_dispatch()

View file

@ -23,8 +23,6 @@ class BedrockAudioTranscriptionRustDispatch:
audio_format: Final = formats.get(processed_audio.content_type) or (
processed_audio.filename.rsplit(".", 1)[-1].lower() if "." in processed_audio.filename else ""
)
if audio_format not in {"wav", "mp3", "flac", "ogg"}:
raise ValueError(f"Unsupported Bedrock audio format for file {processed_audio.filename!r}")
return {
"data": base64.b64encode(processed_audio.file_content).decode("ascii"),
"format": audio_format,

View file

@ -17,7 +17,6 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
)
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
@ -196,7 +195,6 @@ class BedrockConverseLLM(BaseAWSLLM):
headers: dict = {},
client: AsyncHTTPHandler | None = None,
api_key: str | None = None,
skip_pre_call_logging: bool = False,
) -> ModelResponse | CustomStreamWrapper:
request_data: Final = await litellm.AmazonConverseConfig()._async_transform_request(
model=model,
@ -222,16 +220,15 @@ class BedrockConverseLLM(BaseAWSLLM):
# The Rust path already logged this request's pre_call before handing
# it here, and it only declines before the provider is called, so this
# is the same attempt continuing rather than a second one.
if not skip_pre_call_logging:
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": prepped.headers,
},
)
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": prepped.headers,
},
)
headers = dict(prepped.headers)
if client is None or not isinstance(client, AsyncHTTPHandler):
@ -401,43 +398,30 @@ class BedrockConverseLLM(BaseAWSLLM):
# Filter beta headers in HTTP headers before making the request
headers = update_headers_with_filtered_beta(headers=headers, provider="bedrock_converse")
# The Rust core owns the whole call for the subset it accepts. Ask
# before transforming so whichever path runs emits pre_call once, and
# hand down the credentials, region and endpoint this handler already
# resolved so both paths sign as the same principal. Bearer-token auth
# resolves no SigV4 principal at all, and each path reads that token
# itself.
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
**optional_params,
**_sigv4_principal(credentials),
"aws_region_name": aws_region_name,
}
serves_via_rust: Final = rust_chat_completions_accepts(
model=model,
messages=messages,
optional_params=rust_optional_params,
custom_llm_provider="bedrock",
litellm_params=litellm_params,
stream=stream,
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
"messages": messages,
**optional_params,
},
"api_base": proxy_endpoint_url,
"headers": headers,
}
log_rust_pre_call: Final = lambda: logging_obj.pre_call(
input=messages, api_key="", additional_args=rust_logging_args
)
if serves_via_rust:
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
"messages": messages,
**optional_params,
},
"api_base": proxy_endpoint_url,
"headers": headers,
}
logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args)
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key="",
additional_args=rust_logging_args,
)
if acompletion:
return rust_chat_completions_bridge.achat_completions_or_fallback(
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key="",
additional_args=rust_logging_args,
)
if acompletion:
return rust_chat_completions_bridge.achat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
@ -447,6 +431,9 @@ class BedrockConverseLLM(BaseAWSLLM):
custom_llm_provider="bedrock",
extra_headers=headers,
timeout=timeout,
stream=stream,
litellm_params=litellm_params,
on_request=log_rust_pre_call,
on_response=log_rust_post_call,
python_fallback=lambda: self.async_completion(
model=model,
@ -464,23 +451,26 @@ class BedrockConverseLLM(BaseAWSLLM):
client=client,
credentials=credentials,
api_key=api_key,
skip_pre_call_logging=True,
),
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
api_base=proxy_endpoint_url,
custom_llm_provider="bedrock",
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
)
if rust_response is not None:
return rust_response
rust_response: Final = rust_chat_completions_bridge.chat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
api_base=proxy_endpoint_url,
custom_llm_provider="bedrock",
extra_headers=headers,
timeout=timeout,
stream=stream,
litellm_params=litellm_params,
on_request=log_rust_pre_call,
on_response=log_rust_post_call,
python_fallback=lambda: None,
)
if rust_response is not None:
return rust_response
### ROUTING (ASYNC, STREAMING, SYNC)
if acompletion:
@ -548,21 +538,15 @@ class BedrockConverseLLM(BaseAWSLLM):
)
## LOGGING
# Reaching here with `serves_via_rust` set means the synchronous Rust
# attempt declined at call time, before the provider was called, and
# already logged this request. That is the same attempt continuing.
# The asynchronous branch above returns before this point, and hands
# its own fallback `skip_pre_call_logging=True` for the same reason.
if not serves_via_rust:
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
if client is None or isinstance(client, AsyncHTTPHandler):
_params: Final = {}
if timeout is not None:

View file

@ -162,15 +162,6 @@ from litellm.utils import (
async_pre_call_deployment_hook,
)
def _rust_responses_websocket_enabled(
custom_llm_provider: str | None,
) -> bool:
from litellm.rust_bridge.configuration import RouteName, rust_enabled
return custom_llm_provider == "openai" and rust_enabled(RouteName.RESPONSES)
from .http_handler import get_shared_realtime_ssl_context
if TYPE_CHECKING:
@ -2454,34 +2445,19 @@ class BaseLLMHTTPHandler:
request_body: dict,
timeout: float | httpx.Timeout | None,
) -> AnthropicMessagesResponse | None:
if custom_llm_provider not in ("azure_ai", "anthropic"):
return None
from litellm.rust_bridge.configuration import RouteName, rust_enabled
if not rust_enabled(RouteName.MESSAGES):
return None
if has_agentic_hook:
return None
from litellm.rust_bridge import messages as rust_messages_bridge
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
try:
rust_response: Final = await rust_messages_bridge.amessages(
model=model,
body=upstream_body,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
)
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
verbose_logger.debug(
"Rust Anthropic messages bridge raised %s; falling back to Python path",
type(rust_error).__name__,
)
return None
rust_response: Final = await rust_messages_bridge.amessages(
model=model,
body=upstream_body,
has_agentic_hook=has_agentic_hook,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
)
if rust_response is None:
return None
@ -6657,17 +6633,21 @@ class BaseLLMHTTPHandler:
@asynccontextmanager
async def _backend_connection():
if _rust_responses_websocket_enabled(custom_llm_provider):
from litellm.rust_bridge.responses import websocket as rust_responses_websocket
from litellm.rust_bridge.responses import websocket as rust_responses_websocket
rust_backend: Final = await rust_responses_websocket.connect(
url=ws_url,
headers={str(key): str(value) for key, value in headers.items()},
timeout=timeout,
)
if rust_backend is not None:
rust_backend: Final = await rust_responses_websocket.connect(
url=ws_url,
headers={str(key): str(value) for key, value in headers.items()},
timeout=timeout,
custom_llm_provider=custom_llm_provider,
model=model,
)
if rust_backend is not None:
try:
yield rust_backend
return
finally:
await rust_backend.close()
return
async with websockets.connect(
ws_url,

View file

@ -4,8 +4,7 @@ from typing import Final, Literal, Protocol, cast # noqa: TID251 # native call
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.configuration import rust_ocr_enabled
from litellm.rust_bridge.ocr.definition import COMPONENT
class FileReader(Protocol):
@ -30,7 +29,7 @@ class NativeMimeType(Protocol):
def __call__(self, file_name: str) -> str: ...
_FILE_DOCUMENT: Final = NativeBinding(
_FILE_DOCUMENT: Final = COMPONENT.bind(
"_ocr_file_document",
validate=lambda value: (
cast( # cast-ok: native export owns the callable signature
@ -40,7 +39,7 @@ _FILE_DOCUMENT: Final = NativeBinding(
else None
),
)
_UPLOAD_DOCUMENT: Final = NativeBinding(
_UPLOAD_DOCUMENT: Final = COMPONENT.bind(
"_ocr_upload_document",
validate=lambda value: (
cast( # cast-ok: native export owns the callable signature
@ -50,10 +49,10 @@ _UPLOAD_DOCUMENT: Final = NativeBinding(
else None
),
)
_MAX_FILE_BYTES: Final = NativeBinding(
_MAX_FILE_BYTES: Final = COMPONENT.bind(
"_OCR_MAX_FILE_BYTES", validate=lambda value: value if isinstance(value, int) and value > 0 else None
)
_MIME_TYPE: Final = NativeBinding(
_MIME_TYPE: Final = COMPONENT.bind(
"_ocr_mime_type",
validate=lambda value: (
cast( # cast-ok: native export owns the callable signature
@ -67,7 +66,7 @@ _PYTHON_MAX_FILE_BYTES: Final = 50 * 1024 * 1024
def get_mime_type(file_path: str) -> str:
native: Final = _MIME_TYPE.load() if rust_ocr_enabled() else None
native: Final = COMPONENT.resolve().select(_MIME_TYPE)
if native is None:
from litellm.ocr import legacy
@ -76,14 +75,14 @@ def get_mime_type(file_path: str) -> str:
def get_max_file_bytes() -> int:
limit: Final = _MAX_FILE_BYTES.load() if rust_ocr_enabled() else None
limit: Final = COMPONENT.resolve().select(_MAX_FILE_BYTES)
if limit is None:
return _PYTHON_MAX_FILE_BYTES
return limit
def convert_file_document_to_url_document(document: FileDocument) -> dict[str, str]:
native: Final = _FILE_DOCUMENT.load() if rust_ocr_enabled() else None
native: Final = COMPONENT.resolve().select(_FILE_DOCUMENT)
if native is None:
from litellm.ocr import legacy
@ -94,7 +93,7 @@ def convert_file_document_to_url_document(document: FileDocument) -> dict[str, s
def convert_upload_to_url_document(
file_content: bytes, filename: str | None, content_type: str | None
) -> dict[str, str]:
native: Final = _UPLOAD_DOCUMENT.load() if rust_ocr_enabled() else None
native: Final = COMPONENT.resolve().select(_UPLOAD_DOCUMENT)
if native is None:
from litellm.ocr import legacy

View file

@ -6,10 +6,11 @@ import httpx
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import legacy
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
from litellm.rust_bridge.bindings import native_exception_types
from litellm.rust_bridge.configuration import rust_ocr_enabled
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
from litellm.rust_bridge.ocr.definition import COMPONENT
from litellm.rust_bridge.ocr.host import HOST
from litellm.rust_bridge.ocr.lifecycle import select
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
__all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr")
@ -48,32 +49,48 @@ def ocr(
**kwargs: object, # kwargs-ok: preserve the public OCR call shape
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
request: Final = _public_request("ocr", args, kwargs)
native: Final = select(request) if rust_ocr_enabled() else None
if native is not None:
try:
return native(request, args, kwargs, False)
except _decline_types():
pass
execution: Final = COMPONENT.resolve()
native: Final = select(request, execution)
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]], legacy.ocr
)
return fallback(*args, **kwargs)
native_call: Final[Callable[[], OCRResponse] | None] = (
(lambda: native(request, args, kwargs, False, HOST)) if native is not None else None
)
return invoke(
execution=execution,
native_call=native_call,
python_fallback=lambda: fallback(*args, **kwargs),
adapt=lambda value: value,
context=BridgeErrorContext(
route=COMPONENT.name.value,
provider=request.custom_llm_provider or "",
model=request.model,
),
)
async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: preserve the public OCR call shape
request: Final = _public_request("aocr", args, kwargs)
native: Final = select(request) if rust_ocr_enabled() else None
if native is not None:
try:
return await native(request, args, kwargs, True)
except _decline_types():
pass
execution: Final = COMPONENT.resolve()
native: Final = select(request, execution)
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
Callable[..., Awaitable[OCRResponse]], legacy.aocr
)
return await fallback(*args, **kwargs)
native_call: Final[Callable[[], Awaitable[OCRResponse]] | None] = (
(lambda: native(request, args, kwargs, True, HOST)) if native is not None else None
)
async def python_fallback() -> OCRResponse:
return await fallback(*args, **kwargs)
def _decline_types() -> tuple[type[BaseException], ...]:
exception_types: Final = native_exception_types()
return (exception_types[0],) if exception_types is not None else ()
return await ainvoke(
execution=execution,
native_call=native_call,
python_fallback=python_fallback,
adapt=lambda value: value,
context=BridgeErrorContext(
route=COMPONENT.name.value,
provider=request.custom_llm_provider or "",
model=request.model,
),
)

View file

@ -1,25 +1,35 @@
# Native route foundation
# Native bridge catalog
Each route has a Python package and a matching Rust module under `crates/python-bridge/src/routes/`. Python `__init__.py` files are thin entrypoints exporting `ROUTE: NativeRoute` and public adapters. Request/response protocols live in `types.py`, value adapters in `value.py`, callback adapters in `callbacks.py` where needed, and full-call bindings in `lifecycle.py`. Unimplemented routes add these files when they gain an implementation
Every SDK API has one `NativeComponent` in the immutable `COMPONENTS` catalog. A component declares its native exports and one `CapabilitySpec`, which resolves implementation availability and rollout from `CapabilityContext(provider, model, delivery)`
WebSocket is a Responses transport: its adapters live in `responses/websocket.py` and Rust `routes/responses/websocket.rs`, using the Responses route policy. Token counting is a utility outside the route registry, in `token_counter.py` and Rust `src/token_counter.rs`
`DeliveryMode` contains `COMPLETED`, `STREAMING`, and `WEBSOCKET`. Lifecycle is an implementation detail, so lifecycle and value entrypoints for the same API share the same completed-delivery policy
`configuration.py` owns release policy. OCR is default-on, Messages and other optional routes are default-off, and transcription is required-native. The process override takes precedence over the environment except for OCR's existing environment opt-out. Required-native execution ignores optional rollout switches. A default is an enablement choice, not a claim that a lifecycle implementation exists
`RustImplementationState` records whether Rust is unimplemented, experimental, or ready. `RolloutPolicy` independently selects unsupported, Python-only, Rust opt-in, Rust opt-out, or Rust-required execution. Optional Rust execution can fall back to Python. Rust-required execution cannot
`NativeRoute.select(binding)` checks policy before discovering the native module. `NativeBinding` handles validation and resettable overrides, with injectable discovery for tests. Native exports keep their existing names; moving a Python module into a package does not change its import path
OCR completed delivery is ready and default-on. Messages, chat completions, token counting, and Responses WebSocket transport are experimental and opt-in. Other completed APIs remain Python-only. Bedrock transcription requires Rust because it has no Python implementation; Python-backed transcription providers remain on Python
## Lifecycle contract
```python
execution = COMPONENT.resolve(
CapabilityContext(
provider=provider,
model=model,
delivery=DeliveryMode.COMPLETED,
)
)
```
Full-call bindings implement `NativeLifecycle[Request, Response]`: `(request, args, kwargs, asynchronous)`. The synchronous form returns a response, while the asynchronous form returns an inline-driven coroutine. Original positional arguments, keyword arguments and Python object identities stay available to the host
Optional capabilities use the `litellm.rust(bool)` process override first, `LITELLM_RUST=1` or `LITELLM_RUST=0` second, then their catalog default. Python-only, Rust-required, and unsupported capabilities ignore overrides
Core owns effect-free admission and callback sequencing. The PyO3 `PythonRoute` implementation retains Python objects, projects consumed fields and executes core-selected hooks. The shared native handle and Python `lifecycle.py` driver preserve caller task/context, error identity, cancellation and cleanup. Python logging continues to select registered integrations and their dispatch modes
## Fallback contract
Only disabled/unavailable native execution or a typed pre-effect admission decline permits Python fallback. Callback failures, projection errors and post-admission failures must not replay the request. Keep success/failure dispatch after fallible response finalization
Each API calls its native entrypoint at most once. Rust performs request admission inside that entrypoint before provider calls or host callbacks. An unavailable binding or `RustBridgeDeclined` selects the supplied Python fallback only when the policy allows it
OCR implements this contract today. Messages, chat completions and transcription retain their existing value-based execution while their new full-call lifecycle slots are unfinished. Embeddings, rerank, image generation/edit, speech, moderation and Responses have lifecycle slots but no public SDK wiring here. The Rust `unimplemented_lifecycle_route!` macro registers each slot and maps a pure core decline to `RustBridgeDeclined`, without inspecting the request. Deliberately avoid `todo!()` in Python-callable paths because it panics instead of providing safe admission fallback
Provider failures, host callback failures, cancellation, conversion failures, and response adaptation failures propagate without replay. Adaptation runs outside the decline-catching boundary
## Extending a route
`invoke` and `ainvoke` return the native result or execute the supplied fallback directly. There is no public admission, prepare, accepts, or can-handle API
Replace the route's Rust lifecycle stub with a typed core call and a `PythonRoute` host, following OCR's `project`, `callbacks` and `lifecycle` split. Give its Python binding concrete request/response types, wire the public entrypoint through admission-only fallback, and prove positive native execution and callback parity before changing its release default
Token counting follows the same component policy. Its one native counting entrypoint validates the tokenizer configuration and request body, obtains and caches the required tokenizer resource, then counts. Unsupported inputs decline, known resource loading failures report native unavailability, and unexpected counting failures propagate
Streaming and WebSocket sessions do not yet use the full-call lifecycle contract. Their follow-up needs explicit chunk delivery, backpressure, final response aggregation, consumer close, cancellation acknowledgement, deferred terminal dispatch and exactly-once cleanup. Returning an iterator or opening a socket is not terminal success. WebSocket uses the Responses policy; token counting keeps the global optional switch. Both use shared loading and retain their own session/utility protocols
## Package layout
Python component packages keep their descriptor in `definition.py`, dynamic call protocols in `types.py`, and entrypoint adapters in `value.py`, `lifecycle.py`, or transport modules. Rust mirrors those APIs below `crates/python-bridge/src/routes/`

View file

@ -1,6 +1,6 @@
from asyncio import Future
from collections.abc import Coroutine
from typing import Literal, overload
from collections.abc import Callable, Coroutine
from typing import Literal, final, overload
from typing_extensions import Never
@ -8,45 +8,55 @@ from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
class RustBridgeDeclined(Exception): ...
class RustBridgeUnavailable(Exception): ...
class RustHostCallbackError(Exception): ...
class RustUpstreamError(Exception): ...
@overload
def _ocr_lifecycle(
request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: Literal[False]
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: dict[str, object],
asynchronous: Literal[False],
host: object,
) -> OCRResponse: ...
@overload
def _ocr_lifecycle(
request: LiteLLMOcrRequest, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: Literal[True]
request: LiteLLMOcrRequest,
args: tuple[object, ...],
kwargs: dict[str, object],
asynchronous: Literal[True],
host: object,
) -> Coroutine[object, object, OCRResponse]: ...
def _messages_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _chat_completions_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _transcription_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _embeddings_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _rerank_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _image_generation_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _image_edit_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _speech_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _moderation_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def _responses_lifecycle(
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool
request: object, args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, host: object
) -> Never: ...
def ocr(
model: str,
@ -98,6 +108,7 @@ def messages(
custom_llm_provider: str | None = None,
extra_headers: object = None,
timeout_seconds: float | None = None,
has_agentic_hook: bool | None = None,
) -> dict[str, object]: ...
def amessages(
model: str,
@ -107,6 +118,7 @@ def amessages(
custom_llm_provider: str | None = None,
extra_headers: object = None,
timeout_seconds: float | None = None,
has_agentic_hook: bool | None = None,
) -> Future[dict[str, object]]: ...
def chat_completions(
model: str,
@ -117,6 +129,8 @@ def chat_completions(
custom_llm_provider: str | None = None,
extra_headers: object = None,
timeout_seconds: float | None = None,
host_facts: object = None,
on_request: Callable[[], None] | None = None,
) -> dict[str, object]: ...
def achat_completions(
model: str,
@ -127,10 +141,9 @@ def achat_completions(
custom_llm_provider: str | None = None,
extra_headers: object = None,
timeout_seconds: float | None = None,
host_facts: object = None,
on_request: Callable[[], None] | None = None,
) -> Future[dict[str, object]]: ...
def chat_completions_decline(
model: str, messages: object, optional_params: object = None, custom_llm_provider: str | None = None
) -> str | None: ...
_OCR_MAX_FILE_BYTES: int
@ -140,21 +153,60 @@ def _ocr_upload_document(
file_content: bytes, file_name: str | None = None, content_type: str | None = None
) -> dict[str, object]: ...
@final
class ResponsesWebSocketConnection:
@classmethod
def connect(
cls, url: str, headers: object = None, timeout_seconds: float | None = None
cls,
url: str,
headers: object = None,
timeout_seconds: float | None = None,
custom_llm_provider: str | None = None,
) -> Future[ResponsesWebSocketConnection]: ...
def send_text(self, text: str) -> Future[None]: ...
def recv_text(self) -> Future[str | None]: ...
def close(self) -> Future[None]: ...
class TokenCounter:
def __init__(self, tokenizer_json: str) -> None: ...
@staticmethod
def from_cl100k_ranks(rank_file: str) -> TokenCounter: ...
@staticmethod
def from_o200k_ranks(rank_file: str) -> TokenCounter: ...
def acount_request(self, body: bytes) -> Future[dict[str, object]]: ...
def count_input_tokens(
body: bytes,
kind: str | None,
encoding: str,
disabled: bool,
legacy_accounting: bool,
resource_loader: Callable[[str], str],
) -> Future[dict[str, object]]: ...
def gil_stats() -> dict[str, int]: ...
__all__ = [
"RustBridgeUnavailable",
"RustBridgeDeclined",
"RustHostCallbackError",
"RustUpstreamError",
"ocr",
"aocr",
"_OCR_MAX_FILE_BYTES",
"_ocr_upload_document",
"_ocr_file_document",
"_ocr_mime_type",
"_ocr_lifecycle",
"_transcription_lifecycle",
"transcription",
"atranscription",
"_messages_lifecycle",
"messages",
"amessages",
"_chat_completions_lifecycle",
"chat_completions",
"achat_completions",
"_embeddings_lifecycle",
"_image_edit_lifecycle",
"_image_generation_lifecycle",
"_moderation_lifecycle",
"_rerank_lifecycle",
"ResponsesWebSocketConnection",
"_responses_lifecycle",
"_speech_lifecycle",
"count_input_tokens",
"gil_stats",
]

View file

@ -63,3 +63,15 @@ def native_exception_types() -> tuple[type[BaseException], type[BaseException]]
if not isinstance(declined, type) or not isinstance(upstream, type):
return None
return declined, upstream
def native_unavailable_exception() -> tuple[type[BaseException], ...]:
native: Final = get_native_bridge()
unavailable: Final = getattr(native, "RustBridgeUnavailable", None)
return (unavailable,) if isinstance(unavailable, type) and issubclass(unavailable, BaseException) else ()
def native_host_callback_exception() -> tuple[type[BaseException], ...]:
native: Final = get_native_bridge()
callback: Final = getattr(native, "RustHostCallbackError", None)
return (callback,) if isinstance(callback, type) and issubclass(callback, BaseException) else ()

View file

@ -0,0 +1,145 @@
from types import MappingProxyType
from typing import Final
from litellm.rust_bridge.configuration import (
CapabilityContext,
CapabilityDefinition,
CapabilitySpec,
DeliveryMode,
RolloutPolicy,
RouteName,
RustImplementationState,
UtilityName,
)
from litellm.rust_bridge.route import NativeComponent
def _unimplemented(*, python_available: bool = True) -> CapabilityDefinition:
return CapabilityDefinition(
rust=RustImplementationState.UNIMPLEMENTED,
python_available=python_available,
rollout=RolloutPolicy.PYTHON_ONLY if python_available else RolloutPolicy.UNSUPPORTED,
)
def _experimental() -> CapabilityDefinition:
return CapabilityDefinition(
rust=RustImplementationState.EXPERIMENTAL,
python_available=True,
rollout=RolloutPolicy.RUST_OPT_IN,
)
def _ready_default() -> CapabilityDefinition:
return CapabilityDefinition(
rust=RustImplementationState.READY,
python_available=True,
rollout=RolloutPolicy.RUST_OPT_OUT,
)
def _completed_only(context: CapabilityContext, completed: CapabilityDefinition) -> CapabilityDefinition:
return completed if context.delivery is DeliveryMode.COMPLETED else _unimplemented()
def _ocr_capability(context: CapabilityContext) -> CapabilityDefinition:
return _completed_only(context, _ready_default())
def _experimental_completed(context: CapabilityContext) -> CapabilityDefinition:
return _completed_only(context, _experimental())
def _python_completed(context: CapabilityContext) -> CapabilityDefinition:
return _completed_only(context, _unimplemented())
def _responses_capability(context: CapabilityContext) -> CapabilityDefinition:
return _experimental() if context.delivery is DeliveryMode.WEBSOCKET else _unimplemented()
def _transcription_capability(context: CapabilityContext) -> CapabilityDefinition:
if context.delivery is not DeliveryMode.COMPLETED:
return _unimplemented()
if context.provider == "bedrock":
return CapabilityDefinition(
rust=RustImplementationState.EXPERIMENTAL,
python_available=False,
rollout=RolloutPolicy.RUST_REQUIRED,
)
import litellm
from litellm.constants import AZURE_OPENAI_AUDIO_PROVIDERS
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
provider: Final = next((provider for provider in LlmProviders if provider.value == context.provider), None)
python_available: Final = provider is not None and (
context.provider in AZURE_OPENAI_AUDIO_PROVIDERS
or context.provider in litellm.openai_compatible_providers
or ProviderConfigManager.get_provider_audio_transcription_config(model=context.model, provider=provider)
is not None
)
return _unimplemented(python_available=python_available)
def _component(
name: RouteName | UtilityName,
capability: CapabilitySpec,
exports: tuple[str, ...],
) -> NativeComponent:
return NativeComponent(name=name, capability=capability, exports=exports)
COMPONENTS: Final = MappingProxyType(
{
RouteName.OCR: _component(
RouteName.OCR,
_ocr_capability,
(
"ocr",
"aocr",
"_ocr_file_document",
"_ocr_upload_document",
"_OCR_MAX_FILE_BYTES",
"_ocr_mime_type",
"_ocr_lifecycle",
),
),
RouteName.MESSAGES: _component(
RouteName.MESSAGES,
_experimental_completed,
("messages", "amessages", "_messages_lifecycle"),
),
RouteName.CHAT_COMPLETIONS: _component(
RouteName.CHAT_COMPLETIONS,
_experimental_completed,
("chat_completions", "achat_completions", "_chat_completions_lifecycle"),
),
RouteName.TRANSCRIPTION: _component(
RouteName.TRANSCRIPTION,
_transcription_capability,
("transcription", "atranscription", "_transcription_lifecycle"),
),
RouteName.EMBEDDINGS: _component(RouteName.EMBEDDINGS, _python_completed, ("_embeddings_lifecycle",)),
RouteName.RERANK: _component(RouteName.RERANK, _python_completed, ("_rerank_lifecycle",)),
RouteName.IMAGE_GENERATION: _component(
RouteName.IMAGE_GENERATION, _python_completed, ("_image_generation_lifecycle",)
),
RouteName.IMAGE_EDIT: _component(RouteName.IMAGE_EDIT, _python_completed, ("_image_edit_lifecycle",)),
RouteName.SPEECH: _component(RouteName.SPEECH, _python_completed, ("_speech_lifecycle",)),
RouteName.MODERATION: _component(RouteName.MODERATION, _python_completed, ("_moderation_lifecycle",)),
RouteName.RESPONSES: _component(
RouteName.RESPONSES,
_responses_capability,
("ResponsesWebSocketConnection", "_responses_lifecycle"),
),
UtilityName.TOKEN_COUNTER: _component(
UtilityName.TOKEN_COUNTER,
_experimental_completed,
("count_input_tokens",),
),
}
)
NATIVE_EXPORTS: Final = frozenset(export for component in COMPONENTS.values() for export in component.exports)

View file

@ -1,39 +1,31 @@
from typing import Final
from litellm.rust_bridge.chat_completions.callbacks import response_logger
from litellm.rust_bridge.chat_completions.definition import COMPONENT
from litellm.rust_bridge.chat_completions.types import (
ResponseObserver,
RustAchatCompletions,
RustChatCompletions,
RustChatCompletionsDecline,
)
from litellm.rust_bridge.chat_completions.value import (
ROUTE,
RUST_CHAT_COMPLETIONS_PROVIDERS,
RUST_RESPONSE_HEADER,
achat_completions,
achat_completions_or_fallback,
chat_completions,
load_rust_achat_completions,
load_rust_chat_completions,
rust_chat_completions_accepts,
set_rust_chat_completions,
)
__all__: Final = (
"ROUTE",
"RUST_CHAT_COMPLETIONS_PROVIDERS",
"COMPONENT",
"RUST_RESPONSE_HEADER",
"ResponseObserver",
"RustAchatCompletions",
"RustChatCompletions",
"RustChatCompletionsDecline",
"achat_completions",
"achat_completions_or_fallback",
"chat_completions",
"load_rust_achat_completions",
"load_rust_chat_completions",
"response_logger",
"rust_chat_completions_accepts",
"set_rust_chat_completions",
)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.CHAT_COMPLETIONS]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.chat_completions import ROUTE
from litellm.rust_bridge.chat_completions.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from collections.abc import Awaitable, Callable, Mapping, Sequence
from typing import Protocol
@ -15,6 +15,8 @@ class RustChatCompletions(Protocol):
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
host_facts: Mapping[str, bool] | None = None,
on_request: Callable[[], None] | None = None,
) -> Mapping[str, object]:
raise NotImplementedError
@ -30,21 +32,12 @@ class RustAchatCompletions(Protocol):
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
host_facts: Mapping[str, bool] | None = None,
on_request: Callable[[], None] | None = None,
) -> Awaitable[Mapping[str, object]]:
raise NotImplementedError
class RustChatCompletionsDecline(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
custom_llm_provider: str | None,
) -> str | None:
raise NotImplementedError
class ResponseObserver(Protocol):
"""Invoked with the payload the core returned, on success only.

View file

@ -1,55 +1,28 @@
"""Native chat completions bindings.
The Rust core owns the conversation translation, the provider call, and the
response normalization for the subset of `/chat/completions` requests it
accepts. This module only marshals inputs and hands the normalized result to
LiteLLM's existing `ModelResponse` builder.
``None`` means the provider was never called, so the caller is free to serve the
request on the Python path. A failure after the call was issued raises instead:
retrying it there would bill the customer for the same work twice.
"""
from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping, Sequence
from types import MappingProxyType
from typing import Final, cast # noqa: TID251 # native callables are validated at load time
import httpx
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.exceptions import APIError
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
convert_to_model_response_object,
)
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset, native_exception_types
from litellm.rust_bridge.chat_completions.types import (
ResponseObserver,
RustAchatCompletions,
RustChatCompletions,
RustChatCompletionsDecline,
)
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
from litellm.rust_bridge.chat_completions.definition import COMPONENT
from litellm.rust_bridge.chat_completions.types import ResponseObserver, RustAchatCompletions, RustChatCompletions
from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.utils import ModelResponse
# Providers whose `/chat/completions` deployments the Rust core can serve. A
# provider outside this set never reaches the bridge.
RUST_CHAT_COMPLETIONS_PROVIDERS: Final = frozenset({"anthropic", "bedrock"})
# `litellm_params` values are `object`, so validate the one this module reads
# rather than narrowing an unparameterized `Mapping` and typing the result Any.
_LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object])
RUST_RESPONSE_HEADER: Final = "x-litellm-rust"
ROUTE: Final = NativeRoute(RouteName.CHAT_COMPLETIONS)
def _as_chat(value: object) -> RustChatCompletions | None:
return cast(RustChatCompletions, value) if callable(value) else None
@ -58,38 +31,25 @@ def _as_achat(value: object) -> RustAchatCompletions | None:
return cast(RustAchatCompletions, value) if callable(value) else None
def _as_decline(value: object) -> RustChatCompletionsDecline | None:
return cast(RustChatCompletionsDecline, value) if callable(value) else None
_CHAT: Final = ROUTE.bind("chat_completions", validate=_as_chat)
_ACHAT: Final = ROUTE.bind("achat_completions", validate=_as_achat)
_DECLINE: Final = ROUTE.bind("chat_completions_decline", validate=_as_decline)
_CHAT: Final = COMPONENT.bind("chat_completions", validate=_as_chat)
_ACHAT: Final = COMPONENT.bind("achat_completions", validate=_as_achat)
def set_rust_chat_completions(
*,
chat_completions: RustChatCompletions | None | BindingUnset = BINDING_UNSET,
achat_completions: RustAchatCompletions | None | BindingUnset = BINDING_UNSET,
decline: RustChatCompletionsDecline | None | BindingUnset = BINDING_UNSET,
) -> None:
"""Inject the native callables, so tests can supply a double instead of
patching module attributes."""
_CHAT.configure(chat_completions)
_ACHAT.configure(achat_completions)
_DECLINE.configure(decline)
def load_rust_chat_completions() -> RustChatCompletions | None:
return ROUTE.select(_CHAT)
return COMPONENT.resolve().select(_CHAT)
def load_rust_achat_completions() -> RustAchatCompletions | None:
return ROUTE.select(_ACHAT)
def _load_rust_decline() -> RustChatCompletionsDecline | None:
return ROUTE.select(_DECLINE)
return COMPONENT.resolve().select(_ACHAT)
def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool:
@ -101,134 +61,21 @@ def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | N
return entries.get("user_id") is not None
def _litellm_metadata_reaches_the_provider(
custom_llm_provider: str | None, litellm_params: Mapping[str, object] | None
) -> bool:
"""Whether the Python transform would promote proxy-owned attribution into the
provider request, below this gate and inside the function the Rust route replaces.
`AnthropicConfig.transform_request` promotes a valid `metadata["user_id"]`
into the Messages body, so the core never sees the key and would send the
request to Anthropic with the abuse-detection attribution missing.
`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the
Converse body whenever the operator armed `bedrock_request_metadata_fields`.
Owning that field also means evicting a caller-supplied one, which the core
cannot do either, so ownership alone is the condition rather than whether
anything resolved.
Deliberately a superset of Python's condition in both cases: declining a
request Python would not have attributed anyway costs only the Rust path,
while missing one loses the attribution silently.
"""
match custom_llm_provider:
case "anthropic":
return _anthropic_user_id_reaches_the_body(litellm_params)
case "bedrock":
return bedrock_request_metadata_is_owned()
case _:
return False
def rust_chat_completions_accepts(
*,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object],
custom_llm_provider: str | None,
litellm_params: Mapping[str, object] | None,
stream: object,
) -> bool:
"""Whether the Rust path will serve this request.
Asked before the caller commits to either path, so pre-call logging is
emitted exactly once, on whichever path actually runs. The core's own
capability gate answers the second half; it resolves no credentials and
performs no I/O.
"""
if custom_llm_provider not in RUST_CHAT_COMPLETIONS_PROVIDERS:
return False
if stream:
return False
if not ROUTE.enabled():
return False
if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params):
verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path")
return False
decline: Final = _load_rust_decline()
if decline is None:
return False
try:
reason: Final = decline(
model=model,
messages=messages,
optional_params=optional_params,
custom_llm_provider=custom_llm_provider,
)
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
verbose_logger.debug(
"Rust chat completions gate raised %s; staying on the Python path",
type(rust_error).__name__,
)
return False
if reason is not None:
verbose_logger.debug("Rust chat completions declined (%s); using the Python path", reason)
return False
return True
def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
"""`(declined, upstream_failed)` from the native module, or None when absent."""
return native_exception_types()
def _reraise_or_decline(
rust_error: BaseException,
*,
model: str,
custom_llm_provider: str | None,
) -> None:
"""Re-raise a failure the provider already saw, or return so the caller declines.
A request that never reached the provider is safe to serve on the Python
path. One that did is not: the provider has already done the work, so a
second attempt bills for it twice. Those surface as an `APIError` carrying
the upstream status, which LiteLLM's exception mapping already understands.
"""
exceptions: Final = _rust_bridge_exceptions()
if exceptions is None:
verbose_logger.debug(
"Rust chat completions bridge raised %s; falling back to Python path",
type(rust_error).__name__,
)
return
declined, upstream_failed = exceptions
if isinstance(rust_error, upstream_failed):
args: Final = rust_error.args
status: Final = args[0] if args else 0
message: Final = args[1] if len(args) > 1 else ""
raise APIError(
status_code=int(status) or 500,
message=f"litellm rust chat completions: {message}",
llm_provider=custom_llm_provider or "",
model=model,
)
if not isinstance(rust_error, declined):
raise rust_error
verbose_logger.debug(
"Rust chat completions declined before calling the provider (%s); using the Python path",
rust_error,
def _host_facts(stream: object, litellm_params: Mapping[str, object] | None) -> Mapping[str, bool]:
return MappingProxyType(
{
"stream": bool(stream),
"anthropic_user_id": _anthropic_user_id_reaches_the_body(litellm_params),
"bedrock_metadata_owned": bedrock_request_metadata_is_owned(),
}
)
def _build_model_response(
rust_response: Mapping[str, object],
model_response: ModelResponse,
) -> ModelResponse:
def _build_model_response(rust_response: Mapping[str, object], model_response: ModelResponse) -> ModelResponse:
built: Final = convert_to_model_response_object(
response_object=dict(rust_response), # mutable-ok: the converter takes a real dict and rewrites it
model_response_object=model_response,
hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: rewritten by the converter
hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: converter rewrites it
)
if not isinstance(built, ModelResponse):
raise TypeError(f"expected a ModelResponse from the rust path, got {type(built).__name__}")
@ -246,13 +93,27 @@ def chat_completions(
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
) -> ModelResponse | None:
rust_chat_completions: Final = load_rust_chat_completions()
if rust_chat_completions is None:
return None
try:
rust_response: Final = rust_chat_completions(
python_fallback: Callable[[], object],
stream: object = False,
litellm_params: Mapping[str, object] | None = None,
on_request: Callable[[], None] = lambda: None,
on_response: ResponseObserver = lambda _response: None,
) -> object:
execution: Final = COMPONENT.resolve(
CapabilityContext(
provider=custom_llm_provider or "",
model=model,
delivery=DeliveryMode.STREAMING if bool(stream) else DeliveryMode.COMPLETED,
)
)
rust_chat_completions: Final = execution.select(_CHAT)
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
on_response(rust_response)
return _build_model_response(rust_response, model_response)
native_call: Final[Callable[[], Mapping[str, object]] | None] = (
lambda: rust_chat_completions(
model=model,
messages=messages,
optional_params=optional_params,
@ -261,12 +122,19 @@ def chat_completions(
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
host_facts=_host_facts(stream, litellm_params),
on_request=on_request,
)
except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw
_reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider)
return None
on_response(rust_response)
return _build_model_response(rust_response, model_response)
if rust_chat_completions is not None
else None
)
return invoke(
execution=execution,
native_call=native_call,
python_fallback=python_fallback,
adapt=adapt,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)
async def achat_completions(
@ -280,13 +148,27 @@ async def achat_completions(
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
) -> ModelResponse | None:
rust_achat_completions: Final = load_rust_achat_completions()
if rust_achat_completions is None:
return None
try:
rust_response: Final = await rust_achat_completions(
python_fallback: Callable[[], Awaitable[object]],
stream: object = False,
litellm_params: Mapping[str, object] | None = None,
on_request: Callable[[], None] = lambda: None,
on_response: ResponseObserver = lambda _response: None,
) -> object:
execution: Final = COMPONENT.resolve(
CapabilityContext(
provider=custom_llm_provider or "",
model=model,
delivery=DeliveryMode.STREAMING if bool(stream) else DeliveryMode.COMPLETED,
)
)
rust_achat_completions: Final = execution.select(_ACHAT)
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
on_response(rust_response)
return _build_model_response(rust_response, model_response)
native_call: Final[Callable[[], Awaitable[Mapping[str, object]]] | None] = (
lambda: rust_achat_completions(
model=model,
messages=messages,
optional_params=optional_params,
@ -295,48 +177,16 @@ async def achat_completions(
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
host_facts=_host_facts(stream, litellm_params),
on_request=on_request,
)
except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw
_reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider)
return None
on_response(rust_response)
return _build_model_response(rust_response, model_response)
async def achat_completions_or_fallback(
*,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object],
model_response: ModelResponse,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
python_fallback: Callable[[], Awaitable[object]],
) -> object:
"""Await the Rust path, falling back to the caller's own Python path when
the bridge is unavailable or the call fails.
The caller supplies the fallback, so the bridge stays free of provider
dispatch. This exists because a caller that dispatches asynchronously has
already returned a coroutine by the time a Rust failure surfaces, and so
cannot fall back on its own.
"""
response: Final = await achat_completions(
model=model,
messages=messages,
optional_params=optional_params,
model_response=model_response,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout=timeout,
on_response=on_response,
if rust_achat_completions is not None
else None
)
return await ainvoke(
execution=execution,
native_call=native_call,
python_fallback=python_fallback,
adapt=adapt,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)
if response is not None:
return response
return await python_fallback()

View file

@ -3,14 +3,10 @@ from __future__ import annotations
import os
from dataclasses import dataclass
from enum import Enum
from types import MappingProxyType
from typing import Final
from pydantic import TypeAdapter
from typing import Final, Protocol, TypeAlias, assert_never
DEFAULT_RUST_ENABLED: Final = False
_GLOBAL_ENV_NAME: Final = "LITELLM_RUST"
_ENV_BOOL: Final = TypeAdapter(bool)
class RouteName(str, Enum):
@ -27,33 +23,77 @@ class RouteName(str, Enum):
RESPONSES = "responses"
class RouteMode(str, Enum):
OPTIONAL = "optional"
REQUIRED = "required"
class UtilityName(str, Enum):
TOKEN_COUNTER = "token_counter"
class RustImplementationState(str, Enum):
UNIMPLEMENTED = "unimplemented"
EXPERIMENTAL = "experimental"
READY = "ready"
class RolloutPolicy(str, Enum):
UNSUPPORTED = "unsupported"
PYTHON_ONLY = "python_only"
RUST_OPT_IN = "rust_opt_in"
RUST_OPT_OUT = "rust_opt_out"
RUST_REQUIRED = "rust_required"
class ExecutionDecision(str, Enum):
UNSUPPORTED = "unsupported"
PYTHON = "python"
RUST_WITH_FALLBACK = "rust_with_fallback"
RUST_REQUIRED = "rust_required"
class DeliveryMode(str, Enum):
COMPLETED = "completed"
STREAMING = "streaming"
WEBSOCKET = "websocket"
@dataclass(frozen=True, slots=True)
class RoutePolicy:
default_enabled: bool = False
mode: RouteMode = RouteMode.OPTIONAL
environment_opt_out: bool = False
class CapabilityContext:
provider: str = ""
model: str = ""
delivery: DeliveryMode = DeliveryMode.COMPLETED
ROUTE_POLICIES: Final = MappingProxyType(
{
RouteName.OCR: RoutePolicy(default_enabled=True, environment_opt_out=True),
RouteName.MESSAGES: RoutePolicy(),
RouteName.CHAT_COMPLETIONS: RoutePolicy(),
RouteName.TRANSCRIPTION: RoutePolicy(mode=RouteMode.REQUIRED),
RouteName.EMBEDDINGS: RoutePolicy(),
RouteName.RERANK: RoutePolicy(),
RouteName.IMAGE_GENERATION: RoutePolicy(),
RouteName.IMAGE_EDIT: RoutePolicy(),
RouteName.SPEECH: RoutePolicy(),
RouteName.MODERATION: RoutePolicy(),
RouteName.RESPONSES: RoutePolicy(),
}
)
@dataclass(frozen=True, slots=True)
class CapabilityDefinition:
rust: RustImplementationState
python_available: bool
rollout: RolloutPolicy
def __post_init__(self) -> None:
if self.rollout is RolloutPolicy.UNSUPPORTED:
if self.rust is not RustImplementationState.UNIMPLEMENTED or self.python_available:
raise ValueError("an unsupported capability must have neither implementation")
return
if self.rust is RustImplementationState.UNIMPLEMENTED:
if self.rollout is not RolloutPolicy.PYTHON_ONLY or not self.python_available:
raise ValueError("an unimplemented Rust capability must use its Python implementation")
return
if self.rollout is RolloutPolicy.PYTHON_ONLY:
raise ValueError("an implemented Rust capability must declare a Rust rollout")
if self.rollout is RolloutPolicy.RUST_OPT_IN or self.rollout is RolloutPolicy.RUST_OPT_OUT:
if not self.python_available:
raise ValueError("an optional Rust capability requires a Python fallback")
return
if self.rollout is RolloutPolicy.RUST_REQUIRED:
if self.python_available:
raise ValueError("a required Rust capability must have no Python implementation")
return
assert_never(self.rollout)
class CapabilityResolver(Protocol):
def __call__(self, context: CapabilityContext, /) -> CapabilityDefinition: ...
CapabilitySpec: TypeAlias = CapabilityDefinition | CapabilityResolver
class _RustConfiguration:
@ -67,38 +107,72 @@ _CONFIGURATION: Final = _RustConfiguration()
def _parse_env_bool(value: str | None) -> bool | None:
if value is None:
return None
return _ENV_BOOL.validate_python(value.strip())
match value.strip():
case "1":
return True
case "0":
return False
case invalid:
raise ValueError(f"{_GLOBAL_ENV_NAME} must be '1' or '0', got {invalid!r}")
def resolve_rust_enabled(
def resolve_capability(
capability: CapabilityDefinition,
*,
process_override: bool | None,
environment_override: bool | None,
release_default: bool = DEFAULT_RUST_ENABLED,
) -> bool:
if process_override is not None:
return process_override
if environment_override is not None:
return environment_override
return release_default
) -> ExecutionDecision:
match capability.rollout:
case RolloutPolicy.UNSUPPORTED:
return ExecutionDecision.UNSUPPORTED
case RolloutPolicy.PYTHON_ONLY:
return ExecutionDecision.PYTHON
case RolloutPolicy.RUST_REQUIRED:
return ExecutionDecision.RUST_REQUIRED
case RolloutPolicy.RUST_OPT_IN | RolloutPolicy.RUST_OPT_OUT:
enabled: Final = (
process_override
if process_override is not None
else environment_override
if environment_override is not None
else capability.rollout is RolloutPolicy.RUST_OPT_OUT
)
return ExecutionDecision.RUST_WITH_FALLBACK if enabled else ExecutionDecision.PYTHON
assert_never(capability.rollout)
def rust_enabled(route: RouteName | None = None) -> bool:
policy: Final = ROUTE_POLICIES[route] if route is not None else RoutePolicy()
if policy.mode is RouteMode.REQUIRED:
return True
environment: Final = _parse_env_bool(os.getenv(_GLOBAL_ENV_NAME))
if policy.environment_opt_out and environment is False:
return False
return resolve_rust_enabled(
def _definition(spec: CapabilitySpec, context: CapabilityContext) -> CapabilityDefinition:
if isinstance(spec, CapabilityDefinition):
return spec
return spec(context)
def capability_decision(spec: CapabilitySpec, *, context: CapabilityContext) -> ExecutionDecision:
capability: Final = _definition(spec, context)
environment_override: Final = (
_parse_env_bool(os.getenv(_GLOBAL_ENV_NAME))
if _CONFIGURATION.override is None
and capability.rollout in (RolloutPolicy.RUST_OPT_IN, RolloutPolicy.RUST_OPT_OUT)
else None
)
return resolve_capability(
capability,
process_override=_CONFIGURATION.override,
environment_override=environment,
release_default=policy.default_enabled,
environment_override=environment_override,
)
def rust_ocr_enabled() -> bool:
return rust_enabled(RouteName.OCR)
def rust_enabled() -> bool:
environment_override: Final = (
_parse_env_bool(os.getenv(_GLOBAL_ENV_NAME)) if _CONFIGURATION.override is None else None
)
return (
_CONFIGURATION.override
if _CONFIGURATION.override is not None
else environment_override
if environment_override is not None
else DEFAULT_RUST_ENABLED
)
def reset_rust_configuration() -> None:
@ -106,8 +180,5 @@ def reset_rust_configuration() -> None:
def rust(enabled: bool) -> None:
"""Set the process override for optional Rust paths.
Rust-only paths, including Bedrock transcription, are not controlled by this switch.
"""
"""Set the process override for optional Rust capabilities."""
_CONFIGURATION.override = enabled

View file

@ -1,6 +1,5 @@
from typing import Final
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.embeddings.definition import COMPONENT
ROUTE: Final = NativeRoute(RouteName.EMBEDDINGS)
__all__: Final = ("COMPONENT",)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.EMBEDDINGS]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.embeddings import ROUTE
from litellm.rust_bridge.embeddings.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -0,0 +1,10 @@
class RustRouteUnavailableError(RuntimeError):
pass
class RustRouteDeclinedError(RuntimeError):
pass
class RustRouteUnsupportedError(NotImplementedError):
pass

View file

@ -1,6 +1,5 @@
from typing import Final
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.image_edit.definition import COMPONENT
ROUTE: Final = NativeRoute(RouteName.IMAGE_EDIT)
__all__: Final = ("COMPONENT",)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.IMAGE_EDIT]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.image_edit import ROUTE
from litellm.rust_bridge.image_edit.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -1,6 +1,5 @@
from typing import Final
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.image_generation.definition import COMPONENT
ROUTE: Final = NativeRoute(RouteName.IMAGE_GENERATION)
__all__: Final = ("COMPONENT",)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.IMAGE_GENERATION]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.image_generation import ROUTE
from litellm.rust_bridge.image_generation.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -1,8 +1,8 @@
from typing import Final
from litellm.rust_bridge.messages.definition import COMPONENT
from litellm.rust_bridge.messages.types import RustAmessages, RustMessages
from litellm.rust_bridge.messages.value import (
ROUTE,
amessages,
load_rust_amessages,
load_rust_messages,
@ -11,7 +11,7 @@ from litellm.rust_bridge.messages.value import (
)
__all__: Final = (
"ROUTE",
"COMPONENT",
"RustAmessages",
"RustMessages",
"amessages",

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.MESSAGES]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.messages import ROUTE
from litellm.rust_bridge.messages.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import Awaitable
from collections.abc import Awaitable, Mapping
from typing import Protocol
@ -8,12 +8,13 @@ class RustMessages(Protocol):
def __call__(
self,
model: str,
body: dict[str, object],
body: Mapping[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
has_agentic_hook: bool = False,
) -> dict[str, object]:
raise NotImplementedError
@ -22,11 +23,12 @@ class RustAmessages(Protocol):
def __call__(
self,
model: str,
body: dict[str, object],
body: Mapping[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
has_agentic_hook: bool = False,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError

View file

@ -2,6 +2,7 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping
from typing import (
Final,
cast, # noqa: TID251 # native callable signatures are checked by bridge contract tests
@ -10,13 +11,12 @@ from typing import (
import httpx
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode
from litellm.rust_bridge.messages.definition import COMPONENT
from litellm.rust_bridge.messages.types import RustAmessages, RustMessages
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
from litellm.rust_bridge.timeouts import timeout_to_seconds
ROUTE: Final = NativeRoute(RouteName.MESSAGES)
def _as_messages(value: object) -> RustMessages | None:
return cast(RustMessages, value) if callable(value) else None # cast-ok: validated callable native binding
@ -26,8 +26,8 @@ def _as_amessages(value: object) -> RustAmessages | None:
return cast(RustAmessages, value) if callable(value) else None # cast-ok: validated callable native binding
_MESSAGES: Final = ROUTE.bind("messages", validate=_as_messages)
_AMESSAGES: Final = ROUTE.bind("amessages", validate=_as_amessages)
_MESSAGES: Final = COMPONENT.bind("messages", validate=_as_messages)
_AMESSAGES: Final = COMPONENT.bind("amessages", validate=_as_amessages)
def set_rust_messages(
@ -40,56 +40,108 @@ def set_rust_messages(
def load_rust_messages() -> RustMessages | None:
return ROUTE.select(_MESSAGES)
return COMPONENT.resolve().select(_MESSAGES)
def load_rust_amessages() -> RustAmessages | None:
return ROUTE.select(_AMESSAGES)
return COMPONENT.resolve().select(_AMESSAGES)
def messages(
*,
model: str,
body: dict[str, object],
body: Mapping[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
has_agentic_hook: bool = False,
) -> dict[str, object] | None:
rust_messages: Final = load_rust_messages()
if rust_messages is None:
return None
return rust_messages(
model=model,
body=body,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
execution: Final = COMPONENT.resolve(
CapabilityContext(
provider=custom_llm_provider or "",
model=model,
delivery=DeliveryMode.STREAMING if body.get("stream") is True else DeliveryMode.COMPLETED,
)
)
rust_messages: Final = execution.select(_MESSAGES)
native_call: Final[Callable[[], dict[str, object]] | None] = (
(
lambda: rust_messages(
model=model,
body=body,
has_agentic_hook=has_agentic_hook,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
)
)
if rust_messages is not None
else None
)
return invoke(
execution=execution,
native_call=native_call,
python_fallback=lambda: None,
adapt=lambda value: value,
context=BridgeErrorContext(
route=COMPONENT.name.value,
provider=custom_llm_provider or "",
model=model,
),
)
async def amessages(
*,
model: str,
body: dict[str, object],
body: Mapping[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
has_agentic_hook: bool = False,
) -> dict[str, object] | None:
rust_amessages: Final = load_rust_amessages()
if rust_amessages is None:
return None
return await rust_amessages(
model=model,
body=body,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
execution: Final = COMPONENT.resolve(
CapabilityContext(
provider=custom_llm_provider or "",
model=model,
delivery=DeliveryMode.STREAMING if body.get("stream") is True else DeliveryMode.COMPLETED,
)
)
rust_amessages: Final = execution.select(_AMESSAGES)
native_call: Final[Callable[[], Awaitable[dict[str, object]]] | None] = (
(
lambda: rust_amessages(
model=model,
body=body,
has_agentic_hook=has_agentic_hook,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
)
)
if rust_amessages is not None
else None
)
async def python_fallback() -> None:
return None
return await ainvoke(
execution=execution,
native_call=native_call,
python_fallback=python_fallback,
adapt=lambda value: value,
context=BridgeErrorContext(
route=COMPONENT.name.value,
provider=custom_llm_provider or "",
model=model,
),
)

View file

@ -1,6 +1,5 @@
from typing import Final
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.moderation.definition import COMPONENT
ROUTE: Final = NativeRoute(RouteName.MODERATION)
__all__: Final = ("COMPONENT",)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.MODERATION]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.moderation import ROUTE
from litellm.rust_bridge.moderation.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -1,8 +1,8 @@
from typing import Final
from litellm.rust_bridge.ocr.definition import COMPONENT
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest, RustAocr, RustOcr
from litellm.rust_bridge.ocr.value import (
ROUTE,
aocr,
load_rust_aocr,
load_rust_ocr,
@ -10,7 +10,7 @@ from litellm.rust_bridge.ocr.value import (
)
__all__: Final = (
"ROUTE",
"COMPONENT",
"LiteLLMOcrRequest",
"RustAocr",
"RustOcr",

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.OCR]

View file

@ -0,0 +1,52 @@
from __future__ import annotations
from collections.abc import Mapping
from typing import Final, Protocol, cast
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest
from litellm.rust_bridge.ocr.value import adapt_response
from litellm.types.utils import CustomPricingLiteLLMParams
class ExceptionMapper(Protocol):
def __call__(
self,
*,
model: str,
custom_llm_provider: str | None,
original_exception: Exception,
completion_kwargs: dict[str, object],
extra_kwargs: dict[str, object],
) -> Exception: ...
class OcrLifecycleHost:
def response(self, response: Mapping[str, object]) -> OCRResponse:
return adapt_response(response)
def custom_pricing_fields(self) -> tuple[str, ...]:
return tuple(CustomPricingLiteLLMParams.model_fields)
def map_failure(
self,
error: Exception,
request: LiteLLMOcrRequest,
request_provider: str,
) -> Exception:
mapper: Final = cast(ExceptionMapper, litellm.exception_type) # cast-ok: legacy public exception mapper
try:
return mapper(
model=request.model.removeprefix(f"{request_provider}/"),
custom_llm_provider=request_provider,
original_exception=error,
completion_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
extra_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
)
except Exception as public_error:
public_error.__context__ = error
return public_error
HOST: Final = OcrLifecycleHost()

View file

@ -1,60 +1,31 @@
from __future__ import annotations
from collections.abc import Mapping
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
from typing import Final, cast # noqa: TID251 # validates dynamically loaded native callables
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge.ocr import ROUTE, LiteLLMOcrRequest
from litellm.rust_bridge.route import NativeLifecycle
from litellm.rust_bridge.ocr.definition import COMPONENT
from litellm.rust_bridge.ocr.host import HOST
from litellm.rust_bridge.ocr.types import LiteLLMOcrRequest
from litellm.rust_bridge.route import ComponentExecution, NativeLifecycle
NativeOcrLifecycle = NativeLifecycle[LiteLLMOcrRequest, OCRResponse]
class ExceptionMapper(Protocol):
def __call__(
self,
*,
model: str,
custom_llm_provider: str | None,
original_exception: Exception,
completion_kwargs: dict[str, object],
extra_kwargs: dict[str, object],
) -> Exception: ...
def _binding(value: object) -> NativeOcrLifecycle | None:
if not callable(value):
return None
return cast("NativeOcrLifecycle", value) # cast-ok: callable validated at the native binding boundary
LIFECYCLE: Final = ROUTE.bind("_ocr_lifecycle", validate=_binding)
LIFECYCLE: Final = COMPONENT.bind("_ocr_lifecycle", validate=_binding)
NATIVE_OCR_LIFECYCLE: Final = LIFECYCLE
def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None:
def select(request: LiteLLMOcrRequest, execution: ComponentExecution) -> NativeOcrLifecycle | None:
if request.kwargs.get("aocr"):
return None
return ROUTE.select(NATIVE_OCR_LIFECYCLE)
def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]:
return request.kwargs
return execution.select(NATIVE_OCR_LIFECYCLE)
def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception:
mapper: Final = cast( # cast-ok: bounded adapter for the legacy public exception mapper
ExceptionMapper, litellm.exception_type
)
try:
return mapper(
model=request.model.removeprefix(f"{request_provider}/"),
custom_llm_provider=request_provider,
original_exception=error,
completion_kwargs=dict(arguments(request)), # mutable-ok: exception mapper requires owned kwargs
extra_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs
)
except Exception as public_error:
public_error.__context__ = error
return public_error
return HOST.map_failure(error, request, request_provider)

View file

@ -9,9 +9,9 @@ from typing import Final, cast # noqa: TID251 # native extension exposes dynam
import httpx
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.ocr.definition import COMPONENT
from litellm.rust_bridge.ocr.types import RustAocr, RustOcr
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
@ -23,20 +23,19 @@ def _as_aocr(value: object) -> RustAocr | None:
return cast(RustAocr, value) if callable(value) else None
ROUTE: Final = NativeRoute(RouteName.OCR)
_OCR: Final = ROUTE.bind("ocr", validate=_as_ocr)
_AOCR: Final = ROUTE.bind("aocr", validate=_as_aocr)
_OCR: Final = COMPONENT.bind("ocr", validate=_as_ocr)
_AOCR: Final = COMPONENT.bind("aocr", validate=_as_aocr)
def load_rust_ocr() -> RustOcr | None:
return ROUTE.select(_OCR)
return COMPONENT.resolve().select(_OCR)
def load_rust_aocr() -> RustAocr | None:
return ROUTE.select(_AOCR)
return COMPONENT.resolve().select(_AOCR)
def _response(response: Mapping[str, object]) -> OCRResponse:
def adapt_response(response: Mapping[str, object]) -> OCRResponse:
provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY)
normalized: Final = OCRResponse.model_validate(
MappingProxyType({key: value for key, value in response.items() if key != PROVIDER_NATIVE_RESPONSE_KEY})
@ -58,19 +57,28 @@ def ocr(
timeout: float | httpx.Timeout | None,
input_sources: Mapping[str, str] | None = None,
) -> dict[str, object] | None:
rust_ocr: Final = load_rust_ocr()
if rust_ocr is None:
return None
return rust_ocr(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
timeout_seconds=_timeout_to_seconds(timeout),
execution: Final = COMPONENT.resolve()
rust_ocr: Final = execution.select(_OCR)
return invoke(
execution=execution,
native_call=(
lambda: rust_ocr(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
timeout_seconds=_timeout_to_seconds(timeout),
)
)
if rust_ocr is not None
else None,
python_fallback=lambda: None,
adapt=lambda value: value,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)
@ -86,17 +94,29 @@ async def aocr(
timeout: float | httpx.Timeout | None,
input_sources: Mapping[str, str] | None = None,
) -> dict[str, object] | None:
rust_aocr: Final = load_rust_aocr()
if rust_aocr is None:
execution: Final = COMPONENT.resolve()
rust_aocr: Final = execution.select(_AOCR)
async def python_fallback() -> None:
return None
return await rust_aocr(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
timeout_seconds=_timeout_to_seconds(timeout),
return await ainvoke(
execution=execution,
native_call=(
lambda: rust_aocr(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict
timeout_seconds=_timeout_to_seconds(timeout),
)
)
if rust_aocr is not None
else None,
python_fallback=python_fallback,
adapt=lambda value: value,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)

View file

@ -1,6 +1,5 @@
from typing import Final
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.rerank.definition import COMPONENT
ROUTE: Final = NativeRoute(RouteName.RERANK)
__all__: Final = ("COMPONENT",)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.RERANK]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.rerank import ROUTE
from litellm.rust_bridge.rerank.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -1,6 +1,5 @@
from typing import Final
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.responses.definition import COMPONENT
ROUTE: Final = NativeRoute(RouteName.RESPONSES)
__all__: Final = ("COMPONENT",)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.RESPONSES]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.responses import ROUTE
from litellm.rust_bridge.responses.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -2,14 +2,16 @@
from __future__ import annotations
from collections.abc import Awaitable
from collections.abc import Awaitable, Callable
from typing import Final, Protocol, cast # noqa: TID251 # native class is validated at load time
import httpx
from websockets.exceptions import ConnectionClosedOK
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
from litellm.rust_bridge.responses import ROUTE
from litellm.rust_bridge.configuration import CapabilityContext, DeliveryMode
from litellm.rust_bridge.responses.definition import COMPONENT
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -28,6 +30,7 @@ class RustResponsesWebSocketConnection(Protocol):
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
custom_llm_provider: str | None,
) -> Awaitable[RustResponsesWebSocket]: ...
@ -37,7 +40,7 @@ def _as_connection(value: object) -> RustResponsesWebSocketConnection | None:
return cast(RustResponsesWebSocketConnection, value) # cast-ok: native connection factory validated above
_CONNECTION: Final = ROUTE.bind("ResponsesWebSocketConnection", validate=_as_connection)
_CONNECTION: Final = COMPONENT.bind("ResponsesWebSocketConnection", validate=_as_connection)
def set_rust_responses_websocket(
@ -48,7 +51,7 @@ def set_rust_responses_websocket(
def load_rust_responses_websocket() -> RustResponsesWebSocketConnection | None:
return ROUTE.select(_CONNECTION)
return COMPONENT.resolve(CapabilityContext(delivery=DeliveryMode.WEBSOCKET)).select(_CONNECTION)
class _ConnectionAdapter:
@ -73,16 +76,40 @@ async def connect(
url: str,
headers: dict[str, str],
timeout: float | httpx.Timeout | None,
custom_llm_provider: str | None,
model: str,
) -> _ConnectionAdapter | None:
connection_type: Final = load_rust_responses_websocket()
if connection_type is None:
return None
try:
connection: Final = await connection_type.connect(
url=url,
headers=headers,
timeout_seconds=timeout_to_seconds(timeout),
execution: Final = COMPONENT.resolve(
CapabilityContext(
provider=custom_llm_provider or "",
model=model,
delivery=DeliveryMode.WEBSOCKET,
)
except Exception: # noqa: BLE001 # bridge failures must fall back to Python
)
connection_type: Final = execution.select(_CONNECTION)
native_call: Final[Callable[[], Awaitable[RustResponsesWebSocket]] | None] = (
(
lambda: connection_type.connect(
url=url,
custom_llm_provider=custom_llm_provider,
headers=headers,
timeout_seconds=timeout_to_seconds(timeout),
)
)
if connection_type is not None
else None
)
async def python_fallback() -> None:
return None
connection: Final = await ainvoke(
execution=execution,
native_call=native_call,
python_fallback=python_fallback,
adapt=lambda value: value,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)
if connection is None:
return None
return _ConnectionAdapter(connection)

View file

@ -3,17 +3,18 @@ from __future__ import annotations
from collections.abc import Callable, Coroutine
from dataclasses import dataclass
from types import ModuleType
from typing import (
Final,
Literal,
Protocol,
TypeVar,
cast, # noqa: TID251 # validate callability at the native boundary
overload,
)
from typing import Final, Literal, Protocol, TypeVar, cast, overload
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.configuration import ROUTE_POLICIES, RouteName, RoutePolicy, rust_enabled
from litellm.rust_bridge.configuration import (
CapabilityContext,
CapabilitySpec,
ExecutionDecision,
RouteName,
UtilityName,
capability_decision,
)
from litellm.rust_bridge.errors import RustRouteUnavailableError, RustRouteUnsupportedError
BindingT = TypeVar("BindingT")
RequestT = TypeVar("RequestT", contravariant=True)
@ -28,6 +29,7 @@ class NativeLifecycle(Protocol[RequestT, ResponseT]):
args: tuple[object, ...],
kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict
asynchronous: Literal[False],
host: object,
) -> ResponseT: ...
@overload
@ -37,6 +39,7 @@ class NativeLifecycle(Protocol[RequestT, ResponseT]):
args: tuple[object, ...],
kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict
asynchronous: Literal[True],
host: object,
) -> Coroutine[object, object, ResponseT]: ...
@overload
@ -46,6 +49,7 @@ class NativeLifecycle(Protocol[RequestT, ResponseT]):
args: tuple[object, ...],
kwargs: dict[str, object], # mutable-ok: PyO3 requires the original concrete dict
asynchronous: bool,
host: object,
) -> ResponseT | Coroutine[object, object, ResponseT]: ...
@ -56,15 +60,39 @@ def _lifecycle(value: object) -> NativeLifecycle[object, object] | None:
@dataclass(frozen=True, slots=True)
class NativeRoute:
name: RouteName
class ComponentExecution:
route_name: RouteName | UtilityName
decision: ExecutionDecision
@property
def policy(self) -> RoutePolicy:
return ROUTE_POLICIES[self.name]
def require_supported(self) -> None:
if self.decision is ExecutionDecision.UNSUPPORTED:
raise RustRouteUnsupportedError(
f"No Python or Rust implementation for {self.route_name.value}"
)
def enabled(self) -> bool:
return rust_enabled(self.name)
def select(self, binding: NativeBinding[BindingT]) -> BindingT | None:
self.require_supported()
if self.decision is ExecutionDecision.PYTHON:
return None
selected: Final = binding.load()
if selected is None and self.decision is ExecutionDecision.RUST_REQUIRED:
raise RustRouteUnavailableError(
f"Rust {self.route_name.value} bridge is unavailable"
)
return selected
@dataclass(frozen=True, slots=True)
class NativeComponent:
name: RouteName | UtilityName
capability: CapabilitySpec
exports: tuple[str, ...]
def resolve(self, context: CapabilityContext = CapabilityContext()) -> ComponentExecution:
return ComponentExecution(
route_name=self.name,
decision=capability_decision(self.capability, context=context),
)
def bind(
self,
@ -73,11 +101,10 @@ class NativeRoute:
validate: Callable[[object], BindingT | None],
module_loader: Callable[[], ModuleType | None] | None = None,
) -> NativeBinding[BindingT]:
if export not in self.exports:
raise ValueError(f"native export {export!r} is not declared for {self.name.value}")
return NativeBinding(export, validate=validate, module_loader=module_loader)
def select(self, binding: NativeBinding[BindingT]) -> BindingT | None:
return binding.load() if self.enabled() else None
def lifecycle(self) -> NativeBinding[NativeLifecycle[object, object]]:
export: Final = f"_{self.name.value}_lifecycle"
return self.bind(export, validate=_lifecycle)

View file

@ -2,39 +2,22 @@ from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from enum import Enum
from typing import Final, Generic, NoReturn, TypeAlias, TypeVar
from typing import Final, NoReturn, TypeVar
from litellm.exceptions import APIError
from litellm.rust_bridge.bindings import native_exception_types
from litellm.rust_bridge.bindings import (
native_exception_types,
native_host_callback_exception,
native_unavailable_exception,
)
from litellm.rust_bridge.configuration import ExecutionDecision
from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError
from litellm.rust_bridge.route import ComponentExecution
NativeT = TypeVar("NativeT")
ResultT = TypeVar("ResultT")
class FallbackMode(Enum):
PYTHON = "python"
RUST_REQUIRED = "rust_required"
@dataclass(frozen=True, slots=True)
class RustHandled(Generic[ResultT]):
value: ResultT
@dataclass(frozen=True, slots=True)
class RustDeclined:
reason: str
@dataclass(frozen=True, slots=True)
class RustUnavailable:
pass
RustAttempt: TypeAlias = RustHandled[ResultT] | RustDeclined | RustUnavailable
@dataclass(frozen=True, slots=True)
class BridgeErrorContext:
route: str
@ -45,95 +28,120 @@ class BridgeErrorContext:
def invoke(
*,
native_call: Callable[[], NativeT] | None,
fallback: Callable[[], ResultT],
python_fallback: Callable[[], ResultT],
adapt: Callable[[NativeT], ResultT],
mode: FallbackMode,
execution: ComponentExecution,
context: BridgeErrorContext,
) -> ResultT:
result: Final = attempt(native_call=native_call, adapt=adapt, context=context)
if isinstance(result, RustHandled):
return result.value
if mode is FallbackMode.PYTHON:
return fallback()
_raise_required(result, context)
execution.require_supported()
if execution.decision is ExecutionDecision.PYTHON:
return python_fallback()
if native_call is None:
return _unavailable_or_fallback(execution, python_fallback)
exceptions: Final = native_exception_types()
if exceptions is None:
return adapt(native_call())
declined, upstream = exceptions
unavailable: Final = native_unavailable_exception()
host_callback: Final = native_host_callback_exception()
try:
value: Final = native_call()
except host_callback as error:
_raise_host_callback(error)
except unavailable:
return _unavailable_or_fallback(execution, python_fallback)
except declined as error:
return _declined_or_fallback(execution, python_fallback, error)
except upstream as error:
_raise_upstream(error, context)
return adapt(value)
async def ainvoke(
*,
native_call: Callable[[], Awaitable[NativeT]] | None,
fallback: Callable[[], Awaitable[ResultT]],
python_fallback: Callable[[], Awaitable[ResultT]],
adapt: Callable[[NativeT], ResultT],
mode: FallbackMode,
execution: ComponentExecution,
context: BridgeErrorContext,
) -> ResultT:
result: Final = await aattempt(native_call=native_call, adapt=adapt, context=context)
if isinstance(result, RustHandled):
return result.value
if mode is FallbackMode.PYTHON:
return await fallback()
_raise_required(result, context)
def attempt(
*,
native_call: Callable[[], NativeT] | None,
adapt: Callable[[NativeT], ResultT],
context: BridgeErrorContext,
) -> RustAttempt[ResultT]:
execution.require_supported()
if execution.decision is ExecutionDecision.PYTHON:
return await python_fallback()
if native_call is None:
return RustUnavailable()
return await _aunavailable_or_fallback(execution, python_fallback)
exceptions: Final = native_exception_types()
if exceptions is None:
return RustHandled(adapt(native_call()))
declined, upstream = exceptions
try:
value: Final = native_call()
except declined as error:
return RustDeclined(reason=_decline_reason(error))
except upstream as error:
_raise_upstream(error, context)
return RustHandled(adapt(value))
async def aattempt(
*,
native_call: Callable[[], Awaitable[NativeT]] | None,
adapt: Callable[[NativeT], ResultT],
context: BridgeErrorContext,
) -> RustAttempt[ResultT]:
if native_call is None:
return RustUnavailable()
exceptions: Final = native_exception_types()
if exceptions is None:
return RustHandled(adapt(await native_call()))
return adapt(await native_call())
declined, upstream = exceptions
unavailable: Final = native_unavailable_exception()
host_callback: Final = native_host_callback_exception()
try:
value: Final = await native_call()
except host_callback as error:
_raise_host_callback(error)
except unavailable:
return await _aunavailable_or_fallback(execution, python_fallback)
except declined as error:
return RustDeclined(reason=_decline_reason(error))
return await _adeclined_or_fallback(execution, python_fallback, error)
except upstream as error:
_raise_upstream(error, context)
return RustHandled(adapt(value))
return adapt(value)
def _decline_reason(error: BaseException) -> str:
reason: Final[object] = error.args[0] if error.args else str(error)
return reason if isinstance(reason, str) else str(reason)
def _unavailable_or_fallback(
execution: ComponentExecution,
python_fallback: Callable[[], ResultT],
) -> ResultT:
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
return python_fallback()
raise RustRouteUnavailableError(f"Rust {execution.route_name.value} bridge is unavailable")
def _raise_required(
result: RustDeclined | RustUnavailable,
context: BridgeErrorContext,
) -> NoReturn:
raise RuntimeError(f"Rust {context.route} bridge {_required_reason(result)}")
async def _aunavailable_or_fallback(
execution: ComponentExecution,
python_fallback: Callable[[], Awaitable[ResultT]],
) -> ResultT:
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
return await python_fallback()
raise RustRouteUnavailableError(f"Rust {execution.route_name.value} bridge is unavailable")
def _required_reason(result: RustDeclined | RustUnavailable) -> str:
match result:
case RustUnavailable():
return "is unavailable"
case RustDeclined(reason=reason):
return f"declined the request: {reason}"
def _declined_or_fallback(
execution: ComponentExecution,
python_fallback: Callable[[], ResultT],
error: BaseException,
) -> ResultT:
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
return python_fallback()
_raise_declined(execution, error)
async def _adeclined_or_fallback(
execution: ComponentExecution,
python_fallback: Callable[[], Awaitable[ResultT]],
error: BaseException,
) -> ResultT:
if execution.decision is ExecutionDecision.RUST_WITH_FALLBACK:
return await python_fallback()
_raise_declined(execution, error)
def _raise_declined(execution: ComponentExecution, error: BaseException) -> NoReturn:
reason_value: Final[object] = error.args[0] if error.args else str(error)
reason: Final = reason_value if isinstance(reason_value, str) else str(reason_value)
raise RustRouteDeclinedError(
f"Rust {execution.route_name.value} bridge declined the request: {reason}"
) from error
def _raise_host_callback(error: BaseException) -> NoReturn:
cause: Final = error.__cause__
if cause is not None:
raise cause
raise error
def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:

View file

@ -1,6 +1,5 @@
from typing import Final
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.speech.definition import COMPONENT
ROUTE: Final = NativeRoute(RouteName.SPEECH)
__all__: Final = ("COMPONENT",)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.SPEECH]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.speech import ROUTE
from litellm.rust_bridge.speech.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -1,114 +0,0 @@
"""Thin Python wrapper for the native Rust input token counter."""
from __future__ import annotations
from collections.abc import Awaitable
from dataclasses import dataclass
from functools import lru_cache
from typing import Final, Literal, Protocol, cast # noqa: TID251 # native extension exposes untyped callables
from pydantic import TypeAdapter
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.default_encoding import cl100k_base_rank_file, o200k_base_rank_file
from litellm.litellm_core_utils.token_counter import openai_tokenizer_encoding, uses_legacy_message_accounting
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.runtime import BridgeErrorContext, RustHandled, aattempt
from litellm.utils import claude_json_str, huggingface_tokenizer_kind
RustTokenizer = Literal["anthropic", "cl100k_base", "o200k_base"]
class RustTokenCounter(Protocol):
def acount_request(self, body: bytes) -> Awaitable[object]:
raise NotImplementedError
class RustTokenCounterFactory(Protocol):
def __call__(self, tokenizer_json: str) -> RustTokenCounter:
raise NotImplementedError
def from_cl100k_ranks(self, rank_file: str) -> RustTokenCounter:
raise NotImplementedError
def from_o200k_ranks(self, rank_file: str) -> RustTokenCounter:
raise NotImplementedError
@dataclass(frozen=True, slots=True)
class InputTokenCount:
model: str | None
input_tokens: int
_INPUT_TOKEN_COUNT: Final = TypeAdapter(InputTokenCount)
def _as_factory(value: object) -> RustTokenCounterFactory | None:
return (
cast( # cast-ok: native extension protocol is runtime-defined
RustTokenCounterFactory, value
)
if callable(value)
else None
)
TOKEN_COUNTER: Final = NativeBinding("TokenCounter", validate=_as_factory)
def rust_tokenizer(model: str) -> RustTokenizer | None:
"""The Rust counter for the tokenizer `litellm.token_counter` selects for `model`, `None` when Python must count.
Mirrors `_select_tokenizer_helper`: the Anthropic tokenizer has a Rust port, the other HuggingFace
downloads do not, and of the tiktoken encodings `cl100k_base` and `o200k_base` do (p50k/r50k do not). Rust
prices every message with the default constants, so the legacy `gpt-3.5-turbo-0301` accounting stays in
Python."""
if litellm.disable_token_counter is True:
return None
kind: Final = None if litellm.disable_hf_tokenizer_download is True else huggingface_tokenizer_kind(model)
if kind == "anthropic":
return "anthropic"
if kind is not None or uses_legacy_message_accounting(model):
return None
match openai_tokenizer_encoding(model).name:
case "cl100k_base":
return "cl100k_base"
case "o200k_base":
return "o200k_base"
case _:
return None
@lru_cache(maxsize=4)
def _counter(factory: RustTokenCounterFactory, tokenizer: RustTokenizer) -> RustTokenCounter:
match tokenizer:
case "anthropic":
return factory(claude_json_str)
case "cl100k_base":
return factory.from_cl100k_ranks(cl100k_base_rank_file())
case "o200k_base":
return factory.from_o200k_ranks(o200k_base_rank_file())
async def count_input_tokens(body: bytes, tokenizer: RustTokenizer) -> InputTokenCount | None:
if not rust_enabled():
return None
factory: Final = TOKEN_COUNTER.load()
if factory is None:
return None
try:
attempt: Final = await aattempt(
native_call=lambda: _counter(factory, tokenizer).acount_request(body),
adapt=_INPUT_TOKEN_COUNT.validate_python,
context=BridgeErrorContext(route="token_counter", provider=tokenizer, model=""),
)
except (RuntimeError, ValueError) as error:
verbose_logger.debug("Rust token counter (%s) failed, counting in Python: %s", tokenizer, error)
return None
if not isinstance(attempt, RustHandled):
return None
verbose_logger.debug("Rust token counter (%s) counted %d input tokens", tokenizer, attempt.value.input_tokens)
return attempt.value

View file

@ -0,0 +1,12 @@
from typing import Final
from litellm.rust_bridge.token_counter.types import InputTokenCount, RustTokenizer
from litellm.rust_bridge.token_counter.value import TOKEN_COUNTER, count_input_tokens, rust_tokenizer
__all__: Final = (
"TOKEN_COUNTER",
"InputTokenCount",
"RustTokenizer",
"count_input_tokens",
"rust_tokenizer",
)

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import UtilityName
COMPONENT: Final = COMPONENTS[UtilityName.TOKEN_COUNTER]

View file

@ -0,0 +1,31 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Protocol
@dataclass(frozen=True, slots=True)
class RustTokenizer:
kind: str | None
encoding: str
disabled: bool
legacy_accounting: bool
class RustTokenCounter(Protocol):
def __call__(
self,
body: bytes,
kind: str | None,
encoding: str,
disabled: bool,
legacy_accounting: bool,
resource_loader: Callable[[str], str],
) -> Awaitable[object]: ...
@dataclass(frozen=True, slots=True)
class InputTokenCount:
model: str | None
input_tokens: int

View file

@ -0,0 +1,75 @@
from __future__ import annotations
from typing import Final, cast # noqa: TID251 # native extension exposes untyped callables
from pydantic import TypeAdapter
import litellm
from litellm.litellm_core_utils.default_encoding import cl100k_base_rank_file, o200k_base_rank_file
from litellm.litellm_core_utils.token_counter import openai_tokenizer_encoding, uses_legacy_message_accounting
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke
from litellm.rust_bridge.token_counter.definition import COMPONENT
from litellm.rust_bridge.token_counter.types import InputTokenCount, RustTokenCounter, RustTokenizer
from litellm.utils import claude_json_str, huggingface_tokenizer_kind
_INPUT_TOKEN_COUNT: Final = TypeAdapter(InputTokenCount)
def _as_counter(value: object) -> RustTokenCounter | None:
return cast(RustTokenCounter, value) if callable(value) else None
TOKEN_COUNTER: Final = COMPONENT.bind("count_input_tokens", validate=_as_counter)
def rust_tokenizer(model: str) -> RustTokenizer | None:
execution: Final = COMPONENT.resolve()
if execution.select(TOKEN_COUNTER) is None:
return None
kind: Final = None if litellm.disable_hf_tokenizer_download is True else huggingface_tokenizer_kind(model)
encoding: Final = openai_tokenizer_encoding(model).name if kind is None else ""
return RustTokenizer(
kind=kind,
encoding=encoding,
disabled=litellm.disable_token_counter is True,
legacy_accounting=uses_legacy_message_accounting(model),
)
def _tokenizer_resource(tokenizer: str) -> str:
match tokenizer:
case "anthropic":
return claude_json_str
case "cl100k_base":
return cl100k_base_rank_file()
case "o200k_base":
return o200k_base_rank_file()
case _:
raise ValueError(f"unsupported Rust tokenizer resource: {tokenizer}")
async def count_input_tokens(body: bytes, tokenizer: RustTokenizer) -> InputTokenCount | None:
execution: Final = COMPONENT.resolve()
counter: Final = execution.select(TOKEN_COUNTER)
async def python_fallback() -> None:
return None
return await ainvoke(
execution=execution,
native_call=(
lambda: counter(
body,
tokenizer.kind,
tokenizer.encoding,
tokenizer.disabled,
tokenizer.legacy_accounting,
_tokenizer_resource,
)
)
if counter is not None
else None,
python_fallback=python_fallback,
adapt=_INPUT_TOKEN_COUNT.validate_python,
context=BridgeErrorContext(route=COMPONENT.name.value, provider="", model=""),
)

View file

@ -1,8 +1,8 @@
from typing import Final
from litellm.rust_bridge.transcription.definition import COMPONENT
from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription
from litellm.rust_bridge.transcription.value import (
ROUTE,
atranscription,
configure_rust_transcription,
load_rust_atranscription,
@ -11,7 +11,7 @@ from litellm.rust_bridge.transcription.value import (
)
__all__: Final = (
"ROUTE",
"COMPONENT",
"RustAtranscription",
"RustTranscription",
"atranscription",

View file

@ -0,0 +1,6 @@
from typing import Final
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import RouteName
COMPONENT: Final = COMPONENTS[RouteName.TRANSCRIPTION]

View file

@ -1,5 +1,5 @@
from typing import Final
from litellm.rust_bridge.transcription import ROUTE
from litellm.rust_bridge.transcription.definition import COMPONENT
LIFECYCLE: Final = ROUTE.lifecycle()
LIFECYCLE: Final = COMPONENT.lifecycle()

View file

@ -8,13 +8,12 @@ from typing import (
import httpx
from litellm.rust_bridge.bindings import BINDING_UNSET, BindingUnset
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.configuration import CapabilityContext
from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke, invoke
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.rust_bridge.transcription.definition import COMPONENT
from litellm.rust_bridge.transcription.types import RustAtranscription, RustTranscription
ROUTE: Final = NativeRoute(RouteName.TRANSCRIPTION)
def _as_transcription(value: object) -> RustTranscription | None:
return cast(RustTranscription, value) if callable(value) else None # cast-ok: validated callable native binding
@ -24,8 +23,8 @@ def _as_atranscription(value: object) -> RustAtranscription | None:
return cast(RustAtranscription, value) if callable(value) else None # cast-ok: validated callable native binding
_TRANSCRIPTION: Final = ROUTE.bind("transcription", validate=_as_transcription)
_ATRANSCRIPTION: Final = ROUTE.bind("atranscription", validate=_as_atranscription)
_TRANSCRIPTION: Final = COMPONENT.bind("transcription", validate=_as_transcription)
_ATRANSCRIPTION: Final = COMPONENT.bind("atranscription", validate=_as_atranscription)
def configure_rust_transcription(
@ -37,12 +36,12 @@ def configure_rust_transcription(
_ATRANSCRIPTION.configure(atranscription)
def load_rust_transcription() -> RustTranscription | None:
return ROUTE.select(_TRANSCRIPTION)
def load_rust_transcription(*, context: CapabilityContext) -> RustTranscription | None:
return COMPONENT.resolve(context).select(_TRANSCRIPTION)
def load_rust_atranscription() -> RustAtranscription | None:
return ROUTE.select(_ATRANSCRIPTION)
def load_rust_atranscription(*, context: CapabilityContext) -> RustAtranscription | None:
return COMPONENT.resolve(context).select(_ATRANSCRIPTION)
def transcription(
@ -56,18 +55,27 @@ def transcription(
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
rust_transcription: Final = load_rust_transcription()
if rust_transcription is None:
return None
return rust_transcription(
model=model,
audio=audio,
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),
execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model))
rust_transcription: Final = execution.select(_TRANSCRIPTION)
return invoke(
execution=execution,
native_call=(
lambda: rust_transcription(
model=model,
audio=audio,
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),
)
)
if rust_transcription is not None
else None,
python_fallback=lambda: None,
adapt=lambda response: response,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)
@ -82,16 +90,28 @@ async def atranscription(
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
rust_atranscription: Final = load_rust_atranscription()
if rust_atranscription is None:
execution: Final = COMPONENT.resolve(CapabilityContext(provider=custom_llm_provider or "", model=model))
rust_atranscription: Final = execution.select(_ATRANSCRIPTION)
async def python_fallback() -> None:
return None
return await rust_atranscription(
model=model,
audio=audio,
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),
return await ainvoke(
execution=execution,
native_call=(
lambda: rust_atranscription(
model=model,
audio=audio,
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),
)
)
if rust_atranscription is not None
else None,
python_fallback=python_fallback,
adapt=lambda response: response,
context=BridgeErrorContext(route=COMPONENT.name.value, provider=custom_llm_provider or "", model=model),
)

View file

@ -181,6 +181,7 @@ dev = [
"hypothesis==6.165.10",
"reportlab==5.0.1",
"basedpyright==1.39.7",
"mypy==1.20.1",
"keyring==25.7.0",
"pytest==9.0.3",
"tomli==2.4.1; python_version < '3.11'",
@ -288,6 +289,7 @@ profile = "release"
editable-profile = "dev"
include = [
"litellm/proxy/_experimental/out/**",
"litellm/rust_bridge/_native.pyi",
"litellm/router_strategy/complexity_router/artifacts/*.json",
]
exclude = [

View file

@ -57,11 +57,11 @@ PUBLIC_RUST_DISPATCH_MAPPINGS: Final = (
mapping(span="public_sdk_entrypoint", python_frame=r"ocr/main\.py:\d+ a?ocr$"),
mapping(span="public_request", python_frame=r"ocr/main\.py:\d+ _public_request$"),
mapping(span="bind_request", python_frame=r"ocr/main\.py:\d+ _bind_request$"),
mapping(span="rust_ocr_enabled", python_frame=r"rust_bridge/configuration\.py:\d+ rust_ocr_enabled$"),
mapping(span="rust_ocr_enabled", python_frame=r"rust_bridge/configuration\.py:\d+ capability_decision$"),
mapping(span="select_native_ocr", python_frame=r"rust_bridge/ocr/lifecycle\.py:\d+ select$"),
mapping(span="load_native_bridge", python_frame=r"rust_bridge/bindings\.py:\d+ NativeBinding\.load$"),
mapping(span="native_call_setup", python_frame=r"rust_bridge/lifecycle\.py:\d+ setup$"),
mapping(span="native_response", python_frame=r"rust_bridge/ocr\.py:\d+ _response$"),
mapping(span="native_response", python_frame=r"rust_bridge/ocr/host\.py:\d+ OcrLifecycleHost\.response$"),
mapping(span="native_call_finalize", python_frame=r"rust_bridge/lifecycle\.py:\d+ finalize$"),
mapping(
span="native_success_bookkeeping",

View file

@ -47,6 +47,7 @@ class RecordingMessages:
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
has_agentic_hook: bool = False,
) -> dict[str, object]:
self.calls.append(
{
@ -75,6 +76,7 @@ class RecordingAsyncMessages:
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
has_agentic_hook: bool = False,
) -> dict[str, object]:
self.calls.append(
{
@ -238,19 +240,19 @@ async def test_gate_invokes_rust_and_marks_response_header():
@pytest.mark.asyncio
async def test_gate_falls_back_to_python_when_bridge_raises():
async def test_gate_propagates_unclassified_bridge_failure():
bridge = RaisingAsyncMessages()
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate()
assert response is None
with pytest.raises(RuntimeError, match="upstream request failed"):
await _gate()
assert bridge.calls == 1
@pytest.mark.asyncio
async def test_gate_skips_rust_when_flag_absent():
async def test_gate_skips_rust_when_flag_absent(monkeypatch):
monkeypatch.delenv("LITELLM_RUST", raising=False)
bridge = ExplodingAsyncMessages()
rust_messages.set_rust_messages(amessages=bridge)
@ -324,26 +326,20 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
@pytest.mark.asyncio
async def test_gate_skips_rust_for_unsupported_provider():
bridge = ExplodingAsyncMessages()
native = pytest.importorskip("litellm.rust_bridge._native")
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate(custom_llm_provider="openai")
rust_messages.set_rust_messages(amessages=native.amessages)
response = await _gate(custom_llm_provider="openai", api_base="http://127.0.0.1:1")
assert response is None
assert bridge.calls == 0
@pytest.mark.asyncio
async def test_gate_skips_rust_for_agentic_hook():
bridge = ExplodingAsyncMessages()
native = pytest.importorskip("litellm.rust_bridge._native")
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate(has_agentic_hook=True)
rust_messages.set_rust_messages(amessages=native.amessages)
response = await _gate(has_agentic_hook=True, api_base="http://127.0.0.1:1")
assert response is None
assert bridge.calls == 0
@pytest.mark.asyncio

View file

@ -13,9 +13,9 @@ from unittest.mock import MagicMock, patch
import boto3
import httpx
import pytest
from botocore.credentials import Credentials
from botocore.exceptions import ClientError
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.rust_bridge import chat_completions as bridge
@ -54,29 +54,27 @@ RESOLVED_CREDENTIALS = Credentials(
@pytest.fixture(autouse=True)
def reset_bridge(monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
bridge.set_rust_chat_completions(
chat_completions=None, achat_completions=None, decline=None
)
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
yield
bridge.set_rust_chat_completions(
chat_completions=None, achat_completions=None, decline=None
)
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
def _inject(*, decline_reason=None, error: Exception | None = None):
seen: dict[str, list[dict]] = {"gate": [], "call": []}
def gate(**kwargs):
seen["gate"].append(kwargs)
return decline_reason
def native(**kwargs):
seen["call"].append(kwargs)
seen["gate"].append(kwargs)
if decline_reason is not None:
from litellm.rust_bridge import _native
raise _native.RustBridgeDeclined(decline_reason)
if error is not None:
raise error
seen["call"].append(kwargs)
kwargs["on_request"]()
return dict(RUST_RESPONSE)
bridge.set_rust_chat_completions(decline=gate, chat_completions=native)
bridge.set_rust_chat_completions(chat_completions=native)
return seen
@ -142,9 +140,7 @@ def test_the_core_receives_the_converse_url_this_handler_already_built():
seen = _inject()
_run()
assert seen["call"][0]["api_base"].endswith(
"/model/anthropic.claude-sonnet-4-5-v1%3A0/converse"
)
assert seen["call"][0]["api_base"].endswith("/model/anthropic.claude-sonnet-4-5-v1%3A0/converse")
assert "bedrock-runtime.us-east-1.amazonaws.com" in seen["call"][0]["api_base"]
@ -215,9 +211,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
async def declining_native(**_kwargs):
raise _Declined("blank message text")
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, achat_completions=declining_native
)
bridge.set_rust_chat_completions(achat_completions=declining_native)
sentinel = object()
@ -225,16 +219,10 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
return sentinel
with (
patch.object(
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
),
patch.object(
BedrockConverseLLM, "async_completion", side_effect=python_path
) as python_call,
patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS),
patch.object(BedrockConverseLLM, "async_completion", side_effect=python_path) as python_call,
):
result = await BedrockConverseLLM().completion(
**_completion_kwargs(acompletion=True)
)
result = await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True))
assert result is sentinel
assert python_call.called, "a failing rust call must re-enter the python path"
@ -242,22 +230,17 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
@pytest.mark.asyncio
async def test_the_async_path_serves_the_rust_response_without_the_fallback():
async def native(**_kwargs):
async def native(**kwargs):
kwargs["on_request"]()
return dict(RUST_RESPONSE)
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, achat_completions=native
)
bridge.set_rust_chat_completions(achat_completions=native)
with (
patch.object(
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
),
patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS),
patch.object(BedrockConverseLLM, "async_completion") as python_call,
):
result = await BedrockConverseLLM().completion(
**_completion_kwargs(acompletion=True)
)
result = await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True))
assert result.choices[0].message.content == "hello from rust"
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
@ -284,26 +267,19 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines():
async def python_path(**kwargs):
served.append(kwargs)
logging_obj.pre_call(input=kwargs["messages"], api_key="", additional_args={})
return ModelResponse()
with (
patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()),
patch.object(
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
),
patch.object(
BedrockConverseLLM, "async_completion", side_effect=python_path
),
patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS),
patch.object(BedrockConverseLLM, "async_completion", side_effect=python_path),
):
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, achat_completions=declining_native
)
await BedrockConverseLLM().completion(
**_completion_kwargs(acompletion=True, logging_obj=logging_obj)
)
bridge.set_rust_chat_completions(achat_completions=declining_native)
await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True, logging_obj=logging_obj))
assert logging_obj.pre_call.call_count == 1
assert served and served[0]["skip_pre_call_logging"] is True
assert served and "skip_pre_call_logging" not in served[0]
CONVERSE_RESPONSE = {
@ -313,9 +289,7 @@ CONVERSE_RESPONSE = {
}
async def _drive_async_completion(
*, skip_pre_call_logging: bool, logging_obj, credentials: Credentials = RESOLVED_CREDENTIALS
):
async def _drive_async_completion(*, logging_obj, credentials: Credentials = RESOLVED_CREDENTIALS):
"""Run the real `async_completion` with a stubbed transport."""
import httpx as _httpx
@ -345,22 +319,14 @@ async def _drive_async_completion(
credentials=credentials,
headers={},
client=client,
skip_pre_call_logging=skip_pre_call_logging,
)
@pytest.mark.asyncio
async def test_async_completion_honors_the_pre_call_suppression():
logging_obj = MagicMock()
await _drive_async_completion(skip_pre_call_logging=True, logging_obj=logging_obj)
assert logging_obj.pre_call.call_count == 0
@pytest.mark.asyncio
async def test_async_completion_logs_pre_call_by_default():
"""The suppression must be opt-in, so every existing caller keeps its log."""
logging_obj = MagicMock()
await _drive_async_completion(skip_pre_call_logging=False, logging_obj=logging_obj)
await _drive_async_completion(logging_obj=logging_obj)
assert logging_obj.pre_call.call_count == 1
@ -373,7 +339,7 @@ async def test_async_completion_signs_off_the_event_loop(monkeypatch):
release = asyncio.create_task(probe.release_refresh_from_the_loop())
response = await _drive_async_completion(
skip_pre_call_logging=False, logging_obj=MagicMock(), credentials=probe.credentials()
logging_obj=MagicMock(), credentials=probe.credentials()
)
await release
@ -414,9 +380,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines():
logging_obj = MagicMock()
with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()):
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, chat_completions=declining_native
)
bridge.set_rust_chat_completions(chat_completions=declining_native)
response = _run(
logging_obj=logging_obj,
client=_sync_client_returning_converse_response(),
@ -462,20 +426,15 @@ async def test_post_call_logging_fires_on_the_async_rust_path():
cannot drift apart the way the pre_call suppression once did."""
import json
async def native(**_kwargs):
async def native(**kwargs):
kwargs["on_request"]()
return dict(RUST_RESPONSE)
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, achat_completions=native
)
bridge.set_rust_chat_completions(achat_completions=native)
logging_obj = MagicMock()
with patch.object(
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
):
await BedrockConverseLLM().completion(
**_completion_kwargs(acompletion=True, logging_obj=logging_obj)
)
with patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS):
await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True, logging_obj=logging_obj))
assert logging_obj.post_call.call_count == 1
logged = logging_obj.post_call.call_args.kwargs["original_response"]
@ -500,9 +459,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines():
logging_obj, calls = _recording_logging_obj()
with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()):
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, chat_completions=declining_native
)
bridge.set_rust_chat_completions(chat_completions=declining_native)
response = _run(
logging_obj=logging_obj,
client=_sync_client_returning_converse_response(),

View file

@ -13,29 +13,28 @@ from botocore.credentials import RefreshableCredentials
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.code_interpreter_interception.handler import (
CodeInterpreterInterceptionLogger,
LITELLM_CODE_EXECUTION_TOOL_NAME,
CodeInterpreterInterceptionLogger,
)
from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
BaseAudioTranscriptionConfig,
)
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.base_llm.search.transformation import BaseSearchConfig, SearchResponse
from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
from litellm.llms.brave.search.transformation import BraveSearchConfig
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
BaseLLMHTTPHandler,
_collect_ws_project_quota_callbacks,
_google_genai_streaming_hidden_params,
_has_pre_call_deployment_hook,
_rust_responses_websocket_enabled,
)
from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
@ -1405,9 +1404,7 @@ def test_sync_delete_responses_sets_json_content_type():
({}, True, None, None),
],
)
def test_resolve_anthropic_messages_timeout(
monkeypatch, litellm_params_kwargs, stream, global_timeout, expected
):
def test_resolve_anthropic_messages_timeout(monkeypatch, litellm_params_kwargs, stream, global_timeout, expected):
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
if global_timeout is None:
@ -1423,9 +1420,7 @@ def test_resolve_anthropic_messages_timeout(
)
else:
monkeypatch.setattr("litellm.request_timeout", global_timeout, raising=False)
monkeypatch.setattr(
"litellm.request_timeout_explicitly_set", True, raising=False
)
monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False)
resolved = BaseLLMHTTPHandler._resolve_anthropic_messages_timeout(
litellm_params=GenericLiteLLMParams(**litellm_params_kwargs),
@ -1450,9 +1445,7 @@ async def test_async_anthropic_messages_handler_forwards_request_timeout(monkeyp
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
)
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude", "messages": []}
)
mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []})
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
mock_config.max_retry_on_anthropic_messages_http_error = 1
@ -1498,9 +1491,7 @@ async def test_async_anthropic_messages_handler_forwards_stream_timeout(monkeypa
return_value=({"x-api-key": "k"}, "https://api.anthropic.com")
)
mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude", "messages": []}
)
mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []})
mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages")
mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None))
mock_config.max_retry_on_anthropic_messages_http_error = 1
@ -1691,6 +1682,7 @@ def _make_responses_handler_call(signed_body):
signing provider (e.g. Bedrock Mantle).
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
@ -1746,6 +1738,7 @@ def test_responses_handler_signs_after_fake_stream_prep_strips_stream():
We snapshot request_data at sign time and assert "stream" is already gone.
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.llms.openai import ResponsesAPIResponse
@ -1809,6 +1802,7 @@ def _make_compact_handler_call(signed_body, is_async):
signing provider (e.g. Bedrock Mantle SigV4 / bearer).
"""
from unittest.mock import MagicMock
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
@ -1910,7 +1904,13 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
)
mock_config.sign_request = Mock(return_value=({}, None))
fake_raw_response = {"id": "msg_1", "type": "message", "role": "assistant", "content": [], "stop_reason": "end_turn"}
fake_raw_response = {
"id": "msg_1",
"type": "message",
"role": "assistant",
"content": [],
"stop_reason": "end_turn",
}
mock_config.transform_anthropic_messages_response = Mock(return_value=fake_raw_response)
mock_logging_obj = Mock()
@ -1930,10 +1930,17 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks(
mock_httpx_response.status_code = 200
with (
patch.object(handler, "_async_post_anthropic_messages_with_http_error_retry", new=AsyncMock(return_value=mock_httpx_response)),
patch.object(
handler,
"_async_post_anthropic_messages_with_http_error_retry",
new=AsyncMock(return_value=mock_httpx_response),
),
patch.object(handler, "_call_agentic_completion_hooks", side_effect=fake_agentic_hooks),
patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client"),
patch("litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers", return_value=None),
patch(
"litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers",
return_value=None,
),
):
result = await handler.async_anthropic_messages_handler(
model="claude-haiku",
@ -2175,7 +2182,9 @@ def test_audio_transcriptions_sends_dict_data_as_json_body():
form-encodes it and silently ignores json=; JSON-body providers (e.g.
Google Speech-to-Text) need an application/json body."""
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured))))
client = HTTPHandler(
client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured)))
)
response = BaseLLMHTTPHandler().audio_transcriptions(
client=client,
@ -2259,9 +2268,7 @@ def _transform_subtitle_response(payload):
def test_subtitle_synthesis_fallback_without_timings_drops_words():
response = _transform_subtitle_response(
{"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]}
)
response = _transform_subtitle_response({"text": "hello world", "words": [{"word": "hello"}, {"word": "world"}]})
assert response.text == "hello world"
assert "words" not in response
@ -2481,9 +2488,7 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques
ok_response = httpx.Response(200, json={"id": "msg_1"}, request=httpx.Request("POST", request_url))
class FakeAsyncClient:
async def post(
self, url, headers, data, stream=False, logging_obj=None, timeout=None
):
async def post(self, url, headers, data, stream=False, logging_obj=None, timeout=None):
posts.append({"headers": dict(headers), "data": data})
return invalid_signature_response if len(posts) == 1 else ok_response
@ -2912,21 +2917,6 @@ async def test_generic_http_handler_async_streaming_forwards_provider_response_h
assert "".join([chunk.choices[0].delta.content or "" for chunk in collected]) == "hi"
@pytest.mark.parametrize(
"custom_llm_provider, enabled, expected",
[("openai", True, True), ("openai", False, False), ("azure", True, False),
("hosted_vllm", True, False), (None, True, False)],
)
def test_the_rust_responses_websocket_needs_openai_and_process_enablement(
custom_llm_provider, enabled, expected, monkeypatch
):
from litellm.rust_bridge import configuration
configuration.reset_rust_configuration()
monkeypatch.setenv("LITELLM_RUST", "1" if enabled else "0")
assert _rust_responses_websocket_enabled(custom_llm_provider) is expected
def test_a_plain_callback_does_not_advertise_a_pre_call_deployment_hook(monkeypatch):
from litellm.integrations.custom_logger import CustomLogger
@ -3044,7 +3034,13 @@ def _capture_video_create_request(captured):
captured["body"] = request.content
return httpx.Response(
200,
json={"id": "video_123", "object": "video", "status": "queued", "created_at": 1712697600, "model": "sora-2"},
json={
"id": "video_123",
"object": "video",
"status": "queued",
"created_at": 1712697600,
"model": "sora-2",
},
)
return respond
@ -3067,7 +3063,9 @@ def test_video_generation_without_file_sends_multipart_form_data():
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(OpenAIVideoConfig()))
result = BaseLLMHTTPHandler().video_generation_handler(
client=client, **_video_create_call_kwargs(OpenAIVideoConfig())
)
assert captured["content_type"].startswith("multipart/form-data")
assert _multipart_text_fields(captured["content_type"], captured["body"]) == {
@ -3107,7 +3105,9 @@ def test_azure_video_generation_without_file_sends_multipart_form_data():
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(AzureVideoConfig()))
result = BaseLLMHTTPHandler().video_generation_handler(
client=client, **_video_create_call_kwargs(AzureVideoConfig())
)
assert captured["content_type"].startswith("multipart/form-data")
assert _multipart_text_fields(captured["content_type"], captured["body"]) == {
@ -3122,7 +3122,9 @@ def test_video_generation_json_provider_keeps_json_body():
captured = {}
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_video_create_request(captured))))
result = BaseLLMHTTPHandler().video_generation_handler(client=client, **_video_create_call_kwargs(_JSONBodyVideoConfig()))
result = BaseLLMHTTPHandler().video_generation_handler(
client=client, **_video_create_call_kwargs(_JSONBodyVideoConfig())
)
assert captured["content_type"] == "application/json"
assert json.loads(captured["body"]) == {"model": "sora-2", "prompt": "a cat surfing", "seconds": "4"}
@ -3152,6 +3154,7 @@ def test_video_generation_with_input_reference_keeps_file_multipart():
AZURE_AI_BASE = "https://myfoundry.services.ai.azure.com"
AZURE_AI_CHAT_COMPLETIONS_URL = f"{AZURE_AI_BASE}/models/chat/completions"
def _a_tool_with_an_unsupported_field() -> dict:
return {
"type": "function",
@ -3159,14 +3162,13 @@ def _a_tool_with_an_unsupported_field() -> dict:
"strict": True,
}
A_COMPLETION = {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 1,
"model": "grok-3",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "sent"}, "finish_reason": "stop"}
],
"choices": [{"index": 0, "message": {"role": "assistant", "content": "sent"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
@ -3210,9 +3212,7 @@ def _call_azure_ai(recorder: _RecordedAzureAI, **overrides):
def test_a_tool_field_the_provider_rejects_is_dropped_and_the_call_retried():
recorder = _RecordedAzureAI(
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
)
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
response = _call_azure_ai(recorder)
@ -3223,9 +3223,7 @@ def test_a_tool_field_the_provider_rejects_is_dropped_and_the_call_retried():
def test_the_retry_changes_only_the_field_the_provider_named():
recorder = _RecordedAzureAI(
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
)
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
_call_azure_ai(recorder)
@ -3264,9 +3262,7 @@ def test_an_extra_input_outside_a_tool_is_not_retried_unless_dropping_params_was
def test_an_extra_input_outside_a_tool_is_retried_when_dropping_params_was_asked_for():
recorder = _RecordedAzureAI(
[_rejection(UNRELATED_REJECTION), httpx.Response(200, json=A_COMPLETION)]
)
recorder = _RecordedAzureAI([_rejection(UNRELATED_REJECTION), httpx.Response(200, json=A_COMPLETION)])
response = _call_azure_ai(recorder, drop_params=True)
@ -3280,9 +3276,7 @@ async def test_a_tool_field_the_provider_rejects_is_dropped_and_retried_on_the_a
):
import respx
recorder = _RecordedAzureAI(
[_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)]
)
recorder = _RecordedAzureAI([_rejection(TOOL_LEVEL_REJECTION), httpx.Response(200, json=A_COMPLETION)])
with respx.mock(assert_all_called=True) as router:
router.post(AZURE_AI_CHAT_COMPLETIONS_URL).mock(side_effect=recorder)
@ -3487,7 +3481,9 @@ def _start_async_completion(config, logging_obj=None):
custom_llm_provider="openai",
model_response=ModelResponse(),
encoding=None,
logging_obj=logging_obj if logging_obj is not None else Mock(dynamic_success_callbacks=None, model_call_details={}),
logging_obj=logging_obj
if logging_obj is not None
else Mock(dynamic_success_callbacks=None, model_call_details={}),
optional_params={},
timeout=10.0,
litellm_params={},

View file

@ -1,7 +1,7 @@
import importlib
from collections.abc import AsyncGenerator
from datetime import datetime
from io import BytesIO
from types import SimpleNamespace
from typing import Final
from unittest.mock import Mock
@ -61,8 +61,11 @@ async def test_python_request_response_and_callbacks(
if dispatch != "disabled":
monkeypatch.setenv("LITELLM_RUST", "1")
NATIVE_OCR_LIFECYCLE.override(Mock(side_effect=Declined()) if dispatch == "declined" else None)
main: Final = importlib.import_module("litellm.ocr.main")
monkeypatch.setattr(main, "native_exception_types", lambda: (Declined, RuntimeError))
native: Final = SimpleNamespace(
RustBridgeDeclined=Declined,
RustUpstreamError=RuntimeError,
)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
logger: Final = Mock(spec=CustomLogger)
monkeypatch.setattr(litellm, "input_callback", [logger])
arguments: Final = {

View file

@ -2,7 +2,6 @@ from __future__ import annotations
import pytest
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
from litellm.rust_bridge import configuration
from litellm.rust_bridge.responses import websocket as responses_websocket
@ -35,6 +34,7 @@ class _FakeNativeBridge:
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
custom_llm_provider: str | None,
) -> _FakeNativeConnection:
return _FakeNativeConnection()
@ -48,14 +48,6 @@ def reset_responses_websocket():
configuration.reset_rust_configuration()
def test_rust_websocket_bridge_uses_process_enablement() -> None:
configuration.rust(False)
assert not _rust_responses_websocket_enabled("openai")
configuration.rust(True)
assert _rust_responses_websocket_enabled("openai")
assert not _rust_responses_websocket_enabled("anthropic")
@pytest.mark.asyncio
async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection())
@ -72,6 +64,8 @@ async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch)
assert (
await responses_websocket.connect(
url="wss://example.test/responses",
custom_llm_provider="openai",
model="test",
headers={},
timeout=None,
)
@ -88,6 +82,8 @@ async def test_enabled_bridge_connects_and_adapts_socket(
connection = await responses_websocket.connect(
url="wss://example.test/responses",
custom_llm_provider="openai",
model="test",
headers={"Authorization": "Bearer key"},
timeout=1.0,
)
@ -96,3 +92,51 @@ async def test_enabled_bridge_connects_and_adapts_socket(
await connection.send("response.create")
assert await connection.recv() == "response.completed"
await connection.close()
@pytest.mark.asyncio
async def test_disabled_websocket_does_not_connect() -> None:
class UnexpectedConnection:
@classmethod
async def connect(cls, **kwargs: object) -> None:
raise AssertionError("disabled Rust must not connect")
responses_websocket.set_rust_responses_websocket(connection=UnexpectedConnection)
configuration.rust(False)
assert (
await responses_websocket.connect(
url="ws://127.0.0.1:1",
headers={},
timeout=0.1,
custom_llm_provider="openai",
model="test",
)
is None
)
@pytest.mark.asyncio
async def test_native_websocket_decline_falls_back_but_connection_failure_does_not() -> None:
from litellm.exceptions import APIError
native = pytest.importorskip("litellm.rust_bridge._native")
responses_websocket.set_rust_responses_websocket(connection=native.ResponsesWebSocketConnection)
configuration.rust(True)
assert (
await responses_websocket.connect(
url="ws://127.0.0.1:1",
headers={},
timeout=0.1,
custom_llm_provider="azure",
model="test",
)
is None
)
with pytest.raises(APIError):
await responses_websocket.connect(
url="ws://127.0.0.1:1",
headers={},
timeout=0.1,
custom_llm_provider="openai",
model="test",
)

View file

@ -0,0 +1,2 @@
[mypy]
follow_imports = skip

View file

@ -1,20 +1,17 @@
"""Tests for the Rust chat completions bridge.
The native callables are dependency-injected through
``set_rust_chat_completions`` rather than patched, so these run without the
compiled extension present.
"""
from __future__ import annotations
from typing import Final
import pytest
import litellm
from litellm.rust_bridge import configuration
from litellm.rust_bridge import chat_completions as bridge
from litellm.rust_bridge import configuration
from litellm.types.utils import ModelResponse
RUST_RESPONSE = {
native = pytest.importorskip("litellm.rust_bridge._native")
RUST_RESPONSE: Final = {
"created": 1_700_000_000,
"model": "claude-sonnet-4-5-20260101",
"choices": [
@ -24,27 +21,12 @@ RUST_RESPONSE = {
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 11,
"completion_tokens": 4,
"total_tokens": 15,
"prompt_tokens_details": {
"cached_tokens": 0,
"cache_creation_tokens": 0,
"text_tokens": 11,
},
},
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
MESSAGES: Final = [{"role": "user", "content": "hi"}]
MESSAGES = [{"role": "user", "content": "hi"}]
class _FakeDeclined(Exception):
"""Stands in for the native `RustBridgeDeclined`."""
class _FakeUpstream(Exception):
"""Stands in for the native `RustUpstreamError`; args are (status, message)."""
_FakeDeclined = native.RustBridgeDeclined
_FakeUpstream = native.RustUpstreamError
class _FakeNative:
@ -52,181 +34,42 @@ class _FakeNative:
RustUpstreamError = _FakeUpstream
def _fake_native_bridge(monkeypatch):
"""Expose the bridge's exception classes without the compiled extension."""
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
def _hide_native_bridge(monkeypatch):
"""Simulate a wheel built without the compiled extension.
There is no injection seam for "the .so is absent", so the loader itself is
replaced; every other case here uses `set_rust_chat_completions`.
"""
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
@pytest.fixture(autouse=True)
def reset_bridge(monkeypatch):
"""Every test starts with no injected callables, and leaves none behind."""
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
configuration.reset_rust_configuration()
monkeypatch.setenv("LITELLM_RUST", "1")
yield
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
configuration.reset_rust_configuration()
class _RecordingDecline:
"""A stand-in for the native gate that records what it was asked."""
def __init__(self, reason: str | None = None):
self.reason = reason
self.calls: list[dict] = []
def __call__(self, **kwargs):
self.calls.append(kwargs)
return self.reason
class _RecordingCall:
def __init__(self, result=None, error: Exception | None = None):
self.result = result if result is not None else dict(RUST_RESPONSE)
self.error = error
self.calls: list[dict] = []
def __init__(self, result: object = RUST_RESPONSE, error: Exception | None = None) -> None:
self.result: Final = result
self.error: Final = error
self.calls: Final[list[dict[str, object]]] = []
def __call__(self, **kwargs):
def __call__(self, **kwargs: object) -> object:
self.calls.append(kwargs)
if self.error is not None:
raise self.error
on_request: Final = kwargs["on_request"]
assert callable(on_request)
on_request()
return self.result
class _RecordingAsyncCall(_RecordingCall):
async def __call__(self, **kwargs):
return _RecordingCall.__call__(self, **kwargs)
async def __call__(self, **kwargs: object) -> object:
return super().__call__(**kwargs)
def _accepts(**overrides) -> bool:
kwargs = {
"model": "claude-sonnet-4-5",
"messages": MESSAGES,
"optional_params": {"max_tokens": 16},
"custom_llm_provider": "anthropic",
"litellm_params": {},
"stream": None,
}
kwargs.update(overrides)
return bridge.rust_chat_completions_accepts(**kwargs)
@pytest.fixture(autouse=True)
def reset_bridge(monkeypatch: pytest.MonkeyPatch):
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
configuration.reset_rust_configuration()
monkeypatch.setenv("LITELLM_RUST", "1")
yield
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None)
configuration.reset_rust_configuration()
class TestGate:
def test_declines_when_the_deployment_did_not_opt_in(self, monkeypatch):
monkeypatch.delenv("LITELLM_RUST", raising=False)
gate = _RecordingDecline()
bridge.set_rust_chat_completions(decline=gate)
assert _accepts(litellm_params={}) is False
assert _accepts(litellm_params=None) is False
assert gate.calls == [], "the gate must not be consulted before opt-in"
def test_accepts_when_the_deployment_opted_in_and_the_core_agrees(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
gate = _RecordingDecline()
bridge.set_rust_chat_completions(decline=gate)
assert _accepts() is True
assert gate.calls[0]["model"] == "claude-sonnet-4-5"
assert gate.calls[0]["custom_llm_provider"] == "anthropic"
def test_process_enable_applies_without_request_override(self):
bridge.set_rust_chat_completions(decline=_RecordingDecline())
configuration.rust(True)
assert _accepts(litellm_params={}) is True
def test_the_env_var_opts_in_without_a_per_model_flag(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "true")
bridge.set_rust_chat_completions(decline=_RecordingDecline())
assert _accepts(litellm_params={}) is True
def test_declines_streaming_and_providers_off_the_path(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
gate = _RecordingDecline()
bridge.set_rust_chat_completions(decline=gate)
assert _accepts(stream=True) is False
assert _accepts(custom_llm_provider="openai") is False
assert _accepts(custom_llm_provider=None) is False
assert gate.calls == []
def test_declines_an_anthropic_request_carrying_a_litellm_metadata_user_id(self, monkeypatch):
"""`AnthropicConfig.transform_request` copies a valid `user_id` into the Messages body.
It does that inside the function the Rust route replaces, and the core is
handed `optional_params` only, so accepting here would send the request
to Anthropic with the abuse-detection attribution silently missing.
"""
monkeypatch.setenv("LITELLM_RUST", "1")
gate = _RecordingDecline()
bridge.set_rust_chat_completions(decline=gate)
assert _accepts(litellm_params={"metadata": {"user_id": "u-123"}}) is False
assert gate.calls == [], "the core must not be consulted for a request it cannot see the key of"
# Bedrock's Converse transform reads no `user_id`, and an Anthropic request
# whose metadata carries none is one Python would not attribute either.
assert (
_accepts(
custom_llm_provider="bedrock",
model="bedrock/us-east-1/anthropic.claude-v2",
litellm_params={"metadata": {"user_id": "u-123"}},
)
is True
)
assert _accepts(litellm_params={"metadata": {"trace_id": "t-1"}}) is True
assert _accepts(litellm_params={"metadata": {"user_id": None}}) is True
assert _accepts(litellm_params={"metadata": None}) is True
def test_declines_a_bedrock_request_while_the_proxy_owns_request_metadata(self, monkeypatch):
"""`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the
Converse body from `litellm_params`, and owning that field also means
evicting a caller-supplied one. The core can do neither, so an operator
who armed `bedrock_request_metadata_fields` keeps the Python path.
"""
monkeypatch.setenv("LITELLM_RUST", "1")
gate = _RecordingDecline()
bridge.set_rust_chat_completions(decline=gate)
bedrock = {
"custom_llm_provider": "bedrock",
"model": "bedrock/us-east-1/anthropic.claude-v2",
}
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["user_api_key_team_id"])
assert _accepts(**bedrock) is False
assert gate.calls == [], "the core must not be consulted for a field it cannot write"
assert _accepts() is True, "arming Bedrock attribution must not decline Anthropic"
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None)
assert _accepts(**bedrock) is True, "the decline follows the operator's opt-in alone"
def test_declines_when_the_core_declines(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
bridge.set_rust_chat_completions(decline=_RecordingDecline("streaming"))
assert _accepts() is False
def test_declines_when_the_bridge_is_unavailable(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
_hide_native_bridge(monkeypatch)
assert _accepts() is False
def test_declines_when_the_gate_itself_raises(self, monkeypatch):
monkeypatch.setenv("LITELLM_RUST", "1")
def exploding(**_kwargs):
raise RuntimeError("boom")
bridge.set_rust_chat_completions(decline=exploding)
assert _accepts() is False
def _fake_native_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
def _call_kwargs(model_response: ModelResponse) -> dict:
def _call_kwargs(model_response: ModelResponse, fallback: object = "python") -> dict[str, object]:
return {
"model": "claude-sonnet-4-5",
"messages": MESSAGES,
@ -237,159 +80,92 @@ def _call_kwargs(model_response: ModelResponse) -> dict:
"custom_llm_provider": "anthropic",
"extra_headers": {},
"timeout": 30.0,
"on_response": lambda _rust_response: None,
"python_fallback": lambda: fallback,
}
class TestSyncCall:
def test_builds_a_model_response_and_stamps_the_rust_header(self):
native = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native)
model_response = ModelResponse()
original_id = model_response.id
def test_sync_native_entrypoint_runs_once_and_logs_once() -> None:
events: Final[list[str]] = []
native_call: Final = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native_call)
kwargs: Final = _call_kwargs(ModelResponse())
kwargs.update({"on_request": lambda: events.append("pre"), "on_response": lambda _value: events.append("post")})
result = bridge.chat_completions(**_call_kwargs(model_response))
result: Final = bridge.chat_completions(**kwargs)
assert result is not None
assert result.choices[0].message.content == "hello from rust"
assert result.choices[0].finish_reason == "stop"
assert result.model == "claude-sonnet-4-5-20260101"
assert result.usage.prompt_tokens == 11
assert result.usage.completion_tokens == 4
assert result.usage.total_tokens == 15
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
assert result.id == original_id, "the rust path must keep the chatcmpl id litellm already minted"
def test_passes_the_timeout_through_as_seconds(self):
native = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native)
bridge.chat_completions(**_call_kwargs(ModelResponse()))
assert native.calls[0]["timeout_seconds"] == 30.0
def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
_hide_native_bridge(monkeypatch)
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
_fake_native_bridge(monkeypatch)
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
assert isinstance(result, ModelResponse)
assert result.choices[0].message.content == "hello from rust"
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
assert len(native_call.calls) == 1
assert events == ["pre", "post"]
class TestAsyncCall:
@pytest.mark.asyncio
async def test_builds_a_model_response(self):
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
result = await bridge.achat_completions(**_call_kwargs(ModelResponse()))
assert result is not None
assert result.choices[0].message.content == "hello from rust"
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
def test_decline_has_no_logging_effect_and_runs_one_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
_fake_native_bridge(monkeypatch)
events: Final[list[str]] = []
native_call: Final = _RecordingCall(error=_FakeDeclined("unsupported"))
bridge.set_rust_chat_completions(chat_completions=native_call)
kwargs: Final = _call_kwargs(ModelResponse(), fallback="python")
kwargs.update({"on_request": lambda: events.append("pre"), "on_response": lambda _value: events.append("post")})
@pytest.mark.asyncio
async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
_hide_native_bridge(monkeypatch)
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
@pytest.mark.asyncio
async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
_fake_native_bridge(monkeypatch)
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
assert bridge.chat_completions(**kwargs) == "python"
assert len(native_call.calls) == 1
assert events == []
class TestAsyncFallbackWrapper:
@pytest.mark.asyncio
async def test_returns_the_rust_response_without_running_the_fallback(self):
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
ran = []
async def fallback():
ran.append(True)
return "python"
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result.choices[0].message.content == "hello from rust"
assert ran == []
@pytest.mark.asyncio
async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch):
_fake_native_bridge(monkeypatch)
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
async def fallback():
return "python"
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"
@pytest.mark.asyncio
async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch):
_hide_native_bridge(monkeypatch)
async def fallback():
return "python"
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"
def test_streaming_uses_python_without_loading_native() -> None:
native_call: Final = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native_call)
kwargs: Final = _call_kwargs(ModelResponse())
kwargs["stream"] = True
assert bridge.chat_completions(**kwargs) == "python"
assert native_call.calls == []
class TestFailureClassification:
"""A failure the provider already saw must not be retried on the Python
path: it would bill the customer for the same work twice."""
def test_host_facts_reach_the_single_native_call(monkeypatch: pytest.MonkeyPatch) -> None:
native_call: Final = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native_call)
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["team_id"])
kwargs: Final = _call_kwargs(ModelResponse())
kwargs["litellm_params"] = {"metadata": {"user_id": "u-1"}}
bridge.chat_completions(**kwargs)
assert native_call.calls[0]["host_facts"] == {
"stream": False,
"anthropic_user_id": True,
"bedrock_metadata_owned": True,
}
@pytest.fixture(autouse=True)
def _native_exceptions(self, monkeypatch):
_fake_native_bridge(monkeypatch)
def test_a_decline_falls_back_because_nothing_was_sent(self):
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
def test_upstream_and_adaptation_failures_never_fall_back(monkeypatch: pytest.MonkeyPatch) -> None:
_fake_native_bridge(monkeypatch)
fallback_calls: Final[list[bool]] = []
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "rate limited")))
kwargs: Final = _call_kwargs(ModelResponse())
kwargs["python_fallback"] = lambda: fallback_calls.append(True)
with pytest.raises(litellm.APIError, match="rate limited"):
bridge.chat_completions(**kwargs)
assert fallback_calls == []
def test_an_upstream_failure_is_surfaced_with_its_status(self):
from litellm.exceptions import APIError
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
kwargs["on_response"] = lambda _value: (_ for _ in ()).throw(RuntimeError("adapt failed"))
with pytest.raises(RuntimeError, match="adapt failed"):
bridge.chat_completions(**kwargs)
assert fallback_calls == []
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited")))
with pytest.raises(APIError) as raised:
bridge.chat_completions(**_call_kwargs(ModelResponse()))
assert raised.value.status_code == 429
assert "rate limited" in str(raised.value)
def test_a_transport_failure_with_no_response_surfaces_as_a_500(self):
from litellm.exceptions import APIError
@pytest.mark.asyncio
async def test_async_native_and_fallback_paths(monkeypatch: pytest.MonkeyPatch) -> None:
native_call: Final = _RecordingAsyncCall()
bridge.set_rust_chat_completions(achat_completions=native_call)
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset")))
with pytest.raises(APIError) as raised:
bridge.chat_completions(**_call_kwargs(ModelResponse()))
assert raised.value.status_code == 500
async def fallback() -> str:
return "python"
def test_an_unrecognized_error_is_not_swallowed(self):
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else")))
with pytest.raises(RuntimeError):
bridge.chat_completions(**_call_kwargs(ModelResponse()))
kwargs: Final = _call_kwargs(ModelResponse())
kwargs["python_fallback"] = fallback
result: Final = await bridge.achat_completions(**kwargs)
assert isinstance(result, ModelResponse)
assert len(native_call.calls) == 1
@pytest.mark.asyncio
async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self):
from litellm.exceptions import APIError
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom")))
ran = []
async def fallback():
ran.append(True)
return "python"
with pytest.raises(APIError):
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert ran == [], "a request the provider already served must not be re-issued"
@pytest.mark.asyncio
async def test_the_async_wrapper_falls_back_on_a_decline(self):
bridge.set_rust_chat_completions(
achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text"))
)
async def fallback():
return "python"
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"
configuration.rust(False)
assert await bridge.achat_completions(**kwargs) == "python"

View file

@ -1,138 +0,0 @@
from __future__ import annotations
import os
import subprocess
import sys
from collections.abc import Generator
from concurrent.futures import ThreadPoolExecutor
from typing import Final
import pytest
from litellm.rust_bridge import configuration
@pytest.fixture(autouse=True)
def _isolated_configuration( # pyright: ignore[reportUnusedFunction] # pytest discovers fixtures dynamically
monkeypatch: pytest.MonkeyPatch,
) -> Generator[None]:
configuration.reset_rust_configuration()
monkeypatch.delenv("LITELLM_RUST", raising=False)
yield
configuration.reset_rust_configuration()
@pytest.mark.parametrize(
("process", "environment", "release_default", "expected"),
(
(False, True, True, False),
(True, False, False, True),
(None, False, True, False),
(None, True, False, True),
(None, None, False, False),
(None, None, True, True),
),
)
def test_resolution_precedence(
process: bool | None,
environment: bool | None,
release_default: bool,
expected: bool,
) -> None:
assert (
configuration.resolve_rust_enabled(
process_override=process,
environment_override=environment,
release_default=release_default,
)
is expected
)
def test_release_default_remains_disabled() -> None:
assert configuration.DEFAULT_RUST_ENABLED is False
assert configuration.rust_enabled() is False
assert configuration.rust_ocr_enabled() is True
@pytest.mark.parametrize("route", tuple(configuration.RouteName))
def test_each_route_has_an_explicit_release_default(route: configuration.RouteName) -> None:
assert configuration.rust_enabled(route) is (
route in {configuration.RouteName.OCR, configuration.RouteName.TRANSCRIPTION}
)
@pytest.mark.parametrize("enabled", (True, False))
def test_required_transcription_ignores_optional_rollout(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
configuration.rust(enabled)
assert configuration.rust_enabled(configuration.RouteName.TRANSCRIPTION) is True
assert configuration.rust_enabled(configuration.RouteName.MESSAGES) is enabled
assert configuration.rust_enabled(configuration.RouteName.OCR) is False
@pytest.mark.parametrize("process", [None, False, True])
@pytest.mark.parametrize("environment", [None, "0", "1", "off"])
def test_ocr_configuration(monkeypatch: pytest.MonkeyPatch, process: bool | None, environment: str | None) -> None:
if environment is not None:
monkeypatch.setenv("LITELLM_RUST", environment)
if process is not None:
configuration.rust(process)
assert configuration.rust_ocr_enabled() is (environment not in {"0", "off"} and process is not False)
def test_process_override_wins_over_environment(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
configuration.rust(True)
assert configuration.rust_enabled() is True
def test_global_environment_accepts_explicit_false(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "off")
assert configuration.rust_enabled() is False
@pytest.mark.parametrize("value", ("", " ", "sometimes", "2"))
def test_invalid_environment_value_disables_rust(monkeypatch: pytest.MonkeyPatch, value: str) -> None:
monkeypatch.setenv("LITELLM_RUST", value)
assert configuration.rust_enabled() is False
def test_process_override_and_reset_apply_to_existing_threads(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "1")
with ThreadPoolExecutor(max_workers=1) as executor:
assert executor.submit(configuration.rust_enabled).result() is True
configuration.rust(False)
assert executor.submit(configuration.rust_enabled).result() is False
configuration.reset_rust_configuration()
assert executor.submit(configuration.rust_enabled).result() is True
def test_explicit_override_precedes_invalid_environment(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "sometimes")
configuration.rust(True)
assert configuration.rust_enabled() is True
@pytest.mark.parametrize(("value", "expected"), (("1", "True"), ("0", "False")))
def test_environment_controls_startup(value: str, expected: str) -> None:
environment: Final = {**os.environ, "LITELLM_RUST": value}
result: Final = subprocess.run(
(
sys.executable,
"-c",
"from litellm.rust_bridge.configuration import rust_enabled; print(rust_enabled())",
),
check=True,
capture_output=True,
text=True,
env=environment,
)
assert result.stdout.strip() == expected

View file

@ -1,27 +1,22 @@
from __future__ import annotations
import pytest
from pydantic import ValidationError
from litellm.rust_bridge.configuration import (
_parse_env_bool, # pyright: ignore[reportPrivateUsage] # directly test env parsing contract
)
@pytest.mark.parametrize("value", ("1", "true", "t", "yes", "y", "on", "TRUE", " yes "))
def test_parse_env_bool_accepts_standard_true_values(value: str) -> None:
assert _parse_env_bool(value) is True
@pytest.mark.parametrize("value", ("0", "false", "f", "no", "n", "off", "FALSE", " no "))
def test_parse_env_bool_accepts_standard_false_values(value: str) -> None:
assert _parse_env_bool(value) is False
@pytest.mark.parametrize(("value", "expected"), (("1", True), ("0", False), (" 1 ", True), (" 0 ", False)))
def test_parse_env_bool_accepts_binary_values(value: str, expected: bool) -> None:
assert _parse_env_bool(value) is expected
def test_parse_env_bool_preserves_unset_value() -> None:
assert _parse_env_bool(None) is None
def test_parse_env_bool_rejects_unknown_value() -> None:
with pytest.raises(ValidationError):
_parse_env_bool("enabled")
@pytest.mark.parametrize("value", ("enabled", "true", "false", "yes", "no", "on", "off", ""))
def test_parse_env_bool_rejects_unknown_value(value: str) -> None:
with pytest.raises(ValueError, match="must be '1' or '0'"):
_parse_env_bool(value)

View file

@ -1,4 +1,5 @@
from collections.abc import Generator, Mapping
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, Mock
@ -51,7 +52,7 @@ def test_admitted_failure_is_returned_without_replay() -> None:
assert caught.value is failure
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
configuration.reset_rust_configuration()
assert native.call_count == 1
@ -64,6 +65,7 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_
args: tuple[object, ...],
kwargs: Mapping[str, object],
asynchronous: bool,
host: object,
) -> OCRResponse:
captured.append((request, args, kwargs, asynchronous))
return OCRResponse(pages=[], model=request.model)
@ -74,7 +76,7 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_
response: Final = litellm.ocr("mistral/mistral-ocr-latest", document)
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
configuration.reset_rust_configuration()
request, call_args, hook_kwargs, asynchronous = captured[0]
assert response.model == "mistral/mistral-ocr-latest"
@ -94,6 +96,7 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs()
args: tuple[object, ...],
kwargs: Mapping[str, object],
asynchronous: bool,
host: object,
) -> OCRResponse:
assert args == ()
captured.append(kwargs)
@ -105,7 +108,7 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs()
litellm.ocr(model="mistral/mistral-ocr-latest", document=document)
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
configuration.reset_rust_configuration()
assert captured[0]["model"] == "mistral/mistral-ocr-latest"
assert captured[0]["document"] is document
@ -123,7 +126,7 @@ def test_public_duplicate_argument_error_does_not_depend_on_native_selection(ena
litellm.ocr("mistral/mistral-ocr-latest", document, model="duplicate")
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
configuration.reset_rust_configuration()
assert native.call_count == 0
@ -137,14 +140,14 @@ def test_public_missing_required_argument_error_does_not_depend_on_native_select
litellm.ocr("mistral/mistral-ocr-latest")
finally:
NATIVE_OCR_LIFECYCLE.reset()
litellm.rust(None)
configuration.reset_rust_configuration()
assert native.call_count == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("enabled", [False, True, None])
async def test_environment_opt_out_never_loads_native(
@pytest.mark.parametrize("enabled", [False, None])
async def test_environment_opt_out_skips_native_without_process_enable(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, enabled: bool | None
) -> None:
monkeypatch.setenv("LITELLM_RUST", "0")
@ -153,7 +156,8 @@ async def test_environment_opt_out_never_loads_native(
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)
load: Final = Mock(side_effect=AssertionError("native must not be loaded"))
monkeypatch.setattr(bindings, "get_native_bridge", load)
litellm.rust(enabled)
if enabled is not None:
litellm.rust(enabled)
document: Final = {"type": "file", "file": b"pdf"}
result: Final = (
@ -196,6 +200,10 @@ class Declined(Exception):
pass
class Upstream(Exception):
pass
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("declined", [False, True])
@ -205,10 +213,11 @@ async def test_only_native_declines_replay_on_legacy(
failure: Final = Declined("unsupported") if declined else RuntimeError("provider already called")
native: Final = AsyncMock(side_effect=failure) if asynchronous else Mock(side_effect=failure)
NATIVE_OCR_LIFECYCLE.override(native)
import importlib
main: Final = importlib.import_module("litellm.ocr.main")
monkeypatch.setattr(main, "native_exception_types", lambda: (Declined, RuntimeError))
monkeypatch.setattr(
bindings,
"get_native_bridge",
lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream),
)
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
fallback: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
monkeypatch.setattr(legacy, "aocr" if asynchronous else "ocr", fallback)

View file

@ -1,5 +1,6 @@
from __future__ import annotations
from collections.abc import Iterator
from types import ModuleType
from typing import Final
@ -7,8 +8,25 @@ import pytest
from litellm.rust_bridge import configuration
from litellm.rust_bridge.bindings import BINDING_UNSET
from litellm.rust_bridge.configuration import RouteName
from litellm.rust_bridge.route import NativeRoute
from litellm.rust_bridge.catalog import COMPONENTS
from litellm.rust_bridge.configuration import (
CapabilityContext,
CapabilityDefinition,
DeliveryMode,
ExecutionDecision,
RolloutPolicy,
RouteName,
RustImplementationState,
)
from litellm.rust_bridge.errors import RustRouteUnavailableError, RustRouteUnsupportedError
from litellm.rust_bridge.route import NativeComponent
@pytest.fixture(autouse=True)
def reset_configuration() -> Iterator[None]:
configuration.reset_rust_configuration()
yield
configuration.reset_rust_configuration()
def _string(value: object) -> str | None:
@ -16,22 +34,30 @@ def _string(value: object) -> str | None:
def _unexpected_load() -> ModuleType:
raise AssertionError("disabled route loaded the extension")
raise AssertionError("Python selection loaded the native extension")
def test_disabled_route_does_not_load_native(monkeypatch: pytest.MonkeyPatch) -> None:
def test_python_decision_does_not_load_native(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("LITELLM_RUST", raising=False)
configuration.reset_rust_configuration()
route: Final = NativeRoute(RouteName.MESSAGES)
binding: Final = route.bind("messages", validate=_string, module_loader=_unexpected_load)
assert route.select(binding) is None
component: Final = COMPONENTS[RouteName.MESSAGES]
binding: Final = component.bind("messages", validate=_string, module_loader=_unexpected_load)
assert component.resolve().select(binding) is None
def test_delivery_mode_uses_one_component_policy(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "1")
component: Final = COMPONENTS[RouteName.MESSAGES]
completed: Final = component.resolve(CapabilityContext(delivery=DeliveryMode.COMPLETED))
streaming: Final = component.resolve(CapabilityContext(delivery=DeliveryMode.STREAMING))
assert completed.decision is ExecutionDecision.RUST_WITH_FALLBACK
assert streaming.decision is ExecutionDecision.PYTHON
def test_binding_discovery_validation_and_override() -> None:
module: Final = ModuleType("fake_native")
setattr(module, "messages", "native")
route: Final = NativeRoute(RouteName.MESSAGES)
binding: Final = route.bind("messages", validate=_string, module_loader=lambda: module)
component: Final = COMPONENTS[RouteName.MESSAGES]
binding: Final = component.bind("messages", validate=_string, module_loader=lambda: module)
assert binding.load() == "native"
binding.configure("override")
binding.configure(BINDING_UNSET)
@ -44,17 +70,94 @@ def test_binding_discovery_validation_and_override() -> None:
assert binding.load() is None
def test_missing_native_is_unavailable() -> None:
route: Final = NativeRoute(RouteName.OCR)
binding: Final = route.bind("ocr", validate=_string, module_loader=lambda: None)
assert binding.load() is None
def test_component_rejects_undeclared_binding() -> None:
with pytest.raises(ValueError, match="not declared"):
COMPONENTS[RouteName.MESSAGES].bind("typo", validate=_string)
def test_explicit_enablement_loads_optional_route(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_RUST", "1")
configuration.reset_rust_configuration()
module: Final = ModuleType("fake_native")
setattr(module, "messages", "native")
route: Final = NativeRoute(RouteName.MESSAGES)
binding: Final = route.bind("messages", validate=_string, module_loader=lambda: module)
assert route.select(binding) == "native"
class ModelCapability:
def __init__(self) -> None:
self.calls: int = 0
def __call__(self, context: CapabilityContext) -> CapabilityDefinition:
self.calls += 1
if context.provider == "provider" and context.model == "new-model":
return CapabilityDefinition(
rust=RustImplementationState.EXPERIMENTAL,
python_available=False,
rollout=RolloutPolicy.RUST_REQUIRED,
)
return CapabilityDefinition(
rust=RustImplementationState.READY,
python_available=True,
rollout=RolloutPolicy.RUST_OPT_IN,
)
def test_dynamic_capability_resolves_once_before_binding_selection() -> None:
resolver: Final = ModelCapability()
component: Final = NativeComponent(
name=RouteName.TRANSCRIPTION,
capability=resolver,
exports=("transcription",),
)
binding: Final = component.bind("transcription", validate=_string, module_loader=_unexpected_load)
binding.override("native")
execution: Final = component.resolve(CapabilityContext(provider="provider", model="new-model"))
assert execution.select(binding) == "native"
assert resolver.calls == 1
@pytest.mark.parametrize("process_override", (None, False, True))
@pytest.mark.parametrize("environment_override", ("0", "1", "invalid"))
def test_bedrock_transcription_requires_rust_regardless_of_overrides(
monkeypatch: pytest.MonkeyPatch, process_override: bool | None, environment_override: str
) -> None:
monkeypatch.setenv("LITELLM_RUST", environment_override)
if process_override is not None:
configuration.rust(process_override)
component: Final = COMPONENTS[RouteName.TRANSCRIPTION]
binding: Final = component.bind("transcription", validate=_string, module_loader=lambda: None)
execution: Final = component.resolve(CapabilityContext(provider="bedrock", model="model"))
assert execution.decision is ExecutionDecision.RUST_REQUIRED
with pytest.raises(RustRouteUnavailableError, match="transcription bridge is unavailable"):
execution.select(binding)
@pytest.mark.parametrize("provider", ("openai", "azure", "azure_ai", "groq", "mistral", "nvidia_riva", "soniox"))
def test_python_transcription_providers_skip_native_discovery(provider: str) -> None:
configuration.rust(True)
component: Final = COMPONENTS[RouteName.TRANSCRIPTION]
binding: Final = component.bind("transcription", validate=_string, module_loader=_unexpected_load)
execution: Final = component.resolve(CapabilityContext(provider=provider, model="model"))
assert execution.decision is ExecutionDecision.PYTHON
assert execution.select(binding) is None
@pytest.mark.parametrize("provider", ("unknown-provider", "anthropic"))
def test_unsupported_transcription_never_selects_an_implementation(provider: str) -> None:
component: Final = COMPONENTS[RouteName.TRANSCRIPTION]
binding: Final = component.bind("transcription", validate=_string, module_loader=_unexpected_load)
execution: Final = component.resolve(CapabilityContext(provider=provider, model="model"))
assert execution.decision is ExecutionDecision.UNSUPPORTED
with pytest.raises(RustRouteUnsupportedError, match="No Python or Rust implementation for transcription"):
execution.select(binding)
@pytest.mark.parametrize(
("rust", "python_available", "rollout"),
(
(RustImplementationState.READY, True, RolloutPolicy.RUST_REQUIRED),
(RustImplementationState.EXPERIMENTAL, False, RolloutPolicy.RUST_OPT_IN),
(RustImplementationState.READY, False, RolloutPolicy.RUST_OPT_OUT),
(RustImplementationState.UNIMPLEMENTED, False, RolloutPolicy.PYTHON_ONLY),
(RustImplementationState.UNIMPLEMENTED, False, RolloutPolicy.RUST_REQUIRED),
(RustImplementationState.UNIMPLEMENTED, True, RolloutPolicy.UNSUPPORTED),
(RustImplementationState.READY, False, RolloutPolicy.UNSUPPORTED),
),
)
def test_invalid_capability_definitions_rejected(
rust: RustImplementationState, python_available: bool, rollout: RolloutPolicy
) -> None:
with pytest.raises(ValueError, match="capability"):
CapabilityDefinition(rust=rust, python_available=python_available, rollout=rollout)

View file

@ -1,95 +1,191 @@
from __future__ import annotations
import asyncio
from collections.abc import Callable
from types import SimpleNamespace
from typing import Final
import pytest
from litellm.exceptions import APIError
from litellm.rust_bridge import bindings, runtime
from litellm.rust_bridge.configuration import ExecutionDecision, RouteName
from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError, RustRouteUnsupportedError
from litellm.rust_bridge.route import ComponentExecution
class RustBridgeDeclined(Exception):
pass
class RustBridgeUnavailable(Exception):
pass
class RustUpstreamError(Exception):
pass
class RustHostCallbackError(Exception):
pass
@pytest.fixture(autouse=True)
def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None:
native = SimpleNamespace(
native: Final = SimpleNamespace(
RustBridgeDeclined=RustBridgeDeclined,
RustBridgeUnavailable=RustBridgeUnavailable,
RustHostCallbackError=RustHostCallbackError,
RustUpstreamError=RustUpstreamError,
)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
def context() -> runtime.BridgeErrorContext:
return runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model")
async def _invoke(
asynchronous: bool,
decision: ExecutionDecision,
native_call: Callable[[], object] | None,
adapt: Callable[[object], object],
fallback: Callable[[], object] = lambda: "python",
) -> object:
execution: Final = ComponentExecution(route_name=RouteName.MESSAGES, decision=decision)
context: Final = runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model")
if not asynchronous:
return runtime.invoke(
execution=execution,
native_call=native_call,
python_fallback=fallback,
adapt=adapt,
context=context,
)
async def call() -> object:
assert native_call is not None
return native_call()
def test_invoke_tags_native_decline_before_running_fallback() -> None:
calls: list[str] = []
async def afallback() -> object:
return fallback()
def decline() -> object:
calls.append("rust")
raise RustBridgeDeclined("unsupported")
value = runtime.invoke(
native_call=decline,
fallback=lambda: calls.append("python") or "fallback",
adapt=str,
mode=runtime.FallbackMode.PYTHON,
context=context(),
return await runtime.ainvoke(
execution=execution,
native_call=call if native_call is not None else None,
python_fallback=afallback,
adapt=adapt,
context=context,
)
assert value == "fallback"
assert calls == ["rust", "python"]
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_python_selection_runs_only_the_fallback(asynchronous: bool) -> None:
calls: Final[list[str]] = []
def native() -> str:
calls.append("native")
return "native"
result: Final = await _invoke(
asynchronous,
ExecutionDecision.PYTHON,
native,
str,
lambda: calls.append("python") or "python",
)
assert result == "python"
assert calls == ["python"]
def test_invoke_translates_upstream_without_fallback() -> None:
def fail() -> object:
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("error", (None, RustBridgeDeclined("unsupported"), RustBridgeUnavailable()))
@pytest.mark.parametrize("required", (False, True))
async def test_unavailable_and_declined_follow_policy(
asynchronous: bool, error: Exception | None, required: bool
) -> None:
def fail() -> str:
assert error is not None
raise error
decision: Final = ExecutionDecision.RUST_REQUIRED if required else ExecutionDecision.RUST_WITH_FALLBACK
native_call: Final = fail if error is not None else None
if required:
expected: Final = RustRouteDeclinedError if isinstance(error, RustBridgeDeclined) else RustRouteUnavailableError
with pytest.raises(expected):
await _invoke(asynchronous, decision, native_call, str)
return
assert await _invoke(asynchronous, decision, native_call, str) == "python"
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("decision", (ExecutionDecision.RUST_WITH_FALLBACK, ExecutionDecision.RUST_REQUIRED))
@pytest.mark.parametrize("value", (None, False, 0, "native"))
async def test_native_success_does_not_run_fallback(
asynchronous: bool, decision: ExecutionDecision, value: None | bool | int | str
) -> None:
calls: Final[list[str]] = []
result: Final = await _invoke(
asynchronous,
decision,
lambda: value,
lambda native: native,
lambda: calls.append("python") or "python",
)
assert result is value
assert calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_upstream_failure_never_falls_back(asynchronous: bool) -> None:
def fail() -> str:
raise RustUpstreamError(429, "rate limited")
with pytest.raises(APIError, match="rate limited") as caught:
runtime.invoke(
native_call=fail,
fallback=lambda: pytest.fail("fallback must not run"),
adapt=str,
mode=runtime.FallbackMode.PYTHON,
context=context(),
)
await _invoke(asynchronous, ExecutionDecision.RUST_WITH_FALLBACK, fail, str)
assert caught.value.status_code == 429
@pytest.mark.asyncio
async def test_ainvoke_handles_native_success() -> None:
async def native() -> int:
return 3
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("error", (RuntimeError("execution failed"), asyncio.CancelledError()))
async def test_execution_failure_never_falls_back(asynchronous: bool, error: BaseException) -> None:
def fail() -> str:
raise error
async def fallback() -> str:
pytest.fail("fallback must not run")
assert (
await runtime.ainvoke(
native_call=native,
fallback=fallback,
adapt=str,
mode=runtime.FallbackMode.PYTHON,
context=context(),
)
== "3"
)
with pytest.raises(type(error)) as caught:
await _invoke(asynchronous, ExecutionDecision.RUST_WITH_FALLBACK, fail, str)
assert caught.value is error
def test_required_mode_rejects_unavailable_bridge() -> None:
with pytest.raises(RuntimeError, match="is unavailable"):
runtime.invoke(
native_call=None,
fallback=lambda: pytest.fail("fallback must not run"),
adapt=str,
mode=runtime.FallbackMode.RUST_REQUIRED,
context=context(),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("error", (RustBridgeDeclined("adapt"), RustBridgeUnavailable(), RuntimeError("adapt")))
async def test_adaptation_failure_never_falls_back(asynchronous: bool, error: Exception) -> None:
def adapt(_value: str) -> str:
raise error
with pytest.raises(type(error)) as caught:
await _invoke(asynchronous, ExecutionDecision.RUST_WITH_FALLBACK, lambda: "native", adapt)
assert caught.value is error
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_host_callback_failure_preserves_its_cause(asynchronous: bool) -> None:
callback_error: Final = RustBridgeDeclined("callback raised native decline type")
def fail() -> str:
wrapped: Final = RustHostCallbackError("callback failed")
wrapped.__cause__ = callback_error
raise wrapped
with pytest.raises(RustBridgeDeclined) as caught:
await _invoke(asynchronous, ExecutionDecision.RUST_WITH_FALLBACK, fail, str)
assert caught.value is callback_error
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_unsupported_execution_runs_nothing(asynchronous: bool) -> None:
with pytest.raises(RustRouteUnsupportedError):
await _invoke(asynchronous, ExecutionDecision.UNSUPPORTED, lambda: "native", str)

Some files were not shown because too many files have changed in this diff Show more