From 92e1a2ded70467ff61e7d06f8a36783965b626c8 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Wed, 16 Sep 2026 19:07:23 -0700 Subject: [PATCH] Revert "refactor(python-bridge): classify callbacks natively instead of function_setup" This reverts commit f8190bbe80ccefd6b24585dbe885be4e19ca6bd2. --- .../crates/core/src/call_lifecycle/mod.rs | 1 - .../core/src/call_lifecycle/registration.rs | 374 -------- .../crates/python-bridge/src/diagnostics.rs | 24 +- .../python-bridge/src/lifecycle/bindings.rs | 251 +++++- .../python-bridge/src/lifecycle/compat.rs | 168 ---- .../python-bridge/src/lifecycle/dispatch.rs | 771 ----------------- .../crates/python-bridge/src/lifecycle/mod.rs | 531 ++++++++---- .../python-bridge/src/lifecycle/setup.rs | 337 -------- .../python-bridge/src/routes/ocr/callbacks.rs | 61 +- .../python-bridge/src/routes/ocr/host.rs | 42 +- .../python-bridge/src/routes/ocr/mod.rs | 17 +- .../python-bridge/src/routes/ocr/project.rs | 19 +- litellm/rust_bridge/_native.pyi | 10 - litellm/rust_bridge/leaves.py | 815 ------------------ litellm/rust_bridge/lifecycle.py | 69 +- litellm/rust_bridge/ocr.py | 30 + litellm/rust_bridge/setup.py | 262 ------ .../rust_bridge/test_ocr_lifecycle.py | 20 +- tests/test_litellm/rust_bridge/test_setup.py | 169 ---- tests/test_litellm_rust/ocr/test_callbacks.py | 248 +----- tests/test_litellm_rust/ocr/test_lifecycle.py | 20 +- 21 files changed, 763 insertions(+), 3476 deletions(-) delete mode 100644 litellm-rust/crates/core/src/call_lifecycle/registration.rs delete mode 100644 litellm-rust/crates/python-bridge/src/lifecycle/compat.rs delete mode 100644 litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs delete mode 100644 litellm-rust/crates/python-bridge/src/lifecycle/setup.rs delete mode 100644 litellm/rust_bridge/leaves.py delete mode 100644 litellm/rust_bridge/setup.py delete mode 100644 tests/test_litellm/rust_bridge/test_setup.py diff --git a/litellm-rust/crates/core/src/call_lifecycle/mod.rs b/litellm-rust/crates/core/src/call_lifecycle/mod.rs index f5625ac6537..992b29ad63e 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/mod.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/mod.rs @@ -3,7 +3,6 @@ use std::time::{Instant, SystemTime, UNIX_EPOCH}; pub mod callbacks; pub mod host; -pub mod registration; pub mod types; pub use callbacks::{ diff --git a/litellm-rust/crates/core/src/call_lifecycle/registration.rs b/litellm-rust/crates/core/src/call_lifecycle/registration.rs deleted file mode 100644 index eac0526c9fa..00000000000 --- a/litellm-rust/crates/core/src/call_lifecycle/registration.rs +++ /dev/null @@ -1,374 +0,0 @@ -use std::collections::HashSet; - -use super::callbacks::CallbackId; - -#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] -pub enum Registry { - Input, - AsyncInput, - Success, - AsyncSuccess, - Failure, - AsyncFailure, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum Registration { - Named { known: bool, async_only: bool }, - Object { asynchronous: bool }, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct Entry { - pub id: CallbackId, - pub registration: Registration, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct Candidate { - pub resolved: Option, - pub duplicate_type: bool, -} - -#[derive(Clone, Debug, PartialEq, Eq, Default)] -pub struct RegistrationFacts { - pub candidates: Vec, - pub input: Vec, - pub success: Vec, - pub failure: Vec, - pub async_success: Vec, - pub async_failure: Vec, - pub bootstrap_pending: bool, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum NamedEvent { - Success, - Failure, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum RegistryMutation { - Append(Registry, CallbackId), - Remove(Registry, CallbackId), - ExpandNamed(NamedEvent, CallbackId), - Bootstrap, -} - -struct Planner { - input: Vec, - success: Vec, - failure: Vec, - async_success: HashSet, - async_failure: HashSet, - mutations: Vec, -} - -impl Planner { - fn contains(&self, registry: Registry, id: CallbackId) -> bool { - match registry { - Registry::Input => self.input.iter().any(|entry| entry.id == id), - Registry::Success => self.success.iter().any(|entry| entry.id == id), - Registry::Failure => self.failure.iter().any(|entry| entry.id == id), - Registry::AsyncSuccess => self.async_success.contains(&id), - Registry::AsyncFailure => self.async_failure.contains(&id), - Registry::AsyncInput => false, - } - } - - fn append(&mut self, registry: Registry, entry: Entry) { - if self.contains(registry, entry.id) { - return; - } - match registry { - Registry::Input => self.input.push(entry), - Registry::Success => self.success.push(entry), - Registry::Failure => self.failure.push(entry), - Registry::AsyncSuccess => { - self.async_success.insert(entry.id); - } - Registry::AsyncFailure => { - self.async_failure.insert(entry.id); - } - Registry::AsyncInput => {} - } - self.mutations - .push(RegistryMutation::Append(registry, entry.id)); - } - - fn record(&mut self, mutation: RegistryMutation) { - self.mutations.push(mutation); - } -} - -fn is_asynchronous(registration: Registration) -> bool { - matches!(registration, Registration::Object { asynchronous: true }) -} - -pub fn plan_registration(facts: &RegistrationFacts) -> Vec { - let mut planner = Planner { - input: facts.input.clone(), - success: facts.success.clone(), - failure: facts.failure.clone(), - async_success: facts.async_success.iter().copied().collect(), - async_failure: facts.async_failure.iter().copied().collect(), - mutations: Vec::new(), - }; - - for candidate in &facts.candidates { - let Some(entry) = candidate.resolved else { - continue; - }; - if candidate.duplicate_type { - continue; - } - planner.append(Registry::Input, entry); - if !is_asynchronous(entry.registration) { - planner.append(Registry::Success, entry); - planner.append(Registry::Failure, entry); - } - planner.append(Registry::AsyncSuccess, entry); - planner.append(Registry::AsyncFailure, entry); - } - - if facts.bootstrap_pending - && !(planner.input.is_empty() && planner.success.is_empty() && planner.failure.is_empty()) - { - planner.record(RegistryMutation::Bootstrap); - } - - let input = planner.input.clone(); - for entry in input - .iter() - .filter(|entry| is_asynchronous(entry.registration)) - { - planner.record(RegistryMutation::Append(Registry::AsyncInput, entry.id)); - planner.record(RegistryMutation::Remove(Registry::Input, entry.id)); - } - - let success = planner.success.clone(); - for entry in &success { - match entry.registration { - Registration::Object { asynchronous: true } - | Registration::Named { - async_only: true, .. - } => { - planner.append(Registry::AsyncSuccess, *entry); - planner.record(RegistryMutation::Remove(Registry::Success, entry.id)); - } - Registration::Named { known: true, .. } => { - planner.record(RegistryMutation::ExpandNamed(NamedEvent::Success, entry.id)); - } - _ => {} - } - } - - let failure = planner.failure.clone(); - for entry in &failure { - match entry.registration { - Registration::Object { asynchronous: true } => { - planner.append(Registry::AsyncFailure, *entry); - planner.record(RegistryMutation::Remove(Registry::Failure, entry.id)); - } - Registration::Named { known: true, .. } => { - planner.record(RegistryMutation::ExpandNamed(NamedEvent::Failure, entry.id)); - } - _ => {} - } - } - - planner.mutations -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum DynamicSuccessSlot { - Sync, - Async, -} - -pub fn classify_dynamic_success(entry: Entry, named_async: bool) -> DynamicSuccessSlot { - match entry.registration { - Registration::Object { asynchronous: true } => DynamicSuccessSlot::Async, - Registration::Named { .. } if named_async => DynamicSuccessSlot::Async, - _ => DynamicSuccessSlot::Sync, - } -} - -#[cfg(test)] -mod tests { - use super::*; - - fn object(id: u64, asynchronous: bool) -> Entry { - Entry { - id: CallbackId(id), - registration: Registration::Object { asynchronous }, - } - } - - fn named(id: u64, known: bool, async_only: bool) -> Entry { - Entry { - id: CallbackId(id), - registration: Registration::Named { known, async_only }, - } - } - - fn candidate(entry: Entry) -> Candidate { - Candidate { - resolved: Some(entry), - duplicate_type: false, - } - } - - #[test] - fn sync_callback_in_callbacks_registers_in_every_list_once() { - let facts = RegistrationFacts { - candidates: vec![candidate(object(1, false)), candidate(object(1, false))], - ..RegistrationFacts::default() - }; - assert_eq!( - plan_registration(&facts), - [ - RegistryMutation::Append(Registry::Input, CallbackId(1)), - RegistryMutation::Append(Registry::Success, CallbackId(1)), - RegistryMutation::Append(Registry::Failure, CallbackId(1)), - RegistryMutation::Append(Registry::AsyncSuccess, CallbackId(1)), - RegistryMutation::Append(Registry::AsyncFailure, CallbackId(1)), - ] - ); - } - - #[test] - fn async_callable_in_callbacks_skips_sync_lists_and_moves_out_of_input() { - let facts = RegistrationFacts { - candidates: vec![candidate(object(2, true))], - ..RegistrationFacts::default() - }; - assert_eq!( - plan_registration(&facts), - [ - RegistryMutation::Append(Registry::Input, CallbackId(2)), - RegistryMutation::Append(Registry::AsyncSuccess, CallbackId(2)), - RegistryMutation::Append(Registry::AsyncFailure, CallbackId(2)), - RegistryMutation::Append(Registry::AsyncInput, CallbackId(2)), - RegistryMutation::Remove(Registry::Input, CallbackId(2)), - ] - ); - } - - #[test] - fn unresolved_and_duplicate_type_named_candidates_are_skipped() { - let facts = RegistrationFacts { - candidates: vec![ - Candidate { - resolved: None, - duplicate_type: false, - }, - Candidate { - resolved: Some(object(3, false)), - duplicate_type: true, - }, - ], - ..RegistrationFacts::default() - }; - assert!(plan_registration(&facts).is_empty()); - } - - #[test] - fn already_registered_callbacks_are_not_appended_again() { - let facts = RegistrationFacts { - candidates: vec![candidate(object(1, false))], - input: vec![object(1, false)], - success: vec![object(1, false)], - failure: vec![object(1, false)], - async_success: vec![CallbackId(1)], - async_failure: vec![CallbackId(1)], - ..RegistrationFacts::default() - }; - assert!(plan_registration(&facts).is_empty()); - } - - #[test] - fn bootstrap_runs_once_when_any_public_list_is_populated() { - let empty = RegistrationFacts { - bootstrap_pending: true, - ..RegistrationFacts::default() - }; - assert!(plan_registration(&empty).is_empty()); - let populated = RegistrationFacts { - bootstrap_pending: true, - candidates: vec![candidate(object(1, false))], - ..RegistrationFacts::default() - }; - assert!(plan_registration(&populated).contains(&RegistryMutation::Bootstrap)); - let already = RegistrationFacts { - bootstrap_pending: false, - success: vec![object(1, false)], - ..RegistrationFacts::default() - }; - assert!(!plan_registration(&already).contains(&RegistryMutation::Bootstrap)); - } - - #[test] - fn success_safety_net_moves_async_and_async_only_names_and_expands_known_names() { - let facts = RegistrationFacts { - success: vec![ - object(1, true), - named(2, false, true), - named(3, true, false), - named(4, false, false), - object(5, false), - ], - ..RegistrationFacts::default() - }; - assert_eq!( - plan_registration(&facts), - [ - RegistryMutation::Append(Registry::AsyncSuccess, CallbackId(1)), - RegistryMutation::Remove(Registry::Success, CallbackId(1)), - RegistryMutation::Append(Registry::AsyncSuccess, CallbackId(2)), - RegistryMutation::Remove(Registry::Success, CallbackId(2)), - RegistryMutation::ExpandNamed(NamedEvent::Success, CallbackId(3)), - ] - ); - } - - #[test] - fn failure_safety_net_ignores_async_only_names() { - let facts = RegistrationFacts { - failure: vec![ - object(1, true), - named(2, false, true), - named(3, true, false), - ], - async_failure: vec![CallbackId(1)], - ..RegistrationFacts::default() - }; - assert_eq!( - plan_registration(&facts), - [ - RegistryMutation::Remove(Registry::Failure, CallbackId(1)), - RegistryMutation::ExpandNamed(NamedEvent::Failure, CallbackId(3)), - ] - ); - } - - #[test] - fn dynamic_success_split_follows_async_callables_and_selected_names() { - assert_eq!( - classify_dynamic_success(object(1, true), false), - DynamicSuccessSlot::Async - ); - assert_eq!( - classify_dynamic_success(object(1, false), false), - DynamicSuccessSlot::Sync - ); - assert_eq!( - classify_dynamic_success(named(2, true, false), true), - DynamicSuccessSlot::Async - ); - assert_eq!( - classify_dynamic_success(named(2, true, true), false), - DynamicSuccessSlot::Sync - ); - } -} diff --git a/litellm-rust/crates/python-bridge/src/diagnostics.rs b/litellm-rust/crates/python-bridge/src/diagnostics.rs index 0c2b4c9e731..cc153a89b8f 100644 --- a/litellm-rust/crates/python-bridge/src/diagnostics.rs +++ b/litellm-rust/crates/python-bridge/src/diagnostics.rs @@ -1,6 +1,6 @@ use litellm_python_interop::release_count; use pyo3::prelude::*; -use pyo3::types::{PyDict, PyTuple}; +use pyo3::types::PyDict; #[pyfunction] fn gil_stats(py: Python<'_>) -> PyResult> { @@ -9,27 +9,6 @@ fn gil_stats(py: Python<'_>) -> PyResult> { Ok(stats.into_any().unbind()) } -#[pyfunction] -fn _debug_setup( - py: Python<'_>, - call_type: String, - args: Bound<'_, PyTuple>, - kwargs: Bound<'_, PyDict>, - start: Py, - asynchronous: bool, -) -> PyResult<(Py, Py)> { - let leaked: &'static str = Box::leak(call_type.into_boxed_str()); - let result = crate::lifecycle::debug_setup( - py, - leaked, - &args.unbind(), - &kwargs.unbind(), - &start, - asynchronous, - )?; - Ok(result) -} - #[cfg(feature = "panic-test")] #[pyfunction] fn _panic_for_test() { @@ -38,7 +17,6 @@ fn _panic_for_test() { pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { module.add_function(wrap_pyfunction!(gil_stats, module)?)?; - module.add_function(wrap_pyfunction!(_debug_setup, module)?)?; #[cfg(feature = "panic-test")] module.add_function(wrap_pyfunction!(_panic_for_test, module)?)?; Ok(()) diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs index f78d4f87c62..23ac0283646 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs @@ -1,7 +1,7 @@ use pyo3::exceptions::PyBaseException; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; -use pyo3::types::PyDict; +use pyo3::types::{PyDict, PyTuple}; #[derive(FromPyObject)] pub(crate) struct PythonLogger(Py); @@ -28,10 +28,128 @@ impl PythonLogger { pub(super) fn defer_success( &self, py: Python<'_>, - pending: Py, + 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( @@ -42,11 +160,9 @@ pub(super) fn finalize( start: &Py, end: &Option>, ) -> PyResult<()> { - let model = kwargs.bind(py).get_item("model")?; - let model = model.filter(|value| value.is_instance_of::()); - py.import("litellm.litellm_core_utils.llm_response_utils.response_metadata")? - .getattr("update_response_metadata")? - .call1((response, logger.object(py), model, kwargs, start, end))?; + py.import("litellm.rust_bridge.lifecycle")? + .getattr("finalize")? + .call1((response, logger.object(py), kwargs, start, end))?; Ok(()) } @@ -95,3 +211,124 @@ impl DeploymentHooks { .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/compat.rs b/litellm-rust/crates/python-bridge/src/lifecycle/compat.rs deleted file mode 100644 index ce8d508a552..00000000000 --- a/litellm-rust/crates/python-bridge/src/lifecycle/compat.rs +++ /dev/null @@ -1,168 +0,0 @@ -use litellm_core::call_lifecycle::{ - CallbackFamily, Delivery, ReleaseGate, SuccessFacts, plan_success, -}; -use pyo3::prelude::*; - -use super::bindings::PythonLogger; -use super::{PythonCallState, dispatch, missing_state}; - -pub(super) fn dispatch_success( - py: Python<'_>, - state: &PythonCallState, - logger: &PythonLogger, -) -> PyResult<()> { - let facts = SuccessFacts { - asynchronous: state.asynchronous, - internal: state.internal, - fallbacks: !state - .kwargs - .bind(py) - .get_item("fallbacks")? - .is_none_or(|value| value.is_none()), - deferred: logger.defers_async_logging(py), - sync_target_kinds: sync_kinds(py, logger)?, - }; - let leaves = dispatch::leaves(py)?; - let object = logger.object(py); - for selected in plan_success(&facts) { - match (selected.family, selected.delivery, selected.gate) { - (CallbackFamily::SyncSuccess, Delivery::Worker, _) => { - let handler = object.getattr("success_handler")?; - let bound = py.import("functools")?.getattr("partial")?.call1(( - handler, - &state.response, - &state.start, - &state.end, - ))?; - leaves.getattr("submit_worker")?.call1((bound,))?; - } - (CallbackFamily::AsyncSuccess, Delivery::Background, ReleaseGate::Immediate) => { - let coroutine = object.call_method1( - "async_success_handler", - (&state.response, &state.start, &state.end), - )?; - let enqueue = leaves.getattr("enqueue_background")?.call1((&coroutine,)); - if enqueue.is_err() - && let Err(error) = coroutine.call_method0("close") - { - error.write_unraisable(py, Some(&coroutine)); - } - enqueue?; - } - (CallbackFamily::AsyncSuccess, Delivery::Background, ReleaseGate::Deferred) => { - let pending = Py::new( - py, - SuppliedDeferred { - logger: Some(logger.clone_ref(py)), - response: state.response.as_ref().map(|value| value.clone_ref(py)), - start: state.start.clone_ref(py), - end: state.end.as_ref().map(|value| value.clone_ref(py)), - }, - )?; - object.setattr("_native_pending_logging", pending)?; - } - _ => return Err(missing_state()), - } - } - Ok(()) -} - -fn sync_kinds( - py: Python<'_>, - logger: &PythonLogger, -) -> PyResult> { - let (targets, ids) = dispatch::family_targets(py, logger, CallbackFamily::SyncSuccess)?; - Ok(targets.kinds(&ids)) -} - -pub(super) fn dispatch_failure( - py: Python<'_>, - state: &PythonCallState, - family: CallbackFamily, -) -> PyResult>> { - let logger = state.logger()?.object(py); - let error = state.error.as_ref().ok_or_else(missing_state)?; - let trace = py - .import("traceback")? - .getattr("format_exception")? - .call1((error,))?; - let trace = pyo3::types::PyString::new(py, "").call_method1("join", (trace,))?; - match family { - CallbackFamily::SyncFailure => { - logger.call_method1("failure_handler", (error, trace, &state.start, &state.end))?; - Ok(None) - } - CallbackFamily::AsyncFailure => Ok(Some( - logger - .call_method1( - "async_failure_handler", - (error, trace, &state.start, &state.end), - )? - .unbind(), - )), - _ => Err(missing_state()), - } -} - -#[pyclass] -struct SuppliedDeferred { - logger: Option, - response: Option>, - start: Py, - end: Option>, -} - -#[pymethods] -impl SuppliedDeferred { - fn __call__(slf: &Bound<'_, Self>, py: Python<'_>) -> PyResult<()> { - let logger = slf.borrow_mut().logger.take(); - let Some(logger) = logger else { - return Ok(()); - }; - let (response, start, end) = { - let this = slf.borrow(); - ( - this.response.as_ref().map(|value| value.clone_ref(py)), - this.start.clone_ref(py), - this.end.as_ref().map(|value| value.clone_ref(py)), - ) - }; - let object = logger.object(py); - let coroutine = object.call_method1("async_success_handler", (response, start, end))?; - let enqueue = dispatch::leaves(py)? - .getattr("enqueue_background")? - .call1((&coroutine,)); - if enqueue.is_err() - && let Err(error) = coroutine.call_method0("close") - { - error.write_unraisable(py, Some(&coroutine)); - } - match enqueue { - Err(error) if error.is_instance_of::(py) => { - error.write_unraisable(py, Some(object)); - Ok(()) - } - result => result.map(|_| ()), - } - } - - fn __traverse__(&self, visit: pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> { - if let Some(logger) = &self.logger { - logger.traverse(&visit)?; - } - visit.call(&self.response)?; - visit.call(&self.start)?; - visit.call(&self.end) - } - - fn close(slf: &Bound<'_, Self>) { - let mut this = slf.borrow_mut(); - this.logger = None; - this.response = None; - this.end = None; - } - - fn __clear__(slf: &Bound<'_, Self>) { - Self::close(slf); - } -} diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs b/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs deleted file mode 100644 index 4215b8e6c2d..00000000000 --- a/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs +++ /dev/null @@ -1,771 +0,0 @@ -use litellm_core::call_lifecycle::{ - CallbackFamily, CallbackId, CallbackInvocation, CallbackKind, CallbackMethod, Delivery, - DispatchCursor, DispatchFacts, DispatchStep, InvocationOutcome, LoggedMarker, -}; -use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; -use pyo3::gc::{PyTraverseError, PyVisit}; -use pyo3::prelude::*; -use pyo3::types::PyString; - -use super::bindings::PythonLogger; - -const LEAVES: &str = "litellm.rust_bridge.leaves"; - -pub(super) fn leaves(py: Python<'_>) -> PyResult> { - py.import(LEAVES) -} - -#[derive(Clone, Copy)] -pub(super) enum Outcome { - Success, - Failure, -} - -pub(super) struct Targets { - objects: Vec>, - kinds: Vec, -} - -impl Targets { - pub(super) fn read( - py: Python<'_>, - lists: &[Bound<'_, PyAny>], - ) -> PyResult<(Self, Vec>)> { - let custom_logger = py - .import("litellm.integrations.custom_logger")? - .getattr("CustomLogger")?; - let known = py - .import("litellm")? - .getattr("_known_custom_logger_compatible_callbacks")?; - let mut targets = Self { - objects: Vec::new(), - kinds: Vec::new(), - }; - let mut ids = Vec::with_capacity(lists.len()); - for list in lists { - let mut family = Vec::new(); - for object in list.try_iter()? { - let object = object?; - family.push(targets.intern(py, &object, &custom_logger, &known)?); - } - ids.push(family); - } - Ok((targets, ids)) - } - - fn intern( - &mut self, - py: Python<'_>, - object: &Bound<'_, PyAny>, - custom_logger: &Bound<'_, PyAny>, - known: &Bound<'_, PyAny>, - ) -> PyResult { - for (index, existing) in self.objects.iter().enumerate() { - let existing = existing.bind(py); - if existing.is(object) || existing.eq(object)? { - return Ok(CallbackId(index as u64)); - } - } - let kind = if object.is_instance(custom_logger)? { - CallbackKind::CustomLogger - } else if let Ok(name) = object.cast::() { - CallbackKind::Named { - known: known.contains(name)?, - } - } else if object.is_callable() { - CallbackKind::Callable { - internal: internal_callable(object)?, - } - } else { - CallbackKind::Opaque - }; - self.objects.push(object.clone().unbind()); - self.kinds.push(kind); - Ok(CallbackId((self.objects.len() - 1) as u64)) - } - - pub(super) fn kinds(&self, ids: &[CallbackId]) -> Vec { - ids.iter().map(|id| self.kinds[id.0 as usize]).collect() - } - - fn object<'py>(&self, py: Python<'py>, id: CallbackId) -> &Bound<'py, PyAny> { - self.objects[id.0 as usize].bind(py) - } - - fn kind(&self, id: CallbackId) -> CallbackKind { - self.kinds[id.0 as usize] - } - - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - for object in &self.objects { - visit.call(object)?; - } - Ok(()) - } -} - -fn internal_callable(object: &Bound<'_, PyAny>) -> PyResult { - let name = if let Ok(name) = object.getattr("__name__") { - name.extract::()? - } else if let Ok(func) = object.getattr("__func__") { - func.getattr("__name__")?.extract::()? - } else { - object.get_type().name()?.to_string() - }; - Ok([ - "_PROXY", - "_service_logger.ServiceLogging", - "sync_deployment_callback_on_success", - ] - .iter() - .any(|prefix| name.contains(prefix))) -} - -pub(super) struct Job { - pub logger: PythonLogger, - pub targets: Targets, - pub ids: Vec, - pub family: CallbackFamily, - pub response: Option>, - pub error: Option>, - pub start: Py, - pub end: Py, - pub stream: bool, -} - -impl Job { - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - self.logger.traverse(visit)?; - self.targets.traverse(visit)?; - visit.call(&self.response)?; - visit.call(&self.error)?; - visit.call(&self.start)?; - visit.call(&self.end) - } - - fn family_name(&self) -> &'static str { - match self.family { - CallbackFamily::SyncSuccess => "sync_success", - CallbackFamily::AsyncSuccess => "async_success", - CallbackFamily::SyncFailure => "sync_failure", - CallbackFamily::AsyncFailure => "async_failure", - _ => "request", - } - } - - fn outcome(&self) -> Outcome { - match self.family { - CallbackFamily::SyncSuccess | CallbackFamily::AsyncSuccess => Outcome::Success, - _ => Outcome::Failure, - } - } -} - -struct Eligibility<'a, 'py> { - py: Python<'py>, - job: &'a Job, - leaves: &'a Bound<'py, PyModule>, -} - -impl DispatchFacts for Eligibility<'_, '_> { - fn eligible(&mut self, target: CallbackId, method: CallbackMethod) -> bool { - let object = self.job.targets.object(self.py, target); - let kind = self.job.targets.kind(target); - let result = match method { - CallbackMethod::LoggingHook | CallbackMethod::AsyncLoggingHook => { - if kind != CallbackKind::CustomLogger { - return false; - } - self.leaves - .getattr("should_run_guardrail_hook") - .and_then(|f| f.call1((self.job.logger.object(self.py), object))) - .and_then(|v| v.extract::()) - } - CallbackMethod::LogSuccessEvent - | CallbackMethod::AsyncLogSuccessEvent - | CallbackMethod::LogFailureEvent - | CallbackMethod::AsyncLogFailureEvent => { - if kind == CallbackKind::Opaque { - return false; - } - self.leaves - .getattr("should_run_callback") - .and_then(|f| { - f.call1(( - self.job.logger.object(self.py), - object, - self.job.family_name(), - )) - }) - .and_then(|v| v.extract::()) - } - _ => Ok(kind != CallbackKind::Opaque), - }; - match result { - Ok(value) => value, - Err(error) => { - error.write_unraisable(self.py, Some(object)); - false - } - } - } -} - -pub(super) enum Step { - Await(Py), - Done, -} - -pub(super) struct Runner { - job: Job, - cursor: DispatchCursor, - result: Option>, - formatted: Option>, - pending: Option, -} - -impl Runner { - pub(super) fn start(py: Python<'_>, job: Job) -> PyResult { - let leaves = leaves(py)?; - let already = match job.family.marker() { - Some(marker) => leaves - .getattr("already_logged")? - .call1((job.logger.object(py), marker.key()))? - .extract::()?, - None => false, - }; - let cursor = DispatchCursor::start(job.family, job.ids.clone(), already, job.stream); - let result = job.response.as_ref().map(|value| value.clone_ref(py)); - Ok(Self { - job, - cursor, - result, - formatted: None, - pending: None, - }) - } - - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - self.job.traverse(visit)?; - visit.call(&self.result)?; - visit.call(&self.formatted) - } - - pub(super) fn logger<'py>(&self, py: Python<'py>) -> &Bound<'py, PyAny> { - self.job.logger.object(py) - } - - pub(super) fn resume( - &mut self, - py: Python<'_>, - awaited: Option>>, - ) -> PyResult { - let leaves = leaves(py)?; - if let Some(awaited) = awaited { - let invocation = self.pending.take().ok_or_else(super::missing_state)?; - let result = awaited.map(|value| { - (invocation.method == CallbackMethod::AsyncLoggingHook).then_some(value) - }); - self.accept(py, &leaves, invocation.target, result)?; - } - loop { - let step = self.cursor.next(&mut Eligibility { - py, - job: &self.job, - leaves: &leaves, - }); - match step { - DispatchStep::PrepareLogging => self.prepare(py, &leaves)?, - DispatchStep::MarkLogged(marker) => self.mark(py, &leaves, marker)?, - DispatchStep::Invoke(invocation) => { - let outcome = self.invoke(py, &leaves, invocation); - match outcome { - Ok(Some(awaitable)) => { - self.pending = Some(invocation); - return Ok(Step::Await(awaitable)); - } - Ok(None) => self.accept(py, &leaves, invocation.target, Ok(None))?, - Err(error) => self.accept(py, &leaves, invocation.target, Err(error))?, - } - } - DispatchStep::Complete { .. } => { - leaves - .getattr("restore_correlation_context")? - .call1((self.job.logger.object(py),))?; - return Ok(Step::Done); - } - } - } - } - - fn prepare(&mut self, py: Python<'_>, leaves: &Bound<'_, PyModule>) -> PyResult<()> { - let logger = self.job.logger.object(py); - let result = match self.job.outcome() { - Outcome::Success => leaves - .getattr("prepare_success_logging")? - .call1((logger, &self.result, &self.job.start, &self.job.end)) - .map(|redacted| self.result = Some(redacted.unbind())), - Outcome::Failure => leaves - .getattr("prepare_failure_logging")? - .call1((logger, &self.job.error, &self.job.start, &self.job.end)) - .map(|formatted| self.formatted = Some(formatted.unbind())), - }; - match result { - Err(error) if error.is_instance_of::(py) => { - error.write_unraisable(py, Some(logger)); - Ok(()) - } - result => result, - } - } - - fn mark( - &self, - py: Python<'_>, - leaves: &Bound<'_, PyModule>, - marker: LoggedMarker, - ) -> PyResult<()> { - leaves - .getattr("mark_logged")? - .call1((self.job.logger.object(py), marker.key()))?; - Ok(()) - } - - fn invoke( - &mut self, - py: Python<'_>, - leaves: &Bound<'_, PyModule>, - invocation: CallbackInvocation, - ) -> PyResult>> { - let logger = self.job.logger.object(py); - let target = self.job.targets.object(py, invocation.target); - let kind = self.job.targets.kind(invocation.target); - let awaits = matches!(invocation.delivery, Delivery::Await | Delivery::Background); - let value = match (invocation.method, kind) { - (CallbackMethod::LoggingHook, CallbackKind::CustomLogger) => { - let replaced = - leaves - .getattr("logging_hook")? - .call1((logger, target, &self.result))?; - self.result = Some(replaced.unbind()); - return Ok(None); - } - (CallbackMethod::AsyncLoggingHook, CallbackKind::CustomLogger) => leaves - .getattr("async_logging_hook")? - .call1((logger, target, &self.result))?, - (CallbackMethod::LogSuccessEvent, CallbackKind::CustomLogger) => leaves - .getattr("log_success_event")? - .call1((logger, target, &self.result, &self.job.start, &self.job.end))?, - (CallbackMethod::AsyncLogSuccessEvent, CallbackKind::CustomLogger) => leaves - .getattr("async_log_success_event")? - .call1((logger, target, &self.result, &self.job.start, &self.job.end))?, - (CallbackMethod::LogFailureEvent, CallbackKind::CustomLogger) => leaves - .getattr("log_failure_event")? - .call1((logger, target, &self.job.start, &self.job.end))?, - (CallbackMethod::AsyncLogFailureEvent, CallbackKind::CustomLogger) => leaves - .getattr("async_log_failure_event")? - .call1((logger, target, &self.job.start, &self.job.end))?, - ( - CallbackMethod::LogSuccessEvent - | CallbackMethod::AsyncLogSuccessEvent - | CallbackMethod::LogFailureEvent - | CallbackMethod::AsyncLogFailureEvent, - CallbackKind::Callable { .. }, - ) => leaves.getattr("dispatch_callable")?.call1(( - logger, - target, - self.job.family_name(), - &self.result, - &self.job.start, - &self.job.end, - ))?, - ( - CallbackMethod::LogSuccessEvent | CallbackMethod::AsyncLogSuccessEvent, - CallbackKind::Named { .. }, - ) => leaves.getattr("dispatch_named_success")?.call1(( - logger, - target, - &self.result, - &self.job.start, - &self.job.end, - ))?, - ( - CallbackMethod::LogFailureEvent | CallbackMethod::AsyncLogFailureEvent, - CallbackKind::Named { .. }, - ) => leaves.getattr("dispatch_named_failure")?.call1(( - logger, - target, - &self.job.error, - &self.formatted, - &self.job.start, - &self.job.end, - ))?, - _ => return Ok(None), - }; - if awaits && !value.is_none() { - return Ok(Some(value.unbind())); - } - Ok(None) - } - - fn accept( - &mut self, - py: Python<'_>, - leaves: &Bound<'_, PyModule>, - target: CallbackId, - result: PyResult>>, - ) -> PyResult<()> { - match result { - Ok(Some(replaced)) => { - self.result = Some(replaced); - self.cursor.accept(InvocationOutcome::Completed); - } - Ok(None) => self.cursor.accept(InvocationOutcome::Completed), - Err(error) if error.is_instance_of::(py) => { - let object = self.job.targets.object(py, target); - if let Err(report) = leaves.getattr("report_target_failure").and_then(|f| { - f.call1(( - self.job.logger.object(py), - object, - self.job.family_name(), - error.value(py), - )) - }) { - report.write_unraisable(py, Some(object)); - } - self.cursor.accept(InvocationOutcome::Failed); - } - Err(error) => return Err(error), - } - Ok(()) - } -} - -#[pyclass] -pub(super) struct WorkerJob { - runner: Option, -} - -impl WorkerJob { - pub(super) fn new(runner: Runner) -> Self { - Self { - runner: Some(runner), - } - } -} - -#[pymethods] -impl WorkerJob { - fn __call__(slf: &Bound<'_, Self>, py: Python<'_>) -> PyResult<()> { - let Some(mut runner) = slf.borrow_mut().runner.take() else { - return Ok(()); - }; - match runner.resume(py, None) { - Ok(Step::Done) => Ok(()), - Ok(Step::Await(_)) => Err(PyRuntimeError::new_err( - "worker dispatch selected an awaiting delivery", - )), - Err(error) if error.is_instance_of::(py) => { - error.write_unraisable(py, Some(runner.job.logger.object(py))); - Ok(()) - } - Err(error) => Err(error), - } - } - - fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { - match &self.runner { - Some(runner) => runner.traverse(&visit), - None => Ok(()), - } - } - - fn __clear__(&mut self) { - self.runner = None; - } -} - -pub(super) struct AwaitingBody { - runner: Option, -} - -impl AwaitingBody { - pub(super) fn new(runner: Runner) -> Self { - Self { - runner: Some(runner), - } - } -} - -impl super::handle::ExecutionBody for AwaitingBody { - fn resume( - &mut self, - result: Option>>, - ) -> PyResult { - Python::attach(|py| { - let runner = self.runner.as_mut().ok_or_else(super::missing_state)?; - match runner.resume(py, result) { - Ok(Step::Await(awaitable)) => Ok(super::handle::ExecutionStep::Await(awaitable)), - Ok(Step::Done) => { - self.runner = None; - Ok(super::handle::ExecutionStep::Return(py.None())) - } - Err(error) if error.is_instance_of::(py) => { - error.write_unraisable(py, Some(runner.job.logger.object(py))); - self.runner = None; - Ok(super::handle::ExecutionStep::Return(py.None())) - } - Err(error) => Err(error), - } - }) - } - - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - match &self.runner { - Some(runner) => runner.traverse(visit), - None => Ok(()), - } - } -} - -pub(super) fn coroutine(py: Python<'_>, runner: Runner) -> PyResult> { - let execution = Py::new(py, super::handle::Execution::new(AwaitingBody::new(runner)))?; - py.import("litellm.rust_bridge.lifecycle")? - .getattr("drive")? - .call1((execution,)) - .map(Bound::unbind) -} - -pub(super) struct RequestJob<'a> { - pub logger: &'a PythonLogger, - pub family: CallbackFamily, -} - -pub(super) fn dispatch_request(py: Python<'_>, job: RequestJob<'_>) -> PyResult<()> { - let leaves = leaves(py)?; - let (targets, ids) = family_targets(py, job.logger, job.family)?; - let mut cursor = DispatchCursor::start(job.family, ids, false, false); - let logger = job.logger.object(py); - let event = match job.family { - CallbackFamily::RequestPreCall => "pre_api_call", - _ => "post_api_call", - }; - let mut facts = RequestEligibility; - loop { - match cursor.next(&mut facts) { - DispatchStep::Invoke(invocation) => { - let target = targets.object(py, invocation.target); - let result = match (invocation.method, targets.kind(invocation.target)) { - (CallbackMethod::LogPreApiCall, CallbackKind::CustomLogger) => { - leaves.getattr("log_pre_api_call")?.call1((logger, target)) - } - (CallbackMethod::LogPostApiCall, CallbackKind::CustomLogger) => { - leaves.getattr("log_post_api_call")?.call1((logger, target)) - } - (_, CallbackKind::Named { .. }) => leaves - .getattr("dispatch_named_request")? - .call1((logger, target, event)), - (CallbackMethod::LogPreApiCall, CallbackKind::Callable { .. }) => leaves - .getattr("dispatch_callable_request")? - .call1((logger, target)), - _ => Ok(py.None().into_bound(py)), - }; - match result { - Ok(_) => cursor.accept(InvocationOutcome::Completed), - Err(error) if error.is_instance_of::(py) => { - if let Err(report) = leaves - .getattr("report_target_failure") - .and_then(|f| f.call1((logger, target, "request", error.value(py)))) - { - report.write_unraisable(py, Some(target)); - } - cursor.accept(InvocationOutcome::Failed); - } - Err(error) => return Err(error), - } - } - DispatchStep::Complete { .. } => return Ok(()), - DispatchStep::PrepareLogging | DispatchStep::MarkLogged(_) => {} - } - } -} - -struct RequestEligibility; - -impl DispatchFacts for RequestEligibility { - fn eligible(&mut self, _: CallbackId, _: CallbackMethod) -> bool { - true - } -} - -pub(super) fn read_lists<'py>( - py: Python<'py>, - logger: &PythonLogger, - family: CallbackFamily, -) -> PyResult<(Bound<'py, PyAny>, Option>)> { - let litellm = py.import("litellm")?; - let logger = logger.object(py); - let (global, dynamic) = match family { - CallbackFamily::RequestPreCall | CallbackFamily::RequestPostCall => { - ("input_callback", "dynamic_input_callbacks") - } - CallbackFamily::SyncSuccess => ("success_callback", "dynamic_success_callbacks"), - CallbackFamily::AsyncSuccess => { - ("_async_success_callback", "dynamic_async_success_callbacks") - } - CallbackFamily::SyncFailure => ("failure_callback", "dynamic_failure_callbacks"), - CallbackFamily::AsyncFailure => { - ("_async_failure_callback", "dynamic_async_failure_callbacks") - } - CallbackFamily::DeploymentPreCall - | CallbackFamily::DeploymentPostCall - | CallbackFamily::DeploymentFailure => ("callbacks", ""), - }; - let global = litellm.getattr(global)?; - let dynamic = if dynamic.is_empty() { - None - } else { - let value = logger.getattr(dynamic)?; - (!value.is_none()).then_some(value) - }; - Ok((global, dynamic)) -} - -pub(super) fn family_targets( - py: Python<'_>, - logger: &PythonLogger, - family: CallbackFamily, -) -> PyResult<(Targets, Vec)> { - let (global, dynamic) = read_lists(py, logger, family)?; - let lists: Vec> = std::iter::once(global).chain(dynamic.clone()).collect(); - let (targets, ids) = Targets::read(py, &lists)?; - let global_ids = ids.first().cloned().unwrap_or_default(); - let dynamic_ids = ids.get(1).map(Vec::as_slice); - let ordered = family.targets( - &global_ids, - dynamic.is_some().then_some(dynamic_ids.unwrap_or(&[])), - ); - Ok((targets, ordered)) -} - -#[cfg(test)] -mod tests { - use super::*; - use pyo3::types::{PyDict, PyList}; - - fn fixture(py: Python<'_>) -> Bound<'_, PyDict> { - let locals = PyDict::new(py); - py.run( - pyo3::ffi::c_str!( - r#" -import sys, types -for name in ("litellm", "litellm.integrations", "litellm.integrations.custom_logger"): - sys.modules.setdefault(name, types.ModuleType(name)) -class CustomLogger: pass -sys.modules["litellm.integrations.custom_logger"].CustomLogger = CustomLogger -sys.modules["litellm"]._known_custom_logger_compatible_callbacks = [] -class Logger: pass -logger = Logger() -target = CustomLogger() -"# - ), - Some(&locals), - Some(&locals), - ) - .unwrap(); - locals - } - - fn runner(py: Python<'_>, locals: &Bound<'_, PyDict>, family: CallbackFamily) -> Runner { - let logger = locals.get_item("logger").unwrap().unwrap(); - let target = locals.get_item("target").unwrap().unwrap(); - let list = PyList::new(py, [&target]).unwrap().into_any(); - let (targets, ids) = Targets::read(py, &[list]).unwrap(); - let ids: Vec = ids.into_iter().flatten().collect(); - Runner { - cursor: DispatchCursor::start(family, ids.clone(), false, false), - job: Job { - logger: logger.extract().unwrap(), - targets, - ids, - family, - response: Some(target.clone().unbind()), - error: None, - start: py.None(), - end: py.None(), - stream: false, - }, - result: Some(target.unbind()), - formatted: None, - pending: None, - } - } - - fn assert_collectable(py: Python<'_>, locals: &Bound<'_, PyDict>, handle: Py) { - locals.set_item("handle", handle).unwrap(); - py.run( - pyo3::ffi::c_str!( - r#" -import gc -import weakref -target.handle = handle -reference = weakref.ref(target) -del logger, target, handle -gc.collect() -assert reference() is None, "cycle through the retained target was not collected" -"# - ), - Some(locals), - Some(locals), - ) - .unwrap(); - } - - #[test] - fn worker_job_collects_cycles_through_logger_targets_and_response() { - Python::initialize(); - Python::attach(|py| { - let locals = fixture(py); - let runner = runner(py, &locals, CallbackFamily::SyncSuccess); - let handle = Py::new(py, WorkerJob::new(runner)).unwrap().into_any(); - assert_collectable(py, &locals, handle); - }); - } - - #[test] - fn deferred_success_collects_cycles_and_close_is_idempotent() { - Python::initialize(); - Python::attach(|py| { - let locals = fixture(py); - let runner = runner(py, &locals, CallbackFamily::AsyncSuccess); - let deferred = Py::new( - py, - super::super::DeferredSuccess { - runner: Some(runner), - }, - ) - .unwrap(); - assert_collectable(py, &locals, deferred.into_any()); - }); - } - - #[test] - fn deferred_success_releases_at_most_once_and_close_prevents_release() { - Python::initialize(); - Python::attach(|py| { - let locals = fixture(py); - let runner = runner(py, &locals, CallbackFamily::AsyncSuccess); - let deferred = Py::new( - py, - super::super::DeferredSuccess { - runner: Some(runner), - }, - ) - .unwrap(); - deferred.call_method0(py, "close").unwrap(); - deferred.call_method0(py, "close").unwrap(); - deferred.call0(py).unwrap(); - assert!(deferred.borrow(py).runner.is_none()); - }); - } -} diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs index ca247138083..3a9e305c9f6 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs @@ -1,30 +1,25 @@ use std::sync::Arc; use std::task::Poll; +use futures_util::future::{AbortHandle, Abortable}; +#[cfg(test)] +use litellm_core::call_lifecycle::host::HostCallFuture; +use litellm_core::call_lifecycle::host::{ + HostCall as NativeCall, HostCallStep as NativeCallStep, HostFailure, HostPhase, HostStep, +}; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; use pyo3::types::{PyDict, PyTuple}; - -use futures_util::future::{AbortHandle, Abortable}; use tokio::sync::Mutex; -use litellm_core::call_lifecycle::host::{ - HostCall as NativeCall, HostCallStep as NativeCallStep, HostFailure, HostPhase, HostStep, -}; -use litellm_core::call_lifecycle::{ - CallbackFamily, Delivery, ReleaseGate, SuccessFacts, plan_failure, plan_success, -}; +use crate::execution::{poll_async_value, run_async_value, run_sync_value}; mod arguments; mod bindings; -mod compat; -mod dispatch; mod handle; mod preparation; -mod setup; -use crate::execution::{poll_async_value, run_async_value, run_sync_value}; pub(crate) use arguments::{BoundArguments, Signature}; use bindings::DeploymentHooks; pub(crate) use bindings::PythonLogger; @@ -98,18 +93,6 @@ pub(crate) fn run_call( } } -pub(crate) fn debug_setup( - py: Python<'_>, - call_type: &'static str, - args: &Py, - kwargs: &Py, - start: &Py, - asynchronous: bool, -) -> PyResult<(Py, Py)> { - let result = setup::setup(py, call_type, args, kwargs, start, asynchronous)?; - Ok((result.logger.object(py).clone().unbind(), result.kwargs)) -} - pub(crate) fn missing_state() -> PyErr { pyo3::exceptions::PyRuntimeError::new_err("missing native call state") } @@ -179,11 +162,7 @@ impl PythonLifecycle { HostFailure::Cancelled(native) }; let state = self.route.state_mut(); - state.retain_first_error( - py, - error, - cancelled && phase != Some(HostPhase::DeploymentFailure), - ); + state.retain_first_error(py, error, cancelled && phase != Some(HostPhase::DeploymentFailure)); let _ = state.finish(py); failure } @@ -304,7 +283,6 @@ pub(crate) struct PythonCallState { pub error: Option>, pub asynchronous: bool, pub internal: bool, - pub supplied: bool, pub call_type: &'static str, } @@ -394,7 +372,6 @@ impl PythonCallState { error: None, asynchronous, internal: false, - supplied: false, call_type, }) } @@ -408,7 +385,7 @@ impl PythonCallState { pub fn setup(&mut self, py: Python<'_>) -> PyResult<()> { self.start = now(py)?; self.internal = bindings::is_internal_call(py)?; - let result = setup::setup( + let result = bindings::setup( py, self.call_type, &self.args, @@ -416,9 +393,8 @@ impl PythonCallState { &self.start, self.asynchronous, )?; - self.logger = Some(result.logger); - self.kwargs = result.kwargs; - self.supplied = result.supplied; + self.logger = Some(result.logger()?); + self.kwargs = result.kwargs()?; Ok(()) } @@ -438,16 +414,6 @@ impl PythonCallState { ) } - pub fn dispatch_request(&self, py: Python<'_>, family: CallbackFamily) -> PyResult<()> { - dispatch::dispatch_request( - py, - dispatch::RequestJob { - logger: self.logger()?, - family, - }, - ) - } - pub fn dispatch_success(&self, py: Python<'_>) -> PyResult<()> { match self.try_dispatch_success(py) { Err(error) if error.is_instance_of::(py) => { @@ -458,71 +424,40 @@ impl PythonCallState { } } - fn job(&self, py: Python<'_>, family: CallbackFamily) -> PyResult { - let logger = self.logger()?; - let (targets, ids) = dispatch::family_targets(py, logger, family)?; - Ok(dispatch::Job { - logger: logger.clone_ref(py), - targets, - ids, - family, - response: self.response.as_ref().map(|value| value.clone_ref(py)), - error: self.error.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)) - .unwrap_or_else(|| py.None()), - stream: false, - }) - } - fn try_dispatch_success(&self, py: Python<'_>) -> PyResult<()> { let logger = self.logger()?; - if self.supplied { - return compat::dispatch_success(py, self, logger); - } - let (sync_targets, sync_ids) = - dispatch::family_targets(py, logger, CallbackFamily::SyncSuccess)?; - let facts = SuccessFacts { - asynchronous: self.asynchronous, - internal: self.internal, - fallbacks: !self - .kwargs - .bind(py) - .get_item("fallbacks")? - .is_none_or(|value| value.is_none()), - deferred: logger.defers_async_logging(py), - sync_target_kinds: sync_targets.kinds(&sync_ids), + let pending = || PendingSuccess { + 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)), }; - for selected in plan_success(&facts) { - let runner = dispatch::Runner::start(py, self.job(py, selected.family)?)?; - match (selected.delivery, selected.gate) { - (Delivery::Worker, _) => { - let job = Py::new(py, dispatch::WorkerJob::new(runner))?; - dispatch::leaves(py)? - .getattr("submit_worker")? - .call1((job,))?; - } - (Delivery::Background, ReleaseGate::Immediate) => { - DeferredSuccess::release(py, runner)?; - } - (Delivery::Background, ReleaseGate::Deferred) => { + if !self.asynchronous { + pending().sync(py) + } else { + if !self.internal + && self + .kwargs + .bind(py) + .get_item("fallbacks")? + .is_none_or(|value| value.is_none()) + { + if logger.defers_async_logging(py) { logger.defer_success( py, Py::new( py, - DeferredSuccess { - runner: Some(runner), + PendingLogging { + pending: Some(pending()), }, )?, )?; + } else { + pending().asynchronous(py)?; } - (Delivery::Inline | Delivery::Await, _) => return Err(missing_state()), } + logger.sync_success_for_async_call(py, &self.response, &self.start, &self.end) } - Ok(()) } pub fn dispatch_failure( @@ -530,39 +465,19 @@ impl PythonCallState { py: Python<'_>, asynchronous: bool, ) -> PyResult>> { - if self.logger.is_none() || self.error.is_none() { + if self.logger.is_none() || (self.asynchronous && self.internal) { return Ok(None); } - let phase = if asynchronous { - HostPhase::AsyncFailure - } else { - HostPhase::Failure - }; - let Some(family) = plan_failure(phase, self.asynchronous, self.internal) else { + let Some(error) = &self.error else { return Ok(None); }; - if self.supplied { - return compat::dispatch_failure(py, self, family); - } - let mut runner = dispatch::Runner::start(py, self.job(py, family)?)?; - match family.delivery() { - Delivery::Inline => match runner.resume(py, None)? { - dispatch::Step::Done => Ok(None), - dispatch::Step::Await(_) => Err(missing_state()), - }, - Delivery::Await => Ok(Some(dispatch::coroutine(py, runner)?)), - Delivery::Worker | Delivery::Background => Err(missing_state()), - } + 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) = dispatch::leaves(py).and_then(|leaves| { - leaves - .getattr("restore_correlation_context")? - .call1((logger.object(py),)) - .map(|_| ()) - }) + && let Err(error) = logger.restore_context(py) { error.write_unraisable(py, None); } @@ -598,53 +513,58 @@ impl PythonCallState { } } -#[pyclass] -struct DeferredSuccess { - runner: Option, +struct PendingSuccess { + logger: PythonLogger, + response: Option>, + start: Py, + end: Option>, } -impl DeferredSuccess { - fn release(py: Python<'_>, runner: dispatch::Runner) -> PyResult<()> { - let coroutine = dispatch::coroutine(py, runner)?; - let enqueue = dispatch::leaves(py)? - .getattr("enqueue_background")? - .call1((&coroutine,)); - if enqueue.is_err() - && let Err(error) = coroutine.call_method0(py, "close") - { - error.write_unraisable(py, Some(coroutine.bind(py))); - } - enqueue.map(|_| ()) +impl PendingSuccess { + fn sync(&self, py: Python<'_>) -> PyResult<()> { + self.logger + .submit_success(py, &self.response, &self.start, &self.end) } + + fn asynchronous(&self, py: Python<'_>) -> PyResult<()> { + self.logger + .enqueue_success(py, &self.response, &self.start, &self.end) + } +} + +#[pyclass] +struct PendingLogging { + pending: Option, } #[pymethods] -impl DeferredSuccess { +impl PendingLogging { fn __call__(slf: &Bound<'_, Self>, py: Python<'_>) -> PyResult<()> { - let runner = slf.borrow_mut().runner.take(); - let Some(runner) = runner else { - return Ok(()); - }; - let logger = runner.logger(py).clone().unbind(); - match Self::release(py, runner) { - Err(error) if error.is_instance_of::(py) => { - error.write_unraisable(py, Some(logger.bind(py))); - Ok(()) + let pending = slf.borrow_mut().pending.take(); + if let Some(pending) = pending { + match pending.asynchronous(py) { + Err(error) if error.is_instance_of::(py) => { + error.write_unraisable(py, Some(pending.logger.object(py))); + } + result => return result, } - result => result, } + Ok(()) } fn __traverse__(&self, visit: pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> { - match &self.runner { - Some(runner) => runner.traverse(&visit), - None => Ok(()), + if let Some(pending) = &self.pending { + pending.logger.traverse(&visit)?; + visit.call(&pending.response)?; + visit.call(&pending.start)?; + visit.call(&pending.end)?; } + Ok(()) } fn close(slf: &Bound<'_, Self>) { - let runner = slf.borrow_mut().runner.take(); - drop(runner); + let pending = slf.borrow_mut().pending.take(); + drop(pending); } fn __clear__(slf: &Bound<'_, Self>) { @@ -655,39 +575,14 @@ impl DeferredSuccess { #[cfg(test)] mod tests { use super::*; - use pyo3::types::PyDict; use std::sync::Mutex; - use litellm_core::call_lifecycle::host::HostCallFuture; - static PYTHON_GLOBALS: Mutex<()> = Mutex::new(()); - fn load_lifecycle_module(py: Python<'_>) -> Bound<'_, PyModule> { - py.run( - pyo3::ffi::c_str!( - r#" -import sys -import types -for name in ("litellm", "litellm.rust_bridge"): - sys.modules.setdefault(name, types.ModuleType(name)) -"# - ), - None, - None, - ) - .unwrap(); - let source = std::ffi::CString::new(include_str!( - "../../../../../litellm/rust_bridge/lifecycle.py" - )) - .unwrap(); - PyModule::from_code( - py, - &source, - pyo3::ffi::c_str!("lifecycle.py"), - pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"), - ) - .unwrap() + fn install_logging_worker(py: Python<'_>, worker: &Bound<'_, PyAny>) -> PyResult<()> { + py.import("litellm.litellm_core_utils.logging_worker")? + .setattr("GLOBAL_LOGGING_WORKER", worker) } struct RetainingHost { @@ -859,7 +754,17 @@ for name in ("litellm", "litellm.rust_bridge"): .unwrap_or_else(|error| error.into_inner()); Python::initialize(); Python::attach(|py| { - load_lifecycle_module(py); + let source = std::ffi::CString::new(include_str!( + "../../../../../litellm/rust_bridge/lifecycle.py" + )) + .unwrap(); + PyModule::from_code( + py, + &source, + pyo3::ffi::c_str!("lifecycle.py"), + pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"), + ) + .unwrap(); let route = SyntheticRoute( PythonCallState::new( py, @@ -895,7 +800,17 @@ for name in ("litellm", "litellm.rust_bridge"): Python::initialize(); Python::attach(|py| { py.import("asyncio").unwrap(); - let module = load_lifecycle_module(py); + let source = std::ffi::CString::new(include_str!( + "../../../../../litellm/rust_bridge/lifecycle.py" + )) + .unwrap(); + let module = PyModule::from_code( + py, + &source, + pyo3::ffi::c_str!("lifecycle.py"), + pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"), + ) + .unwrap(); let locals = PyDict::new(py); locals .set_item("drive", module.getattr("drive").unwrap()) @@ -1003,11 +918,64 @@ assert reference() is None error: None, asynchronous, internal: false, - supplied: false, call_type: "test", } } + #[test] + fn success_dispatch_reports_ordinary_failures_without_replacing_response() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!( + r#" +import sys + +response = object() +failure = ValueError('terminal diagnostic') +diagnostics = [] +old_hook = sys.unraisablehook +sys.unraisablehook = lambda event: diagnostics.append(event.exc_value) + +class Logger: + def handle_sync_success_callbacks_for_async_calls(self, *args): + raise failure + +logger = Logger() +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + let response = locals.get_item("response").unwrap().unwrap().unbind(); + let mut lifecycle_state = state( + py, + locals.get_item("logger").unwrap().unwrap().unbind(), + response.clone_ref(py), + true, + ); + lifecycle_state.internal = true; + lifecycle_state.dispatch_success(py).unwrap(); + assert!(lifecycle_state.response.as_ref().unwrap().is(&response)); + py.run( + pyo3::ffi::c_str!( + r#" +assert diagnostics == [failure] +sys.unraisablehook = old_hook +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + #[test] fn retained_failure_preserves_exception_identity() { Python::initialize(); @@ -1023,6 +991,201 @@ assert reference() is None }); } + #[test] + fn deferred_logging_uses_call_context_and_allows_reentry_once() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!( + r#" +import sys +import types +from contextvars import ContextVar + +litellm = types.ModuleType('litellm') +core_utils = types.ModuleType('litellm.litellm_core_utils') +logging_worker = types.ModuleType('litellm.litellm_core_utils.logging_worker') +litellm.litellm_core_utils = core_utils +core_utils.logging_worker = logging_worker +sys.modules['litellm'] = litellm +sys.modules['litellm.litellm_core_utils'] = core_utils +sys.modules['litellm.litellm_core_utils.logging_worker'] = logging_worker + +marker = ContextVar('marker', default='unset') +observed = [] + +class Coroutine: + def close(self): + observed.append('closed') + +class Worker: + def ensure_initialized_and_enqueue(self, coroutine): + observed.append(marker.get()) + pending() + coroutine.close() + +class Logger: + def async_success_handler(self, *args): + observed.append('created') + return Coroutine() + +worker = Worker() +logger = Logger() +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + install_logging_worker(py, &locals.get_item("worker").unwrap().unwrap()).unwrap(); + let pending = Py::new( + py, + PendingLogging { + pending: Some(PendingSuccess { + logger: locals + .get_item("logger") + .unwrap() + .unwrap() + .extract() + .unwrap(), + response: Some(py.None()), + start: py.None(), + end: Some(py.None()), + }), + }, + ) + .unwrap(); + locals.set_item("pending", &pending).unwrap(); + py.run( + pyo3::ffi::c_str!( + r#" +marker.set('call') +pending() +pending() +assert observed == ['created', 'call', 'closed'] +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + #[test] + fn deferred_logging_close_is_reentry_safe_and_invalidates_aliases() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!( + r#" +observed = [] + +class Retained: + def __del__(self): + observed.append('finalized') + alias() + +class Logger: + def async_success_handler(self, *args): + observed.append('enqueued') + +logger = Logger() +retained = Retained() +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + let pending = Py::new( + py, + PendingLogging { + pending: Some(PendingSuccess { + logger: locals + .get_item("logger") + .unwrap() + .unwrap() + .extract() + .unwrap(), + response: Some(locals.get_item("retained").unwrap().unwrap().unbind()), + start: py.None(), + end: None, + }), + }, + ) + .unwrap(); + locals.set_item("pending", &pending).unwrap(); + locals.set_item("alias", &pending).unwrap(); + locals.del_item("retained").unwrap(); + py.run( + pyo3::ffi::c_str!( + r#" +pending.close() +alias() +assert observed == ['finalized'] +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + #[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/setup.rs b/litellm-rust/crates/python-bridge/src/lifecycle/setup.rs deleted file mode 100644 index c7370be628f..00000000000 --- a/litellm-rust/crates/python-bridge/src/lifecycle/setup.rs +++ /dev/null @@ -1,337 +0,0 @@ -use litellm_core::call_lifecycle::CallbackId; -use litellm_core::call_lifecycle::registration::{ - Candidate, DynamicSuccessSlot, Entry, NamedEvent, Registration, RegistrationFacts, Registry, - RegistryMutation, classify_dynamic_success, plan_registration, -}; -use pyo3::prelude::*; -use pyo3::types::{PyDict, PyList, PyString, PyTuple}; - -use super::bindings::PythonLogger; - -const SETUP_MODULE: &str = "litellm.rust_bridge.setup"; - -struct Targets<'py> { - objects: Vec>, -} - -impl<'py> Targets<'py> { - fn new() -> Self { - Self { - objects: Vec::new(), - } - } - - fn intern(&mut self, object: &Bound<'py, PyAny>) -> CallbackId { - if let Some(index) = self - .objects - .iter() - .position(|existing| existing.is(object) || existing.eq(object).unwrap_or(false)) - { - return CallbackId(index as u64); - } - self.objects.push(object.clone()); - CallbackId((self.objects.len() - 1) as u64) - } - - fn get(&self, id: CallbackId) -> PyResult<&Bound<'py, PyAny>> { - self.objects - .get(id.0 as usize) - .ok_or_else(super::missing_state) - } -} - -fn registration<'py>( - setup: &Bound<'py, PyModule>, - object: &Bound<'py, PyAny>, -) -> PyResult { - if let Ok(name) = object.cast::() { - let name = name.to_str()?; - let known = setup.getattr("is_known_name")?.call1((name,))?.extract()?; - return Ok(Registration::Named { - known, - async_only: matches!(name, "dynamodb" | "openmeter"), - }); - } - let asynchronous = setup - .getattr("is_async_callable")? - .call1((object,))? - .extract()?; - Ok(Registration::Object { asynchronous }) -} - -fn entries<'py>( - setup: &Bound<'py, PyModule>, - targets: &mut Targets<'py>, - list: &Bound<'py, PyAny>, -) -> PyResult> { - list.try_iter()? - .map(|object| { - let object = object?; - Ok(Entry { - id: targets.intern(&object), - registration: registration(setup, &object)?, - }) - }) - .collect() -} - -fn registry_name(registry: Registry) -> &'static str { - match registry { - Registry::Input => "input", - Registry::AsyncInput => "async_input", - Registry::Success => "success", - Registry::AsyncSuccess => "async_success", - Registry::Failure => "failure", - Registry::AsyncFailure => "async_failure", - } -} - -fn read_registry<'py>(setup: &Bound<'py, PyModule>, name: &str) -> PyResult> { - setup.getattr("registry")?.call1((name,)) -} - -fn read_candidates<'py>( - setup: &Bound<'py, PyModule>, - targets: &mut Targets<'py>, - dynamic: Option>, -) -> PyResult> { - let mut candidates = Vec::new(); - let global = read_registry(setup, "callbacks")?; - let sources = std::iter::once(global).chain(dynamic); - for source in sources { - for object in source.try_iter()? { - let object = object?; - let candidate = if object.is_instance_of::() { - let resolved = setup - .getattr("resolve_named_integration")? - .call1((&object,))?; - if resolved.is_none() { - Candidate { - resolved: None, - duplicate_type: false, - } - } else { - let duplicate_type = setup - .getattr("async_success_registry_has_type")? - .call1((&resolved,))? - .extract()?; - Candidate { - resolved: Some(Entry { - id: targets.intern(&resolved), - registration: registration(setup, &resolved)?, - }), - duplicate_type, - } - } - } else { - Candidate { - resolved: Some(Entry { - id: targets.intern(&object), - registration: registration(setup, &object)?, - }), - duplicate_type: false, - } - }; - candidates.push(candidate); - } - } - Ok(candidates) -} - -fn apply<'py>( - setup: &Bound<'py, PyModule>, - targets: &Targets<'py>, - mutations: &[RegistryMutation], - function_id: Option<&Bound<'py, PyAny>>, -) -> PyResult<()> { - for mutation in mutations { - match mutation { - RegistryMutation::Append(registry, id) => { - setup - .getattr("append_registry")? - .call1((registry_name(*registry), targets.get(*id)?))?; - } - RegistryMutation::Remove(registry, id) => { - setup - .getattr("remove_registry")? - .call1((registry_name(*registry), targets.get(*id)?))?; - } - RegistryMutation::ExpandNamed(event, id) => { - let event = match event { - NamedEvent::Success => "success", - NamedEvent::Failure => "failure", - }; - setup - .getattr("expand_named")? - .call1((targets.get(*id)?, event))?; - } - RegistryMutation::Bootstrap => { - setup.getattr("bootstrap")?.call1((function_id,))?; - } - } - } - Ok(()) -} - -struct DynamicLists<'py> { - success: Option>, - async_success: Option>, - failure: Option>, -} - -fn split_dynamic<'py>( - py: Python<'py>, - setup: &Bound<'py, PyModule>, - targets: &mut Targets<'py>, - kwargs: &Bound<'py, PyDict>, -) -> PyResult> { - let success = match kwargs.get_item("success_callback")? { - Some(value) if value.is_instance_of::() => { - let list = value.cast_into::()?; - let sync = PyList::empty(py); - let asynchronous = PyList::empty(py); - for object in list.iter() { - let entry = Entry { - id: targets.intern(&object), - registration: registration(setup, &object)?, - }; - let named_async = object - .cast::() - .ok() - .and_then(|name| { - name.to_str() - .ok() - .map(|name| matches!(name, "dynamodb" | "s3")) - }) - .unwrap_or(false); - match classify_dynamic_success(entry, named_async) { - DynamicSuccessSlot::Sync => sync.append(&object)?, - DynamicSuccessSlot::Async => asynchronous.append(&object)?, - } - } - kwargs.del_item("success_callback")?; - Some((sync, (!asynchronous.is_empty()).then_some(asynchronous))) - } - _ => None, - }; - let failure = match kwargs.get_item("failure_callback")? { - Some(value) if value.is_instance_of::() => { - kwargs.del_item("failure_callback")?; - Some(value.cast_into::()?) - } - _ => None, - }; - let (success, async_success) = match success { - Some((sync, asynchronous)) => (Some(sync), asynchronous), - None => (None, None), - }; - Ok(DynamicLists { - success, - async_success, - failure, - }) -} - -pub(super) struct Setup { - pub logger: PythonLogger, - pub kwargs: Py, - pub supplied: bool, -} - -pub(super) fn setup( - py: Python<'_>, - call_type: &str, - args: &Py, - kwargs: &Py, - start: &Py, - asynchronous: bool, -) -> PyResult { - let setup = py.import(SETUP_MODULE)?; - let kwargs = kwargs.bind(py).copy()?; - if !kwargs.contains("litellm_call_id")? { - let call_id = py.import("uuid")?.call_method0("uuid4")?.str()?; - kwargs.set_item("litellm_call_id", call_id)?; - } - if let Some(supplied) = kwargs.get_item("litellm_logging_obj")? { - let logging_class = py - .import("litellm.litellm_core_utils.litellm_logging")? - .getattr("Logging")?; - if supplied.is_instance(&logging_class)? { - return Ok(Setup { - logger: supplied.extract()?, - kwargs: kwargs.unbind(), - supplied: true, - }); - } - } - - setup.getattr("prepare_environment")?.call0()?; - let guardrails = setup.getattr("applied_guardrails")?.call1((&kwargs,))?; - let function_id = kwargs.get_item("id")?; - - let mut targets = Targets::new(); - let dynamic = match kwargs.get_item("callbacks")? { - Some(value) => { - kwargs.del_item("callbacks")?; - (!value.is_none()).then_some(value) - } - None => None, - }; - let candidates = read_candidates(&setup, &mut targets, dynamic)?; - let facts = RegistrationFacts { - candidates, - input: entries(&setup, &mut targets, &read_registry(&setup, "input")?)?, - success: entries(&setup, &mut targets, &read_registry(&setup, "success")?)?, - failure: entries(&setup, &mut targets, &read_registry(&setup, "failure")?)?, - async_success: entries( - &setup, - &mut targets, - &read_registry(&setup, "async_success")?, - )? - .into_iter() - .map(|entry| entry.id) - .collect(), - async_failure: entries( - &setup, - &mut targets, - &read_registry(&setup, "async_failure")?, - )? - .into_iter() - .map(|entry| entry.id) - .collect(), - bootstrap_pending: setup.getattr("bootstrap_pending")?.call0()?.extract()?, - }; - apply( - &setup, - &targets, - &plan_registration(&facts), - function_id.as_ref(), - )?; - - let dynamic = split_dynamic(py, &setup, &mut targets, &kwargs)?; - setup.getattr("breadcrumb")?.call1((&kwargs,))?; - if let Some(logger_fn) = kwargs.get_item("logger_fn")? { - setup.getattr("logger_fn")?.call1((logger_fn,))?; - } - let model = match args.bind(py).get_item(0) { - Ok(model) => Some(model), - Err(_) => kwargs.get_item("model")?, - }; - let build = setup.getattr("build_logging")?; - let build_kwargs = PyDict::new(py); - build_kwargs.set_item("call_type", call_type)?; - build_kwargs.set_item("model", model)?; - build_kwargs.set_item("kwargs", &kwargs)?; - build_kwargs.set_item("start_time", start)?; - build_kwargs.set_item("asynchronous", asynchronous)?; - build_kwargs.set_item("dynamic_success", dynamic.success)?; - build_kwargs.set_item("dynamic_async_success", dynamic.async_success)?; - build_kwargs.set_item("dynamic_failure", dynamic.failure)?; - build_kwargs.set_item("guardrails", guardrails)?; - let logger = build.call((), Some(&build_kwargs))?; - Ok(Setup { - logger: logger.extract()?, - kwargs: kwargs.unbind(), - supplied: false, - }) -} 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 217f69f592d..de70f7ce6f0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -6,10 +6,8 @@ use litellm_core::ocr::LiteLLMOcrResponse; use litellm_core::ocr::hooks::OcrDuringCallRequest; use litellm_python_interop::to_py_preserving_errors as to_py; -use litellm_core::call_lifecycle::CallbackFamily; - use super::host::PythonPayload; -use crate::lifecycle::{PythonCallState, PythonLogger}; +use crate::lifecycle::PythonLogger; pub(super) fn update_logging( py: Python<'_>, @@ -34,63 +32,26 @@ pub(super) fn update_logging( pub(super) fn pre_call( py: Python<'_>, - state: &PythonCallState, + logger: &PythonLogger, request: &OcrDuringCallRequest, payload: &PythonPayload, ) -> PyResult<()> { - let logger = state.logger()?; - if state.supplied { - let additional_args = PyDict::new(py); - additional_args.set_item("complete_input_dict", &payload.body)?; - additional_args.set_item("headers", &payload.headers)?; - additional_args.set_item("api_base", &request.url)?; - let kwargs = PyDict::new(py); - kwargs.set_item("input", "OCR document processing")?; - kwargs.set_item("api_key", request.api_key.as_deref())?; - kwargs.set_item("additional_args", additional_args)?; - logger - .object(py) - .call_method("pre_call", (), Some(&kwargs))?; - return Ok(()); - } - let kwargs = PyDict::new(py); - kwargs.set_item("api_key", request.api_key.as_deref())?; - kwargs.set_item("body", &payload.body)?; - kwargs.set_item("headers", &payload.headers)?; - kwargs.set_item("url", &request.url)?; - py.import("litellm.rust_bridge.leaves")? - .getattr("record_pre_call")? - .call((logger.object(py),), Some(&kwargs))?; - state.dispatch_request(py, CallbackFamily::RequestPreCall) + py.import("litellm.rust_bridge.ocr")? + .getattr("pre_call")? + .call1((logger.object(py), request.api_key.as_deref(), &payload.body, &payload.headers, &request.url))?; + Ok(()) } pub(super) fn post_call( py: Python<'_>, - state: &PythonCallState, + logger: &PythonLogger, original_response: &Value, payload: &PythonPayload, ) -> PyResult<()> { - let logger = state.logger()?; - if state.supplied { - let additional_args = PyDict::new(py); - additional_args.set_item("complete_input_dict", &payload.body)?; - additional_args.set_item("headers", &payload.headers)?; - let kwargs = PyDict::new(py); - kwargs.set_item("original_response", to_py(py, original_response)?)?; - kwargs.set_item("additional_args", additional_args)?; - logger - .object(py) - .call_method("post_call", (), Some(&kwargs))?; - return Ok(()); - } - let kwargs = PyDict::new(py); - kwargs.set_item("original_response", to_py(py, original_response)?)?; - kwargs.set_item("body", &payload.body)?; - kwargs.set_item("headers", &payload.headers)?; - py.import("litellm.rust_bridge.leaves")? - .getattr("record_post_call")? - .call((logger.object(py),), Some(&kwargs))?; - state.dispatch_request(py, CallbackFamily::RequestPostCall) + py.import("litellm.rust_bridge.ocr")? + .getattr("post_call")? + .call1((logger.object(py), to_py(py, original_response)?, &payload.body, &payload.headers))?; + Ok(()) } pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult> { diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 9fcd53569bc..5ccb65780df 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -40,28 +40,17 @@ pub(super) struct PythonPayload { impl PythonPayload { fn from_request(py: Python<'_>, request: &OcrDuringCallRequest) -> PyResult { - let body = to_py(py, &request.body)? - .into_bound(py) - .cast_into::()?; + let body = to_py(py, &request.body)?.into_bound(py).cast_into::()?; let headers = PyDict::new(py); for (name, value) in &request.headers { headers.set_item(name, value)?; } - Ok(Self { - body: body.unbind(), - headers: headers.unbind(), - }) + Ok(Self { body: body.unbind(), headers: headers.unbind() }) } - fn write_back( - &self, - py: Python<'_>, - mut request: OcrDuringCallRequest, - ) -> PyResult { + fn write_back(&self, py: Python<'_>, mut request: OcrDuringCallRequest) -> PyResult { request.body = from_py(self.body.bind(py))?; - request.headers = self - .headers - .bind(py) + request.headers = self.headers.bind(py) .iter() .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) .collect::>>()?; @@ -92,9 +81,7 @@ impl PythonOcrHost { } fn project(&mut self, py: Python<'_>) -> PyResult { - let arguments = self - .signature - .bind(self.state.args.bind(py), self.state.kwargs.bind(py))?; + let arguments = self.signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?; let Projection { native, retained } = project(py, &arguments)?; let host_token_provider = retained.azure_ad_token_provider.is_some(); self.retained = Some(retained); @@ -128,7 +115,7 @@ impl PythonOcrHost { &retained.secret_fields, )?; let payload = PythonPayload::from_request(py, &request)?; - callbacks::pre_call(py, &self.state, &request, &payload)?; + callbacks::pre_call(py, logger, &request, &payload)?; let request = payload.write_back(py, request)?; self.retained_mut()?.payload = Some(payload); Ok(request) @@ -139,12 +126,13 @@ impl PythonOcrHost { py: Python<'_>, request: OcrPostCallRequest, ) -> PyResult { - let payload = self - .retained()? - .payload - .as_ref() - .ok_or_else(missing_state)?; - callbacks::post_call(py, &self.state, &request.original_response, payload)?; + let payload = self.retained()?.payload.as_ref().ok_or_else(missing_state)?; + callbacks::post_call( + py, + self.state.logger()?, + &request.original_response, + payload, + )?; Ok(request) } @@ -214,7 +202,9 @@ impl PythonRoute for PythonOcrHost { OcrHostOperation::AcquireAzureAdToken => { OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?)) } - OcrHostOperation::PreCall(request) => OcrHostResult::PreCall(Ok(request)), + OcrHostOperation::PreCall(request) => { + OcrHostResult::PreCall(Ok(request)) + } OcrHostOperation::DuringCall(request) => { OcrHostResult::DuringCall(Ok(self.during_call(py, request)?)) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index d1fd977c5b6..cbd32792843 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -62,16 +62,13 @@ fn call( ..OcrAdmission::all() }, ))?; - let host = PythonOcrHost::new( - PythonCallState::new( - py, - args.unbind(), - kwargs.copy()?.unbind(), - asynchronous, - signature.name, - )?, - signature, - ); + let host = PythonOcrHost::new(PythonCallState::new( + py, + args.unbind(), + kwargs.copy()?.unbind(), + asynchronous, + signature.name, + )?, signature); run_call(py, call, host) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 5157f7683ed..57cddac48cb 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -20,18 +20,14 @@ use crate::marshal::{BoundRouteInputs, Projection}; /// `optional_params`. const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"]; -fn project_document( - document: &Bound<'_, PyAny>, -) -> PyResult> { +fn project_document(document: &Bound<'_, PyAny>) -> PyResult> { let kind: String = document.get_item("type")?.extract()?; if kind != "file" { let value: serde_json::Value = from_py(document)?; - return Ok( - OcrDocument::try_from(value).map(|document| FileDocumentInput { - input: document.into(), - reader: None, - }), - ); + return Ok(OcrDocument::try_from(value).map(|document| FileDocumentInput { + input: document.into(), + reader: None, + })); } document.extract().map(Ok) } @@ -158,7 +154,10 @@ mod tests { let document = py .eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None) .unwrap(); - let error = project_document(&document).unwrap().err().unwrap(); + let error = project_document(&document) + .unwrap() + .err() + .unwrap(); assert!(error.to_string().contains("document")); }); } diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index c65bdb79b1b..3f183db9c10 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -1,4 +1,3 @@ -import datetime from asyncio import Future from collections.abc import Coroutine, Mapping from typing import Never, final @@ -116,20 +115,11 @@ class TokenCounter: def acount_request(self, body: bytes) -> Future[dict[str, object]]: ... def gil_stats() -> dict[str, int]: ... -def _debug_setup( - call_type: str, - args: tuple[object, ...], - kwargs: dict[str, object], - start: datetime.datetime, - asynchronous: bool, -) -> tuple[object, dict[str, object]]: ... - __all__ = [ "ResponsesWebSocketConnection", "RustBridgeDeclined", "RustUpstreamError", "TokenCounter", - "_debug_setup", "achat_completions", "amessages", "aocr", diff --git a/litellm/rust_bridge/leaves.py b/litellm/rust_bridge/leaves.py deleted file mode 100644 index 64e239d1615..00000000000 --- a/litellm/rust_bridge/leaves.py +++ /dev/null @@ -1,815 +0,0 @@ -"""Leaf helpers invoked by the native callback dispatcher. - -Rust selects every target, delivery and sequence. Nothing here chooses a callback, -reads a registry, or decides whether an event fires. Each function performs one -integration-specific or interop-specific action on values Rust hands it. - -Labeled leaf helpers pending native migration (see rust-callback-inventory.md): -`prepare_success_logging`, `prepare_failure_logging`, `dispatch_named_success`, -`dispatch_named_failure`, `dispatch_callable`. They are deleted when cost and -payload construction move into core and when each string integration becomes a -CustomLogger. -""" - -from __future__ import annotations - -import datetime -import json -import traceback -from collections.abc import Awaitable, Callable, Mapping -from typing import ( # noqa: TID251 # narrows the untyped legacy Logging and CustomLogger surfaces once - TYPE_CHECKING, - Final, - Literal, - Protocol, - cast, -) - -from litellm.integrations.custom_logger import CustomLogger - -if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import Logging - -TerminalFamily = Literal["sync_success", "async_success", "sync_failure", "async_failure"] -Family = Literal["request", TerminalFamily] -Details = dict[str, object] -Timestamp = datetime.datetime -LegacyCall = Callable[..., object] -LegacyAsyncCall = Callable[..., Awaitable[None]] - - -class LoggerView(Protocol): - model: str | None - messages: object - call_type: str - start_time: Timestamp - litellm_call_id: str - completion_start_time: Timestamp | None - model_call_details: Details - log_raw_request_response: bool - standard_callback_dynamic_params: object - standard_built_in_tools_params: object - - def record_api_call_start_time(self) -> None: ... - - def record_post_call( - self, original_response: object, input: object, api_key: object, additional_args: Details - ) -> None: ... - - def should_run_callback(self, callback: object, litellm_params: Details, event_hook: str) -> bool: ... - - def _pre_call(self, input: str, api_key: str | None, model: str | None, additional_args: Details) -> None: ... - - def _print_llm_call_debugging_log(self, api_base: str, headers: Details, additional_args: Details) -> None: ... - - def _get_request_curl_command( - self, api_base: str, headers: Details | None, additional_args: Details, data: object - ) -> str: ... - - def _get_masked_api_base(self, api_base: str) -> str: ... - - def _get_raw_request_body(self, data: object) -> Details: ... - - def _get_masked_headers(self, headers: Details) -> Details: ... - - def _response_cost_calculator(self, result: object) -> float | None: ... - - def _build_standard_logging_payload( - self, init_response_obj: object, start_time: Timestamp, end_time: Timestamp - ) -> object: ... - - def _handle_callback_failure(self, callback: object) -> None: ... - - -class IntegrationView(Protocol): - def log_pre_api_call(self, model: str | None, messages: object, kwargs: Details) -> None: ... - - def log_post_api_call( - self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp | None - ) -> None: ... - - def log_success_event( - self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp - ) -> None: ... - - def log_failure_event( - self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp - ) -> None: ... - - def async_log_success_event( - self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp - ) -> Awaitable[None]: ... - - def async_log_failure_event( - self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp - ) -> Awaitable[None]: ... - - def logging_hook(self, kwargs: Details, result: object, call_type: str) -> tuple[Details, object]: ... - - def async_logging_hook( - self, kwargs: Details, result: object, call_type: str - ) -> Awaitable[tuple[Details, object]]: ... - - def redact_standard_logging_payload_from_model_call_details(self, model_call_details: Details) -> Details: ... - - def log_input_event( - self, model: str | None, messages: object, kwargs: Details, print_verbose: LegacyCall, callback_func: LegacyCall - ) -> None: ... - - def log_event( - self, - kwargs: Details, - response_obj: object, - start_time: Timestamp, - end_time: Timestamp, - print_verbose: LegacyCall, - callback_func: LegacyCall, - ) -> None: ... - - def async_log_event( - self, - kwargs: Details, - response_obj: object, - start_time: Timestamp, - end_time: Timestamp, - print_verbose: LegacyCall, - callback_func: LegacyCall, - ) -> Awaitable[None]: ... - - -def _logger(logger: Logging) -> LoggerView: - return cast(LoggerView, logger) # cast-ok: legacy Logging is untyped; this protocol names the attributes we read - - -def _integration(callback: CustomLogger) -> IntegrationView: - return cast(IntegrationView, callback) # cast-ok: legacy CustomLogger methods are untyped - - -def _legacy_module() -> Mapping[str, object]: - from litellm.litellm_core_utils import litellm_logging - - return cast( - Mapping[str, object], vars(litellm_logging) - ) # cast-ok: module globals hold the legacy integration singletons - - -def _print_verbose() -> LegacyCall: - from litellm.litellm_core_utils import litellm_logging - - return cast(LegacyCall, litellm_logging.print_verbose) # cast-ok: legacy debug printer is untyped - - -def _method(target: object, name: str) -> LegacyCall: - return cast(LegacyCall, getattr(target, name)) # cast-ok: legacy integration singletons are untyped - - -def _async_method(target: object, name: str) -> LegacyAsyncCall: - return cast(LegacyAsyncCall, getattr(target, name)) # cast-ok: legacy integration singletons are untyped - - -def _redact_string(value: str) -> str: - from litellm.litellm_core_utils import litellm_logging - - return cast(Callable[[str], str], litellm_logging._redact_string)(value) # pyright: ignore[reportPrivateUsage] # cast-ok: legacy helper - - -def _redact_result(details: Details, result: object) -> object: - from litellm.litellm_core_utils import redact_messages - - redact: Final = cast(LegacyCall, redact_messages.redact_message_input_output_from_logging) # cast-ok: legacy helper - return redact(model_call_details=details, result=result) - - -def record_pre_call( - logger: Logging, - *, - api_key: str | None, - body: Details, - headers: dict[str, str], - url: str, -) -> None: - view: Final = _logger(logger) - additional_args: Final[Details] = {"complete_input_dict": body, "headers": headers, "api_base": url} - view._pre_call(input="OCR document processing", api_key=api_key, model=None, additional_args=additional_args) # pyright: ignore[reportPrivateUsage] # legacy state writer - view._print_llm_call_debugging_log(api_base=url, headers=dict(headers), additional_args=additional_args) # pyright: ignore[reportPrivateUsage] # legacy debug output - _capture_raw_request(view, additional_args) - _run_logger_fn(logger) - view.record_api_call_start_time() - - -def _capture_raw_request(view: LoggerView, additional_args: Details) -> None: - import litellm - from litellm.types.utils import RawRequestTypedDict - - if not (view.log_raw_request_response or litellm.log_raw_request_response): - return - details: Final = view.model_call_details - params: Final = cast(Details, details.get("litellm_params") or {}) # cast-ok: legacy nested dict - metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict - params.setdefault("metadata", metadata) - if litellm.turn_off_message_logging: - metadata["raw_request"] = "redacted by litellm. 'litellm.turn_off_message_logging=True'" - return - api_base: Final = str(additional_args.get("api_base") or "") - headers: Final = cast(Details, additional_args.get("headers") or {}) # cast-ok: legacy nested dict - body: Final = additional_args.get("complete_input_dict", {}) - try: - curl: Final = view._get_request_curl_command( # pyright: ignore[reportPrivateUsage] # legacy debug formatter - api_base=api_base, headers=headers, additional_args=additional_args, data=body - ) - metadata["raw_request"] = _redact_string(str(curl)) - details["raw_request_typed_dict"] = RawRequestTypedDict( - raw_request_api_base=view._get_masked_api_base(api_base), # pyright: ignore[reportPrivateUsage] # legacy masking - raw_request_body=view._get_raw_request_body(body), # pyright: ignore[reportPrivateUsage] # legacy masking - raw_request_headers=view._get_masked_headers(headers), # pyright: ignore[reportPrivateUsage] # legacy masking - error=None, - ) - except Exception as error: # noqa: BLE001 # raw-request capture is best effort by contract - details["raw_request_typed_dict"] = RawRequestTypedDict(error=str(error)) - metadata["raw_request"] = _redact_string(f"Unable to Log raw request: {error}") - - -def _run_logger_fn(logger: Logging) -> None: - from litellm._logging import verbose_logger - - logger_fn: Final = cast( - Callable[[Details], object] | None, getattr(logger, "logger_fn", None) - ) # cast-ok: user hook is untyped - if not callable(logger_fn): - return - try: - logger_fn(_logger(logger).model_call_details) - except Exception as error: # noqa: BLE001 # user logger_fn failures never fail the request - verbose_logger.exception("LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging %s", error) - - -def record_post_call(logger: Logging, *, original_response: object, body: Details, headers: dict[str, str]) -> None: - view: Final = _logger(logger) - serialized: Final = ( - json.dumps(original_response, default=str) if isinstance(original_response, dict) else original_response - ) - view.record_post_call( - original_response=serialized, - input=None, - api_key=None, - additional_args={"complete_input_dict": body, "headers": headers}, - ) - _run_logger_fn(logger) - _redact_result(view.model_call_details, serialized) - - -def log_pre_api_call(logger: Logging, callback: CustomLogger) -> None: - view: Final = _logger(logger) - _integration(callback).log_pre_api_call(model=view.model, messages=view.messages, kwargs=view.model_call_details) - - -def log_post_api_call(logger: Logging, callback: CustomLogger) -> None: - view: Final = _logger(logger) - _integration(callback).log_post_api_call( - kwargs=view.model_call_details, response_obj=None, start_time=view.start_time, end_time=None - ) - - -def dispatch_named_request(logger: Logging, name: str, event: Literal["pre_api_call", "post_api_call"]) -> None: - view: Final = _logger(logger) - module: Final = _legacy_module() - if name == "supabase" and event == "pre_api_call" and (client := module.get("supabaseClient")) is not None: - details: Final = view.model_call_details - _method(client, "input_log_event")( - model=view.model, - messages=view.messages, - end_user=details.get("user", "default"), - litellm_call_id=details["litellm_call_id"], - print_verbose=_print_verbose(), - ) - if name == "sentry" and (add_breadcrumb := module.get("add_breadcrumb")) is not None: - cast(LegacyCall, add_breadcrumb)( # cast-ok: legacy sentry hook - category="litellm.llm_call", message=f"Model Call Details {event}: {view.model_call_details}", level="info" - ) - - -def dispatch_callable_request(logger: Logging, callback: LegacyCall) -> None: - custom: Final = _legacy_module().get("customLogger") - if not isinstance(custom, CustomLogger): - return - view: Final = _logger(logger) - _integration(custom).log_input_event( - model=view.model, - messages=view.messages, - kwargs=view.model_call_details, - print_verbose=_print_verbose(), - callback_func=callback, - ) - - -def report_target_failure(logger: Logging, callback: object, family: Family, error: BaseException) -> None: - from litellm._logging import verbose_logger - - verbose_logger.error( - "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while %s logging with %s: %s", - family, - callback, - "".join(traceback.format_exception(error)), - ) - capture: Final = _legacy_module().get("capture_exception") - if capture is not None and family in ("request", "sync_success", "sync_failure"): - cast(LegacyCall, capture)(error) # cast-ok: legacy sentry hook - if family not in ("request", "sync_failure"): - _logger(logger)._handle_callback_failure(callback=callback) # pyright: ignore[reportPrivateUsage] # legacy prometheus counter - - -def prepare_success_logging(logger: Logging, response: object, start_time: Timestamp, end_time: Timestamp) -> object: - from litellm.litellm_core_utils.litellm_logging import emit_standard_logging_payload - from litellm.types.utils import StandardLoggingPayload - - view: Final = _logger(logger) - details: Final = view.model_call_details - if view.completion_start_time is None: - view.completion_start_time = end_time - details["completion_start_time"] = end_time - details["log_event_type"] = "successful_api_call" - details["end_time"] = end_time - details["cache_hit"] = None - hidden: Final = cast(Details, getattr(response, "_hidden_params", None) or {}) # cast-ok: legacy response attribute - params: Final = cast(Details | None, details.get("litellm_params")) # cast-ok: legacy nested dict - if hidden and params is not None: - metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict - params["metadata"] = metadata - metadata["hidden_params"] = hidden - existing: Final = details.get("response_cost") - if "response_cost" in hidden: - details["response_cost"] = hidden["response_cost"] - elif existing is None or existing == 0: - details["response_cost"] = view._response_cost_calculator(result=response) # pyright: ignore[reportPrivateUsage] # labeled leaf: native cost pending - payload: Final = view._build_standard_logging_payload(response, start_time, end_time) # pyright: ignore[reportPrivateUsage] # labeled leaf: native payload pending - details["standard_logging_object"] = payload - if payload is not None: - emit_standard_logging_payload( - cast(StandardLoggingPayload, payload) - ) # cast-ok: legacy builder returns the payload TypedDict - return _redact_result(details, response) - - -def prepare_failure_logging( - logger: Logging, exception: BaseException, start_time: Timestamp, end_time: Timestamp -) -> str: - from litellm.litellm_core_utils import litellm_logging - - formatted: Final = "".join(traceback.format_exception(exception)) - view: Final = _logger(logger) - details: Final = view.model_call_details - if details.get("exception") is exception and details.get("standard_logging_object") is not None: - return formatted - details["log_event_type"] = "failed_api_call" - details["exception"] = exception - details["traceback_exception"] = _redact_string(formatted) - details["end_time"] = end_time - details.setdefault("original_response", None) - if details.get("combined_usage_object") is None: - details["response_cost"] = 0 - headers: Final = getattr(exception, "headers", None) - if isinstance(headers, dict): - params: Final = cast(Details, details.setdefault("litellm_params", {})) # cast-ok: legacy nested dict - metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict - metadata.update(cast(Details, headers)) # cast-ok: exception headers are a plain dict - build: Final = cast( - LegacyCall, litellm_logging.get_standard_logging_object_payload - ) # cast-ok: labeled leaf: native payload pending - details["standard_logging_object"] = build( - kwargs=details, - init_response_obj={}, - start_time=start_time, - end_time=end_time, - logging_obj=logger, - status="failure", - error_str=_redact_string(str(exception)), - original_exception=exception, - standard_built_in_tools_params=view.standard_built_in_tools_params, - ) - return formatted - - -_EVENT_HOOKS: Final[Mapping[TerminalFamily, str]] = { - "sync_success": "success_handler", - "async_success": "async_success_handler", - "sync_failure": "failure_handler", - "async_failure": "async_failure_handler", -} - - -def should_run_callback(logger: Logging, callback: object, family: TerminalFamily) -> bool: - view: Final = _logger(logger) - params: Final = cast(Details, view.model_call_details.get("litellm_params") or {}) # cast-ok: legacy nested dict - return view.should_run_callback(callback=callback, litellm_params=params, event_hook=_EVENT_HOOKS[family]) - - -def should_run_guardrail_hook(logger: Logging, callback: object) -> bool: - from litellm.integrations.custom_guardrail import CustomGuardrail - from litellm.types.guardrails import GuardrailEventHooks - - if not isinstance(callback, CustomGuardrail): - return True - decide: Final = cast(LegacyCall, callback.should_run_guardrail) # cast-ok: legacy guardrail method is untyped - return decide(data=_logger(logger).model_call_details, event_type=GuardrailEventHooks.logging_only) is True - - -def logging_hook(logger: Logging, callback: CustomLogger, result: object) -> object: - view: Final = _logger(logger) - details, replaced = _integration(callback).logging_hook( - kwargs=view.model_call_details, result=result, call_type=view.call_type - ) - view.model_call_details = details - return replaced - - -async def async_logging_hook(logger: Logging, callback: CustomLogger, result: object) -> object: - from litellm.integrations.custom_guardrail import CustomGuardrail - from litellm.litellm_core_utils import redact_messages - - view: Final = _logger(logger) - redact: Final = cast( - LegacyCall, redact_messages.redact_message_input_output_from_custom_logger - ) # cast-ok: legacy helper - redacted: Final = ( - result - if isinstance(callback, CustomGuardrail) - else redact(result=result, litellm_logging_obj=logger, custom_logger=callback) - ) - details, replaced = await _integration(callback).async_logging_hook( - kwargs=view.model_call_details, result=redacted, call_type=view.call_type - ) - view.model_call_details = details - return replaced - - -def mark_logged(logger: Logging, marker: str) -> None: - _logger(logger).model_call_details[marker] = True - - -def already_logged(logger: Logging, marker: str) -> bool: - return _logger(logger).model_call_details.get(marker, False) is True - - -def log_success_event( - logger: Logging, callback: CustomLogger, response: object, start_time: Timestamp, end_time: Timestamp -) -> None: - _integration(callback).log_success_event( - kwargs=_logger(logger).model_call_details, response_obj=response, start_time=start_time, end_time=end_time - ) - - -def async_log_success_event( - logger: Logging, callback: CustomLogger, response: object, start_time: Timestamp, end_time: Timestamp -) -> Awaitable[None]: - from litellm.litellm_core_utils import redact_messages - - integration: Final = _integration(callback) - details: Final = integration.redact_standard_logging_payload_from_model_call_details( - model_call_details=_logger(logger).model_call_details - ) - redact: Final = cast( - Callable[..., Details], redact_messages.redact_streaming_responses_for_custom_logger - ) # cast-ok: legacy helper - view: Final = redact(model_call_details=details, custom_logger=callback) - return integration.async_log_success_event( - kwargs=view, response_obj=response, start_time=start_time, end_time=end_time - ) - - -def log_failure_event(logger: Logging, callback: CustomLogger, start_time: Timestamp, end_time: Timestamp) -> None: - _integration(callback).log_failure_event( - kwargs=_logger(logger).model_call_details, response_obj=None, start_time=start_time, end_time=end_time - ) - - -def async_log_failure_event( - logger: Logging, callback: CustomLogger, start_time: Timestamp, end_time: Timestamp -) -> Awaitable[None]: - return _integration(callback).async_log_failure_event( - kwargs=_logger(logger).model_call_details, response_obj=None, start_time=start_time, end_time=end_time - ) - - -def _custom_logger_singleton() -> IntegrationView: - from litellm.litellm_core_utils import litellm_logging - - existing: Final = _legacy_module().get("customLogger") - if isinstance(existing, CustomLogger): - return _integration(existing) - created: Final = CustomLogger() - litellm_logging.customLogger = created # pyright: ignore[reportAttributeAccessIssue] # legacy module global - return _integration(created) - - -def dispatch_callable( - logger: Logging, - callback: LegacyCall, - family: TerminalFamily, - response: object, - start_time: Timestamp, - end_time: Timestamp, -) -> Awaitable[None] | None: - custom: Final = _custom_logger_singleton() - details: Final = _logger(logger).model_call_details - match family: - case "sync_success" | "sync_failure": - custom.log_event( - kwargs=details, - response_obj=response, - start_time=start_time, - end_time=end_time, - print_verbose=_print_verbose(), - callback_func=callback, - ) - return None - case "async_success" | "async_failure": - return custom.async_log_event( - kwargs=details, - response_obj=response, - start_time=start_time, - end_time=end_time, - print_verbose=_print_verbose(), - callback_func=callback, - ) - - -_SUCCESS_SINGLETONS: Final[Mapping[str, str]] = { - "promptlayer": "promptLayerLogger", - "supabase": "supabaseClient", - "wandb": "weightsBiasesLogger", - "logfire": "logfireLogger", - "lunary": "lunaryLogger", - "helicone": "heliconeLogger", - "greenscale": "greenscaleLogger", - "athina": "athinaLogger", - "traceloop": "traceloopLogger", - "s3": "s3Logger", - "openmeter": "openMeterLogger", -} - - -def dispatch_named_success( - logger: Logging, name: str, response: object, start_time: Timestamp, end_time: Timestamp -) -> Awaitable[None] | None: - view: Final = _logger(logger) - details: Final = view.model_call_details - print_verbose: Final = _print_verbose() - integration: Final = _legacy_module().get(_SUCCESS_SINGLETONS.get(name, "")) - without_response: Final = {key: value for key, value in details.items() if key != "original_response"} - match name: - case "promptlayer" | "wandb" | "athina" if integration is not None: - _method(integration, "log_event")( - kwargs=details, - response_obj=response, - start_time=start_time, - end_time=end_time, - print_verbose=print_verbose, - ) - case "logfire" if integration is not None: - from litellm.integrations.logfire_logger import LogfireLevel - - _method(integration, "log_event")( - kwargs=without_response, - response_obj=response, - start_time=start_time, - end_time=end_time, - print_verbose=print_verbose, - level=LogfireLevel.INFO.value, - ) - case "greenscale" if integration is not None: - _method(integration, "log_event")( - kwargs=without_response, - response_obj=response, - start_time=start_time, - end_time=end_time, - print_verbose=print_verbose, - ) - case "supabase" if integration is not None: - _method(integration, "log_event")( - model=view.model, - messages=view.messages, - end_user=details.get("user", "default"), - response_obj=response, - start_time=start_time, - end_time=end_time, - litellm_call_id=details["litellm_call_id"], - print_verbose=print_verbose, - ) - case "lunary" if integration is not None: - _method(integration, "log_event")( - kwargs=details, - type="llm", - event="end", - model=view.model, - input=details["input"], - user_id=details.get("user", "default"), - response_obj=response, - start_time=start_time, - end_time=end_time, - run_id=view.litellm_call_id, - print_verbose=print_verbose, - ) - case "helicone" if integration is not None: - _method(integration, "log_success")( - model=view.model, - messages=view.messages, - response_obj=response, - start_time=start_time, - end_time=end_time, - print_verbose=print_verbose, - kwargs=details, - ) - case "langfuse": - _langfuse( - logger, response=response, start_time=start_time, end_time=end_time, level=None, status_message=None - ) - case "traceloop" if integration is not None: - _method(integration, "log_event")( - kwargs=details, - response_obj=response, - start_time=start_time, - end_time=end_time, - user_id=details.get("user", None), - print_verbose=print_verbose, - ) - case "s3" if integration is not None: - _method(integration, "log_event")( - kwargs=details, - response_obj=response, - start_time=start_time, - end_time=end_time, - print_verbose=print_verbose, - ) - case "openmeter" if integration is not None: - return _async_method(integration, "async_log_success_event")( - kwargs=details, response_obj=response, start_time=start_time, end_time=end_time - ) - case "dynamodb": - return _dynamodb(details, response, start_time, end_time, print_verbose) - case _: - return None - return None - - -def _dynamodb( - details: Details, response: object, start_time: Timestamp, end_time: Timestamp, print_verbose: LegacyCall -) -> Awaitable[None]: - from litellm.integrations.dynamodb import DyanmoDBLogger - from litellm.litellm_core_utils import litellm_logging - - existing: Final = _legacy_module().get("dynamoLogger") - dynamo: Final = existing if isinstance(existing, DyanmoDBLogger) else DyanmoDBLogger() - litellm_logging.dynamoLogger = dynamo # pyright: ignore[reportAttributeAccessIssue] # legacy module global - return _async_method(dynamo, "_async_log_event")( - kwargs=details, response_obj=response, start_time=start_time, end_time=end_time, print_verbose=print_verbose - ) - - -def _langfuse( - logger: Logging, - *, - response: object, - start_time: Timestamp, - end_time: Timestamp, - level: str | None, - status_message: str | None, -) -> None: - from litellm.integrations.langfuse import langfuse_handler - - view: Final = _logger(logger) - module: Final = _legacy_module() - kwargs: Final = {key: value for key, value in view.model_call_details.items() if key != "original_response"} - select: Final = cast( - LegacyCall, langfuse_handler.LangFuseHandler.get_langfuse_logger_for_request - ) # cast-ok: legacy factory - handler: Final = select( - globalLangfuseLogger=module.get("langFuseLogger"), - standard_callback_dynamic_params=view.standard_callback_dynamic_params, - in_memory_dynamic_logger_cache=module["in_memory_dynamic_logger_cache"], - ) - if handler is None: - return - extra: Final[Details] = {"level": level, "status_message": status_message} if level is not None else {} - result: Final = _method(handler, "log_event_on_langfuse")( - kwargs=kwargs, - response_obj=response, - start_time=start_time, - end_time=end_time, - user_id=kwargs.get("user", None), - **extra, - ) - trace_id: Final = ( - cast(Details, result).get("trace_id") if isinstance(result, dict) else None - ) # cast-ok: legacy response dict - if trace_id is not None: - _method(module["in_memory_trace_id_cache"], "set_cache")( - litellm_call_id=view.litellm_call_id, service_name="langfuse", trace_id=trace_id - ) - - -def dispatch_named_failure( - logger: Logging, - name: str, - exception: BaseException, - formatted: str, - start_time: Timestamp, - end_time: Timestamp, -) -> None: - view: Final = _logger(logger) - module: Final = _legacy_module() - details: Final = view.model_call_details - print_verbose: Final = _print_verbose() - match name: - case "lunary" if (lunary := module.get("lunaryLogger")) is not None: - _method(lunary, "log_event")( - kwargs=details, - type="llm", - event="error", - user_id=details.get("user", "default"), - model=view.model, - input=details["input"], - error=formatted, - run_id=view.litellm_call_id, - start_time=start_time, - end_time=end_time, - print_verbose=print_verbose, - ) - case "sentry" if (capture := module.get("capture_exception")) is not None: - cast(LegacyCall, capture)(exception) # cast-ok: legacy sentry hook - case "supabase" if (supabase := module.get("supabaseClient")) is not None: - _method(supabase, "log_event")( - model=view.model, - messages=view.messages, - end_user=details.get("user", "default"), - response_obj=None, - start_time=start_time, - end_time=end_time, - litellm_call_id=details["litellm_call_id"], - print_verbose=print_verbose, - ) - case "langfuse": - _langfuse( - logger, - response=None, - start_time=start_time, - end_time=end_time, - level="ERROR", - status_message=str(exception), - ) - case "traceloop" if (traceloop := module.get("traceloopLogger")) is not None: - _method(traceloop, "log_event")( - start_time=start_time, - end_time=end_time, - response_obj=None, - user_id=details.get("user", None), - print_verbose=print_verbose, - status_message=str(exception), - level="ERROR", - kwargs=details, - ) - case "logfire" if (logfire := module.get("logfireLogger")) is not None: - from litellm.integrations.logfire_logger import LogfireLevel - - _method(logfire, "log_event")( - kwargs={ - **{key: value for key, value in details.items() if key != "original_response"}, - "exception": exception, - }, - response_obj=None, - start_time=start_time, - end_time=end_time, - level=LogfireLevel.ERROR.value, - print_verbose=print_verbose, - ) - case _: - return - - -def restore_correlation_context(logger: object) -> None: - from litellm import utils - - utils._restore_correlation_context_if_supported(logger) # pyright: ignore[reportPrivateUsage] # legacy interop helper - - -def submit_worker(job: Callable[[], None]) -> None: - import contextvars - - from litellm.litellm_core_utils.thread_pool_executor import executor - - context: Final = contextvars.copy_context() - _ = executor.submit(context.run, job) - - -def enqueue_background(coroutine: Awaitable[None]) -> None: - import contextvars - - from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER - - enqueue: Final = cast( - Callable[[Awaitable[None]], None], GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue - ) # cast-ok: legacy worker accepts any coroutine - contextvars.copy_context().run(enqueue, coroutine) - - -def now() -> Timestamp: - return datetime.datetime.now() diff --git a/litellm/rust_bridge/lifecycle.py b/litellm/rust_bridge/lifecycle.py index c3ab962a5cb..35429ea7ccc 100644 --- a/litellm/rust_bridge/lifecycle.py +++ b/litellm/rust_bridge/lifecycle.py @@ -1,8 +1,18 @@ from __future__ import annotations +import datetime +import uuid from collections.abc import Awaitable, Mapping from dataclasses import dataclass -from typing import Final, Protocol +from typing import ( + TYPE_CHECKING, + Final, + Protocol, + cast, # noqa: TID251 # bounded compatibility calls into legacy Python integrations +) + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging @dataclass(frozen=True, slots=True) @@ -42,6 +52,47 @@ async def drive(execution: Execution) -> object: execution.close() +class MetadataUpdater(Protocol): + def __call__( + self, + result: object, + logging_obj: Logging, + model: str | None, + kwargs: dict[str, object], + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + +@dataclass(frozen=True, slots=True) +class CallSetup: + logger: Logging + kwargs: dict[str, object] + + +def setup( + call_type: str, + args: tuple[object, ...], + kwargs: Mapping[str, object], + start_time: datetime.datetime, + asynchronous: bool, +) -> CallSetup: + from litellm import utils + from litellm.litellm_core_utils.litellm_logging import Logging + + arguments: Final = { # mutable-ok: function_setup consumes an owned kwargs dict + "litellm_call_id": str(uuid.uuid4()), + **kwargs, + } + supplied: Final = arguments.get("litellm_logging_obj") + if isinstance(supplied, Logging): + return CallSetup(supplied, arguments) + logger, prepared = utils.function_setup( + call_type, utils.Rules(), start_time, *args, is_async_call=asynchronous, **arguments + ) + return CallSetup(logger, prepared) + + def check_limits(kwargs: Mapping[str, object]) -> None: import litellm from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit @@ -51,3 +102,19 @@ def check_limits(kwargs: Mapping[str, object]) -> None: raise litellm.BudgetExceededError(current_cost=current_cost, max_budget=litellm.max_budget) if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request): raise RuntimeError("Max retries per request hit!") + + +def finalize( + response: object, + logger: Logging, + kwargs: dict[str, object], + start_time: datetime.datetime, + end_time: datetime.datetime, +) -> None: + from litellm.litellm_core_utils.llm_response_utils import response_metadata + + model: Final = kwargs.get("model") + update: Final = cast( # cast-ok: legacy metadata function accepts concrete kwargs + MetadataUpdater, response_metadata.update_response_metadata + ) + update(response, logger, model if isinstance(model, str) else None, kwargs, start_time, end_time) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 069ec7a4ae6..35c399de73b 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -22,6 +22,10 @@ class OcrLoggingProtocol(Protocol): custom_llm_provider: str, ) -> object: ... + def pre_call(self, *, input: str, api_key: str | None, additional_args: dict[str, object]) -> object: ... + + def post_call(self, *, original_response: object, additional_args: dict[str, object]) -> object: ... + def _redact(params: Mapping[str, object], secret_fields: Sequence[str]) -> dict[str, object]: return { # mutable-ok: update_from_kwargs takes dict @@ -60,6 +64,32 @@ def update_logging( ) +def pre_call( + logger: OcrLoggingProtocol, + api_key: str | None, + body: dict[str, object], + headers: dict[str, str], + url: str, +) -> None: + logger.pre_call( + input="OCR document processing", + api_key=api_key, + additional_args={"complete_input_dict": body, "headers": headers, "api_base": url}, + ) + + +def post_call( + logger: OcrLoggingProtocol, + original_response: object, + body: dict[str, object], + headers: dict[str, str], +) -> None: + logger.post_call( + original_response=original_response, + additional_args={"complete_input_dict": body, "headers": headers}, + ) + + class RustOcr(Protocol): def __call__(self, *args: object, **kwargs: object) -> OCRResponse: ... diff --git a/litellm/rust_bridge/setup.py b/litellm/rust_bridge/setup.py deleted file mode 100644 index 820d39c7b3e..00000000000 --- a/litellm/rust_bridge/setup.py +++ /dev/null @@ -1,262 +0,0 @@ -from __future__ import annotations - -import datetime -from collections.abc import Callable, Mapping, MutableSequence, Sequence -from typing import ( # noqa: TID251 # narrows legacy untyped registries at the boundary - TYPE_CHECKING, - Final, - Literal, - Protocol, - cast, -) - -from litellm.integrations.custom_logger import CustomLogger - -if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import Logging - -CallbackTarget = str | Callable[..., object] | CustomLogger -RegistryName = Literal["input", "async_input", "success", "async_success", "failure", "async_failure", "callbacks"] - -_REGISTRY_ATTRIBUTES: Final[Mapping[RegistryName, str]] = { - "input": "input_callback", - "async_input": "_async_input_callback", - "success": "success_callback", - "async_success": "_async_success_callback", - "failure": "failure_callback", - "async_failure": "_async_failure_callback", - "callbacks": "callbacks", -} - - -class _CallbackManager(Protocol): - def add_litellm_success_callback(self, callback: CallbackTarget) -> None: ... - - def add_litellm_failure_callback(self, callback: CallbackTarget) -> None: ... - - def add_litellm_async_success_callback(self, callback: CallbackTarget) -> None: ... - - def add_litellm_async_failure_callback(self, callback: CallbackTarget) -> None: ... - - -class _LoggingFactory(Protocol): - def __call__( - self, - *, - model: str | None, - messages: object, - stream: bool, - litellm_call_id: str, - litellm_trace_id: str | None, - function_id: str, - call_type: str, - start_time: datetime.datetime, - dynamic_success_callbacks: list[CallbackTarget] | None, - dynamic_failure_callbacks: list[CallbackTarget] | None, - dynamic_async_success_callbacks: list[CallbackTarget] | None, - dynamic_async_failure_callbacks: list[CallbackTarget] | None, - kwargs: dict[str, object], - applied_guardrails: list[str], - supports_correlation_logging: bool, - ) -> Logging: ... - - -class _EnvironmentUpdater(Protocol): - def __call__( - self, - *, - model: str | None, - user: str, - optional_params: dict[str, object], - litellm_params: dict[str, object], - stream_options: object, - ) -> None: ... - - -def registry(name: RegistryName) -> MutableSequence[CallbackTarget]: - import litellm - - return cast( - MutableSequence[CallbackTarget], getattr(litellm, _REGISTRY_ATTRIBUTES[name]) - ) # cast-ok: legacy module-level lists are untyped - - -def is_async_callable(callback: object) -> bool: - from litellm.litellm_core_utils.cached_imports import get_coroutine_checker - - return get_coroutine_checker().is_async_callable(callback) - - -def is_known_name(callback: str) -> bool: - import litellm - - known: Final = cast(Sequence[str], litellm._known_custom_logger_compatible_callbacks) # pyright: ignore[reportPrivateUsage] # cast-ok: registry list has no public typed accessor - return callback in known - - -def resolve_named_integration(callback: str) -> CustomLogger | None: - from litellm.litellm_core_utils import litellm_logging - - resolve: Final = cast( # cast-ok: legacy factory is untyped at its definition - Callable[..., CustomLogger | None], - litellm_logging._init_custom_logger_compatible_class, # pyright: ignore[reportPrivateUsage] # legacy factory - ) - return resolve(callback, internal_usage_cache=None, llm_router=None) - - -def async_success_registry_has_type(callback: object) -> bool: - return any(type(existing) is type(callback) for existing in registry("async_success")) - - -def bootstrap_pending() -> bool: - from litellm import utils - - return not utils.callback_list - - -def bootstrap(function_id: str | None) -> None: - from litellm import utils - from litellm.litellm_core_utils import cached_imports - - combined: Final = list({*registry("input"), *registry("success"), *registry("failure")}) - utils.callback_list = cast( - list[str], combined - ) # rebind-ok: legacy module global consumed by set_callbacks # cast-ok: legacy list annotation is narrower than its contents - set_callbacks: Final = cast(Callable[..., None], cached_imports.get_set_callbacks()) # pyright: ignore[reportUnknownMemberType] # cast-ok: cached import is untyped - set_callbacks(callback_list=combined, function_id=function_id) - - -def expand_named(callback: str, event: Literal["success", "failure"]) -> None: - from litellm import utils - - utils._add_custom_logger_callback_to_specific_event(callback, event) # pyright: ignore[reportPrivateUsage] # legacy expansion helper - - -def append_registry(name: RegistryName, callback: CallbackTarget) -> None: - import litellm - - manager: Final = cast(_CallbackManager, litellm.logging_callback_manager) # cast-ok: legacy manager is untyped - match name: - case "input" | "async_input": - registry(name).append(callback) - case "success": - manager.add_litellm_success_callback(callback) - case "async_success": - manager.add_litellm_async_success_callback(callback) - case "failure": - manager.add_litellm_failure_callback(callback) - case "async_failure": - manager.add_litellm_async_failure_callback(callback) - case "callbacks": - raise KeyError(name) - - -def remove_registry(name: RegistryName, callback: CallbackTarget) -> None: - target: Final = registry(name) - for index in range(len(target) - 1, -1, -1): - if target[index] is callback or target[index] == callback: - del target[index] - return - - -def logger_fn(callback: object) -> None: - from litellm import utils - - utils.user_logger_fn = callback # rebind-ok: legacy module global read by pre_call - - -def breadcrumb(kwargs: Mapping[str, object]) -> None: - from litellm import utils - - add_breadcrumb: Final = cast( - Callable[..., None] | None, utils.add_breadcrumb - ) # cast-ok: legacy sentry hook is untyped - if add_breadcrumb is None: - return - import litellm - from litellm.litellm_core_utils import core_helpers - - deep_copy: Final = cast( # cast-ok: legacy helper is untyped - Callable[[dict[str, object]], dict[str, object]], - core_helpers.safe_deep_copy, # pyright: ignore[reportUnknownMemberType] # legacy helper - ) - try: - copied: dict[str, object] = deep_copy(dict(kwargs)) - except Exception: # noqa: BLE001 # legacy breadcrumb falls back to the live mapping - copied = dict(kwargs) - hidden: Final = frozenset(("messages", "input", "prompt")) if litellm.turn_off_message_logging else frozenset[str]() - details: Final = {key: value for key, value in copied.items() if key not in hidden} - add_breadcrumb(category="litellm.llm_call", message=f"Keyword Args: {details}", level="info") - - -def prepare_environment() -> None: - from litellm import utils - - utils.custom_llm_setup() - - -def applied_guardrails(kwargs: Mapping[str, object]) -> list[str]: - from litellm.utils import get_applied_guardrails - - return get_applied_guardrails(dict(kwargs)) - - -def build_logging( - *, - call_type: str, - model: str | None, - kwargs: dict[str, object], - start_time: datetime.datetime, - asynchronous: bool, - dynamic_success: Sequence[CallbackTarget] | None, - dynamic_async_success: Sequence[CallbackTarget] | None, - dynamic_failure: Sequence[CallbackTarget] | None, - guardrails: Sequence[str], -) -> Logging: - from litellm.litellm_core_utils.cached_imports import get_litellm_logging_class - - function_id: Final = kwargs.get("id") - metadata: Final = kwargs.get("metadata") - trace_id: Final = kwargs.get("litellm_trace_id") - factory: Final = cast(_LoggingFactory, get_litellm_logging_class()) # cast-ok: legacy constructor is untyped - logger: Final = factory( - model=model, - messages="default-message-value", - stream=False, - litellm_call_id=str(kwargs["litellm_call_id"]), - litellm_trace_id=trace_id if isinstance(trace_id, str) else None, - function_id=function_id if isinstance(function_id, str) else "", - call_type=call_type, - start_time=start_time, - dynamic_success_callbacks=list(dynamic_success) if dynamic_success is not None else None, - dynamic_failure_callbacks=list(dynamic_failure) if dynamic_failure is not None else None, - dynamic_async_success_callbacks=list(dynamic_async_success) if dynamic_async_success is not None else None, - dynamic_async_failure_callbacks=None, - kwargs=kwargs, - applied_guardrails=list(guardrails), - supports_correlation_logging=asynchronous, - ) - litellm_metadata: Final = kwargs.get("litellm_metadata") - litellm_params: Final[dict[str, object]] = { - "api_base": "", - **({"metadata": kwargs["metadata"]} if "metadata" in kwargs else {}), - **( - { - "litellm_metadata": litellm_metadata, - **( - {} if metadata else {"metadata": dict(cast(Mapping[str, object], litellm_metadata))} - ), # cast-ok: isinstance narrows only to dict[Unknown, Unknown] - } - if isinstance(litellm_metadata, dict) - else {} - ), - } - update: Final = cast(_EnvironmentUpdater, logger.update_environment_variables) # cast-ok: legacy method is untyped - update( - model=model, - user="", - optional_params={}, - litellm_params=litellm_params, - stream_options=kwargs.get("stream_options"), - ) - return logger diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index 272974b2689..96983968bee 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -8,8 +8,7 @@ import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.ocr import legacy from litellm.rust_bridge import bindings, configuration -from litellm.rust_bridge.leaves import record_post_call, record_pre_call -from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, update_logging +from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, post_call, pre_call, update_logging @pytest.fixture(autouse=True) @@ -66,26 +65,25 @@ def test_logging_redacts_views_and_preserves_opaque_arguments_and_pricing() -> N def test_logging_callbacks_receive_captured_payload_roots_and_propagate_errors() -> None: - logger: Final = Mock(log_raw_request_response=False, logger_fn=None, model_call_details={}) + logger: Final = Mock() body: Final[dict[str, object]] = {"document": "original"} headers: Final = {"authorization": "key"} response: Final = object() - record_pre_call(logger, api_key="key", body=body, headers=headers, url="https://provider") - record_post_call(logger, original_response=response, body=body, headers=headers) - logger._pre_call.assert_called_once_with( + pre_call(logger, "key", body, headers, "https://provider") + post_call(logger, response, body, headers) + logger.pre_call.assert_called_once_with( input="OCR document processing", api_key="key", - model=None, additional_args={"complete_input_dict": body, "headers": headers, "api_base": "https://provider"}, ) - for callback in (logger._pre_call, logger.record_post_call): + for callback in (logger.pre_call, logger.post_call): assert callback.call_args.kwargs["additional_args"]["complete_input_dict"] is body assert callback.call_args.kwargs["additional_args"]["headers"] is headers - assert logger.record_post_call.call_args.kwargs["original_response"] is response + assert logger.post_call.call_args.kwargs["original_response"] is response failure: Final = RuntimeError("callback failed") - failing_logger: Final = Mock(_pre_call=Mock(side_effect=failure)) + failing_logger: Final = Mock(pre_call=Mock(side_effect=failure)) with pytest.raises(RuntimeError) as caught: - record_pre_call(failing_logger, api_key=None, body=body, headers=headers, url="https://provider") + pre_call(failing_logger, None, body, headers, "https://provider") assert caught.value is failure diff --git a/tests/test_litellm/rust_bridge/test_setup.py b/tests/test_litellm/rust_bridge/test_setup.py deleted file mode 100644 index ecfd68349c8..00000000000 --- a/tests/test_litellm/rust_bridge/test_setup.py +++ /dev/null @@ -1,169 +0,0 @@ -import datetime -from collections.abc import Iterator -from typing import Final - -import pytest - -import litellm -from litellm import utils -from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.rust_bridge.loader import native_bridge_available - -pytestmark = pytest.mark.skipif(not native_bridge_available(), reason="requires the Rust extension") - -REGISTRIES: Final = ( - "input_callback", - "_async_input_callback", - "success_callback", - "_async_success_callback", - "failure_callback", - "_async_failure_callback", - "callbacks", -) - - -@pytest.fixture -def clean_registries(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: - for name in REGISTRIES: - monkeypatch.setattr(litellm, name, []) - monkeypatch.setattr(utils, "callback_list", []) - monkeypatch.setattr(utils, "user_logger_fn", None) - monkeypatch.setenv("OPENMETER_API_KEY", "test") - monkeypatch.setenv("OPENMETER_API_ENDPOINT", "http://127.0.0.1:9") - yield - - -class SyncLogger(CustomLogger): - pass - - -class AsyncOnly(CustomLogger): - pass - - -def sync_fn(*args: object, **kwargs: object) -> None: - del args, kwargs - - -async def async_fn(*args: object, **kwargs: object) -> None: - del args, kwargs - - -def snapshot() -> dict[str, list[object]]: - return {name: list(getattr(litellm, name)) for name in REGISTRIES} - - -def run_legacy(kwargs: dict[str, object]) -> tuple[Logging, dict[str, object]]: - logger, prepared = utils.function_setup( - "ocr", - utils.Rules(), - datetime.datetime.now(), - is_async_call=False, - **{"litellm_call_id": "legacy", **kwargs}, - ) - assert isinstance(logger, Logging) - return logger, prepared - - -def run_native(kwargs: dict[str, object]) -> tuple[Logging, dict[str, object]]: - from litellm.rust_bridge import _native - - logger, prepared = _native._debug_setup( - "ocr", (), {"litellm_call_id": "native", **kwargs}, datetime.datetime.now(), False - ) - assert isinstance(logger, Logging) - return logger, prepared - - -@pytest.mark.parametrize( - "globals_before, kwargs", - [ - ({}, {}), - ({"callbacks": [SyncLogger()]}, {}), - ({"callbacks": [async_fn]}, {}), - ({}, {"callbacks": [SyncLogger(), sync_fn]}), - ({"success_callback": [async_fn, "openmeter", sync_fn]}, {}), - ({"failure_callback": [async_fn, sync_fn]}, {}), - ({"input_callback": [async_fn, sync_fn]}, {}), - ({}, {"success_callback": [sync_fn, async_fn, "s3", "dynamodb"], "failure_callback": [sync_fn]}), - ({"callbacks": [SyncLogger()]}, {"callbacks": [SyncLogger()], "success_callback": [async_fn]}), - ], - ids=[ - "empty", - "global-custom-logger", - "global-async-callable", - "dynamic-callbacks", - "success-safety-net", - "failure-safety-net", - "input-safety-net", - "per-call-success-failure-split", - "mixed", - ], -) -def test_native_setup_registry_side_effects_match_function_setup( - clean_registries: None, globals_before: dict[str, list[object]], kwargs: dict[str, object] -) -> None: - for name, values in globals_before.items(): - getattr(litellm, name).extend(values) - legacy_logger, legacy_kwargs = run_legacy( - {key: list(value) if isinstance(value, list) else value for key, value in kwargs.items()} - ) - legacy_snapshot: Final = snapshot() - legacy_bootstrap: Final = list(utils.callback_list or []) - - for name in REGISTRIES: - getattr(litellm, name).clear() - utils.callback_list = [] - for name, values in globals_before.items(): - getattr(litellm, name).extend(values) - native_logger, native_kwargs = run_native( - {key: list(value) if isinstance(value, list) else value for key, value in kwargs.items()} - ) - - assert snapshot() == legacy_snapshot - assert sorted(map(repr, utils.callback_list or [])) == sorted(map(repr, legacy_bootstrap)) - assert set(native_kwargs) - {"litellm_call_id"} == set(legacy_kwargs) - {"litellm_call_id"} - for attribute in ( - "dynamic_success_callbacks", - "dynamic_async_success_callbacks", - "dynamic_failure_callbacks", - "dynamic_async_failure_callbacks", - "call_type", - "stream", - "model", - ): - assert getattr(native_logger, attribute) == getattr(legacy_logger, attribute), attribute - - -def test_native_setup_honours_caller_supplied_logging_object(clean_registries: None) -> None: - class Supplied(Logging): - pass - - supplied: Final = Supplied( - model="mistral/mistral-ocr-latest", - messages=[], - stream=False, - call_type="ocr", - start_time=datetime.datetime.now(), - litellm_call_id="supplied", - function_id="", - ) - litellm.callbacks.append(SyncLogger()) - logger, kwargs = run_native({"litellm_logging_obj": supplied, "callbacks": [SyncLogger()]}) - assert logger is supplied - assert kwargs["callbacks"] is not None - assert litellm.success_callback == [] - - -def test_native_setup_records_logger_fn_and_metadata(clean_registries: None) -> None: - def logger_fn(details: object) -> None: - del details - - logger, kwargs = run_native( - {"model": "mistral/mistral-ocr-latest", "logger_fn": logger_fn, "metadata": {"source": "test"}} - ) - assert utils.user_logger_fn is logger_fn - assert logger.litellm_params["metadata"] == {"source": "test"} - assert logger.model == "mistral/mistral-ocr-latest" - assert kwargs["metadata"] == {"source": "test"} diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index 9487550c80e..66e2bea0186 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -10,7 +10,6 @@ import litellm from litellm.integrations.custom_logger import CustomLogger from litellm.llms.base_llm.ocr.transformation import OCRResponse from tests.test_litellm_rust.support.callback_recorder import RecordingLogger -from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec from tests.test_litellm_rust.support.requests import ( OCR_DOCUMENT, OCR_RESPONSE, @@ -19,6 +18,7 @@ from tests.test_litellm_rust.support.requests import ( request_body, request_headers, ) +from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec pytestmark = pytest.mark.requires_rust_extension @@ -309,7 +309,6 @@ async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context asynchronous: bool, ) -> None: from contextvars import ContextVar - context: Final = ContextVar("azure-token-context", default="missing") context.set("caller") caller_thread: Final = threading.current_thread() @@ -338,7 +337,9 @@ async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context "callbacks": [Edit()], } response: Final = ( - await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) + await call_native_aocr(ocr_server, **arguments) + if asynchronous + else call_native_ocr(ocr_server, **arguments) ) assert response.pages[0].markdown == "native OCR response" assert observations == ["token", "pre_call"] @@ -367,7 +368,9 @@ async def test_native_azure_ocr_token_provider_can_make_nested_native_ocr_call( "azure_ad_token_provider": provider, } response: Final = ( - await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) + await call_native_aocr(ocr_server, **arguments) + if asynchronous + else call_native_ocr(ocr_server, **arguments) ) assert response.pages[0].markdown == "native OCR response" assert calls == ["token"] @@ -422,9 +425,7 @@ async def test_native_azure_ocr_releases_token_provider_after_terminal_outcome( ) -> None: import gc import weakref - from tests.test_litellm_rust.support.callback_recorder import drain_logging - class Provider: def __call__(self) -> str: if outcome == "failure": @@ -470,238 +471,3 @@ async def test_native_azure_ocr_releases_token_provider_after_terminal_outcome( await asyncio.sleep(0) gc.collect() assert reference() is None - - -FORBIDDEN_ORCHESTRATION: Final = ( - "pre_call", - "post_call", - "success_handler", - "async_success_handler", - "failure_handler", - "async_failure_handler", - "_success_handler_body", - "_async_success_handler_body", - "_failure_handler_body", - "_async_failure_handler_body", - "dispatch_success_handlers", - "dispatch_failure_handlers", - "handle_sync_success_callbacks_for_async_calls", -) - - -@pytest.fixture -def legacy_orchestration_disabled(monkeypatch: pytest.MonkeyPatch) -> list[str]: - from litellm import utils - from litellm.litellm_core_utils.litellm_logging import Logging - - reached: Final[list[str]] = [] - - def forbid(name: str): - def method(self, *args, **kwargs): - reached.append(name) - raise AssertionError(f"legacy orchestration reached: {name}") - - return method - - for name in FORBIDDEN_ORCHESTRATION: - monkeypatch.setattr(Logging, name, forbid(name)) - - def forbidden_setup(*args, **kwargs): - reached.append("function_setup") - raise AssertionError("legacy orchestration reached: function_setup") - - monkeypatch.setattr(utils, "function_setup", forbidden_setup) - return reached - - -@pytest.mark.asyncio -@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) -async def test_native_ocr_success_runs_integrations_without_legacy_orchestration( - ocr_server: RecordingServer, legacy_orchestration_disabled: list[str], asynchronous: bool -) -> None: - recorder: Final = RecordingLogger() - arguments: Final = {"callbacks": [recorder]} - response: Final = ( - await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) - ) - assert response.pages[0].markdown == "native OCR response" - success_event: Final = "async_log_success_event" if asynchronous else "log_success_event" - events: Final = await recorder.wait_for_async(success_event) - assert legacy_orchestration_disabled == [] - assert recorder.names.count("log_pre_api_call") == 1 - assert events[0].kwargs["standard_logging_object"]["status"] == "success" - assert events[0].kwargs["response_cost"] is not None - assert events[0].kwargs["litellm_params"]["api_base"] == f"{ocr_server.base_url}/v1/ocr" - - -@pytest.mark.asyncio -@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) -async def test_native_ocr_failure_runs_integrations_without_legacy_orchestration( - ocr_server: RecordingServer, legacy_orchestration_disabled: list[str], asynchronous: bool -) -> None: - ocr_server.enqueue(ResponseSpec(body={"message": "provider unavailable"}, status=500)) - recorder: Final = RecordingLogger() - arguments: Final = {"callbacks": [recorder]} - with pytest.raises(litellm.InternalServerError) as caught: - await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) - assert legacy_orchestration_disabled == [] - failures: Final = tuple(event for event in recorder.events if event.name.endswith("log_failure_event")) - assert [event.name for event in failures] == ( - ["log_failure_event", "async_log_failure_event"] if asynchronous else ["log_failure_event"] - ) - assert all(event.kwargs["exception"] is caught.value for event in failures) - assert all(event.kwargs["standard_logging_object"]["status"] == "failure" for event in failures) - assert "log_success_event" not in recorder.names - - -def test_native_ocr_success_hooks_all_run_before_any_success_dispatch(ocr_server: RecordingServer) -> None: - order: Final = [] - finished: Final = threading.Event() - - class Hooked(CustomLogger): - def __init__(self, name: str) -> None: - super().__init__() - self.name = name - - def logging_hook(self, kwargs, result, call_type): - order.append(("hook", self.name)) - kwargs[f"seen-by-{self.name}"] = True - return kwargs, result - - def log_success_event(self, kwargs, response_obj, start_time, end_time): - order.append(("log", self.name, kwargs.get("seen-by-a"), kwargs.get("seen-by-b"))) - if self.name == "b": - finished.set() - - call_native_ocr_with_callbacks(ocr_server, [Hooked("a"), Hooked("b")]) - - assert finished.wait(10) - assert order == [("hook", "a"), ("hook", "b"), ("log", "a", True, True), ("log", "b", True, True)] - - -def test_native_ocr_hook_failure_is_contained_and_target_still_dispatches(ocr_server: RecordingServer) -> None: - order: Final = [] - finished: Final = threading.Event() - - class Broken(CustomLogger): - def logging_hook(self, kwargs, result, call_type): - raise RuntimeError("hook failed") - - def log_success_event(self, kwargs, response_obj, start_time, end_time): - order.append("broken-log") - - class Healthy(CustomLogger): - def logging_hook(self, kwargs, result, call_type): - order.append("healthy-hook") - return kwargs, result - - def log_success_event(self, kwargs, response_obj, start_time, end_time): - order.append("healthy-log") - finished.set() - - call_native_ocr_with_callbacks(ocr_server, [Broken(), Healthy()]) - - assert finished.wait(10) - assert order == ["healthy-hook", "broken-log", "healthy-log"] - - -@pytest.mark.asyncio -async def test_native_aocr_hook_replacement_of_result_reaches_success_dispatch(ocr_server: RecordingServer) -> None: - replacement: Final = object() - observed: Final = [] - - class Replace(CustomLogger): - async def async_logging_hook(self, kwargs, result, call_type): - return kwargs, replacement - - class Observe(CustomLogger): - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - observed.append(response_obj) - - recorder: Final = RecordingLogger() - await call_native_aocr_with_callbacks(ocr_server, [Replace(), Observe(), recorder]) - await recorder.wait_for_async("async_log_success_event") - - assert observed == [replacement] - - -@pytest.mark.asyncio -async def test_native_aocr_shared_logging_object_dispatches_success_once(ocr_server: RecordingServer) -> None: - from litellm.litellm_core_utils.litellm_logging import Logging - - ocr_server.expected_requests = 2 - recorder: Final = RecordingLogger() - litellm.callbacks.append(recorder) - logger: Final = Logging( - model="mistral-ocr-latest", - messages=[], - stream=False, - call_type="aocr", - start_time=__import__("datetime").datetime.now(), - litellm_call_id="shared", - function_id="shared", - ) - logger.dynamic_async_success_callbacks = [recorder] - - await call_native_aocr(ocr_server, litellm_logging_obj=logger) - await call_native_aocr(ocr_server, litellm_logging_obj=logger) - await recorder.wait_for_async("async_log_success_event") - from tests.test_litellm_rust.support.callback_recorder import drain_logging - - await drain_logging() - - assert logger.model_call_details["has_logged_async_success"] is True - assert recorder.names.count("async_log_success_event") == 1 - - -def test_native_ocr_logging_preparation_failure_does_not_fail_request( - ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch -) -> None: - from litellm.litellm_core_utils import litellm_logging - - def broken_payload(*args, **kwargs): - raise RuntimeError("payload unavailable") - - monkeypatch.setattr(litellm_logging, "get_standard_logging_object_payload", broken_payload) - unraisable: Final = [] - monkeypatch.setattr(__import__("sys"), "unraisablehook", lambda event: unraisable.append(event.exc_value)) - recorder: Final = RecordingLogger() - - response: Final = call_native_ocr_with_callbacks(ocr_server, [recorder]) - events: Final = recorder.wait_for("log_success_event") - - assert response.pages[0].markdown == "native OCR response" - assert len(events) == 1 - assert events[0].thread is not threading.current_thread() - assert any(str(error) == "payload unavailable" for error in unraisable) - - -def test_native_ocr_writes_success_marker_and_honours_existing_marker( - ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch -) -> None: - from litellm.rust_bridge import setup as native_setup - - ocr_server.expected_requests = 2 - loggers: Final = [] - original_build: Final = native_setup.build_logging - - def build_logging(**kwargs): - logger = original_build(**kwargs) - if loggers: - logger.model_call_details["has_logged_sync_success"] = True - loggers.append(logger) - return logger - - monkeypatch.setattr(native_setup, "build_logging", build_logging) - recorder: Final = RecordingLogger() - - call_native_ocr_with_callbacks(ocr_server, [recorder]) - recorder.wait_for("log_success_event") - assert loggers[0].model_call_details["has_logged_sync_success"] is True - - call_native_ocr_with_callbacks(ocr_server, [recorder]) - from litellm.litellm_core_utils.thread_pool_executor import executor - - executor.submit(lambda: None).result(10) - assert recorder.names.count("log_success_event") == 1 - assert recorder.names.count("logging_hook") == 1 diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 04cc6eb41fc..2ca9e77db4f 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -967,18 +967,26 @@ async def test_terminal_registration_added_during_http_is_observed( @pytest.fixture def created_loggers(monkeypatch: pytest.MonkeyPatch) -> list[Logging]: - from litellm.rust_bridge import setup as native_setup + from litellm import utils - original_build: Final = native_setup.build_logging + original_setup: Final = utils.function_setup loggers: Final[list[Logging]] = [] - def build_logging(**kwargs: object) -> Logging: - logger: Final = original_build(**kwargs) # pyright: ignore[reportArgumentType] # passthrough of the factory signature + def setup( + call_type: str, + rules: utils.Rules, + start: datetime.datetime, + *args: object, + is_async_call: bool = True, + **kwargs: object, + ) -> tuple[Logging, dict[str, object]]: + logger, prepared = original_setup(call_type, rules, start, *args, is_async_call=is_async_call, **kwargs) + assert isinstance(logger, Logging) setattr(logger, "_defer_async_logging", True) loggers.append(logger) - return logger + return logger, prepared - monkeypatch.setattr(native_setup, "build_logging", build_logging) + monkeypatch.setattr(utils, "function_setup", setup) return loggers