mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
f0b877532c
commit
316700aeb8
3 changed files with 128 additions and 41 deletions
|
|
@ -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(())
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue