Revert "refactor(python-bridge): classify callbacks natively instead of function_setup"

This reverts commit f8190bbe80.
This commit is contained in:
Yujong Lee 2026-09-16 19:07:23 -07:00
parent 58be89ffba
commit 92e1a2ded7
21 changed files with 763 additions and 3476 deletions

View file

@ -3,7 +3,6 @@ use std::time::{Instant, SystemTime, UNIX_EPOCH};
pub mod callbacks;
pub mod host;
pub mod registration;
pub mod types;
pub use callbacks::{

View file

@ -1,374 +0,0 @@
use std::collections::HashSet;
use super::callbacks::CallbackId;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum Registry {
Input,
AsyncInput,
Success,
AsyncSuccess,
Failure,
AsyncFailure,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Registration {
Named { known: bool, async_only: bool },
Object { asynchronous: bool },
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Entry {
pub id: CallbackId,
pub registration: Registration,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Candidate {
pub resolved: Option<Entry>,
pub duplicate_type: bool,
}
#[derive(Clone, Debug, PartialEq, Eq, Default)]
pub struct RegistrationFacts {
pub candidates: Vec<Candidate>,
pub input: Vec<Entry>,
pub success: Vec<Entry>,
pub failure: Vec<Entry>,
pub async_success: Vec<CallbackId>,
pub async_failure: Vec<CallbackId>,
pub bootstrap_pending: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum NamedEvent {
Success,
Failure,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RegistryMutation {
Append(Registry, CallbackId),
Remove(Registry, CallbackId),
ExpandNamed(NamedEvent, CallbackId),
Bootstrap,
}
struct Planner {
input: Vec<Entry>,
success: Vec<Entry>,
failure: Vec<Entry>,
async_success: HashSet<CallbackId>,
async_failure: HashSet<CallbackId>,
mutations: Vec<RegistryMutation>,
}
impl Planner {
fn contains(&self, registry: Registry, id: CallbackId) -> bool {
match registry {
Registry::Input => self.input.iter().any(|entry| entry.id == id),
Registry::Success => self.success.iter().any(|entry| entry.id == id),
Registry::Failure => self.failure.iter().any(|entry| entry.id == id),
Registry::AsyncSuccess => self.async_success.contains(&id),
Registry::AsyncFailure => self.async_failure.contains(&id),
Registry::AsyncInput => false,
}
}
fn append(&mut self, registry: Registry, entry: Entry) {
if self.contains(registry, entry.id) {
return;
}
match registry {
Registry::Input => self.input.push(entry),
Registry::Success => self.success.push(entry),
Registry::Failure => self.failure.push(entry),
Registry::AsyncSuccess => {
self.async_success.insert(entry.id);
}
Registry::AsyncFailure => {
self.async_failure.insert(entry.id);
}
Registry::AsyncInput => {}
}
self.mutations
.push(RegistryMutation::Append(registry, entry.id));
}
fn record(&mut self, mutation: RegistryMutation) {
self.mutations.push(mutation);
}
}
fn is_asynchronous(registration: Registration) -> bool {
matches!(registration, Registration::Object { asynchronous: true })
}
pub fn plan_registration(facts: &RegistrationFacts) -> Vec<RegistryMutation> {
let mut planner = Planner {
input: facts.input.clone(),
success: facts.success.clone(),
failure: facts.failure.clone(),
async_success: facts.async_success.iter().copied().collect(),
async_failure: facts.async_failure.iter().copied().collect(),
mutations: Vec::new(),
};
for candidate in &facts.candidates {
let Some(entry) = candidate.resolved else {
continue;
};
if candidate.duplicate_type {
continue;
}
planner.append(Registry::Input, entry);
if !is_asynchronous(entry.registration) {
planner.append(Registry::Success, entry);
planner.append(Registry::Failure, entry);
}
planner.append(Registry::AsyncSuccess, entry);
planner.append(Registry::AsyncFailure, entry);
}
if facts.bootstrap_pending
&& !(planner.input.is_empty() && planner.success.is_empty() && planner.failure.is_empty())
{
planner.record(RegistryMutation::Bootstrap);
}
let input = planner.input.clone();
for entry in input
.iter()
.filter(|entry| is_asynchronous(entry.registration))
{
planner.record(RegistryMutation::Append(Registry::AsyncInput, entry.id));
planner.record(RegistryMutation::Remove(Registry::Input, entry.id));
}
let success = planner.success.clone();
for entry in &success {
match entry.registration {
Registration::Object { asynchronous: true }
| Registration::Named {
async_only: true, ..
} => {
planner.append(Registry::AsyncSuccess, *entry);
planner.record(RegistryMutation::Remove(Registry::Success, entry.id));
}
Registration::Named { known: true, .. } => {
planner.record(RegistryMutation::ExpandNamed(NamedEvent::Success, entry.id));
}
_ => {}
}
}
let failure = planner.failure.clone();
for entry in &failure {
match entry.registration {
Registration::Object { asynchronous: true } => {
planner.append(Registry::AsyncFailure, *entry);
planner.record(RegistryMutation::Remove(Registry::Failure, entry.id));
}
Registration::Named { known: true, .. } => {
planner.record(RegistryMutation::ExpandNamed(NamedEvent::Failure, entry.id));
}
_ => {}
}
}
planner.mutations
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DynamicSuccessSlot {
Sync,
Async,
}
pub fn classify_dynamic_success(entry: Entry, named_async: bool) -> DynamicSuccessSlot {
match entry.registration {
Registration::Object { asynchronous: true } => DynamicSuccessSlot::Async,
Registration::Named { .. } if named_async => DynamicSuccessSlot::Async,
_ => DynamicSuccessSlot::Sync,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn object(id: u64, asynchronous: bool) -> Entry {
Entry {
id: CallbackId(id),
registration: Registration::Object { asynchronous },
}
}
fn named(id: u64, known: bool, async_only: bool) -> Entry {
Entry {
id: CallbackId(id),
registration: Registration::Named { known, async_only },
}
}
fn candidate(entry: Entry) -> Candidate {
Candidate {
resolved: Some(entry),
duplicate_type: false,
}
}
#[test]
fn sync_callback_in_callbacks_registers_in_every_list_once() {
let facts = RegistrationFacts {
candidates: vec![candidate(object(1, false)), candidate(object(1, false))],
..RegistrationFacts::default()
};
assert_eq!(
plan_registration(&facts),
[
RegistryMutation::Append(Registry::Input, CallbackId(1)),
RegistryMutation::Append(Registry::Success, CallbackId(1)),
RegistryMutation::Append(Registry::Failure, CallbackId(1)),
RegistryMutation::Append(Registry::AsyncSuccess, CallbackId(1)),
RegistryMutation::Append(Registry::AsyncFailure, CallbackId(1)),
]
);
}
#[test]
fn async_callable_in_callbacks_skips_sync_lists_and_moves_out_of_input() {
let facts = RegistrationFacts {
candidates: vec![candidate(object(2, true))],
..RegistrationFacts::default()
};
assert_eq!(
plan_registration(&facts),
[
RegistryMutation::Append(Registry::Input, CallbackId(2)),
RegistryMutation::Append(Registry::AsyncSuccess, CallbackId(2)),
RegistryMutation::Append(Registry::AsyncFailure, CallbackId(2)),
RegistryMutation::Append(Registry::AsyncInput, CallbackId(2)),
RegistryMutation::Remove(Registry::Input, CallbackId(2)),
]
);
}
#[test]
fn unresolved_and_duplicate_type_named_candidates_are_skipped() {
let facts = RegistrationFacts {
candidates: vec![
Candidate {
resolved: None,
duplicate_type: false,
},
Candidate {
resolved: Some(object(3, false)),
duplicate_type: true,
},
],
..RegistrationFacts::default()
};
assert!(plan_registration(&facts).is_empty());
}
#[test]
fn already_registered_callbacks_are_not_appended_again() {
let facts = RegistrationFacts {
candidates: vec![candidate(object(1, false))],
input: vec![object(1, false)],
success: vec![object(1, false)],
failure: vec![object(1, false)],
async_success: vec![CallbackId(1)],
async_failure: vec![CallbackId(1)],
..RegistrationFacts::default()
};
assert!(plan_registration(&facts).is_empty());
}
#[test]
fn bootstrap_runs_once_when_any_public_list_is_populated() {
let empty = RegistrationFacts {
bootstrap_pending: true,
..RegistrationFacts::default()
};
assert!(plan_registration(&empty).is_empty());
let populated = RegistrationFacts {
bootstrap_pending: true,
candidates: vec![candidate(object(1, false))],
..RegistrationFacts::default()
};
assert!(plan_registration(&populated).contains(&RegistryMutation::Bootstrap));
let already = RegistrationFacts {
bootstrap_pending: false,
success: vec![object(1, false)],
..RegistrationFacts::default()
};
assert!(!plan_registration(&already).contains(&RegistryMutation::Bootstrap));
}
#[test]
fn success_safety_net_moves_async_and_async_only_names_and_expands_known_names() {
let facts = RegistrationFacts {
success: vec![
object(1, true),
named(2, false, true),
named(3, true, false),
named(4, false, false),
object(5, false),
],
..RegistrationFacts::default()
};
assert_eq!(
plan_registration(&facts),
[
RegistryMutation::Append(Registry::AsyncSuccess, CallbackId(1)),
RegistryMutation::Remove(Registry::Success, CallbackId(1)),
RegistryMutation::Append(Registry::AsyncSuccess, CallbackId(2)),
RegistryMutation::Remove(Registry::Success, CallbackId(2)),
RegistryMutation::ExpandNamed(NamedEvent::Success, CallbackId(3)),
]
);
}
#[test]
fn failure_safety_net_ignores_async_only_names() {
let facts = RegistrationFacts {
failure: vec![
object(1, true),
named(2, false, true),
named(3, true, false),
],
async_failure: vec![CallbackId(1)],
..RegistrationFacts::default()
};
assert_eq!(
plan_registration(&facts),
[
RegistryMutation::Remove(Registry::Failure, CallbackId(1)),
RegistryMutation::ExpandNamed(NamedEvent::Failure, CallbackId(3)),
]
);
}
#[test]
fn dynamic_success_split_follows_async_callables_and_selected_names() {
assert_eq!(
classify_dynamic_success(object(1, true), false),
DynamicSuccessSlot::Async
);
assert_eq!(
classify_dynamic_success(object(1, false), false),
DynamicSuccessSlot::Sync
);
assert_eq!(
classify_dynamic_success(named(2, true, false), true),
DynamicSuccessSlot::Async
);
assert_eq!(
classify_dynamic_success(named(2, true, true), false),
DynamicSuccessSlot::Sync
);
}
}

View file

@ -1,6 +1,6 @@
use litellm_python_interop::release_count;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use pyo3::types::PyDict;
#[pyfunction]
fn gil_stats(py: Python<'_>) -> PyResult<Py<PyAny>> {
@ -9,27 +9,6 @@ fn gil_stats(py: Python<'_>) -> PyResult<Py<PyAny>> {
Ok(stats.into_any().unbind())
}
#[pyfunction]
fn _debug_setup(
py: Python<'_>,
call_type: String,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
start: Py<PyAny>,
asynchronous: bool,
) -> PyResult<(Py<PyAny>, Py<PyDict>)> {
let leaked: &'static str = Box::leak(call_type.into_boxed_str());
let result = crate::lifecycle::debug_setup(
py,
leaked,
&args.unbind(),
&kwargs.unbind(),
&start,
asynchronous,
)?;
Ok(result)
}
#[cfg(feature = "panic-test")]
#[pyfunction]
fn _panic_for_test() {
@ -38,7 +17,6 @@ fn _panic_for_test() {
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(gil_stats, module)?)?;
module.add_function(wrap_pyfunction!(_debug_setup, module)?)?;
#[cfg(feature = "panic-test")]
module.add_function(wrap_pyfunction!(_panic_for_test, module)?)?;
Ok(())

View file

@ -1,7 +1,7 @@
use pyo3::exceptions::PyBaseException;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::types::PyDict;
use pyo3::types::{PyDict, PyTuple};
#[derive(FromPyObject)]
pub(crate) struct PythonLogger(Py<PyAny>);
@ -28,10 +28,128 @@ impl PythonLogger {
pub(super) fn defer_success(
&self,
py: Python<'_>,
pending: Py<super::DeferredSuccess>,
pending: Py<super::PendingLogging>,
) -> PyResult<()> {
self.object(py).setattr("_native_pending_logging", pending)
}
pub(super) fn sync_success_for_async_call(
&self,
py: Python<'_>,
response: &Option<Py<PyAny>>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> 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<PyBaseException>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
asynchronous: bool,
) -> PyResult<Option<Py<PyAny>>> {
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<Py<PyAny>>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> 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<Py<PyAny>>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
let context = py.import("contextvars")?.call_method0("copy_context")?;
let worker = py
.import("litellm.litellm_core_utils.logging_worker")?
.getattr("GLOBAL_LOGGING_WORKER")?
.getattr("ensure_initialized_and_enqueue")?;
let coroutine = self
.object(py)
.call_method1("async_success_handler", (response, start, end))?;
let enqueue = context.call_method1("run", (worker, &coroutine));
if enqueue.is_err()
&& let Err(error) = coroutine.call_method0("close")
{
error.write_unraisable(py, Some(&coroutine));
}
enqueue.map(|_| ())
}
}
pub(super) struct SetupResult<'py>(Bound<'py, PyAny>);
impl SetupResult<'_> {
pub(super) fn logger(&self) -> PyResult<PythonLogger> {
self.0.getattr("logger")?.extract()
}
pub(super) fn kwargs(&self) -> PyResult<Py<PyDict>> {
Ok(self.0.getattr("kwargs")?.extract()?)
}
}
pub(super) fn setup<'py>(
py: Python<'py>,
call_type: &str,
args: &Py<PyTuple>,
kwargs: &Py<PyDict>,
start: &Py<PyAny>,
asynchronous: bool,
) -> PyResult<SetupResult<'py>> {
py.import("litellm.rust_bridge.lifecycle")?
.getattr("setup")?
.call1((call_type, args, kwargs, start, asynchronous))
.map(SetupResult)
}
pub(super) fn finalize(
@ -42,11 +160,9 @@ pub(super) fn finalize(
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
let model = kwargs.bind(py).get_item("model")?;
let model = model.filter(|value| value.is_instance_of::<pyo3::types::PyString>());
py.import("litellm.litellm_core_utils.llm_response_utils.response_metadata")?
.getattr("update_response_metadata")?
.call1((response, logger.object(py), model, kwargs, start, end))?;
py.import("litellm.rust_bridge.lifecycle")?
.getattr("finalize")?
.call1((response, logger.object(py), kwargs, start, end))?;
Ok(())
}
@ -95,3 +211,124 @@ impl DeploymentHooks {
.map(Bound::unbind)
}
}
#[cfg(test)]
mod tests {
use super::*;
use pyo3::exceptions::PyTypeError;
#[test]
fn setup_fields_are_checked_in_order_without_eager_logger_method_reads() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
reads = []
class Logger:
def __getattribute__(self, name):
reads.append(name)
raise AssertionError('logger methods must remain lazy')
logger = Logger()
class Setup:
@property
def logger(self):
reads.append('logger')
return logger
@property
def kwargs(self):
reads.append('kwargs')
return []
result = Setup()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
let result = SetupResult(locals.get_item("result").unwrap().unwrap());
let logger = result.logger().unwrap();
assert!(
logger
.object(py)
.is(locals.get_item("logger").unwrap().unwrap())
);
assert_eq!(
locals
.get_item("reads")
.unwrap()
.unwrap()
.extract::<Vec<String>>()
.unwrap(),
["logger"]
);
assert!(
result
.kwargs()
.unwrap_err()
.is_instance_of::<PyTypeError>(py)
);
assert_eq!(
locals
.get_item("reads")
.unwrap()
.unwrap()
.extract::<Vec<String>>()
.unwrap(),
["logger", "kwargs"]
);
});
}
#[test]
fn logger_resolves_each_callback_at_invocation_and_preserves_arguments() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
calls = []
response, start, end = object(), object(), object()
class Logger:
@property
def handle_sync_success_callbacks_for_async_calls(self):
generation = len(calls)
def callback(*args):
assert args == (response, start, end)
calls.append(generation)
return callback
logger = Logger()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
let logger: PythonLogger = locals
.get_item("logger")
.unwrap()
.unwrap()
.extract()
.unwrap();
let response = Some(locals.get_item("response").unwrap().unwrap().unbind());
let start = locals.get_item("start").unwrap().unwrap().unbind();
let end = Some(locals.get_item("end").unwrap().unwrap().unbind());
for _ in 0..2 {
logger
.sync_success_for_async_call(py, &response, &start, &end)
.unwrap();
}
assert_eq!(
locals
.get_item("calls")
.unwrap()
.unwrap()
.extract::<Vec<usize>>()
.unwrap(),
[0, 1]
);
});
}
}

View file

@ -1,168 +0,0 @@
use litellm_core::call_lifecycle::{
CallbackFamily, Delivery, ReleaseGate, SuccessFacts, plan_success,
};
use pyo3::prelude::*;
use super::bindings::PythonLogger;
use super::{PythonCallState, dispatch, missing_state};
pub(super) fn dispatch_success(
py: Python<'_>,
state: &PythonCallState,
logger: &PythonLogger,
) -> PyResult<()> {
let facts = SuccessFacts {
asynchronous: state.asynchronous,
internal: state.internal,
fallbacks: !state
.kwargs
.bind(py)
.get_item("fallbacks")?
.is_none_or(|value| value.is_none()),
deferred: logger.defers_async_logging(py),
sync_target_kinds: sync_kinds(py, logger)?,
};
let leaves = dispatch::leaves(py)?;
let object = logger.object(py);
for selected in plan_success(&facts) {
match (selected.family, selected.delivery, selected.gate) {
(CallbackFamily::SyncSuccess, Delivery::Worker, _) => {
let handler = object.getattr("success_handler")?;
let bound = py.import("functools")?.getattr("partial")?.call1((
handler,
&state.response,
&state.start,
&state.end,
))?;
leaves.getattr("submit_worker")?.call1((bound,))?;
}
(CallbackFamily::AsyncSuccess, Delivery::Background, ReleaseGate::Immediate) => {
let coroutine = object.call_method1(
"async_success_handler",
(&state.response, &state.start, &state.end),
)?;
let enqueue = leaves.getattr("enqueue_background")?.call1((&coroutine,));
if enqueue.is_err()
&& let Err(error) = coroutine.call_method0("close")
{
error.write_unraisable(py, Some(&coroutine));
}
enqueue?;
}
(CallbackFamily::AsyncSuccess, Delivery::Background, ReleaseGate::Deferred) => {
let pending = Py::new(
py,
SuppliedDeferred {
logger: Some(logger.clone_ref(py)),
response: state.response.as_ref().map(|value| value.clone_ref(py)),
start: state.start.clone_ref(py),
end: state.end.as_ref().map(|value| value.clone_ref(py)),
},
)?;
object.setattr("_native_pending_logging", pending)?;
}
_ => return Err(missing_state()),
}
}
Ok(())
}
fn sync_kinds(
py: Python<'_>,
logger: &PythonLogger,
) -> PyResult<Vec<litellm_core::call_lifecycle::CallbackKind>> {
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<Option<Py<PyAny>>> {
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<PythonLogger>,
response: Option<Py<PyAny>>,
start: Py<PyAny>,
end: Option<Py<PyAny>>,
}
#[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::<pyo3::exceptions::PyException>(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);
}
}

View file

@ -1,771 +0,0 @@
use litellm_core::call_lifecycle::{
CallbackFamily, CallbackId, CallbackInvocation, CallbackKind, CallbackMethod, Delivery,
DispatchCursor, DispatchFacts, DispatchStep, InvocationOutcome, LoggedMarker,
};
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError};
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::types::PyString;
use super::bindings::PythonLogger;
const LEAVES: &str = "litellm.rust_bridge.leaves";
pub(super) fn leaves(py: Python<'_>) -> PyResult<Bound<'_, PyModule>> {
py.import(LEAVES)
}
#[derive(Clone, Copy)]
pub(super) enum Outcome {
Success,
Failure,
}
pub(super) struct Targets {
objects: Vec<Py<PyAny>>,
kinds: Vec<CallbackKind>,
}
impl Targets {
pub(super) fn read(
py: Python<'_>,
lists: &[Bound<'_, PyAny>],
) -> PyResult<(Self, Vec<Vec<CallbackId>>)> {
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<CallbackId> {
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::<PyString>() {
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<CallbackKind> {
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<bool> {
let name = if let Ok(name) = object.getattr("__name__") {
name.extract::<String>()?
} else if let Ok(func) = object.getattr("__func__") {
func.getattr("__name__")?.extract::<String>()?
} 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<CallbackId>,
pub family: CallbackFamily,
pub response: Option<Py<PyAny>>,
pub error: Option<Py<PyBaseException>>,
pub start: Py<PyAny>,
pub end: Py<PyAny>,
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::<bool>())
}
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::<bool>())
}
_ => 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<PyAny>),
Done,
}
pub(super) struct Runner {
job: Job,
cursor: DispatchCursor,
result: Option<Py<PyAny>>,
formatted: Option<Py<PyAny>>,
pending: Option<CallbackInvocation>,
}
impl Runner {
pub(super) fn start(py: Python<'_>, job: Job) -> PyResult<Self> {
let leaves = leaves(py)?;
let already = match job.family.marker() {
Some(marker) => leaves
.getattr("already_logged")?
.call1((job.logger.object(py), marker.key()))?
.extract::<bool>()?,
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<Py<PyAny>>>,
) -> PyResult<Step> {
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::<PyException>(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<Option<Py<PyAny>>> {
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<Option<Py<PyAny>>>,
) -> 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::<PyException>(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<Runner>,
}
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::<PyException>(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<Runner>,
}
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<Py<PyAny>>>,
) -> PyResult<super::handle::ExecutionStep> {
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::<PyException>(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<Py<PyAny>> {
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::<PyException>(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<Bound<'py, PyAny>>)> {
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<CallbackId>)> {
let (global, dynamic) = read_lists(py, logger, family)?;
let lists: Vec<Bound<'_, PyAny>> = std::iter::once(global).chain(dynamic.clone()).collect();
let (targets, ids) = Targets::read(py, &lists)?;
let global_ids = ids.first().cloned().unwrap_or_default();
let dynamic_ids = ids.get(1).map(Vec::as_slice);
let ordered = family.targets(
&global_ids,
dynamic.is_some().then_some(dynamic_ids.unwrap_or(&[])),
);
Ok((targets, ordered))
}
#[cfg(test)]
mod tests {
use super::*;
use pyo3::types::{PyDict, PyList};
fn fixture(py: Python<'_>) -> Bound<'_, PyDict> {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
import sys, types
for name in ("litellm", "litellm.integrations", "litellm.integrations.custom_logger"):
sys.modules.setdefault(name, types.ModuleType(name))
class CustomLogger: pass
sys.modules["litellm.integrations.custom_logger"].CustomLogger = CustomLogger
sys.modules["litellm"]._known_custom_logger_compatible_callbacks = []
class Logger: pass
logger = Logger()
target = CustomLogger()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
locals
}
fn runner(py: Python<'_>, locals: &Bound<'_, PyDict>, family: CallbackFamily) -> Runner {
let logger = locals.get_item("logger").unwrap().unwrap();
let target = locals.get_item("target").unwrap().unwrap();
let list = PyList::new(py, [&target]).unwrap().into_any();
let (targets, ids) = Targets::read(py, &[list]).unwrap();
let ids: Vec<CallbackId> = ids.into_iter().flatten().collect();
Runner {
cursor: DispatchCursor::start(family, ids.clone(), false, false),
job: Job {
logger: logger.extract().unwrap(),
targets,
ids,
family,
response: Some(target.clone().unbind()),
error: None,
start: py.None(),
end: py.None(),
stream: false,
},
result: Some(target.unbind()),
formatted: None,
pending: None,
}
}
fn assert_collectable(py: Python<'_>, locals: &Bound<'_, PyDict>, handle: Py<PyAny>) {
locals.set_item("handle", handle).unwrap();
py.run(
pyo3::ffi::c_str!(
r#"
import gc
import weakref
target.handle = handle
reference = weakref.ref(target)
del logger, target, handle
gc.collect()
assert reference() is None, "cycle through the retained target was not collected"
"#
),
Some(locals),
Some(locals),
)
.unwrap();
}
#[test]
fn worker_job_collects_cycles_through_logger_targets_and_response() {
Python::initialize();
Python::attach(|py| {
let locals = fixture(py);
let runner = runner(py, &locals, CallbackFamily::SyncSuccess);
let handle = Py::new(py, WorkerJob::new(runner)).unwrap().into_any();
assert_collectable(py, &locals, handle);
});
}
#[test]
fn deferred_success_collects_cycles_and_close_is_idempotent() {
Python::initialize();
Python::attach(|py| {
let locals = fixture(py);
let runner = runner(py, &locals, CallbackFamily::AsyncSuccess);
let deferred = Py::new(
py,
super::super::DeferredSuccess {
runner: Some(runner),
},
)
.unwrap();
assert_collectable(py, &locals, deferred.into_any());
});
}
#[test]
fn deferred_success_releases_at_most_once_and_close_prevents_release() {
Python::initialize();
Python::attach(|py| {
let locals = fixture(py);
let runner = runner(py, &locals, CallbackFamily::AsyncSuccess);
let deferred = Py::new(
py,
super::super::DeferredSuccess {
runner: Some(runner),
},
)
.unwrap();
deferred.call_method0(py, "close").unwrap();
deferred.call_method0(py, "close").unwrap();
deferred.call0(py).unwrap();
assert!(deferred.borrow(py).runner.is_none());
});
}
}

View file

@ -1,30 +1,25 @@
use std::sync::Arc;
use std::task::Poll;
use futures_util::future::{AbortHandle, Abortable};
#[cfg(test)]
use litellm_core::call_lifecycle::host::HostCallFuture;
use litellm_core::call_lifecycle::host::{
HostCall as NativeCall, HostCallStep as NativeCallStep, HostFailure, HostPhase, HostStep,
};
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError};
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use futures_util::future::{AbortHandle, Abortable};
use tokio::sync::Mutex;
use litellm_core::call_lifecycle::host::{
HostCall as NativeCall, HostCallStep as NativeCallStep, HostFailure, HostPhase, HostStep,
};
use litellm_core::call_lifecycle::{
CallbackFamily, Delivery, ReleaseGate, SuccessFacts, plan_failure, plan_success,
};
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
mod arguments;
mod bindings;
mod compat;
mod dispatch;
mod handle;
mod preparation;
mod setup;
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
pub(crate) use arguments::{BoundArguments, Signature};
use bindings::DeploymentHooks;
pub(crate) use bindings::PythonLogger;
@ -98,18 +93,6 @@ pub(crate) fn run_call<R: PythonRoute + 'static>(
}
}
pub(crate) fn debug_setup(
py: Python<'_>,
call_type: &'static str,
args: &Py<PyTuple>,
kwargs: &Py<PyDict>,
start: &Py<PyAny>,
asynchronous: bool,
) -> PyResult<(Py<PyAny>, Py<PyDict>)> {
let result = setup::setup(py, call_type, args, kwargs, start, asynchronous)?;
Ok((result.logger.object(py).clone().unbind(), result.kwargs))
}
pub(crate) fn missing_state() -> PyErr {
pyo3::exceptions::PyRuntimeError::new_err("missing native call state")
}
@ -179,11 +162,7 @@ impl<R: PythonRoute> PythonLifecycle<R> {
HostFailure::Cancelled(native)
};
let state = self.route.state_mut();
state.retain_first_error(
py,
error,
cancelled && phase != Some(HostPhase::DeploymentFailure),
);
state.retain_first_error(py, error, cancelled && phase != Some(HostPhase::DeploymentFailure));
let _ = state.finish(py);
failure
}
@ -304,7 +283,6 @@ pub(crate) struct PythonCallState {
pub error: Option<Py<PyBaseException>>,
pub asynchronous: bool,
pub internal: bool,
pub supplied: bool,
pub call_type: &'static str,
}
@ -394,7 +372,6 @@ impl PythonCallState {
error: None,
asynchronous,
internal: false,
supplied: false,
call_type,
})
}
@ -408,7 +385,7 @@ impl PythonCallState {
pub fn setup(&mut self, py: Python<'_>) -> PyResult<()> {
self.start = now(py)?;
self.internal = bindings::is_internal_call(py)?;
let result = setup::setup(
let result = bindings::setup(
py,
self.call_type,
&self.args,
@ -416,9 +393,8 @@ impl PythonCallState {
&self.start,
self.asynchronous,
)?;
self.logger = Some(result.logger);
self.kwargs = result.kwargs;
self.supplied = result.supplied;
self.logger = Some(result.logger()?);
self.kwargs = result.kwargs()?;
Ok(())
}
@ -438,16 +414,6 @@ impl PythonCallState {
)
}
pub fn dispatch_request(&self, py: Python<'_>, family: CallbackFamily) -> PyResult<()> {
dispatch::dispatch_request(
py,
dispatch::RequestJob {
logger: self.logger()?,
family,
},
)
}
pub fn dispatch_success(&self, py: Python<'_>) -> PyResult<()> {
match self.try_dispatch_success(py) {
Err(error) if error.is_instance_of::<PyException>(py) => {
@ -458,71 +424,40 @@ impl PythonCallState {
}
}
fn job(&self, py: Python<'_>, family: CallbackFamily) -> PyResult<dispatch::Job> {
let logger = self.logger()?;
let (targets, ids) = dispatch::family_targets(py, logger, family)?;
Ok(dispatch::Job {
logger: logger.clone_ref(py),
targets,
ids,
family,
response: self.response.as_ref().map(|value| value.clone_ref(py)),
error: self.error.as_ref().map(|value| value.clone_ref(py)),
start: self.start.clone_ref(py),
end: self
.end
.as_ref()
.map(|value| value.clone_ref(py))
.unwrap_or_else(|| py.None()),
stream: false,
})
}
fn try_dispatch_success(&self, py: Python<'_>) -> PyResult<()> {
let logger = self.logger()?;
if self.supplied {
return compat::dispatch_success(py, self, logger);
}
let (sync_targets, sync_ids) =
dispatch::family_targets(py, logger, CallbackFamily::SyncSuccess)?;
let facts = SuccessFacts {
asynchronous: self.asynchronous,
internal: self.internal,
fallbacks: !self
.kwargs
.bind(py)
.get_item("fallbacks")?
.is_none_or(|value| value.is_none()),
deferred: logger.defers_async_logging(py),
sync_target_kinds: sync_targets.kinds(&sync_ids),
let pending = || PendingSuccess {
logger: logger.clone_ref(py),
response: self.response.as_ref().map(|value| value.clone_ref(py)),
start: self.start.clone_ref(py),
end: self.end.as_ref().map(|value| value.clone_ref(py)),
};
for selected in plan_success(&facts) {
let runner = dispatch::Runner::start(py, self.job(py, selected.family)?)?;
match (selected.delivery, selected.gate) {
(Delivery::Worker, _) => {
let job = Py::new(py, dispatch::WorkerJob::new(runner))?;
dispatch::leaves(py)?
.getattr("submit_worker")?
.call1((job,))?;
}
(Delivery::Background, ReleaseGate::Immediate) => {
DeferredSuccess::release(py, runner)?;
}
(Delivery::Background, ReleaseGate::Deferred) => {
if !self.asynchronous {
pending().sync(py)
} else {
if !self.internal
&& self
.kwargs
.bind(py)
.get_item("fallbacks")?
.is_none_or(|value| value.is_none())
{
if logger.defers_async_logging(py) {
logger.defer_success(
py,
Py::new(
py,
DeferredSuccess {
runner: Some(runner),
PendingLogging {
pending: Some(pending()),
},
)?,
)?;
} else {
pending().asynchronous(py)?;
}
(Delivery::Inline | Delivery::Await, _) => return Err(missing_state()),
}
logger.sync_success_for_async_call(py, &self.response, &self.start, &self.end)
}
Ok(())
}
pub fn dispatch_failure(
@ -530,39 +465,19 @@ impl PythonCallState {
py: Python<'_>,
asynchronous: bool,
) -> PyResult<Option<Py<PyAny>>> {
if self.logger.is_none() || self.error.is_none() {
if self.logger.is_none() || (self.asynchronous && self.internal) {
return Ok(None);
}
let phase = if asynchronous {
HostPhase::AsyncFailure
} else {
HostPhase::Failure
};
let Some(family) = plan_failure(phase, self.asynchronous, self.internal) else {
let Some(error) = &self.error else {
return Ok(None);
};
if self.supplied {
return compat::dispatch_failure(py, self, family);
}
let mut runner = dispatch::Runner::start(py, self.job(py, family)?)?;
match family.delivery() {
Delivery::Inline => match runner.resume(py, None)? {
dispatch::Step::Done => Ok(None),
dispatch::Step::Await(_) => Err(missing_state()),
},
Delivery::Await => Ok(Some(dispatch::coroutine(py, runner)?)),
Delivery::Worker | Delivery::Background => Err(missing_state()),
}
self.logger()?
.failure(py, error, &self.start, &self.end, asynchronous)
}
pub fn cleanup(&mut self, py: Python<'_>) {
if let Some(logger) = self.logger.take()
&& let Err(error) = dispatch::leaves(py).and_then(|leaves| {
leaves
.getattr("restore_correlation_context")?
.call1((logger.object(py),))
.map(|_| ())
})
&& let Err(error) = logger.restore_context(py)
{
error.write_unraisable(py, None);
}
@ -598,53 +513,58 @@ impl PythonCallState {
}
}
#[pyclass]
struct DeferredSuccess {
runner: Option<dispatch::Runner>,
struct PendingSuccess {
logger: PythonLogger,
response: Option<Py<PyAny>>,
start: Py<PyAny>,
end: Option<Py<PyAny>>,
}
impl DeferredSuccess {
fn release(py: Python<'_>, runner: dispatch::Runner) -> PyResult<()> {
let coroutine = dispatch::coroutine(py, runner)?;
let enqueue = dispatch::leaves(py)?
.getattr("enqueue_background")?
.call1((&coroutine,));
if enqueue.is_err()
&& let Err(error) = coroutine.call_method0(py, "close")
{
error.write_unraisable(py, Some(coroutine.bind(py)));
}
enqueue.map(|_| ())
impl PendingSuccess {
fn sync(&self, py: Python<'_>) -> PyResult<()> {
self.logger
.submit_success(py, &self.response, &self.start, &self.end)
}
fn asynchronous(&self, py: Python<'_>) -> PyResult<()> {
self.logger
.enqueue_success(py, &self.response, &self.start, &self.end)
}
}
#[pyclass]
struct PendingLogging {
pending: Option<PendingSuccess>,
}
#[pymethods]
impl DeferredSuccess {
impl PendingLogging {
fn __call__(slf: &Bound<'_, Self>, py: Python<'_>) -> PyResult<()> {
let runner = slf.borrow_mut().runner.take();
let Some(runner) = runner else {
return Ok(());
};
let logger = runner.logger(py).clone().unbind();
match Self::release(py, runner) {
Err(error) if error.is_instance_of::<PyException>(py) => {
error.write_unraisable(py, Some(logger.bind(py)));
Ok(())
let pending = slf.borrow_mut().pending.take();
if let Some(pending) = pending {
match pending.asynchronous(py) {
Err(error) if error.is_instance_of::<PyException>(py) => {
error.write_unraisable(py, Some(pending.logger.object(py)));
}
result => return result,
}
result => result,
}
Ok(())
}
fn __traverse__(&self, visit: pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> {
match &self.runner {
Some(runner) => runner.traverse(&visit),
None => Ok(()),
if let Some(pending) = &self.pending {
pending.logger.traverse(&visit)?;
visit.call(&pending.response)?;
visit.call(&pending.start)?;
visit.call(&pending.end)?;
}
Ok(())
}
fn close(slf: &Bound<'_, Self>) {
let runner = slf.borrow_mut().runner.take();
drop(runner);
let pending = slf.borrow_mut().pending.take();
drop(pending);
}
fn __clear__(slf: &Bound<'_, Self>) {
@ -655,39 +575,14 @@ impl DeferredSuccess {
#[cfg(test)]
mod tests {
use super::*;
use pyo3::types::PyDict;
use std::sync::Mutex;
use litellm_core::call_lifecycle::host::HostCallFuture;
static PYTHON_GLOBALS: Mutex<()> = Mutex::new(());
fn load_lifecycle_module(py: Python<'_>) -> Bound<'_, PyModule> {
py.run(
pyo3::ffi::c_str!(
r#"
import sys
import types
for name in ("litellm", "litellm.rust_bridge"):
sys.modules.setdefault(name, types.ModuleType(name))
"#
),
None,
None,
)
.unwrap();
let source = std::ffi::CString::new(include_str!(
"../../../../../litellm/rust_bridge/lifecycle.py"
))
.unwrap();
PyModule::from_code(
py,
&source,
pyo3::ffi::c_str!("lifecycle.py"),
pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"),
)
.unwrap()
fn install_logging_worker(py: Python<'_>, worker: &Bound<'_, PyAny>) -> PyResult<()> {
py.import("litellm.litellm_core_utils.logging_worker")?
.setattr("GLOBAL_LOGGING_WORKER", worker)
}
struct RetainingHost {
@ -859,7 +754,17 @@ for name in ("litellm", "litellm.rust_bridge"):
.unwrap_or_else(|error| error.into_inner());
Python::initialize();
Python::attach(|py| {
load_lifecycle_module(py);
let source = std::ffi::CString::new(include_str!(
"../../../../../litellm/rust_bridge/lifecycle.py"
))
.unwrap();
PyModule::from_code(
py,
&source,
pyo3::ffi::c_str!("lifecycle.py"),
pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"),
)
.unwrap();
let route = SyntheticRoute(
PythonCallState::new(
py,
@ -895,7 +800,17 @@ for name in ("litellm", "litellm.rust_bridge"):
Python::initialize();
Python::attach(|py| {
py.import("asyncio").unwrap();
let module = load_lifecycle_module(py);
let source = std::ffi::CString::new(include_str!(
"../../../../../litellm/rust_bridge/lifecycle.py"
))
.unwrap();
let module = PyModule::from_code(
py,
&source,
pyo3::ffi::c_str!("lifecycle.py"),
pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"),
)
.unwrap();
let locals = PyDict::new(py);
locals
.set_item("drive", module.getattr("drive").unwrap())
@ -1003,11 +918,64 @@ assert reference() is None
error: None,
asynchronous,
internal: false,
supplied: false,
call_type: "test",
}
}
#[test]
fn success_dispatch_reports_ordinary_failures_without_replacing_response() {
let _guard = PYTHON_GLOBALS
.lock()
.unwrap_or_else(|error| error.into_inner());
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
import sys
response = object()
failure = ValueError('terminal diagnostic')
diagnostics = []
old_hook = sys.unraisablehook
sys.unraisablehook = lambda event: diagnostics.append(event.exc_value)
class Logger:
def handle_sync_success_callbacks_for_async_calls(self, *args):
raise failure
logger = Logger()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
let response = locals.get_item("response").unwrap().unwrap().unbind();
let mut lifecycle_state = state(
py,
locals.get_item("logger").unwrap().unwrap().unbind(),
response.clone_ref(py),
true,
);
lifecycle_state.internal = true;
lifecycle_state.dispatch_success(py).unwrap();
assert!(lifecycle_state.response.as_ref().unwrap().is(&response));
py.run(
pyo3::ffi::c_str!(
r#"
assert diagnostics == [failure]
sys.unraisablehook = old_hook
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
});
}
#[test]
fn retained_failure_preserves_exception_identity() {
Python::initialize();
@ -1023,6 +991,201 @@ assert reference() is None
});
}
#[test]
fn deferred_logging_uses_call_context_and_allows_reentry_once() {
let _guard = PYTHON_GLOBALS
.lock()
.unwrap_or_else(|error| error.into_inner());
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
import sys
import types
from contextvars import ContextVar
litellm = types.ModuleType('litellm')
core_utils = types.ModuleType('litellm.litellm_core_utils')
logging_worker = types.ModuleType('litellm.litellm_core_utils.logging_worker')
litellm.litellm_core_utils = core_utils
core_utils.logging_worker = logging_worker
sys.modules['litellm'] = litellm
sys.modules['litellm.litellm_core_utils'] = core_utils
sys.modules['litellm.litellm_core_utils.logging_worker'] = logging_worker
marker = ContextVar('marker', default='unset')
observed = []
class Coroutine:
def close(self):
observed.append('closed')
class Worker:
def ensure_initialized_and_enqueue(self, coroutine):
observed.append(marker.get())
pending()
coroutine.close()
class Logger:
def async_success_handler(self, *args):
observed.append('created')
return Coroutine()
worker = Worker()
logger = Logger()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
install_logging_worker(py, &locals.get_item("worker").unwrap().unwrap()).unwrap();
let pending = Py::new(
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: locals
.get_item("logger")
.unwrap()
.unwrap()
.extract()
.unwrap(),
response: Some(py.None()),
start: py.None(),
end: Some(py.None()),
}),
},
)
.unwrap();
locals.set_item("pending", &pending).unwrap();
py.run(
pyo3::ffi::c_str!(
r#"
marker.set('call')
pending()
pending()
assert observed == ['created', 'call', 'closed']
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
});
}
#[test]
fn deferred_logging_close_is_reentry_safe_and_invalidates_aliases() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
observed = []
class Retained:
def __del__(self):
observed.append('finalized')
alias()
class Logger:
def async_success_handler(self, *args):
observed.append('enqueued')
logger = Logger()
retained = Retained()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
let pending = Py::new(
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: locals
.get_item("logger")
.unwrap()
.unwrap()
.extract()
.unwrap(),
response: Some(locals.get_item("retained").unwrap().unwrap().unbind()),
start: py.None(),
end: None,
}),
},
)
.unwrap();
locals.set_item("pending", &pending).unwrap();
locals.set_item("alias", &pending).unwrap();
locals.del_item("retained").unwrap();
py.run(
pyo3::ffi::c_str!(
r#"
pending.close()
alias()
assert observed == ['finalized']
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
});
}
#[test]
fn deferred_logging_collects_cycles_through_typed_logger() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!("class Logger: pass\nlogger = Logger()"),
Some(&locals),
Some(&locals),
)
.unwrap();
let pending = Py::new(
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: locals
.get_item("logger")
.unwrap()
.unwrap()
.extract()
.unwrap(),
response: None,
start: py.None(),
end: None,
}),
},
)
.unwrap();
locals.set_item("pending", pending).unwrap();
py.run(
pyo3::ffi::c_str!(
r#"
import gc
import weakref
logger.pending = pending
reference = weakref.ref(logger)
del logger, pending
gc.collect()
assert reference() is None
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
});
}
#[test]
fn coroutine_collects_cycles_retained_by_bridge_host() {
Python::initialize();

View file

@ -1,337 +0,0 @@
use litellm_core::call_lifecycle::CallbackId;
use litellm_core::call_lifecycle::registration::{
Candidate, DynamicSuccessSlot, Entry, NamedEvent, Registration, RegistrationFacts, Registry,
RegistryMutation, classify_dynamic_success, plan_registration,
};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyList, PyString, PyTuple};
use super::bindings::PythonLogger;
const SETUP_MODULE: &str = "litellm.rust_bridge.setup";
struct Targets<'py> {
objects: Vec<Bound<'py, PyAny>>,
}
impl<'py> Targets<'py> {
fn new() -> Self {
Self {
objects: Vec::new(),
}
}
fn intern(&mut self, object: &Bound<'py, PyAny>) -> CallbackId {
if let Some(index) = self
.objects
.iter()
.position(|existing| existing.is(object) || existing.eq(object).unwrap_or(false))
{
return CallbackId(index as u64);
}
self.objects.push(object.clone());
CallbackId((self.objects.len() - 1) as u64)
}
fn get(&self, id: CallbackId) -> PyResult<&Bound<'py, PyAny>> {
self.objects
.get(id.0 as usize)
.ok_or_else(super::missing_state)
}
}
fn registration<'py>(
setup: &Bound<'py, PyModule>,
object: &Bound<'py, PyAny>,
) -> PyResult<Registration> {
if let Ok(name) = object.cast::<PyString>() {
let name = name.to_str()?;
let known = setup.getattr("is_known_name")?.call1((name,))?.extract()?;
return Ok(Registration::Named {
known,
async_only: matches!(name, "dynamodb" | "openmeter"),
});
}
let asynchronous = setup
.getattr("is_async_callable")?
.call1((object,))?
.extract()?;
Ok(Registration::Object { asynchronous })
}
fn entries<'py>(
setup: &Bound<'py, PyModule>,
targets: &mut Targets<'py>,
list: &Bound<'py, PyAny>,
) -> PyResult<Vec<Entry>> {
list.try_iter()?
.map(|object| {
let object = object?;
Ok(Entry {
id: targets.intern(&object),
registration: registration(setup, &object)?,
})
})
.collect()
}
fn registry_name(registry: Registry) -> &'static str {
match registry {
Registry::Input => "input",
Registry::AsyncInput => "async_input",
Registry::Success => "success",
Registry::AsyncSuccess => "async_success",
Registry::Failure => "failure",
Registry::AsyncFailure => "async_failure",
}
}
fn read_registry<'py>(setup: &Bound<'py, PyModule>, name: &str) -> PyResult<Bound<'py, PyAny>> {
setup.getattr("registry")?.call1((name,))
}
fn read_candidates<'py>(
setup: &Bound<'py, PyModule>,
targets: &mut Targets<'py>,
dynamic: Option<Bound<'py, PyAny>>,
) -> PyResult<Vec<Candidate>> {
let mut candidates = Vec::new();
let global = read_registry(setup, "callbacks")?;
let sources = std::iter::once(global).chain(dynamic);
for source in sources {
for object in source.try_iter()? {
let object = object?;
let candidate = if object.is_instance_of::<PyString>() {
let resolved = setup
.getattr("resolve_named_integration")?
.call1((&object,))?;
if resolved.is_none() {
Candidate {
resolved: None,
duplicate_type: false,
}
} else {
let duplicate_type = setup
.getattr("async_success_registry_has_type")?
.call1((&resolved,))?
.extract()?;
Candidate {
resolved: Some(Entry {
id: targets.intern(&resolved),
registration: registration(setup, &resolved)?,
}),
duplicate_type,
}
}
} else {
Candidate {
resolved: Some(Entry {
id: targets.intern(&object),
registration: registration(setup, &object)?,
}),
duplicate_type: false,
}
};
candidates.push(candidate);
}
}
Ok(candidates)
}
fn apply<'py>(
setup: &Bound<'py, PyModule>,
targets: &Targets<'py>,
mutations: &[RegistryMutation],
function_id: Option<&Bound<'py, PyAny>>,
) -> PyResult<()> {
for mutation in mutations {
match mutation {
RegistryMutation::Append(registry, id) => {
setup
.getattr("append_registry")?
.call1((registry_name(*registry), targets.get(*id)?))?;
}
RegistryMutation::Remove(registry, id) => {
setup
.getattr("remove_registry")?
.call1((registry_name(*registry), targets.get(*id)?))?;
}
RegistryMutation::ExpandNamed(event, id) => {
let event = match event {
NamedEvent::Success => "success",
NamedEvent::Failure => "failure",
};
setup
.getattr("expand_named")?
.call1((targets.get(*id)?, event))?;
}
RegistryMutation::Bootstrap => {
setup.getattr("bootstrap")?.call1((function_id,))?;
}
}
}
Ok(())
}
struct DynamicLists<'py> {
success: Option<Bound<'py, PyList>>,
async_success: Option<Bound<'py, PyList>>,
failure: Option<Bound<'py, PyList>>,
}
fn split_dynamic<'py>(
py: Python<'py>,
setup: &Bound<'py, PyModule>,
targets: &mut Targets<'py>,
kwargs: &Bound<'py, PyDict>,
) -> PyResult<DynamicLists<'py>> {
let success = match kwargs.get_item("success_callback")? {
Some(value) if value.is_instance_of::<PyList>() => {
let list = value.cast_into::<PyList>()?;
let sync = PyList::empty(py);
let asynchronous = PyList::empty(py);
for object in list.iter() {
let entry = Entry {
id: targets.intern(&object),
registration: registration(setup, &object)?,
};
let named_async = object
.cast::<PyString>()
.ok()
.and_then(|name| {
name.to_str()
.ok()
.map(|name| matches!(name, "dynamodb" | "s3"))
})
.unwrap_or(false);
match classify_dynamic_success(entry, named_async) {
DynamicSuccessSlot::Sync => sync.append(&object)?,
DynamicSuccessSlot::Async => asynchronous.append(&object)?,
}
}
kwargs.del_item("success_callback")?;
Some((sync, (!asynchronous.is_empty()).then_some(asynchronous)))
}
_ => None,
};
let failure = match kwargs.get_item("failure_callback")? {
Some(value) if value.is_instance_of::<PyList>() => {
kwargs.del_item("failure_callback")?;
Some(value.cast_into::<PyList>()?)
}
_ => None,
};
let (success, async_success) = match success {
Some((sync, asynchronous)) => (Some(sync), asynchronous),
None => (None, None),
};
Ok(DynamicLists {
success,
async_success,
failure,
})
}
pub(super) struct Setup {
pub logger: PythonLogger,
pub kwargs: Py<PyDict>,
pub supplied: bool,
}
pub(super) fn setup(
py: Python<'_>,
call_type: &str,
args: &Py<PyTuple>,
kwargs: &Py<PyDict>,
start: &Py<PyAny>,
asynchronous: bool,
) -> PyResult<Setup> {
let setup = py.import(SETUP_MODULE)?;
let kwargs = kwargs.bind(py).copy()?;
if !kwargs.contains("litellm_call_id")? {
let call_id = py.import("uuid")?.call_method0("uuid4")?.str()?;
kwargs.set_item("litellm_call_id", call_id)?;
}
if let Some(supplied) = kwargs.get_item("litellm_logging_obj")? {
let logging_class = py
.import("litellm.litellm_core_utils.litellm_logging")?
.getattr("Logging")?;
if supplied.is_instance(&logging_class)? {
return Ok(Setup {
logger: supplied.extract()?,
kwargs: kwargs.unbind(),
supplied: true,
});
}
}
setup.getattr("prepare_environment")?.call0()?;
let guardrails = setup.getattr("applied_guardrails")?.call1((&kwargs,))?;
let function_id = kwargs.get_item("id")?;
let mut targets = Targets::new();
let dynamic = match kwargs.get_item("callbacks")? {
Some(value) => {
kwargs.del_item("callbacks")?;
(!value.is_none()).then_some(value)
}
None => None,
};
let candidates = read_candidates(&setup, &mut targets, dynamic)?;
let facts = RegistrationFacts {
candidates,
input: entries(&setup, &mut targets, &read_registry(&setup, "input")?)?,
success: entries(&setup, &mut targets, &read_registry(&setup, "success")?)?,
failure: entries(&setup, &mut targets, &read_registry(&setup, "failure")?)?,
async_success: entries(
&setup,
&mut targets,
&read_registry(&setup, "async_success")?,
)?
.into_iter()
.map(|entry| entry.id)
.collect(),
async_failure: entries(
&setup,
&mut targets,
&read_registry(&setup, "async_failure")?,
)?
.into_iter()
.map(|entry| entry.id)
.collect(),
bootstrap_pending: setup.getattr("bootstrap_pending")?.call0()?.extract()?,
};
apply(
&setup,
&targets,
&plan_registration(&facts),
function_id.as_ref(),
)?;
let dynamic = split_dynamic(py, &setup, &mut targets, &kwargs)?;
setup.getattr("breadcrumb")?.call1((&kwargs,))?;
if let Some(logger_fn) = kwargs.get_item("logger_fn")? {
setup.getattr("logger_fn")?.call1((logger_fn,))?;
}
let model = match args.bind(py).get_item(0) {
Ok(model) => Some(model),
Err(_) => kwargs.get_item("model")?,
};
let build = setup.getattr("build_logging")?;
let build_kwargs = PyDict::new(py);
build_kwargs.set_item("call_type", call_type)?;
build_kwargs.set_item("model", model)?;
build_kwargs.set_item("kwargs", &kwargs)?;
build_kwargs.set_item("start_time", start)?;
build_kwargs.set_item("asynchronous", asynchronous)?;
build_kwargs.set_item("dynamic_success", dynamic.success)?;
build_kwargs.set_item("dynamic_async_success", dynamic.async_success)?;
build_kwargs.set_item("dynamic_failure", dynamic.failure)?;
build_kwargs.set_item("guardrails", guardrails)?;
let logger = build.call((), Some(&build_kwargs))?;
Ok(Setup {
logger: logger.extract()?,
kwargs: kwargs.unbind(),
supplied: false,
})
}

View file

@ -6,10 +6,8 @@ use litellm_core::ocr::LiteLLMOcrResponse;
use litellm_core::ocr::hooks::OcrDuringCallRequest;
use litellm_python_interop::to_py_preserving_errors as to_py;
use litellm_core::call_lifecycle::CallbackFamily;
use super::host::PythonPayload;
use crate::lifecycle::{PythonCallState, PythonLogger};
use crate::lifecycle::PythonLogger;
pub(super) fn update_logging(
py: Python<'_>,
@ -34,63 +32,26 @@ pub(super) fn update_logging(
pub(super) fn pre_call(
py: Python<'_>,
state: &PythonCallState,
logger: &PythonLogger,
request: &OcrDuringCallRequest,
payload: &PythonPayload,
) -> PyResult<()> {
let logger = state.logger()?;
if state.supplied {
let additional_args = PyDict::new(py);
additional_args.set_item("complete_input_dict", &payload.body)?;
additional_args.set_item("headers", &payload.headers)?;
additional_args.set_item("api_base", &request.url)?;
let kwargs = PyDict::new(py);
kwargs.set_item("input", "OCR document processing")?;
kwargs.set_item("api_key", request.api_key.as_deref())?;
kwargs.set_item("additional_args", additional_args)?;
logger
.object(py)
.call_method("pre_call", (), Some(&kwargs))?;
return Ok(());
}
let kwargs = PyDict::new(py);
kwargs.set_item("api_key", request.api_key.as_deref())?;
kwargs.set_item("body", &payload.body)?;
kwargs.set_item("headers", &payload.headers)?;
kwargs.set_item("url", &request.url)?;
py.import("litellm.rust_bridge.leaves")?
.getattr("record_pre_call")?
.call((logger.object(py),), Some(&kwargs))?;
state.dispatch_request(py, CallbackFamily::RequestPreCall)
py.import("litellm.rust_bridge.ocr")?
.getattr("pre_call")?
.call1((logger.object(py), request.api_key.as_deref(), &payload.body, &payload.headers, &request.url))?;
Ok(())
}
pub(super) fn post_call(
py: Python<'_>,
state: &PythonCallState,
logger: &PythonLogger,
original_response: &Value,
payload: &PythonPayload,
) -> PyResult<()> {
let logger = state.logger()?;
if state.supplied {
let additional_args = PyDict::new(py);
additional_args.set_item("complete_input_dict", &payload.body)?;
additional_args.set_item("headers", &payload.headers)?;
let kwargs = PyDict::new(py);
kwargs.set_item("original_response", to_py(py, original_response)?)?;
kwargs.set_item("additional_args", additional_args)?;
logger
.object(py)
.call_method("post_call", (), Some(&kwargs))?;
return Ok(());
}
let kwargs = PyDict::new(py);
kwargs.set_item("original_response", to_py(py, original_response)?)?;
kwargs.set_item("body", &payload.body)?;
kwargs.set_item("headers", &payload.headers)?;
py.import("litellm.rust_bridge.leaves")?
.getattr("record_post_call")?
.call((logger.object(py),), Some(&kwargs))?;
state.dispatch_request(py, CallbackFamily::RequestPostCall)
py.import("litellm.rust_bridge.ocr")?
.getattr("post_call")?
.call1((logger.object(py), to_py(py, original_response)?, &payload.body, &payload.headers))?;
Ok(())
}
pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {

View file

@ -40,28 +40,17 @@ pub(super) struct PythonPayload {
impl PythonPayload {
fn from_request(py: Python<'_>, request: &OcrDuringCallRequest) -> PyResult<Self> {
let body = to_py(py, &request.body)?
.into_bound(py)
.cast_into::<PyDict>()?;
let body = to_py(py, &request.body)?.into_bound(py).cast_into::<PyDict>()?;
let headers = PyDict::new(py);
for (name, value) in &request.headers {
headers.set_item(name, value)?;
}
Ok(Self {
body: body.unbind(),
headers: headers.unbind(),
})
Ok(Self { body: body.unbind(), headers: headers.unbind() })
}
fn write_back(
&self,
py: Python<'_>,
mut request: OcrDuringCallRequest,
) -> PyResult<OcrDuringCallRequest> {
fn write_back(&self, py: Python<'_>, mut request: OcrDuringCallRequest) -> PyResult<OcrDuringCallRequest> {
request.body = from_py(self.body.bind(py))?;
request.headers = self
.headers
.bind(py)
request.headers = self.headers.bind(py)
.iter()
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
.collect::<PyResult<Vec<_>>>()?;
@ -92,9 +81,7 @@ impl PythonOcrHost {
}
fn project(&mut self, py: Python<'_>) -> PyResult<OcrHostResult> {
let arguments = self
.signature
.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?;
let arguments = self.signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?;
let Projection { native, retained } = project(py, &arguments)?;
let host_token_provider = retained.azure_ad_token_provider.is_some();
self.retained = Some(retained);
@ -128,7 +115,7 @@ impl PythonOcrHost {
&retained.secret_fields,
)?;
let payload = PythonPayload::from_request(py, &request)?;
callbacks::pre_call(py, &self.state, &request, &payload)?;
callbacks::pre_call(py, logger, &request, &payload)?;
let request = payload.write_back(py, request)?;
self.retained_mut()?.payload = Some(payload);
Ok(request)
@ -139,12 +126,13 @@ impl PythonOcrHost {
py: Python<'_>,
request: OcrPostCallRequest,
) -> PyResult<OcrPostCallRequest> {
let payload = self
.retained()?
.payload
.as_ref()
.ok_or_else(missing_state)?;
callbacks::post_call(py, &self.state, &request.original_response, payload)?;
let payload = self.retained()?.payload.as_ref().ok_or_else(missing_state)?;
callbacks::post_call(
py,
self.state.logger()?,
&request.original_response,
payload,
)?;
Ok(request)
}
@ -214,7 +202,9 @@ impl PythonRoute for PythonOcrHost {
OcrHostOperation::AcquireAzureAdToken => {
OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?))
}
OcrHostOperation::PreCall(request) => OcrHostResult::PreCall(Ok(request)),
OcrHostOperation::PreCall(request) => {
OcrHostResult::PreCall(Ok(request))
}
OcrHostOperation::DuringCall(request) => {
OcrHostResult::DuringCall(Ok(self.during_call(py, request)?))
}

View file

@ -62,16 +62,13 @@ fn call(
..OcrAdmission::all()
},
))?;
let host = PythonOcrHost::new(
PythonCallState::new(
py,
args.unbind(),
kwargs.copy()?.unbind(),
asynchronous,
signature.name,
)?,
signature,
);
let host = PythonOcrHost::new(PythonCallState::new(
py,
args.unbind(),
kwargs.copy()?.unbind(),
asynchronous,
signature.name,
)?, signature);
run_call(py, call, host)
}

View file

@ -20,18 +20,14 @@ use crate::marshal::{BoundRouteInputs, Projection};
/// `optional_params`.
const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"];
fn project_document(
document: &Bound<'_, PyAny>,
) -> PyResult<Result<FileDocumentInput, litellm_core::ocr::Error>> {
fn project_document(document: &Bound<'_, PyAny>) -> PyResult<Result<FileDocumentInput, litellm_core::ocr::Error>> {
let kind: String = document.get_item("type")?.extract()?;
if kind != "file" {
let value: serde_json::Value = from_py(document)?;
return Ok(
OcrDocument::try_from(value).map(|document| FileDocumentInput {
input: document.into(),
reader: None,
}),
);
return Ok(OcrDocument::try_from(value).map(|document| FileDocumentInput {
input: document.into(),
reader: None,
}));
}
document.extract().map(Ok)
}
@ -158,7 +154,10 @@ mod tests {
let document = py
.eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None)
.unwrap();
let error = project_document(&document).unwrap().err().unwrap();
let error = project_document(&document)
.unwrap()
.err()
.unwrap();
assert!(error.to_string().contains("document"));
});
}

View file

@ -1,4 +1,3 @@
import datetime
from asyncio import Future
from collections.abc import Coroutine, Mapping
from typing import Never, final
@ -116,20 +115,11 @@ class TokenCounter:
def acount_request(self, body: bytes) -> Future[dict[str, object]]: ...
def gil_stats() -> dict[str, int]: ...
def _debug_setup(
call_type: str,
args: tuple[object, ...],
kwargs: dict[str, object],
start: datetime.datetime,
asynchronous: bool,
) -> tuple[object, dict[str, object]]: ...
__all__ = [
"ResponsesWebSocketConnection",
"RustBridgeDeclined",
"RustUpstreamError",
"TokenCounter",
"_debug_setup",
"achat_completions",
"amessages",
"aocr",

View file

@ -1,815 +0,0 @@
"""Leaf helpers invoked by the native callback dispatcher.
Rust selects every target, delivery and sequence. Nothing here chooses a callback,
reads a registry, or decides whether an event fires. Each function performs one
integration-specific or interop-specific action on values Rust hands it.
Labeled leaf helpers pending native migration (see rust-callback-inventory.md):
`prepare_success_logging`, `prepare_failure_logging`, `dispatch_named_success`,
`dispatch_named_failure`, `dispatch_callable`. They are deleted when cost and
payload construction move into core and when each string integration becomes a
CustomLogger.
"""
from __future__ import annotations
import datetime
import json
import traceback
from collections.abc import Awaitable, Callable, Mapping
from typing import ( # noqa: TID251 # narrows the untyped legacy Logging and CustomLogger surfaces once
TYPE_CHECKING,
Final,
Literal,
Protocol,
cast,
)
from litellm.integrations.custom_logger import CustomLogger
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
TerminalFamily = Literal["sync_success", "async_success", "sync_failure", "async_failure"]
Family = Literal["request", TerminalFamily]
Details = dict[str, object]
Timestamp = datetime.datetime
LegacyCall = Callable[..., object]
LegacyAsyncCall = Callable[..., Awaitable[None]]
class LoggerView(Protocol):
model: str | None
messages: object
call_type: str
start_time: Timestamp
litellm_call_id: str
completion_start_time: Timestamp | None
model_call_details: Details
log_raw_request_response: bool
standard_callback_dynamic_params: object
standard_built_in_tools_params: object
def record_api_call_start_time(self) -> None: ...
def record_post_call(
self, original_response: object, input: object, api_key: object, additional_args: Details
) -> None: ...
def should_run_callback(self, callback: object, litellm_params: Details, event_hook: str) -> bool: ...
def _pre_call(self, input: str, api_key: str | None, model: str | None, additional_args: Details) -> None: ...
def _print_llm_call_debugging_log(self, api_base: str, headers: Details, additional_args: Details) -> None: ...
def _get_request_curl_command(
self, api_base: str, headers: Details | None, additional_args: Details, data: object
) -> str: ...
def _get_masked_api_base(self, api_base: str) -> str: ...
def _get_raw_request_body(self, data: object) -> Details: ...
def _get_masked_headers(self, headers: Details) -> Details: ...
def _response_cost_calculator(self, result: object) -> float | None: ...
def _build_standard_logging_payload(
self, init_response_obj: object, start_time: Timestamp, end_time: Timestamp
) -> object: ...
def _handle_callback_failure(self, callback: object) -> None: ...
class IntegrationView(Protocol):
def log_pre_api_call(self, model: str | None, messages: object, kwargs: Details) -> None: ...
def log_post_api_call(
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp | None
) -> None: ...
def log_success_event(
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp
) -> None: ...
def log_failure_event(
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp
) -> None: ...
def async_log_success_event(
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp
) -> Awaitable[None]: ...
def async_log_failure_event(
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp
) -> Awaitable[None]: ...
def logging_hook(self, kwargs: Details, result: object, call_type: str) -> tuple[Details, object]: ...
def async_logging_hook(
self, kwargs: Details, result: object, call_type: str
) -> Awaitable[tuple[Details, object]]: ...
def redact_standard_logging_payload_from_model_call_details(self, model_call_details: Details) -> Details: ...
def log_input_event(
self, model: str | None, messages: object, kwargs: Details, print_verbose: LegacyCall, callback_func: LegacyCall
) -> None: ...
def log_event(
self,
kwargs: Details,
response_obj: object,
start_time: Timestamp,
end_time: Timestamp,
print_verbose: LegacyCall,
callback_func: LegacyCall,
) -> None: ...
def async_log_event(
self,
kwargs: Details,
response_obj: object,
start_time: Timestamp,
end_time: Timestamp,
print_verbose: LegacyCall,
callback_func: LegacyCall,
) -> Awaitable[None]: ...
def _logger(logger: Logging) -> LoggerView:
return cast(LoggerView, logger) # cast-ok: legacy Logging is untyped; this protocol names the attributes we read
def _integration(callback: CustomLogger) -> IntegrationView:
return cast(IntegrationView, callback) # cast-ok: legacy CustomLogger methods are untyped
def _legacy_module() -> Mapping[str, object]:
from litellm.litellm_core_utils import litellm_logging
return cast(
Mapping[str, object], vars(litellm_logging)
) # cast-ok: module globals hold the legacy integration singletons
def _print_verbose() -> LegacyCall:
from litellm.litellm_core_utils import litellm_logging
return cast(LegacyCall, litellm_logging.print_verbose) # cast-ok: legacy debug printer is untyped
def _method(target: object, name: str) -> LegacyCall:
return cast(LegacyCall, getattr(target, name)) # cast-ok: legacy integration singletons are untyped
def _async_method(target: object, name: str) -> LegacyAsyncCall:
return cast(LegacyAsyncCall, getattr(target, name)) # cast-ok: legacy integration singletons are untyped
def _redact_string(value: str) -> str:
from litellm.litellm_core_utils import litellm_logging
return cast(Callable[[str], str], litellm_logging._redact_string)(value) # pyright: ignore[reportPrivateUsage] # cast-ok: legacy helper
def _redact_result(details: Details, result: object) -> object:
from litellm.litellm_core_utils import redact_messages
redact: Final = cast(LegacyCall, redact_messages.redact_message_input_output_from_logging) # cast-ok: legacy helper
return redact(model_call_details=details, result=result)
def record_pre_call(
logger: Logging,
*,
api_key: str | None,
body: Details,
headers: dict[str, str],
url: str,
) -> None:
view: Final = _logger(logger)
additional_args: Final[Details] = {"complete_input_dict": body, "headers": headers, "api_base": url}
view._pre_call(input="OCR document processing", api_key=api_key, model=None, additional_args=additional_args) # pyright: ignore[reportPrivateUsage] # legacy state writer
view._print_llm_call_debugging_log(api_base=url, headers=dict(headers), additional_args=additional_args) # pyright: ignore[reportPrivateUsage] # legacy debug output
_capture_raw_request(view, additional_args)
_run_logger_fn(logger)
view.record_api_call_start_time()
def _capture_raw_request(view: LoggerView, additional_args: Details) -> None:
import litellm
from litellm.types.utils import RawRequestTypedDict
if not (view.log_raw_request_response or litellm.log_raw_request_response):
return
details: Final = view.model_call_details
params: Final = cast(Details, details.get("litellm_params") or {}) # cast-ok: legacy nested dict
metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict
params.setdefault("metadata", metadata)
if litellm.turn_off_message_logging:
metadata["raw_request"] = "redacted by litellm. 'litellm.turn_off_message_logging=True'"
return
api_base: Final = str(additional_args.get("api_base") or "")
headers: Final = cast(Details, additional_args.get("headers") or {}) # cast-ok: legacy nested dict
body: Final = additional_args.get("complete_input_dict", {})
try:
curl: Final = view._get_request_curl_command( # pyright: ignore[reportPrivateUsage] # legacy debug formatter
api_base=api_base, headers=headers, additional_args=additional_args, data=body
)
metadata["raw_request"] = _redact_string(str(curl))
details["raw_request_typed_dict"] = RawRequestTypedDict(
raw_request_api_base=view._get_masked_api_base(api_base), # pyright: ignore[reportPrivateUsage] # legacy masking
raw_request_body=view._get_raw_request_body(body), # pyright: ignore[reportPrivateUsage] # legacy masking
raw_request_headers=view._get_masked_headers(headers), # pyright: ignore[reportPrivateUsage] # legacy masking
error=None,
)
except Exception as error: # noqa: BLE001 # raw-request capture is best effort by contract
details["raw_request_typed_dict"] = RawRequestTypedDict(error=str(error))
metadata["raw_request"] = _redact_string(f"Unable to Log raw request: {error}")
def _run_logger_fn(logger: Logging) -> None:
from litellm._logging import verbose_logger
logger_fn: Final = cast(
Callable[[Details], object] | None, getattr(logger, "logger_fn", None)
) # cast-ok: user hook is untyped
if not callable(logger_fn):
return
try:
logger_fn(_logger(logger).model_call_details)
except Exception as error: # noqa: BLE001 # user logger_fn failures never fail the request
verbose_logger.exception("LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging %s", error)
def record_post_call(logger: Logging, *, original_response: object, body: Details, headers: dict[str, str]) -> None:
view: Final = _logger(logger)
serialized: Final = (
json.dumps(original_response, default=str) if isinstance(original_response, dict) else original_response
)
view.record_post_call(
original_response=serialized,
input=None,
api_key=None,
additional_args={"complete_input_dict": body, "headers": headers},
)
_run_logger_fn(logger)
_redact_result(view.model_call_details, serialized)
def log_pre_api_call(logger: Logging, callback: CustomLogger) -> None:
view: Final = _logger(logger)
_integration(callback).log_pre_api_call(model=view.model, messages=view.messages, kwargs=view.model_call_details)
def log_post_api_call(logger: Logging, callback: CustomLogger) -> None:
view: Final = _logger(logger)
_integration(callback).log_post_api_call(
kwargs=view.model_call_details, response_obj=None, start_time=view.start_time, end_time=None
)
def dispatch_named_request(logger: Logging, name: str, event: Literal["pre_api_call", "post_api_call"]) -> None:
view: Final = _logger(logger)
module: Final = _legacy_module()
if name == "supabase" and event == "pre_api_call" and (client := module.get("supabaseClient")) is not None:
details: Final = view.model_call_details
_method(client, "input_log_event")(
model=view.model,
messages=view.messages,
end_user=details.get("user", "default"),
litellm_call_id=details["litellm_call_id"],
print_verbose=_print_verbose(),
)
if name == "sentry" and (add_breadcrumb := module.get("add_breadcrumb")) is not None:
cast(LegacyCall, add_breadcrumb)( # cast-ok: legacy sentry hook
category="litellm.llm_call", message=f"Model Call Details {event}: {view.model_call_details}", level="info"
)
def dispatch_callable_request(logger: Logging, callback: LegacyCall) -> None:
custom: Final = _legacy_module().get("customLogger")
if not isinstance(custom, CustomLogger):
return
view: Final = _logger(logger)
_integration(custom).log_input_event(
model=view.model,
messages=view.messages,
kwargs=view.model_call_details,
print_verbose=_print_verbose(),
callback_func=callback,
)
def report_target_failure(logger: Logging, callback: object, family: Family, error: BaseException) -> None:
from litellm._logging import verbose_logger
verbose_logger.error(
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while %s logging with %s: %s",
family,
callback,
"".join(traceback.format_exception(error)),
)
capture: Final = _legacy_module().get("capture_exception")
if capture is not None and family in ("request", "sync_success", "sync_failure"):
cast(LegacyCall, capture)(error) # cast-ok: legacy sentry hook
if family not in ("request", "sync_failure"):
_logger(logger)._handle_callback_failure(callback=callback) # pyright: ignore[reportPrivateUsage] # legacy prometheus counter
def prepare_success_logging(logger: Logging, response: object, start_time: Timestamp, end_time: Timestamp) -> object:
from litellm.litellm_core_utils.litellm_logging import emit_standard_logging_payload
from litellm.types.utils import StandardLoggingPayload
view: Final = _logger(logger)
details: Final = view.model_call_details
if view.completion_start_time is None:
view.completion_start_time = end_time
details["completion_start_time"] = end_time
details["log_event_type"] = "successful_api_call"
details["end_time"] = end_time
details["cache_hit"] = None
hidden: Final = cast(Details, getattr(response, "_hidden_params", None) or {}) # cast-ok: legacy response attribute
params: Final = cast(Details | None, details.get("litellm_params")) # cast-ok: legacy nested dict
if hidden and params is not None:
metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict
params["metadata"] = metadata
metadata["hidden_params"] = hidden
existing: Final = details.get("response_cost")
if "response_cost" in hidden:
details["response_cost"] = hidden["response_cost"]
elif existing is None or existing == 0:
details["response_cost"] = view._response_cost_calculator(result=response) # pyright: ignore[reportPrivateUsage] # labeled leaf: native cost pending
payload: Final = view._build_standard_logging_payload(response, start_time, end_time) # pyright: ignore[reportPrivateUsage] # labeled leaf: native payload pending
details["standard_logging_object"] = payload
if payload is not None:
emit_standard_logging_payload(
cast(StandardLoggingPayload, payload)
) # cast-ok: legacy builder returns the payload TypedDict
return _redact_result(details, response)
def prepare_failure_logging(
logger: Logging, exception: BaseException, start_time: Timestamp, end_time: Timestamp
) -> str:
from litellm.litellm_core_utils import litellm_logging
formatted: Final = "".join(traceback.format_exception(exception))
view: Final = _logger(logger)
details: Final = view.model_call_details
if details.get("exception") is exception and details.get("standard_logging_object") is not None:
return formatted
details["log_event_type"] = "failed_api_call"
details["exception"] = exception
details["traceback_exception"] = _redact_string(formatted)
details["end_time"] = end_time
details.setdefault("original_response", None)
if details.get("combined_usage_object") is None:
details["response_cost"] = 0
headers: Final = getattr(exception, "headers", None)
if isinstance(headers, dict):
params: Final = cast(Details, details.setdefault("litellm_params", {})) # cast-ok: legacy nested dict
metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict
metadata.update(cast(Details, headers)) # cast-ok: exception headers are a plain dict
build: Final = cast(
LegacyCall, litellm_logging.get_standard_logging_object_payload
) # cast-ok: labeled leaf: native payload pending
details["standard_logging_object"] = build(
kwargs=details,
init_response_obj={},
start_time=start_time,
end_time=end_time,
logging_obj=logger,
status="failure",
error_str=_redact_string(str(exception)),
original_exception=exception,
standard_built_in_tools_params=view.standard_built_in_tools_params,
)
return formatted
_EVENT_HOOKS: Final[Mapping[TerminalFamily, str]] = {
"sync_success": "success_handler",
"async_success": "async_success_handler",
"sync_failure": "failure_handler",
"async_failure": "async_failure_handler",
}
def should_run_callback(logger: Logging, callback: object, family: TerminalFamily) -> bool:
view: Final = _logger(logger)
params: Final = cast(Details, view.model_call_details.get("litellm_params") or {}) # cast-ok: legacy nested dict
return view.should_run_callback(callback=callback, litellm_params=params, event_hook=_EVENT_HOOKS[family])
def should_run_guardrail_hook(logger: Logging, callback: object) -> bool:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
if not isinstance(callback, CustomGuardrail):
return True
decide: Final = cast(LegacyCall, callback.should_run_guardrail) # cast-ok: legacy guardrail method is untyped
return decide(data=_logger(logger).model_call_details, event_type=GuardrailEventHooks.logging_only) is True
def logging_hook(logger: Logging, callback: CustomLogger, result: object) -> object:
view: Final = _logger(logger)
details, replaced = _integration(callback).logging_hook(
kwargs=view.model_call_details, result=result, call_type=view.call_type
)
view.model_call_details = details
return replaced
async def async_logging_hook(logger: Logging, callback: CustomLogger, result: object) -> object:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils import redact_messages
view: Final = _logger(logger)
redact: Final = cast(
LegacyCall, redact_messages.redact_message_input_output_from_custom_logger
) # cast-ok: legacy helper
redacted: Final = (
result
if isinstance(callback, CustomGuardrail)
else redact(result=result, litellm_logging_obj=logger, custom_logger=callback)
)
details, replaced = await _integration(callback).async_logging_hook(
kwargs=view.model_call_details, result=redacted, call_type=view.call_type
)
view.model_call_details = details
return replaced
def mark_logged(logger: Logging, marker: str) -> None:
_logger(logger).model_call_details[marker] = True
def already_logged(logger: Logging, marker: str) -> bool:
return _logger(logger).model_call_details.get(marker, False) is True
def log_success_event(
logger: Logging, callback: CustomLogger, response: object, start_time: Timestamp, end_time: Timestamp
) -> None:
_integration(callback).log_success_event(
kwargs=_logger(logger).model_call_details, response_obj=response, start_time=start_time, end_time=end_time
)
def async_log_success_event(
logger: Logging, callback: CustomLogger, response: object, start_time: Timestamp, end_time: Timestamp
) -> Awaitable[None]:
from litellm.litellm_core_utils import redact_messages
integration: Final = _integration(callback)
details: Final = integration.redact_standard_logging_payload_from_model_call_details(
model_call_details=_logger(logger).model_call_details
)
redact: Final = cast(
Callable[..., Details], redact_messages.redact_streaming_responses_for_custom_logger
) # cast-ok: legacy helper
view: Final = redact(model_call_details=details, custom_logger=callback)
return integration.async_log_success_event(
kwargs=view, response_obj=response, start_time=start_time, end_time=end_time
)
def log_failure_event(logger: Logging, callback: CustomLogger, start_time: Timestamp, end_time: Timestamp) -> None:
_integration(callback).log_failure_event(
kwargs=_logger(logger).model_call_details, response_obj=None, start_time=start_time, end_time=end_time
)
def async_log_failure_event(
logger: Logging, callback: CustomLogger, start_time: Timestamp, end_time: Timestamp
) -> Awaitable[None]:
return _integration(callback).async_log_failure_event(
kwargs=_logger(logger).model_call_details, response_obj=None, start_time=start_time, end_time=end_time
)
def _custom_logger_singleton() -> IntegrationView:
from litellm.litellm_core_utils import litellm_logging
existing: Final = _legacy_module().get("customLogger")
if isinstance(existing, CustomLogger):
return _integration(existing)
created: Final = CustomLogger()
litellm_logging.customLogger = created # pyright: ignore[reportAttributeAccessIssue] # legacy module global
return _integration(created)
def dispatch_callable(
logger: Logging,
callback: LegacyCall,
family: TerminalFamily,
response: object,
start_time: Timestamp,
end_time: Timestamp,
) -> Awaitable[None] | None:
custom: Final = _custom_logger_singleton()
details: Final = _logger(logger).model_call_details
match family:
case "sync_success" | "sync_failure":
custom.log_event(
kwargs=details,
response_obj=response,
start_time=start_time,
end_time=end_time,
print_verbose=_print_verbose(),
callback_func=callback,
)
return None
case "async_success" | "async_failure":
return custom.async_log_event(
kwargs=details,
response_obj=response,
start_time=start_time,
end_time=end_time,
print_verbose=_print_verbose(),
callback_func=callback,
)
_SUCCESS_SINGLETONS: Final[Mapping[str, str]] = {
"promptlayer": "promptLayerLogger",
"supabase": "supabaseClient",
"wandb": "weightsBiasesLogger",
"logfire": "logfireLogger",
"lunary": "lunaryLogger",
"helicone": "heliconeLogger",
"greenscale": "greenscaleLogger",
"athina": "athinaLogger",
"traceloop": "traceloopLogger",
"s3": "s3Logger",
"openmeter": "openMeterLogger",
}
def dispatch_named_success(
logger: Logging, name: str, response: object, start_time: Timestamp, end_time: Timestamp
) -> Awaitable[None] | None:
view: Final = _logger(logger)
details: Final = view.model_call_details
print_verbose: Final = _print_verbose()
integration: Final = _legacy_module().get(_SUCCESS_SINGLETONS.get(name, ""))
without_response: Final = {key: value for key, value in details.items() if key != "original_response"}
match name:
case "promptlayer" | "wandb" | "athina" if integration is not None:
_method(integration, "log_event")(
kwargs=details,
response_obj=response,
start_time=start_time,
end_time=end_time,
print_verbose=print_verbose,
)
case "logfire" if integration is not None:
from litellm.integrations.logfire_logger import LogfireLevel
_method(integration, "log_event")(
kwargs=without_response,
response_obj=response,
start_time=start_time,
end_time=end_time,
print_verbose=print_verbose,
level=LogfireLevel.INFO.value,
)
case "greenscale" if integration is not None:
_method(integration, "log_event")(
kwargs=without_response,
response_obj=response,
start_time=start_time,
end_time=end_time,
print_verbose=print_verbose,
)
case "supabase" if integration is not None:
_method(integration, "log_event")(
model=view.model,
messages=view.messages,
end_user=details.get("user", "default"),
response_obj=response,
start_time=start_time,
end_time=end_time,
litellm_call_id=details["litellm_call_id"],
print_verbose=print_verbose,
)
case "lunary" if integration is not None:
_method(integration, "log_event")(
kwargs=details,
type="llm",
event="end",
model=view.model,
input=details["input"],
user_id=details.get("user", "default"),
response_obj=response,
start_time=start_time,
end_time=end_time,
run_id=view.litellm_call_id,
print_verbose=print_verbose,
)
case "helicone" if integration is not None:
_method(integration, "log_success")(
model=view.model,
messages=view.messages,
response_obj=response,
start_time=start_time,
end_time=end_time,
print_verbose=print_verbose,
kwargs=details,
)
case "langfuse":
_langfuse(
logger, response=response, start_time=start_time, end_time=end_time, level=None, status_message=None
)
case "traceloop" if integration is not None:
_method(integration, "log_event")(
kwargs=details,
response_obj=response,
start_time=start_time,
end_time=end_time,
user_id=details.get("user", None),
print_verbose=print_verbose,
)
case "s3" if integration is not None:
_method(integration, "log_event")(
kwargs=details,
response_obj=response,
start_time=start_time,
end_time=end_time,
print_verbose=print_verbose,
)
case "openmeter" if integration is not None:
return _async_method(integration, "async_log_success_event")(
kwargs=details, response_obj=response, start_time=start_time, end_time=end_time
)
case "dynamodb":
return _dynamodb(details, response, start_time, end_time, print_verbose)
case _:
return None
return None
def _dynamodb(
details: Details, response: object, start_time: Timestamp, end_time: Timestamp, print_verbose: LegacyCall
) -> Awaitable[None]:
from litellm.integrations.dynamodb import DyanmoDBLogger
from litellm.litellm_core_utils import litellm_logging
existing: Final = _legacy_module().get("dynamoLogger")
dynamo: Final = existing if isinstance(existing, DyanmoDBLogger) else DyanmoDBLogger()
litellm_logging.dynamoLogger = dynamo # pyright: ignore[reportAttributeAccessIssue] # legacy module global
return _async_method(dynamo, "_async_log_event")(
kwargs=details, response_obj=response, start_time=start_time, end_time=end_time, print_verbose=print_verbose
)
def _langfuse(
logger: Logging,
*,
response: object,
start_time: Timestamp,
end_time: Timestamp,
level: str | None,
status_message: str | None,
) -> None:
from litellm.integrations.langfuse import langfuse_handler
view: Final = _logger(logger)
module: Final = _legacy_module()
kwargs: Final = {key: value for key, value in view.model_call_details.items() if key != "original_response"}
select: Final = cast(
LegacyCall, langfuse_handler.LangFuseHandler.get_langfuse_logger_for_request
) # cast-ok: legacy factory
handler: Final = select(
globalLangfuseLogger=module.get("langFuseLogger"),
standard_callback_dynamic_params=view.standard_callback_dynamic_params,
in_memory_dynamic_logger_cache=module["in_memory_dynamic_logger_cache"],
)
if handler is None:
return
extra: Final[Details] = {"level": level, "status_message": status_message} if level is not None else {}
result: Final = _method(handler, "log_event_on_langfuse")(
kwargs=kwargs,
response_obj=response,
start_time=start_time,
end_time=end_time,
user_id=kwargs.get("user", None),
**extra,
)
trace_id: Final = (
cast(Details, result).get("trace_id") if isinstance(result, dict) else None
) # cast-ok: legacy response dict
if trace_id is not None:
_method(module["in_memory_trace_id_cache"], "set_cache")(
litellm_call_id=view.litellm_call_id, service_name="langfuse", trace_id=trace_id
)
def dispatch_named_failure(
logger: Logging,
name: str,
exception: BaseException,
formatted: str,
start_time: Timestamp,
end_time: Timestamp,
) -> None:
view: Final = _logger(logger)
module: Final = _legacy_module()
details: Final = view.model_call_details
print_verbose: Final = _print_verbose()
match name:
case "lunary" if (lunary := module.get("lunaryLogger")) is not None:
_method(lunary, "log_event")(
kwargs=details,
type="llm",
event="error",
user_id=details.get("user", "default"),
model=view.model,
input=details["input"],
error=formatted,
run_id=view.litellm_call_id,
start_time=start_time,
end_time=end_time,
print_verbose=print_verbose,
)
case "sentry" if (capture := module.get("capture_exception")) is not None:
cast(LegacyCall, capture)(exception) # cast-ok: legacy sentry hook
case "supabase" if (supabase := module.get("supabaseClient")) is not None:
_method(supabase, "log_event")(
model=view.model,
messages=view.messages,
end_user=details.get("user", "default"),
response_obj=None,
start_time=start_time,
end_time=end_time,
litellm_call_id=details["litellm_call_id"],
print_verbose=print_verbose,
)
case "langfuse":
_langfuse(
logger,
response=None,
start_time=start_time,
end_time=end_time,
level="ERROR",
status_message=str(exception),
)
case "traceloop" if (traceloop := module.get("traceloopLogger")) is not None:
_method(traceloop, "log_event")(
start_time=start_time,
end_time=end_time,
response_obj=None,
user_id=details.get("user", None),
print_verbose=print_verbose,
status_message=str(exception),
level="ERROR",
kwargs=details,
)
case "logfire" if (logfire := module.get("logfireLogger")) is not None:
from litellm.integrations.logfire_logger import LogfireLevel
_method(logfire, "log_event")(
kwargs={
**{key: value for key, value in details.items() if key != "original_response"},
"exception": exception,
},
response_obj=None,
start_time=start_time,
end_time=end_time,
level=LogfireLevel.ERROR.value,
print_verbose=print_verbose,
)
case _:
return
def restore_correlation_context(logger: object) -> None:
from litellm import utils
utils._restore_correlation_context_if_supported(logger) # pyright: ignore[reportPrivateUsage] # legacy interop helper
def submit_worker(job: Callable[[], None]) -> None:
import contextvars
from litellm.litellm_core_utils.thread_pool_executor import executor
context: Final = contextvars.copy_context()
_ = executor.submit(context.run, job)
def enqueue_background(coroutine: Awaitable[None]) -> None:
import contextvars
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
enqueue: Final = cast(
Callable[[Awaitable[None]], None], GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue
) # cast-ok: legacy worker accepts any coroutine
contextvars.copy_context().run(enqueue, coroutine)
def now() -> Timestamp:
return datetime.datetime.now()

View file

@ -1,8 +1,18 @@
from __future__ import annotations
import datetime
import uuid
from collections.abc import Awaitable, Mapping
from dataclasses import dataclass
from typing import Final, Protocol
from typing import (
TYPE_CHECKING,
Final,
Protocol,
cast, # noqa: TID251 # bounded compatibility calls into legacy Python integrations
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
@dataclass(frozen=True, slots=True)
@ -42,6 +52,47 @@ async def drive(execution: Execution) -> object:
execution.close()
class MetadataUpdater(Protocol):
def __call__(
self,
result: object,
logging_obj: Logging,
model: str | None,
kwargs: dict[str, object],
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
@dataclass(frozen=True, slots=True)
class CallSetup:
logger: Logging
kwargs: dict[str, object]
def setup(
call_type: str,
args: tuple[object, ...],
kwargs: Mapping[str, object],
start_time: datetime.datetime,
asynchronous: bool,
) -> CallSetup:
from litellm import utils
from litellm.litellm_core_utils.litellm_logging import Logging
arguments: Final = { # mutable-ok: function_setup consumes an owned kwargs dict
"litellm_call_id": str(uuid.uuid4()),
**kwargs,
}
supplied: Final = arguments.get("litellm_logging_obj")
if isinstance(supplied, Logging):
return CallSetup(supplied, arguments)
logger, prepared = utils.function_setup(
call_type, utils.Rules(), start_time, *args, is_async_call=asynchronous, **arguments
)
return CallSetup(logger, prepared)
def check_limits(kwargs: Mapping[str, object]) -> None:
import litellm
from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit
@ -51,3 +102,19 @@ def check_limits(kwargs: Mapping[str, object]) -> None:
raise litellm.BudgetExceededError(current_cost=current_cost, max_budget=litellm.max_budget)
if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request):
raise RuntimeError("Max retries per request hit!")
def finalize(
response: object,
logger: Logging,
kwargs: dict[str, object],
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None:
from litellm.litellm_core_utils.llm_response_utils import response_metadata
model: Final = kwargs.get("model")
update: Final = cast( # cast-ok: legacy metadata function accepts concrete kwargs
MetadataUpdater, response_metadata.update_response_metadata
)
update(response, logger, model if isinstance(model, str) else None, kwargs, start_time, end_time)

View file

@ -22,6 +22,10 @@ class OcrLoggingProtocol(Protocol):
custom_llm_provider: str,
) -> object: ...
def pre_call(self, *, input: str, api_key: str | None, additional_args: dict[str, object]) -> object: ...
def post_call(self, *, original_response: object, additional_args: dict[str, object]) -> object: ...
def _redact(params: Mapping[str, object], secret_fields: Sequence[str]) -> dict[str, object]:
return { # mutable-ok: update_from_kwargs takes dict
@ -60,6 +64,32 @@ def update_logging(
)
def pre_call(
logger: OcrLoggingProtocol,
api_key: str | None,
body: dict[str, object],
headers: dict[str, str],
url: str,
) -> None:
logger.pre_call(
input="OCR document processing",
api_key=api_key,
additional_args={"complete_input_dict": body, "headers": headers, "api_base": url},
)
def post_call(
logger: OcrLoggingProtocol,
original_response: object,
body: dict[str, object],
headers: dict[str, str],
) -> None:
logger.post_call(
original_response=original_response,
additional_args={"complete_input_dict": body, "headers": headers},
)
class RustOcr(Protocol):
def __call__(self, *args: object, **kwargs: object) -> OCRResponse: ...

View file

@ -1,262 +0,0 @@
from __future__ import annotations
import datetime
from collections.abc import Callable, Mapping, MutableSequence, Sequence
from typing import ( # noqa: TID251 # narrows legacy untyped registries at the boundary
TYPE_CHECKING,
Final,
Literal,
Protocol,
cast,
)
from litellm.integrations.custom_logger import CustomLogger
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
CallbackTarget = str | Callable[..., object] | CustomLogger
RegistryName = Literal["input", "async_input", "success", "async_success", "failure", "async_failure", "callbacks"]
_REGISTRY_ATTRIBUTES: Final[Mapping[RegistryName, str]] = {
"input": "input_callback",
"async_input": "_async_input_callback",
"success": "success_callback",
"async_success": "_async_success_callback",
"failure": "failure_callback",
"async_failure": "_async_failure_callback",
"callbacks": "callbacks",
}
class _CallbackManager(Protocol):
def add_litellm_success_callback(self, callback: CallbackTarget) -> None: ...
def add_litellm_failure_callback(self, callback: CallbackTarget) -> None: ...
def add_litellm_async_success_callback(self, callback: CallbackTarget) -> None: ...
def add_litellm_async_failure_callback(self, callback: CallbackTarget) -> None: ...
class _LoggingFactory(Protocol):
def __call__(
self,
*,
model: str | None,
messages: object,
stream: bool,
litellm_call_id: str,
litellm_trace_id: str | None,
function_id: str,
call_type: str,
start_time: datetime.datetime,
dynamic_success_callbacks: list[CallbackTarget] | None,
dynamic_failure_callbacks: list[CallbackTarget] | None,
dynamic_async_success_callbacks: list[CallbackTarget] | None,
dynamic_async_failure_callbacks: list[CallbackTarget] | None,
kwargs: dict[str, object],
applied_guardrails: list[str],
supports_correlation_logging: bool,
) -> Logging: ...
class _EnvironmentUpdater(Protocol):
def __call__(
self,
*,
model: str | None,
user: str,
optional_params: dict[str, object],
litellm_params: dict[str, object],
stream_options: object,
) -> None: ...
def registry(name: RegistryName) -> MutableSequence[CallbackTarget]:
import litellm
return cast(
MutableSequence[CallbackTarget], getattr(litellm, _REGISTRY_ATTRIBUTES[name])
) # cast-ok: legacy module-level lists are untyped
def is_async_callable(callback: object) -> bool:
from litellm.litellm_core_utils.cached_imports import get_coroutine_checker
return get_coroutine_checker().is_async_callable(callback)
def is_known_name(callback: str) -> bool:
import litellm
known: Final = cast(Sequence[str], litellm._known_custom_logger_compatible_callbacks) # pyright: ignore[reportPrivateUsage] # cast-ok: registry list has no public typed accessor
return callback in known
def resolve_named_integration(callback: str) -> CustomLogger | None:
from litellm.litellm_core_utils import litellm_logging
resolve: Final = cast( # cast-ok: legacy factory is untyped at its definition
Callable[..., CustomLogger | None],
litellm_logging._init_custom_logger_compatible_class, # pyright: ignore[reportPrivateUsage] # legacy factory
)
return resolve(callback, internal_usage_cache=None, llm_router=None)
def async_success_registry_has_type(callback: object) -> bool:
return any(type(existing) is type(callback) for existing in registry("async_success"))
def bootstrap_pending() -> bool:
from litellm import utils
return not utils.callback_list
def bootstrap(function_id: str | None) -> None:
from litellm import utils
from litellm.litellm_core_utils import cached_imports
combined: Final = list({*registry("input"), *registry("success"), *registry("failure")})
utils.callback_list = cast(
list[str], combined
) # rebind-ok: legacy module global consumed by set_callbacks # cast-ok: legacy list annotation is narrower than its contents
set_callbacks: Final = cast(Callable[..., None], cached_imports.get_set_callbacks()) # pyright: ignore[reportUnknownMemberType] # cast-ok: cached import is untyped
set_callbacks(callback_list=combined, function_id=function_id)
def expand_named(callback: str, event: Literal["success", "failure"]) -> None:
from litellm import utils
utils._add_custom_logger_callback_to_specific_event(callback, event) # pyright: ignore[reportPrivateUsage] # legacy expansion helper
def append_registry(name: RegistryName, callback: CallbackTarget) -> None:
import litellm
manager: Final = cast(_CallbackManager, litellm.logging_callback_manager) # cast-ok: legacy manager is untyped
match name:
case "input" | "async_input":
registry(name).append(callback)
case "success":
manager.add_litellm_success_callback(callback)
case "async_success":
manager.add_litellm_async_success_callback(callback)
case "failure":
manager.add_litellm_failure_callback(callback)
case "async_failure":
manager.add_litellm_async_failure_callback(callback)
case "callbacks":
raise KeyError(name)
def remove_registry(name: RegistryName, callback: CallbackTarget) -> None:
target: Final = registry(name)
for index in range(len(target) - 1, -1, -1):
if target[index] is callback or target[index] == callback:
del target[index]
return
def logger_fn(callback: object) -> None:
from litellm import utils
utils.user_logger_fn = callback # rebind-ok: legacy module global read by pre_call
def breadcrumb(kwargs: Mapping[str, object]) -> None:
from litellm import utils
add_breadcrumb: Final = cast(
Callable[..., None] | None, utils.add_breadcrumb
) # cast-ok: legacy sentry hook is untyped
if add_breadcrumb is None:
return
import litellm
from litellm.litellm_core_utils import core_helpers
deep_copy: Final = cast( # cast-ok: legacy helper is untyped
Callable[[dict[str, object]], dict[str, object]],
core_helpers.safe_deep_copy, # pyright: ignore[reportUnknownMemberType] # legacy helper
)
try:
copied: dict[str, object] = deep_copy(dict(kwargs))
except Exception: # noqa: BLE001 # legacy breadcrumb falls back to the live mapping
copied = dict(kwargs)
hidden: Final = frozenset(("messages", "input", "prompt")) if litellm.turn_off_message_logging else frozenset[str]()
details: Final = {key: value for key, value in copied.items() if key not in hidden}
add_breadcrumb(category="litellm.llm_call", message=f"Keyword Args: {details}", level="info")
def prepare_environment() -> None:
from litellm import utils
utils.custom_llm_setup()
def applied_guardrails(kwargs: Mapping[str, object]) -> list[str]:
from litellm.utils import get_applied_guardrails
return get_applied_guardrails(dict(kwargs))
def build_logging(
*,
call_type: str,
model: str | None,
kwargs: dict[str, object],
start_time: datetime.datetime,
asynchronous: bool,
dynamic_success: Sequence[CallbackTarget] | None,
dynamic_async_success: Sequence[CallbackTarget] | None,
dynamic_failure: Sequence[CallbackTarget] | None,
guardrails: Sequence[str],
) -> Logging:
from litellm.litellm_core_utils.cached_imports import get_litellm_logging_class
function_id: Final = kwargs.get("id")
metadata: Final = kwargs.get("metadata")
trace_id: Final = kwargs.get("litellm_trace_id")
factory: Final = cast(_LoggingFactory, get_litellm_logging_class()) # cast-ok: legacy constructor is untyped
logger: Final = factory(
model=model,
messages="default-message-value",
stream=False,
litellm_call_id=str(kwargs["litellm_call_id"]),
litellm_trace_id=trace_id if isinstance(trace_id, str) else None,
function_id=function_id if isinstance(function_id, str) else "",
call_type=call_type,
start_time=start_time,
dynamic_success_callbacks=list(dynamic_success) if dynamic_success is not None else None,
dynamic_failure_callbacks=list(dynamic_failure) if dynamic_failure is not None else None,
dynamic_async_success_callbacks=list(dynamic_async_success) if dynamic_async_success is not None else None,
dynamic_async_failure_callbacks=None,
kwargs=kwargs,
applied_guardrails=list(guardrails),
supports_correlation_logging=asynchronous,
)
litellm_metadata: Final = kwargs.get("litellm_metadata")
litellm_params: Final[dict[str, object]] = {
"api_base": "",
**({"metadata": kwargs["metadata"]} if "metadata" in kwargs else {}),
**(
{
"litellm_metadata": litellm_metadata,
**(
{} if metadata else {"metadata": dict(cast(Mapping[str, object], litellm_metadata))}
), # cast-ok: isinstance narrows only to dict[Unknown, Unknown]
}
if isinstance(litellm_metadata, dict)
else {}
),
}
update: Final = cast(_EnvironmentUpdater, logger.update_environment_variables) # cast-ok: legacy method is untyped
update(
model=model,
user="",
optional_params={},
litellm_params=litellm_params,
stream_options=kwargs.get("stream_options"),
)
return logger

View file

@ -8,8 +8,7 @@ import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import legacy
from litellm.rust_bridge import bindings, configuration
from litellm.rust_bridge.leaves import record_post_call, record_pre_call
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, update_logging
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, post_call, pre_call, update_logging
@pytest.fixture(autouse=True)
@ -66,26 +65,25 @@ def test_logging_redacts_views_and_preserves_opaque_arguments_and_pricing() -> N
def test_logging_callbacks_receive_captured_payload_roots_and_propagate_errors() -> None:
logger: Final = Mock(log_raw_request_response=False, logger_fn=None, model_call_details={})
logger: Final = Mock()
body: Final[dict[str, object]] = {"document": "original"}
headers: Final = {"authorization": "key"}
response: Final = object()
record_pre_call(logger, api_key="key", body=body, headers=headers, url="https://provider")
record_post_call(logger, original_response=response, body=body, headers=headers)
logger._pre_call.assert_called_once_with(
pre_call(logger, "key", body, headers, "https://provider")
post_call(logger, response, body, headers)
logger.pre_call.assert_called_once_with(
input="OCR document processing",
api_key="key",
model=None,
additional_args={"complete_input_dict": body, "headers": headers, "api_base": "https://provider"},
)
for callback in (logger._pre_call, logger.record_post_call):
for callback in (logger.pre_call, logger.post_call):
assert callback.call_args.kwargs["additional_args"]["complete_input_dict"] is body
assert callback.call_args.kwargs["additional_args"]["headers"] is headers
assert logger.record_post_call.call_args.kwargs["original_response"] is response
assert logger.post_call.call_args.kwargs["original_response"] is response
failure: Final = RuntimeError("callback failed")
failing_logger: Final = Mock(_pre_call=Mock(side_effect=failure))
failing_logger: Final = Mock(pre_call=Mock(side_effect=failure))
with pytest.raises(RuntimeError) as caught:
record_pre_call(failing_logger, api_key=None, body=body, headers=headers, url="https://provider")
pre_call(failing_logger, None, body, headers, "https://provider")
assert caught.value is failure

View file

@ -1,169 +0,0 @@
import datetime
from collections.abc import Iterator
from typing import Final
import pytest
import litellm
from litellm import utils
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.rust_bridge.loader import native_bridge_available
pytestmark = pytest.mark.skipif(not native_bridge_available(), reason="requires the Rust extension")
REGISTRIES: Final = (
"input_callback",
"_async_input_callback",
"success_callback",
"_async_success_callback",
"failure_callback",
"_async_failure_callback",
"callbacks",
)
@pytest.fixture
def clean_registries(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
for name in REGISTRIES:
monkeypatch.setattr(litellm, name, [])
monkeypatch.setattr(utils, "callback_list", [])
monkeypatch.setattr(utils, "user_logger_fn", None)
monkeypatch.setenv("OPENMETER_API_KEY", "test")
monkeypatch.setenv("OPENMETER_API_ENDPOINT", "http://127.0.0.1:9")
yield
class SyncLogger(CustomLogger):
pass
class AsyncOnly(CustomLogger):
pass
def sync_fn(*args: object, **kwargs: object) -> None:
del args, kwargs
async def async_fn(*args: object, **kwargs: object) -> None:
del args, kwargs
def snapshot() -> dict[str, list[object]]:
return {name: list(getattr(litellm, name)) for name in REGISTRIES}
def run_legacy(kwargs: dict[str, object]) -> tuple[Logging, dict[str, object]]:
logger, prepared = utils.function_setup(
"ocr",
utils.Rules(),
datetime.datetime.now(),
is_async_call=False,
**{"litellm_call_id": "legacy", **kwargs},
)
assert isinstance(logger, Logging)
return logger, prepared
def run_native(kwargs: dict[str, object]) -> tuple[Logging, dict[str, object]]:
from litellm.rust_bridge import _native
logger, prepared = _native._debug_setup(
"ocr", (), {"litellm_call_id": "native", **kwargs}, datetime.datetime.now(), False
)
assert isinstance(logger, Logging)
return logger, prepared
@pytest.mark.parametrize(
"globals_before, kwargs",
[
({}, {}),
({"callbacks": [SyncLogger()]}, {}),
({"callbacks": [async_fn]}, {}),
({}, {"callbacks": [SyncLogger(), sync_fn]}),
({"success_callback": [async_fn, "openmeter", sync_fn]}, {}),
({"failure_callback": [async_fn, sync_fn]}, {}),
({"input_callback": [async_fn, sync_fn]}, {}),
({}, {"success_callback": [sync_fn, async_fn, "s3", "dynamodb"], "failure_callback": [sync_fn]}),
({"callbacks": [SyncLogger()]}, {"callbacks": [SyncLogger()], "success_callback": [async_fn]}),
],
ids=[
"empty",
"global-custom-logger",
"global-async-callable",
"dynamic-callbacks",
"success-safety-net",
"failure-safety-net",
"input-safety-net",
"per-call-success-failure-split",
"mixed",
],
)
def test_native_setup_registry_side_effects_match_function_setup(
clean_registries: None, globals_before: dict[str, list[object]], kwargs: dict[str, object]
) -> None:
for name, values in globals_before.items():
getattr(litellm, name).extend(values)
legacy_logger, legacy_kwargs = run_legacy(
{key: list(value) if isinstance(value, list) else value for key, value in kwargs.items()}
)
legacy_snapshot: Final = snapshot()
legacy_bootstrap: Final = list(utils.callback_list or [])
for name in REGISTRIES:
getattr(litellm, name).clear()
utils.callback_list = []
for name, values in globals_before.items():
getattr(litellm, name).extend(values)
native_logger, native_kwargs = run_native(
{key: list(value) if isinstance(value, list) else value for key, value in kwargs.items()}
)
assert snapshot() == legacy_snapshot
assert sorted(map(repr, utils.callback_list or [])) == sorted(map(repr, legacy_bootstrap))
assert set(native_kwargs) - {"litellm_call_id"} == set(legacy_kwargs) - {"litellm_call_id"}
for attribute in (
"dynamic_success_callbacks",
"dynamic_async_success_callbacks",
"dynamic_failure_callbacks",
"dynamic_async_failure_callbacks",
"call_type",
"stream",
"model",
):
assert getattr(native_logger, attribute) == getattr(legacy_logger, attribute), attribute
def test_native_setup_honours_caller_supplied_logging_object(clean_registries: None) -> None:
class Supplied(Logging):
pass
supplied: Final = Supplied(
model="mistral/mistral-ocr-latest",
messages=[],
stream=False,
call_type="ocr",
start_time=datetime.datetime.now(),
litellm_call_id="supplied",
function_id="",
)
litellm.callbacks.append(SyncLogger())
logger, kwargs = run_native({"litellm_logging_obj": supplied, "callbacks": [SyncLogger()]})
assert logger is supplied
assert kwargs["callbacks"] is not None
assert litellm.success_callback == []
def test_native_setup_records_logger_fn_and_metadata(clean_registries: None) -> None:
def logger_fn(details: object) -> None:
del details
logger, kwargs = run_native(
{"model": "mistral/mistral-ocr-latest", "logger_fn": logger_fn, "metadata": {"source": "test"}}
)
assert utils.user_logger_fn is logger_fn
assert logger.litellm_params["metadata"] == {"source": "test"}
assert logger.model == "mistral/mistral-ocr-latest"
assert kwargs["metadata"] == {"source": "test"}

View file

@ -10,7 +10,6 @@ import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import (
OCR_DOCUMENT,
OCR_RESPONSE,
@ -19,6 +18,7 @@ from tests.test_litellm_rust.support.requests import (
request_body,
request_headers,
)
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
pytestmark = pytest.mark.requires_rust_extension
@ -309,7 +309,6 @@ async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context
asynchronous: bool,
) -> None:
from contextvars import ContextVar
context: Final = ContextVar("azure-token-context", default="missing")
context.set("caller")
caller_thread: Final = threading.current_thread()
@ -338,7 +337,9 @@ async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context
"callbacks": [Edit()],
}
response: Final = (
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
await call_native_aocr(ocr_server, **arguments)
if asynchronous
else call_native_ocr(ocr_server, **arguments)
)
assert response.pages[0].markdown == "native OCR response"
assert observations == ["token", "pre_call"]
@ -367,7 +368,9 @@ async def test_native_azure_ocr_token_provider_can_make_nested_native_ocr_call(
"azure_ad_token_provider": provider,
}
response: Final = (
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
await call_native_aocr(ocr_server, **arguments)
if asynchronous
else call_native_ocr(ocr_server, **arguments)
)
assert response.pages[0].markdown == "native OCR response"
assert calls == ["token"]
@ -422,9 +425,7 @@ async def test_native_azure_ocr_releases_token_provider_after_terminal_outcome(
) -> None:
import gc
import weakref
from tests.test_litellm_rust.support.callback_recorder import drain_logging
class Provider:
def __call__(self) -> str:
if outcome == "failure":
@ -470,238 +471,3 @@ async def test_native_azure_ocr_releases_token_provider_after_terminal_outcome(
await asyncio.sleep(0)
gc.collect()
assert reference() is None
FORBIDDEN_ORCHESTRATION: Final = (
"pre_call",
"post_call",
"success_handler",
"async_success_handler",
"failure_handler",
"async_failure_handler",
"_success_handler_body",
"_async_success_handler_body",
"_failure_handler_body",
"_async_failure_handler_body",
"dispatch_success_handlers",
"dispatch_failure_handlers",
"handle_sync_success_callbacks_for_async_calls",
)
@pytest.fixture
def legacy_orchestration_disabled(monkeypatch: pytest.MonkeyPatch) -> list[str]:
from litellm import utils
from litellm.litellm_core_utils.litellm_logging import Logging
reached: Final[list[str]] = []
def forbid(name: str):
def method(self, *args, **kwargs):
reached.append(name)
raise AssertionError(f"legacy orchestration reached: {name}")
return method
for name in FORBIDDEN_ORCHESTRATION:
monkeypatch.setattr(Logging, name, forbid(name))
def forbidden_setup(*args, **kwargs):
reached.append("function_setup")
raise AssertionError("legacy orchestration reached: function_setup")
monkeypatch.setattr(utils, "function_setup", forbidden_setup)
return reached
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_ocr_success_runs_integrations_without_legacy_orchestration(
ocr_server: RecordingServer, legacy_orchestration_disabled: list[str], asynchronous: bool
) -> None:
recorder: Final = RecordingLogger()
arguments: Final = {"callbacks": [recorder]}
response: Final = (
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
)
assert response.pages[0].markdown == "native OCR response"
success_event: Final = "async_log_success_event" if asynchronous else "log_success_event"
events: Final = await recorder.wait_for_async(success_event)
assert legacy_orchestration_disabled == []
assert recorder.names.count("log_pre_api_call") == 1
assert events[0].kwargs["standard_logging_object"]["status"] == "success"
assert events[0].kwargs["response_cost"] is not None
assert events[0].kwargs["litellm_params"]["api_base"] == f"{ocr_server.base_url}/v1/ocr"
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_ocr_failure_runs_integrations_without_legacy_orchestration(
ocr_server: RecordingServer, legacy_orchestration_disabled: list[str], asynchronous: bool
) -> None:
ocr_server.enqueue(ResponseSpec(body={"message": "provider unavailable"}, status=500))
recorder: Final = RecordingLogger()
arguments: Final = {"callbacks": [recorder]}
with pytest.raises(litellm.InternalServerError) as caught:
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
assert legacy_orchestration_disabled == []
failures: Final = tuple(event for event in recorder.events if event.name.endswith("log_failure_event"))
assert [event.name for event in failures] == (
["log_failure_event", "async_log_failure_event"] if asynchronous else ["log_failure_event"]
)
assert all(event.kwargs["exception"] is caught.value for event in failures)
assert all(event.kwargs["standard_logging_object"]["status"] == "failure" for event in failures)
assert "log_success_event" not in recorder.names
def test_native_ocr_success_hooks_all_run_before_any_success_dispatch(ocr_server: RecordingServer) -> None:
order: Final = []
finished: Final = threading.Event()
class Hooked(CustomLogger):
def __init__(self, name: str) -> None:
super().__init__()
self.name = name
def logging_hook(self, kwargs, result, call_type):
order.append(("hook", self.name))
kwargs[f"seen-by-{self.name}"] = True
return kwargs, result
def log_success_event(self, kwargs, response_obj, start_time, end_time):
order.append(("log", self.name, kwargs.get("seen-by-a"), kwargs.get("seen-by-b")))
if self.name == "b":
finished.set()
call_native_ocr_with_callbacks(ocr_server, [Hooked("a"), Hooked("b")])
assert finished.wait(10)
assert order == [("hook", "a"), ("hook", "b"), ("log", "a", True, True), ("log", "b", True, True)]
def test_native_ocr_hook_failure_is_contained_and_target_still_dispatches(ocr_server: RecordingServer) -> None:
order: Final = []
finished: Final = threading.Event()
class Broken(CustomLogger):
def logging_hook(self, kwargs, result, call_type):
raise RuntimeError("hook failed")
def log_success_event(self, kwargs, response_obj, start_time, end_time):
order.append("broken-log")
class Healthy(CustomLogger):
def logging_hook(self, kwargs, result, call_type):
order.append("healthy-hook")
return kwargs, result
def log_success_event(self, kwargs, response_obj, start_time, end_time):
order.append("healthy-log")
finished.set()
call_native_ocr_with_callbacks(ocr_server, [Broken(), Healthy()])
assert finished.wait(10)
assert order == ["healthy-hook", "broken-log", "healthy-log"]
@pytest.mark.asyncio
async def test_native_aocr_hook_replacement_of_result_reaches_success_dispatch(ocr_server: RecordingServer) -> None:
replacement: Final = object()
observed: Final = []
class Replace(CustomLogger):
async def async_logging_hook(self, kwargs, result, call_type):
return kwargs, replacement
class Observe(CustomLogger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
observed.append(response_obj)
recorder: Final = RecordingLogger()
await call_native_aocr_with_callbacks(ocr_server, [Replace(), Observe(), recorder])
await recorder.wait_for_async("async_log_success_event")
assert observed == [replacement]
@pytest.mark.asyncio
async def test_native_aocr_shared_logging_object_dispatches_success_once(ocr_server: RecordingServer) -> None:
from litellm.litellm_core_utils.litellm_logging import Logging
ocr_server.expected_requests = 2
recorder: Final = RecordingLogger()
litellm.callbacks.append(recorder)
logger: Final = Logging(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr",
start_time=__import__("datetime").datetime.now(),
litellm_call_id="shared",
function_id="shared",
)
logger.dynamic_async_success_callbacks = [recorder]
await call_native_aocr(ocr_server, litellm_logging_obj=logger)
await call_native_aocr(ocr_server, litellm_logging_obj=logger)
await recorder.wait_for_async("async_log_success_event")
from tests.test_litellm_rust.support.callback_recorder import drain_logging
await drain_logging()
assert logger.model_call_details["has_logged_async_success"] is True
assert recorder.names.count("async_log_success_event") == 1
def test_native_ocr_logging_preparation_failure_does_not_fail_request(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.litellm_core_utils import litellm_logging
def broken_payload(*args, **kwargs):
raise RuntimeError("payload unavailable")
monkeypatch.setattr(litellm_logging, "get_standard_logging_object_payload", broken_payload)
unraisable: Final = []
monkeypatch.setattr(__import__("sys"), "unraisablehook", lambda event: unraisable.append(event.exc_value))
recorder: Final = RecordingLogger()
response: Final = call_native_ocr_with_callbacks(ocr_server, [recorder])
events: Final = recorder.wait_for("log_success_event")
assert response.pages[0].markdown == "native OCR response"
assert len(events) == 1
assert events[0].thread is not threading.current_thread()
assert any(str(error) == "payload unavailable" for error in unraisable)
def test_native_ocr_writes_success_marker_and_honours_existing_marker(
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.rust_bridge import setup as native_setup
ocr_server.expected_requests = 2
loggers: Final = []
original_build: Final = native_setup.build_logging
def build_logging(**kwargs):
logger = original_build(**kwargs)
if loggers:
logger.model_call_details["has_logged_sync_success"] = True
loggers.append(logger)
return logger
monkeypatch.setattr(native_setup, "build_logging", build_logging)
recorder: Final = RecordingLogger()
call_native_ocr_with_callbacks(ocr_server, [recorder])
recorder.wait_for("log_success_event")
assert loggers[0].model_call_details["has_logged_sync_success"] is True
call_native_ocr_with_callbacks(ocr_server, [recorder])
from litellm.litellm_core_utils.thread_pool_executor import executor
executor.submit(lambda: None).result(10)
assert recorder.names.count("log_success_event") == 1
assert recorder.names.count("logging_hook") == 1

View file

@ -967,18 +967,26 @@ async def test_terminal_registration_added_during_http_is_observed(
@pytest.fixture
def created_loggers(monkeypatch: pytest.MonkeyPatch) -> list[Logging]:
from litellm.rust_bridge import setup as native_setup
from litellm import utils
original_build: Final = native_setup.build_logging
original_setup: Final = utils.function_setup
loggers: Final[list[Logging]] = []
def build_logging(**kwargs: object) -> Logging:
logger: Final = original_build(**kwargs) # pyright: ignore[reportArgumentType] # passthrough of the factory signature
def setup(
call_type: str,
rules: utils.Rules,
start: datetime.datetime,
*args: object,
is_async_call: bool = True,
**kwargs: object,
) -> tuple[Logging, dict[str, object]]:
logger, prepared = original_setup(call_type, rules, start, *args, is_async_call=is_async_call, **kwargs)
assert isinstance(logger, Logging)
setattr(logger, "_defer_async_logging", True)
loggers.append(logger)
return logger
return logger, prepared
monkeypatch.setattr(native_setup, "build_logging", build_logging)
monkeypatch.setattr(utils, "function_setup", setup)
return loggers