mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(python-bridge): dispatch OCR callbacks through the native cursor
Success, failure and request callbacks on the native OCR path no longer go through Logging.pre_call, post_call, success_handler or failure_handler. Core's DispatchCursor selects each target and leaf method; the bridge invokes it directly or through one labeled leaf in litellm.rust_bridge.leaves. Delivery follows the family: request and sync failure run inline, deployment hooks and async failure are awaited by drive(), sync success runs as one grouped WorkerJob on the executor, async success is a coroutine enqueued on the logging worker with the deferred gate held natively in DeferredSuccess. PrepareLogging (cost, standard payload, redaction) runs inside the dispatch delivery and never fails the request; Finalize keeps only public response metadata. Dedup markers stay on the shared logging object and are read as an eligibility fact. Caller-supplied Logging instances take a separate compat path that calls their own handlers. Tests prove the two-pass hook order, per-target hook containment, hook result replacement, marker write and honour, best-effort preparation, and a positive run with every legacy orchestration entry point patched to raise.
This commit is contained in:
parent
f8190bbe80
commit
f0b877532c
11 changed files with 2105 additions and 569 deletions
|
|
@ -28,102 +28,10 @@ impl PythonLogger {
|
|||
pub(super) fn defer_success(
|
||||
&self,
|
||||
py: Python<'_>,
|
||||
pending: Py<super::PendingLogging>,
|
||||
pending: Py<super::DeferredSuccess>,
|
||||
) -> 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) fn finalize(
|
||||
|
|
@ -185,59 +93,3 @@ impl DeploymentHooks {
|
|||
.map(Bound::unbind)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn logger_resolves_each_callback_at_invocation_and_preserves_arguments() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
calls = []
|
||||
response, start, end = object(), object(), object()
|
||||
class Logger:
|
||||
@property
|
||||
def handle_sync_success_callbacks_for_async_calls(self):
|
||||
generation = len(calls)
|
||||
def callback(*args):
|
||||
assert args == (response, start, end)
|
||||
calls.append(generation)
|
||||
return callback
|
||||
logger = Logger()
|
||||
"#
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let logger: PythonLogger = locals
|
||||
.get_item("logger")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract()
|
||||
.unwrap();
|
||||
let response = Some(locals.get_item("response").unwrap().unwrap().unbind());
|
||||
let start = locals.get_item("start").unwrap().unwrap().unbind();
|
||||
let end = Some(locals.get_item("end").unwrap().unwrap().unbind());
|
||||
for _ in 0..2 {
|
||||
logger
|
||||
.sync_success_for_async_call(py, &response, &start, &end)
|
||||
.unwrap();
|
||||
}
|
||||
assert_eq!(
|
||||
locals
|
||||
.get_item("calls")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract::<Vec<usize>>()
|
||||
.unwrap(),
|
||||
[0, 1]
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
168
litellm-rust/crates/python-bridge/src/lifecycle/compat.rs
Normal file
168
litellm-rust/crates/python-bridge/src/lifecycle/compat.rs
Normal file
|
|
@ -0,0 +1,168 @@
|
|||
use litellm_core::call_lifecycle::{
|
||||
CallbackFamily, Delivery, ReleaseGate, SuccessFacts, plan_success,
|
||||
};
|
||||
use pyo3::prelude::*;
|
||||
|
||||
use super::bindings::PythonLogger;
|
||||
use super::{PythonCallState, dispatch, missing_state};
|
||||
|
||||
pub(super) fn dispatch_success(
|
||||
py: Python<'_>,
|
||||
state: &PythonCallState,
|
||||
logger: &PythonLogger,
|
||||
) -> PyResult<()> {
|
||||
let facts = SuccessFacts {
|
||||
asynchronous: state.asynchronous,
|
||||
internal: state.internal,
|
||||
fallbacks: !state
|
||||
.kwargs
|
||||
.bind(py)
|
||||
.get_item("fallbacks")?
|
||||
.is_none_or(|value| value.is_none()),
|
||||
deferred: logger.defers_async_logging(py),
|
||||
sync_target_kinds: sync_kinds(py, logger)?,
|
||||
};
|
||||
let leaves = dispatch::leaves(py)?;
|
||||
let object = logger.object(py);
|
||||
for selected in plan_success(&facts) {
|
||||
match (selected.family, selected.delivery, selected.gate) {
|
||||
(CallbackFamily::SyncSuccess, Delivery::Worker, _) => {
|
||||
let handler = object.getattr("success_handler")?;
|
||||
let bound = py.import("functools")?.getattr("partial")?.call1((
|
||||
handler,
|
||||
&state.response,
|
||||
&state.start,
|
||||
&state.end,
|
||||
))?;
|
||||
leaves.getattr("submit_worker")?.call1((bound,))?;
|
||||
}
|
||||
(CallbackFamily::AsyncSuccess, Delivery::Background, ReleaseGate::Immediate) => {
|
||||
let coroutine = object.call_method1(
|
||||
"async_success_handler",
|
||||
(&state.response, &state.start, &state.end),
|
||||
)?;
|
||||
let enqueue = leaves.getattr("enqueue_background")?.call1((&coroutine,));
|
||||
if enqueue.is_err()
|
||||
&& let Err(error) = coroutine.call_method0("close")
|
||||
{
|
||||
error.write_unraisable(py, Some(&coroutine));
|
||||
}
|
||||
enqueue?;
|
||||
}
|
||||
(CallbackFamily::AsyncSuccess, Delivery::Background, ReleaseGate::Deferred) => {
|
||||
let pending = Py::new(
|
||||
py,
|
||||
SuppliedDeferred {
|
||||
logger: Some(logger.clone_ref(py)),
|
||||
response: state.response.as_ref().map(|value| value.clone_ref(py)),
|
||||
start: state.start.clone_ref(py),
|
||||
end: state.end.as_ref().map(|value| value.clone_ref(py)),
|
||||
},
|
||||
)?;
|
||||
object.setattr("_native_pending_logging", pending)?;
|
||||
}
|
||||
_ => return Err(missing_state()),
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn sync_kinds(
|
||||
py: Python<'_>,
|
||||
logger: &PythonLogger,
|
||||
) -> PyResult<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);
|
||||
}
|
||||
}
|
||||
649
litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs
Normal file
649
litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs
Normal file
|
|
@ -0,0 +1,649 @@
|
|||
use litellm_core::call_lifecycle::{
|
||||
CallbackFamily, CallbackId, CallbackInvocation, CallbackKind, CallbackMethod, Delivery,
|
||||
DispatchCursor, DispatchFacts, DispatchStep, InvocationOutcome, LoggedMarker,
|
||||
};
|
||||
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError};
|
||||
use pyo3::gc::{PyTraverseError, PyVisit};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyString;
|
||||
|
||||
use super::bindings::PythonLogger;
|
||||
|
||||
const LEAVES: &str = "litellm.rust_bridge.leaves";
|
||||
|
||||
pub(super) fn leaves(py: Python<'_>) -> PyResult<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))
|
||||
}
|
||||
|
|
@ -1,26 +1,30 @@
|
|||
use std::sync::Arc;
|
||||
use std::task::Poll;
|
||||
|
||||
use futures_util::future::{AbortHandle, Abortable};
|
||||
#[cfg(test)]
|
||||
use litellm_core::call_lifecycle::host::HostCallFuture;
|
||||
use litellm_core::call_lifecycle::host::{
|
||||
HostCall as NativeCall, HostCallStep as NativeCallStep, HostFailure, HostPhase, HostStep,
|
||||
};
|
||||
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError};
|
||||
use pyo3::gc::{PyTraverseError, PyVisit};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
|
||||
use futures_util::future::{AbortHandle, Abortable};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
|
||||
use litellm_core::call_lifecycle::host::{
|
||||
HostCall as NativeCall, HostCallStep as NativeCallStep, HostFailure, HostPhase, HostStep,
|
||||
};
|
||||
use litellm_core::call_lifecycle::{
|
||||
CallbackFamily, Delivery, ReleaseGate, SuccessFacts, plan_failure, plan_success,
|
||||
};
|
||||
|
||||
mod arguments;
|
||||
mod bindings;
|
||||
mod compat;
|
||||
mod dispatch;
|
||||
mod handle;
|
||||
mod preparation;
|
||||
mod setup;
|
||||
|
||||
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
|
||||
pub(crate) use arguments::{BoundArguments, Signature};
|
||||
use bindings::DeploymentHooks;
|
||||
pub(crate) use bindings::PythonLogger;
|
||||
|
|
@ -300,6 +304,7 @@ pub(crate) struct PythonCallState {
|
|||
pub error: Option<Py<PyBaseException>>,
|
||||
pub asynchronous: bool,
|
||||
pub internal: bool,
|
||||
pub supplied: bool,
|
||||
pub call_type: &'static str,
|
||||
}
|
||||
|
||||
|
|
@ -389,6 +394,7 @@ impl PythonCallState {
|
|||
error: None,
|
||||
asynchronous,
|
||||
internal: false,
|
||||
supplied: false,
|
||||
call_type,
|
||||
})
|
||||
}
|
||||
|
|
@ -412,6 +418,7 @@ impl PythonCallState {
|
|||
)?;
|
||||
self.logger = Some(result.logger);
|
||||
self.kwargs = result.kwargs;
|
||||
self.supplied = result.supplied;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
|
@ -431,6 +438,16 @@ impl PythonCallState {
|
|||
)
|
||||
}
|
||||
|
||||
pub fn dispatch_request(&self, py: Python<'_>, family: CallbackFamily) -> PyResult<()> {
|
||||
dispatch::dispatch_request(
|
||||
py,
|
||||
dispatch::RequestJob {
|
||||
logger: self.logger()?,
|
||||
family,
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub fn dispatch_success(&self, py: Python<'_>) -> PyResult<()> {
|
||||
match self.try_dispatch_success(py) {
|
||||
Err(error) if error.is_instance_of::<PyException>(py) => {
|
||||
|
|
@ -441,40 +458,71 @@ 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()?;
|
||||
let pending = || PendingSuccess {
|
||||
logger: logger.clone_ref(py),
|
||||
response: self.response.as_ref().map(|value| value.clone_ref(py)),
|
||||
start: self.start.clone_ref(py),
|
||||
end: self.end.as_ref().map(|value| value.clone_ref(py)),
|
||||
if self.supplied {
|
||||
return compat::dispatch_success(py, self, logger);
|
||||
}
|
||||
let (sync_targets, sync_ids) =
|
||||
dispatch::family_targets(py, logger, CallbackFamily::SyncSuccess)?;
|
||||
let facts = SuccessFacts {
|
||||
asynchronous: self.asynchronous,
|
||||
internal: self.internal,
|
||||
fallbacks: !self
|
||||
.kwargs
|
||||
.bind(py)
|
||||
.get_item("fallbacks")?
|
||||
.is_none_or(|value| value.is_none()),
|
||||
deferred: logger.defers_async_logging(py),
|
||||
sync_target_kinds: sync_targets.kinds(&sync_ids),
|
||||
};
|
||||
if !self.asynchronous {
|
||||
pending().sync(py)
|
||||
} else {
|
||||
if !self.internal
|
||||
&& self
|
||||
.kwargs
|
||||
.bind(py)
|
||||
.get_item("fallbacks")?
|
||||
.is_none_or(|value| value.is_none())
|
||||
{
|
||||
if logger.defers_async_logging(py) {
|
||||
for selected in plan_success(&facts) {
|
||||
let runner = dispatch::Runner::start(py, self.job(py, selected.family)?)?;
|
||||
match (selected.delivery, selected.gate) {
|
||||
(Delivery::Worker, _) => {
|
||||
let job = Py::new(py, dispatch::WorkerJob::new(runner))?;
|
||||
dispatch::leaves(py)?
|
||||
.getattr("submit_worker")?
|
||||
.call1((job,))?;
|
||||
}
|
||||
(Delivery::Background, ReleaseGate::Immediate) => {
|
||||
DeferredSuccess::release(py, runner)?;
|
||||
}
|
||||
(Delivery::Background, ReleaseGate::Deferred) => {
|
||||
logger.defer_success(
|
||||
py,
|
||||
Py::new(
|
||||
py,
|
||||
PendingLogging {
|
||||
pending: Some(pending()),
|
||||
DeferredSuccess {
|
||||
runner: Some(runner),
|
||||
},
|
||||
)?,
|
||||
)?;
|
||||
} else {
|
||||
pending().asynchronous(py)?;
|
||||
}
|
||||
(Delivery::Inline | Delivery::Await, _) => return Err(missing_state()),
|
||||
}
|
||||
logger.sync_success_for_async_call(py, &self.response, &self.start, &self.end)
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn dispatch_failure(
|
||||
|
|
@ -482,19 +530,39 @@ impl PythonCallState {
|
|||
py: Python<'_>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Option<Py<PyAny>>> {
|
||||
if self.logger.is_none() || (self.asynchronous && self.internal) {
|
||||
if self.logger.is_none() || self.error.is_none() {
|
||||
return Ok(None);
|
||||
}
|
||||
let Some(error) = &self.error else {
|
||||
let phase = if asynchronous {
|
||||
HostPhase::AsyncFailure
|
||||
} else {
|
||||
HostPhase::Failure
|
||||
};
|
||||
let Some(family) = plan_failure(phase, self.asynchronous, self.internal) else {
|
||||
return Ok(None);
|
||||
};
|
||||
self.logger()?
|
||||
.failure(py, error, &self.start, &self.end, asynchronous)
|
||||
if self.supplied {
|
||||
return compat::dispatch_failure(py, self, family);
|
||||
}
|
||||
let mut runner = dispatch::Runner::start(py, self.job(py, family)?)?;
|
||||
match family.delivery() {
|
||||
Delivery::Inline => match runner.resume(py, None)? {
|
||||
dispatch::Step::Done => Ok(None),
|
||||
dispatch::Step::Await(_) => Err(missing_state()),
|
||||
},
|
||||
Delivery::Await => Ok(Some(dispatch::coroutine(py, runner)?)),
|
||||
Delivery::Worker | Delivery::Background => Err(missing_state()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn cleanup(&mut self, py: Python<'_>) {
|
||||
if let Some(logger) = self.logger.take()
|
||||
&& let Err(error) = logger.restore_context(py)
|
||||
&& let Err(error) = dispatch::leaves(py).and_then(|leaves| {
|
||||
leaves
|
||||
.getattr("restore_correlation_context")?
|
||||
.call1((logger.object(py),))
|
||||
.map(|_| ())
|
||||
})
|
||||
{
|
||||
error.write_unraisable(py, None);
|
||||
}
|
||||
|
|
@ -530,58 +598,53 @@ impl PythonCallState {
|
|||
}
|
||||
}
|
||||
|
||||
struct PendingSuccess {
|
||||
logger: PythonLogger,
|
||||
response: Option<Py<PyAny>>,
|
||||
start: Py<PyAny>,
|
||||
end: Option<Py<PyAny>>,
|
||||
}
|
||||
|
||||
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>,
|
||||
struct DeferredSuccess {
|
||||
runner: Option<dispatch::Runner>,
|
||||
}
|
||||
|
||||
impl DeferredSuccess {
|
||||
fn release(py: Python<'_>, runner: dispatch::Runner) -> PyResult<()> {
|
||||
let coroutine = dispatch::coroutine(py, runner)?;
|
||||
let enqueue = dispatch::leaves(py)?
|
||||
.getattr("enqueue_background")?
|
||||
.call1((&coroutine,));
|
||||
if enqueue.is_err()
|
||||
&& let Err(error) = coroutine.call_method0(py, "close")
|
||||
{
|
||||
error.write_unraisable(py, Some(coroutine.bind(py)));
|
||||
}
|
||||
enqueue.map(|_| ())
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl PendingLogging {
|
||||
impl DeferredSuccess {
|
||||
fn __call__(slf: &Bound<'_, Self>, py: Python<'_>) -> PyResult<()> {
|
||||
let pending = slf.borrow_mut().pending.take();
|
||||
if let Some(pending) = pending {
|
||||
match pending.asynchronous(py) {
|
||||
Err(error) if error.is_instance_of::<PyException>(py) => {
|
||||
error.write_unraisable(py, Some(pending.logger.object(py)));
|
||||
}
|
||||
result => return result,
|
||||
let runner = slf.borrow_mut().runner.take();
|
||||
let Some(runner) = runner else {
|
||||
return Ok(());
|
||||
};
|
||||
let logger = runner.logger(py).clone().unbind();
|
||||
match Self::release(py, runner) {
|
||||
Err(error) if error.is_instance_of::<PyException>(py) => {
|
||||
error.write_unraisable(py, Some(logger.bind(py)));
|
||||
Ok(())
|
||||
}
|
||||
result => result,
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn __traverse__(&self, visit: pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> {
|
||||
if let Some(pending) = &self.pending {
|
||||
pending.logger.traverse(&visit)?;
|
||||
visit.call(&pending.response)?;
|
||||
visit.call(&pending.start)?;
|
||||
visit.call(&pending.end)?;
|
||||
match &self.runner {
|
||||
Some(runner) => runner.traverse(&visit),
|
||||
None => Ok(()),
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn close(slf: &Bound<'_, Self>) {
|
||||
let pending = slf.borrow_mut().pending.take();
|
||||
drop(pending);
|
||||
let runner = slf.borrow_mut().runner.take();
|
||||
drop(runner);
|
||||
}
|
||||
|
||||
fn __clear__(slf: &Bound<'_, Self>) {
|
||||
|
|
@ -592,14 +655,39 @@ impl PendingLogging {
|
|||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
use pyo3::types::PyDict;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use litellm_core::call_lifecycle::host::HostCallFuture;
|
||||
|
||||
static PYTHON_GLOBALS: Mutex<()> = Mutex::new(());
|
||||
|
||||
fn install_logging_worker(py: Python<'_>, worker: &Bound<'_, PyAny>) -> PyResult<()> {
|
||||
py.import("litellm.litellm_core_utils.logging_worker")?
|
||||
.setattr("GLOBAL_LOGGING_WORKER", worker)
|
||||
fn load_lifecycle_module(py: Python<'_>) -> Bound<'_, PyModule> {
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
import sys
|
||||
import types
|
||||
for name in ("litellm", "litellm.rust_bridge"):
|
||||
sys.modules.setdefault(name, types.ModuleType(name))
|
||||
"#
|
||||
),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
let source = std::ffi::CString::new(include_str!(
|
||||
"../../../../../litellm/rust_bridge/lifecycle.py"
|
||||
))
|
||||
.unwrap();
|
||||
PyModule::from_code(
|
||||
py,
|
||||
&source,
|
||||
pyo3::ffi::c_str!("lifecycle.py"),
|
||||
pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
struct RetainingHost {
|
||||
|
|
@ -771,17 +859,7 @@ mod tests {
|
|||
.unwrap_or_else(|error| error.into_inner());
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let source = std::ffi::CString::new(include_str!(
|
||||
"../../../../../litellm/rust_bridge/lifecycle.py"
|
||||
))
|
||||
.unwrap();
|
||||
PyModule::from_code(
|
||||
py,
|
||||
&source,
|
||||
pyo3::ffi::c_str!("lifecycle.py"),
|
||||
pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"),
|
||||
)
|
||||
.unwrap();
|
||||
load_lifecycle_module(py);
|
||||
let route = SyntheticRoute(
|
||||
PythonCallState::new(
|
||||
py,
|
||||
|
|
@ -817,17 +895,7 @@ mod tests {
|
|||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
py.import("asyncio").unwrap();
|
||||
let source = std::ffi::CString::new(include_str!(
|
||||
"../../../../../litellm/rust_bridge/lifecycle.py"
|
||||
))
|
||||
.unwrap();
|
||||
let module = PyModule::from_code(
|
||||
py,
|
||||
&source,
|
||||
pyo3::ffi::c_str!("lifecycle.py"),
|
||||
pyo3::ffi::c_str!("litellm.rust_bridge.lifecycle"),
|
||||
)
|
||||
.unwrap();
|
||||
let module = load_lifecycle_module(py);
|
||||
let locals = PyDict::new(py);
|
||||
locals
|
||||
.set_item("drive", module.getattr("drive").unwrap())
|
||||
|
|
@ -935,64 +1003,11 @@ assert reference() is None
|
|||
error: None,
|
||||
asynchronous,
|
||||
internal: false,
|
||||
supplied: false,
|
||||
call_type: "test",
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn success_dispatch_reports_ordinary_failures_without_replacing_response() {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
import sys
|
||||
|
||||
response = object()
|
||||
failure = ValueError('terminal diagnostic')
|
||||
diagnostics = []
|
||||
old_hook = sys.unraisablehook
|
||||
sys.unraisablehook = lambda event: diagnostics.append(event.exc_value)
|
||||
|
||||
class Logger:
|
||||
def handle_sync_success_callbacks_for_async_calls(self, *args):
|
||||
raise failure
|
||||
|
||||
logger = Logger()
|
||||
"#
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let response = locals.get_item("response").unwrap().unwrap().unbind();
|
||||
let mut lifecycle_state = state(
|
||||
py,
|
||||
locals.get_item("logger").unwrap().unwrap().unbind(),
|
||||
response.clone_ref(py),
|
||||
true,
|
||||
);
|
||||
lifecycle_state.internal = true;
|
||||
lifecycle_state.dispatch_success(py).unwrap();
|
||||
assert!(lifecycle_state.response.as_ref().unwrap().is(&response));
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
assert diagnostics == [failure]
|
||||
sys.unraisablehook = old_hook
|
||||
"#
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn retained_failure_preserves_exception_identity() {
|
||||
Python::initialize();
|
||||
|
|
@ -1008,201 +1023,6 @@ sys.unraisablehook = old_hook
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn deferred_logging_uses_call_context_and_allows_reentry_once() {
|
||||
let _guard = PYTHON_GLOBALS
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
import sys
|
||||
import types
|
||||
from contextvars import ContextVar
|
||||
|
||||
litellm = types.ModuleType('litellm')
|
||||
core_utils = types.ModuleType('litellm.litellm_core_utils')
|
||||
logging_worker = types.ModuleType('litellm.litellm_core_utils.logging_worker')
|
||||
litellm.litellm_core_utils = core_utils
|
||||
core_utils.logging_worker = logging_worker
|
||||
sys.modules['litellm'] = litellm
|
||||
sys.modules['litellm.litellm_core_utils'] = core_utils
|
||||
sys.modules['litellm.litellm_core_utils.logging_worker'] = logging_worker
|
||||
|
||||
marker = ContextVar('marker', default='unset')
|
||||
observed = []
|
||||
|
||||
class Coroutine:
|
||||
def close(self):
|
||||
observed.append('closed')
|
||||
|
||||
class Worker:
|
||||
def ensure_initialized_and_enqueue(self, coroutine):
|
||||
observed.append(marker.get())
|
||||
pending()
|
||||
coroutine.close()
|
||||
|
||||
class Logger:
|
||||
def async_success_handler(self, *args):
|
||||
observed.append('created')
|
||||
return Coroutine()
|
||||
|
||||
worker = Worker()
|
||||
logger = Logger()
|
||||
"#
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
install_logging_worker(py, &locals.get_item("worker").unwrap().unwrap()).unwrap();
|
||||
let pending = Py::new(
|
||||
py,
|
||||
PendingLogging {
|
||||
pending: Some(PendingSuccess {
|
||||
logger: locals
|
||||
.get_item("logger")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract()
|
||||
.unwrap(),
|
||||
response: Some(py.None()),
|
||||
start: py.None(),
|
||||
end: Some(py.None()),
|
||||
}),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
locals.set_item("pending", &pending).unwrap();
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
marker.set('call')
|
||||
pending()
|
||||
pending()
|
||||
assert observed == ['created', 'call', 'closed']
|
||||
"#
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn deferred_logging_close_is_reentry_safe_and_invalidates_aliases() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
observed = []
|
||||
|
||||
class Retained:
|
||||
def __del__(self):
|
||||
observed.append('finalized')
|
||||
alias()
|
||||
|
||||
class Logger:
|
||||
def async_success_handler(self, *args):
|
||||
observed.append('enqueued')
|
||||
|
||||
logger = Logger()
|
||||
retained = Retained()
|
||||
"#
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let pending = Py::new(
|
||||
py,
|
||||
PendingLogging {
|
||||
pending: Some(PendingSuccess {
|
||||
logger: locals
|
||||
.get_item("logger")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract()
|
||||
.unwrap(),
|
||||
response: Some(locals.get_item("retained").unwrap().unwrap().unbind()),
|
||||
start: py.None(),
|
||||
end: None,
|
||||
}),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
locals.set_item("pending", &pending).unwrap();
|
||||
locals.set_item("alias", &pending).unwrap();
|
||||
locals.del_item("retained").unwrap();
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
pending.close()
|
||||
alias()
|
||||
assert observed == ['finalized']
|
||||
"#
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn deferred_logging_collects_cycles_through_typed_logger() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
pyo3::ffi::c_str!("class Logger: pass\nlogger = Logger()"),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let pending = Py::new(
|
||||
py,
|
||||
PendingLogging {
|
||||
pending: Some(PendingSuccess {
|
||||
logger: locals
|
||||
.get_item("logger")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract()
|
||||
.unwrap(),
|
||||
response: None,
|
||||
start: py.None(),
|
||||
end: None,
|
||||
}),
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
locals.set_item("pending", pending).unwrap();
|
||||
py.run(
|
||||
pyo3::ffi::c_str!(
|
||||
r#"
|
||||
import gc
|
||||
import weakref
|
||||
logger.pending = pending
|
||||
reference = weakref.ref(logger)
|
||||
del logger, pending
|
||||
gc.collect()
|
||||
assert reference() is None
|
||||
"#
|
||||
),
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn coroutine_collects_cycles_retained_by_bridge_host() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
|
|
@ -235,6 +235,7 @@ fn split_dynamic<'py>(
|
|||
pub(super) struct Setup {
|
||||
pub logger: PythonLogger,
|
||||
pub kwargs: Py<PyDict>,
|
||||
pub supplied: bool,
|
||||
}
|
||||
|
||||
pub(super) fn setup(
|
||||
|
|
@ -259,6 +260,7 @@ pub(super) fn setup(
|
|||
return Ok(Setup {
|
||||
logger: supplied.extract()?,
|
||||
kwargs: kwargs.unbind(),
|
||||
supplied: true,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
@ -330,5 +332,6 @@ pub(super) fn setup(
|
|||
Ok(Setup {
|
||||
logger: logger.extract()?,
|
||||
kwargs: kwargs.unbind(),
|
||||
supplied: false,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,8 +6,10 @@ use litellm_core::ocr::LiteLLMOcrResponse;
|
|||
use litellm_core::ocr::hooks::OcrDuringCallRequest;
|
||||
use litellm_python_interop::to_py_preserving_errors as to_py;
|
||||
|
||||
use litellm_core::call_lifecycle::CallbackFamily;
|
||||
|
||||
use super::host::PythonPayload;
|
||||
use crate::lifecycle::PythonLogger;
|
||||
use crate::lifecycle::{PythonCallState, PythonLogger};
|
||||
|
||||
pub(super) fn update_logging(
|
||||
py: Python<'_>,
|
||||
|
|
@ -32,37 +34,63 @@ pub(super) fn update_logging(
|
|||
|
||||
pub(super) fn pre_call(
|
||||
py: Python<'_>,
|
||||
logger: &PythonLogger,
|
||||
state: &PythonCallState,
|
||||
request: &OcrDuringCallRequest,
|
||||
payload: &PythonPayload,
|
||||
) -> PyResult<()> {
|
||||
py.import("litellm.rust_bridge.ocr")?
|
||||
.getattr("pre_call")?
|
||||
.call1((
|
||||
logger.object(py),
|
||||
request.api_key.as_deref(),
|
||||
&payload.body,
|
||||
&payload.headers,
|
||||
&request.url,
|
||||
))?;
|
||||
Ok(())
|
||||
let logger = state.logger()?;
|
||||
if state.supplied {
|
||||
let additional_args = PyDict::new(py);
|
||||
additional_args.set_item("complete_input_dict", &payload.body)?;
|
||||
additional_args.set_item("headers", &payload.headers)?;
|
||||
additional_args.set_item("api_base", &request.url)?;
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("input", "OCR document processing")?;
|
||||
kwargs.set_item("api_key", request.api_key.as_deref())?;
|
||||
kwargs.set_item("additional_args", additional_args)?;
|
||||
logger
|
||||
.object(py)
|
||||
.call_method("pre_call", (), Some(&kwargs))?;
|
||||
return Ok(());
|
||||
}
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("api_key", request.api_key.as_deref())?;
|
||||
kwargs.set_item("body", &payload.body)?;
|
||||
kwargs.set_item("headers", &payload.headers)?;
|
||||
kwargs.set_item("url", &request.url)?;
|
||||
py.import("litellm.rust_bridge.leaves")?
|
||||
.getattr("record_pre_call")?
|
||||
.call((logger.object(py),), Some(&kwargs))?;
|
||||
state.dispatch_request(py, CallbackFamily::RequestPreCall)
|
||||
}
|
||||
|
||||
pub(super) fn post_call(
|
||||
py: Python<'_>,
|
||||
logger: &PythonLogger,
|
||||
state: &PythonCallState,
|
||||
original_response: &Value,
|
||||
payload: &PythonPayload,
|
||||
) -> PyResult<()> {
|
||||
py.import("litellm.rust_bridge.ocr")?
|
||||
.getattr("post_call")?
|
||||
.call1((
|
||||
logger.object(py),
|
||||
to_py(py, original_response)?,
|
||||
&payload.body,
|
||||
&payload.headers,
|
||||
))?;
|
||||
Ok(())
|
||||
let logger = state.logger()?;
|
||||
if state.supplied {
|
||||
let additional_args = PyDict::new(py);
|
||||
additional_args.set_item("complete_input_dict", &payload.body)?;
|
||||
additional_args.set_item("headers", &payload.headers)?;
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("original_response", to_py(py, original_response)?)?;
|
||||
kwargs.set_item("additional_args", additional_args)?;
|
||||
logger
|
||||
.object(py)
|
||||
.call_method("post_call", (), Some(&kwargs))?;
|
||||
return Ok(());
|
||||
}
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("original_response", to_py(py, original_response)?)?;
|
||||
kwargs.set_item("body", &payload.body)?;
|
||||
kwargs.set_item("headers", &payload.headers)?;
|
||||
py.import("litellm.rust_bridge.leaves")?
|
||||
.getattr("record_post_call")?
|
||||
.call((logger.object(py),), Some(&kwargs))?;
|
||||
state.dispatch_request(py, CallbackFamily::RequestPostCall)
|
||||
}
|
||||
|
||||
pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
|
||||
|
|
|
|||
|
|
@ -128,7 +128,7 @@ impl PythonOcrHost {
|
|||
&retained.secret_fields,
|
||||
)?;
|
||||
let payload = PythonPayload::from_request(py, &request)?;
|
||||
callbacks::pre_call(py, logger, &request, &payload)?;
|
||||
callbacks::pre_call(py, &self.state, &request, &payload)?;
|
||||
let request = payload.write_back(py, request)?;
|
||||
self.retained_mut()?.payload = Some(payload);
|
||||
Ok(request)
|
||||
|
|
@ -144,12 +144,7 @@ impl PythonOcrHost {
|
|||
.payload
|
||||
.as_ref()
|
||||
.ok_or_else(missing_state)?;
|
||||
callbacks::post_call(
|
||||
py,
|
||||
self.state.logger()?,
|
||||
&request.original_response,
|
||||
payload,
|
||||
)?;
|
||||
callbacks::post_call(py, &self.state, &request.original_response, payload)?;
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
|
|
|
|||
815
litellm/rust_bridge/leaves.py
Normal file
815
litellm/rust_bridge/leaves.py
Normal file
|
|
@ -0,0 +1,815 @@
|
|||
"""Leaf helpers invoked by the native callback dispatcher.
|
||||
|
||||
Rust selects every target, delivery and sequence. Nothing here chooses a callback,
|
||||
reads a registry, or decides whether an event fires. Each function performs one
|
||||
integration-specific or interop-specific action on values Rust hands it.
|
||||
|
||||
Labeled leaf helpers pending native migration (see rust-callback-inventory.md):
|
||||
`prepare_success_logging`, `prepare_failure_logging`, `dispatch_named_success`,
|
||||
`dispatch_named_failure`, `dispatch_callable`. They are deleted when cost and
|
||||
payload construction move into core and when each string integration becomes a
|
||||
CustomLogger.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import json
|
||||
import traceback
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import ( # noqa: TID251 # narrows the untyped legacy Logging and CustomLogger surfaces once
|
||||
TYPE_CHECKING,
|
||||
Final,
|
||||
Literal,
|
||||
Protocol,
|
||||
cast,
|
||||
)
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
TerminalFamily = Literal["sync_success", "async_success", "sync_failure", "async_failure"]
|
||||
Family = Literal["request", TerminalFamily]
|
||||
Details = dict[str, object]
|
||||
Timestamp = datetime.datetime
|
||||
LegacyCall = Callable[..., object]
|
||||
LegacyAsyncCall = Callable[..., Awaitable[None]]
|
||||
|
||||
|
||||
class LoggerView(Protocol):
|
||||
model: str | None
|
||||
messages: object
|
||||
call_type: str
|
||||
start_time: Timestamp
|
||||
litellm_call_id: str
|
||||
completion_start_time: Timestamp | None
|
||||
model_call_details: Details
|
||||
log_raw_request_response: bool
|
||||
standard_callback_dynamic_params: object
|
||||
standard_built_in_tools_params: object
|
||||
|
||||
def record_api_call_start_time(self) -> None: ...
|
||||
|
||||
def record_post_call(
|
||||
self, original_response: object, input: object, api_key: object, additional_args: Details
|
||||
) -> None: ...
|
||||
|
||||
def should_run_callback(self, callback: object, litellm_params: Details, event_hook: str) -> bool: ...
|
||||
|
||||
def _pre_call(self, input: str, api_key: str | None, model: str | None, additional_args: Details) -> None: ...
|
||||
|
||||
def _print_llm_call_debugging_log(self, api_base: str, headers: Details, additional_args: Details) -> None: ...
|
||||
|
||||
def _get_request_curl_command(
|
||||
self, api_base: str, headers: Details | None, additional_args: Details, data: object
|
||||
) -> str: ...
|
||||
|
||||
def _get_masked_api_base(self, api_base: str) -> str: ...
|
||||
|
||||
def _get_raw_request_body(self, data: object) -> Details: ...
|
||||
|
||||
def _get_masked_headers(self, headers: Details) -> Details: ...
|
||||
|
||||
def _response_cost_calculator(self, result: object) -> float | None: ...
|
||||
|
||||
def _build_standard_logging_payload(
|
||||
self, init_response_obj: object, start_time: Timestamp, end_time: Timestamp
|
||||
) -> object: ...
|
||||
|
||||
def _handle_callback_failure(self, callback: object) -> None: ...
|
||||
|
||||
|
||||
class IntegrationView(Protocol):
|
||||
def log_pre_api_call(self, model: str | None, messages: object, kwargs: Details) -> None: ...
|
||||
|
||||
def log_post_api_call(
|
||||
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp | None
|
||||
) -> None: ...
|
||||
|
||||
def log_success_event(
|
||||
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp
|
||||
) -> None: ...
|
||||
|
||||
def log_failure_event(
|
||||
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp
|
||||
) -> None: ...
|
||||
|
||||
def async_log_success_event(
|
||||
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp
|
||||
) -> Awaitable[None]: ...
|
||||
|
||||
def async_log_failure_event(
|
||||
self, kwargs: Details, response_obj: object, start_time: Timestamp, end_time: Timestamp
|
||||
) -> Awaitable[None]: ...
|
||||
|
||||
def logging_hook(self, kwargs: Details, result: object, call_type: str) -> tuple[Details, object]: ...
|
||||
|
||||
def async_logging_hook(
|
||||
self, kwargs: Details, result: object, call_type: str
|
||||
) -> Awaitable[tuple[Details, object]]: ...
|
||||
|
||||
def redact_standard_logging_payload_from_model_call_details(self, model_call_details: Details) -> Details: ...
|
||||
|
||||
def log_input_event(
|
||||
self, model: str | None, messages: object, kwargs: Details, print_verbose: LegacyCall, callback_func: LegacyCall
|
||||
) -> None: ...
|
||||
|
||||
def log_event(
|
||||
self,
|
||||
kwargs: Details,
|
||||
response_obj: object,
|
||||
start_time: Timestamp,
|
||||
end_time: Timestamp,
|
||||
print_verbose: LegacyCall,
|
||||
callback_func: LegacyCall,
|
||||
) -> None: ...
|
||||
|
||||
def async_log_event(
|
||||
self,
|
||||
kwargs: Details,
|
||||
response_obj: object,
|
||||
start_time: Timestamp,
|
||||
end_time: Timestamp,
|
||||
print_verbose: LegacyCall,
|
||||
callback_func: LegacyCall,
|
||||
) -> Awaitable[None]: ...
|
||||
|
||||
|
||||
def _logger(logger: Logging) -> LoggerView:
|
||||
return cast(LoggerView, logger) # cast-ok: legacy Logging is untyped; this protocol names the attributes we read
|
||||
|
||||
|
||||
def _integration(callback: CustomLogger) -> IntegrationView:
|
||||
return cast(IntegrationView, callback) # cast-ok: legacy CustomLogger methods are untyped
|
||||
|
||||
|
||||
def _legacy_module() -> Mapping[str, object]:
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
return cast(
|
||||
Mapping[str, object], vars(litellm_logging)
|
||||
) # cast-ok: module globals hold the legacy integration singletons
|
||||
|
||||
|
||||
def _print_verbose() -> LegacyCall:
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
return cast(LegacyCall, litellm_logging.print_verbose) # cast-ok: legacy debug printer is untyped
|
||||
|
||||
|
||||
def _method(target: object, name: str) -> LegacyCall:
|
||||
return cast(LegacyCall, getattr(target, name)) # cast-ok: legacy integration singletons are untyped
|
||||
|
||||
|
||||
def _async_method(target: object, name: str) -> LegacyAsyncCall:
|
||||
return cast(LegacyAsyncCall, getattr(target, name)) # cast-ok: legacy integration singletons are untyped
|
||||
|
||||
|
||||
def _redact_string(value: str) -> str:
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
return cast(Callable[[str], str], litellm_logging._redact_string)(value) # pyright: ignore[reportPrivateUsage] # cast-ok: legacy helper
|
||||
|
||||
|
||||
def _redact_result(details: Details, result: object) -> object:
|
||||
from litellm.litellm_core_utils import redact_messages
|
||||
|
||||
redact: Final = cast(LegacyCall, redact_messages.redact_message_input_output_from_logging) # cast-ok: legacy helper
|
||||
return redact(model_call_details=details, result=result)
|
||||
|
||||
|
||||
def record_pre_call(
|
||||
logger: Logging,
|
||||
*,
|
||||
api_key: str | None,
|
||||
body: Details,
|
||||
headers: dict[str, str],
|
||||
url: str,
|
||||
) -> None:
|
||||
view: Final = _logger(logger)
|
||||
additional_args: Final[Details] = {"complete_input_dict": body, "headers": headers, "api_base": url}
|
||||
view._pre_call(input="OCR document processing", api_key=api_key, model=None, additional_args=additional_args) # pyright: ignore[reportPrivateUsage] # legacy state writer
|
||||
view._print_llm_call_debugging_log(api_base=url, headers=dict(headers), additional_args=additional_args) # pyright: ignore[reportPrivateUsage] # legacy debug output
|
||||
_capture_raw_request(view, additional_args)
|
||||
_run_logger_fn(logger)
|
||||
view.record_api_call_start_time()
|
||||
|
||||
|
||||
def _capture_raw_request(view: LoggerView, additional_args: Details) -> None:
|
||||
import litellm
|
||||
from litellm.types.utils import RawRequestTypedDict
|
||||
|
||||
if not (view.log_raw_request_response or litellm.log_raw_request_response):
|
||||
return
|
||||
details: Final = view.model_call_details
|
||||
params: Final = cast(Details, details.get("litellm_params") or {}) # cast-ok: legacy nested dict
|
||||
metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict
|
||||
params.setdefault("metadata", metadata)
|
||||
if litellm.turn_off_message_logging:
|
||||
metadata["raw_request"] = "redacted by litellm. 'litellm.turn_off_message_logging=True'"
|
||||
return
|
||||
api_base: Final = str(additional_args.get("api_base") or "")
|
||||
headers: Final = cast(Details, additional_args.get("headers") or {}) # cast-ok: legacy nested dict
|
||||
body: Final = additional_args.get("complete_input_dict", {})
|
||||
try:
|
||||
curl: Final = view._get_request_curl_command( # pyright: ignore[reportPrivateUsage] # legacy debug formatter
|
||||
api_base=api_base, headers=headers, additional_args=additional_args, data=body
|
||||
)
|
||||
metadata["raw_request"] = _redact_string(str(curl))
|
||||
details["raw_request_typed_dict"] = RawRequestTypedDict(
|
||||
raw_request_api_base=view._get_masked_api_base(api_base), # pyright: ignore[reportPrivateUsage] # legacy masking
|
||||
raw_request_body=view._get_raw_request_body(body), # pyright: ignore[reportPrivateUsage] # legacy masking
|
||||
raw_request_headers=view._get_masked_headers(headers), # pyright: ignore[reportPrivateUsage] # legacy masking
|
||||
error=None,
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 # raw-request capture is best effort by contract
|
||||
details["raw_request_typed_dict"] = RawRequestTypedDict(error=str(error))
|
||||
metadata["raw_request"] = _redact_string(f"Unable to Log raw request: {error}")
|
||||
|
||||
|
||||
def _run_logger_fn(logger: Logging) -> None:
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
logger_fn: Final = cast(
|
||||
Callable[[Details], object] | None, getattr(logger, "logger_fn", None)
|
||||
) # cast-ok: user hook is untyped
|
||||
if not callable(logger_fn):
|
||||
return
|
||||
try:
|
||||
logger_fn(_logger(logger).model_call_details)
|
||||
except Exception as error: # noqa: BLE001 # user logger_fn failures never fail the request
|
||||
verbose_logger.exception("LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging %s", error)
|
||||
|
||||
|
||||
def record_post_call(logger: Logging, *, original_response: object, body: Details, headers: dict[str, str]) -> None:
|
||||
view: Final = _logger(logger)
|
||||
serialized: Final = (
|
||||
json.dumps(original_response, default=str) if isinstance(original_response, dict) else original_response
|
||||
)
|
||||
view.record_post_call(
|
||||
original_response=serialized,
|
||||
input=None,
|
||||
api_key=None,
|
||||
additional_args={"complete_input_dict": body, "headers": headers},
|
||||
)
|
||||
_run_logger_fn(logger)
|
||||
_redact_result(view.model_call_details, serialized)
|
||||
|
||||
|
||||
def log_pre_api_call(logger: Logging, callback: CustomLogger) -> None:
|
||||
view: Final = _logger(logger)
|
||||
_integration(callback).log_pre_api_call(model=view.model, messages=view.messages, kwargs=view.model_call_details)
|
||||
|
||||
|
||||
def log_post_api_call(logger: Logging, callback: CustomLogger) -> None:
|
||||
view: Final = _logger(logger)
|
||||
_integration(callback).log_post_api_call(
|
||||
kwargs=view.model_call_details, response_obj=None, start_time=view.start_time, end_time=None
|
||||
)
|
||||
|
||||
|
||||
def dispatch_named_request(logger: Logging, name: str, event: Literal["pre_api_call", "post_api_call"]) -> None:
|
||||
view: Final = _logger(logger)
|
||||
module: Final = _legacy_module()
|
||||
if name == "supabase" and event == "pre_api_call" and (client := module.get("supabaseClient")) is not None:
|
||||
details: Final = view.model_call_details
|
||||
_method(client, "input_log_event")(
|
||||
model=view.model,
|
||||
messages=view.messages,
|
||||
end_user=details.get("user", "default"),
|
||||
litellm_call_id=details["litellm_call_id"],
|
||||
print_verbose=_print_verbose(),
|
||||
)
|
||||
if name == "sentry" and (add_breadcrumb := module.get("add_breadcrumb")) is not None:
|
||||
cast(LegacyCall, add_breadcrumb)( # cast-ok: legacy sentry hook
|
||||
category="litellm.llm_call", message=f"Model Call Details {event}: {view.model_call_details}", level="info"
|
||||
)
|
||||
|
||||
|
||||
def dispatch_callable_request(logger: Logging, callback: LegacyCall) -> None:
|
||||
custom: Final = _legacy_module().get("customLogger")
|
||||
if not isinstance(custom, CustomLogger):
|
||||
return
|
||||
view: Final = _logger(logger)
|
||||
_integration(custom).log_input_event(
|
||||
model=view.model,
|
||||
messages=view.messages,
|
||||
kwargs=view.model_call_details,
|
||||
print_verbose=_print_verbose(),
|
||||
callback_func=callback,
|
||||
)
|
||||
|
||||
|
||||
def report_target_failure(logger: Logging, callback: object, family: Family, error: BaseException) -> None:
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
verbose_logger.error(
|
||||
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while %s logging with %s: %s",
|
||||
family,
|
||||
callback,
|
||||
"".join(traceback.format_exception(error)),
|
||||
)
|
||||
capture: Final = _legacy_module().get("capture_exception")
|
||||
if capture is not None and family in ("request", "sync_success", "sync_failure"):
|
||||
cast(LegacyCall, capture)(error) # cast-ok: legacy sentry hook
|
||||
if family not in ("request", "sync_failure"):
|
||||
_logger(logger)._handle_callback_failure(callback=callback) # pyright: ignore[reportPrivateUsage] # legacy prometheus counter
|
||||
|
||||
|
||||
def prepare_success_logging(logger: Logging, response: object, start_time: Timestamp, end_time: Timestamp) -> object:
|
||||
from litellm.litellm_core_utils.litellm_logging import emit_standard_logging_payload
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
view: Final = _logger(logger)
|
||||
details: Final = view.model_call_details
|
||||
if view.completion_start_time is None:
|
||||
view.completion_start_time = end_time
|
||||
details["completion_start_time"] = end_time
|
||||
details["log_event_type"] = "successful_api_call"
|
||||
details["end_time"] = end_time
|
||||
details["cache_hit"] = None
|
||||
hidden: Final = cast(Details, getattr(response, "_hidden_params", None) or {}) # cast-ok: legacy response attribute
|
||||
params: Final = cast(Details | None, details.get("litellm_params")) # cast-ok: legacy nested dict
|
||||
if hidden and params is not None:
|
||||
metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict
|
||||
params["metadata"] = metadata
|
||||
metadata["hidden_params"] = hidden
|
||||
existing: Final = details.get("response_cost")
|
||||
if "response_cost" in hidden:
|
||||
details["response_cost"] = hidden["response_cost"]
|
||||
elif existing is None or existing == 0:
|
||||
details["response_cost"] = view._response_cost_calculator(result=response) # pyright: ignore[reportPrivateUsage] # labeled leaf: native cost pending
|
||||
payload: Final = view._build_standard_logging_payload(response, start_time, end_time) # pyright: ignore[reportPrivateUsage] # labeled leaf: native payload pending
|
||||
details["standard_logging_object"] = payload
|
||||
if payload is not None:
|
||||
emit_standard_logging_payload(
|
||||
cast(StandardLoggingPayload, payload)
|
||||
) # cast-ok: legacy builder returns the payload TypedDict
|
||||
return _redact_result(details, response)
|
||||
|
||||
|
||||
def prepare_failure_logging(
|
||||
logger: Logging, exception: BaseException, start_time: Timestamp, end_time: Timestamp
|
||||
) -> str:
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
formatted: Final = "".join(traceback.format_exception(exception))
|
||||
view: Final = _logger(logger)
|
||||
details: Final = view.model_call_details
|
||||
if details.get("exception") is exception and details.get("standard_logging_object") is not None:
|
||||
return formatted
|
||||
details["log_event_type"] = "failed_api_call"
|
||||
details["exception"] = exception
|
||||
details["traceback_exception"] = _redact_string(formatted)
|
||||
details["end_time"] = end_time
|
||||
details.setdefault("original_response", None)
|
||||
if details.get("combined_usage_object") is None:
|
||||
details["response_cost"] = 0
|
||||
headers: Final = getattr(exception, "headers", None)
|
||||
if isinstance(headers, dict):
|
||||
params: Final = cast(Details, details.setdefault("litellm_params", {})) # cast-ok: legacy nested dict
|
||||
metadata: Final = cast(Details, params.get("metadata") or {}) # cast-ok: legacy nested dict
|
||||
metadata.update(cast(Details, headers)) # cast-ok: exception headers are a plain dict
|
||||
build: Final = cast(
|
||||
LegacyCall, litellm_logging.get_standard_logging_object_payload
|
||||
) # cast-ok: labeled leaf: native payload pending
|
||||
details["standard_logging_object"] = build(
|
||||
kwargs=details,
|
||||
init_response_obj={},
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
logging_obj=logger,
|
||||
status="failure",
|
||||
error_str=_redact_string(str(exception)),
|
||||
original_exception=exception,
|
||||
standard_built_in_tools_params=view.standard_built_in_tools_params,
|
||||
)
|
||||
return formatted
|
||||
|
||||
|
||||
_EVENT_HOOKS: Final[Mapping[TerminalFamily, str]] = {
|
||||
"sync_success": "success_handler",
|
||||
"async_success": "async_success_handler",
|
||||
"sync_failure": "failure_handler",
|
||||
"async_failure": "async_failure_handler",
|
||||
}
|
||||
|
||||
|
||||
def should_run_callback(logger: Logging, callback: object, family: TerminalFamily) -> bool:
|
||||
view: Final = _logger(logger)
|
||||
params: Final = cast(Details, view.model_call_details.get("litellm_params") or {}) # cast-ok: legacy nested dict
|
||||
return view.should_run_callback(callback=callback, litellm_params=params, event_hook=_EVENT_HOOKS[family])
|
||||
|
||||
|
||||
def should_run_guardrail_hook(logger: Logging, callback: object) -> bool:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
if not isinstance(callback, CustomGuardrail):
|
||||
return True
|
||||
decide: Final = cast(LegacyCall, callback.should_run_guardrail) # cast-ok: legacy guardrail method is untyped
|
||||
return decide(data=_logger(logger).model_call_details, event_type=GuardrailEventHooks.logging_only) is True
|
||||
|
||||
|
||||
def logging_hook(logger: Logging, callback: CustomLogger, result: object) -> object:
|
||||
view: Final = _logger(logger)
|
||||
details, replaced = _integration(callback).logging_hook(
|
||||
kwargs=view.model_call_details, result=result, call_type=view.call_type
|
||||
)
|
||||
view.model_call_details = details
|
||||
return replaced
|
||||
|
||||
|
||||
async def async_logging_hook(logger: Logging, callback: CustomLogger, result: object) -> object:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils import redact_messages
|
||||
|
||||
view: Final = _logger(logger)
|
||||
redact: Final = cast(
|
||||
LegacyCall, redact_messages.redact_message_input_output_from_custom_logger
|
||||
) # cast-ok: legacy helper
|
||||
redacted: Final = (
|
||||
result
|
||||
if isinstance(callback, CustomGuardrail)
|
||||
else redact(result=result, litellm_logging_obj=logger, custom_logger=callback)
|
||||
)
|
||||
details, replaced = await _integration(callback).async_logging_hook(
|
||||
kwargs=view.model_call_details, result=redacted, call_type=view.call_type
|
||||
)
|
||||
view.model_call_details = details
|
||||
return replaced
|
||||
|
||||
|
||||
def mark_logged(logger: Logging, marker: str) -> None:
|
||||
_logger(logger).model_call_details[marker] = True
|
||||
|
||||
|
||||
def already_logged(logger: Logging, marker: str) -> bool:
|
||||
return _logger(logger).model_call_details.get(marker, False) is True
|
||||
|
||||
|
||||
def log_success_event(
|
||||
logger: Logging, callback: CustomLogger, response: object, start_time: Timestamp, end_time: Timestamp
|
||||
) -> None:
|
||||
_integration(callback).log_success_event(
|
||||
kwargs=_logger(logger).model_call_details, response_obj=response, start_time=start_time, end_time=end_time
|
||||
)
|
||||
|
||||
|
||||
def async_log_success_event(
|
||||
logger: Logging, callback: CustomLogger, response: object, start_time: Timestamp, end_time: Timestamp
|
||||
) -> Awaitable[None]:
|
||||
from litellm.litellm_core_utils import redact_messages
|
||||
|
||||
integration: Final = _integration(callback)
|
||||
details: Final = integration.redact_standard_logging_payload_from_model_call_details(
|
||||
model_call_details=_logger(logger).model_call_details
|
||||
)
|
||||
redact: Final = cast(
|
||||
Callable[..., Details], redact_messages.redact_streaming_responses_for_custom_logger
|
||||
) # cast-ok: legacy helper
|
||||
view: Final = redact(model_call_details=details, custom_logger=callback)
|
||||
return integration.async_log_success_event(
|
||||
kwargs=view, response_obj=response, start_time=start_time, end_time=end_time
|
||||
)
|
||||
|
||||
|
||||
def log_failure_event(logger: Logging, callback: CustomLogger, start_time: Timestamp, end_time: Timestamp) -> None:
|
||||
_integration(callback).log_failure_event(
|
||||
kwargs=_logger(logger).model_call_details, response_obj=None, start_time=start_time, end_time=end_time
|
||||
)
|
||||
|
||||
|
||||
def async_log_failure_event(
|
||||
logger: Logging, callback: CustomLogger, start_time: Timestamp, end_time: Timestamp
|
||||
) -> Awaitable[None]:
|
||||
return _integration(callback).async_log_failure_event(
|
||||
kwargs=_logger(logger).model_call_details, response_obj=None, start_time=start_time, end_time=end_time
|
||||
)
|
||||
|
||||
|
||||
def _custom_logger_singleton() -> IntegrationView:
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
existing: Final = _legacy_module().get("customLogger")
|
||||
if isinstance(existing, CustomLogger):
|
||||
return _integration(existing)
|
||||
created: Final = CustomLogger()
|
||||
litellm_logging.customLogger = created # pyright: ignore[reportAttributeAccessIssue] # legacy module global
|
||||
return _integration(created)
|
||||
|
||||
|
||||
def dispatch_callable(
|
||||
logger: Logging,
|
||||
callback: LegacyCall,
|
||||
family: TerminalFamily,
|
||||
response: object,
|
||||
start_time: Timestamp,
|
||||
end_time: Timestamp,
|
||||
) -> Awaitable[None] | None:
|
||||
custom: Final = _custom_logger_singleton()
|
||||
details: Final = _logger(logger).model_call_details
|
||||
match family:
|
||||
case "sync_success" | "sync_failure":
|
||||
custom.log_event(
|
||||
kwargs=details,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=_print_verbose(),
|
||||
callback_func=callback,
|
||||
)
|
||||
return None
|
||||
case "async_success" | "async_failure":
|
||||
return custom.async_log_event(
|
||||
kwargs=details,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=_print_verbose(),
|
||||
callback_func=callback,
|
||||
)
|
||||
|
||||
|
||||
_SUCCESS_SINGLETONS: Final[Mapping[str, str]] = {
|
||||
"promptlayer": "promptLayerLogger",
|
||||
"supabase": "supabaseClient",
|
||||
"wandb": "weightsBiasesLogger",
|
||||
"logfire": "logfireLogger",
|
||||
"lunary": "lunaryLogger",
|
||||
"helicone": "heliconeLogger",
|
||||
"greenscale": "greenscaleLogger",
|
||||
"athina": "athinaLogger",
|
||||
"traceloop": "traceloopLogger",
|
||||
"s3": "s3Logger",
|
||||
"openmeter": "openMeterLogger",
|
||||
}
|
||||
|
||||
|
||||
def dispatch_named_success(
|
||||
logger: Logging, name: str, response: object, start_time: Timestamp, end_time: Timestamp
|
||||
) -> Awaitable[None] | None:
|
||||
view: Final = _logger(logger)
|
||||
details: Final = view.model_call_details
|
||||
print_verbose: Final = _print_verbose()
|
||||
integration: Final = _legacy_module().get(_SUCCESS_SINGLETONS.get(name, ""))
|
||||
without_response: Final = {key: value for key, value in details.items() if key != "original_response"}
|
||||
match name:
|
||||
case "promptlayer" | "wandb" | "athina" if integration is not None:
|
||||
_method(integration, "log_event")(
|
||||
kwargs=details,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case "logfire" if integration is not None:
|
||||
from litellm.integrations.logfire_logger import LogfireLevel
|
||||
|
||||
_method(integration, "log_event")(
|
||||
kwargs=without_response,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=print_verbose,
|
||||
level=LogfireLevel.INFO.value,
|
||||
)
|
||||
case "greenscale" if integration is not None:
|
||||
_method(integration, "log_event")(
|
||||
kwargs=without_response,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case "supabase" if integration is not None:
|
||||
_method(integration, "log_event")(
|
||||
model=view.model,
|
||||
messages=view.messages,
|
||||
end_user=details.get("user", "default"),
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
litellm_call_id=details["litellm_call_id"],
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case "lunary" if integration is not None:
|
||||
_method(integration, "log_event")(
|
||||
kwargs=details,
|
||||
type="llm",
|
||||
event="end",
|
||||
model=view.model,
|
||||
input=details["input"],
|
||||
user_id=details.get("user", "default"),
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
run_id=view.litellm_call_id,
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case "helicone" if integration is not None:
|
||||
_method(integration, "log_success")(
|
||||
model=view.model,
|
||||
messages=view.messages,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=print_verbose,
|
||||
kwargs=details,
|
||||
)
|
||||
case "langfuse":
|
||||
_langfuse(
|
||||
logger, response=response, start_time=start_time, end_time=end_time, level=None, status_message=None
|
||||
)
|
||||
case "traceloop" if integration is not None:
|
||||
_method(integration, "log_event")(
|
||||
kwargs=details,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
user_id=details.get("user", None),
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case "s3" if integration is not None:
|
||||
_method(integration, "log_event")(
|
||||
kwargs=details,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case "openmeter" if integration is not None:
|
||||
return _async_method(integration, "async_log_success_event")(
|
||||
kwargs=details, response_obj=response, start_time=start_time, end_time=end_time
|
||||
)
|
||||
case "dynamodb":
|
||||
return _dynamodb(details, response, start_time, end_time, print_verbose)
|
||||
case _:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _dynamodb(
|
||||
details: Details, response: object, start_time: Timestamp, end_time: Timestamp, print_verbose: LegacyCall
|
||||
) -> Awaitable[None]:
|
||||
from litellm.integrations.dynamodb import DyanmoDBLogger
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
existing: Final = _legacy_module().get("dynamoLogger")
|
||||
dynamo: Final = existing if isinstance(existing, DyanmoDBLogger) else DyanmoDBLogger()
|
||||
litellm_logging.dynamoLogger = dynamo # pyright: ignore[reportAttributeAccessIssue] # legacy module global
|
||||
return _async_method(dynamo, "_async_log_event")(
|
||||
kwargs=details, response_obj=response, start_time=start_time, end_time=end_time, print_verbose=print_verbose
|
||||
)
|
||||
|
||||
|
||||
def _langfuse(
|
||||
logger: Logging,
|
||||
*,
|
||||
response: object,
|
||||
start_time: Timestamp,
|
||||
end_time: Timestamp,
|
||||
level: str | None,
|
||||
status_message: str | None,
|
||||
) -> None:
|
||||
from litellm.integrations.langfuse import langfuse_handler
|
||||
|
||||
view: Final = _logger(logger)
|
||||
module: Final = _legacy_module()
|
||||
kwargs: Final = {key: value for key, value in view.model_call_details.items() if key != "original_response"}
|
||||
select: Final = cast(
|
||||
LegacyCall, langfuse_handler.LangFuseHandler.get_langfuse_logger_for_request
|
||||
) # cast-ok: legacy factory
|
||||
handler: Final = select(
|
||||
globalLangfuseLogger=module.get("langFuseLogger"),
|
||||
standard_callback_dynamic_params=view.standard_callback_dynamic_params,
|
||||
in_memory_dynamic_logger_cache=module["in_memory_dynamic_logger_cache"],
|
||||
)
|
||||
if handler is None:
|
||||
return
|
||||
extra: Final[Details] = {"level": level, "status_message": status_message} if level is not None else {}
|
||||
result: Final = _method(handler, "log_event_on_langfuse")(
|
||||
kwargs=kwargs,
|
||||
response_obj=response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
user_id=kwargs.get("user", None),
|
||||
**extra,
|
||||
)
|
||||
trace_id: Final = (
|
||||
cast(Details, result).get("trace_id") if isinstance(result, dict) else None
|
||||
) # cast-ok: legacy response dict
|
||||
if trace_id is not None:
|
||||
_method(module["in_memory_trace_id_cache"], "set_cache")(
|
||||
litellm_call_id=view.litellm_call_id, service_name="langfuse", trace_id=trace_id
|
||||
)
|
||||
|
||||
|
||||
def dispatch_named_failure(
|
||||
logger: Logging,
|
||||
name: str,
|
||||
exception: BaseException,
|
||||
formatted: str,
|
||||
start_time: Timestamp,
|
||||
end_time: Timestamp,
|
||||
) -> None:
|
||||
view: Final = _logger(logger)
|
||||
module: Final = _legacy_module()
|
||||
details: Final = view.model_call_details
|
||||
print_verbose: Final = _print_verbose()
|
||||
match name:
|
||||
case "lunary" if (lunary := module.get("lunaryLogger")) is not None:
|
||||
_method(lunary, "log_event")(
|
||||
kwargs=details,
|
||||
type="llm",
|
||||
event="error",
|
||||
user_id=details.get("user", "default"),
|
||||
model=view.model,
|
||||
input=details["input"],
|
||||
error=formatted,
|
||||
run_id=view.litellm_call_id,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case "sentry" if (capture := module.get("capture_exception")) is not None:
|
||||
cast(LegacyCall, capture)(exception) # cast-ok: legacy sentry hook
|
||||
case "supabase" if (supabase := module.get("supabaseClient")) is not None:
|
||||
_method(supabase, "log_event")(
|
||||
model=view.model,
|
||||
messages=view.messages,
|
||||
end_user=details.get("user", "default"),
|
||||
response_obj=None,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
litellm_call_id=details["litellm_call_id"],
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case "langfuse":
|
||||
_langfuse(
|
||||
logger,
|
||||
response=None,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
level="ERROR",
|
||||
status_message=str(exception),
|
||||
)
|
||||
case "traceloop" if (traceloop := module.get("traceloopLogger")) is not None:
|
||||
_method(traceloop, "log_event")(
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
response_obj=None,
|
||||
user_id=details.get("user", None),
|
||||
print_verbose=print_verbose,
|
||||
status_message=str(exception),
|
||||
level="ERROR",
|
||||
kwargs=details,
|
||||
)
|
||||
case "logfire" if (logfire := module.get("logfireLogger")) is not None:
|
||||
from litellm.integrations.logfire_logger import LogfireLevel
|
||||
|
||||
_method(logfire, "log_event")(
|
||||
kwargs={
|
||||
**{key: value for key, value in details.items() if key != "original_response"},
|
||||
"exception": exception,
|
||||
},
|
||||
response_obj=None,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
level=LogfireLevel.ERROR.value,
|
||||
print_verbose=print_verbose,
|
||||
)
|
||||
case _:
|
||||
return
|
||||
|
||||
|
||||
def restore_correlation_context(logger: object) -> None:
|
||||
from litellm import utils
|
||||
|
||||
utils._restore_correlation_context_if_supported(logger) # pyright: ignore[reportPrivateUsage] # legacy interop helper
|
||||
|
||||
|
||||
def submit_worker(job: Callable[[], None]) -> None:
|
||||
import contextvars
|
||||
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
|
||||
context: Final = contextvars.copy_context()
|
||||
_ = executor.submit(context.run, job)
|
||||
|
||||
|
||||
def enqueue_background(coroutine: Awaitable[None]) -> None:
|
||||
import contextvars
|
||||
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
enqueue: Final = cast(
|
||||
Callable[[Awaitable[None]], None], GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue
|
||||
) # cast-ok: legacy worker accepts any coroutine
|
||||
contextvars.copy_context().run(enqueue, coroutine)
|
||||
|
||||
|
||||
def now() -> Timestamp:
|
||||
return datetime.datetime.now()
|
||||
|
|
@ -22,10 +22,6 @@ class OcrLoggingProtocol(Protocol):
|
|||
custom_llm_provider: str,
|
||||
) -> object: ...
|
||||
|
||||
def pre_call(self, *, input: str, api_key: str | None, additional_args: dict[str, object]) -> object: ...
|
||||
|
||||
def post_call(self, *, original_response: object, additional_args: dict[str, object]) -> object: ...
|
||||
|
||||
|
||||
def _redact(params: Mapping[str, object], secret_fields: Sequence[str]) -> dict[str, object]:
|
||||
return {
|
||||
|
|
@ -62,32 +58,6 @@ def update_logging(
|
|||
)
|
||||
|
||||
|
||||
def pre_call(
|
||||
logger: OcrLoggingProtocol,
|
||||
api_key: str | None,
|
||||
body: dict[str, object],
|
||||
headers: dict[str, str],
|
||||
url: str,
|
||||
) -> None:
|
||||
logger.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": body, "headers": headers, "api_base": url},
|
||||
)
|
||||
|
||||
|
||||
def post_call(
|
||||
logger: OcrLoggingProtocol,
|
||||
original_response: object,
|
||||
body: dict[str, object],
|
||||
headers: dict[str, str],
|
||||
) -> None:
|
||||
logger.post_call(
|
||||
original_response=original_response,
|
||||
additional_args={"complete_input_dict": body, "headers": headers},
|
||||
)
|
||||
|
||||
|
||||
class RustOcr(Protocol):
|
||||
def __call__(self, *args: object, **kwargs: object) -> OCRResponse: ...
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,8 @@ import litellm
|
|||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.ocr import legacy
|
||||
from litellm.rust_bridge import bindings, configuration
|
||||
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, post_call, pre_call, update_logging
|
||||
from litellm.rust_bridge.leaves import record_post_call, record_pre_call
|
||||
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, update_logging
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
@ -65,25 +66,26 @@ def test_logging_redacts_views_and_preserves_opaque_arguments_and_pricing() -> N
|
|||
|
||||
|
||||
def test_logging_callbacks_receive_captured_payload_roots_and_propagate_errors() -> None:
|
||||
logger: Final = Mock()
|
||||
logger: Final = Mock(log_raw_request_response=False, logger_fn=None, model_call_details={})
|
||||
body: Final[dict[str, object]] = {"document": "original"}
|
||||
headers: Final = {"authorization": "key"}
|
||||
response: Final = object()
|
||||
pre_call(logger, "key", body, headers, "https://provider")
|
||||
post_call(logger, response, body, headers)
|
||||
logger.pre_call.assert_called_once_with(
|
||||
record_pre_call(logger, api_key="key", body=body, headers=headers, url="https://provider")
|
||||
record_post_call(logger, original_response=response, body=body, headers=headers)
|
||||
logger._pre_call.assert_called_once_with(
|
||||
input="OCR document processing",
|
||||
api_key="key",
|
||||
model=None,
|
||||
additional_args={"complete_input_dict": body, "headers": headers, "api_base": "https://provider"},
|
||||
)
|
||||
for callback in (logger.pre_call, logger.post_call):
|
||||
for callback in (logger._pre_call, logger.record_post_call):
|
||||
assert callback.call_args.kwargs["additional_args"]["complete_input_dict"] is body
|
||||
assert callback.call_args.kwargs["additional_args"]["headers"] is headers
|
||||
assert logger.post_call.call_args.kwargs["original_response"] is response
|
||||
assert logger.record_post_call.call_args.kwargs["original_response"] is response
|
||||
failure: Final = RuntimeError("callback failed")
|
||||
failing_logger: Final = Mock(pre_call=Mock(side_effect=failure))
|
||||
failing_logger: Final = Mock(_pre_call=Mock(side_effect=failure))
|
||||
with pytest.raises(RuntimeError) as caught:
|
||||
pre_call(failing_logger, None, body, headers, "https://provider")
|
||||
record_pre_call(failing_logger, api_key=None, body=body, headers=headers, url="https://provider")
|
||||
assert caught.value is failure
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import litellm
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger
|
||||
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
|
||||
from tests.test_litellm_rust.support.requests import (
|
||||
OCR_DOCUMENT,
|
||||
OCR_RESPONSE,
|
||||
|
|
@ -18,7 +19,6 @@ from tests.test_litellm_rust.support.requests import (
|
|||
request_body,
|
||||
request_headers,
|
||||
)
|
||||
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
|
||||
|
|
@ -309,6 +309,7 @@ async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context
|
|||
asynchronous: bool,
|
||||
) -> None:
|
||||
from contextvars import ContextVar
|
||||
|
||||
context: Final = ContextVar("azure-token-context", default="missing")
|
||||
context.set("caller")
|
||||
caller_thread: Final = threading.current_thread()
|
||||
|
|
@ -337,9 +338,7 @@ async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context
|
|||
"callbacks": [Edit()],
|
||||
}
|
||||
response: Final = (
|
||||
await call_native_aocr(ocr_server, **arguments)
|
||||
if asynchronous
|
||||
else call_native_ocr(ocr_server, **arguments)
|
||||
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
|
||||
)
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert observations == ["token", "pre_call"]
|
||||
|
|
@ -368,9 +367,7 @@ async def test_native_azure_ocr_token_provider_can_make_nested_native_ocr_call(
|
|||
"azure_ad_token_provider": provider,
|
||||
}
|
||||
response: Final = (
|
||||
await call_native_aocr(ocr_server, **arguments)
|
||||
if asynchronous
|
||||
else call_native_ocr(ocr_server, **arguments)
|
||||
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
|
||||
)
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert calls == ["token"]
|
||||
|
|
@ -425,7 +422,9 @@ async def test_native_azure_ocr_releases_token_provider_after_terminal_outcome(
|
|||
) -> None:
|
||||
import gc
|
||||
import weakref
|
||||
|
||||
from tests.test_litellm_rust.support.callback_recorder import drain_logging
|
||||
|
||||
class Provider:
|
||||
def __call__(self) -> str:
|
||||
if outcome == "failure":
|
||||
|
|
@ -471,3 +470,238 @@ async def test_native_azure_ocr_releases_token_provider_after_terminal_outcome(
|
|||
await asyncio.sleep(0)
|
||||
gc.collect()
|
||||
assert reference() is None
|
||||
|
||||
|
||||
FORBIDDEN_ORCHESTRATION: Final = (
|
||||
"pre_call",
|
||||
"post_call",
|
||||
"success_handler",
|
||||
"async_success_handler",
|
||||
"failure_handler",
|
||||
"async_failure_handler",
|
||||
"_success_handler_body",
|
||||
"_async_success_handler_body",
|
||||
"_failure_handler_body",
|
||||
"_async_failure_handler_body",
|
||||
"dispatch_success_handlers",
|
||||
"dispatch_failure_handlers",
|
||||
"handle_sync_success_callbacks_for_async_calls",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def legacy_orchestration_disabled(monkeypatch: pytest.MonkeyPatch) -> list[str]:
|
||||
from litellm import utils
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
reached: Final[list[str]] = []
|
||||
|
||||
def forbid(name: str):
|
||||
def method(self, *args, **kwargs):
|
||||
reached.append(name)
|
||||
raise AssertionError(f"legacy orchestration reached: {name}")
|
||||
|
||||
return method
|
||||
|
||||
for name in FORBIDDEN_ORCHESTRATION:
|
||||
monkeypatch.setattr(Logging, name, forbid(name))
|
||||
|
||||
def forbidden_setup(*args, **kwargs):
|
||||
reached.append("function_setup")
|
||||
raise AssertionError("legacy orchestration reached: function_setup")
|
||||
|
||||
monkeypatch.setattr(utils, "function_setup", forbidden_setup)
|
||||
return reached
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
|
||||
async def test_native_ocr_success_runs_integrations_without_legacy_orchestration(
|
||||
ocr_server: RecordingServer, legacy_orchestration_disabled: list[str], asynchronous: bool
|
||||
) -> None:
|
||||
recorder: Final = RecordingLogger()
|
||||
arguments: Final = {"callbacks": [recorder]}
|
||||
response: Final = (
|
||||
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
|
||||
)
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
success_event: Final = "async_log_success_event" if asynchronous else "log_success_event"
|
||||
events: Final = await recorder.wait_for_async(success_event)
|
||||
assert legacy_orchestration_disabled == []
|
||||
assert recorder.names.count("log_pre_api_call") == 1
|
||||
assert events[0].kwargs["standard_logging_object"]["status"] == "success"
|
||||
assert events[0].kwargs["response_cost"] is not None
|
||||
assert events[0].kwargs["litellm_params"]["api_base"] == f"{ocr_server.base_url}/v1/ocr"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
|
||||
async def test_native_ocr_failure_runs_integrations_without_legacy_orchestration(
|
||||
ocr_server: RecordingServer, legacy_orchestration_disabled: list[str], asynchronous: bool
|
||||
) -> None:
|
||||
ocr_server.enqueue(ResponseSpec(body={"message": "provider unavailable"}, status=500))
|
||||
recorder: Final = RecordingLogger()
|
||||
arguments: Final = {"callbacks": [recorder]}
|
||||
with pytest.raises(litellm.InternalServerError) as caught:
|
||||
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
|
||||
assert legacy_orchestration_disabled == []
|
||||
failures: Final = tuple(event for event in recorder.events if event.name.endswith("log_failure_event"))
|
||||
assert [event.name for event in failures] == (
|
||||
["log_failure_event", "async_log_failure_event"] if asynchronous else ["log_failure_event"]
|
||||
)
|
||||
assert all(event.kwargs["exception"] is caught.value for event in failures)
|
||||
assert all(event.kwargs["standard_logging_object"]["status"] == "failure" for event in failures)
|
||||
assert "log_success_event" not in recorder.names
|
||||
|
||||
|
||||
def test_native_ocr_success_hooks_all_run_before_any_success_dispatch(ocr_server: RecordingServer) -> None:
|
||||
order: Final = []
|
||||
finished: Final = threading.Event()
|
||||
|
||||
class Hooked(CustomLogger):
|
||||
def __init__(self, name: str) -> None:
|
||||
super().__init__()
|
||||
self.name = name
|
||||
|
||||
def logging_hook(self, kwargs, result, call_type):
|
||||
order.append(("hook", self.name))
|
||||
kwargs[f"seen-by-{self.name}"] = True
|
||||
return kwargs, result
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
order.append(("log", self.name, kwargs.get("seen-by-a"), kwargs.get("seen-by-b")))
|
||||
if self.name == "b":
|
||||
finished.set()
|
||||
|
||||
call_native_ocr_with_callbacks(ocr_server, [Hooked("a"), Hooked("b")])
|
||||
|
||||
assert finished.wait(10)
|
||||
assert order == [("hook", "a"), ("hook", "b"), ("log", "a", True, True), ("log", "b", True, True)]
|
||||
|
||||
|
||||
def test_native_ocr_hook_failure_is_contained_and_target_still_dispatches(ocr_server: RecordingServer) -> None:
|
||||
order: Final = []
|
||||
finished: Final = threading.Event()
|
||||
|
||||
class Broken(CustomLogger):
|
||||
def logging_hook(self, kwargs, result, call_type):
|
||||
raise RuntimeError("hook failed")
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
order.append("broken-log")
|
||||
|
||||
class Healthy(CustomLogger):
|
||||
def logging_hook(self, kwargs, result, call_type):
|
||||
order.append("healthy-hook")
|
||||
return kwargs, result
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
order.append("healthy-log")
|
||||
finished.set()
|
||||
|
||||
call_native_ocr_with_callbacks(ocr_server, [Broken(), Healthy()])
|
||||
|
||||
assert finished.wait(10)
|
||||
assert order == ["healthy-hook", "broken-log", "healthy-log"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_aocr_hook_replacement_of_result_reaches_success_dispatch(ocr_server: RecordingServer) -> None:
|
||||
replacement: Final = object()
|
||||
observed: Final = []
|
||||
|
||||
class Replace(CustomLogger):
|
||||
async def async_logging_hook(self, kwargs, result, call_type):
|
||||
return kwargs, replacement
|
||||
|
||||
class Observe(CustomLogger):
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
observed.append(response_obj)
|
||||
|
||||
recorder: Final = RecordingLogger()
|
||||
await call_native_aocr_with_callbacks(ocr_server, [Replace(), Observe(), recorder])
|
||||
await recorder.wait_for_async("async_log_success_event")
|
||||
|
||||
assert observed == [replacement]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_aocr_shared_logging_object_dispatches_success_once(ocr_server: RecordingServer) -> None:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
ocr_server.expected_requests = 2
|
||||
recorder: Final = RecordingLogger()
|
||||
litellm.callbacks.append(recorder)
|
||||
logger: Final = Logging(
|
||||
model="mistral-ocr-latest",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="aocr",
|
||||
start_time=__import__("datetime").datetime.now(),
|
||||
litellm_call_id="shared",
|
||||
function_id="shared",
|
||||
)
|
||||
logger.dynamic_async_success_callbacks = [recorder]
|
||||
|
||||
await call_native_aocr(ocr_server, litellm_logging_obj=logger)
|
||||
await call_native_aocr(ocr_server, litellm_logging_obj=logger)
|
||||
await recorder.wait_for_async("async_log_success_event")
|
||||
from tests.test_litellm_rust.support.callback_recorder import drain_logging
|
||||
|
||||
await drain_logging()
|
||||
|
||||
assert logger.model_call_details["has_logged_async_success"] is True
|
||||
assert recorder.names.count("async_log_success_event") == 1
|
||||
|
||||
|
||||
def test_native_ocr_logging_preparation_failure_does_not_fail_request(
|
||||
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
def broken_payload(*args, **kwargs):
|
||||
raise RuntimeError("payload unavailable")
|
||||
|
||||
monkeypatch.setattr(litellm_logging, "get_standard_logging_object_payload", broken_payload)
|
||||
unraisable: Final = []
|
||||
monkeypatch.setattr(__import__("sys"), "unraisablehook", lambda event: unraisable.append(event.exc_value))
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
response: Final = call_native_ocr_with_callbacks(ocr_server, [recorder])
|
||||
events: Final = recorder.wait_for("log_success_event")
|
||||
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert len(events) == 1
|
||||
assert events[0].thread is not threading.current_thread()
|
||||
assert any(str(error) == "payload unavailable" for error in unraisable)
|
||||
|
||||
|
||||
def test_native_ocr_writes_success_marker_and_honours_existing_marker(
|
||||
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from litellm.rust_bridge import setup as native_setup
|
||||
|
||||
ocr_server.expected_requests = 2
|
||||
loggers: Final = []
|
||||
original_build: Final = native_setup.build_logging
|
||||
|
||||
def build_logging(**kwargs):
|
||||
logger = original_build(**kwargs)
|
||||
if loggers:
|
||||
logger.model_call_details["has_logged_sync_success"] = True
|
||||
loggers.append(logger)
|
||||
return logger
|
||||
|
||||
monkeypatch.setattr(native_setup, "build_logging", build_logging)
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
call_native_ocr_with_callbacks(ocr_server, [recorder])
|
||||
recorder.wait_for("log_success_event")
|
||||
assert loggers[0].model_call_details["has_logged_sync_success"] is True
|
||||
|
||||
call_native_ocr_with_callbacks(ocr_server, [recorder])
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
|
||||
executor.submit(lambda: None).result(10)
|
||||
assert recorder.names.count("log_success_event") == 1
|
||||
assert recorder.names.count("logging_hook") == 1
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue