diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs index 3850b827093..f78d4f87c62 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs @@ -42,9 +42,11 @@ pub(super) fn finalize( start: &Py, end: &Option>, ) -> PyResult<()> { - py.import("litellm.rust_bridge.lifecycle")? - .getattr("finalize")? - .call1((response, logger.object(py), kwargs, start, end))?; + let model = kwargs.bind(py).get_item("model")?; + let model = model.filter(|value| value.is_instance_of::()); + py.import("litellm.litellm_core_utils.llm_response_utils.response_metadata")? + .getattr("update_response_metadata")? + .call1((response, logger.object(py), model, kwargs, start, end))?; Ok(()) } diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs b/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs index 5fae1a01d7f..4215b8e6c2d 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/dispatch.rs @@ -647,3 +647,125 @@ pub(super) fn family_targets( ); Ok((targets, ordered)) } + +#[cfg(test)] +mod tests { + use super::*; + use pyo3::types::{PyDict, PyList}; + + fn fixture(py: Python<'_>) -> Bound<'_, PyDict> { + let locals = PyDict::new(py); + py.run( + pyo3::ffi::c_str!( + r#" +import sys, types +for name in ("litellm", "litellm.integrations", "litellm.integrations.custom_logger"): + sys.modules.setdefault(name, types.ModuleType(name)) +class CustomLogger: pass +sys.modules["litellm.integrations.custom_logger"].CustomLogger = CustomLogger +sys.modules["litellm"]._known_custom_logger_compatible_callbacks = [] +class Logger: pass +logger = Logger() +target = CustomLogger() +"# + ), + Some(&locals), + Some(&locals), + ) + .unwrap(); + locals + } + + fn runner(py: Python<'_>, locals: &Bound<'_, PyDict>, family: CallbackFamily) -> Runner { + let logger = locals.get_item("logger").unwrap().unwrap(); + let target = locals.get_item("target").unwrap().unwrap(); + let list = PyList::new(py, [&target]).unwrap().into_any(); + let (targets, ids) = Targets::read(py, &[list]).unwrap(); + let ids: Vec = ids.into_iter().flatten().collect(); + Runner { + cursor: DispatchCursor::start(family, ids.clone(), false, false), + job: Job { + logger: logger.extract().unwrap(), + targets, + ids, + family, + response: Some(target.clone().unbind()), + error: None, + start: py.None(), + end: py.None(), + stream: false, + }, + result: Some(target.unbind()), + formatted: None, + pending: None, + } + } + + fn assert_collectable(py: Python<'_>, locals: &Bound<'_, PyDict>, handle: Py) { + locals.set_item("handle", handle).unwrap(); + py.run( + pyo3::ffi::c_str!( + r#" +import gc +import weakref +target.handle = handle +reference = weakref.ref(target) +del logger, target, handle +gc.collect() +assert reference() is None, "cycle through the retained target was not collected" +"# + ), + Some(locals), + Some(locals), + ) + .unwrap(); + } + + #[test] + fn worker_job_collects_cycles_through_logger_targets_and_response() { + Python::initialize(); + Python::attach(|py| { + let locals = fixture(py); + let runner = runner(py, &locals, CallbackFamily::SyncSuccess); + let handle = Py::new(py, WorkerJob::new(runner)).unwrap().into_any(); + assert_collectable(py, &locals, handle); + }); + } + + #[test] + fn deferred_success_collects_cycles_and_close_is_idempotent() { + Python::initialize(); + Python::attach(|py| { + let locals = fixture(py); + let runner = runner(py, &locals, CallbackFamily::AsyncSuccess); + let deferred = Py::new( + py, + super::super::DeferredSuccess { + runner: Some(runner), + }, + ) + .unwrap(); + assert_collectable(py, &locals, deferred.into_any()); + }); + } + + #[test] + fn deferred_success_releases_at_most_once_and_close_prevents_release() { + Python::initialize(); + Python::attach(|py| { + let locals = fixture(py); + let runner = runner(py, &locals, CallbackFamily::AsyncSuccess); + let deferred = Py::new( + py, + super::super::DeferredSuccess { + runner: Some(runner), + }, + ) + .unwrap(); + deferred.call_method0(py, "close").unwrap(); + deferred.call_method0(py, "close").unwrap(); + deferred.call0(py).unwrap(); + assert!(deferred.borrow(py).runner.is_none()); + }); + } +} diff --git a/litellm/rust_bridge/lifecycle.py b/litellm/rust_bridge/lifecycle.py index 432ece7e093..c3ab962a5cb 100644 --- a/litellm/rust_bridge/lifecycle.py +++ b/litellm/rust_bridge/lifecycle.py @@ -1,17 +1,8 @@ from __future__ import annotations -import datetime from collections.abc import Awaitable, Mapping from dataclasses import dataclass -from typing import ( - TYPE_CHECKING, - Final, - Protocol, - cast, # noqa: TID251 # bounded compatibility calls into legacy Python integrations -) - -if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import Logging +from typing import Final, Protocol @dataclass(frozen=True, slots=True) @@ -51,18 +42,6 @@ async def drive(execution: Execution) -> object: execution.close() -class MetadataUpdater(Protocol): - def __call__( - self, - result: object, - logging_obj: Logging, - model: str | None, - kwargs: dict[str, object], - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - def check_limits(kwargs: Mapping[str, object]) -> None: import litellm from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit @@ -72,19 +51,3 @@ def check_limits(kwargs: Mapping[str, object]) -> None: raise litellm.BudgetExceededError(current_cost=current_cost, max_budget=litellm.max_budget) if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request): raise RuntimeError("Max retries per request hit!") - - -def finalize( - response: object, - logger: Logging, - kwargs: dict[str, object], - start_time: datetime.datetime, - end_time: datetime.datetime, -) -> None: - from litellm.litellm_core_utils.llm_response_utils import response_metadata - - model: Final = kwargs.get("model") - update: Final = cast( # cast-ok: legacy metadata function accepts concrete kwargs - MetadataUpdater, response_metadata.update_response_metadata - ) - update(response, logger, model if isinstance(model, str) else None, kwargs, start_time, end_time)