diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs index 6f95444fc8c..3850b827093 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs @@ -28,102 +28,10 @@ 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) fn finalize( @@ -185,59 +93,3 @@ impl DeploymentHooks { .map(Bound::unbind) } } - -#[cfg(test)] -mod tests { - use super::*; - - #[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 new file mode 100644 index 00000000000..ce8d508a552 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/lifecycle/compat.rs @@ -0,0 +1,168 @@ +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 new file mode 100644 index 00000000000..5fae1a01d7f --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs @@ -0,0 +1,649 @@ +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)) +} diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs index 5aede7baa13..ca247138083 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs @@ -1,26 +1,30 @@ 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 crate::execution::{poll_async_value, run_async_value, run_sync_value}; +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, +}; 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; @@ -300,6 +304,7 @@ pub(crate) struct PythonCallState { pub error: Option>, pub asynchronous: bool, pub internal: bool, + pub supplied: bool, pub call_type: &'static str, } @@ -389,6 +394,7 @@ impl PythonCallState { error: None, asynchronous, internal: false, + supplied: false, call_type, }) } @@ -412,6 +418,7 @@ impl PythonCallState { )?; self.logger = Some(result.logger); self.kwargs = result.kwargs; + self.supplied = result.supplied; Ok(()) } @@ -431,6 +438,16 @@ 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) => { @@ -441,40 +458,71 @@ 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()?; - 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)), + 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), }; - 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) { + 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) => { logger.defer_success( py, Py::new( py, - PendingLogging { - pending: Some(pending()), + DeferredSuccess { + runner: Some(runner), }, )?, )?; - } 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( @@ -482,19 +530,39 @@ impl PythonCallState { py: Python<'_>, asynchronous: bool, ) -> PyResult>> { - if self.logger.is_none() || (self.asynchronous && self.internal) { + if self.logger.is_none() || self.error.is_none() { return Ok(None); } - let Some(error) = &self.error else { + let phase = if asynchronous { + HostPhase::AsyncFailure + } else { + HostPhase::Failure + }; + let Some(family) = plan_failure(phase, self.asynchronous, self.internal) else { return Ok(None); }; - self.logger()? - .failure(py, error, &self.start, &self.end, asynchronous) + 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()), + } } pub fn cleanup(&mut self, py: Python<'_>) { if let Some(logger) = self.logger.take() - && let Err(error) = logger.restore_context(py) + && let Err(error) = dispatch::leaves(py).and_then(|leaves| { + leaves + .getattr("restore_correlation_context")? + .call1((logger.object(py),)) + .map(|_| ()) + }) { error.write_unraisable(py, None); } @@ -530,58 +598,53 @@ impl PythonCallState { } } -struct PendingSuccess { - logger: PythonLogger, - response: Option>, - start: Py, - end: Option>, -} - -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, +struct DeferredSuccess { + runner: 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(|_| ()) + } } #[pymethods] -impl PendingLogging { +impl DeferredSuccess { fn __call__(slf: &Bound<'_, Self>, py: Python<'_>) -> PyResult<()> { - 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, + 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(()) } + result => result, } - Ok(()) } fn __traverse__(&self, visit: pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> { - if let Some(pending) = &self.pending { - pending.logger.traverse(&visit)?; - visit.call(&pending.response)?; - visit.call(&pending.start)?; - visit.call(&pending.end)?; + match &self.runner { + Some(runner) => runner.traverse(&visit), + None => Ok(()), } - Ok(()) } fn close(slf: &Bound<'_, Self>) { - let pending = slf.borrow_mut().pending.take(); - drop(pending); + let runner = slf.borrow_mut().runner.take(); + drop(runner); } fn __clear__(slf: &Bound<'_, Self>) { @@ -592,14 +655,39 @@ impl PendingLogging { #[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 install_logging_worker(py: Python<'_>, worker: &Bound<'_, PyAny>) -> PyResult<()> { - py.import("litellm.litellm_core_utils.logging_worker")? - .setattr("GLOBAL_LOGGING_WORKER", worker) + 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() } struct RetainingHost { @@ -771,17 +859,7 @@ mod tests { .unwrap_or_else(|error| error.into_inner()); Python::initialize(); Python::attach(|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(); + load_lifecycle_module(py); let route = SyntheticRoute( PythonCallState::new( py, @@ -817,17 +895,7 @@ mod tests { Python::initialize(); Python::attach(|py| { py.import("asyncio").unwrap(); - 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 module = load_lifecycle_module(py); let locals = PyDict::new(py); locals .set_item("drive", module.getattr("drive").unwrap()) @@ -935,64 +1003,11 @@ 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(); @@ -1008,201 +1023,6 @@ sys.unraisablehook = old_hook }); } - #[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 index db90a001c52..c7370be628f 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/setup.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/setup.rs @@ -235,6 +235,7 @@ fn split_dynamic<'py>( pub(super) struct Setup { pub logger: PythonLogger, pub kwargs: Py, + pub supplied: bool, } pub(super) fn setup( @@ -259,6 +260,7 @@ pub(super) fn setup( return Ok(Setup { logger: supplied.extract()?, kwargs: kwargs.unbind(), + supplied: true, }); } } @@ -330,5 +332,6 @@ pub(super) fn setup( 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 deb0a4e7f4c..217f69f592d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -6,8 +6,10 @@ 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::PythonLogger; +use crate::lifecycle::{PythonCallState, PythonLogger}; pub(super) fn update_logging( py: Python<'_>, @@ -32,37 +34,63 @@ pub(super) fn update_logging( pub(super) fn pre_call( py: Python<'_>, - logger: &PythonLogger, + state: &PythonCallState, request: &OcrDuringCallRequest, payload: &PythonPayload, ) -> PyResult<()> { - 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(()) + 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) } pub(super) fn post_call( py: Python<'_>, - logger: &PythonLogger, + state: &PythonCallState, original_response: &Value, payload: &PythonPayload, ) -> PyResult<()> { - py.import("litellm.rust_bridge.ocr")? - .getattr("post_call")? - .call1(( - logger.object(py), - to_py(py, original_response)?, - &payload.body, - &payload.headers, - ))?; - Ok(()) + 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) } 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 72971bce2d4..9fcd53569bc 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -128,7 +128,7 @@ impl PythonOcrHost { &retained.secret_fields, )?; let payload = PythonPayload::from_request(py, &request)?; - callbacks::pre_call(py, logger, &request, &payload)?; + callbacks::pre_call(py, &self.state, &request, &payload)?; let request = payload.write_back(py, request)?; self.retained_mut()?.payload = Some(payload); Ok(request) @@ -144,12 +144,7 @@ impl PythonOcrHost { .payload .as_ref() .ok_or_else(missing_state)?; - callbacks::post_call( - py, - self.state.logger()?, - &request.original_response, - payload, - )?; + callbacks::post_call(py, &self.state, &request.original_response, payload)?; Ok(request) } diff --git a/litellm/rust_bridge/leaves.py b/litellm/rust_bridge/leaves.py new file mode 100644 index 00000000000..64e239d1615 --- /dev/null +++ b/litellm/rust_bridge/leaves.py @@ -0,0 +1,815 @@ +"""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/ocr.py b/litellm/rust_bridge/ocr.py index fc3c5523856..04f842c8dcc 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -22,10 +22,6 @@ 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 { @@ -62,32 +58,6 @@ 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/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index 96983968bee..272974b2689 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -8,7 +8,8 @@ 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.ocr import NATIVE_AOCR, NATIVE_OCR, post_call, pre_call, update_logging +from litellm.rust_bridge.leaves import record_post_call, record_pre_call +from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, update_logging @pytest.fixture(autouse=True) @@ -65,25 +66,26 @@ 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() + logger: Final = Mock(log_raw_request_response=False, logger_fn=None, model_call_details={}) body: Final[dict[str, object]] = {"document": "original"} headers: Final = {"authorization": "key"} response: Final = object() - pre_call(logger, "key", body, headers, "https://provider") - post_call(logger, response, body, headers) - logger.pre_call.assert_called_once_with( + 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( 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.post_call): + for callback in (logger._pre_call, logger.record_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.post_call.call_args.kwargs["original_response"] is response + assert logger.record_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: - pre_call(failing_logger, None, body, headers, "https://provider") + record_pre_call(failing_logger, api_key=None, body=body, headers=headers, url="https://provider") assert caught.value is failure diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index 66e2bea0186..9487550c80e 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -10,6 +10,7 @@ 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, @@ -18,7 +19,6 @@ 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,6 +309,7 @@ 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() @@ -337,9 +338,7 @@ 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"] @@ -368,9 +367,7 @@ 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"] @@ -425,7 +422,9 @@ 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": @@ -471,3 +470,238 @@ 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