refactor(python-bridge): drop the finalize shim and prove native handles collect

Finalize calls update_response_metadata directly instead of through a
lifecycle.py passthrough. GC tests cover WorkerJob and DeferredSuccess
traversal through logger, targets and response, and DeferredSuccess close
and release idempotence.
This commit is contained in:
Yujong Lee 2026-09-16 09:18:07 -07:00
parent f0b877532c
commit 316700aeb8
3 changed files with 128 additions and 41 deletions

View file

@ -42,9 +42,11 @@ pub(super) fn finalize(
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> 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::<pyo3::types::PyString>());
py.import("litellm.litellm_core_utils.llm_response_utils.response_metadata")?
.getattr("update_response_metadata")?
.call1((response, logger.object(py), model, kwargs, start, end))?;
Ok(())
}

View file

@ -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<CallbackId> = ids.into_iter().flatten().collect();
Runner {
cursor: DispatchCursor::start(family, ids.clone(), false, false),
job: Job {
logger: logger.extract().unwrap(),
targets,
ids,
family,
response: Some(target.clone().unbind()),
error: None,
start: py.None(),
end: py.None(),
stream: false,
},
result: Some(target.unbind()),
formatted: None,
pending: None,
}
}
fn assert_collectable(py: Python<'_>, locals: &Bound<'_, PyDict>, handle: Py<PyAny>) {
locals.set_item("handle", handle).unwrap();
py.run(
pyo3::ffi::c_str!(
r#"
import gc
import weakref
target.handle = handle
reference = weakref.ref(target)
del logger, target, handle
gc.collect()
assert reference() is None, "cycle through the retained target was not collected"
"#
),
Some(locals),
Some(locals),
)
.unwrap();
}
#[test]
fn worker_job_collects_cycles_through_logger_targets_and_response() {
Python::initialize();
Python::attach(|py| {
let locals = fixture(py);
let runner = runner(py, &locals, CallbackFamily::SyncSuccess);
let handle = Py::new(py, WorkerJob::new(runner)).unwrap().into_any();
assert_collectable(py, &locals, handle);
});
}
#[test]
fn deferred_success_collects_cycles_and_close_is_idempotent() {
Python::initialize();
Python::attach(|py| {
let locals = fixture(py);
let runner = runner(py, &locals, CallbackFamily::AsyncSuccess);
let deferred = Py::new(
py,
super::super::DeferredSuccess {
runner: Some(runner),
},
)
.unwrap();
assert_collectable(py, &locals, deferred.into_any());
});
}
#[test]
fn deferred_success_releases_at_most_once_and_close_prevents_release() {
Python::initialize();
Python::attach(|py| {
let locals = fixture(py);
let runner = runner(py, &locals, CallbackFamily::AsyncSuccess);
let deferred = Py::new(
py,
super::super::DeferredSuccess {
runner: Some(runner),
},
)
.unwrap();
deferred.call_method0(py, "close").unwrap();
deferred.call_method0(py, "close").unwrap();
deferred.call0(py).unwrap();
assert!(deferred.borrow(py).runner.is_none());
});
}
}

View file

@ -1,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)