From 9d507383a9eb1df67a6f4c0c70e87adaf2108d00 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 31 Aug 2026 16:48:29 -0700 Subject: [PATCH] fix(python-bridge): harden route runtime --- litellm-rust/Cargo.lock | 1 + .../core/src/chat_completions/handler.rs | 31 +- .../crates/core/src/chat_completions/tests.rs | 2 +- litellm-rust/crates/core/src/error.rs | 8 + litellm-rust/crates/core/src/http_utils.rs | 8 + .../crates/core/src/messages/handler.rs | 13 +- .../crates/core/src/messages/tests.rs | 91 ++++- litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../crates/python-bridge/src/errors.rs | 59 ++- litellm-rust/crates/python-bridge/src/lib.rs | 37 +- .../crates/python-bridge/src/marshal.rs | 31 +- .../src/routes/audio_transcription.rs | 60 ++-- .../src/routes/chat_completions.rs | 64 ++-- .../python-bridge/src/routes/messages.rs | 57 ++- .../crates/python-bridge/src/routes/mod.rs | 337 ++++++++++++++++-- .../crates/python-bridge/src/routes/ocr.rs | 60 ++-- .../python-bridge/src/routes/runtime.rs | 229 +++++++++++- litellm-rust/crates/python-interop/src/lib.rs | 2 +- .../crates/python-interop/src/marshal.rs | 72 ++++ litellm/llms/custom_httpx/llm_http_handler.py | 24 +- .../test_rust_bridge_messages.py | 52 ++- 21 files changed, 985 insertions(+), 254 deletions(-) diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 8a618188ea8..6e461f394f1 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1442,6 +1442,7 @@ name = "litellm-python-bridge" version = "0.1.0" dependencies = [ "criterion", + "futures-util", "litellm-ai-gateway", "litellm-core", "litellm-python-interop", diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index afc4529fd26..713938fad51 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -1,7 +1,7 @@ use serde_json::Value; -use crate::error::{CoreError, CoreResult}; -use crate::http_utils::truncate_error_body; +use crate::error::{CoreError, CoreResult, as_response_error}; +use crate::http_utils::{classify_send_error, truncate_error_body}; use super::client::http_client; use super::transformation::ChatCompletionsAuth; @@ -27,16 +27,7 @@ pub(super) async fn execute_chat_completions_provider_call( request_builder = request_builder.timeout(duration); } - let response = request_builder.send().await.map_err(|err| { - // Failing to establish the connection means the request never went out, - // so the host can still serve it. Everything else here, a timeout - // above all, may have reached the provider and been answered. - if err.is_connect() || err.is_builder() { - CoreError::Connect(err.to_string()) - } else { - CoreError::Network(err.to_string()) - } - })?; + let response = request_builder.send().await.map_err(classify_send_error)?; let status = response.status(); let text = response @@ -60,22 +51,6 @@ pub(super) async fn execute_chat_completions_provider_call( .map_err(as_response_error) } -/// Re-tag an error raised while normalizing a response the provider already -/// returned. -/// -/// A config reports the same variants on either side of the call: a missing -/// field or an unsupported block can mean "this request cannot be translated" -/// during prepare and "this response cannot be normalized" here. Only the -/// second kind has already been billed, and a host that keeps a reference -/// implementation must not retry those, so collapse them to one variant that -/// can only mean the provider was already called. -pub(super) fn as_response_error(err: CoreError) -> CoreError { - match err { - already @ (CoreError::InvalidResponse(_) | CoreError::Http { .. }) => already, - other => CoreError::InvalidResponse(other.to_string()), - } -} - #[cfg(feature = "bedrock-auth")] pub(super) async fn signed_headers( request: &ProviderChatCompletionsRequest, diff --git a/litellm-rust/crates/core/src/chat_completions/tests.rs b/litellm-rust/crates/core/src/chat_completions/tests.rs index e2383723cb0..0f89cc9df96 100644 --- a/litellm-rust/crates/core/src/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/chat_completions/tests.rs @@ -794,7 +794,7 @@ mod round_trip { #[test] fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() { - use crate::chat_completions::handler::as_response_error; + use crate::error::as_response_error; for original in [ CoreError::MissingField("usage"), diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index 739532f8cb5..dc23d440233 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -38,6 +38,14 @@ pub enum CoreError { Unsupported(&'static str), } +/// Re-tag an error raised after the provider has already returned a response. +pub(crate) fn as_response_error(err: CoreError) -> CoreError { + match err { + already @ (CoreError::InvalidResponse(_) | CoreError::Http { .. }) => already, + other => CoreError::InvalidResponse(other.to_string()), + } +} + pub fn json_type_name(value: &serde_json::Value) -> &'static str { match value { serde_json::Value::Null => "null", diff --git a/litellm-rust/crates/core/src/http_utils.rs b/litellm-rust/crates/core/src/http_utils.rs index c541f50275b..d79f178a124 100644 --- a/litellm-rust/crates/core/src/http_utils.rs +++ b/litellm-rust/crates/core/src/http_utils.rs @@ -5,6 +5,14 @@ use serde_json::{Map, Value}; use crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS; use crate::error::{CoreError, CoreResult, json_type_name}; +pub(crate) fn classify_send_error(error: reqwest::Error) -> CoreError { + if error.is_connect() || error.is_builder() { + CoreError::Connect(error.to_string()) + } else { + CoreError::Network(error.to_string()) + } +} + /// Bound an upstream error body before it crosses a host boundary, so provider /// bodies stay data-minimized. pub fn truncate_error_body(body: &str) -> String { diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index 1c895f66eba..92c67058dca 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,5 +1,6 @@ use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; -use crate::error::{CoreError, CoreResult}; +use crate::error::{CoreError, CoreResult, as_response_error}; +use crate::http_utils::classify_send_error; use super::client::http_client; use super::common_utils::truncate_error_body; @@ -16,10 +17,7 @@ pub(super) async fn execute_messages_provider_call( request_builder = request_builder.timeout(duration); } - let response = request_builder - .send() - .await - .map_err(|err| CoreError::Network(err.to_string()))?; + let response = request_builder.send().await.map_err(classify_send_error)?; let status = response.status(); let text = response @@ -37,7 +35,10 @@ pub(super) async fn execute_messages_provider_call( let response = serde_json::from_str(&text).map_err(|err| { CoreError::InvalidResponse(format!("invalid messages response JSON: {err}")) })?; - request.config.transform_response(&request.model, response) + request + .config + .transform_response(&request.model, response) + .map_err(as_response_error) } pub(super) async fn execute_messages_provider_stream( diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index 9fc1763683b..6a59b6c7061 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -4,13 +4,46 @@ use serde_json::{Map, Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; -use crate::error::CoreError; +use crate::error::{CoreError, CoreResult}; use super::common_utils::{ has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body, }; +use super::handler::execute_messages_provider_call; use super::messages; -use super::types::MessagesRequest; +use super::transformation::AnthropicMessagesProviderConfig; +use super::types::{AnthropicMessagesResponse, MessagesRequest, ProviderMessagesRequest}; + +struct RejectingResponseConfig; + +impl AnthropicMessagesProviderConfig for RejectingResponseConfig { + fn complete_url( + &self, + _api_base: Option<&str>, + _model: &str, + _env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + unreachable!() + } + + fn resolve_api_key( + &self, + _api_key: Option<&str>, + _env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + unreachable!() + } + + fn transform_response( + &self, + _model: &str, + _response: AnthropicMessagesResponse, + ) -> CoreResult { + Err(CoreError::MissingField("normalized_content")) + } +} + +static REJECTING_RESPONSE_CONFIG: RejectingResponseConfig = RejectingResponseConfig; async fn read_http_request(socket: &mut TcpStream) -> String { let mut request = Vec::new(); @@ -178,6 +211,39 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through() ); } +#[tokio::test] +async fn post_response_transform_errors_are_non_retryable() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let addr = listener.local_addr().expect("addr"); + + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let _ = read_http_request(&mut socket).await; + let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-test"}"#; + socket + .write_all(write_response(response_body).as_bytes()) + .await + .expect("writes response"); + }); + + let error = execute_messages_provider_call(ProviderMessagesRequest { + provider: "anthropic".to_string(), + model: "claude-test".to_string(), + config: &REJECTING_RESPONSE_CONFIG, + url: format!("http://{addr}/v1/messages"), + body: json!({}), + upstream_headers: Vec::new(), + timeout: Some(Duration::from_secs(5)), + }) + .await + .expect_err("response transform should fail"); + + server.await.expect("server task completes"); + assert!( + matches!(error, CoreError::InvalidResponse(message) if message.contains("normalized_content")) + ); +} + #[tokio::test] async fn messages_round_trip_builds_native_anthropic_request() { let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); @@ -439,3 +505,24 @@ async fn messages_rejects_unsupported_provider() { assert!(matches!(err, CoreError::InvalidProvider(provider) if provider == "openai")); } + +#[tokio::test] +async fn messages_classifies_a_refused_connection_as_safe_to_fallback() { + let port = { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + listener.local_addr().expect("has an address").port() + }; + let error = messages(MessagesRequest { + model: "claude-test", + body: json!({"model": "claude-test", "max_tokens": 8, "messages": []}), + api_key: Some("sk"), + api_base: Some(&format!("http://127.0.0.1:{port}")), + custom_llm_provider: Some("anthropic"), + extra_headers: None, + timeout: Some(Duration::from_secs(1)), + }) + .await + .expect_err("nothing is listening"); + + assert!(matches!(error, CoreError::Connect(_))); +} diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index dbeca5e56e3..c3cfb968d7e 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -16,6 +16,7 @@ extension-module = ["pyo3/extension-module"] panic-test = [] [dependencies] +futures-util.workspace = true litellm-core = { workspace = true, features = ["bedrock-auth"] } litellm-ai-gateway = { workspace = true, default-features = false } litellm-python-interop.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index ec68cf7bfa3..b1601930d2e 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -33,7 +33,7 @@ pub(crate) fn core_error_to_pyerr(err: CoreError) -> PyErr { /// Everything raised before the request goes out is safe for the host to retry /// on its own path; anything after it is not, because the provider has already /// done the work and billed for it. -pub(crate) fn chat_completions_error_to_pyerr(err: CoreError) -> PyErr { +pub(crate) fn fallback_route_error_to_pyerr(err: CoreError) -> PyErr { match err { CoreError::Unsupported(_) | CoreError::Auth(_) @@ -59,3 +59,60 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { module.add("RustBridgeDeclined", py.get_type::())?; module.add("RustUpstreamError", py.get_type::()) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn fallback_routes_distinguish_declines_from_upstream_failures() { + Python::initialize(); + Python::attach(|py| { + let declines = [ + CoreError::Unsupported("unsupported"), + CoreError::Auth("missing key".to_string()), + CoreError::InvalidProvider("unsupported".to_string()), + CoreError::InvalidRequest("invalid".to_string()), + CoreError::InvalidType { + expected: "string", + actual: "number", + }, + CoreError::MissingField("model"), + CoreError::Routing("no route".to_string()), + CoreError::Connect("connection refused".to_string()), + ]; + for error in declines { + let mapped = fallback_route_error_to_pyerr(error); + assert!(mapped.is_instance_of::(py)); + } + + let upstream_failures = [ + ( + CoreError::Http { + status: 429, + body: "rate limited".to_string(), + }, + (429, "429: rate limited"), + ), + ( + CoreError::Network("request timed out".to_string()), + (0, "request timed out"), + ), + ( + CoreError::InvalidResponse("bad JSON".to_string()), + (0, "bad JSON"), + ), + ]; + for (error, expected) in upstream_failures { + let mapped = fallback_route_error_to_pyerr(error); + assert!(mapped.is_instance_of::(py)); + let args: (u16, String) = mapped + .value(py) + .getattr("args") + .and_then(|args| args.extract()) + .expect("upstream error should carry status and message"); + assert_eq!(args, (expected.0, expected.1.to_string())); + } + }); + } +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index f628c987220..9f2c471ef25 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -3,7 +3,10 @@ mod errors; mod marshal; mod routes; +use std::panic::{AssertUnwindSafe, catch_unwind}; + use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; +use litellm_python_interop::panic_to_pyerr; use pyo3::prelude::*; use pyo3::types::PyAny; @@ -15,6 +18,20 @@ struct ResponsesWebSocketConnection { inner: RustResponsesWebSocketConnection, } +struct NewResponsesWebSocketConnection(ResponsesWebSocketConnection); + +impl<'py> IntoPyObject<'py> for NewResponsesWebSocketConnection { + type Target = ResponsesWebSocketConnection; + type Output = Bound<'py, ResponsesWebSocketConnection>; + type Error = PyErr; + + fn into_pyobject(self, py: Python<'py>) -> PyResult { + catch_unwind(AssertUnwindSafe(|| Py::new(py, self.0))) + .map_err(panic_to_pyerr)? + .map(|value| value.into_bound(py)) + } +} + #[pymethods] impl ResponsesWebSocketConnection { #[classmethod] @@ -32,7 +49,9 @@ impl ResponsesWebSocketConnection { let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout) .await .map_err(core_error_to_pyerr)?; - Python::attach(|py| Py::new(py, ResponsesWebSocketConnection { inner })) + Ok(NewResponsesWebSocketConnection( + ResponsesWebSocketConnection { inner }, + )) }) } @@ -93,13 +112,15 @@ mod tests { "gil_stats", ]; - for name in expected { - assert!( - module - .hasattr(name) - .expect("attribute lookup should succeed") - ); - } + let public_names: Vec = module + .dict() + .keys() + .extract::>() + .expect("module names should be strings") + .into_iter() + .filter(|name| !name.starts_with("__")) + .collect(); + assert_eq!(public_names, expected); }); } } diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 0ff88d0ada1..a89dffc539a 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -15,23 +15,24 @@ pub(crate) struct RouteOptions { pub(crate) timeout: Option, } +pub(crate) struct RouteOptionsInputs { + pub(crate) model: String, + pub(crate) api_key: Option, + pub(crate) api_base: Option, + pub(crate) custom_llm_provider: Option, + pub(crate) extra_headers: Option>, + pub(crate) timeout_seconds: Option, +} + impl RouteOptions { - pub(crate) fn from_python( - py: Python<'_>, - model: String, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - extra_headers: Option>, - timeout_seconds: Option, - ) -> PyResult { + pub(crate) fn from_python(py: Python<'_>, inputs: RouteOptionsInputs) -> PyResult { Ok(Self { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers: optional_object(py, "extra_headers", extra_headers)?, - timeout: optional_timeout(timeout_seconds), + model: inputs.model, + api_key: inputs.api_key, + api_base: inputs.api_base, + custom_llm_provider: inputs.custom_llm_provider, + extra_headers: optional_object(py, "extra_headers", inputs.extra_headers)?, + timeout: optional_timeout(inputs.timeout_seconds), }) } } diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index e39f4beadb0..b1d6fb5958a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -1,41 +1,35 @@ +use std::future::Future; + use litellm_ai_gateway::io::audio_transcription::{ AudioTranscriptionRequest, audio_transcription as run_audio_transcription, }; use litellm_core::error::CoreResult; use litellm_python_interop::from_py; use pyo3::prelude::*; -use serde_json::{Map, Value}; +use serde_json::Value; use crate::errors::core_error_to_pyerr; -use crate::marshal::{RouteOptions, object_or_empty}; -use crate::routes::BridgeRoute; +use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty}; -struct AudioTranscriptionCall { - options: RouteOptions, - audio: Value, - optional_params: Map, -} +fn prepare_transcription( + py: Python<'_>, + inputs: AudioTranscriptionInputs, +) -> PyResult> + Send + 'static> { + let audio = from_py(inputs.audio.bind(py))?; + let options = RouteOptions::from_python( + py, + RouteOptionsInputs { + model: inputs.model, + api_key: inputs.api_key, + api_base: inputs.api_base, + custom_llm_provider: inputs.custom_llm_provider, + extra_headers: inputs.extra_headers, + timeout_seconds: inputs.timeout_seconds, + }, + )?; + let optional_params = object_or_empty(py, "optional_params", inputs.optional_params)?; -impl BridgeRoute for AudioTranscriptionCall { - type Output = Value; - - fn from_python(py: Python<'_>, inputs: AudioTranscriptionInputs) -> PyResult { - Ok(Self { - options: RouteOptions::from_python( - py, - inputs.model, - inputs.api_key, - inputs.api_base, - inputs.custom_llm_provider, - inputs.extra_headers, - inputs.timeout_seconds, - )?, - audio: from_py(inputs.audio.bind(py))?, - optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?, - }) - } - - async fn run(self) -> CoreResult { + Ok(async move { let RouteOptions { model, api_key, @@ -43,15 +37,15 @@ impl BridgeRoute for AudioTranscriptionCall { custom_llm_provider, extra_headers, timeout, - } = self.options; + } = options; run_audio_transcription(AudioTranscriptionRequest { model: &model, - audio: self.audio, + audio, api_key: api_key.as_deref(), api_base: api_base.as_deref(), custom_llm_provider: custom_llm_provider.as_deref(), extra_headers, - optional_params: self.optional_params, + optional_params, timeout, callbacks: Vec::new(), guardrails: Vec::new(), @@ -59,7 +53,7 @@ impl BridgeRoute for AudioTranscriptionCall { litellm_call_id: None, }) .await - } + }) } bridge_route! { @@ -78,6 +72,6 @@ bridge_route! { optional_params: Option>, timeout_seconds: Option, }, - call = AudioTranscriptionCall, + prepare = prepare_transcription, errors = core_error_to_pyerr, } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index c05e1fcd4cd..8ee4b3e8ff1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -1,3 +1,5 @@ +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, @@ -5,38 +7,30 @@ use litellm_core::chat_completions::{ use litellm_core::error::CoreResult; use litellm_python_interop::from_py; use pyo3::prelude::*; -use serde_json::{Map, Value}; +use serde_json::Value; -use crate::errors::chat_completions_error_to_pyerr; -use crate::marshal::{RouteOptions, object_or_empty, required_value}; -use crate::routes::BridgeRoute; +use crate::errors::fallback_route_error_to_pyerr; +use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_value}; -struct ChatCompletionsCall { - options: RouteOptions, - messages: Value, - optional_params: Map, -} +fn prepare_chat_completions( + py: Python<'_>, + inputs: ChatCompletionsInputs, +) -> PyResult> + Send + 'static> { + let messages = required_value(py, "messages", inputs.messages, Value::is_array, "list")?; + let optional_params = object_or_empty(py, "optional_params", inputs.optional_params)?; + let options = RouteOptions::from_python( + py, + RouteOptionsInputs { + model: inputs.model, + api_key: inputs.api_key, + api_base: inputs.api_base, + custom_llm_provider: inputs.custom_llm_provider, + extra_headers: inputs.extra_headers, + timeout_seconds: inputs.timeout_seconds, + }, + )?; -impl BridgeRoute for ChatCompletionsCall { - type Output = ChatCompletionsResponse; - - fn from_python(py: Python<'_>, inputs: ChatCompletionsInputs) -> PyResult { - Ok(Self { - options: RouteOptions::from_python( - py, - inputs.model, - inputs.api_key, - inputs.api_base, - inputs.custom_llm_provider, - inputs.extra_headers, - inputs.timeout_seconds, - )?, - messages: required_value(py, "messages", inputs.messages, Value::is_array, "list")?, - optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?, - }) - } - - async fn run(self) -> CoreResult { + Ok(async move { let RouteOptions { model, api_key, @@ -44,11 +38,11 @@ impl BridgeRoute for ChatCompletionsCall { custom_llm_provider, extra_headers, timeout, - } = self.options; + } = options; run_chat_completions(ChatCompletionsRequest { model: &model, - messages: self.messages, - optional_params: self.optional_params, + messages, + optional_params, api_key: api_key.as_deref(), api_base: api_base.as_deref(), custom_llm_provider: custom_llm_provider.as_deref(), @@ -56,7 +50,7 @@ impl BridgeRoute for ChatCompletionsCall { timeout, }) .await - } + }) } #[pyfunction] @@ -95,7 +89,7 @@ bridge_route! { extra_headers: Option>, timeout_seconds: Option, }, - call = ChatCompletionsCall, - errors = chat_completions_error_to_pyerr, + prepare = prepare_chat_completions, + errors = fallback_route_error_to_pyerr, extra = [chat_completions_decline], } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages.rs b/litellm-rust/crates/python-bridge/src/routes/messages.rs index 4f8e52967e0..accba5cf0e9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages.rs @@ -1,37 +1,32 @@ +use std::future::Future; + use litellm_core::error::CoreResult; use litellm_core::messages::messages as run_messages; use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest}; use pyo3::prelude::*; use serde_json::Value; -use crate::errors::core_error_to_pyerr; -use crate::marshal::{RouteOptions, required_value}; -use crate::routes::BridgeRoute; +use crate::errors::fallback_route_error_to_pyerr; +use crate::marshal::{RouteOptions, RouteOptionsInputs, required_value}; -struct MessagesCall { - options: RouteOptions, - body: Value, -} +fn prepare_messages( + py: Python<'_>, + inputs: MessagesInputs, +) -> PyResult> + Send + 'static> { + let body = required_value(py, "body", inputs.body, Value::is_object, "dict")?; + let options = RouteOptions::from_python( + py, + RouteOptionsInputs { + model: inputs.model, + api_key: inputs.api_key, + api_base: inputs.api_base, + custom_llm_provider: inputs.custom_llm_provider, + extra_headers: inputs.extra_headers, + timeout_seconds: inputs.timeout_seconds, + }, + )?; -impl BridgeRoute for MessagesCall { - type Output = AnthropicMessagesResponse; - - fn from_python(py: Python<'_>, inputs: MessagesInputs) -> PyResult { - Ok(Self { - options: RouteOptions::from_python( - py, - inputs.model, - inputs.api_key, - inputs.api_base, - inputs.custom_llm_provider, - inputs.extra_headers, - inputs.timeout_seconds, - )?, - body: required_value(py, "body", inputs.body, Value::is_object, "dict")?, - }) - } - - async fn run(self) -> CoreResult { + Ok(async move { let RouteOptions { model, api_key, @@ -39,10 +34,10 @@ impl BridgeRoute for MessagesCall { custom_llm_provider, extra_headers, timeout, - } = self.options; + } = options; run_messages(MessagesRequest { model: &model, - body: self.body, + body, api_key: api_key.as_deref(), api_base: api_base.as_deref(), custom_llm_provider: custom_llm_provider.as_deref(), @@ -50,7 +45,7 @@ impl BridgeRoute for MessagesCall { timeout, }) .await - } + }) } bridge_route! { @@ -68,6 +63,6 @@ bridge_route! { extra_headers: Option>, timeout_seconds: Option, }, - call = MessagesCall, - errors = core_error_to_pyerr, + prepare = prepare_messages, + errors = fallback_route_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 d17c98c9dc7..5fe94b9123f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -1,29 +1,19 @@ -use std::future::Future; - -use litellm_core::error::CoreResult; +use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; -use serde::Serialize; +use pyo3::types::PyCFunction; mod runtime; use runtime::{run_async, run_sync}; -trait BridgeRoute: Sized { - type Output: Serialize + Send + 'static; - - fn from_python(py: Python<'_>, inputs: I) -> PyResult; - - fn run(self) -> impl Future> + Send + 'static; -} - macro_rules! bridge_route { ( sync = $sync_name:ident, asynchronous = $async_name:ident, inputs = $inputs:ident, - required = { $($required_name:ident: $required_type:ty),* $(,)? }, + required = { $($required_name:ident: $required_type:ty),+ $(,)? }, optional = { $($optional_name:ident: $optional_type:ty),* $(,)? }, - call = $call:ty, + prepare = $prepare:path, errors = $map_error:path $(, extra = [$($extra:ident),* $(,)?])? $(,)? @@ -41,15 +31,11 @@ macro_rules! bridge_route { $($required_name: $required_type,)* $($optional_name: $optional_type),* ) -> pyo3::PyResult> { - let call = <$call as crate::routes::BridgeRoute<$inputs>>::from_python(py, $inputs { + let future = $prepare(py, $inputs { $($required_name,)* $($optional_name),* })?; - crate::routes::run_sync( - py, - <$call as crate::routes::BridgeRoute<$inputs>>::run(call), - $map_error, - ) + crate::routes::run_sync(py, future, $map_error) } #[pyfunction] @@ -60,47 +46,119 @@ macro_rules! bridge_route { $($required_name: $required_type,)* $($optional_name: $optional_type),* ) -> pyo3::PyResult> { - let call = <$call as crate::routes::BridgeRoute<$inputs>>::from_python(py, $inputs { + let future = $prepare(py, $inputs { $($required_name,)* $($optional_name),* })?; - crate::routes::run_async( - py, - <$call as crate::routes::BridgeRoute<$inputs>>::run(call), - $map_error, - ) + crate::routes::run_async(py, future, $map_error) } pub(super) fn register( module: &pyo3::Bound<'_, pyo3::types::PyModule>, ) -> pyo3::PyResult<()> { - module.add_function(pyo3::wrap_pyfunction!($sync_name, module)?)?; - module.add_function(pyo3::wrap_pyfunction!($async_name, module)?)?; - $($(module.add_function(pyo3::wrap_pyfunction!($extra, module)?)?;)*)? + $($(crate::routes::add_function(module, pyo3::wrap_pyfunction!($extra, module)?)?;)*)? + crate::routes::add_function(module, pyo3::wrap_pyfunction!($sync_name, module)?)?; + crate::routes::add_function(module, pyo3::wrap_pyfunction!($async_name, module)?)?; Ok(()) } }; } -macro_rules! routes { - ($($route:ident),* $(,)?) => { - $(mod $route;)* - - pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - $($route::register(module)?;)* - Ok(()) - } - }; +fn add_function(module: &Bound<'_, PyModule>, function: Bound<'_, PyCFunction>) -> PyResult<()> { + let name: String = function.getattr("__name__")?.extract()?; + if module.hasattr(&name)? { + return Err(PyRuntimeError::new_err(format!( + "duplicate native route: {name}" + ))); + } + module.add_function(function) } -routes!(ocr, audio_transcription, messages, chat_completions); +mod audio_transcription; +mod chat_completions; +mod messages; +mod ocr; + +pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + ocr::register(module)?; + audio_transcription::register(module)?; + messages::register(module)?; + chat_completions::register(module) +} #[cfg(test)] mod tests { + use std::ffi::CString; + use std::sync::atomic::{AtomicBool, Ordering}; + + use litellm_core::error::{CoreError, CoreResult}; + use pyo3::exceptions::PyLookupError; use pyo3::types::{PyDict, PyList}; use super::*; + mod synthetic { + use std::future::{Future, pending}; + + use super::*; + + static FUTURE_DROPPED: AtomicBool = AtomicBool::new(false); + + struct DropGuard; + + impl Drop for DropGuard { + fn drop(&mut self) { + FUTURE_DROPPED.store(true, Ordering::SeqCst); + } + } + + #[pyfunction] + fn future_dropped() -> bool { + FUTURE_DROPPED.load(Ordering::SeqCst) + } + + bridge_route! { + sync = echo, + asynchronous = aecho, + inputs = EchoInputs, + required = { value: String }, + optional = {}, + prepare = prepare_echo, + errors = map_error, + extra = [future_dropped], + } + + fn prepare_echo( + _py: Python<'_>, + inputs: EchoInputs, + ) -> PyResult> + Send + 'static> { + FUTURE_DROPPED.store(false, Ordering::SeqCst); + let drop_guard = (inputs.value == "pending").then_some(DropGuard); + Ok(async move { + let _drop_guard = drop_guard; + tokio::task::yield_now().await; + match inputs.value.as_str() { + "error" => Err(CoreError::InvalidRequest("synthetic error".to_string())), + "map_panic" => Err(CoreError::InvalidRequest("panic in mapper".to_string())), + "panic" => panic!("synthetic panic"), + "pending" => { + pending::<()>().await; + unreachable!() + } + _ => Ok(inputs.value), + } + }) + } + + fn map_error(error: CoreError) -> PyErr { + if matches!(&error, CoreError::InvalidRequest(message) if message == "panic in mapper") + { + panic!("synthetic mapper panic") + } + PyLookupError::new_err(error.to_string()) + } + } + #[test] fn sync_and_async_route_signatures_match_the_python_contract() { Python::initialize(); @@ -215,4 +273,205 @@ mod tests { } }); } + + #[test] + fn route_input_validation_preserves_left_to_right_order() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "routes").expect("module should be created"); + register(&module).expect("routes should register"); + let invalid = PyList::empty(py); + + let chat_kwargs = PyDict::new(py); + chat_kwargs + .set_item("optional_params", &invalid) + .expect("kwargs should accept optional_params"); + chat_kwargs + .set_item("extra_headers", &invalid) + .expect("kwargs should accept extra_headers"); + let invalid_messages = PyDict::new(py); + let error = module + .getattr("chat_completions") + .and_then(|function| { + function.call(("model", &invalid_messages), Some(&chat_kwargs)) + }) + .expect_err("messages should be validated first"); + assert_eq!(error.to_string(), "ValueError: messages must be a list"); + + let valid_messages = PyList::empty(py); + let error = module + .getattr("chat_completions") + .and_then(|function| function.call(("model", &valid_messages), Some(&chat_kwargs))) + .expect_err("optional_params should be validated before headers"); + assert_eq!( + error.to_string(), + "ValueError: optional_params must be a dict" + ); + + let headers_kwargs = PyDict::new(py); + headers_kwargs + .set_item("extra_headers", &invalid) + .expect("kwargs should accept extra_headers"); + let invalid_body = PyList::empty(py); + let error = module + .getattr("messages") + .and_then(|function| function.call(("model", &invalid_body), Some(&headers_kwargs))) + .expect_err("body should be validated before headers"); + assert_eq!(error.to_string(), "ValueError: body must be a dict"); + + let invalid_payload = + PyModule::new(py, "invalid_payload").expect("invalid payload should be created"); + for name in ["ocr", "transcription"] { + let error = module + .getattr(name) + .and_then(|function| { + function.call(("model", &invalid_payload), Some(&headers_kwargs)) + }) + .expect_err("payload should be validated before headers"); + assert!(!error.to_string().contains("extra_headers")); + } + }); + } + + #[test] + fn generated_routes_execute_sync_and_async_contracts() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "synthetic").expect("module should be created"); + synthetic::register(&module).expect("routes should register"); + + let sync_value: String = module + .getattr("echo") + .and_then(|function| function.call1(("sync",))) + .and_then(|value| value.extract()) + .expect("sync route should return its value"); + assert_eq!(sync_value, "sync"); + + let sync_error = module + .getattr("echo") + .and_then(|function| function.call1(("error",))) + .expect_err("sync route should map its error"); + assert!(sync_error.is_instance_of::(py)); + assert_eq!( + sync_error.to_string(), + "LookupError: invalid request: synthetic error" + ); + + let locals = PyDict::new(py); + locals + .set_item("routes", &module) + .expect("module should enter Python locals"); + let code = CString::new( + r#" +import asyncio + +async def exercise(): + assert await routes.aecho("async") == "async" + + try: + await routes.aecho("error") + except LookupError as error: + assert str(error) == "invalid request: synthetic error" + else: + raise AssertionError("mapped error was not raised") + + try: + await routes.aecho("panic") + except BaseException as error: + assert type(error).__name__ == "PanicException" + assert str(error) == "synthetic panic" + else: + raise AssertionError("panic was not raised") + + try: + await routes.aecho("map_panic") + except BaseException as error: + assert type(error).__name__ == "PanicException" + assert str(error) == "synthetic mapper panic" + else: + raise AssertionError("mapper panic was not raised") + + task = asyncio.ensure_future(routes.aecho("pending")) + await asyncio.sleep(0) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + else: + raise AssertionError("cancelled route completed") + + for _ in range(100): + if routes.future_dropped(): + break + await asyncio.sleep(0.001) + assert routes.future_dropped() + +asyncio.run(exercise()) +"#, + ) + .expect("Python source should not contain null bytes"); + py.run(&code, Some(&locals), Some(&locals)) + .expect("async route contract should hold"); + }); + } + + #[test] + fn messages_routes_map_declines_before_python_fallback() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "routes").expect("module should be created"); + register(&module).expect("routes should register"); + let body = PyDict::new(py); + let kwargs = PyDict::new(py); + kwargs + .set_item("custom_llm_provider", "openai") + .expect("kwargs should accept provider"); + + let sync_error = module + .getattr("messages") + .and_then(|function| function.call(("model", &body), Some(&kwargs))) + .expect_err("unsupported provider should decline"); + assert!(sync_error.is_instance_of::(py)); + + let locals = PyDict::new(py); + locals + .set_item("routes", &module) + .expect("module should enter Python locals"); + let code = CString::new( + r#" +import asyncio + +async def exercise(): + try: + await routes.amessages("model", {}, custom_llm_provider="openai") + except Exception as error: + assert type(error).__name__ == "RustBridgeDeclined" + else: + raise AssertionError("unsupported provider did not decline") + +asyncio.run(exercise()) +"#, + ) + .expect("Python source should not contain null bytes"); + py.run(&code, Some(&locals), Some(&locals)) + .expect("async route should preserve the decline contract"); + }); + } + + #[test] + fn route_registration_rejects_duplicate_python_names() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "synthetic").expect("module should be created"); + synthetic::register(&module).expect("first registration should succeed"); + let error = synthetic::register(&module) + .expect_err("duplicate registration should be rejected"); + + assert_eq!( + error.to_string(), + "RuntimeError: duplicate native route: future_dropped" + ); + }); + } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index fd21ecce7d9..2afc5078549 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -1,39 +1,33 @@ +use std::future::Future; + use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; use litellm_core::error::CoreResult; use litellm_python_interop::from_py; use pyo3::prelude::*; -use serde_json::{Map, Value}; +use serde_json::Value; use crate::errors::core_error_to_pyerr; -use crate::marshal::{RouteOptions, object_or_empty}; -use crate::routes::BridgeRoute; +use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty}; -struct OcrCall { - options: RouteOptions, - document: Value, - optional_params: Map, -} +fn prepare_ocr( + py: Python<'_>, + inputs: OcrInputs, +) -> PyResult> + Send + 'static> { + let document = from_py(inputs.document.bind(py))?; + let options = RouteOptions::from_python( + py, + RouteOptionsInputs { + model: inputs.model, + api_key: inputs.api_key, + api_base: inputs.api_base, + custom_llm_provider: inputs.custom_llm_provider, + extra_headers: inputs.extra_headers, + timeout_seconds: inputs.timeout_seconds, + }, + )?; + let optional_params = object_or_empty(py, "optional_params", inputs.optional_params)?; -impl BridgeRoute for OcrCall { - type Output = Value; - - fn from_python(py: Python<'_>, inputs: OcrInputs) -> PyResult { - Ok(Self { - options: RouteOptions::from_python( - py, - inputs.model, - inputs.api_key, - inputs.api_base, - inputs.custom_llm_provider, - inputs.extra_headers, - inputs.timeout_seconds, - )?, - document: from_py(inputs.document.bind(py))?, - optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?, - }) - } - - async fn run(self) -> CoreResult { + Ok(async move { let RouteOptions { model, api_key, @@ -41,15 +35,15 @@ impl BridgeRoute for OcrCall { custom_llm_provider, extra_headers, timeout, - } = self.options; + } = options; run_ocr(OcrRequest { model: &model, - document: self.document, + document, api_key: api_key.as_deref(), api_base: api_base.as_deref(), custom_llm_provider: custom_llm_provider.as_deref(), extra_headers, - optional_params: self.optional_params, + optional_params, timeout, callbacks: Vec::new(), guardrails: Vec::new(), @@ -57,7 +51,7 @@ impl BridgeRoute for OcrCall { litellm_call_id: None, }) .await - } + }) } bridge_route! { @@ -76,6 +70,6 @@ bridge_route! { optional_params: Option>, timeout_seconds: Option, }, - call = OcrCall, + prepare = prepare_ocr, errors = core_error_to_pyerr, } diff --git a/litellm-rust/crates/python-bridge/src/routes/runtime.rs b/litellm-rust/crates/python-bridge/src/routes/runtime.rs index da5a7124a58..4fc55d2965d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/runtime.rs +++ b/litellm-rust/crates/python-bridge/src/routes/runtime.rs @@ -1,11 +1,15 @@ use std::future::Future; -use std::sync::mpsc::sync_channel; +use std::panic::AssertUnwindSafe; +use std::time::Duration; +use futures_util::FutureExt; use litellm_core::error::{CoreError, CoreResult}; -use litellm_python_interop::{release_gil, to_py}; +use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil, to_py}; use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; use serde::Serialize; +use tokio::runtime::{Handle, Runtime}; +use tokio::time::{self, MissedTickBehavior}; pub(super) fn run_sync( py: Python<'_>, @@ -16,13 +20,32 @@ where T: Serialize + Send + 'static, F: Future> + Send + 'static, { - let (sender, receiver) = sync_channel(1); - pyo3_async_runtimes::tokio::get_runtime().spawn(async move { - let _ = sender.send(future.await); - }); - let result = release_gil(py, move || receiver.recv()) - .map_err(|_| PyRuntimeError::new_err("native route task terminated"))? - .map_err(map_error)?; + run_sync_on( + py, + pyo3_async_runtimes::tokio::get_runtime(), + future, + map_error, + ) +} + +fn run_sync_on( + py: Python<'_>, + runtime: &Runtime, + future: F, + map_error: fn(CoreError) -> PyErr, +) -> PyResult> +where + T: Serialize + Send + 'static, + F: Future> + Send + 'static, +{ + if Handle::try_current().is_ok() { + return Err(PyRuntimeError::new_err( + "synchronous native routes cannot run from a Tokio context; use the async route", + )); + } + + let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?; + let result = map_core_result(result, map_error)?; to_py(py, &result) } @@ -36,14 +59,61 @@ where F: Future> + Send + 'static, { pyo3_async_runtimes::tokio::future_into_py(py, async move { - let result = future.await.map_err(map_error)?; - Python::attach(|py| to_py(py, &result)) + let result = catch_route_panic(future).await?; + let result = map_core_result(result, map_error)?; + Ok(Pythonized(result)) }) } +fn map_core_result(result: CoreResult, map_error: fn(CoreError) -> PyErr) -> PyResult { + match result { + Ok(value) => Ok(value), + Err(error) => Err( + std::panic::catch_unwind(AssertUnwindSafe(|| map_error(error))) + .map_err(panic_to_pyerr)?, + ), + } +} + +async fn catch_route_panic(future: F) -> PyResult> +where + F: Future>, +{ + AssertUnwindSafe(future) + .catch_unwind() + .await + .map_err(panic_to_pyerr) +} + +async fn wait_for_sync_result(future: F) -> PyResult> +where + F: Future>, +{ + let future = catch_route_panic(future); + tokio::pin!(future); + + let signal_interval = Duration::from_millis(50); + let mut signal_checks = + time::interval_at(time::Instant::now() + signal_interval, signal_interval); + signal_checks.set_missed_tick_behavior(MissedTickBehavior::Delay); + loop { + tokio::select! { + result = &mut future => return result, + _ = signal_checks.tick() => Python::attach(|py| py.check_signals())?, + } + } +} + #[cfg(test)] mod tests { - use std::time::Duration; + use std::ffi::CString; + use std::future::poll_fn; + use std::task::Poll; + + use pyo3::panic::PanicException; + use pyo3::types::{PyDict, PyModule}; + use serde::Serializer; + use tokio::runtime::Builder; use super::*; @@ -51,6 +121,26 @@ mod tests { PyRuntimeError::new_err(error.to_string()) } + fn panicking_error_mapper(_error: CoreError) -> PyErr { + panic!("error mapper panicked") + } + + struct PanickingOutput; + + impl Serialize for PanickingOutput { + fn serialize(&self, _serializer: S) -> Result + where + S: Serializer, + { + panic!("async serializer panicked") + } + } + + #[pyfunction] + fn async_serialization_panic(py: Python<'_>) -> PyResult> { + run_async(py, async { Ok(PanickingOutput) }, runtime_error) + } + fn extract_bool(py: Python<'_>, result: PyResult>) -> bool { result .expect("route should complete") @@ -60,13 +150,13 @@ mod tests { } #[test] - fn sync_runner_polls_future_on_tokio_worker() { + fn sync_runner_polls_future_on_the_caller_thread() { Python::initialize(); Python::attach(|py| { let caller_thread = std::thread::current().id(); let result = run_sync( py, - async move { Ok(std::thread::current().id() != caller_thread) }, + async move { Ok(std::thread::current().id() == caller_thread) }, runtime_error, ); @@ -94,4 +184,115 @@ mod tests { assert!(extract_bool(py, result)); }); } + + #[test] + fn sync_runner_rejects_calls_from_a_tokio_context() { + Python::initialize(); + let runtime = Builder::new_current_thread() + .enable_all() + .build() + .expect("runtime should build"); + + let error = runtime.block_on(async { + Python::attach(|py| { + run_sync::(py, async { Ok(true) }, runtime_error) + .expect_err("sync route should reject a nested Tokio runtime") + }) + }); + + assert_eq!( + error.to_string(), + "RuntimeError: synchronous native routes cannot run from a Tokio context; use the async route" + ); + } + + #[test] + fn sync_runner_can_drive_a_current_thread_runtime() { + Python::initialize(); + let runtime = Builder::new_current_thread() + .enable_all() + .build() + .expect("runtime should build"); + Python::attach(|py| { + let result = run_sync_on( + py, + &runtime, + async { + tokio::task::yield_now().await; + Ok(true) + }, + runtime_error, + ); + assert!(extract_bool(py, result)); + }); + } + + #[test] + fn sync_runner_maps_a_panicked_future() { + Python::initialize(); + Python::attach(|py| { + let error = run_sync::( + py, + poll_fn(|_| -> Poll> { panic!("route future panicked") }), + runtime_error, + ) + .expect_err("panicked route should become a Python exception"); + + assert!(error.is_instance_of::(py)); + assert_eq!(error.to_string(), "PanicException: route future panicked"); + }); + } + + #[test] + fn sync_runner_maps_a_panicked_error_mapper() { + Python::initialize(); + Python::attach(|py| { + let error = run_sync::( + py, + async { Err(CoreError::InvalidRequest("invalid".to_string())) }, + panicking_error_mapper, + ) + .expect_err("panicked mapper should become a Python exception"); + + assert!(error.is_instance_of::(py)); + assert_eq!(error.to_string(), "PanicException: error mapper panicked"); + }); + } + + #[test] + fn async_runner_surfaces_serializer_panics() { + Python::initialize(); + Python::attach(|py| { + let module = PyModule::new(py, "runtime").expect("module should be created"); + module + .add_function( + wrap_pyfunction!(async_serialization_panic, &module) + .expect("function should wrap"), + ) + .expect("function should register"); + let locals = PyDict::new(py); + locals + .set_item("runtime", &module) + .expect("module should enter Python locals"); + let code = CString::new( + r#" +import asyncio + +async def exercise(): + try: + await runtime.async_serialization_panic() + except BaseException as error: + assert type(error).__name__ == "PanicException" + assert str(error) == "async serializer panicked" + else: + raise AssertionError("serializer panic was not raised") + +asyncio.run(exercise()) +"#, + ) + .expect("Python source should not contain null bytes"); + py.run(&code, Some(&locals), Some(&locals)) + .expect("serializer panic should reach the Python awaiter"); + }); + } } diff --git a/litellm-rust/crates/python-interop/src/lib.rs b/litellm-rust/crates/python-interop/src/lib.rs index df2bd260fdb..2e562bdae70 100644 --- a/litellm-rust/crates/python-interop/src/lib.rs +++ b/litellm-rust/crates/python-interop/src/lib.rs @@ -2,4 +2,4 @@ mod gil; mod marshal; pub use gil::{release_count, release_gil}; -pub use marshal::{from_py, to_py}; +pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py}; diff --git a/litellm-rust/crates/python-interop/src/marshal.rs b/litellm-rust/crates/python-interop/src/marshal.rs index c3d0638427c..cf34a64a89d 100644 --- a/litellm-rust/crates/python-interop/src/marshal.rs +++ b/litellm-rust/crates/python-interop/src/marshal.rs @@ -1,4 +1,8 @@ +use std::any::Any; +use std::panic::{AssertUnwindSafe, catch_unwind}; + use pyo3::exceptions::PyValueError; +use pyo3::panic::PanicException; use pyo3::prelude::*; use serde::Serialize; use serde::de::DeserializeOwned; @@ -18,3 +22,71 @@ where .map(Bound::unbind) .map_err(|error| PyValueError::new_err(error.to_string())) } + +pub struct Pythonized(pub T); + +impl<'py, T> IntoPyObject<'py> for Pythonized +where + T: Serialize, +{ + type Target = PyAny; + type Output = Bound<'py, PyAny>; + type Error = PyErr; + + fn into_pyobject(self, py: Python<'py>) -> PyResult { + catch_unwind(AssertUnwindSafe(|| to_py(py, &self.0))) + .map_err(panic_to_pyerr)? + .map(|value| value.into_bound(py)) + } +} + +pub fn panic_to_pyerr(payload: Box) -> PyErr { + let message = payload + .downcast_ref::() + .map(String::as_str) + .or_else(|| payload.downcast_ref::<&str>().copied()) + .unwrap_or("panic from Rust code"); + PanicException::new_err(message.to_string()) +} + +#[cfg(test)] +mod tests { + use serde::Serializer; + + use super::*; + + struct PanickingSerializer; + + impl Serialize for PanickingSerializer { + fn serialize(&self, _serializer: S) -> Result + where + S: Serializer, + { + panic!("serializer panicked") + } + } + + #[test] + fn pythonized_converts_on_the_attached_thread() { + Python::initialize(); + Python::attach(|py| { + let value: Vec = Pythonized(vec![1, 2, 3]) + .into_pyobject(py) + .and_then(|value| value.extract()) + .expect("value should convert"); + assert_eq!(value, vec![1, 2, 3]); + }); + } + + #[test] + fn pythonized_maps_serializer_panics_to_a_base_exception() { + Python::initialize(); + Python::attach(|py| { + let error = Pythonized(PanickingSerializer) + .into_pyobject(py) + .expect_err("serializer panic should become a Python exception"); + assert!(error.is_instance_of::(py)); + assert_eq!(error.to_string(), "PanicException: serializer panicked"); + }); + } +} diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 834f7d564a2..e220c449c5d 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -20,6 +20,7 @@ import litellm.types.utils from litellm._logging import _redact_string, verbose_logger from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES +from litellm.exceptions import APIError from litellm.litellm_core_utils.agentic_loop_settings import ( DEFAULT_MAX_AGENTIC_LOOPS, validated_max_agentic_loops, @@ -2400,10 +2401,27 @@ class BaseLLMHTTPHandler: 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 + except Exception as rust_error: # noqa: BLE001 + from litellm.rust_bridge import get_native_bridge + + native_bridge: Final = get_native_bridge() + declined = getattr(native_bridge, "RustBridgeDeclined", None) + upstream_failed = getattr(native_bridge, "RustUpstreamError", None) + if isinstance(upstream_failed, type) and 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 messages: {message}", + llm_provider=custom_llm_provider, + model=model, + ) from rust_error + if not isinstance(declined, type) or not isinstance(rust_error, declined): + raise verbose_logger.debug( - "Rust Anthropic messages bridge raised %s; falling back to Python path", - type(rust_error).__name__, + "Rust Anthropic messages bridge declined before calling the provider (%s); falling back to Python path", + rust_error, ) return None if rust_response is None: 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 fbd7e36e298..688028ba933 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -1,12 +1,14 @@ """Tests for the optional Rust-backed Anthropic Messages path.""" import importlib +from types import ModuleType from typing import cast import httpx import pytest import litellm +from litellm.exceptions import APIError from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, @@ -99,12 +101,28 @@ class ExplodingAsyncMessages: class RaisingAsyncMessages: - def __init__(self) -> None: + def __init__(self, error: Exception) -> None: self.calls = 0 + self.error = error async def __call__(self, **kwargs: object) -> dict[str, object]: self.calls += 1 - raise RuntimeError("upstream request failed with status 400: bad request") + raise self.error + + +class FakeBridgeDeclined(Exception): + pass + + +class FakeUpstreamError(Exception): + pass + + +def _install_fake_bridge_exceptions(monkeypatch) -> None: + native_bridge = ModuleType("_native") + native_bridge.RustBridgeDeclined = FakeBridgeDeclined + native_bridge.RustUpstreamError = FakeUpstreamError + monkeypatch.setattr(rust_bridge_loader, "_cached_bridge", native_bridge) @pytest.fixture(autouse=True) @@ -251,8 +269,9 @@ async def test_gate_invokes_rust_and_marks_response_header(): @pytest.mark.asyncio -async def test_gate_falls_back_to_python_when_bridge_raises(): - bridge = RaisingAsyncMessages() +async def test_gate_falls_back_only_when_bridge_declines(monkeypatch): + _install_fake_bridge_exceptions(monkeypatch) + bridge = RaisingAsyncMessages(FakeBridgeDeclined("unsupported request")) litellm.use_litellm_rust(True, amessages=bridge) response = await _gate() @@ -261,6 +280,31 @@ async def test_gate_falls_back_to_python_when_bridge_raises(): assert bridge.calls == 1 +@pytest.mark.asyncio +async def test_gate_surfaces_an_upstream_failure_without_fallback(monkeypatch): + _install_fake_bridge_exceptions(monkeypatch) + bridge = RaisingAsyncMessages(FakeUpstreamError(429, "429: rate limited")) + litellm.use_litellm_rust(True, amessages=bridge) + + with pytest.raises(APIError) as exc_info: + await _gate() + + assert exc_info.value.status_code == 429 + assert "429: rate limited" in str(exc_info.value) + assert bridge.calls == 1 + + +@pytest.mark.asyncio +async def test_gate_reraises_an_unknown_bridge_failure(): + bridge = RaisingAsyncMessages(RuntimeError("unknown bridge failure")) + litellm.use_litellm_rust(True, amessages=bridge) + + with pytest.raises(RuntimeError, match="unknown bridge failure"): + await _gate() + + assert bridge.calls == 1 + + @pytest.mark.asyncio async def test_gate_skips_rust_when_flag_absent(): bridge = ExplodingAsyncMessages()