diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 0b04319d966..c998c71ebf0 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -64,6 +64,7 @@ async fn execute_ocr_provider_call( let decoded = adapter .read_response(client, response, &url, &headers, &request) .await?; + request.hooks.post_call(&decoded.text).await?; let response = adapter.transform_ocr_response(&request, decoded.data)?; Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs index 3ad706648d8..204ca4cf8b3 100644 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -28,7 +28,7 @@ pub struct OcrDuringCallRequest { pub body: Value, } -pub struct OcrPreparedRequest { +pub struct OcrRequestDraft { pub model: String, pub url: String, pub headers: Vec<(String, String)>, @@ -36,12 +36,12 @@ pub struct OcrPreparedRequest { } pub trait OcrHooks: Send + Sync { - fn prepared_request( - &self, - request: OcrPreparedRequest, - ) -> OcrHookFuture<'_, OcrPreparedRequest> { + fn before_send(&self, request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> { Box::pin(async move { Ok(request) }) } + fn post_call<'a>(&'a self, _response: &'a str) -> OcrHookFuture<'a, ()> { + Box::pin(async { Ok(()) }) + } fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> { Box::pin(async move { Ok(request) }) } diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 64c4191bb12..3c5ea40e899 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -3,7 +3,7 @@ use serde_json::{Map, Value}; use super::OcrClient; use super::error::{OcrError, OcrRequestError}; -use super::hooks::{OcrDuringCallRequest, OcrPreparedRequest}; +use super::hooks::{OcrDuringCallRequest, OcrRequestDraft}; use super::types::{LiteLLMOcrRequest, OcrDocument}; #[derive(Debug, Deserialize)] @@ -94,9 +94,9 @@ pub(crate) async fn build_http_request( where B: Serialize + DeserializeOwned, { - let prepared = request + let draft = request .hooks - .prepared_request(OcrPreparedRequest { + .before_send(OcrRequestDraft { model: request.model.clone(), url: url.into(), headers: headers.to_vec(), @@ -105,15 +105,15 @@ where })?, }) .await?; - let body: B = super::wire::decode_request_value(prepared.body, "guardrail.body")?; + let body: B = super::wire::decode_request_value(draft.body, "guardrail.body")?; let builder = client .provider_http() - .post(&prepared.url) + .post(&draft.url) .json(&body) .timeout(request.connection.timeout); crate::http_utils::with_headers( builder, - &prepared.headers, + &draft.headers, crate::http_utils::HeaderPolicy::All, ) .build() diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index 215c23735eb..7c2f903899a 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -17,6 +17,7 @@ use serde_json::{Map, Value}; pub struct DecodedOcrResponse { pub data: T, pub native: Option, + pub text: String, } #[derive(Deserialize)] @@ -135,7 +136,11 @@ pub fn decode_response( } else { None }; - Ok(DecodedOcrResponse { data, native }) + Ok(DecodedOcrResponse { + data, + native, + text: String::from_utf8_lossy(bytes).into_owned(), + }) } pub fn decode_pre_call_result( diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index 1ceb12df64a..d60af7303e0 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -3,7 +3,7 @@ use std::sync::{Arc, Mutex}; use serde_json::{Value, json}; use super::OcrClient; -use super::hooks::{OcrHookFuture, OcrHooks, OcrLogFuture, OcrPreCallRequest, OcrPreparedRequest}; +use super::hooks::{OcrHookFuture, OcrHooks, OcrLogFuture, OcrPreCallRequest, OcrRequestDraft}; use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; use super::wire::{OcrWireRequest, decode_request}; use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming}; @@ -123,13 +123,10 @@ struct RecordingHooks { block: bool, } -struct EditPreparedRequest; +struct EditRequestDraft; -impl OcrHooks for EditPreparedRequest { - fn prepared_request( - &self, - mut request: OcrPreparedRequest, - ) -> OcrHookFuture<'_, OcrPreparedRequest> { +impl OcrHooks for EditRequestDraft { + fn before_send(&self, mut request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> { Box::pin(async move { assert_eq!(request.model, "model"); assert!(request.url.ends_with("/v1/ocr")); @@ -148,14 +145,14 @@ impl OcrHooks for EditPreparedRequest { } #[tokio::test] -async fn prepared_request_hook_edits_wire_body_and_headers_without_guardrails() { +async fn before_send_hook_edits_wire_body_and_headers_without_guardrails() { let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; let request = wire_request( "mistral/model", &base, json!({"include_image_base64":false}), ) - .with_host_hooks(Arc::new(EditPreparedRequest), None); + .with_host_hooks(Arc::new(EditRequestDraft), None); perform_ocr(request).await.unwrap(); server.await.unwrap(); let requests = seen.lock().unwrap(); @@ -166,12 +163,9 @@ async fn prepared_request_hook_edits_wire_body_and_headers_without_guardrails() } impl OcrHooks for RecordingHooks { - fn prepared_request( - &self, - request: OcrPreparedRequest, - ) -> OcrHookFuture<'_, OcrPreparedRequest> { + fn before_send(&self, request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> { Box::pin(async move { - self.events.lock().unwrap().push("prepared"); + self.events.lock().unwrap().push("before_send"); Ok(request) }) } @@ -235,7 +229,7 @@ async fn lifecycle_orders_hooks_and_emits_one_success() { server.await.unwrap(); assert_eq!( *events.lock().unwrap(), - ["pre", "during", "prepared", "success"] + ["pre", "during", "before_send", "success"] ); assert_eq!(seen.lock().unwrap().len(), 1); } @@ -277,7 +271,7 @@ async fn upstream_failure_emits_one_terminal_failure() { server.await.unwrap(); assert_eq!( *events.lock().unwrap(), - ["pre", "during", "prepared", "failure"] + ["pre", "during", "before_send", "failure"] ); assert_eq!(seen.lock().unwrap().len(), 1); } @@ -333,7 +327,7 @@ async fn every_adapter_runs_the_complete_lifecycle() { server.await.unwrap(); assert_eq!( *events.lock().unwrap(), - ["pre", "during", "prepared", "success"], + ["pre", "during", "before_send", "success"], "{model}" ); assert_eq!(seen.lock().unwrap().len(), 1, "{model}"); @@ -343,7 +337,7 @@ async fn every_adapter_runs_the_complete_lifecycle() { #[derive(Clone, Copy, Debug)] enum FailureStage { During, - Prepared, + BeforeSend, InvalidBody, Preparation, Response, @@ -372,17 +366,14 @@ impl OcrHooks for FailingHooks { }) } - fn prepared_request( - &self, - request: OcrPreparedRequest, - ) -> OcrHookFuture<'_, OcrPreparedRequest> { + fn before_send(&self, request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> { Box::pin(async move { - let request = self.recording.prepared_request(request).await?; + let request = self.recording.before_send(request).await?; match self.stage { - FailureStage::Prepared => Err(crate::Error::InvalidRequest( - "blocked prepared request".into(), + FailureStage::BeforeSend => Err(crate::Error::InvalidRequest( + "blocked before_send request".into(), )), - FailureStage::InvalidBody => Ok(OcrPreparedRequest { + FailureStage::InvalidBody => Ok(OcrRequestDraft { body: json!({"document":null}), ..request }), @@ -416,7 +407,7 @@ impl OcrHooks for FailingHooks { #[rstest::rstest] #[case(FailureStage::During)] -#[case(FailureStage::Prepared)] +#[case(FailureStage::BeforeSend)] #[case(FailureStage::InvalidBody)] #[case(FailureStage::Preparation)] #[case(FailureStage::Response)] @@ -454,7 +445,7 @@ async fn lifecycle_reports_failures_once_at_each_boundary(#[case] stage: Failure let expected = match stage { FailureStage::Preparation => vec!["pre", "failure"], FailureStage::During => vec!["pre", "during", "failure"], - _ => vec!["pre", "during", "prepared", "failure"], + _ => vec!["pre", "during", "before_send", "failure"], }; assert_eq!(*events.lock().unwrap(), expected); assert_eq!( diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index dc04f15ebfc..ce8b2239534 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -89,9 +89,9 @@ pub(crate) fn ocr_route_error(err: Error) -> BridgeError { Error::MissingField("document_url" | "image_url") => { BridgeError::InvalidArgument("Document URL is required".into()) } - Error::Http { status, .. } => BridgeError::Upstream { + Error::Http { status, body } => BridgeError::Upstream { status: Some(status), - message: "OCR provider request failed".into(), + message: body, }, other => required_route_error(other), } @@ -147,7 +147,7 @@ mod ocr_error_tests { } #[test] - fn ocr_errors_preserve_status_without_provider_body() { + fn ocr_errors_preserve_provider_body() { Python::initialize(); Python::attach(|py| { for field in ["document_url", "image_url"] { @@ -166,7 +166,7 @@ mod ocr_error_tests { .getattr("args") .and_then(|args| args.extract()) .expect("OCR failures retain status and unprefixed provider message"); - assert_eq!(args, (429, "OCR provider request failed".to_string())); + assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string())); }); } } diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 0a2227b0110..9855e3e2ac6 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -225,7 +225,7 @@ mod tests { ( "ocr", "aocr", - "(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, input_sources=None, timeout_seconds=None, logging_obj=None, callback_loop=None, token_provider=None)", + "(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, input_sources=None, timeout_seconds=None, logging_obj=None, callback_loop=None, token_provider=None, call_completion=None)", ), ( "transcription", diff --git a/litellm-rust/crates/python-bridge/src/routes/messages.rs b/litellm-rust/crates/python-bridge/src/routes/messages.rs index cdd5cb6c8f2..495990cac43 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages.rs @@ -5,7 +5,7 @@ use pyo3::prelude::*; use serde_json::Value; use std::future::Future; -use crate::errors::fallback_route_error; +use crate::errors::required_route_error; use crate::marshal::{RouteOptions, RouteOptionsInputs, required_value}; fn prepare_messages( @@ -61,5 +61,5 @@ bridge_route! { timeout_seconds: Option, }, prepare = prepare_messages, - errors = fallback_route_error, + errors = required_route_error, } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index 1f7458bcb5a..ea146b2e67a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -42,6 +42,7 @@ fn prepare_ocr( }) .transpose()?, error: Mutex::new(None), + execution_body: Mutex::new(None), })) }) }) @@ -84,6 +85,10 @@ fn prepare_ocr( timeout_seconds: timeout.map(|value| value.as_secs_f64()), }) .map_err(ocr_route_error)?; + if let Some(call_completion) = &inputs.call_completion { + callbacks::NativeOcrCompletion::attach(call_completion) + .map_err(BridgeError::Host)?; + } if let Some(hooks) = &hooks && hooks.token_provider.is_some() { @@ -151,6 +156,7 @@ bridge_route! { logging_obj: Option>, callback_loop: Option>, token_provider: Option>, + call_completion: Option>, }, prepare = prepare_ocr, errors = std::convert::identity, 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 0556907d230..e8f0cfcfe38 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs @@ -2,7 +2,7 @@ use std::sync::Mutex; use litellm_core::Error; use litellm_core::auth::{ResolvedCredential, SecretValue, TokenFuture, TokenProvider}; -use litellm_core::ocr::hooks::{OcrHookFuture, OcrHooks, OcrPreparedRequest}; +use litellm_core::ocr::hooks::{OcrHookFuture, OcrHooks, OcrRequestDraft}; use litellm_python_interop::{from_py, to_py}; use pyo3::prelude::*; use pyo3::types::{PyDict, PyTuple}; @@ -30,6 +30,77 @@ pub(super) struct PythonOcrHooks { pub api_key: Option, pub locals: Option, pub error: Mutex>, + pub execution_body: Mutex>>, +} + +#[pyclass] +pub(super) struct NativeOcrCompletion { + python_completion: Py, +} + +impl NativeOcrCompletion { + pub(super) fn attach(call_completion: &Py) -> PyResult<()> { + Python::attach(|py| { + let python_completion = call_completion.getattr(py, "python_implementation")?; + let native = Py::new(py, Self { python_completion })?; + let attached: bool = call_completion + .call_method1(py, "attach", (native,))? + .extract(py)?; + if attached { + Ok(()) + } else { + Err(pyo3::exceptions::PyRuntimeError::new_err( + "OCR completion was already attached", + )) + } + }) + } +} + +#[pymethods] +impl NativeOcrCompletion { + fn success( + &self, + py: Python<'_>, + result: Py, + start_time: Py, + end_time: Py, + ) -> PyResult<()> { + self.python_completion + .call_method1(py, "success", (result, start_time, end_time))?; + Ok(()) + } + + fn failure( + &self, + py: Python<'_>, + exception: Py, + traceback_exception: String, + start_time: Py, + end_time: Py, + ) -> PyResult<()> { + self.python_completion.call_method1( + py, + "failure", + (exception, traceback_exception, start_time, end_time), + )?; + Ok(()) + } + + fn async_failure( + &self, + py: Python<'_>, + exception: Py, + traceback_exception: String, + start_time: Py, + end_time: Py, + ) -> PyResult> { + self.python_completion.call_method1( + py, + "async_failure", + (exception, traceback_exception, start_time, end_time), + ) + } } impl PythonOcrHooks { @@ -75,7 +146,7 @@ impl PythonOcrHooks { fn callback( &self, py: Python<'_>, - request: &OcrPreparedRequest, + request: &OcrRequestDraft, ) -> PyResult<(Py, Py, Py)> { let body = to_py(py, &request.body)? .into_bound(py) @@ -113,19 +184,23 @@ impl PythonOcrHooks { } impl OcrHooks for PythonOcrHooks { - fn prepared_request( - &self, - request: OcrPreparedRequest, - ) -> OcrHookFuture<'_, OcrPreparedRequest> { + fn before_send(&self, request: OcrRequestDraft) -> OcrHookFuture<'_, OcrRequestDraft> { Box::pin(async move { if self.logger.is_none() { return Ok(request); } - let result: PyResult = async { + let result: PyResult = async { let (callback, body, headers) = Python::attach(|py| self.callback(py, &request))?; self.invoke(callback).await?; Python::attach(|py| { - Ok(OcrPreparedRequest { + *self + .execution_body + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) = + Some(body.clone_ref(py)); + }); + Python::attach(|py| { + Ok(OcrRequestDraft { body: from_py(body.bind(py).as_any())?, headers: headers .bind(py) @@ -143,6 +218,46 @@ impl OcrHooks for PythonOcrHooks { }) }) } + + fn post_call<'a>(&'a self, response: &'a str) -> OcrHookFuture<'a, ()> { + Box::pin(async move { + let Some(logger) = &self.logger else { + return Ok(()); + }; + let result: PyResult<()> = async { + let callback = Python::attach(|py| { + let kwargs = PyDict::new(py); + kwargs.set_item("api_key", &self.api_key)?; + kwargs.set_item("original_response", response)?; + let additional = PyDict::new(py); + let body = self + .execution_body + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + additional.set_item( + "complete_input_dict", + body.as_ref().ok_or_else(|| { + pyo3::exceptions::PyRuntimeError::new_err("missing OCR execution body") + })?, + )?; + kwargs.set_item("additional_args", additional)?; + Ok::<_, PyErr>( + py.import("functools")? + .getattr("partial")? + .call((logger.bind(py).getattr("post_call")?,), Some(&kwargs))? + .unbind(), + ) + })?; + self.invoke(callback).await?; + Ok(()) + } + .await; + result.map_err(|error| { + self.retain_error(error); + Error::InvalidRequest("OCR post-call hook failed".into()) + }) + }) + } } impl std::fmt::Debug for PythonOcrHooks { diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 145d9dd2304..7109e6942d1 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2476,21 +2476,12 @@ class BaseLLMHTTPHandler: extra_headers=headers, timeout=timeout, ) - except Exception as rust_error: # noqa: BLE001 # only explicit pre-dispatch declines permit fallback - from litellm.rust_bridge.bindings import native_exception_types, upstream_error_details - - exceptions: Final = native_exception_types() - if exceptions is not None and isinstance(rust_error, exceptions[0]): - return None - if exceptions is not None and isinstance(rust_error, exceptions[1]): - status, message = upstream_error_details(rust_error) - raise litellm.APIError( - status_code=status, - message=message, - llm_provider=custom_llm_provider, - model=model, - ) from rust_error - raise + 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 if rust_response is None: return None diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index a75fa0de5f0..bef09cacfcd 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -26,7 +26,6 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned -from litellm.rust_bridge.bindings import native_exception_types, upstream_error_details from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.loader import get_native_bridge from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -275,6 +274,17 @@ def rust_chat_completions_accepts( return True +def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None: + native_bridge: Final = get_native_bridge() + if native_bridge is None: + return None + declined: Final = getattr(native_bridge, "RustBridgeDeclined", None) + upstream: Final = getattr(native_bridge, "RustUpstreamError", None) + if declined is None or upstream is None: + return None + return declined, upstream + + def _reraise_or_decline( rust_error: BaseException, *, @@ -288,14 +298,20 @@ def _reraise_or_decline( second attempt bills for it twice. Those surface as an `APIError` carrying the upstream status, which LiteLLM's exception mapping already understands. """ - exceptions: Final = native_exception_types() + exceptions: Final = _rust_bridge_exceptions() if exceptions is None: - raise rust_error + 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): - status, message = upstream_error_details(rust_error) + 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=status, + status_code=int(status) or 500, message=f"litellm rust chat completions: {message}", llm_provider=custom_llm_provider or "", model=model, diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index d5f576e3c7c..6fbeb46c272 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -73,6 +73,7 @@ class RustOcr(Protocol): logging_obj: _OCRLogging | None, callback_loop: asyncio.AbstractEventLoop | None, token_provider: object, + call_completion: object, ) -> dict[str, object]: raise NotImplementedError @@ -92,6 +93,7 @@ class RustAocr(Protocol): logging_obj: _OCRLogging | None, callback_loop: asyncio.AbstractEventLoop | None, token_provider: object, + call_completion: object, ) -> Awaitable[dict[str, object]]: raise NotImplementedError @@ -347,10 +349,11 @@ def run( optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict input_sources=marshalled.input_sources, timeout=marshalled.timeout, - logging_obj=cast( + logging_obj=cast( # cast-ok: client decorator injects Logging _OCRLogging, request.kwargs["litellm_logging_obj"] - ), # cast-ok: client decorator injects Logging + ), token_provider=request.kwargs.get("azure_ad_token_provider"), + call_completion=request.call_completion, ) except Exception as error: mapped: Final = _map_error(error, request) @@ -379,10 +382,11 @@ async def arun( optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict input_sources=marshalled.input_sources, timeout=marshalled.timeout, - logging_obj=cast( + logging_obj=cast( # cast-ok: client decorator injects Logging _OCRLogging, request.kwargs["litellm_logging_obj"] - ), # cast-ok: client decorator injects Logging + ), token_provider=request.kwargs.get("azure_ad_token_provider"), + call_completion=request.call_completion, ) except Exception as error: mapped: Final = _map_error(error, request) @@ -405,6 +409,7 @@ def ocr( input_sources: Mapping[str, str] | None = None, logging_obj: _OCRLogging | None = None, token_provider: object = None, + call_completion: object = None, ) -> dict[str, object] | None: rust_ocr: Final = load_rust_ocr() if rust_ocr is None: @@ -422,6 +427,7 @@ def ocr( logging_obj=logging_obj, callback_loop=None, token_provider=token_provider, + call_completion=call_completion, ) @@ -438,6 +444,7 @@ async def aocr( input_sources: Mapping[str, str] | None = None, logging_obj: _OCRLogging | None = None, token_provider: object = None, + call_completion: object = None, ) -> dict[str, object] | None: rust_aocr: Final = load_rust_aocr() if rust_aocr is None: @@ -455,4 +462,5 @@ async def aocr( logging_obj=logging_obj, callback_loop=asyncio.get_running_loop(), token_provider=token_provider, + call_completion=call_completion, ) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py index 9ee6cf10717..d411673439f 100644 --- a/litellm/rust_bridge/runtime.py +++ b/litellm/rust_bridge/runtime.py @@ -6,7 +6,7 @@ from enum import Enum from typing import Final, Generic, NoReturn, TypeAlias, TypeVar from litellm.exceptions import APIError -from litellm.rust_bridge.bindings import native_exception_types, upstream_error_details +from litellm.rust_bridge.bindings import native_exception_types NativeT = TypeVar("NativeT") ResultT = TypeVar("ResultT") @@ -137,9 +137,13 @@ def _required_reason(result: RustDeclined | RustUnavailable) -> str: def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn: - status, message = upstream_error_details(error) + args: Final[tuple[object, ...]] = error.args + status_value: Final = args[0] if args else 0 + message_value: Final = args[1] if len(args) > 1 else str(error) + status: Final = status_value if isinstance(status_value, int) else 0 + message: Final = message_value if isinstance(message_value, str) else str(message_value) raise APIError( - status_code=status, + status_code=status or 500, message=f"litellm rust {context.route}: {message}", llm_provider=context.provider, model=context.model, 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 0b68dda027e..a30474245c6 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -238,66 +238,17 @@ async def test_gate_invokes_rust_and_marks_response_header(): @pytest.mark.asyncio -async def test_gate_propagates_unclassified_bridge_failure(): +async def test_gate_falls_back_to_python_when_bridge_raises(): bridge = RaisingAsyncMessages() litellm.rust(True) rust_messages.set_rust_messages(amessages=bridge) - with pytest.raises(RuntimeError, match="upstream request failed"): - await _gate() + response = await _gate() + + assert response is None assert bridge.calls == 1 -@pytest.mark.asyncio -@pytest.mark.parametrize("status", [0, 400, 429, 500]) -async def test_gate_never_falls_back_after_possible_dispatch(monkeypatch, status): - from types import SimpleNamespace - from litellm.rust_bridge import bindings - - class Declined(Exception): - pass - - class Upstream(Exception): - pass - - error = Upstream(status, "provider failure") - - async def bridge(**kwargs): - raise error - - monkeypatch.setattr( - bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) - ) - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - with pytest.raises(litellm.APIError) as caught: - await _gate() - assert caught.value.status_code == (status or 500) - assert caught.value.__cause__ is error - - -@pytest.mark.asyncio -async def test_gate_falls_back_only_for_explicit_decline(monkeypatch): - from types import SimpleNamespace - from litellm.rust_bridge import bindings - - class Declined(Exception): - pass - - class Upstream(Exception): - pass - - async def bridge(**kwargs): - raise Declined("unsupported before dispatch") - - monkeypatch.setattr( - bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) - ) - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - assert await _gate() is None - - @pytest.mark.asyncio async def test_gate_skips_rust_when_flag_absent(): bridge = ExplodingAsyncMessages() diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index f946d489ecc..29bd5ceb51e 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -65,6 +65,7 @@ class RecordingBridge: logging_obj: object = None, callback_loop: asyncio.AbstractEventLoop | None = None, token_provider: object = None, + call_completion: object = None, ) -> dict[str, object]: self.logging_obj = logging_obj self.calls.append( @@ -103,6 +104,7 @@ class RecordingAsyncBridge: logging_obj: object = None, callback_loop: asyncio.AbstractEventLoop | None = None, token_provider: object = None, + call_completion: object = None, ) -> dict[str, object]: self.calls.append( { @@ -135,6 +137,7 @@ class RaisingBridge: logging_obj: object = None, callback_loop: asyncio.AbstractEventLoop | None = None, token_provider: object = None, + call_completion: object = None, ) -> dict[str, object]: raise RuntimeError("bridge failed") @@ -154,6 +157,7 @@ class RaisingAsyncBridge: logging_obj: object = None, callback_loop: asyncio.AbstractEventLoop | None = None, token_provider: object = None, + call_completion: object = None, ) -> dict[str, object]: raise RuntimeError("bridge failed") diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index 36ec234ba19..b2fd2e6dcc0 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -55,9 +55,6 @@ class _FakeNative: def _fake_native_bridge(monkeypatch): """Expose the bridge's exception classes without the compiled extension.""" monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) - from litellm.rust_bridge import bindings - - monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) def _hide_native_bridge(monkeypatch): diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index 5098d2f4ad8..3580ea061ad 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -149,7 +149,7 @@ def test_native_ocr_normalizes_provider_response_model_and_usage(ocr_server: Rec assert response.usage_info.pages_processed == 1 -def test_native_ocr_maps_provider_400_without_exposing_response_body(ocr_server: RecordingServer) -> None: +def test_native_ocr_maps_provider_400_with_response_body(ocr_server: RecordingServer) -> None: ocr_server.enqueue(ResponseSpec(body={"message": "invalid OCR request"}, status=400)) with pytest.raises(litellm.BadRequestError) as caught: @@ -158,7 +158,7 @@ def test_native_ocr_maps_provider_400_without_exposing_response_body(ocr_server: assert caught.value.status_code == 400 assert caught.value.model == "mistral-ocr-latest" assert caught.value.llm_provider == "mistral" - assert "invalid OCR request" not in str(caught.value) + assert "invalid OCR request" in str(caught.value) def test_native_ocr_raises_transport_error_when_request_exceeds_timeout(ocr_server: RecordingServer) -> None: