From 6276022f34262e57e7a4fe40f08a9e449cbbc0da Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 14 Sep 2026 16:52:42 -0700 Subject: [PATCH] wip --- Makefile | 3 + .../core/src/audio_transcription/mod.rs | 20 + .../core/src/audio_transcription/prepare.rs | 4 +- .../core/src/call_lifecycle/admission.rs | 10 + .../crates/core/src/chat_completions/mod.rs | 50 ++ litellm-rust/crates/core/src/messages/mod.rs | 21 + litellm-rust/crates/core/src/ocr/mod.rs | 9 + .../crates/core/src/responses/websocket.rs | 9 + .../crates/python-bridge/src/errors.rs | 37 +- litellm-rust/crates/python-bridge/src/lib.rs | 9 +- .../src/routes/chat_completions/value.rs | 49 +- .../python-bridge/src/routes/definition.rs | 49 +- .../src/routes/messages/value.rs | 11 +- .../crates/python-bridge/src/routes/mod.rs | 2 + .../python-bridge/src/routes/ocr/callbacks.rs | 31 +- .../python-bridge/src/routes/ocr/lifecycle.rs | 10 +- .../python-bridge/src/routes/ocr/value.rs | 5 + .../src/routes/responses/websocket.rs | 19 +- .../src/routes/token_counter/mod.rs | 127 +++++ .../src/routes/transcription/value.rs | 10 +- .../crates/python-bridge/src/token_counter.rs | 105 ---- .../crates/token-counter/src/counter.rs | 11 + .../crates/token-counter/src/error.rs | 2 + litellm-rust/crates/token-counter/src/lib.rs | 17 + litellm/constants.py | 2 +- litellm/llms/anthropic/chat/handler.py | 113 ++-- .../bedrock/audio_transcription/__init__.py | 2 - litellm/llms/bedrock/chat/converse_handler.py | 130 ++--- litellm/llms/custom_httpx/llm_http_handler.py | 66 +-- litellm/ocr/input.py | 19 +- litellm/ocr/main.py | 57 +- litellm/rust_bridge/README.md | 36 +- litellm/rust_bridge/_native.pyi | 102 +++- litellm/rust_bridge/bindings.py | 12 + litellm/rust_bridge/catalog.py | 145 ++++++ .../rust_bridge/chat_completions/__init__.py | 12 +- .../chat_completions/definition.py | 6 + .../rust_bridge/chat_completions/lifecycle.py | 4 +- litellm/rust_bridge/chat_completions/types.py | 17 +- litellm/rust_bridge/chat_completions/value.py | 318 +++--------- litellm/rust_bridge/configuration.py | 175 +++++-- litellm/rust_bridge/embeddings/__init__.py | 5 +- litellm/rust_bridge/embeddings/definition.py | 6 + litellm/rust_bridge/embeddings/lifecycle.py | 4 +- litellm/rust_bridge/errors.py | 10 + litellm/rust_bridge/image_edit/__init__.py | 5 +- litellm/rust_bridge/image_edit/definition.py | 6 + litellm/rust_bridge/image_edit/lifecycle.py | 4 +- .../rust_bridge/image_generation/__init__.py | 5 +- .../image_generation/definition.py | 6 + .../rust_bridge/image_generation/lifecycle.py | 4 +- litellm/rust_bridge/messages/__init__.py | 4 +- litellm/rust_bridge/messages/definition.py | 6 + litellm/rust_bridge/messages/lifecycle.py | 4 +- litellm/rust_bridge/messages/types.py | 12 +- litellm/rust_bridge/messages/value.py | 120 +++-- litellm/rust_bridge/moderation/__init__.py | 5 +- litellm/rust_bridge/moderation/definition.py | 6 + litellm/rust_bridge/moderation/lifecycle.py | 4 +- litellm/rust_bridge/ocr/__init__.py | 4 +- litellm/rust_bridge/ocr/definition.py | 6 + litellm/rust_bridge/ocr/host.py | 52 ++ litellm/rust_bridge/ocr/lifecycle.py | 47 +- litellm/rust_bridge/ocr/value.py | 86 ++-- litellm/rust_bridge/rerank/__init__.py | 5 +- litellm/rust_bridge/rerank/definition.py | 6 + litellm/rust_bridge/rerank/lifecycle.py | 4 +- litellm/rust_bridge/responses/__init__.py | 5 +- litellm/rust_bridge/responses/definition.py | 6 + litellm/rust_bridge/responses/lifecycle.py | 4 +- litellm/rust_bridge/responses/websocket.py | 53 +- litellm/rust_bridge/route.py | 65 ++- litellm/rust_bridge/runtime.py | 186 +++---- litellm/rust_bridge/speech/__init__.py | 5 +- litellm/rust_bridge/speech/definition.py | 6 + litellm/rust_bridge/speech/lifecycle.py | 4 +- litellm/rust_bridge/token_counter.py | 114 ---- litellm/rust_bridge/token_counter/__init__.py | 12 + .../rust_bridge/token_counter/definition.py | 6 + litellm/rust_bridge/token_counter/types.py | 31 ++ litellm/rust_bridge/token_counter/value.py | 75 +++ litellm/rust_bridge/transcription/__init__.py | 4 +- .../rust_bridge/transcription/definition.py | 6 + .../rust_bridge/transcription/lifecycle.py | 4 +- litellm/rust_bridge/transcription/value.py | 86 ++-- pyproject.toml | 2 + .../strategies/trace_parity/sdk/ocr/case.py | 4 +- .../test_rust_bridge_messages.py | 30 +- .../chat/test_anthropic_chat_handler.py | 485 +++++------------- .../chat/test_bedrock_converse_handler.py | 117 ++--- .../custom_httpx/test_llm_http_handler.py | 124 +++-- tests/test_litellm/ocr/test_legacy.py | 9 +- .../responses/test_rust_bridge_websocket.py | 62 ++- tests/test_litellm/rust_bridge/stubtest.ini | 2 + .../rust_bridge/test_chat_completions.py | 424 ++++----------- .../rust_bridge/test_configuration.py | 138 ----- .../rust_bridge/test_configuration_env.py | 19 +- .../rust_bridge/test_ocr_lifecycle.py | 33 +- tests/test_litellm/rust_bridge/test_route.py | 147 +++++- .../test_litellm/rust_bridge/test_runtime.py | 202 ++++++-- .../rust_bridge/test_token_counter.py | 437 +++------------- .../test_audio_transcription_rust_bridge.py | 138 ++++- tests/test_litellm_rust/ocr/test_lifecycle.py | 4 +- .../test_route_foundation.py | 146 ++++-- uv.lock | 170 +++++- 105 files changed, 2955 insertions(+), 2672 deletions(-) create mode 100644 litellm-rust/crates/python-bridge/src/routes/token_counter/mod.rs delete mode 100644 litellm-rust/crates/python-bridge/src/token_counter.rs create mode 100644 litellm/rust_bridge/catalog.py create mode 100644 litellm/rust_bridge/chat_completions/definition.py create mode 100644 litellm/rust_bridge/embeddings/definition.py create mode 100644 litellm/rust_bridge/errors.py create mode 100644 litellm/rust_bridge/image_edit/definition.py create mode 100644 litellm/rust_bridge/image_generation/definition.py create mode 100644 litellm/rust_bridge/messages/definition.py create mode 100644 litellm/rust_bridge/moderation/definition.py create mode 100644 litellm/rust_bridge/ocr/definition.py create mode 100644 litellm/rust_bridge/ocr/host.py create mode 100644 litellm/rust_bridge/rerank/definition.py create mode 100644 litellm/rust_bridge/responses/definition.py create mode 100644 litellm/rust_bridge/speech/definition.py delete mode 100644 litellm/rust_bridge/token_counter.py create mode 100644 litellm/rust_bridge/token_counter/__init__.py create mode 100644 litellm/rust_bridge/token_counter/definition.py create mode 100644 litellm/rust_bridge/token_counter/types.py create mode 100644 litellm/rust_bridge/token_counter/value.py create mode 100644 litellm/rust_bridge/transcription/definition.py create mode 100644 tests/test_litellm/rust_bridge/stubtest.ini delete mode 100644 tests/test_litellm/rust_bridge/test_configuration.py diff --git a/Makefile b/Makefile index d360074ea4e..0e9d2bbf82c 100644 --- a/Makefile +++ b/Makefile @@ -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 diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index 31b6de4b3e4..795d0a72f56 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -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; diff --git a/litellm-rust/crates/core/src/audio_transcription/prepare.rs b/litellm-rust/crates/core/src/audio_transcription/prepare.rs index bbef97341a9..91489ef01aa 100644 --- a/litellm-rust/crates/core/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/core/src/audio_transcription/prepare.rs @@ -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); diff --git a/litellm-rust/crates/core/src/call_lifecycle/admission.rs b/litellm-rust/crates/core/src/call_lifecycle/admission.rs index 022600758bd..c4d6807484e 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/admission.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/admission.rs @@ -18,3 +18,13 @@ pub enum UnimplementedRoute { pub fn admit_unimplemented(route: UnimplementedRoute) -> Result { 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), +} diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 32dea17d202..7441f09d367 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -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, + headers: Option<&Map>, + 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; diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index cfa8bda1104..4af86ad43bb 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -27,5 +27,26 @@ pub async fn messages_stream(request: MessagesRequest<'_>) -> Result, + 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; diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index e29fd6ac572..985844f47b0 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -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; diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs index 34213e5f6c4..bf45742a1fb 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -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), + } +} diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index e625b720044..216838dde6c 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -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::(), + )?; module.add("RustBridgeDeclined", py.get_type::())?; + module.add( + "RustHostCallbackError", + py.get_type::(), + )?; module.add("RustUpstreamError", py.get_type::()) } +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())) +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 3bd243a3b3f..90ba5b65b6f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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" diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/value.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/value.rs index 867a879ab91..c40935da576 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/value.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/value.rs @@ -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> + 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, - custom_llm_provider: Option, -) -> PyResult> { - 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, timeout_seconds: Option, + #[pyo3(from_py_with = litellm_python_interop::from_py)] + host_facts: Option, + on_request: Option>, }, prepare = prepare_chat_completions, errors = execution_error_to_pyerr, - extra = [chat_completions_decline], } diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index d1b06fd0924..2a8e4d62668 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -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> { 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 = 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 = 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 = 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(); diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/value.rs b/litellm-rust/crates/python-bridge/src/routes/messages/value.rs index b741e54f0ca..a436905d171 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/value.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/value.rs @@ -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, timeout_seconds: Option, + has_agentic_hook: Option, }, prepare = prepare_messages, - errors = core_error_to_pyerr, + errors = execution_error_to_pyerr, } diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 2c53ff32af9..3effd8c3f8a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -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")] { diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs index 2ea853cd64a..034c7b1d681 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -29,6 +29,7 @@ impl PythonLogger { pub(super) fn update_ocr( &self, py: Python<'_>, + host: &Py, kwargs: &Py, 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> { - py.import("litellm.types.utils")? - .getattr("CustomPricingLiteLLMParams")? - .getattr("model_fields")? - .cast_into::()? - .keys() - .iter() - .map(|name| name.extract::()) - .collect() +fn custom_pricing_fields(py: Python<'_>, host: &Py) -> PyResult> { + 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.import("litellm.rust_bridge.ocr.value")? - .getattr("_response")? +pub(super) fn response( + py: Python<'_>, + host: &Py, + response: &LiteLLMOcrResponse, +) -> PyResult> { + host.bind(py) + .getattr("response")? .call1((to_py(py, response)?,)) .map(Bound::unbind) } pub(super) fn map_failure( py: Python<'_>, + host: &Py, error: &Py, request: &Bound<'_, PyAny>, provider: &str, ) -> PyResult> { - Ok(py - .import("litellm.rust_bridge.ocr.lifecycle")? + Ok(host + .bind(py) .getattr("map_failure")? .call1((error, request, provider))? .extract()?) diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs index e655290d8fb..f5ede86a64b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs @@ -17,6 +17,7 @@ use crate::lifecycle::{ struct PythonOcrHost { state: PythonCallState, + adapter: Py, 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> { 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(), }, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs index 051ac19d4fb..713093128ba 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs @@ -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, diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs b/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs index 6fe8fe858f0..abe963c5960 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/websocket.rs @@ -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, timeout_seconds: Option, + custom_llm_provider: Option, ) -> PyResult> { + 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> { 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> { 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> { 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) }) } } diff --git a/litellm-rust/crates/python-bridge/src/routes/token_counter/mod.rs b/litellm-rust/crates/python-bridge/src/routes/token_counter/mod.rs new file mode 100644 index 00000000000..2a5808fecee --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/token_counter/mod.rs @@ -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, + encode_slots: Arc, +} + +static COUNTERS: OnceLock>>> = 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, +) -> PyResult> { + 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, +) -> PyResult> { + 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 { + 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 { + 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)?) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/transcription/value.rs b/litellm-rust/crates/python-bridge/src/routes/transcription/value.rs index af60515b0e2..20d1dcbdd58 100644 --- a/litellm-rust/crates/python-bridge/src/routes/transcription/value.rs +++ b/litellm-rust/crates/python-bridge/src/routes/transcription/value.rs @@ -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, }, prepare = prepare_transcription, - errors = core_error_to_pyerr, + errors = execution_error_to_pyerr, } diff --git a/litellm-rust/crates/python-bridge/src/token_counter.rs b/litellm-rust/crates/python-bridge/src/token_counter.rs deleted file mode 100644 index b4de50c5f1a..00000000000 --- a/litellm-rust/crates/python-bridge/src/token_counter.rs +++ /dev/null @@ -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, - encode_slots: Arc, -} - -#[pymethods] -impl TokenCounter { - #[new] - fn new(py: Python<'_>, tokenizer_json: &str) -> PyResult { - Self::load(py, || CoreTokenCounter::from_json(tokenizer_json)) - } - - #[staticmethod] - fn from_cl100k_ranks(py: Python<'_>, rank_file: &str) -> PyResult { - Self::load(py, || CoreTokenCounter::from_cl100k_ranks(rank_file)) - } - - #[staticmethod] - fn from_o200k_ranks(py: Python<'_>, rank_file: &str) -> PyResult { - Self::load(py, || CoreTokenCounter::from_o200k_ranks(rank_file)) - } - - fn acount_request<'py>(&self, py: Python<'py>, body: &[u8]) -> PyResult> { - 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 + Send, - ) -> PyResult { - 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 { - 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::() -} diff --git a/litellm-rust/crates/token-counter/src/counter.rs b/litellm-rust/crates/token-counter/src/counter.rs index 7eedc449dd1..10b3eabb4ef 100644 --- a/litellm-rust/crates/token-counter/src/counter.rs +++ b/litellm-rust/crates/token-counter/src/counter.rs @@ -25,6 +25,7 @@ pub struct InputTokenCount { } enum Encoder { + Admission, HuggingFace { tokenizer: Box, byte_level: Option, @@ -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 { let tokenizer = tokenizer_json @@ -74,6 +84,7 @@ impl TokenCounter { pub fn count_text(&self, text: &str) -> Result { match &self.encoder { + Encoder::Admission => Ok(0), Encoder::Tiktoken(counter) => Ok(counter.count(text)), Encoder::HuggingFace { tokenizer, diff --git a/litellm-rust/crates/token-counter/src/error.rs b/litellm-rust/crates/token-counter/src/error.rs index 6b8668fe182..fa91b8b79d0 100644 --- a/litellm-rust/crates/token-counter/src/error.rs +++ b/litellm-rust/crates/token-counter/src/error.rs @@ -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}")] diff --git a/litellm-rust/crates/token-counter/src/lib.rs b/litellm-rust/crates/token-counter/src/lib.rs index fa0014e2bad..c03090aeb7d 100644 --- a/litellm-rust/crates/token-counter/src/lib.rs +++ b/litellm-rust/crates/token-counter/src/lib.rs @@ -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), + } +} diff --git a/litellm/constants.py b/litellm/constants.py index c106be688e4..e8dcc634b4d 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -911,7 +911,7 @@ openai_compatible_endpoints: Final[list] = [ ] -openai_compatible_providers: Final[list] = [ +openai_compatible_providers: Final[list[str]] = [ "anyscale", "groq", "nvidia_nim", diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 5d3ae444b42..b9e3a6800ab 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -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() diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index b1f8c957ff4..90297b3b492 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -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, diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index d397420cb17..4928a29306c 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -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: diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 4e88c69bde6..ee468535c2a 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/litellm/ocr/input.py b/litellm/ocr/input.py index bcb448371c4..0844bf7d4ee 100644 --- a/litellm/ocr/input.py +++ b/litellm/ocr/input.py @@ -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 diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 60d0ff7fb3b..2548bf10a14 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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, + ), + ) diff --git a/litellm/rust_bridge/README.md b/litellm/rust_bridge/README.md index 693af7ad46c..dab88fa55c3 100644 --- a/litellm/rust_bridge/README.md +++ b/litellm/rust_bridge/README.md @@ -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/` diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 7ab3b16b31f..7912d7c1a11 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -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", +] diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index 0d9f396c087..caea6f1b788 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -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 () diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py new file mode 100644 index 00000000000..706d8e28b51 --- /dev/null +++ b/litellm/rust_bridge/catalog.py @@ -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) diff --git a/litellm/rust_bridge/chat_completions/__init__.py b/litellm/rust_bridge/chat_completions/__init__.py index eb5f7d80b57..e0488760305 100644 --- a/litellm/rust_bridge/chat_completions/__init__.py +++ b/litellm/rust_bridge/chat_completions/__init__.py @@ -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", ) diff --git a/litellm/rust_bridge/chat_completions/definition.py b/litellm/rust_bridge/chat_completions/definition.py new file mode 100644 index 00000000000..8a26cfcfc55 --- /dev/null +++ b/litellm/rust_bridge/chat_completions/definition.py @@ -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] diff --git a/litellm/rust_bridge/chat_completions/lifecycle.py b/litellm/rust_bridge/chat_completions/lifecycle.py index 3e7c7e3fc44..dd2f115658c 100644 --- a/litellm/rust_bridge/chat_completions/lifecycle.py +++ b/litellm/rust_bridge/chat_completions/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/chat_completions/types.py b/litellm/rust_bridge/chat_completions/types.py index 5e1c457dff3..d1481346da6 100644 --- a/litellm/rust_bridge/chat_completions/types.py +++ b/litellm/rust_bridge/chat_completions/types.py @@ -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. diff --git a/litellm/rust_bridge/chat_completions/value.py b/litellm/rust_bridge/chat_completions/value.py index aedca760811..710abb93353 100644 --- a/litellm/rust_bridge/chat_completions/value.py +++ b/litellm/rust_bridge/chat_completions/value.py @@ -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() diff --git a/litellm/rust_bridge/configuration.py b/litellm/rust_bridge/configuration.py index 065431a51a4..dafadef0bb0 100644 --- a/litellm/rust_bridge/configuration.py +++ b/litellm/rust_bridge/configuration.py @@ -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 diff --git a/litellm/rust_bridge/embeddings/__init__.py b/litellm/rust_bridge/embeddings/__init__.py index ff6689bc26d..95ac64fe443 100644 --- a/litellm/rust_bridge/embeddings/__init__.py +++ b/litellm/rust_bridge/embeddings/__init__.py @@ -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",) diff --git a/litellm/rust_bridge/embeddings/definition.py b/litellm/rust_bridge/embeddings/definition.py new file mode 100644 index 00000000000..889bee08fae --- /dev/null +++ b/litellm/rust_bridge/embeddings/definition.py @@ -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] diff --git a/litellm/rust_bridge/embeddings/lifecycle.py b/litellm/rust_bridge/embeddings/lifecycle.py index 51a8372b048..bcdf8a6d070 100644 --- a/litellm/rust_bridge/embeddings/lifecycle.py +++ b/litellm/rust_bridge/embeddings/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/errors.py b/litellm/rust_bridge/errors.py new file mode 100644 index 00000000000..5afbc52fc8a --- /dev/null +++ b/litellm/rust_bridge/errors.py @@ -0,0 +1,10 @@ +class RustRouteUnavailableError(RuntimeError): + pass + + +class RustRouteDeclinedError(RuntimeError): + pass + + +class RustRouteUnsupportedError(NotImplementedError): + pass diff --git a/litellm/rust_bridge/image_edit/__init__.py b/litellm/rust_bridge/image_edit/__init__.py index 443c310e740..a827aeb674e 100644 --- a/litellm/rust_bridge/image_edit/__init__.py +++ b/litellm/rust_bridge/image_edit/__init__.py @@ -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",) diff --git a/litellm/rust_bridge/image_edit/definition.py b/litellm/rust_bridge/image_edit/definition.py new file mode 100644 index 00000000000..f02440ef2db --- /dev/null +++ b/litellm/rust_bridge/image_edit/definition.py @@ -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] diff --git a/litellm/rust_bridge/image_edit/lifecycle.py b/litellm/rust_bridge/image_edit/lifecycle.py index fecf2e93792..d2a9ba8c47d 100644 --- a/litellm/rust_bridge/image_edit/lifecycle.py +++ b/litellm/rust_bridge/image_edit/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/image_generation/__init__.py b/litellm/rust_bridge/image_generation/__init__.py index e3144c5494a..62ca9099edf 100644 --- a/litellm/rust_bridge/image_generation/__init__.py +++ b/litellm/rust_bridge/image_generation/__init__.py @@ -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",) diff --git a/litellm/rust_bridge/image_generation/definition.py b/litellm/rust_bridge/image_generation/definition.py new file mode 100644 index 00000000000..fd705d5f90b --- /dev/null +++ b/litellm/rust_bridge/image_generation/definition.py @@ -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] diff --git a/litellm/rust_bridge/image_generation/lifecycle.py b/litellm/rust_bridge/image_generation/lifecycle.py index d3b93845bb5..49f21d5203b 100644 --- a/litellm/rust_bridge/image_generation/lifecycle.py +++ b/litellm/rust_bridge/image_generation/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/messages/__init__.py b/litellm/rust_bridge/messages/__init__.py index 4d21e25ed2b..3020125aee6 100644 --- a/litellm/rust_bridge/messages/__init__.py +++ b/litellm/rust_bridge/messages/__init__.py @@ -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", diff --git a/litellm/rust_bridge/messages/definition.py b/litellm/rust_bridge/messages/definition.py new file mode 100644 index 00000000000..8559bf7b022 --- /dev/null +++ b/litellm/rust_bridge/messages/definition.py @@ -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] diff --git a/litellm/rust_bridge/messages/lifecycle.py b/litellm/rust_bridge/messages/lifecycle.py index 74973621e92..1a6a18a9b05 100644 --- a/litellm/rust_bridge/messages/lifecycle.py +++ b/litellm/rust_bridge/messages/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/messages/types.py b/litellm/rust_bridge/messages/types.py index 009fe832dea..47c51a01382 100644 --- a/litellm/rust_bridge/messages/types.py +++ b/litellm/rust_bridge/messages/types.py @@ -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 diff --git a/litellm/rust_bridge/messages/value.py b/litellm/rust_bridge/messages/value.py index d52ec18ecc9..2542995a27c 100644 --- a/litellm/rust_bridge/messages/value.py +++ b/litellm/rust_bridge/messages/value.py @@ -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, + ), ) diff --git a/litellm/rust_bridge/moderation/__init__.py b/litellm/rust_bridge/moderation/__init__.py index e1e0360cd43..36070ff0f91 100644 --- a/litellm/rust_bridge/moderation/__init__.py +++ b/litellm/rust_bridge/moderation/__init__.py @@ -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",) diff --git a/litellm/rust_bridge/moderation/definition.py b/litellm/rust_bridge/moderation/definition.py new file mode 100644 index 00000000000..fece2e5dd04 --- /dev/null +++ b/litellm/rust_bridge/moderation/definition.py @@ -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] diff --git a/litellm/rust_bridge/moderation/lifecycle.py b/litellm/rust_bridge/moderation/lifecycle.py index 1115f1b3f94..f6f7402983a 100644 --- a/litellm/rust_bridge/moderation/lifecycle.py +++ b/litellm/rust_bridge/moderation/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/ocr/__init__.py b/litellm/rust_bridge/ocr/__init__.py index 2d70c77c370..5140e733266 100644 --- a/litellm/rust_bridge/ocr/__init__.py +++ b/litellm/rust_bridge/ocr/__init__.py @@ -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", diff --git a/litellm/rust_bridge/ocr/definition.py b/litellm/rust_bridge/ocr/definition.py new file mode 100644 index 00000000000..9e2e03214ae --- /dev/null +++ b/litellm/rust_bridge/ocr/definition.py @@ -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] diff --git a/litellm/rust_bridge/ocr/host.py b/litellm/rust_bridge/ocr/host.py new file mode 100644 index 00000000000..9f4489b2d0e --- /dev/null +++ b/litellm/rust_bridge/ocr/host.py @@ -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() diff --git a/litellm/rust_bridge/ocr/lifecycle.py b/litellm/rust_bridge/ocr/lifecycle.py index ea4fdb48a82..8c5d1f80f7a 100644 --- a/litellm/rust_bridge/ocr/lifecycle.py +++ b/litellm/rust_bridge/ocr/lifecycle.py @@ -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) diff --git a/litellm/rust_bridge/ocr/value.py b/litellm/rust_bridge/ocr/value.py index 7e36ddac37f..ae952014894 100644 --- a/litellm/rust_bridge/ocr/value.py +++ b/litellm/rust_bridge/ocr/value.py @@ -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), ) diff --git a/litellm/rust_bridge/rerank/__init__.py b/litellm/rust_bridge/rerank/__init__.py index e1a4c4e130b..159758c259d 100644 --- a/litellm/rust_bridge/rerank/__init__.py +++ b/litellm/rust_bridge/rerank/__init__.py @@ -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",) diff --git a/litellm/rust_bridge/rerank/definition.py b/litellm/rust_bridge/rerank/definition.py new file mode 100644 index 00000000000..e8107e8fa4f --- /dev/null +++ b/litellm/rust_bridge/rerank/definition.py @@ -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] diff --git a/litellm/rust_bridge/rerank/lifecycle.py b/litellm/rust_bridge/rerank/lifecycle.py index 54b607c5426..46e4e76ecaa 100644 --- a/litellm/rust_bridge/rerank/lifecycle.py +++ b/litellm/rust_bridge/rerank/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/responses/__init__.py b/litellm/rust_bridge/responses/__init__.py index b3e0eb3c858..1a8a8247d3e 100644 --- a/litellm/rust_bridge/responses/__init__.py +++ b/litellm/rust_bridge/responses/__init__.py @@ -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",) diff --git a/litellm/rust_bridge/responses/definition.py b/litellm/rust_bridge/responses/definition.py new file mode 100644 index 00000000000..94a389b8ff1 --- /dev/null +++ b/litellm/rust_bridge/responses/definition.py @@ -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] diff --git a/litellm/rust_bridge/responses/lifecycle.py b/litellm/rust_bridge/responses/lifecycle.py index 98f777efbf5..f0a7bcbde48 100644 --- a/litellm/rust_bridge/responses/lifecycle.py +++ b/litellm/rust_bridge/responses/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/responses/websocket.py b/litellm/rust_bridge/responses/websocket.py index 671ab5475dc..75bb4d10a6a 100644 --- a/litellm/rust_bridge/responses/websocket.py +++ b/litellm/rust_bridge/responses/websocket.py @@ -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) diff --git a/litellm/rust_bridge/route.py b/litellm/rust_bridge/route.py index 5f72e6d447c..f7eaf249f1b 100644 --- a/litellm/rust_bridge/route.py +++ b/litellm/rust_bridge/route.py @@ -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) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index d411673439f..5cd8fa0ca56 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -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: diff --git a/litellm/rust_bridge/speech/__init__.py b/litellm/rust_bridge/speech/__init__.py index 2c39ff42f4b..fa519435ce2 100644 --- a/litellm/rust_bridge/speech/__init__.py +++ b/litellm/rust_bridge/speech/__init__.py @@ -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",) diff --git a/litellm/rust_bridge/speech/definition.py b/litellm/rust_bridge/speech/definition.py new file mode 100644 index 00000000000..604a93920ff --- /dev/null +++ b/litellm/rust_bridge/speech/definition.py @@ -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] diff --git a/litellm/rust_bridge/speech/lifecycle.py b/litellm/rust_bridge/speech/lifecycle.py index 1f1ddc8435f..d12c9b0d009 100644 --- a/litellm/rust_bridge/speech/lifecycle.py +++ b/litellm/rust_bridge/speech/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/token_counter.py b/litellm/rust_bridge/token_counter.py deleted file mode 100644 index d36234f56c1..00000000000 --- a/litellm/rust_bridge/token_counter.py +++ /dev/null @@ -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 diff --git a/litellm/rust_bridge/token_counter/__init__.py b/litellm/rust_bridge/token_counter/__init__.py new file mode 100644 index 00000000000..1734a1799f7 --- /dev/null +++ b/litellm/rust_bridge/token_counter/__init__.py @@ -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", +) diff --git a/litellm/rust_bridge/token_counter/definition.py b/litellm/rust_bridge/token_counter/definition.py new file mode 100644 index 00000000000..8aaadc1c316 --- /dev/null +++ b/litellm/rust_bridge/token_counter/definition.py @@ -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] diff --git a/litellm/rust_bridge/token_counter/types.py b/litellm/rust_bridge/token_counter/types.py new file mode 100644 index 00000000000..5868c664050 --- /dev/null +++ b/litellm/rust_bridge/token_counter/types.py @@ -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 diff --git a/litellm/rust_bridge/token_counter/value.py b/litellm/rust_bridge/token_counter/value.py new file mode 100644 index 00000000000..68dc9203fa0 --- /dev/null +++ b/litellm/rust_bridge/token_counter/value.py @@ -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=""), + ) diff --git a/litellm/rust_bridge/transcription/__init__.py b/litellm/rust_bridge/transcription/__init__.py index eeb4699755f..42d18c7c1e8 100644 --- a/litellm/rust_bridge/transcription/__init__.py +++ b/litellm/rust_bridge/transcription/__init__.py @@ -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", diff --git a/litellm/rust_bridge/transcription/definition.py b/litellm/rust_bridge/transcription/definition.py new file mode 100644 index 00000000000..4307ee1395b --- /dev/null +++ b/litellm/rust_bridge/transcription/definition.py @@ -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] diff --git a/litellm/rust_bridge/transcription/lifecycle.py b/litellm/rust_bridge/transcription/lifecycle.py index 8516aa43c58..b8aad1b8297 100644 --- a/litellm/rust_bridge/transcription/lifecycle.py +++ b/litellm/rust_bridge/transcription/lifecycle.py @@ -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() diff --git a/litellm/rust_bridge/transcription/value.py b/litellm/rust_bridge/transcription/value.py index fe6693888a4..072e69f1978 100644 --- a/litellm/rust_bridge/transcription/value.py +++ b/litellm/rust_bridge/transcription/value.py @@ -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), ) diff --git a/pyproject.toml b/pyproject.toml index 62ce4b4fd61..7d29964b649 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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 = [ diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py b/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py index 671e660cfbe..3763e0f3ec4 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/ocr/case.py @@ -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", diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index e8126600467..8bf10fa9414 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -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 diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py index 4854ebc61de..cf5fa576aa4 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -7,9 +7,9 @@ import httpx import pytest import litellm +from litellm._uuid import uuid from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.llms.anthropic.chat.handler import ModelResponseIterator, make_call -from litellm._uuid import uuid from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.llms.openai import ( ChatCompletionToolCallChunk, @@ -24,9 +24,7 @@ async def test_make_call_passes_logging_obj_to_client_post(): mock_client = AsyncMock() mock_response = MagicMock() mock_response.aiter_lines = MagicMock( - return_value=iter( - [b'data: {"type":"message_start"}\n', b'data: {"type":"message_delta"}\n'] - ) + return_value=iter([b'data: {"type":"message_start"}\n', b'data: {"type":"message_delta"}\n']) ) mock_client.post.return_value = mock_response @@ -143,9 +141,7 @@ def test_redacted_thinking_content_block_delta(): "data": "EuoBCoYBGAIiQJ/SxkPAgqxhKok29YrpJHRUJ0OT8ahCHKAwyhmRuUhtdmDX9+mn4gDzKNv3fVpQdB01zEPMzNY3QuTCd+1bdtEqQK6JuKHqdndbwpr81oVWb4wxd1GqF/7Jkw74IlQa27oobX+KuRkopr9Dllt/RDe7Se0sI1IkU7tJIAQCoP46OAwSDF51P09q67xhHlQ3ihoM2aOVlkghq/X0w8NlIjBMNvXYNbjhyrOcIg6kPFn2ed/KK7Cm5prYAtXCwkb4Wr5tUSoSHu9T5hKdJRbr6WsqEc7Lle7FULqMLZGkhqXyc3BA", }, } - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) model_response = model_response_iterator.chunk_parser(chunk=chunk) print(f"\n\nmodel_response: {model_response}\n\n") assert model_response.choices[0].delta.thinking_blocks is not None @@ -153,19 +149,14 @@ def test_redacted_thinking_content_block_delta(): print( f"\n\nmodel_response.choices[0].delta.thinking_blocks[0]: {model_response.choices[0].delta.thinking_blocks[0]}\n\n" ) - assert ( - model_response.choices[0].delta.thinking_blocks[0]["type"] - == "redacted_thinking" - ) + assert model_response.choices[0].delta.thinking_blocks[0]["type"] == "redacted_thinking" assert model_response.choices[0].delta.provider_specific_fields is not None assert "thinking_blocks" in model_response.choices[0].delta.provider_specific_fields def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -189,17 +180,12 @@ def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): }, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) expected_delta_blocks = ( {"type": "thinking", "thinking": "Step 1. "}, @@ -213,18 +199,12 @@ def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): assert reasoning_content == "Step 1. Step 2." assert thinking_blocks == (*expected_delta_blocks, expected_thinking_block) - assert parsed_chunks[1].choices[0].delta.provider_specific_fields == { - "thinking_blocks": [expected_delta_blocks[0]] - } - assert parsed_chunks[-1].choices[0].delta.provider_specific_fields == { - "thinking_blocks": [expected_thinking_block] - } + assert parsed_chunks[1].choices[0].delta.provider_specific_fields == {"thinking_blocks": [expected_delta_blocks[0]]} + assert parsed_chunks[-1].choices[0].delta.provider_specific_fields == {"thinking_blocks": [expected_thinking_block]} def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -244,17 +224,12 @@ def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): {"type": "content_block_stop", "index": 0}, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) assert reasoning_content == "Step 1. Step 2." @@ -265,9 +240,7 @@ def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -286,17 +259,12 @@ def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): }, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) assert reasoning_content == "Step 1. Step 2." @@ -307,9 +275,7 @@ def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): def test_handle_json_mode_chunk_response_format_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) response_format_tool = ChatCompletionToolCallChunk( id="tool_123", type="function", @@ -320,9 +286,7 @@ def test_handle_json_mode_chunk_response_format_tool(): index=0, ) - text, tool_use = model_response_iterator._handle_json_mode_chunk( - "", response_format_tool - ) + text, tool_use = model_response_iterator._handle_json_mode_chunk("", response_format_tool) print(f"\n\nresponse_format_tool text: {text}\n\n") print(f"\n\nresponse_format_tool tool_use: {tool_use}\n\n") @@ -331,15 +295,11 @@ def test_handle_json_mode_chunk_response_format_tool(): def test_handle_json_mode_chunk_regular_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) regular_tool = ChatCompletionToolCallChunk( id="tool_456", type="function", - function=ChatCompletionToolCallFunctionChunk( - name="get_weather", arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name="get_weather", arguments='{"location": "San Francisco, CA"}'), index=0, ) @@ -353,17 +313,13 @@ def test_handle_json_mode_chunk_regular_tool(): def test_handle_json_mode_chunk_streaming_response_format_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: response_format tool with id and name, but no arguments first_chunk = ChatCompletionToolCallChunk( id="tool_123", type="function", - function=ChatCompletionToolCallFunctionChunk( - name=RESPONSE_FORMAT_TOOL_NAME, arguments="" - ), + function=ChatCompletionToolCallFunctionChunk(name=RESPONSE_FORMAT_TOOL_NAME, arguments=""), index=0, ) @@ -371,9 +327,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): second_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments='{"question": "What is the weather?"' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments='{"question": "What is the weather?"'), index=0, ) @@ -381,9 +335,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): third_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments=', "answer": "It is sunny"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments=', "answer": "It is sunny"}'), index=0, ) @@ -414,9 +366,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): def test_handle_json_mode_chunk_streaming_regular_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: regular tool with id and name, but no arguments first_chunk = ChatCompletionToolCallChunk( @@ -430,9 +380,7 @@ def test_handle_json_mode_chunk_streaming_regular_tool(): second_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments='{"location": "San Francisco, CA"}'), index=0, ) @@ -457,27 +405,19 @@ def test_handle_json_mode_chunk_streaming_regular_tool(): def test_response_format_tool_finish_reason(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: response_format tool response_format_tool = ChatCompletionToolCallChunk( id="tool_123", type="function", - function=ChatCompletionToolCallFunctionChunk( - name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": "test"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": "test"}'), index=0, ) # Process the tool call (should set converted_response_format_tool flag) - text, tool_use = model_response_iterator._handle_json_mode_chunk( - "", response_format_tool - ) - print( - f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n" - ) + text, tool_use = model_response_iterator._handle_json_mode_chunk("", response_format_tool) + print(f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n") # Simulate message_delta chunk with tool_use stop_reason message_delta_chunk = { @@ -496,25 +436,19 @@ def test_response_format_tool_finish_reason(): def test_regular_tool_finish_reason(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: regular tool (not response_format) regular_tool = ChatCompletionToolCallChunk( id="tool_456", type="function", - function=ChatCompletionToolCallFunctionChunk( - name="get_weather", arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name="get_weather", arguments='{"location": "San Francisco, CA"}'), index=0, ) # Process the tool call (should NOT set converted_response_format_tool flag) text, tool_use = model_response_iterator._handle_json_mode_chunk("", regular_tool) - print( - f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n" - ) + print(f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n") # Simulate message_delta chunk with tool_use stop_reason message_delta_chunk = { @@ -574,9 +508,7 @@ def test_text_only_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, f"Expected index=0, got {parsed.choices[0].index}" def test_streaming_thinking_deltas_count_reasoning_tokens_in_usage(): @@ -753,9 +685,7 @@ def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinkin ] self._write_response( content_type="text/event-stream", - body="".join( - f"data: {json.dumps(event)}\n\n" for event in events - ).encode("utf-8"), + body="".join(f"data: {json.dumps(event)}\n\n" for event in events).encode("utf-8"), ) return @@ -836,13 +766,9 @@ def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinkin assert content_chunks == [answer_text] assert stream_usage is not None stream_completion_details = stream_usage["completion_tokens_details"] - assert ( - stream_completion_details["reasoning_tokens"] - == non_stream_details.reasoning_tokens - ) + assert stream_completion_details["reasoning_tokens"] == non_stream_details.reasoning_tokens assert stream_completion_details["text_tokens"] == ( - stream_usage["completion_tokens"] - - stream_completion_details["reasoning_tokens"] + stream_usage["completion_tokens"] - stream_completion_details["reasoning_tokens"] ) assert requests_seen == [ { @@ -934,9 +860,9 @@ def test_text_and_tool_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0 for chunk type {chunk.get('type')}, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, ( + f"Expected index=0 for chunk type {chunk.get('type')}, got {parsed.choices[0].index}" + ) def test_multiple_tools_streaming_has_index_zero(): @@ -989,15 +915,11 @@ def test_multiple_tools_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, f"Expected index=0, got {parsed.choices[0].index}" def test_streaming_chunks_have_stable_ids(): - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) first_chunk = { "type": "content_block_delta", "index": 0, @@ -1022,9 +944,7 @@ def test_partial_json_chunk_accumulation(): This tests the fix for https://github.com/BerriAI/litellm/issues/17473 where network fragmentation can cause SSE data to arrive in partial chunks. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) partial_chunk_1 = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hel' partial_chunk_2 = 'lo"}}' @@ -1032,31 +952,21 @@ def test_partial_json_chunk_accumulation(): # First partial chunk should return None (still accumulating) result1 = iterator._parse_sse_data(f"data:{partial_chunk_1}") assert result1 is None, "First partial chunk should return None while accumulating" - assert ( - iterator.chunk_type == "accumulated_json" - ), "Should switch to accumulated_json mode" - assert ( - iterator.accumulated_json == partial_chunk_1 - ), "Should have accumulated first part" + assert iterator.chunk_type == "accumulated_json", "Should switch to accumulated_json mode" + assert iterator.accumulated_json == partial_chunk_1, "Should have accumulated first part" # Second partial chunk should complete the JSON and return a parsed result result2 = iterator._parse_sse_data(f"data:{partial_chunk_2}") assert result2 is not None, "Second chunk should return parsed result" - assert ( - iterator.accumulated_json == "" - ), "Buffer should be cleared after successful parse" - assert ( - result2.choices[0].delta.content == "Hello" - ), f"Expected 'Hello', got '{result2.choices[0].delta.content}'" + assert iterator.accumulated_json == "", "Buffer should be cleared after successful parse" + assert result2.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result2.choices[0].delta.content}'" def test_complete_json_chunk_no_accumulation(): """ Test that complete JSON chunks are parsed immediately without accumulation. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) complete_chunk = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}' @@ -1064,18 +974,14 @@ def test_complete_json_chunk_no_accumulation(): assert result is not None, "Complete chunk should return parsed result immediately" assert iterator.chunk_type == "valid_json", "Should remain in valid_json mode" assert iterator.accumulated_json == "", "Buffer should remain empty" - assert ( - result.choices[0].delta.content == "Hello" - ), f"Expected 'Hello', got '{result.choices[0].delta.content}'" + assert result.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result.choices[0].delta.content}'" def test_multiple_partial_chunks_accumulation(): """ Test that multiple partial chunks can be accumulated across several iterations. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Split a JSON chunk into three parts part1 = '{"type":"content_block_del' @@ -1103,17 +1009,11 @@ def test_accumulated_json_partial_fragment_returns_none_without_parsing(): unlike Vertex which already deferred parsing until the buffer could close. A fragment that can't close a JSON value must not trigger a decode attempt. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" - with patch.object( - json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode - ) as spy: - result = iterator._handle_accumulated_json_chunk( - '{"type":"content_block_delta","index":0,"delta":' - ) + with patch.object(json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode) as spy: + result = iterator._handle_accumulated_json_chunk('{"type":"content_block_delta","index":0,"delta":') assert result is None assert spy.call_count == 0, "incomplete buffer should not be parsed" @@ -1125,21 +1025,15 @@ def test_accumulated_json_does_not_reparse_every_fragment(): fragment. """ text = "x" * 200_000 - blob = json.dumps( - {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}} - ) + blob = json.dumps({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}) fragments = [blob[i : i + 4096] for i in range(0, len(blob), 4096)] assert len(fragments) > 10, "need a multi-fragment payload to exercise the bug" - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" parsed = None - with patch.object( - json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode - ) as spy: + with patch.object(json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode) as spy: for fragment in fragments: out = iterator._handle_accumulated_json_chunk(fragment) if out is not None: @@ -1163,9 +1057,7 @@ def test_accumulated_json_concatenated_envelopes_do_not_wedge(): and keeps the remainder, so both values surface across two calls. """ obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" first = iterator._handle_accumulated_json_chunk(obj + obj) @@ -1186,9 +1078,7 @@ def test_accumulated_json_heuristic_passes_but_value_still_incomplete(): heuristic must let the parse attempt through, and pop_next_value finding nothing must propagate as None rather than raising. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" result = iterator._handle_accumulated_json_chunk('{"type": {"nested": 1}') @@ -1202,9 +1092,7 @@ def test_accumulated_json_setter_and_sync_end_of_stream_drain(): underlying stream ends, instead of being silently dropped. """ obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator( - streaming_response=iter([]), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=iter([]), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" iterator.accumulated_json = obj # exercises the setter @@ -1219,9 +1107,7 @@ def test_accumulated_json_async_end_of_stream_drain(): import asyncio obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) iterator.chunk_type = "accumulated_json" iterator.accumulated_json = obj mock_async_iterator = MagicMock() @@ -1243,9 +1129,7 @@ def test_web_search_tool_result_no_extra_tool_calls(): The issue was that web_search_tool_result blocks have input_json_delta events with {} that were incorrectly being converted to tool calls. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence: # 1. server_tool_use block starts (web_search) @@ -1320,9 +1204,7 @@ def test_web_search_tool_result_no_extra_tool_calls(): # Should have exactly 2 tool calls: # 1. From content_block_start (server_tool_use) with id and name # 2. From content_block_delta with the actual query - assert ( - len(tool_calls_emitted) == 2 - ), f"Expected 2 tool calls, got {len(tool_calls_emitted)}" + assert len(tool_calls_emitted) == 2, f"Expected 2 tool calls, got {len(tool_calls_emitted)}" # First tool call should have the id and name assert tool_calls_emitted[0]["id"] == "srvtoolu_01ABC123" @@ -1338,9 +1220,7 @@ def test_current_content_block_type_tracking(): """ Test that current_content_block_type is properly tracked and reset. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Initially should be None assert iterator.current_content_block_type is None @@ -1421,7 +1301,9 @@ def test_web_search_calls_are_cumulative_through_incomplete_search(): ) assert list(first_start.choices[0].delta.provider_specific_fields["web_search_calls"]) == ["srvtoolu_A"] - assert first_result.choices[0].delta.provider_specific_fields["web_search_calls"]["srvtoolu_A"].status == "completed" + assert ( + first_result.choices[0].delta.provider_specific_fields["web_search_calls"]["srvtoolu_A"].status == "completed" + ) calls = second_start.choices[0].delta.provider_specific_fields["web_search_calls"] assert list(calls) == ["srvtoolu_A", "srvtoolu_B"] assert calls["srvtoolu_A"].status == "completed" @@ -1439,9 +1321,7 @@ def test_web_search_tool_result_captured_in_provider_specific_fields(): The web_search_tool_result content comes ALL AT ONCE in content_block_start, not in deltas, so we need to capture it there. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence with web_search_tool_result chunks = [ @@ -1512,23 +1392,15 @@ def test_web_search_tool_result_captured_in_provider_specific_fields(): and parsed.choices[0].delta.provider_specific_fields and "web_search_results" in parsed.choices[0].delta.provider_specific_fields ): - web_search_results = parsed.choices[0].delta.provider_specific_fields[ - "web_search_results" - ] + web_search_results = parsed.choices[0].delta.provider_specific_fields["web_search_results"] # Verify web_search_results was captured assert web_search_results is not None, "web_search_results should be captured" assert len(web_search_results) == 1, "Should have 1 web_search_tool_result block" - assert ( - web_search_results[0]["type"] == "web_search_tool_result" - ), "Block type should be web_search_tool_result" - assert ( - web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123" - ), "tool_use_id should match" + assert web_search_results[0]["type"] == "web_search_tool_result", "Block type should be web_search_tool_result" + assert web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123", "tool_use_id should match" assert len(web_search_results[0]["content"]) == 2, "Should have 2 search results" - assert ( - web_search_results[0]["content"][0]["title"] == "Fun Otter Facts" - ), "First result title should match" + assert web_search_results[0]["content"][0]["title"] == "Fun Otter Facts", "First result title should match" def test_web_fetch_tool_result_captured_in_provider_specific_fields(): @@ -1542,9 +1414,7 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields(): The web_fetch_tool_result content comes ALL AT ONCE in content_block_start, not in deltas, so we need to capture it there. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence with web_fetch_tool_result chunks = [ @@ -1615,25 +1485,15 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields(): and parsed.choices[0].delta.provider_specific_fields and "web_search_results" in parsed.choices[0].delta.provider_specific_fields ): - web_search_results = parsed.choices[0].delta.provider_specific_fields[ - "web_search_results" - ] + web_search_results = parsed.choices[0].delta.provider_specific_fields["web_search_results"] # Verify web_fetch_tool_result was captured (stored in web_search_results list) assert web_search_results is not None, "web_search_results should be captured" assert len(web_search_results) == 1, "Should have 1 web_fetch_tool_result block" - assert ( - web_search_results[0]["type"] == "web_fetch_tool_result" - ), "Block type should be web_fetch_tool_result" - assert ( - web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123" - ), "tool_use_id should match" - assert ( - web_search_results[0]["content"]["url"] == "https://example.com" - ), "URL should match" - assert ( - web_search_results[0]["content"]["content"]["title"] == "Example Page" - ), "Title should match" + assert web_search_results[0]["type"] == "web_fetch_tool_result", "Block type should be web_fetch_tool_result" + assert web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123", "tool_use_id should match" + assert web_search_results[0]["content"]["url"] == "https://example.com", "URL should match" + assert web_search_results[0]["content"]["content"]["title"] == "Example Page", "Title should match" def test_web_fetch_tool_result_no_extra_tool_calls(): @@ -1646,9 +1506,7 @@ def test_web_fetch_tool_result_no_extra_tool_calls(): The issue was that web_fetch_tool_result blocks have input_json_delta events with {} that were incorrectly being converted to tool calls. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # to verify it doesn't emit tool calls chunks = [ @@ -1692,9 +1550,9 @@ def test_web_fetch_tool_result_no_extra_tool_calls(): tool_call_count += 1 # Should have 0 tool calls - web_fetch_tool_result should not emit tool calls - assert ( - tool_call_count == 0 - ), f"Expected 0 tool calls, got {tool_call_count}. web_fetch_tool_result should not emit tool calls" + assert tool_call_count == 0, ( + f"Expected 0 tool calls, got {tool_call_count}. web_fetch_tool_result should not emit tool calls" + ) def test_container_in_provider_specific_fields_streaming(): @@ -1704,9 +1562,7 @@ def test_container_in_provider_specific_fields_streaming(): When container with skills is used, the container field should be present in the provider_specific_fields of the message_delta chunk. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate streaming chunks chunks = [ @@ -1774,20 +1630,12 @@ def test_container_in_provider_specific_fields_streaming(): and parsed.choices[0].delta.provider_specific_fields and "container" in parsed.choices[0].delta.provider_specific_fields ): - container_field = parsed.choices[0].delta.provider_specific_fields[ - "container" - ] + container_field = parsed.choices[0].delta.provider_specific_fields["container"] # Verify container was captured - assert ( - container_field is not None - ), "container should be captured in provider_specific_fields" - assert ( - container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p" - ), "container id should match" - assert ( - container_field["expires_at"] == "2025-12-16T04:57:16.913181Z" - ), "expires_at should match" + assert container_field is not None, "container should be captured in provider_specific_fields" + assert container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p", "container id should match" + assert container_field["expires_at"] == "2025-12-16T04:57:16.913181Z", "expires_at should match" assert len(container_field["skills"]) == 1, "Should have 1 skill" assert container_field["skills"][0]["skill_id"] == "pptx", "skill_id should be pptx" assert container_field["skills"][0]["version"] == "20251013", "version should match" @@ -1800,9 +1648,7 @@ def test_container_in_provider_specific_fields_non_streaming(): When container with skills is used in non-streaming, the container field should be present in the provider_specific_fields of the response. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) # Simulate a message_delta chunk with container (as it would appear in non-streaming) message_delta_chunk = { @@ -1838,21 +1684,13 @@ def test_container_in_provider_specific_fields_non_streaming(): # Verify container is in provider_specific_fields assert model_response.choices[0].delta.provider_specific_fields is not None assert "container" in model_response.choices[0].delta.provider_specific_fields - container_field = model_response.choices[0].delta.provider_specific_fields[ - "container" - ] + container_field = model_response.choices[0].delta.provider_specific_fields["container"] assert container_field["id"] == "container_abc123xyz", "container id should match" - assert ( - container_field["expires_at"] == "2025-12-20T10:30:00.000000Z" - ), "expires_at should match" + assert container_field["expires_at"] == "2025-12-20T10:30:00.000000Z", "expires_at should match" assert len(container_field["skills"]) == 2, "Should have 2 skills" - assert ( - container_field["skills"][0]["skill_id"] == "code_execution" - ), "First skill_id should be code_execution" - assert ( - container_field["skills"][1]["skill_id"] == "pptx" - ), "Second skill_id should be pptx" + assert container_field["skills"][0]["skill_id"] == "code_execution", "First skill_id should be code_execution" + assert container_field["skills"][1]["skill_id"] == "pptx", "Second skill_id should be pptx" def test_container_absent_when_not_provided(): @@ -1861,9 +1699,7 @@ def test_container_absent_when_not_provided(): This ensures we don't add empty or None container fields. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) # message_delta without container message_delta_chunk = { @@ -1882,9 +1718,9 @@ def test_container_absent_when_not_provided(): # Verify container is NOT in provider_specific_fields when not provided if model_response.choices[0].delta.provider_specific_fields: - assert ( - "container" not in model_response.choices[0].delta.provider_specific_fields - ), "container should not be present when not provided in delta" + assert "container" not in model_response.choices[0].delta.provider_specific_fields, ( + "container should not be present when not provided in delta" + ) def test_streaming_code_execution_produces_code_interpreter_results(): @@ -2080,8 +1916,7 @@ def test_streaming_multiple_code_executions_no_duplicates(): # Second (final) emission: cumulative list with BOTH results # This is what stream_chunk_builder will pick as "last value wins" assert len(emissions[1]) == 2, ( - f"Expected final emission to have 2 results, got {len(emissions[1])}. " - f"IDs: {[r.id for r in emissions[1]]}" + f"Expected final emission to have 2 results, got {len(emissions[1])}. IDs: {[r.id for r in emissions[1]]}" ) assert emissions[1][0].id == "srvtoolu_01AAA" assert emissions[1][0].code == "echo first" @@ -2245,9 +2080,7 @@ def test_empty_output_produces_null_outputs(): assert code_results is not None, "No code_interpreter_results emitted" assert len(code_results) == 1 assert code_results[0].id == "srvtoolu_01AAA" - assert ( - code_results[0].outputs is None - ), f"Expected outputs=None for empty execution, got {code_results[0].outputs}" + assert code_results[0].outputs is None, f"Expected outputs=None for empty execution, got {code_results[0].outputs}" def test_non_bash_tool_result_skipped(): @@ -2310,12 +2143,10 @@ def test_non_bash_tool_result_skipped(): code_results = psf["code_interpreter_results"] # code_interpreter_results should be emitted but empty (no bash results) - assert ( - code_results is not None - ), "Expected code_interpreter_results key to be emitted" - assert ( - len(code_results) == 0 - ), f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" + assert code_results is not None, "Expected code_interpreter_results key to be emitted" + assert len(code_results) == 0, ( + f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" + ) class TestRustChatCompletionsHook: @@ -2352,13 +2183,9 @@ class TestRustChatCompletionsHook: from litellm.rust_bridge import chat_completions as bridge 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) @staticmethod def _completion_kwargs(**overrides): @@ -2400,17 +2227,19 @@ class TestRustChatCompletionsHook: seen = {"gate": [], "call": []} - def gate(**kwargs): - seen["gate"].append(kwargs) - return decline_reason - def native(**kwargs): + seen["gate"].append(kwargs) + if decline_reason is not None: + from litellm.rust_bridge import _native + + raise _native.RustBridgeDeclined(decline_reason) seen["call"].append(kwargs) if sync_error is not None: raise sync_error + kwargs["on_request"]() return dict(sync_result if sync_result is not None else self.RUST_RESPONSE) - bridge.set_rust_chat_completions(decline=gate, chat_completions=native) + bridge.set_rust_chat_completions(chat_completions=native) return seen def test_rust_true_serves_the_call_and_stamps_the_header(self): @@ -2456,9 +2285,7 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion seen = self._inject() - AnthropicChatCompletion().completion( - **self._completion_kwargs(optional_params={"max_tokens": 7}) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(optional_params={"max_tokens": 7})) assert seen["call"][0]["optional_params"]["max_tokens"] == 7 def test_without_the_opt_in_the_core_is_never_consulted(self, monkeypatch): @@ -2467,15 +2294,14 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ) as transform, patch.object( - AnthropicChatCompletion, "acompletion_function" + with ( + patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ) as transform, + patch.object(AnthropicChatCompletion, "acompletion_function"), ): try: - AnthropicChatCompletion().completion( - **self._completion_kwargs(litellm_params={}) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(litellm_params={})) except Exception: # The Python path goes on to make an HTTP call; reaching it is # the assertion, so the network failure below is expected. @@ -2489,9 +2315,7 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject(decline_reason="unrecognized request parameter") - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: AnthropicChatCompletion().completion(**self._completion_kwargs()) except Exception: @@ -2504,9 +2328,7 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: AnthropicChatCompletion().completion( **self._completion_kwargs(optional_params={"max_tokens": 16, "stream": True}) @@ -2514,15 +2336,14 @@ class TestRustChatCompletionsHook: except Exception: pass assert seen["gate"] == [] + assert seen["call"] == [] def test_pre_call_logging_fires_exactly_once_on_the_rust_path(self): from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion seen = self._inject() logging_obj = MagicMock() - AnthropicChatCompletion().completion( - **self._completion_kwargs(logging_obj=logging_obj) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) assert logging_obj.pre_call.call_count == 1 assert len(seen["call"]) == 1 @@ -2536,9 +2357,7 @@ class TestRustChatCompletionsHook: self._inject() logging_obj = MagicMock() - AnthropicChatCompletion().completion( - **self._completion_kwargs(logging_obj=logging_obj) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) assert logging_obj.post_call.call_count == 1 logged = logging_obj.post_call.call_args.kwargs["original_response"] @@ -2563,18 +2382,12 @@ class TestRustChatCompletionsHook: raise _Declined("blank message text") monkeypatch.setattr("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) logging_obj, calls = self._recording_logging_obj() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: - AnthropicChatCompletion().completion( - **self._completion_kwargs(logging_obj=logging_obj) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) except Exception: # The Python path goes on to make an HTTP call; the log count is # the assertion, so a failure past this point is expected. @@ -2599,21 +2412,15 @@ class TestRustChatCompletionsHook: 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() async def python_path(**_kwargs): return sentinel - with patch.object( - AnthropicChatCompletion, "acompletion_function", side_effect=python_path - ) as python_call: - result = await AnthropicChatCompletion().completion( - **self._completion_kwargs(acompletion=True) - ) + with patch.object(AnthropicChatCompletion, "acompletion_function", side_effect=python_path) as python_call: + result = await AnthropicChatCompletion().completion(**self._completion_kwargs(acompletion=True)) assert result is sentinel assert python_call.called, "a failing rust call must re-enter the python path" @@ -2623,23 +2430,19 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion from litellm.rust_bridge import chat_completions as bridge - async def native(**_kwargs): + async def native(**kwargs): + kwargs["on_request"]() return dict(self.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(AnthropicChatCompletion, "acompletion_function") as python_call: - result = await AnthropicChatCompletion().completion( - **self._completion_kwargs(acompletion=True) - ) + result = await AnthropicChatCompletion().completion(**self._completion_kwargs(acompletion=True)) assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} assert not python_call.called - def test_pre_call_logging_fires_once_when_the_sync_rust_call_declines(self, monkeypatch): """One request, one pre_call, on the synchronous path too. Without the suppression the Python path logs a second time for the same attempt.""" @@ -2659,27 +2462,19 @@ class TestRustChatCompletionsHook: def declining_native(**_kwargs): raise _Declined("blank message text") - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, chat_completions=declining_native - ) + bridge.set_rust_chat_completions(chat_completions=declining_native) logging_obj, calls = self._recording_logging_obj() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: - AnthropicChatCompletion().completion( - **self._completion_kwargs(logging_obj=logging_obj) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) except Exception: # The Python path goes on to make an HTTP call; the log count is # the assertion, so a failure past this point is expected. pass assert len(calls["pre_call"]) == 1 - assert calls["pre_call"][0]["additional_args"]["complete_input_dict"]["model"] == ( - "claude-sonnet-4-5" - ) + assert calls["pre_call"][0]["additional_args"]["complete_input_dict"]["model"] == "m" def test_pre_call_logging_still_fires_when_rust_is_not_involved(self, monkeypatch): """The suppression must not swallow the log on the ordinary path.""" @@ -2689,9 +2484,7 @@ class TestRustChatCompletionsHook: self._inject() logging_obj, calls = self._recording_logging_obj() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: AnthropicChatCompletion().completion( **self._completion_kwargs(litellm_params={}, logging_obj=logging_obj) diff --git a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py index 0e93da3d0c9..1679d1f8d30 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -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(), diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index c39779972c0..98741c177c5 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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={}, diff --git a/tests/test_litellm/ocr/test_legacy.py b/tests/test_litellm/ocr/test_legacy.py index 2313efc92f1..c046b839071 100644 --- a/tests/test_litellm/ocr/test_legacy.py +++ b/tests/test_litellm/ocr/test_legacy.py @@ -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 = { diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index b267c81ab8e..9f07d542b72 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -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", + ) diff --git a/tests/test_litellm/rust_bridge/stubtest.ini b/tests/test_litellm/rust_bridge/stubtest.ini new file mode 100644 index 00000000000..06eab31680e --- /dev/null +++ b/tests/test_litellm/rust_bridge/stubtest.ini @@ -0,0 +1,2 @@ +[mypy] +follow_imports = skip diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index ec69241a4a0..3d20d6a249f 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -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" diff --git a/tests/test_litellm/rust_bridge/test_configuration.py b/tests/test_litellm/rust_bridge/test_configuration.py deleted file mode 100644 index ab2ba946ea9..00000000000 --- a/tests/test_litellm/rust_bridge/test_configuration.py +++ /dev/null @@ -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 diff --git a/tests/test_litellm/rust_bridge/test_configuration_env.py b/tests/test_litellm/rust_bridge/test_configuration_env.py index 0fab10e2bed..8db7fb6355d 100644 --- a/tests/test_litellm/rust_bridge/test_configuration_env.py +++ b/tests/test_litellm/rust_bridge/test_configuration_env.py @@ -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) diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index 5be30a8046e..602927554dc 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -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) diff --git a/tests/test_litellm/rust_bridge/test_route.py b/tests/test_litellm/rust_bridge/test_route.py index 565421a84d1..bc1e1e7fe1e 100644 --- a/tests/test_litellm/rust_bridge/test_route.py +++ b/tests/test_litellm/rust_bridge/test_route.py @@ -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) diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index b0fa510069b..c3b24546e5c 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -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) diff --git a/tests/test_litellm/rust_bridge/test_token_counter.py b/tests/test_litellm/rust_bridge/test_token_counter.py index 71aa79cc4bb..fd9febad083 100644 --- a/tests/test_litellm/rust_bridge/test_token_counter.py +++ b/tests/test_litellm/rust_bridge/test_token_counter.py @@ -1,444 +1,165 @@ -"""Tests for the Rust input token counter bridge. - -The native factory is dependency-injected through ``TOKEN_COUNTER.override`` -so the fallback cases run without the compiled extension present. The parity -cases need the extension and are skipped when it is not built. -""" - from __future__ import annotations import json -from types import MappingProxyType +from collections.abc import Callable, Iterator from typing import Final import pytest -import tiktoken -from tokenizers import Tokenizer import litellm -from litellm.constants import TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS -from litellm.litellm_core_utils.token_counter import openai_tokenizer_encoding -from litellm.proxy.spend_tracking.budget_reservation import _count_input_tokens from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import token_counter as bridge -from litellm.utils import claude_json_str + +native = pytest.importorskip("litellm.rust_bridge._native") MODEL: Final = "claude-sonnet-4-5-20250929" CL100K_MODEL: Final = "gpt-4" O200K_MODEL: Final = "gpt-4o" -TOKENIZERS: Final[tuple[bridge.RustTokenizer, ...]] = ("anthropic", "cl100k_base", "o200k_base") -RANK_FILE_LINES: Final = MappingProxyType({"cl100k_base": 100_256, "o200k_base": 199_998}) BODY: Final = json.dumps({"model": MODEL, "messages": [{"role": "user", "content": "hello"}]}).encode() +ANTHROPIC: Final = bridge.RustTokenizer(kind="anthropic", encoding="", disabled=False, legacy_accounting=False) +CL100K: Final = bridge.RustTokenizer(kind=None, encoding="cl100k_base", disabled=False, legacy_accounting=False) +O200K: Final = bridge.RustTokenizer(kind=None, encoding="o200k_base", disabled=False, legacy_accounting=False) +TOKENIZERS: Final = (ANTHROPIC, CL100K, O200K) -class _FakeDeclined(Exception): - pass - - -class _FakeUpstream(Exception): - pass +_FakeDeclined = native.RustBridgeDeclined +_FakeUnavailable = native.RustBridgeUnavailable +_FakeUpstream = native.RustUpstreamError class _FakeNative: + RustBridgeUnavailable = _FakeUnavailable RustBridgeDeclined = _FakeDeclined RustUpstreamError = _FakeUpstream class _RecordingCounter: - def __init__(self, tokenizer_json: str) -> None: - self.tokenizer_json = tokenizer_json - self.bodies: list[bytes] = [] + def __init__(self, error: Exception | None = None) -> None: + self.error: Final = error + self.calls: Final[list[tuple[bytes, bridge.RustTokenizer, str]]] = [] - async def acount_request(self, body: bytes) -> object: - self.bodies.append(body) + async def __call__( + self, + body: bytes, + kind: str | None, + encoding: str, + disabled: bool, + legacy_accounting: bool, + resource_loader: Callable[[str], str], + ) -> object: + tokenizer: Final = bridge.RustTokenizer(kind, encoding, disabled, legacy_accounting) + resource_name: Final = kind or encoding + self.calls.append((body, tokenizer, resource_name)) + if self.error is not None: + raise self.error + resource_loader(resource_name) return {"model": MODEL, "input_tokens": 42} -class _RecordingFactory: - """Stands in for the native `TokenCounter` class: callable for tokenizer JSON, `from_*_ranks` for rank files.""" - - def __init__(self) -> None: - self.counters: list[_RecordingCounter] = [] - self.rank_files: list[str] = [] - - def __call__(self, tokenizer_json: str) -> _RecordingCounter: - counter = _RecordingCounter(tokenizer_json) - self.counters.append(counter) - return counter - - def from_cl100k_ranks(self, rank_file: str) -> _RecordingCounter: - self.rank_files.append(rank_file) - return self("cl100k_base") - - def from_o200k_ranks(self, rank_file: str) -> _RecordingCounter: - self.rank_files.append(rank_file) - return self("o200k_base") - - -class _RaisingCounter: - def __init__(self, error: Exception) -> None: - self.error = error - - async def acount_request(self, body: bytes) -> object: - raise self.error - - -class _RaisingFactory: - """Every counter it builds, for either tokenizer, raises `error` on count.""" - - def __init__(self, error: Exception) -> None: - self.error = error - - def __call__(self, tokenizer_json: str) -> _RaisingCounter: - return _RaisingCounter(self.error) - - def from_cl100k_ranks(self, rank_file: str) -> _RaisingCounter: - return _RaisingCounter(self.error) - - def from_o200k_ranks(self, rank_file: str) -> _RaisingCounter: - return _RaisingCounter(self.error) - - @pytest.fixture(autouse=True) -def _reset_bridge(monkeypatch: pytest.MonkeyPatch): +def reset_bridge(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: bridge.TOKEN_COUNTER.reset() - bridge._counter.cache_clear() configuration.reset_rust_configuration() + configuration.rust(True) monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) yield bridge.TOKEN_COUNTER.reset() - bridge._counter.cache_clear() configuration.reset_rust_configuration() @pytest.mark.asyncio @pytest.mark.parametrize("tokenizer", TOKENIZERS) -async def test_disabled_bridge_never_constructs_a_counter(tokenizer: bridge.RustTokenizer) -> None: - factory: Final = _RecordingFactory() +async def test_disabled_bridge_never_calls_native(tokenizer: bridge.RustTokenizer) -> None: + counter: Final = _RecordingCounter() + bridge.TOKEN_COUNTER.override(counter) litellm.rust(False) - bridge.TOKEN_COUNTER.override(factory) - assert await bridge.count_input_tokens(BODY, tokenizer) is None - assert factory.counters == [] - - -@pytest.mark.asyncio -async def test_enabled_bridge_returns_typed_count_and_reuses_one_counter() -> None: - factory: Final = _RecordingFactory() - litellm.rust(True) - bridge.TOKEN_COUNTER.override(factory) - - first: Final = await bridge.count_input_tokens(BODY, "anthropic") - second: Final = await bridge.count_input_tokens(BODY, "anthropic") - - assert first == bridge.InputTokenCount(model=MODEL, input_tokens=42) - assert second == first - assert len(factory.counters) == 1 - assert factory.counters[0].bodies == [BODY, BODY] - assert json.loads(factory.counters[0].tokenizer_json)["model"]["type"] == "BPE" - - -@pytest.mark.asyncio -@pytest.mark.parametrize("tokenizer", ("cl100k_base", "o200k_base")) -async def test_tiktoken_counter_is_built_from_the_vendored_rank_file_once(tokenizer: bridge.RustTokenizer) -> None: - factory: Final = _RecordingFactory() - litellm.rust(True) - bridge.TOKEN_COUNTER.override(factory) - - first: Final = await bridge.count_input_tokens(BODY, tokenizer) - second: Final = await bridge.count_input_tokens(BODY, tokenizer) - - assert first == second == bridge.InputTokenCount(model=MODEL, input_tokens=42) - assert len(factory.rank_files) == 1 - assert factory.rank_files[0].startswith("IQ== 0\n") - assert factory.rank_files[0].count("\n") == RANK_FILE_LINES[tokenizer] - assert factory.counters[0].tokenizer_json == tokenizer - assert factory.counters[0].bodies == [BODY, BODY] - - -@pytest.mark.asyncio -async def test_each_tokenizer_gets_its_own_cached_counter() -> None: - factory: Final = _RecordingFactory() - litellm.rust(True) - bridge.TOKEN_COUNTER.override(factory) - - await bridge.count_input_tokens(BODY, "anthropic") - await bridge.count_input_tokens(BODY, "cl100k_base") - await bridge.count_input_tokens(BODY, "o200k_base") - await bridge.count_input_tokens(BODY, "anthropic") - await bridge.count_input_tokens(BODY, "o200k_base") - - assert [counter.tokenizer_json for counter in factory.counters][1:] == ["cl100k_base", "o200k_base"] - assert [len(counter.bodies) for counter in factory.counters] == [2, 1, 2] - - -@pytest.mark.asyncio -async def test_missing_native_module_falls_back(monkeypatch: pytest.MonkeyPatch) -> None: - litellm.rust(True) - monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) - - assert [await bridge.count_input_tokens(BODY, tokenizer) for tokenizer in TOKENIZERS] == [None, None, None] + assert counter.calls == [] @pytest.mark.asyncio @pytest.mark.parametrize("tokenizer", TOKENIZERS) -async def test_declined_request_falls_back(tokenizer: bridge.RustTokenizer) -> None: - litellm.rust(True) - bridge.TOKEN_COUNTER.override(_RaisingFactory(_FakeDeclined("request has no messages"))) - - assert await bridge.count_input_tokens(BODY, tokenizer) is None +async def test_one_native_count_entrypoint_receives_configuration_and_body(tokenizer: bridge.RustTokenizer) -> None: + counter: Final = _RecordingCounter() + bridge.TOKEN_COUNTER.override(counter) + result: Final = await bridge.count_input_tokens(BODY, tokenizer) + assert result == bridge.InputTokenCount(model=MODEL, input_tokens=42) + assert counter.calls == [(BODY, tokenizer, tokenizer.kind or tokenizer.encoding)] @pytest.mark.asyncio -@pytest.mark.parametrize("tokenizer", TOKENIZERS) -async def test_runtime_failure_falls_back(tokenizer: bridge.RustTokenizer) -> None: - litellm.rust(True) - bridge.TOKEN_COUNTER.override(_RaisingFactory(RuntimeError("encode failed"))) +@pytest.mark.parametrize("error", (_FakeDeclined("unsupported"), _FakeUnavailable("resource"))) +async def test_decline_and_resource_unavailability_fall_back(error: Exception) -> None: + bridge.TOKEN_COUNTER.override(_RecordingCounter(error)) + assert await bridge.count_input_tokens(BODY, ANTHROPIC) is None - assert await bridge.count_input_tokens(BODY, tokenizer) is None + +@pytest.mark.asyncio +async def test_unexpected_counting_failure_propagates() -> None: + bridge.TOKEN_COUNTER.override(_RecordingCounter(RuntimeError("encode failed"))) + with pytest.raises(RuntimeError, match="encode failed"): + await bridge.count_input_tokens(BODY, ANTHROPIC) @pytest.mark.parametrize( ("model", "expected"), ( - (MODEL, "anthropic"), - ("claude-3-5-sonnet-20241022", "cl100k_base"), - ("gpt-4", "cl100k_base"), - ("gpt-4-turbo", "cl100k_base"), - ("gpt-3.5-turbo", "cl100k_base"), - ("azure/gpt-35-turbo", "cl100k_base"), - ("gemini/gemini-2.5-pro", "cl100k_base"), - ("mistral/mistral-large-latest", "cl100k_base"), - ("my-router-alias", "cl100k_base"), - ("azure/gpt-4o", "cl100k_base"), - ("command-r-plus", "cl100k_base"), - ("gpt-4o", "o200k_base"), - ("gpt-4o-mini", "o200k_base"), - ("gpt-4o-2024-08-06", "o200k_base"), - ("chatgpt-4o-latest", "o200k_base"), - ("gpt-4.1", "o200k_base"), - ("gpt-5", "o200k_base"), - ("gpt-5-mini", "o200k_base"), - ("o1", "o200k_base"), - ("o3", "o200k_base"), - ("o3-mini", "o200k_base"), - ("o4-mini", "o200k_base"), - ("replicate/meta/llama-2-70b-chat", None), - ("meta-llama/Llama-3-8b", None), + (MODEL, ANTHROPIC), + ("gpt-4", CL100K), + ("gpt-4o", O200K), + ("replicate/meta/llama-2-70b-chat", bridge.RustTokenizer("llama2", "", False, False)), ), ) -def test_rust_tokenizer_mirrors_python_tokenizer_selection(model: str, expected: bridge.RustTokenizer | None) -> None: +def test_tokenizer_configuration_matches_python_selection( + model: str, expected: bridge.RustTokenizer +) -> None: + bridge.TOKEN_COUNTER.override(_RecordingCounter()) assert bridge.rust_tokenizer(model) == expected -@pytest.mark.parametrize( - ("model", "python_encoding"), - (("text-davinci-003", "p50k_base"), ("gpt-oss-120b", "o200k_harmony")), -) -def test_rust_tokenizer_declines_tiktoken_encodings_rust_does_not_have( - monkeypatch: pytest.MonkeyPatch, model: str, python_encoding: str -) -> None: - monkeypatch.setattr(litellm, "open_ai_chat_completion_models", litellm.open_ai_chat_completion_models | {model}) - - assert openai_tokenizer_encoding(model).name == python_encoding - assert bridge.rust_tokenizer(model) is None - - -def test_rust_tokenizer_declines_the_cohere_tokenizer_download(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(litellm, "cohere_models", litellm.cohere_models | {"command-r-plus"}) - - assert bridge.rust_tokenizer("command-r-plus") is None - - -@pytest.mark.parametrize("legacy_model", ("gpt-3.5-turbo-0301", "gpt-35-turbo-0301")) -def test_rust_tokenizer_declines_legacy_message_accounting_python_prices_differently( - monkeypatch: pytest.MonkeyPatch, legacy_model: str -) -> None: - monkeypatch.setattr( - litellm, "open_ai_chat_completion_models", litellm.open_ai_chat_completion_models | {"gpt-3.5-turbo-0301"} - ) - monkeypatch.setattr(litellm, "azure_llms", {**litellm.azure_llms, "gpt-35-turbo-0301": "azure"}) - messages: Final = [{"role": "user", "name": "bob", "content": "hello there"}] - - assert litellm.token_counter(model=legacy_model, messages=messages) != litellm.token_counter( - model=CL100K_MODEL, messages=messages - ) - assert bridge.rust_tokenizer(legacy_model) is None - assert bridge.rust_tokenizer(CL100K_MODEL) == "cl100k_base" - - -@pytest.mark.parametrize("model", (MODEL, CL100K_MODEL, O200K_MODEL, "gpt-5", "o3")) -def test_rust_tokenizer_names_the_encoding_python_actually_counts_with(model: str) -> None: - text: Final = ( - "Hello, world! camelCase ABCdef \u00e9\u00e8 12345 \u3053\u3093\u306b\u3061\u306f <|endoftext|>\r\n" * 9 - ) - python_count: Final = litellm.token_counter(model=model, text=text) - cl100k_count: Final = len(tiktoken.get_encoding("cl100k_base").encode(text, disallowed_special=())) - o200k_count: Final = len(tiktoken.get_encoding("o200k_base").encode(text, disallowed_special=())) - assert cl100k_count != o200k_count - match bridge.rust_tokenizer(model): - case "cl100k_base": - assert python_count == cl100k_count - case "o200k_base": - assert python_count == o200k_count - case "anthropic": - assert python_count == len(Tokenizer.from_str(claude_json_str).encode(text).ids) - assert python_count not in {cl100k_count, o200k_count} - case None: - pytest.fail(f"{model} must have a Rust tokenizer") - - -def test_disabled_hf_download_routes_anthropic_models_to_cl100k_like_python(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(litellm, "disable_hf_tokenizer_download", True) - - assert bridge.rust_tokenizer(MODEL) == "cl100k_base" - assert bridge.rust_tokenizer("meta-llama/Llama-3-8b") == "cl100k_base" - assert bridge.rust_tokenizer(O200K_MODEL) == "o200k_base" - - -def test_disabled_token_counter_declines_every_model(monkeypatch: pytest.MonkeyPatch) -> None: +def test_unsupported_configuration_is_left_for_native_admission(monkeypatch: pytest.MonkeyPatch) -> None: + bridge.TOKEN_COUNTER.override(_RecordingCounter()) monkeypatch.setattr(litellm, "disable_token_counter", True) - - assert bridge.rust_tokenizer(MODEL) is None - assert bridge.rust_tokenizer(CL100K_MODEL) is None - assert bridge.rust_tokenizer(O200K_MODEL) is None + tokenizer: Final = bridge.rust_tokenizer(MODEL) + assert tokenizer == bridge.RustTokenizer("anthropic", "", True, False) PARITY_REQUESTS: Final[tuple[dict[str, object], ...]] = ( {"model": MODEL, "messages": [{"role": "user", "content": "Hello, how are you today?"}]}, - { - "model": MODEL, - "messages": [ - {"role": "system", "content": "You are terse."}, - {"role": "user", "name": "bob", "content": [{"type": "text", "text": "Summarize this."}]}, - {"role": "assistant", "content": "Sure."}, - ], - }, - { - "model": MODEL, - "messages": [{"role": "user", "content": "weather in sf?"}], - "tools": [ - { - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather", - "parameters": { - "type": "object", - "properties": { - "city": {"type": "string", "description": "City"}, - "unit": {"type": "string", "enum": ["c", "f"]}, - }, - "required": ["city"], - }, - }, - } - ], - "tool_choice": {"type": "function", "function": {"name": "get_weather"}}, - }, - { - "model": MODEL, - "messages": [{"role": "user", "content": "x " * 500}], - }, - { - "model": MODEL, - "messages": [ - { - "role": "user", - "content": "I'VE got 1234567 things; it's \"fine\"...\r\n\r\n caf\u00e9 \u0645\u0631\u062d\u0628\u0627 \U0001f600 <|endoftext|>", - } - ], - }, - {"model": MODEL, "prompt": "Write a haiku about ships.", "max_tokens": 20}, {"model": MODEL, "prompt": ["first prompt", "second prompt"]}, - { - "model": MODEL, - "instructions": "be terse", - "input": [ - {"role": "user", "content": [{"type": "input_text", "text": 'Summarise caf\u00e9 menus \u2014 "ok"?\n'}]}, - {"role": "assistant", "content": "Sure."}, - ], - }, {"model": MODEL, "input": "a single embedding string"}, - {"model": MODEL, "input": [[101, 2023, 5], [7]], "encoding_format": "float"}, - {"model": MODEL, "query": "best harbour", "documents": ["doc one", {"text": "doc two", "title": "T", "n": 3}]}, - {"model": MODEL, "messages": None, "prompt": "messages key wins even when null"}, - {"prompt": "model comes from the route"}, + {"model": MODEL, "query": "best harbour", "documents": ["doc one", {"text": "doc two"}]}, ) - -PARITY_MODELS: Final[tuple[tuple[str, bridge.RustTokenizer], ...]] = ( - (MODEL, "anthropic"), - (CL100K_MODEL, "cl100k_base"), - (O200K_MODEL, "o200k_base"), - ("gpt-5", "o200k_base"), -) +PARITY_MODELS: Final = ((MODEL, ANTHROPIC), (CL100K_MODEL, CL100K), (O200K_MODEL, O200K)) +@pytest.mark.requires_rust_extension @pytest.mark.asyncio @pytest.mark.parametrize(("model", "tokenizer"), PARITY_MODELS) @pytest.mark.parametrize("request_body", PARITY_REQUESTS) async def test_native_count_matches_python_budget_counter( - monkeypatch: pytest.MonkeyPatch, request_body: dict[str, object], model: str, tokenizer: bridge.RustTokenizer + monkeypatch: pytest.MonkeyPatch, + request_body: dict[str, object], + model: str, + tokenizer: bridge.RustTokenizer, ) -> None: - native: Final = pytest.importorskip("litellm.rust_bridge._native") - monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - litellm.rust(True) - body: Final = json.dumps(request_body).replace(MODEL, model) + from litellm.proxy.spend_tracking.budget_reservation import _count_input_tokens + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + bridge.TOKEN_COUNTER.reset() + body: Final = json.dumps(request_body).replace(MODEL, model) rust_count: Final = await bridge.count_input_tokens(body.encode(), tokenizer) python_count: Final = _count_input_tokens(request_body=json.loads(body), model=model) - assert rust_count is not None - assert rust_count.model == json.loads(body).get("model") assert rust_count.input_tokens == python_count +@pytest.mark.requires_rust_extension @pytest.mark.asyncio -@pytest.mark.parametrize(("model", "tokenizer"), ((CL100K_MODEL, "cl100k_base"), (O200K_MODEL, "o200k_base"))) -async def test_tiktoken_counts_long_text_exactly_where_python_chunks( - monkeypatch: pytest.MonkeyPatch, model: str, tokenizer: bridge.RustTokenizer +async def test_native_declines_unsupported_request_without_loading_resources( + monkeypatch: pytest.MonkeyPatch, ) -> None: - """Python encodes tiktoken text in fixed-size chunks (drift of up to one token per chunk boundary); Rust does not.""" - native: Final = pytest.importorskip("litellm.rust_bridge._native") monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - litellm.rust(True) - text: Final = "x " * 20_000 - body: Final = {"model": model, "messages": [{"role": "user", "content": text}]} - encoding: Final = tiktoken.get_encoding(tokenizer) - exact: Final = 3 + len(encoding.encode("user")) + len(encoding.encode(text)) + 3 - chunks: Final = -(-len(text) // TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS) - - rust_count: Final = await bridge.count_input_tokens(json.dumps(body).encode(), tokenizer) - python_count: Final = _count_input_tokens(request_body=body, model=model) - - assert rust_count is not None - assert rust_count.input_tokens == exact - assert python_count is not None - assert exact < python_count <= exact + chunks - - -DECLINED_REQUESTS: Final[tuple[dict[str, object], ...]] = ( - { - "model": MODEL, - "messages": [ - {"role": "user", "content": [{"type": "image_url", "image_url": {"url": "data:image/png;base64,AA"}}]} - ], - }, - {"model": MODEL, "prompt": 1.5}, - {"model": MODEL, "documents": [{"score": 0.5}]}, - {"model": MODEL, "file": "audio.mp3"}, -) - - -@pytest.mark.asyncio -@pytest.mark.parametrize("tokenizer", TOKENIZERS) -@pytest.mark.parametrize("request_body", DECLINED_REQUESTS) -async def test_native_declines_shapes_python_prices_differently( - monkeypatch: pytest.MonkeyPatch, request_body: dict[str, object], tokenizer: bridge.RustTokenizer -) -> None: - native: Final = pytest.importorskip("litellm.rust_bridge._native") - monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - litellm.rust(True) - - assert await bridge.count_input_tokens(json.dumps(request_body).encode(), tokenizer) is None + bridge.TOKEN_COUNTER.reset() + assert await bridge.count_input_tokens(b'{"input":1.5}', ANTHROPIC) is None diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 1b0fcfacd1d..7cd92107ae9 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -1,13 +1,28 @@ import importlib +from collections.abc import Iterator +from types import ModuleType +from typing import Final import pytest import litellm +from litellm.exceptions import APIError from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch +from litellm.rust_bridge import configuration +from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") +@pytest.fixture(autouse=True) +def reset_bridge() -> Iterator[None]: + configuration.reset_rust_configuration() + rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) + yield + configuration.reset_rust_configuration() + rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) + + class SyncBridge: def __init__(self) -> None: self.calls: list[dict[str, object]] = [] @@ -42,7 +57,9 @@ class AsyncBridge: return {"text": "async"} -def test_enabled_sync_bridge_receives_audio() -> None: +@pytest.mark.parametrize("enabled", (False, True)) +def test_enabled_sync_bridge_receives_audio(enabled: bool) -> None: + configuration.rust(enabled) bridge = SyncBridge() rust_bridge.configure_rust_transcription(transcription=bridge) result = rust_bridge.transcription( @@ -60,7 +77,9 @@ def test_enabled_sync_bridge_receives_audio() -> None: @pytest.mark.asyncio -async def test_enabled_async_bridge() -> None: +@pytest.mark.parametrize("enabled", (False, True)) +async def test_enabled_async_bridge(enabled: bool) -> None: + configuration.rust(enabled) rust_bridge.configure_rust_transcription(atranscription=AsyncBridge()) result = await rust_bridge.atranscription( model="mistral.voxtral-mini-3b-2507", @@ -78,14 +97,20 @@ async def test_enabled_async_bridge() -> None: def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) - assert rust_bridge.load_rust_transcription() is None - assert rust_bridge.load_rust_atranscription() is None + assert ( + rust_bridge.load_rust_transcription(context=configuration.CapabilityContext(provider="openai", model="test")) + is None + ) + assert ( + rust_bridge.load_rust_atranscription(context=configuration.CapabilityContext(provider="openai", model="test")) + is None + ) def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) - with pytest.raises(RuntimeError, match="bridge is unavailable"): + with pytest.raises(RustRouteUnavailableError, match="bridge is unavailable"): BedrockAudioTranscriptionRustDispatch().audio_transcriptions( model="bedrock/mistral.voxtral-mini-3b-2507", audio_file=("audio.wav", b"audio", "audio/wav"), @@ -100,12 +125,9 @@ def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> @pytest.mark.asyncio async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - async def unavailable(**_: object) -> None: - return None + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) - monkeypatch.setattr(rust_bridge, "atranscription", unavailable) - - with pytest.raises(RuntimeError, match="bridge is unavailable"): + with pytest.raises(RustRouteUnavailableError, match="bridge is unavailable"): await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions( model="bedrock/mistral.voxtral-mini-3b-2507", audio_file=("audio.wav", b"audio", "audio/wav"), @@ -149,3 +171,97 @@ async def test_bedrock_atranscription_uses_rust_only_path() -> None: rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) assert response.text == "rust" + + +class RustBridgeDeclined(Exception): + pass + + +class RustUpstreamError(Exception): + pass + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("error", "expected", "message"), + ( + (RustBridgeDeclined("unsupported model"), RustRouteDeclinedError, "declined the request: unsupported model"), + (RustUpstreamError(429, "rate limited"), APIError, "rate limited"), + ), +) +async def test_bedrock_transcription_errors_never_fall_back( + monkeypatch: pytest.MonkeyPatch, error: Exception, expected: type[Exception], message: str +) -> None: + native: Final = ModuleType("native") + setattr(native, "RustBridgeDeclined", RustBridgeDeclined) + setattr(native, "RustUpstreamError", RustUpstreamError) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: native) + + def fail(**_: object) -> dict[str, object]: + raise error + + async def afail(**_: object) -> dict[str, object]: + raise error + + rust_bridge.configure_rust_transcription(transcription=fail, atranscription=afail) + with pytest.raises(expected, match=message): + rust_bridge.transcription( + model="model", + audio={}, + api_key=None, + api_base=None, + custom_llm_provider="bedrock", + extra_headers=None, + optional_params={}, + timeout=None, + ) + with pytest.raises(expected, match=message): + await rust_bridge.atranscription( + model="model", + audio={}, + api_key=None, + api_base=None, + custom_llm_provider="bedrock", + extra_headers=None, + optional_params={}, + timeout=None, + ) + + +@pytest.mark.asyncio +async def test_python_transcription_skips_rust_when_enabled() -> None: + configuration.rust(True) + + def unexpected(**_: object) -> dict[str, object]: + pytest.fail("Python provider must not call Rust") + + async def aunexpected(**_: object) -> dict[str, object]: + pytest.fail("Python provider must not call Rust") + + rust_bridge.configure_rust_transcription(transcription=unexpected, atranscription=aunexpected) + assert ( + rust_bridge.transcription( + model="model", + audio={}, + api_key=None, + api_base=None, + custom_llm_provider="openai", + extra_headers=None, + optional_params={}, + timeout=None, + ) + is None + ) + assert ( + await rust_bridge.atranscription( + model="model", + audio={}, + api_key=None, + api_base=None, + custom_llm_provider="openai", + extra_headers=None, + optional_params={}, + timeout=None, + ) + is None + ) diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 77d9ef167d0..2c12f10e310 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -594,7 +594,9 @@ def test_unstarted_native_coroutine_releases_input_without_reading_file(ocr_serv def create(): file: Final = File() kwargs: Final = {"model": "mistral/mistral-ocr-latest", "document": {"type": "file", "file": file}} - coroutine: Final = _native._ocr_lifecycle(_public_request("aocr", (), kwargs), (), kwargs, True) + from litellm.rust_bridge.ocr.host import HOST + + coroutine: Final = _native._ocr_lifecycle(_public_request("aocr", (), kwargs), (), kwargs, True, HOST) file.owner = coroutine coroutine.close() return weakref.ref(file) diff --git a/tests/test_litellm_rust/test_route_foundation.py b/tests/test_litellm_rust/test_route_foundation.py index f8051746b78..a7b483854f3 100644 --- a/tests/test_litellm_rust/test_route_foundation.py +++ b/tests/test_litellm_rust/test_route_foundation.py @@ -5,24 +5,36 @@ from typing import Final import pytest from litellm.rust_bridge import _native -from litellm.rust_bridge.configuration import RouteName -from litellm.rust_bridge.route import NativeRoute -from litellm.rust_bridge.runtime import BridgeErrorContext, FallbackMode, invoke +from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.catalog import NATIVE_EXPORTS +from litellm.rust_bridge.chat_completions.lifecycle import LIFECYCLE as CHAT_COMPLETIONS +from litellm.rust_bridge.configuration import ExecutionDecision, RouteName +from litellm.rust_bridge.embeddings.lifecycle import LIFECYCLE as EMBEDDINGS +from litellm.rust_bridge.image_edit.lifecycle import LIFECYCLE as IMAGE_EDIT +from litellm.rust_bridge.image_generation.lifecycle import LIFECYCLE as IMAGE_GENERATION +from litellm.rust_bridge.messages.lifecycle import LIFECYCLE as MESSAGES +from litellm.rust_bridge.moderation.lifecycle import LIFECYCLE as MODERATION +from litellm.rust_bridge.rerank.lifecycle import LIFECYCLE as RERANK +from litellm.rust_bridge.responses.lifecycle import LIFECYCLE as RESPONSES +from litellm.rust_bridge.route import ComponentExecution, NativeLifecycle +from litellm.rust_bridge.runtime import BridgeErrorContext, invoke +from litellm.rust_bridge.speech.lifecycle import LIFECYCLE as SPEECH +from litellm.rust_bridge.transcription.lifecycle import LIFECYCLE as TRANSCRIPTION pytestmark = pytest.mark.requires_rust_extension -UNIMPLEMENTED: Final = ( - RouteName.MESSAGES, - RouteName.CHAT_COMPLETIONS, - RouteName.TRANSCRIPTION, - RouteName.EMBEDDINGS, - RouteName.RERANK, - RouteName.IMAGE_GENERATION, - RouteName.IMAGE_EDIT, - RouteName.SPEECH, - RouteName.MODERATION, - RouteName.RESPONSES, -) +UNIMPLEMENTED: Final[dict[RouteName, NativeBinding[NativeLifecycle[object, object]]]] = { + RouteName.MESSAGES: MESSAGES, + RouteName.CHAT_COMPLETIONS: CHAT_COMPLETIONS, + RouteName.TRANSCRIPTION: TRANSCRIPTION, + RouteName.EMBEDDINGS: EMBEDDINGS, + RouteName.RERANK: RERANK, + RouteName.IMAGE_GENERATION: IMAGE_GENERATION, + RouteName.IMAGE_EDIT: IMAGE_EDIT, + RouteName.SPEECH: SPEECH, + RouteName.MODERATION: MODERATION, + RouteName.RESPONSES: RESPONSES, +} class UntouchedInput: @@ -30,30 +42,102 @@ class UntouchedInput: raise AssertionError(f"unimplemented route inspected {name}") -@pytest.mark.parametrize("route_name", UNIMPLEMENTED) +def test_catalog_exports_are_registered() -> None: + assert all(hasattr(_native, export) for export in NATIVE_EXPORTS) + + +@pytest.mark.parametrize(("route_name", "binding"), tuple(UNIMPLEMENTED.items())) @pytest.mark.parametrize("asynchronous", (False, True)) -def test_unimplemented_lifecycle_declines_without_input_reads(route_name: RouteName, asynchronous: bool) -> None: - route: Final = NativeRoute(route_name) - native: Final = route.select(route.lifecycle()) +def test_package_lifecycle_binding_declines_without_input_reads( + route_name: RouteName, + binding: NativeBinding[NativeLifecycle[object, object]], + asynchronous: bool, +) -> None: + native: Final = binding.load() assert native is not None request: Final = UntouchedInput() with pytest.raises(_native.RustBridgeDeclined, match=f"^{route_name.value} native lifecycle is not implemented$"): - native(request, (request,), {"callback": request, "file": request}, asynchronous) + native(request, (request,), {"callback": request, "file": request}, asynchronous, request) -@pytest.mark.parametrize("route_name", UNIMPLEMENTED) -def test_stub_decline_enters_fallback_once(route_name: RouteName) -> None: - route: Final = NativeRoute(route_name) - native: Final = route.select(route.lifecycle()) +@pytest.mark.parametrize(("route_name", "binding"), tuple(UNIMPLEMENTED.items())) +def test_package_stub_decline_selects_python( + route_name: RouteName, + binding: NativeBinding[NativeLifecycle[object, object]], +) -> None: + native: Final[NativeLifecycle[object, object] | None] = binding.load() assert native is not None - fallback_results: Final = iter(("python result",)) + execution: Final = ComponentExecution( + route_name=route_name, + decision=ExecutionDecision.RUST_WITH_FALLBACK, + ) result: Final = invoke( - native_call=lambda: native(UntouchedInput(), (), {}, False), - fallback=lambda: next(fallback_results), + execution=execution, + native_call=lambda: native(UntouchedInput(), (), {}, False, UntouchedInput()), + python_fallback=lambda: "python", adapt=str, - mode=FallbackMode.PYTHON, context=BridgeErrorContext(route=route_name.value, model="unused", provider="unused"), ) - assert result == "python result" - with pytest.raises(StopIteration): - next(fallback_results) + assert result == "python" + + +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("route", ("messages", "chat_completions", "transcription", "ocr")) +def test_value_admission_declines_unsupported_provider_before_credentials(route: str, asynchronous: bool) -> None: + binding: Final = getattr(_native, ("a" if asynchronous else "") + route) + payload: Final = [{"role": "user", "content": "hi"}] if route == "chat_completions" else {} + with pytest.raises(_native.RustBridgeDeclined): + binding("model", payload, custom_llm_provider="unsupported", api_base="http://127.0.0.1:1") + + +@pytest.mark.parametrize("asynchronous", (False, True)) +def test_messages_declines_required_host_hook_before_preparation(asynchronous: bool) -> None: + binding: Final = _native.amessages if asynchronous else _native.messages + with pytest.raises(_native.RustBridgeDeclined, match="host operations"): + binding("model", {}, custom_llm_provider="anthropic", has_agentic_hook=True) + + +@pytest.mark.parametrize( + ("provider", "facts", "headers"), + ( + ("anthropic", {"stream": True}, {}), + ("anthropic", {"anthropic_user_id": True}, {}), + ("bedrock", {"bedrock_metadata_owned": True}, {}), + ("bedrock", {}, {"x-amz-date": "forwarded"}), + ), +) +def test_chat_entrypoints_decline_before_the_host_callback(provider: str, facts: dict, headers: dict) -> None: + calls: Final[list[bool]] = [] + for binding in (_native.chat_completions, _native.achat_completions): + with pytest.raises(_native.RustBridgeDeclined): + binding( + "model", + [{"role": "user", "content": "hi"}], + custom_llm_provider=provider, + host_facts=facts, + extra_headers=headers, + on_request=lambda: calls.append(True), + ) + assert calls == [] + + +def test_transcription_declines_audio_format_before_credentials() -> None: + with pytest.raises(_native.RustBridgeDeclined, match="audio format"): + _native.transcription("model", {"format": "unsupported", "data": "YQ=="}, custom_llm_provider="bedrock") + + +def test_websocket_declines_before_parsing_or_dialing_url() -> None: + with pytest.raises(_native.RustBridgeDeclined): + _native.ResponsesWebSocketConnection.connect("not a URL", custom_llm_provider="azure") + + +def test_tokenizer_initialization_unavailability_is_not_a_request_decline() -> None: + with pytest.raises(_native.RustBridgeUnavailable): + _native.count_input_tokens( + b'{"input":"hello"}', + "anthropic", + "", + False, + False, + lambda _tokenizer: "{}", + ) diff --git a/uv.lock b/uv.lock index eb4cdef76f1..6364c794be0 100644 --- a/uv.lock +++ b/uv.lock @@ -10,7 +10,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-09-09T21:39:49.468411Z" +exclude-newer = "0001-01-01T00:00:00Z" # This has no effect and is included for backwards compatibility when using relative exclude-newer values. exclude-newer-span = "P3D" [manifest] @@ -4356,6 +4356,105 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/41/a0/b91504515c1f9a299fc157967ffbd2f0321bce0516a3d5b89f6f4cad0355/lazy_object_proxy-1.12.0-pp39.pp310.pp311.graalpy311-none-any.whl", hash = "sha256:c3b2e0af1f7f77c4263759c4824316ce458fabe0fceadcd24ef8ca08b2d1e402", size = 15072, upload-time = "2025-08-22T13:50:05.498Z" }, ] +[[package]] +name = "librt" +version = "0.15.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/36/9b/356320fbae2ac8467e21c5e73e1389c80468e4998c62cc7d3536cc51b614/librt-0.15.0.tar.gz", hash = "sha256:4e66cbe84437497d951b799d3e1551291b6fb3d643820a7014b3655d57a59162", size = 214338, upload-time = "2026-08-07T10:49:42.663Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/48/12/e2e9ca532cf5a0e08c9489826c4a35c6958c92ba0313fda70e8c6c3912be/librt-0.15.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:e1a49adf16a7c9d9646816c2946135527197b6fcf4347c7b8b761cf1bfbf4489", size = 148673, upload-time = "2026-08-07T10:46:22.569Z" }, + { url = "https://files.pythonhosted.org/packages/6d/7c/02005e23478bd5950618d9712e0fd2b4c511657857f3efd8ba6a5feabcdd/librt-0.15.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:81a398f45b45a59200e13cd5ad1ae1d3f44334de98b148331afe2cdfee701c52", size = 153547, upload-time = "2026-08-07T10:46:23.931Z" }, + { url = "https://files.pythonhosted.org/packages/a0/90/d8848a735f5642077fc4b3b4bebcdb08edf10178e3add45597f5201a368f/librt-0.15.0-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:4eafbaff06b9563f8b1c850621ce51605de05208e09d4d71ce490bc972b7b9e8", size = 494355, upload-time = "2026-08-07T10:46:25.122Z" }, + { url = "https://files.pythonhosted.org/packages/e1/0b/8604f41ea02feace490e9e405a338a15f9905369f55b239a9ce31c946f24/librt-0.15.0-cp310-cp310-manylinux2014_i686.manylinux_2_17_i686.manylinux_2_28_i686.whl", hash = "sha256:b0411b4066db926b80258c60dcb0e6db4c9cee312eab45b7e8866b17ddf9ada1", size = 485459, upload-time = "2026-08-07T10:46:26.447Z" }, + { url = "https://files.pythonhosted.org/packages/a0/ac/84153bda1ce0da609182527ab92b40d961809e544eefdc5a1c2422971416/librt-0.15.0-cp310-cp310-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:febb1ce6cac545a54e6b769982824e955a700fdd9fbf3a08a3d82c990968b57d", size = 498398, upload-time = "2026-08-07T10:46:27.701Z" }, + { url = "https://files.pythonhosted.org/packages/2c/3a/5ca6cd282b2c244bec8ec84102e09773264e9c02891d56ab3a8f0e4d7083/librt-0.15.0-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b230acc1c3bfe2d6f2627ba2b95dc92e58aa494600e9722d0e6ccbc931e59702", size = 515474, upload-time = "2026-08-07T10:46:28.9Z" }, + { url = "https://files.pythonhosted.org/packages/73/d3/bd34110234779eb843c6ed66aba7c9b2091d3dd85989f1fb9922f564cb7a/librt-0.15.0-cp310-cp310-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:6da110e5f314c19ab8478464d02ae18808ae73d522c15260fa4918acdcd64da9", size = 509484, upload-time = "2026-08-07T10:46:30.124Z" }, + { url = "https://files.pythonhosted.org/packages/1b/6c/43c3f7f071d71631a7daa3b835ef2168ea39f20692d81464d4e47fbaa6d6/librt-0.15.0-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:eab9208b00ca55bf75983ec99f7bf13acc746a36102e98953addaad7f7ea1e1b", size = 532534, upload-time = "2026-08-07T10:46:31.511Z" }, + { url = "https://files.pythonhosted.org/packages/c5/1c/b854adf036ea817c40408873a5b794d65a91d9f0f39826f2ad2a2d5d7f48/librt-0.15.0-cp310-cp310-musllinux_1_2_i686.whl", hash = "sha256:6c013cd3a1721e69e14380ada97eaa4b7b0cdf1c6b96fa765d4ea47c875088db", size = 537087, upload-time = "2026-08-07T10:46:32.734Z" }, + { url = "https://files.pythonhosted.org/packages/25/5c/c9a890e244e7dd725d3bd8b560e41f0aec787eaf343b46956a290ab7b841/librt-0.15.0-cp310-cp310-musllinux_1_2_ppc64le.whl", hash = "sha256:567b1c430f8bd560e689421468278ac5941bab4a05303b5d95b6ae10db03f451", size = 536575, upload-time = "2026-08-07T10:46:33.965Z" }, + { url = "https://files.pythonhosted.org/packages/5f/c5/c8e70b60b704299555f55db468eb46b1c81bfc60201ffbfe20407d89870c/librt-0.15.0-cp310-cp310-musllinux_1_2_riscv64.whl", hash = "sha256:29c4cab9df457b19672c39be7f384ebb2bc925c4e2684b8780c222b43eb36389", size = 517142, upload-time = "2026-08-07T10:46:35.577Z" }, + { url = "https://files.pythonhosted.org/packages/56/d1/767a90c41f5d381b3195bc88ac0ec4afda35777c9c781e1f9848fedd965e/librt-0.15.0-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:bccbd8e5b0bffb7106cf18eb1baa3d7194b1cebb3b4b1cdbd4bdb19382a6ee6c", size = 558714, upload-time = "2026-08-07T10:46:36.829Z" }, + { url = "https://files.pythonhosted.org/packages/f9/b4/3c0624b8dc8301ab808f2b3a910995bcabe28df070fb9a0e5505ae997dae/librt-0.15.0-cp310-cp310-win32.whl", hash = "sha256:8ae493ed5f659a7761c43d42f183db514536073ded9bcf671d2d1df47e29a07e", size = 104426, upload-time = "2026-08-07T10:46:38.594Z" }, + { url = "https://files.pythonhosted.org/packages/31/98/e91c0382304bedb2db9c6801897319a9dcb68daac5e975819b562362f20d/librt-0.15.0-cp310-cp310-win_amd64.whl", hash = "sha256:bc25fb356d0c7810bb49ff3df908ad1fda6995d660ab099ded69244ed7ab6053", size = 125057, upload-time = "2026-08-07T10:46:40.052Z" }, + { url = "https://files.pythonhosted.org/packages/59/52/06790ced2ac7117f890c21bda43c39c958ec82aa665c0718e821d33ff939/librt-0.15.0-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:823b92cf3c18ecd08afc70c42473888b41b6e8ef5046f3b82c05c154a2fa3d22", size = 148039, upload-time = "2026-08-07T10:46:41.165Z" }, + { url = "https://files.pythonhosted.org/packages/e7/1d/8e150b7fc449a1f33c8a760965cc1f43b14fc1577d9d0b50ab2701420e74/librt-0.15.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:c70bc1b602cf59917e8f0c7a2cbc8bcc6fbc14d5486136b00707a79619121d63", size = 153067, upload-time = "2026-08-07T10:46:42.418Z" }, + { url = "https://files.pythonhosted.org/packages/51/87/a162bc5a66a35599dc619ecb215145f4de7d68e886b479b6d12593139f7c/librt-0.15.0-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:814ff83a25b5fce8b9c80c4dd803153fb5c5599fc74db9e022466938368957ef", size = 493087, upload-time = "2026-08-07T10:46:43.657Z" }, + { url = "https://files.pythonhosted.org/packages/e5/3a/aeea1fc620cf48060d3065b37614edbf97043c099d0f50782bc8ca61d897/librt-0.15.0-cp311-cp311-manylinux2014_i686.manylinux_2_17_i686.manylinux_2_28_i686.whl", hash = "sha256:57f5eeb6ad4c180de583b1038e61fe5fbd9796bb69a8a1c1a0c7ddbec4c8c60f", size = 485608, upload-time = "2026-08-07T10:46:45.038Z" }, + { url = "https://files.pythonhosted.org/packages/52/ff/fe571ad416f0856fd0d5578ffc2e6dc531891e586e36b647bcf50569cab8/librt-0.15.0-cp311-cp311-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:82909c8f7eb9952656b65d3147afde4cf8e6d5a991eebc86418b5e65843b0ab8", size = 498723, upload-time = "2026-08-07T10:46:46.35Z" }, + { url = "https://files.pythonhosted.org/packages/0f/e1/7a65eb5dedb1f00aebd948cdd8e17add48bf066cab3514e9daf84ab45a6c/librt-0.15.0-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f779070399f991400fc451719e0ea388eb7de313388bada2c127a35de05f798a", size = 516002, upload-time = "2026-08-07T10:46:47.599Z" }, + { url = "https://files.pythonhosted.org/packages/5f/45/59832b0ebfbd08c2742e6ece372ceb53f18bf1faef5d33c8daf3abebf749/librt-0.15.0-cp311-cp311-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:bac89069bc496ebdf4f79ebb57bbd10d0b214c8454225deb672d91002bd17e18", size = 508607, upload-time = "2026-08-07T10:46:48.873Z" }, + { url = "https://files.pythonhosted.org/packages/ea/0d/37fa73f3b43ebd8259f91ae9102a15e5a54e65d581e48dea72df3e81d7a4/librt-0.15.0-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:e0d00c708fb2f5822b152429b1ac80a58dbbbc3f6c232c4d13a3f7fcf2ea5b4c", size = 530422, upload-time = "2026-08-07T10:46:50.45Z" }, + { url = "https://files.pythonhosted.org/packages/26/02/e046c6fe7a5881ac34623242192f484426ba8a75595fd18f22c53a3f530f/librt-0.15.0-cp311-cp311-musllinux_1_2_i686.whl", hash = "sha256:6c6624fe268625869485553dd7cc1daf30d22558215bb2a4ff16f67a9801a31a", size = 534303, upload-time = "2026-08-07T10:46:51.693Z" }, + { url = "https://files.pythonhosted.org/packages/95/32/d5e6d861ab0366f3edf74f887ab0c9eb9f535aaf01d32b80b4f734daa179/librt-0.15.0-cp311-cp311-musllinux_1_2_ppc64le.whl", hash = "sha256:f56b397858a23dacf35ede366ed2212fdc03a6a57a1ad36468ad6e9dc5fac091", size = 536084, upload-time = "2026-08-07T10:46:52.951Z" }, + { url = "https://files.pythonhosted.org/packages/2a/de/d69d725513fe53fc90c6d7a1f86e4428939bad2fb905b17fe4c18d413dde/librt-0.15.0-cp311-cp311-musllinux_1_2_riscv64.whl", hash = "sha256:4388184646efe2054911c5b00a1077d6d1ee86a95b7e8ba96dc7850a809f3f40", size = 514307, upload-time = "2026-08-07T10:46:54.194Z" }, + { url = "https://files.pythonhosted.org/packages/36/93/f8aded0d6682b4f25820fa86e0690f87f01df9fd7bd09ddb04d9167ad021/librt-0.15.0-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:97335f59082f9fe2ce6c2a9cc6433a0114bbb6cd4d5c09dd76c95c68b9f9a8b0", size = 557686, upload-time = "2026-08-07T10:46:55.443Z" }, + { url = "https://files.pythonhosted.org/packages/74/09/ffeb6bdeb6cd862b4272fddc8ad05f938dd25d020ed517e631813917d80a/librt-0.15.0-cp311-cp311-win32.whl", hash = "sha256:83380ffde38062a2e9bb55d83e74474f6614665528b98a6928720fc006dfffbb", size = 104917, upload-time = "2026-08-07T10:46:56.605Z" }, + { url = "https://files.pythonhosted.org/packages/96/28/7e2313a3ffbf0b4de7ba3da58a09e488507b4bd1ea2b5e69378354a23415/librt-0.15.0-cp311-cp311-win_amd64.whl", hash = "sha256:f75720477ee05d509a310e856cacc8d909adc182f7b91193c207bcc26d7ee6db", size = 125886, upload-time = "2026-08-07T10:46:57.729Z" }, + { url = "https://files.pythonhosted.org/packages/39/9e/04b8c3cde014ef255ee785730425268354543acc38902093a40afa0dc164/librt-0.15.0-cp311-cp311-win_arm64.whl", hash = "sha256:256237037a3ab001ae8d9803b2d43562a4c3aa38739843694349e4d5ebb0fd56", size = 111885, upload-time = "2026-08-07T10:46:58.787Z" }, + { url = "https://files.pythonhosted.org/packages/ba/39/99c25030e782bdfb7a21be8c05254806a2e4bbb05c8d50c2a2130acbfa05/librt-0.15.0-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:e87bc679f86a99aa3b26e3c78eeb821a247c9a28eae48eaafcc32c3bf4c3bb9e", size = 151021, upload-time = "2026-08-07T10:47:00.057Z" }, + { url = "https://files.pythonhosted.org/packages/14/43/f4b1bd1b2888798a1409808889a25ea1ba49eaabce7d681ed27734c2df9d/librt-0.15.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:71599e011ac880e8e45d46047d714871894c7d4ab6f25626f8d4f89da21f368d", size = 155267, upload-time = "2026-08-07T10:47:01.311Z" }, + { url = "https://files.pythonhosted.org/packages/0c/db/3ad9c965c72f1e1d6beeec44ec10a54e17be8ae042fbb4baade16cbadced/librt-0.15.0-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c802434092b769b1d613ed2e13fac15fbfce1934a74bd10283b03c0fae231cd1", size = 503136, upload-time = "2026-08-07T10:47:02.45Z" }, + { url = "https://files.pythonhosted.org/packages/4b/07/5888a6d76acd62ebce66c61b74d94e9370b9c32929f111e487bb6546f8ed/librt-0.15.0-cp312-cp312-manylinux2014_i686.manylinux_2_17_i686.manylinux_2_28_i686.whl", hash = "sha256:5500eeae393a184d14e1f35645962c27129d20c81afa4069e6ef826ebc2b3aaa", size = 496670, upload-time = "2026-08-07T10:47:03.675Z" }, + { url = "https://files.pythonhosted.org/packages/29/39/ab57cc2f5b276156da02bb7f5a8921bada1cb1993ffec99acf811c602c23/librt-0.15.0-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:6ecfc32dfb46fb7b565bcd6abf9412acf978775a998273d22888a6d7953730dd", size = 513688, upload-time = "2026-08-07T10:47:04.981Z" }, + { url = "https://files.pythonhosted.org/packages/a7/b9/bdbb0b648b5c2befb031f4c6f3b1dd857415e8fb492a25a3c764a6681e6c/librt-0.15.0-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:89cc46cfd15022e35084355478c9ac809d90b1152222706ac9a7655ec21df6fa", size = 531904, upload-time = "2026-08-07T10:47:06.211Z" }, + { url = "https://files.pythonhosted.org/packages/93/26/473c2e4b6c104e9e58e27ce95fc8005c8bd4fc36cae4f254371125a92db8/librt-0.15.0-cp312-cp312-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:d5f51401d102c885b9ca509e62c79b1dbff286e1b9b047fde6f763780789356d", size = 524427, upload-time = "2026-08-07T10:47:07.592Z" }, + { url = "https://files.pythonhosted.org/packages/26/60/03b3abb82b41714671b907bf6989b228e31e6a8af52dec82b5b0728dc250/librt-0.15.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:cc30523e3f1a23fb7511cc659834a0d01a1042bb9de359bc1c131cc4ec6c9656", size = 543155, upload-time = "2026-08-07T10:47:08.866Z" }, + { url = "https://files.pythonhosted.org/packages/f2/0e/9bb1f0a4affbd0a1888f4f79dc03ed2a299d9a2c26c59ab2a97dcbf11903/librt-0.15.0-cp312-cp312-musllinux_1_2_i686.whl", hash = "sha256:59fe030d8ae4a57e3fb7756bf35a858de74e04066fc8555c53d0af979132af81", size = 546890, upload-time = "2026-08-07T10:47:10.327Z" }, + { url = "https://files.pythonhosted.org/packages/dc/84/6937a280d461f7de6e031ffb02edc2b7c3c90d49d630565ce8ff27cbc5f2/librt-0.15.0-cp312-cp312-musllinux_1_2_ppc64le.whl", hash = "sha256:5a6526a2a956bbb1e4ae3568c82e650fc99119c66bb011ea60715744955a2b4d", size = 555163, upload-time = "2026-08-07T10:47:11.798Z" }, + { url = "https://files.pythonhosted.org/packages/bc/95/2a2853c1ee014bf102116e7f897a04beeaeb2461b45b79af98bdfb95f1ef/librt-0.15.0-cp312-cp312-musllinux_1_2_riscv64.whl", hash = "sha256:85ea21ec6730194d67156b0e0b5430ccb1d61f8b8b907e39b37f9812b74a13f0", size = 535812, upload-time = "2026-08-07T10:47:13.279Z" }, + { url = "https://files.pythonhosted.org/packages/c9/4c/cf9601c1b4c5f09280acd5d83abdb2e68527a2be8257136eb42304218622/librt-0.15.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:1e47b8ba865d7ede071a91a7163073bbaeb72541f1ef8a07d512c45c7b5007f2", size = 573688, upload-time = "2026-08-07T10:47:14.727Z" }, + { url = "https://files.pythonhosted.org/packages/47/6d/9ac7cbec46189a7625af4b5acbd25f10d827f4141b2002181848c8418923/librt-0.15.0-cp312-cp312-win32.whl", hash = "sha256:a5207ec414d1c4a2a7231b2086970dc036f94293cdf338190984958a013a42f1", size = 106138, upload-time = "2026-08-07T10:47:15.973Z" }, + { url = "https://files.pythonhosted.org/packages/38/d0/2ae99c83be86ce23f925ac1aeeedc777e97f427c4a8d190c70d0a16e9a87/librt-0.15.0-cp312-cp312-win_amd64.whl", hash = "sha256:73b30cfa976659b3917c8f6153bdb0591c6a9ec6583599fd24a689b690622022", size = 126974, upload-time = "2026-08-07T10:47:17.049Z" }, + { url = "https://files.pythonhosted.org/packages/5d/ef/dd24f9635c730b86b87587967dda7516b1845e8b17684603d31607fed598/librt-0.15.0-cp312-cp312-win_arm64.whl", hash = "sha256:a54cf9e0ef47b96af580849db5471142200568ce1e02cbf416addab551369570", size = 112292, upload-time = "2026-08-07T10:47:18.222Z" }, + { url = "https://files.pythonhosted.org/packages/e7/42/467b53a601b406ccd7b97c1fd54b59cb34f9185ad5ce7e9d5c3c4e8961c8/librt-0.15.0-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:db13ca398005abcbe538deda87b686d9bd08b7001cf40c4c06b444960ae10a26", size = 151029, upload-time = "2026-08-07T10:47:19.312Z" }, + { url = "https://files.pythonhosted.org/packages/3e/e6/36c2299b7a94b84fdd01220d8a777a71be5be0925bb0dbdf71c0a06a34d9/librt-0.15.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:aa1f1995789dca3698bc550aaceb09a51bd5df0a057ff84ff15296cd1975b801", size = 155194, upload-time = "2026-08-07T10:47:20.398Z" }, + { url = "https://files.pythonhosted.org/packages/c9/b6/ed5071f9325845e670bd36012757419767fbf56af77ed483077b9e4db541/librt-0.15.0-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:55456ea87d8df21808446d03817be2f65e20391c1c615d9187440dff28cd08dc", size = 502568, upload-time = "2026-08-07T10:47:21.652Z" }, + { url = "https://files.pythonhosted.org/packages/7f/81/6450c67c3615d87704bcbc21323fafc69c799b06a044c447529f725d4b01/librt-0.15.0-cp313-cp313-manylinux2014_i686.manylinux_2_17_i686.manylinux_2_28_i686.whl", hash = "sha256:5a86a5a08c2235316bdb359d5dbb6ce0abfca7fac06363103e2c5af571d92f95", size = 496153, upload-time = "2026-08-07T10:47:22.925Z" }, + { url = "https://files.pythonhosted.org/packages/e1/d6/5f52b722bc75076954b3bfd49be15ea362df4d580c6fb315d0f617100d30/librt-0.15.0-cp313-cp313-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:e56b6a368529bed262da40ce13f8fef590db0479819cca84f16a1f01ac356d0b", size = 513336, upload-time = "2026-08-07T10:47:24.213Z" }, + { url = "https://files.pythonhosted.org/packages/8d/e2/c08fd1d36ce63ea5a12b85c5d37f4550b5f86a692167e41e5a74222607ae/librt-0.15.0-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:234d8d394721fa0d786af15ebf1f3fb7f3ed82fd1cd0cde45c2f247b5d4281d2", size = 531661, upload-time = "2026-08-07T10:47:25.507Z" }, + { url = "https://files.pythonhosted.org/packages/3f/d8/d9482fcbeb177b9eb87bb3899eeb3b42be690313c652f9e146b1d0681fb2/librt-0.15.0-cp313-cp313-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:d8363d7accb0286ac3a0e633f396e93800dafb8150494505daf9515bbda591f3", size = 524487, upload-time = "2026-08-07T10:47:26.79Z" }, + { url = "https://files.pythonhosted.org/packages/10/cc/075171517b41f861753034fbb151b42cfc83bcc853849f24f5e66fd60ccf/librt-0.15.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:0f0ee3644d951f31055ad07d77d92520e84505dd7a432cc4cd501dd70ee06785", size = 543201, upload-time = "2026-08-07T10:47:27.999Z" }, + { url = "https://files.pythonhosted.org/packages/b0/03/42c2330f37eeb475b6affeedd06518f60035f323af3a839335e3fc9fef2d/librt-0.15.0-cp313-cp313-musllinux_1_2_i686.whl", hash = "sha256:2cfd1a81a648806e6a7717be4cc4d1bb392fa229752bf8444ba365e381e984d6", size = 546467, upload-time = "2026-08-07T10:47:29.396Z" }, + { url = "https://files.pythonhosted.org/packages/57/1e/1ad4c5638f7e64d8560328bd25c54b409a661bdb6ff254b38ff90744288d/librt-0.15.0-cp313-cp313-musllinux_1_2_ppc64le.whl", hash = "sha256:a6cd22c9da0d866558e46a041f1cc0c2bbb26b61b137b2347fa834c332e1d101", size = 555139, upload-time = "2026-08-07T10:47:30.815Z" }, + { url = "https://files.pythonhosted.org/packages/49/41/39fa7d15db1204cd1cbe6514680fbdc243adf754a0885061308f43afc013/librt-0.15.0-cp313-cp313-musllinux_1_2_riscv64.whl", hash = "sha256:6d5225ef8801e4ea5e482fa9b5dfb891dd9ef6f6d870f1f25d449ca2c70ac218", size = 536050, upload-time = "2026-08-07T10:47:32.222Z" }, + { url = "https://files.pythonhosted.org/packages/1e/88/c6dcf0dd8e26dc0c9a499a2abab8646c86dcaf9ecea9524cb46d3686331a/librt-0.15.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:6d28a05796b99f749bf8794f17ba9ba1612d0076b802e9cfc62c554634e9ce3b", size = 573700, upload-time = "2026-08-07T10:47:33.527Z" }, + { url = "https://files.pythonhosted.org/packages/1b/9b/ab54c71a7918a7c34fa5327fb61390a77446a07a146fbfb1165250a61035/librt-0.15.0-cp313-cp313-pyemscripten_2025_0_wasm32.whl", hash = "sha256:2067ff438048cead9d223ca5675bae2a25e520a7c3e6c1498bf9c6892d22caab", size = 82194, upload-time = "2026-08-07T10:47:34.835Z" }, + { url = "https://files.pythonhosted.org/packages/8d/b2/4f9a243bb892395f3becb80789ade13771701091f9f07ab8230247953ba8/librt-0.15.0-cp313-cp313-win32.whl", hash = "sha256:1cd3b721f24c206398b9e26da3c3a9c011e6e89d06f318ba8ebefc30f1003890", size = 106231, upload-time = "2026-08-07T10:47:36.251Z" }, + { url = "https://files.pythonhosted.org/packages/bf/af/64aff4885a40b93132382f2c314647d722574605416504379184ef3045ea/librt-0.15.0-cp313-cp313-win_amd64.whl", hash = "sha256:f395a4a9a03ac062dbe9a9f82e0c720502e590a38feee6a757bc82e9c63afbd8", size = 126996, upload-time = "2026-08-07T10:47:37.453Z" }, + { url = "https://files.pythonhosted.org/packages/27/83/335bccf6c7cb9028cb0b54aead27d9ece3f01f83bc6baa2abace5da655c1/librt-0.15.0-cp313-cp313-win_arm64.whl", hash = "sha256:0a15cb554761247d84a3ec0cbdf4078d70725384f0e4662c0fa3b26266eb60ad", size = 112188, upload-time = "2026-08-07T10:47:38.729Z" }, + { url = "https://files.pythonhosted.org/packages/a8/93/949053fb462eecc4a9a5ee770a81f4b40be7b79538b245545d4aebc6b58b/librt-0.15.0-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:f5de7feedc56337a088eb15cd9fafa9938367362221d8cc62c642b7f94821993", size = 149833, upload-time = "2026-08-07T10:47:39.86Z" }, + { url = "https://files.pythonhosted.org/packages/61/ca/8281aa6cd560a3420e4497729f6b704b53be3eeaaef82d5aeadddaf7441f/librt-0.15.0-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:6c0eb900c0e91f4aebe680845242e614f1864edfd44106380d0752ac29522bf8", size = 154088, upload-time = "2026-08-07T10:47:41.065Z" }, + { url = "https://files.pythonhosted.org/packages/dd/02/1a1662dceaba6a086360891448d5ce9a7d3555976cae59a31a39d744b9c7/librt-0.15.0-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:e8c9a650a188e38bac005048cbe6342e81407782944d01934540ab75e417df21", size = 494215, upload-time = "2026-08-07T10:47:42.388Z" }, + { url = "https://files.pythonhosted.org/packages/69/84/99211619dc656370a3740c33d2b0b6d5a3fb1e73689314f6ed477a397dc4/librt-0.15.0-cp314-cp314-manylinux2014_i686.manylinux_2_17_i686.manylinux_2_28_i686.whl", hash = "sha256:92bfed8deec93df30286b9fe9e3b1dd17329cc076a192b4ee5ec223841d54953", size = 491173, upload-time = "2026-08-07T10:47:43.683Z" }, + { url = "https://files.pythonhosted.org/packages/d4/aa/5448d0b05f4579b635d3899176817ebf561af0e57bacd425b5b1887264c1/librt-0.15.0-cp314-cp314-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:ec4b19788f835711a2072f9dbe6b03b3bf32ed1f0fb30cf399bdd59d9f0c33fa", size = 505512, upload-time = "2026-08-07T10:47:45.314Z" }, + { url = "https://files.pythonhosted.org/packages/95/82/01940e40b83c43a546c4a3c896cf34ca272a9690899d55914e4827b3dcce/librt-0.15.0-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:d4c7bacb70930f3d0a56f4ecf1be474a1f0d941b01dd73b756f3c256d42cb879", size = 523073, upload-time = "2026-08-07T10:47:46.66Z" }, + { url = "https://files.pythonhosted.org/packages/88/fa/759c0030f3ee371439eb26de34fc745807caf0abb878af7af4b8b7c3dd3d/librt-0.15.0-cp314-cp314-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:3e79f05e4a08b4d880342673312bbc895b56df7765605796f15902eb5367d3ae", size = 515080, upload-time = "2026-08-07T10:47:48.319Z" }, + { url = "https://files.pythonhosted.org/packages/0b/27/894e072228fcb159703c655da69f8cd10dbed489c36e3df7dd032a2483be/librt-0.15.0-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:a417149c0cba4d50b61e992e5a15e69eaf96746609b461cc4ed168aeef6b79dd", size = 534164, upload-time = "2026-08-07T10:47:49.875Z" }, + { url = "https://files.pythonhosted.org/packages/98/a3/0078e91c1f36f8815db17827de15650b9a3fe56c55fbf998c854b34e40d3/librt-0.15.0-cp314-cp314-musllinux_1_2_i686.whl", hash = "sha256:da7a94d6a3411f579d72aa3e3bc5fbca7ed4549f3dbd7e5de3aa567333374285", size = 540616, upload-time = "2026-08-07T10:47:51.408Z" }, + { url = "https://files.pythonhosted.org/packages/86/33/81a29b796dd52a45e9ef7974c7732926e8f10f15b8d2be505665979f896d/librt-0.15.0-cp314-cp314-musllinux_1_2_ppc64le.whl", hash = "sha256:856f743ae607f2c1380eccb566c0038a9fb3eabf0fc2be2704d76d9f73557239", size = 545890, upload-time = "2026-08-07T10:47:52.818Z" }, + { url = "https://files.pythonhosted.org/packages/05/82/8be1baa1350e5d30cfd70ae79d0a6f4dc5862ef47f7bb2808aabc9bb86e5/librt-0.15.0-cp314-cp314-musllinux_1_2_riscv64.whl", hash = "sha256:779a6e7c894737e5983e7790a9c78c4000c30e23c9aada08081bdbea53b0fa60", size = 523287, upload-time = "2026-08-07T10:47:54.165Z" }, + { url = "https://files.pythonhosted.org/packages/c6/4f/d1be6a01a35c20ef734e0e44113f87d4af756a9354a89dcfbe3b4f8af5e1/librt-0.15.0-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:96bb17dbe8bab3c0954fbebfc69ed395599de75b6bbc35e3270a878e15d4dd65", size = 565868, upload-time = "2026-08-07T10:47:55.566Z" }, + { url = "https://files.pythonhosted.org/packages/67/88/649cfa33f5825927b160610f670bdab012a64d627eddb94fa795ea4292fd/librt-0.15.0-cp314-cp314-pyemscripten_2026_0_wasm32.whl", hash = "sha256:7220697efaa6e5348fc3d18ee7f8563d4bfecd9872b37ffb915bfc1d08840622", size = 81619, upload-time = "2026-08-07T10:47:56.886Z" }, + { url = "https://files.pythonhosted.org/packages/22/31/8e88a8d5e48fc8d1a817787fb6811dfff6499acd6c8683dd83934aa6ede0/librt-0.15.0-cp314-cp314-win32.whl", hash = "sha256:f54598964d357b1c5ab77cf5d92f21e598fe0e23cdbe9618480807f81b4eba15", size = 100138, upload-time = "2026-08-07T10:47:58.093Z" }, + { url = "https://files.pythonhosted.org/packages/80/92/20fd6c4b6a1b1a564b076d55cd3d427d8428217d7638dc25a654cc4791d4/librt-0.15.0-cp314-cp314-win_amd64.whl", hash = "sha256:3ff5893a2c23d886aa9ce786de5ac6ddc74aeeaf90743682b74d920e117d2e28", size = 121258, upload-time = "2026-08-07T10:47:59.564Z" }, + { url = "https://files.pythonhosted.org/packages/fc/28/6af430b44d9ebb897b865a3c363b6dcace51357be2347cc0f8f869656a86/librt-0.15.0-cp314-cp314-win_arm64.whl", hash = "sha256:3722a099730704c9a3d70c879fc0f51daec25fe5f1555672d97bc595abeafb95", size = 106467, upload-time = "2026-08-07T10:48:01.097Z" }, + { url = "https://files.pythonhosted.org/packages/7e/aa/b42bb798942ced219f6d63b27e07f91237887a8d0bd0921666db79a13790/librt-0.15.0-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:38c0c7d4b6fc06c3324b3f9162c8391bfc4fd9dde53afe1033ce7edb48d5a714", size = 159523, upload-time = "2026-08-07T10:48:02.442Z" }, + { url = "https://files.pythonhosted.org/packages/75/03/1b53cd4ef904e73b1d828a5f90143bf94a2967d7cfff0b9ccf93e12aa9b4/librt-0.15.0-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:8b2fdd7ead3c995c37940a790690660d0ca006c302db26cc51933f6766866fc3", size = 161638, upload-time = "2026-08-07T10:48:03.725Z" }, + { url = "https://files.pythonhosted.org/packages/ac/c4/9f9c9fba097d49e9e694c2b4dc331df31884645ecbc58a93b4b5fc69d2c5/librt-0.15.0-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:2fde98cf1fc4bac144ce23c2c4c017b924ba714509ea9334977b0b27050c837d", size = 701795, upload-time = "2026-08-07T10:48:05.135Z" }, + { url = "https://files.pythonhosted.org/packages/4c/05/0966840bda0380c8ae167b9043c6230202941cc90ea29c48e096964c765e/librt-0.15.0-cp314-cp314t-manylinux2014_i686.manylinux_2_17_i686.manylinux_2_28_i686.whl", hash = "sha256:e3b461183c5fa7681b48560f91515f53a953122fb30c71e07abc67d7ddf58c38", size = 682147, upload-time = "2026-08-07T10:48:06.555Z" }, + { url = "https://files.pythonhosted.org/packages/18/af/1c47ca573c30ea47d195aec26133af522fea1104afaace028d7b32247ea8/librt-0.15.0-cp314-cp314t-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:4bbcc257e3babea20a91715c361b24554ec4e8f51aa578568afc230799fe1a19", size = 696397, upload-time = "2026-08-07T10:48:08.03Z" }, + { url = "https://files.pythonhosted.org/packages/2e/0f/1aed6223d4f9f9d1171a8596ff100ea4c3f7699fea7a4ba657c3e60daa6c/librt-0.15.0-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b845b8d48088fad0cadc84be4b8fda63203be7e9237b71015b3925443c1f35ab", size = 722542, upload-time = "2026-08-07T10:48:09.569Z" }, + { url = "https://files.pythonhosted.org/packages/c6/22/9e3a929aea456c97d69e6ef3884efea56d4807f97399471cc946baebd8af/librt-0.15.0-cp314-cp314t-manylinux_2_34_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:b30e600e8f337b9bd7f39b86d9fdfedc73cc46e3d0f745931a23a234220bb7e2", size = 729709, upload-time = "2026-08-07T10:48:11.129Z" }, + { url = "https://files.pythonhosted.org/packages/e9/1b/c327ef6018e3a9ca0b8e7c5eddeeb331ba8f9b76c24e126d37d0f6d62faf/librt-0.15.0-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:64b0c8c35aa4c4ed79896359f3e0b285cbe4e610042106500da4811c322cc108", size = 752891, upload-time = "2026-08-07T10:48:12.558Z" }, + { url = "https://files.pythonhosted.org/packages/d7/d1/d5f1ea02c56930087009e39db9b70660a663e76c730b27b925d786718457/librt-0.15.0-cp314-cp314t-musllinux_1_2_i686.whl", hash = "sha256:0da0d94cb802f32a0524653e7201f2cef72d5f700a5407678f5290483d4fcd08", size = 745301, upload-time = "2026-08-07T10:48:14.55Z" }, + { url = "https://files.pythonhosted.org/packages/d9/3c/5f7c585d15ebb2250c73e7c0ee4e9e47be72c65d520c07ddbcdc62037674/librt-0.15.0-cp314-cp314t-musllinux_1_2_ppc64le.whl", hash = "sha256:4a6369168d371207339b1e50d4532b06a7121586141f82599505a3f315751d47", size = 747921, upload-time = "2026-08-07T10:48:16.453Z" }, + { url = "https://files.pythonhosted.org/packages/7f/52/1443a446486eba966bcbca1696b472e4f210320ec42f490a47f48fbf0fdc/librt-0.15.0-cp314-cp314t-musllinux_1_2_riscv64.whl", hash = "sha256:c434e072557ade9cbc642d052c89d031efe47d5c9614523619d0d74a02378e81", size = 727561, upload-time = "2026-08-07T10:48:18.089Z" }, + { url = "https://files.pythonhosted.org/packages/79/91/2270a9380f11725cf83ce1925a5e32dd1dde2be9bba597f25c10a38644e7/librt-0.15.0-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:c7eec6a42018bc1d45763b1c162d3d2bf7c3b9a1b0ed30d3e91dcba390efefcc", size = 774417, upload-time = "2026-08-07T10:48:19.611Z" }, + { url = "https://files.pythonhosted.org/packages/9e/3b/f4b1548d4f5b99186737fe27aec238e9823e8d5d23bf4df007c030689dc5/librt-0.15.0-cp314-cp314t-win32.whl", hash = "sha256:6912fa5e635d74529ac7cdb1bdf6ca3af4453da8d1edbe0110ee1cb4ad407ebf", size = 104381, upload-time = "2026-08-07T10:48:21.048Z" }, + { url = "https://files.pythonhosted.org/packages/80/b6/134afad262def1de04c0843c376d02135f1168af43f22e09a52bd8394727/librt-0.15.0-cp314-cp314t-win_amd64.whl", hash = "sha256:8e11699ed745931c395acd3621b07062e0f840efa6935aad87a64ed0995f0915", size = 127034, upload-time = "2026-08-07T10:48:22.561Z" }, + { url = "https://files.pythonhosted.org/packages/99/5f/1b6846b20572bd699c9e9ec321a5f781845bee477df2aa2a43b28bc40119/librt-0.15.0-cp314-cp314t-win_arm64.whl", hash = "sha256:5d2a91724463bfed4f573cd7a9fdc856d2e230d0c0e5a61416a93481dccd8605", size = 110827, upload-time = "2026-08-07T10:48:23.804Z" }, +] + [[package]] name = "litellm" version = "1.102.0" @@ -4527,6 +4626,7 @@ dev = [ { name = "hypothesis" }, { name = "keyring" }, { name = "langfuse" }, + { name = "mypy" }, { name = "openapi-core" }, { name = "opentelemetry-api" }, { name = "opentelemetry-exporter-otlp" }, @@ -4718,6 +4818,7 @@ dev = [ { name = "hypothesis", specifier = "==6.165.10" }, { name = "keyring", specifier = "==25.7.0" }, { name = "langfuse", specifier = "==2.59.7" }, + { name = "mypy", specifier = "==1.20.1" }, { name = "openapi-core", specifier = "==0.22.0" }, { name = "opentelemetry-api", specifier = "==1.28.0" }, { name = "opentelemetry-exporter-otlp", specifier = "==1.28.0" }, @@ -5635,6 +5736,64 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/81/08/7036c080d7117f28a4af526d794aab6a84463126db031b007717c1a6676e/multidict-6.7.1-py3-none-any.whl", hash = "sha256:55d97cc6dae627efa6a6e548885712d4864b81110ac76fa4e534c03819fa4a56", size = 12319, upload-time = "2026-01-26T02:46:44.004Z" }, ] +[[package]] +name = "mypy" +version = "1.20.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "librt", marker = "platform_python_implementation != 'PyPy'" }, + { name = "mypy-extensions" }, + { name = "pathspec" }, + { name = "tomli", marker = "python_full_version < '3.11'" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/0b/3d/5b373635b3146264eb7a68d09e5ca11c305bbb058dfffbb47c47daf4f632/mypy-1.20.1.tar.gz", hash = "sha256:6fc3f4ecd52de81648fed1945498bf42fa2993ddfad67c9056df36ae5757f804", size = 3815892, upload-time = "2026-04-13T02:46:51.474Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/21/4b/b1fa23297c8a5c403aabaac0649549efc5a0af7095f3dd33e7482863f973/mypy-1.20.1-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:3ba5d1e712ada9c3b6223dcbc5a31dac334ed62991e5caa17bcf5a4ddc349af0", size = 14426426, upload-time = "2026-04-13T02:46:37.828Z" }, + { url = "https://files.pythonhosted.org/packages/22/53/82923480aee5507a46df22428316e28b2b710d08506a128b2acef81ab18e/mypy-1.20.1-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:2e731284c117b0987fb1e6c5013a56f33e7faa1fce594066ab83876183ce1c66", size = 13307651, upload-time = "2026-04-13T02:46:22.676Z" }, + { url = "https://files.pythonhosted.org/packages/4e/0c/91905b393c790440fa273f0903ee2b07cce95bb6deccac87e6eb343d077a/mypy-1.20.1-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:f8e945b872a05f4fbefabe2249c0b07b6b194e5e11a86ebee9edf855de09806c", size = 13746066, upload-time = "2026-04-13T02:45:15.345Z" }, + { url = "https://files.pythonhosted.org/packages/88/b9/8a7017270438e34544e19dd6284cad54fd65dde3c35418a2ce07a1897804/mypy-1.20.1-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:2fc88acef0dc9b15246502b418980478c1bfc9702057a0e1e7598d01a7af8937", size = 14617944, upload-time = "2026-04-13T02:45:44.954Z" }, + { url = "https://files.pythonhosted.org/packages/0c/cf/5a61ceec3fc133e0f559d1e1f9adf4150abdbc2ad8eb831ec26fc8459196/mypy-1.20.1-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:14911a115c73608f155f648b978c5055d16ff974e6b1b5512d7fedf4fa8b15c6", size = 14918205, upload-time = "2026-04-13T02:45:42.653Z" }, + { url = "https://files.pythonhosted.org/packages/6f/80/afb1c665e9c426c78e4711cce04e446b645867bfb97936158886103c1648/mypy-1.20.1-cp310-cp310-win_amd64.whl", hash = "sha256:76d9b4c992cca3331d9793ef197ae360ea44953cf35beb2526e95b9e074f2866", size = 10823344, upload-time = "2026-04-13T02:46:07.607Z" }, + { url = "https://files.pythonhosted.org/packages/11/68/7ad64b49b7663c88fef76a2ac689ea73e17804832ac4cb5416bcff17775b/mypy-1.20.1-cp310-cp310-win_arm64.whl", hash = "sha256:b408722f80be44845da555671a5ef3a0c63f51ca5752b0c20e992dc9c0fbd3cd", size = 9760694, upload-time = "2026-04-13T02:46:49.369Z" }, + { url = "https://files.pythonhosted.org/packages/82/0d/555ab7453cc4a4a8643b7f21c842b1a84c36b15392061ae7b052ee119320/mypy-1.20.1-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:c01eb9bac2c6a962d00f9d23421cd2913840e65bba365167d057bd0b4171a92e", size = 14336012, upload-time = "2026-04-13T02:45:39.935Z" }, + { url = "https://files.pythonhosted.org/packages/57/26/85a28893f7db8a16ebb41d1e9dfcb4475844d06a88480b6639e32a74d6ef/mypy-1.20.1-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:55d12ddbd8a9cac5b276878bd534fa39fff5bf543dc6ae18f25d30c8d7d27fca", size = 13224636, upload-time = "2026-04-13T02:45:49.659Z" }, + { url = "https://files.pythonhosted.org/packages/93/41/bd4cd3c2caeb6c448b669222b8cfcbdee4a03b89431527b56fca9e56b6f3/mypy-1.20.1-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c0aa322c1468b6cdfc927a44ce130f79bb44bcd34eb4a009eb9f96571fd80955", size = 13663471, upload-time = "2026-04-13T02:46:20.276Z" }, + { url = "https://files.pythonhosted.org/packages/3e/56/7ee8c471e10402d64b6517ae10434541baca053cffd81090e4097d5609d4/mypy-1.20.1-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:3f8bc95899cf676b6e2285779a08a998cc3a7b26f1026752df9d2741df3c79e8", size = 14532344, upload-time = "2026-04-13T02:46:44.205Z" }, + { url = "https://files.pythonhosted.org/packages/b5/95/b37d1fa859a433f6156742e12f62b0bb75af658544fb6dada9363918743a/mypy-1.20.1-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:47c2b90191a870a04041e910277494b0d92f0711be9e524d45c074fe60c00b65", size = 14776670, upload-time = "2026-04-13T02:45:52.481Z" }, + { url = "https://files.pythonhosted.org/packages/03/77/b302e4cb0b80d2bdf6bf4fce5864bb4cbfa461f7099cea544eaf2457df78/mypy-1.20.1-cp311-cp311-win_amd64.whl", hash = "sha256:9857dc8d2ec1a392ffbda518075beb00ac58859979c79f9e6bdcb7277082c2f2", size = 10816524, upload-time = "2026-04-13T02:45:37.711Z" }, + { url = "https://files.pythonhosted.org/packages/7f/21/d969d7a68eb964993ebcc6170d5ecaf0cf65830c58ac3344562e16dc42a9/mypy-1.20.1-cp311-cp311-win_arm64.whl", hash = "sha256:09d8df92bb25b6065ab91b178da843dda67b33eb819321679a6e98a907ce0e10", size = 9750419, upload-time = "2026-04-13T02:45:08.542Z" }, + { url = "https://files.pythonhosted.org/packages/69/1b/75a7c825a02781ca10bc2f2f12fba2af5202f6d6005aad8d2d1f264d8d78/mypy-1.20.1-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:36ee2b9c6599c230fea89bbd79f401f9f9f8e9fcf0c777827789b19b7da90f51", size = 14494077, upload-time = "2026-04-13T02:45:55.085Z" }, + { url = "https://files.pythonhosted.org/packages/b0/54/5e5a569ea5c2b4d48b729fb32aa936eeb4246e4fc3e6f5b3d36a2dfbefb9/mypy-1.20.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:fba3fb0968a7b48806b0c90f38d39296f10766885a94c83bd21399de1e14eb28", size = 13319495, upload-time = "2026-04-13T02:45:29.674Z" }, + { url = "https://files.pythonhosted.org/packages/6f/a4/a1945b19f33e91721b59deee3abb484f2fa5922adc33bb166daf5325d76d/mypy-1.20.1-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:ef1415a637cd3627d6304dfbeddbadd21079dafc2a8a753c477ce4fc0c2af54f", size = 13696948, upload-time = "2026-04-13T02:46:15.006Z" }, + { url = "https://files.pythonhosted.org/packages/b2/c6/75e969781c2359b2f9c15b061f28ec6d67c8b61865ceda176e85c8e7f2de/mypy-1.20.1-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:ef3461b1ad5cd446e540016e90b5984657edda39f982f4cc45ca317b628f5a37", size = 14706744, upload-time = "2026-04-13T02:46:00.482Z" }, + { url = "https://files.pythonhosted.org/packages/a8/6e/b221b1de981fc4262fe3e0bf9ec272d292dfe42394a689c2d49765c144c4/mypy-1.20.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:542dd63c9e1339b6092eb25bd515f3a32a1453aee8c9521d2ddb17dacd840237", size = 14949035, upload-time = "2026-04-13T02:45:06.021Z" }, + { url = "https://files.pythonhosted.org/packages/ca/4b/298ba2de0aafc0da3ff2288da06884aae7ba6489bc247c933f87847c41b3/mypy-1.20.1-cp312-cp312-win_amd64.whl", hash = "sha256:1d55c7cd8ca22e31f93af2a01160a9e95465b5878de23dba7e48116052f20a8d", size = 10883216, upload-time = "2026-04-13T02:45:47.232Z" }, + { url = "https://files.pythonhosted.org/packages/c7/f9/5e25b8f0b8cb92f080bfed9c21d3279b2a0b6a601cdca369a039ba84789d/mypy-1.20.1-cp312-cp312-win_arm64.whl", hash = "sha256:f5b84a79070586e0d353ee07b719d9d0a4aa7c8ee90c0ea97747e98cbe193019", size = 9814299, upload-time = "2026-04-13T02:45:21.934Z" }, + { url = "https://files.pythonhosted.org/packages/21/e8/ef0991aa24c8f225df10b034f3c2681213cb54cf247623c6dec9a5744e70/mypy-1.20.1-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:8f3886c03e40afefd327bd70b3f634b39ea82e87f314edaa4d0cce4b927ddcc1", size = 14500739, upload-time = "2026-04-13T02:46:05.442Z" }, + { url = "https://files.pythonhosted.org/packages/23/73/416ebec3047636ed89fa871dc8c54bf05e9e20aa9499da59790d7adb312d/mypy-1.20.1-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:e860eb3904f9764e83bafd70c8250bdffdc7dde6b82f486e8156348bf7ceb184", size = 13314735, upload-time = "2026-04-13T02:46:47.154Z" }, + { url = "https://files.pythonhosted.org/packages/10/1e/1505022d9c9ac2e014a384eb17638fb37bf8e9d0a833ea60605b66f8f7ba/mypy-1.20.1-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:a4b5aac6e785719da51a84f5d09e9e843d473170a9045b1ea7ea1af86225df4b", size = 13704356, upload-time = "2026-04-13T02:45:19.773Z" }, + { url = "https://files.pythonhosted.org/packages/98/91/275b01f5eba5c467a3318ec214dd865abb66e9c811231c8587287b92876a/mypy-1.20.1-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f37b6cd0fe2ad3a20f05ace48ca3523fc52ff86940e34937b439613b6854472e", size = 14696420, upload-time = "2026-04-13T02:45:24.205Z" }, + { url = "https://files.pythonhosted.org/packages/a1/57/b3779e134e1b7250d05f874252780d0a88c068bc054bcff99ca20a3a2986/mypy-1.20.1-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:e4bbb0f6b54ce7cc350ef4a770650d15fa70edd99ad5267e227133eda9c94218", size = 14936093, upload-time = "2026-04-13T02:45:32.087Z" }, + { url = "https://files.pythonhosted.org/packages/be/33/81b64991b0f3f278c3b55c335888794af190b2d59031a5ad1401bcb69f1e/mypy-1.20.1-cp313-cp313-win_amd64.whl", hash = "sha256:c3dc20f8ec76eecd77148cdd2f1542ed496e51e185713bf488a414f862deb8f2", size = 10889659, upload-time = "2026-04-13T02:46:02.926Z" }, + { url = "https://files.pythonhosted.org/packages/1b/fd/7adcb8053572edf5ef8f3db59599dfeeee3be9cc4c8c97e2d28f66f42ac5/mypy-1.20.1-cp313-cp313-win_arm64.whl", hash = "sha256:a9d62bbac5d6d46718e2b0330b25e6264463ed832722b8f7d4440ff1be3ca895", size = 9815515, upload-time = "2026-04-13T02:46:32.103Z" }, + { url = "https://files.pythonhosted.org/packages/40/cd/db831e84c81d57d4886d99feee14e372f64bbec6a9cb1a88a19e243f2ef5/mypy-1.20.1-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:12927b9c0ed794daedcf1dab055b6c613d9d5659ac511e8d936d96f19c087d12", size = 14483064, upload-time = "2026-04-13T02:45:26.901Z" }, + { url = "https://files.pythonhosted.org/packages/d5/82/74e62e7097fa67da328ac8ece8de09133448c04d20ddeaeba251a3000f01/mypy-1.20.1-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:752507dd481e958b2c08fc966d3806c962af5a9433b5bf8f3bdd7175c20e34fe", size = 13335694, upload-time = "2026-04-13T02:46:12.514Z" }, + { url = "https://files.pythonhosted.org/packages/74/c4/97e9a0abe4f3cdbbf4d079cb87a03b786efeccf5bf2b89fe4f96939ab2e6/mypy-1.20.1-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c614655b5a065e56274c6cbbe405f7cf7e96c0654db7ba39bc680238837f7b08", size = 13726365, upload-time = "2026-04-13T02:45:17.422Z" }, + { url = "https://files.pythonhosted.org/packages/d7/aa/a19d884a8d28fcd3c065776323029f204dbc774e70ec9c85eba228b680de/mypy-1.20.1-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:2c3f6221a76f34d5100c6d35b3ef6b947054123c3f8d6938a4ba00b1308aa572", size = 14693472, upload-time = "2026-04-13T02:46:41.253Z" }, + { url = "https://files.pythonhosted.org/packages/84/44/cc9324bd21cf786592b44bf3b5d224b3923c1230ec9898d508d00241d465/mypy-1.20.1-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:4bdfc06303ac06500af71ea0cdbe995c502b3c9ba32f3f8313523c137a25d1b6", size = 14919266, upload-time = "2026-04-13T02:46:28.37Z" }, + { url = "https://files.pythonhosted.org/packages/6e/dc/779abb25a8c63e8f44bf5a336217fa92790fa17e0c40e0c725d10cb01bbd/mypy-1.20.1-cp314-cp314-win_amd64.whl", hash = "sha256:0131edd7eba289973d1ba1003d1a37c426b85cdef76650cd02da6420898a5eb3", size = 11049713, upload-time = "2026-04-13T02:45:57.673Z" }, + { url = "https://files.pythonhosted.org/packages/28/08/4172be2ad7de9119b5a92ca36abbf641afdc5cb1ef4ae0c3a8182f29674f/mypy-1.20.1-cp314-cp314-win_arm64.whl", hash = "sha256:33f02904feb2c07e1fdf7909026206396c9deeb9e6f34d466b4cfedb0aadbbe4", size = 9999819, upload-time = "2026-04-13T02:46:35.039Z" }, + { url = "https://files.pythonhosted.org/packages/2d/af/af9e46b0c8eabbce9fc04a477564170f47a1c22b308822282a59b7ff315f/mypy-1.20.1-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:168472149dd8cc505c98cefd21ad77e4257ed6022cd5ed2fe2999bed56977a5a", size = 15547508, upload-time = "2026-04-13T02:46:25.588Z" }, + { url = "https://files.pythonhosted.org/packages/a7/cd/39c9e4ad6ba33e069e5837d772a9e6c304b4a5452a14a975d52b36444650/mypy-1.20.1-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:eb674600309a8f22790cca883a97c90299f948183ebb210fbef6bcee07cb1986", size = 14399557, upload-time = "2026-04-13T02:46:10.021Z" }, + { url = "https://files.pythonhosted.org/packages/83/c1/3fd71bdc118ffc502bf57559c909927bb7e011f327f7bb8e0488e98a5870/mypy-1.20.1-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:ef2b2e4cc464ba9795459f2586923abd58a0055487cbe558cb538ea6e6bc142a", size = 15045789, upload-time = "2026-04-13T02:45:10.81Z" }, + { url = "https://files.pythonhosted.org/packages/8e/73/6f07ff8b57a7d7b3e6e5bf34685d17632382395c8bb53364ec331661f83e/mypy-1.20.1-cp314-cp314t-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:dee461d396dd46b3f0ed5a098dbc9b8860c81c46ad44fa071afcfbc149f167c9", size = 15850795, upload-time = "2026-04-13T02:45:03.349Z" }, + { url = "https://files.pythonhosted.org/packages/ec/e2/f7dffec1c7767078f9e9adf0c786d1fe0ff30964a77eb213c09b8b58cb76/mypy-1.20.1-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:e364926308b3e66f1361f81a566fc1b2f8cd47fc8525e8136d4058a65a4b4f02", size = 16088539, upload-time = "2026-04-13T02:46:17.841Z" }, + { url = "https://files.pythonhosted.org/packages/1a/76/e0dee71035316e75a69d73aec2f03c39c21c967b97e277fd0ef8fd6aec66/mypy-1.20.1-cp314-cp314t-win_amd64.whl", hash = "sha256:a0c17fbd746d38c70cbc42647cfd884f845a9708a4b160a8b4f7e70d41f4d7fa", size = 12575567, upload-time = "2026-04-13T02:45:34.795Z" }, + { url = "https://files.pythonhosted.org/packages/22/a8/7ed43c9d9c3d1468f86605e323a5d97e411a448790a00f07e779f3211a46/mypy-1.20.1-cp314-cp314t-win_arm64.whl", hash = "sha256:db2cb89654626a912efda69c0d5c1d22d948265e2069010d3dde3abf751c7d08", size = 10378823, upload-time = "2026-04-13T02:45:13.35Z" }, + { url = "https://files.pythonhosted.org/packages/d8/28/926bd972388e65a39ee98e188ccf67e81beb3aacfd5d6b310051772d974b/mypy-1.20.1-py3-none-any.whl", hash = "sha256:1aae28507f253fe82d883790d1c0a0d35798a810117c88184097fe8881052f06", size = 2636553, upload-time = "2026-04-13T02:46:30.45Z" }, +] + [[package]] name = "mypy-extensions" version = "1.1.0" @@ -6749,6 +6908,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/7d/eb/b6260b31b1a96386c0a880edebe26f89669098acea8e0318bff6adb378fd/pathable-0.4.4-py3-none-any.whl", hash = "sha256:5ae9e94793b6ef5a4cbe0a7ce9dbbefc1eec38df253763fd0aeeacf2762dbbc2", size = 9592, upload-time = "2025-01-10T18:43:11.88Z" }, ] +[[package]] +name = "pathspec" +version = "1.1.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/5a/82/42f767fc1c1143d6fd36efb827202a2d997a375e160a71eb2888a925aac1/pathspec-1.1.1.tar.gz", hash = "sha256:17db5ecd524104a120e173814c90367a96a98d07c45b2e10c2f3919fff91bf5a", size = 135180, upload-time = "2026-04-27T01:46:08.907Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f1/d9/7fb5aa316bc299258e68c73ba3bddbc499654a07f151cba08f6153988714/pathspec-1.1.1-py3-none-any.whl", hash = "sha256:a00ce642f577bf7f473932318056212bc4f8bfdf53128c78bbd5af0b9b20b189", size = 57328, upload-time = "2026-04-27T01:46:07.06Z" }, +] + [[package]] name = "pfzy" version = "0.3.4"