From a67e94162a51ae3cfece67cc5bc3ffb2d3e1570b Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 12 Sep 2026 08:21:05 -0700 Subject: [PATCH] refactor --- .../crates/python-bridge/src/lifecycle.rs | 252 +++++++------ .../python-bridge/src/lifecycle/bindings.rs | 334 ++++++++++++++++++ .../src/lifecycle/preparation.rs | 4 +- .../crates/python-bridge/src/routes/mod.rs | 1 + .../python-bridge/src/routes/ocr_callbacks.rs | 275 ++++++++++++++ .../python-bridge/src/routes/ocr_lifecycle.rs | 89 ++--- litellm/rust_bridge/ocr_lifecycle.py | 14 - 7 files changed, 753 insertions(+), 216 deletions(-) create mode 100644 litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs diff --git a/litellm-rust/crates/python-bridge/src/lifecycle.rs b/litellm-rust/crates/python-bridge/src/lifecycle.rs index 36efcfa041d..6dd7b9f32d4 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle.rs @@ -14,8 +14,12 @@ use tokio::sync::Mutex; use crate::errors::ocr_error_to_pyerr; use crate::execution::{run_async_value, run_sync_value}; +mod bindings; mod preparation; +use bindings::DeploymentHooks; +pub(crate) use bindings::PythonLogger; + pub(crate) trait PythonRoute: Send + Sync { fn state(&self) -> &PythonCallState; fn state_mut(&mut self) -> &mut PythonCallState; @@ -401,7 +405,7 @@ impl Drop for PythonLifecycle { pub(crate) struct PythonCallState { pub args: Py, pub kwargs: Py, - pub logger: Option>, + pub logger: Option, pub start: Py, pub end: Option>, pub response: Option>, @@ -427,32 +431,31 @@ impl PythonCallState { match phase { HostPhase::Setup => self.setup(py)?, HostPhase::DeploymentPreCall => { - return Ok(HostStep::Suspend( - py.import("litellm.utils")? - .getattr("async_pre_call_deployment_hook")? - .call1((&self.kwargs, self.call_type))? - .unbind(), - )); + return Ok(HostStep::Suspend(DeploymentHooks::before_call( + py, + &self.kwargs, + self.call_type, + )?)); } HostPhase::Prepare => self.prepare(py)?, HostPhase::DeploymentPostCall => { - return Ok(HostStep::Suspend( - py.import("litellm.utils")? - .getattr("async_post_call_success_deployment_hook")? - .call1((&self.kwargs, &self.response, self.call_type))? - .unbind(), - )); + return Ok(HostStep::Suspend(DeploymentHooks::after_success( + py, + &self.kwargs, + &self.response, + self.call_type, + )?)); } HostPhase::Finalize => self.finalize(py)?, HostPhase::Success => self.dispatch_success(py)?, HostPhase::DeploymentFailure => { if let Some(error) = &self.error { - return Ok(HostStep::Suspend( - py.import("litellm.utils")? - .getattr("async_post_call_failure_deployment_hook")? - .call1((&self.kwargs, error, self.call_type))? - .unbind(), - )); + return Ok(HostStep::Suspend(DeploymentHooks::after_failure( + py, + &self.kwargs, + error, + self.call_type, + )?)); } } HostPhase::Failure | HostPhase::AsyncFailure => { @@ -502,59 +505,48 @@ impl PythonCallState { }) } - pub fn logger<'py>(&self, py: Python<'py>) -> PyResult> { - self.logger - .as_ref() - .map(|value| value.bind(py).clone()) - .ok_or_else(|| { - pyo3::exceptions::PyRuntimeError::new_err("call logging is not initialized") - }) + pub fn logger(&self) -> PyResult<&PythonLogger> { + self.logger.as_ref().ok_or_else(|| { + pyo3::exceptions::PyRuntimeError::new_err("call logging is not initialized") + }) } pub fn setup(&mut self, py: Python<'_>) -> PyResult<()> { self.start = now(py)?; - self.internal = py - .import("litellm._internal_context")? - .getattr("is_internal_call")? - .call_method0("get")? - .extract()?; - let result = py - .import("litellm.rust_bridge.lifecycle")? - .getattr("setup")? - .call1(( - self.call_type, - &self.args, - &self.kwargs, - &self.start, - self.asynchronous, - ))?; - self.logger = Some(result.getattr("logger")?.unbind()); - self.kwargs = result.getattr("kwargs")?.cast_into::()?.unbind(); + self.internal = bindings::is_internal_call(py)?; + let result = bindings::setup( + py, + self.call_type, + &self.args, + &self.kwargs, + &self.start, + self.asynchronous, + )?; + self.logger = Some(result.logger()?); + self.kwargs = result.kwargs()?; Ok(()) } pub fn prepare(&mut self, py: Python<'_>) -> PyResult<()> { - self.kwargs = preparation::prepare(py, self.kwargs.bind(py), &self.logger(py)?)?.unbind(); + self.kwargs = preparation::prepare(py, self.kwargs.bind(py), self.logger()?)?.unbind(); Ok(()) } pub fn finalize(&self, py: Python<'_>) -> PyResult<()> { - py.import("litellm.rust_bridge.lifecycle")? - .getattr("finalize")? - .call1(( - &self.response, - self.logger(py)?, - &self.kwargs, - &self.start, - &self.end, - ))?; - Ok(()) + bindings::finalize( + py, + &self.response, + self.logger()?, + &self.kwargs, + &self.start, + &self.end, + ) } pub fn dispatch_success(&self, py: Python<'_>) -> PyResult<()> { match self.try_dispatch_success(py) { Err(error) if error.is_instance_of::(py) => { - error.write_unraisable(py, self.logger.as_ref().map(|logger| logger.bind(py))); + error.write_unraisable(py, self.logger.as_ref().map(|logger| logger.object(py))); Ok(()) } result => result, @@ -562,9 +554,9 @@ impl PythonCallState { } fn try_dispatch_success(&self, py: Python<'_>) -> PyResult<()> { - let logger = self.logger(py)?; + let logger = self.logger()?; let pending = PendingSuccess { - logger: logger.clone().unbind(), + logger: logger.clone_ref(py), response: self.response.as_ref().map(|value| value.clone_ref(py)), start: self.start.clone_ref(py), end: self.end.as_ref().map(|value| value.clone_ref(py)), @@ -579,12 +571,9 @@ impl PythonCallState { .get_item("fallbacks")? .is_none_or(|value| value.is_none()) { - if logger - .getattr("_defer_async_logging") - .is_ok_and(|value| value.is_truthy().unwrap_or(false)) - { - logger.setattr( - "_native_pending_logging", + if logger.defers_async_logging(py) { + logger.defer_success( + py, Py::new( py, PendingLogging { @@ -596,12 +585,7 @@ impl PythonCallState { pending.asynchronous(py)?; } } - logger - .call_method1( - "handle_sync_success_callbacks_for_async_calls", - (&self.response, &self.start, &self.end), - ) - .map(|_| ()) + logger.sync_success_for_async_call(py, &self.response, &self.start, &self.end) } } @@ -616,29 +600,13 @@ impl PythonCallState { let Some(error) = &self.error else { return Ok(None); }; - let trace = py - .import("traceback")? - .getattr("format_exception")? - .call1((error,))?; - let trace = pyo3::types::PyString::new(py, "").call_method1("join", (trace,))?; - let value = self.logger(py)?.call_method1( - if asynchronous { - "async_failure_handler" - } else { - "failure_handler" - }, - (error, trace, &self.start, &self.end), - )?; - Ok(asynchronous.then(|| value.unbind())) + self.logger()? + .failure(py, error, &self.start, &self.end, asynchronous) } pub fn cleanup(&mut self, py: Python<'_>) { if let Some(logger) = self.logger.take() - && let Err(error) = py.import("litellm.utils").and_then(|utils| { - utils - .getattr("_restore_correlation_context_if_supported")? - .call1((logger,)) - }) + && let Err(error) = logger.restore_context(py) { error.write_unraisable(py, None); } @@ -651,7 +619,9 @@ impl PythonCallState { pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.args)?; visit.call(&self.kwargs)?; - visit.call(&self.logger)?; + if let Some(logger) = &self.logger { + logger.traverse(visit)?; + } visit.call(&self.start)?; visit.call(&self.end)?; visit.call(&self.response)?; @@ -660,54 +630,21 @@ impl PythonCallState { } struct PendingSuccess { - logger: Py, + logger: PythonLogger, response: Option>, start: Py, end: Option>, } impl PendingSuccess { - fn context(py: Python<'_>) -> PyResult> { - py.import("contextvars")? - .call_method0("copy_context") - .map(Bound::unbind) - } - fn sync(&self, py: Python<'_>) -> PyResult<()> { - let context = Self::context(py)?; - py.import("litellm.litellm_core_utils.litellm_logging")? - .getattr("executor")? - .call_method1( - "submit", - ( - context.getattr(py, "run")?, - self.logger.getattr(py, "success_handler")?, - &self.response, - &self.start, - &self.end, - ), - )?; - Ok(()) + self.logger + .submit_success(py, &self.response, &self.start, &self.end) } fn asynchronous(&self, py: Python<'_>) -> PyResult<()> { - let context = Self::context(py)?; - let worker = py - .import("litellm.litellm_core_utils.logging_worker")? - .getattr("GLOBAL_LOGGING_WORKER")? - .getattr("ensure_initialized_and_enqueue")?; - let coroutine = self.logger.call_method1( - py, - "async_success_handler", - (&self.response, &self.start, &self.end), - )?; - let enqueue = context.call_method1(py, "run", (worker, &coroutine)); - if enqueue.is_err() - && let Err(error) = coroutine.call_method0(py, "close") - { - error.write_unraisable(py, Some(coroutine.bind(py))); - } - enqueue.map(|_| ()) + self.logger + .enqueue_success(py, &self.response, &self.start, &self.end) } } @@ -725,7 +662,7 @@ impl PendingLogging { { match pending.asynchronous(py) { Err(error) if error.is_instance_of::(py) => { - error.write_unraisable(py, Some(pending.logger.bind(py))); + error.write_unraisable(py, Some(pending.logger.object(py))); } result => return result, } @@ -735,7 +672,7 @@ impl PendingLogging { fn __traverse__(&self, visit: pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> { if let Some(pending) = &self.pending { - visit.call(&pending.logger)?; + pending.logger.traverse(&visit)?; visit.call(&pending.response)?; visit.call(&pending.start)?; visit.call(&pending.end)?; @@ -944,7 +881,7 @@ assert reference() is None PythonCallState { args: PyTuple::empty(py).unbind(), kwargs: PyDict::new(py).unbind(), - logger: Some(logger), + logger: Some(logger.extract(py).unwrap()), start: py.None(), end: Some(py.None()), response: Some(response), @@ -1079,7 +1016,12 @@ logger = Logger() py, PendingLogging { pending: Some(PendingSuccess { - logger: locals.get_item("logger").unwrap().unwrap().unbind(), + logger: locals + .get_item("logger") + .unwrap() + .unwrap() + .extract() + .unwrap(), response: Some(py.None()), start: py.None(), end: Some(py.None()), @@ -1104,6 +1046,54 @@ assert observed == ['created', 'release', 'closed'] }); } + #[test] + fn deferred_logging_collects_cycles_through_typed_logger() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!("class Logger: pass\nlogger = Logger()"), + Some(&locals), + Some(&locals), + ) + .unwrap(); + let pending = Py::new( + py, + PendingLogging { + pending: Some(PendingSuccess { + logger: locals + .get_item("logger") + .unwrap() + .unwrap() + .extract() + .unwrap(), + response: None, + start: py.None(), + end: None, + }), + }, + ) + .unwrap(); + locals.set_item("pending", pending).unwrap(); + py.run( + pyo3::ffi::c_str!( + r#" +import gc +import weakref +logger.pending = pending +reference = weakref.ref(logger) +del logger, pending +gc.collect() +assert reference() is None +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + #[test] fn coroutine_collects_cycles_retained_by_bridge_host() { Python::initialize(); diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs new file mode 100644 index 00000000000..23ac0283646 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs @@ -0,0 +1,334 @@ +use pyo3::exceptions::PyBaseException; +use pyo3::gc::{PyTraverseError, PyVisit}; +use pyo3::prelude::*; +use pyo3::types::{PyDict, PyTuple}; + +#[derive(FromPyObject)] +pub(crate) struct PythonLogger(Py); + +impl PythonLogger { + pub(crate) fn object<'py>(&self, py: Python<'py>) -> &Bound<'py, PyAny> { + self.0.bind(py) + } + + pub(crate) fn clone_ref(&self, py: Python<'_>) -> Self { + Self(self.0.clone_ref(py)) + } + + pub(crate) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.0) + } + + pub(super) fn defers_async_logging(&self, py: Python<'_>) -> bool { + self.object(py) + .getattr("_defer_async_logging") + .is_ok_and(|value| value.is_truthy().unwrap_or(false)) + } + + pub(super) fn defer_success( + &self, + py: Python<'_>, + pending: Py, + ) -> PyResult<()> { + self.object(py).setattr("_native_pending_logging", pending) + } + + pub(super) fn sync_success_for_async_call( + &self, + py: Python<'_>, + response: &Option>, + start: &Py, + end: &Option>, + ) -> PyResult<()> { + self.object(py).call_method1( + "handle_sync_success_callbacks_for_async_calls", + (response, start, end), + )?; + Ok(()) + } + + pub(super) fn failure( + &self, + py: Python<'_>, + error: &Py, + start: &Py, + end: &Option>, + asynchronous: bool, + ) -> PyResult>> { + let trace = py + .import("traceback")? + .getattr("format_exception")? + .call1((error,))?; + let trace = pyo3::types::PyString::new(py, "").call_method1("join", (trace,))?; + let value = self.object(py).call_method1( + if asynchronous { + "async_failure_handler" + } else { + "failure_handler" + }, + (error, trace, start, end), + )?; + Ok(asynchronous.then(|| value.unbind())) + } + + pub(super) fn restore_context(&self, py: Python<'_>) -> PyResult<()> { + py.import("litellm.utils")? + .getattr("_restore_correlation_context_if_supported")? + .call1((self.object(py),))?; + Ok(()) + } + + pub(super) fn submit_success( + &self, + py: Python<'_>, + response: &Option>, + start: &Py, + end: &Option>, + ) -> PyResult<()> { + let context = py.import("contextvars")?.call_method0("copy_context")?; + py.import("litellm.litellm_core_utils.litellm_logging")? + .getattr("executor")? + .call_method1( + "submit", + ( + context.getattr("run")?, + self.object(py).getattr("success_handler")?, + response, + start, + end, + ), + )?; + Ok(()) + } + + pub(super) fn enqueue_success( + &self, + py: Python<'_>, + response: &Option>, + start: &Py, + end: &Option>, + ) -> PyResult<()> { + let context = py.import("contextvars")?.call_method0("copy_context")?; + let worker = py + .import("litellm.litellm_core_utils.logging_worker")? + .getattr("GLOBAL_LOGGING_WORKER")? + .getattr("ensure_initialized_and_enqueue")?; + let coroutine = self + .object(py) + .call_method1("async_success_handler", (response, start, end))?; + let enqueue = context.call_method1("run", (worker, &coroutine)); + if enqueue.is_err() + && let Err(error) = coroutine.call_method0("close") + { + error.write_unraisable(py, Some(&coroutine)); + } + enqueue.map(|_| ()) + } +} + +pub(super) struct SetupResult<'py>(Bound<'py, PyAny>); + +impl SetupResult<'_> { + pub(super) fn logger(&self) -> PyResult { + self.0.getattr("logger")?.extract() + } + + pub(super) fn kwargs(&self) -> PyResult> { + Ok(self.0.getattr("kwargs")?.extract()?) + } +} + +pub(super) fn setup<'py>( + py: Python<'py>, + call_type: &str, + args: &Py, + kwargs: &Py, + start: &Py, + asynchronous: bool, +) -> PyResult> { + py.import("litellm.rust_bridge.lifecycle")? + .getattr("setup")? + .call1((call_type, args, kwargs, start, asynchronous)) + .map(SetupResult) +} + +pub(super) fn finalize( + py: Python<'_>, + response: &Option>, + logger: &PythonLogger, + kwargs: &Py, + start: &Py, + end: &Option>, +) -> PyResult<()> { + py.import("litellm.rust_bridge.lifecycle")? + .getattr("finalize")? + .call1((response, logger.object(py), kwargs, start, end))?; + Ok(()) +} + +pub(super) fn is_internal_call(py: Python<'_>) -> PyResult { + py.import("litellm._internal_context")? + .getattr("is_internal_call")? + .call_method0("get")? + .extract() +} + +pub(super) struct DeploymentHooks; + +impl DeploymentHooks { + pub(super) fn before_call( + py: Python<'_>, + kwargs: &Py, + call_type: &str, + ) -> PyResult> { + py.import("litellm.utils")? + .getattr("async_pre_call_deployment_hook")? + .call1((kwargs, call_type)) + .map(Bound::unbind) + } + + pub(super) fn after_success( + py: Python<'_>, + kwargs: &Py, + response: &Option>, + call_type: &str, + ) -> PyResult> { + py.import("litellm.utils")? + .getattr("async_post_call_success_deployment_hook")? + .call1((kwargs, response, call_type)) + .map(Bound::unbind) + } + + pub(super) fn after_failure( + py: Python<'_>, + kwargs: &Py, + error: &Py, + call_type: &str, + ) -> PyResult> { + py.import("litellm.utils")? + .getattr("async_post_call_failure_deployment_hook")? + .call1((kwargs, error, call_type)) + .map(Bound::unbind) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use pyo3::exceptions::PyTypeError; + + #[test] + fn setup_fields_are_checked_in_order_without_eager_logger_method_reads() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!( + r#" +reads = [] +class Logger: + def __getattribute__(self, name): + reads.append(name) + raise AssertionError('logger methods must remain lazy') +logger = Logger() +class Setup: + @property + def logger(self): + reads.append('logger') + return logger + @property + def kwargs(self): + reads.append('kwargs') + return [] +result = Setup() +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + let result = SetupResult(locals.get_item("result").unwrap().unwrap()); + let logger = result.logger().unwrap(); + assert!( + logger + .object(py) + .is(locals.get_item("logger").unwrap().unwrap()) + ); + assert_eq!( + locals + .get_item("reads") + .unwrap() + .unwrap() + .extract::>() + .unwrap(), + ["logger"] + ); + assert!( + result + .kwargs() + .unwrap_err() + .is_instance_of::(py) + ); + assert_eq!( + locals + .get_item("reads") + .unwrap() + .unwrap() + .extract::>() + .unwrap(), + ["logger", "kwargs"] + ); + }); + } + + #[test] + fn logger_resolves_each_callback_at_invocation_and_preserves_arguments() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!( + r#" +calls = [] +response, start, end = object(), object(), object() +class Logger: + @property + def handle_sync_success_callbacks_for_async_calls(self): + generation = len(calls) + def callback(*args): + assert args == (response, start, end) + calls.append(generation) + return callback +logger = Logger() +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + let logger: PythonLogger = locals + .get_item("logger") + .unwrap() + .unwrap() + .extract() + .unwrap(); + let response = Some(locals.get_item("response").unwrap().unwrap().unbind()); + let start = locals.get_item("start").unwrap().unwrap().unbind(); + let end = Some(locals.get_item("end").unwrap().unwrap().unbind()); + for _ in 0..2 { + logger + .sync_success_for_async_call(py, &response, &start, &end) + .unwrap(); + } + assert_eq!( + locals + .get_item("calls") + .unwrap() + .unwrap() + .extract::>() + .unwrap(), + [0, 1] + ); + }); + } +} diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/preparation.rs b/litellm-rust/crates/python-bridge/src/lifecycle/preparation.rs index e2a34c1bb9b..dd07b46ca3b 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/preparation.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/preparation.rs @@ -5,10 +5,10 @@ use pyo3::types::{PyDict, PyList}; pub(super) fn prepare<'py>( py: Python<'py>, kwargs: &Bound<'py, PyDict>, - logger: &Bound<'py, PyAny>, + logger: &super::PythonLogger, ) -> PyResult> { let arguments = kwargs.copy()?; - arguments.set_item("litellm_logging_obj", logger)?; + arguments.set_item("litellm_logging_obj", logger.object(py))?; let litellm = py.import("litellm")?; inherit_credentials(py, &litellm, &arguments)?; py.import("litellm.rust_bridge.lifecycle")? diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 153e4ad4bac..29d48aaf824 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -10,6 +10,7 @@ mod audio_transcription; mod chat_completions; mod messages; mod ocr; +mod ocr_callbacks; mod ocr_document; mod ocr_lifecycle; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs b/litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs new file mode 100644 index 00000000000..b29a50c17e6 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs @@ -0,0 +1,275 @@ +use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError, PyTypeError}; +use pyo3::gc::{PyTraverseError, PyVisit}; +use pyo3::prelude::*; +use pyo3::types::{PyDict, PyString}; +use serde_json::Value; + +use litellm_core::auth::{ResolvedCredential, SecretValue}; +use litellm_core::ocr::LiteLLMOcrResponse; +use litellm_core::ocr::hooks::OcrPreCallRequest; +use litellm_python_interop::to_py_preserving_errors as to_py; + +use crate::lifecycle::PythonLogger; + +pub(super) struct AzureAdTokenProvider(Py); + +impl AzureAdTokenProvider { + pub(super) fn select(provider: Bound<'_, PyAny>) -> Option { + (provider.is_callable() && provider.is_truthy().unwrap_or(false)) + .then(|| Self(provider.unbind())) + } + + pub(super) fn acquire(&self, py: Python<'_>) -> PyResult { + let provider = self.0.bind(py); + if !provider.is_callable() { + return Err(PyTypeError::new_err( + "Azure AD token provider must be callable", + )); + } + let token = (|| { + let token = provider.call0()?; + if !token.is_instance_of::() { + let message = PyString::new(py, "Azure AD token must be a string, got {}") + .call_method1("format", (token.get_type(),))?; + return Err(PyTypeError::new_err(message.unbind())); + } + Ok(token) + })() + .map_err(|error| { + if error.is_instance_of::(py) || !error.is_instance_of::(py) { + return error; + } + match PyString::new(py, "Failed to get Azure AD token: {}") + .call_method1("format", (error.value(py),)) + { + Ok(message) => { + let wrapped = PyRuntimeError::new_err(message.unbind()); + wrapped.set_context(py, Some(error.clone_ref(py))); + wrapped.set_cause(py, Some(error)); + wrapped + } + Err(format_error) => { + format_error.set_context(py, Some(error)); + format_error + } + } + })?; + Ok(ResolvedCredential::AccessToken { + token: SecretValue::new(token.extract::()?), + expires_on: None, + }) + } + + pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.0) + } +} + +impl PythonLogger { + pub(crate) fn update_ocr( + &self, + py: Python<'_>, + kwargs: &Py, + pre_call: &OcrPreCallRequest, + url: &str, + ) -> PyResult<()> { + let redact = py + .import("litellm.rust_bridge.ocr")? + .getattr("redact_logging_params")?; + let update = PyDict::new(py); + update.set_item("kwargs", redact.call1((kwargs,))?.cast_into::()?)?; + update.set_item("model", &pre_call.model)?; + update.set_item( + "optional_params", + redact + .call1((to_py(py, &pre_call.optional_params)?,))? + .cast_into::()?, + )?; + let params = PyDict::new(py); + params.set_item( + "litellm_call_id", + kwargs.bind(py).get_item("litellm_call_id")?, + )?; + params.set_item("api_base", url)?; + update.set_item("litellm_params", params)?; + update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?; + self.object(py) + .call_method("update_from_kwargs", (), Some(&update))?; + Ok(()) + } + + pub(crate) fn pre_ocr( + &self, + py: Python<'_>, + api_key: &Option>, + body: &Bound<'_, PyDict>, + headers: &Bound<'_, PyDict>, + url: &str, + ) -> PyResult<()> { + let additional = PyDict::new(py); + additional.set_item("complete_input_dict", body)?; + additional.set_item("headers", headers)?; + additional.set_item("api_base", url)?; + let kwargs = PyDict::new(py); + kwargs.set_item("input", "OCR document processing")?; + kwargs.set_item("api_key", api_key)?; + kwargs.set_item("additional_args", additional)?; + self.object(py).call_method("pre_call", (), Some(&kwargs))?; + Ok(()) + } + + pub(crate) fn post_ocr( + &self, + py: Python<'_>, + original_response: &Value, + body: &Option>, + headers: &Option>, + ) -> PyResult<()> { + let kwargs = PyDict::new(py); + kwargs.set_item("original_response", to_py(py, original_response)?)?; + let additional = PyDict::new(py); + additional.set_item("complete_input_dict", body)?; + additional.set_item("headers", headers)?; + kwargs.set_item("additional_args", additional)?; + self.object(py) + .call_method("post_call", (), Some(&kwargs))?; + Ok(()) + } +} + +pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult> { + py.import("litellm.rust_bridge.ocr")? + .getattr("_response")? + .call1((to_py(py, response)?,)) + .map(Bound::unbind) +} + +pub(super) fn map_failure( + py: Python<'_>, + error: &Py, + request: &Bound<'_, PyAny>, + provider: &str, +) -> PyResult> { + Ok(py + .import("litellm.rust_bridge.ocr_lifecycle")? + .getattr("map_failure")? + .call1((error, request, provider))? + .extract()?) +} + +pub(super) fn timeout_seconds(py: Python<'_>, timeout: Py) -> PyResult> { + py.import("litellm.rust_bridge.timeouts")? + .getattr("timeout_to_seconds")? + .call1((timeout,))? + .extract() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn token_callback_preserves_exception_identity_and_explicit_chaining() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!( + r#" +class ProviderError(Exception): + def __format__(self, specification): + return 'unavailable' +ordinary = ProviderError('must use __format__') +type_error = TypeError('signature') +abort = KeyboardInterrupt('cancelled') +def provider(error): + def acquire(): + raise error + return acquire +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + for name in ["ordinary", "type_error", "abort"] { + let original = locals.get_item(name).unwrap().unwrap(); + let callback = locals + .get_item("provider") + .unwrap() + .unwrap() + .call1((&original,)) + .unwrap(); + let provider = AzureAdTokenProvider::select(callback).unwrap(); + let error = provider.acquire(py).unwrap_err(); + if name == "ordinary" { + assert!(error.is_instance_of::(py)); + assert!(error.cause(py).unwrap().value(py).is(&original)); + assert!( + error + .value(py) + .getattr("__context__") + .unwrap() + .is(&original) + ); + assert_eq!( + error.value(py).str().unwrap().to_str().unwrap(), + "Failed to get Azure AD token: unavailable" + ); + } else { + assert!(error.value(py).is(&original)); + } + } + }); + } + + #[test] + fn invalid_token_type_formatting_preserves_python_failure_semantics() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!( + r#" +failure = ValueError('formatting failed') +class TokenType(type): + def __format__(cls, specification): + raise failure +class Token(metaclass=TokenType): + pass +def provider(): + return Token() +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + let provider = + AzureAdTokenProvider::select(locals.get_item("provider").unwrap().unwrap()) + .unwrap(); + let error = provider.acquire(py).unwrap_err(); + assert!(error.is_instance_of::(py)); + assert!( + error + .cause(py) + .unwrap() + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + }); + } + + #[test] + fn token_string_extraction_errors_are_not_wrapped_as_callback_failures() { + Python::initialize(); + Python::attach(|py| { + let callback = py + .eval(pyo3::ffi::c_str!("lambda: '\\ud800'"), None, None) + .unwrap(); + let provider = AzureAdTokenProvider::select(callback).unwrap(); + let error = provider.acquire(py).unwrap_err(); + assert!(error.is_instance_of::(py)); + }); + } +} 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 e1c1214fbbc..0b03f626010 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs @@ -4,7 +4,7 @@ use std::sync::Arc; use pyo3::prelude::*; use pyo3::types::{PyDict, PyTuple}; -use litellm_core::auth::{ResolvedCredential, SecretValue}; +use litellm_core::auth::ResolvedCredential; use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest, OcrPreCallRequest}; use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_param_names, decode_request}; use litellm_core::ocr::{ @@ -14,6 +14,7 @@ use litellm_python_interop::{ from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, }; +use super::ocr_callbacks::{self, AzureAdTokenProvider}; use crate::errors::{RustBridgeDeclined, ocr_error_to_pyerr}; use crate::lifecycle::{PythonCallState, PythonRoute, missing_state, now, run_call}; @@ -23,7 +24,7 @@ struct PythonOcrHost { pre_call: Option, document: Option>, api_key: Option>, - azure_ad_token_provider: Option>, + azure_ad_token_provider: Option, provider: String, retained_fields: Option>, body: Option>, @@ -34,7 +35,7 @@ struct AdmittedOcrCall { request: litellm_core::ocr::LiteLLMOcrRequest, document: Py, api_key: Py, - azure_ad_token_provider: Option>, + azure_ad_token_provider: Option, provider: String, } @@ -70,15 +71,7 @@ impl PythonOcrHost { .azure_ad_token_provider .as_ref() .ok_or_else(missing_state)?; - let token: String = py - .import("litellm.rust_bridge.ocr_lifecycle")? - .getattr("call_azure_ad_token_provider")? - .call1((provider,))? - .extract()?; - Ok(ResolvedCredential::AccessToken { - token: SecretValue::new(token), - expires_on: None, - }) + provider.acquire(py) } fn python_pre_call( @@ -108,35 +101,9 @@ impl PythonOcrHost { } self.body = Some(body.clone().unbind()); self.headers = Some(headers.clone().unbind()); - let logger = self.state.logger(py)?; - let redact = py - .import("litellm.rust_bridge.ocr")? - .getattr("redact_logging_params")?; - let update = PyDict::new(py); - update.set_item("kwargs", redact.call1((&self.state.kwargs,))?)?; - update.set_item("model", &pre_call.model)?; - update.set_item( - "optional_params", - redact.call1((to_py(py, &pre_call.optional_params)?,))?, - )?; - let params = PyDict::new(py); - params.set_item( - "litellm_call_id", - self.state.kwargs.bind(py).get_item("litellm_call_id")?, - )?; - params.set_item("api_base", &request.url)?; - update.set_item("litellm_params", params)?; - update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?; - logger.call_method("update_from_kwargs", (), Some(&update))?; - let additional = PyDict::new(py); - additional.set_item("complete_input_dict", &body)?; - additional.set_item("headers", &headers)?; - additional.set_item("api_base", &request.url)?; - let kwargs = PyDict::new(py); - kwargs.set_item("input", "OCR document processing")?; - kwargs.set_item("api_key", &self.api_key)?; - kwargs.set_item("additional_args", additional)?; - logger.call_method("pre_call", (), Some(&kwargs))?; + let logger = self.state.logger()?; + logger.update_ocr(py, &self.state.kwargs, pre_call, &request.url)?; + logger.pre_ocr(py, &self.api_key, &body, &headers, &request.url)?; let headers = headers .iter() .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) @@ -151,14 +118,9 @@ impl PythonOcrHost { py: Python<'_>, request: OcrPostCallRequest, ) -> PyResult { - let logger = self.state.logger(py)?; - let kwargs = PyDict::new(py); - kwargs.set_item("original_response", to_py(py, &request.original_response)?)?; - let additional = PyDict::new(py); - additional.set_item("complete_input_dict", &self.body)?; - additional.set_item("headers", &self.headers)?; - kwargs.set_item("additional_args", additional)?; - logger.call_method("post_call", (), Some(&kwargs))?; + self.state + .logger()? + .post_ocr(py, &request.original_response, &self.body, &self.headers)?; Ok(request) } } @@ -203,12 +165,7 @@ impl PythonRoute for PythonOcrHost { } OcrHostOperation::ConstructResponse(response) => { self.state.end = Some(now(py)?); - self.state.response = Some( - py.import("litellm.rust_bridge.ocr")? - .getattr("_response")? - .call1((to_py(py, response.as_ref())?,))? - .unbind(), - ); + self.state.response = Some(ocr_callbacks::response(py, response.as_ref())?); OcrHostResult::Lifecycle(Ok(())) } OcrHostOperation::MapFailure(error) => { @@ -220,11 +177,9 @@ impl PythonRoute for PythonOcrHost { } let error = self.state.error.as_ref().ok_or_else(missing_state)?; let request = self.request.as_ref().ok_or_else(missing_state)?.bind(py); - let mapped = py - .import("litellm.rust_bridge.ocr_lifecycle")? - .getattr("map_failure")? - .call1((error, request, &self.provider))?; - self.state.retain_error(py, PyErr::from_value(mapped)); + let mapped = ocr_callbacks::map_failure(py, error, request, &self.provider)?; + self.state + .retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any())); OcrHostResult::Lifecycle(Ok(())) } OcrHostOperation::Lifecycle(_) @@ -247,7 +202,9 @@ impl PythonRoute for PythonOcrHost { visit.call(&self.request)?; visit.call(&self.document)?; visit.call(&self.api_key)?; - visit.call(&self.azure_ad_token_provider)?; + if let Some(provider) = &self.azure_ad_token_provider { + provider.traverse(visit)?; + } visit.call(&self.retained_fields)?; visit.call(&self.body)?; visit.call(&self.headers) @@ -278,8 +235,7 @@ fn project_request( let input_sources = extract_input_sources(request_kwargs, &consumed)?; let azure_ad_token_provider = request_kwargs .get_item("azure_ad_token_provider")? - .filter(|provider| provider.is_callable() && provider.is_truthy().unwrap_or(false)) - .map(Bound::unbind); + .and_then(AzureAdTokenProvider::select); let wire = OcrWireRequest { model, document: wire_document, @@ -294,12 +250,7 @@ fn project_request( input_sources, timeout_seconds: argument("timeout")? .extract::>>()? - .map(|value| { - py.import("litellm.rust_bridge.timeouts")? - .getattr("timeout_to_seconds")? - .call1((value,))? - .extract() - }) + .map(|value| ocr_callbacks::timeout_seconds(py, value)) .transpose()? .flatten(), }; diff --git a/litellm/rust_bridge/ocr_lifecycle.py b/litellm/rust_bridge/ocr_lifecycle.py index 138932c85c2..5ca584e1c11 100644 --- a/litellm/rust_bridge/ocr_lifecycle.py +++ b/litellm/rust_bridge/ocr_lifecycle.py @@ -50,20 +50,6 @@ def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]: return request.kwargs -def call_azure_ad_token_provider(provider: object) -> str: - if not callable(provider): - raise TypeError("Azure AD token provider must be callable") - try: - token: Final = provider() - if not isinstance(token, str): - raise TypeError(f"Azure AD token must be a string, got {type(token)}") - return token - except TypeError: - raise - except Exception as error: - raise RuntimeError(f"Failed to get Azure AD token: {error}") from error - - 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