mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Revert "refactor(python-bridge): classify callbacks natively instead of function_setup"
This reverts commit f8190bbe80.
This commit is contained in:
parent
58be89ffba
commit
92e1a2ded7
21 changed files with 763 additions and 3476 deletions
|
|
@ -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::{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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(())
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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());
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
}
|
||||
|
|
@ -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>> {
|
||||
|
|
|
|||
|
|
@ -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)?))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue