refactor(rust): complete retained callback lifecycle foundation

This commit is contained in:
Yujong Lee 2026-09-06 16:07:34 -07:00
parent e399ec52f1
commit fc1877718d
8 changed files with 1139 additions and 35 deletions

View file

@ -138,4 +138,4 @@ jobs:
LITELLM_LOCAL_MODEL_COST_MAP: "True"
run: >-
cargo test --manifest-path litellm-rust/Cargo.toml
-p litellm-python-interop --test prepared_call --locked -- --include-ignored
-p litellm-python-interop --tests --locked -- --include-ignored

View file

@ -1,26 +1,91 @@
use pyo3::class::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::sync::PyOnceLock;
use pyo3::types::{PyDict, PyTuple};
static AWAIT_CALL: PyOnceLock<Py<PyAny>> = PyOnceLock::new();
#[derive(Clone, Copy)]
pub enum InvocationMode {
Direct,
Await,
}
#[derive(Debug)]
pub enum InvocationOutcome {
Returned(Py<PyAny>),
Awaitable(Py<PyAny>),
}
pub struct PreparedCall {
mode: InvocationMode,
callable: Py<PyAny>,
positional: Py<PyTuple>,
keywords: Option<Py<PyDict>>,
}
impl PreparedCall {
pub fn new(callable: Py<PyAny>, positional: Py<PyTuple>, keywords: Option<Py<PyDict>>) -> Self {
pub fn new(
mode: InvocationMode,
callable: Py<PyAny>,
positional: Py<PyTuple>,
keywords: Option<Py<PyDict>>,
) -> Self {
Self {
mode,
callable,
positional,
keywords,
}
}
pub fn invoke(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
self.callable.call(
py,
self.positional.bind(py),
self.keywords.as_ref().map(|kwargs| kwargs.bind(py)),
)
pub fn invoke(&self, py: Python<'_>) -> PyResult<InvocationOutcome> {
match self.mode {
InvocationMode::Direct => self
.callable
.call(
py,
self.positional.bind(py),
self.keywords.as_ref().map(|kwargs| kwargs.bind(py)),
)
.map(InvocationOutcome::Returned),
InvocationMode::Await => {
let adapter = AWAIT_CALL.get_or_try_init(py, || {
PyModule::from_code(
py,
c"async def invoke_awaited(callable, positional, keywords):
if keywords is None:
return await callable(*positional)
return await callable(*positional, **keywords)
",
c"retained_callback.py",
c"_retained_callback",
)?
.getattr("invoke_awaited")
.map(Bound::unbind)
})?;
adapter
.call1(py, (&self.callable, &self.positional, &self.keywords))
.map(InvocationOutcome::Awaitable)
}
}
}
pub fn clone_ref(&self, py: Python<'_>) -> Self {
Self {
mode: self.mode,
callable: self.callable.clone_ref(py),
positional: self.positional.clone_ref(py),
keywords: self.keywords.as_ref().map(|value| value.clone_ref(py)),
}
}
pub fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.callable)?;
visit.call(&self.positional)?;
if let Some(keywords) = &self.keywords {
visit.call(keywords)?;
}
Ok(())
}
}

View file

@ -2,6 +2,6 @@ mod callback;
mod gil;
mod marshal;
pub use callback::PreparedCall;
pub use callback::{InvocationMode, InvocationOutcome, PreparedCall};
pub use gil::{release_count, release_gil};
pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py};

View file

@ -0,0 +1,153 @@
use std::ffi::CString;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::{fixture, rstest};
#[path = "support/callback_owner.rs"]
mod callback_owner;
struct InitializedPython;
#[fixture]
#[once]
fn initialized_python() -> InitializedPython {
Python::initialize();
InitializedPython
}
#[fixture]
fn scenario_scope(initialized_python: &InitializedPython) -> Py<PyDict> {
let _ = initialized_python;
Python::attach(|py| {
let globals = PyDict::new(py);
globals
.set_item(
"factory",
Py::new(py, callback_owner::OwnerFactory::default()).unwrap(),
)
.unwrap();
let source = CString::new(include_str!("fixtures/callback_lifecycle.py")).unwrap();
py.run(&source, Some(&globals), None).unwrap();
globals.unbind()
})
}
#[rstest]
#[case::awaitable_kinds("awaitable_kinds")]
#[case::identity_and_context("identity_and_context")]
#[case::exceptions("exceptions")]
#[case::exception_ownership("exception_ownership")]
#[case::cancellation_before_start("cancellation_before_start")]
#[case::cancellation_unwinds("cancellation_unwinds")]
#[case::cancellation_during_cleanup("cancellation_during_cleanup")]
#[case::cancellation_suppressed("cancellation_suppressed")]
#[case::registration_and_gc("registration_and_gc")]
#[case::background_and_session("background_and_session")]
#[case::stream_lifecycle("stream_lifecycle")]
#[case::sync_stream_lifecycle("sync_stream_lifecycle")]
#[case::repeated_ownership("repeated_ownership")]
fn lifecycle_contract(
scenario_scope: Py<PyDict>,
#[case] scenario: &str,
#[values(false, true)] retained: bool,
) -> PyResult<()> {
Python::attach(|py| {
scenario_scope
.bind(py)
.get_item("run_scenario")?
.unwrap()
.call1((
scenario,
retained,
scenario_scope.bind(py).get_item("factory")?.unwrap(),
))?;
Ok(())
})
}
#[rstest]
#[case::real_async_logging("real_async_logging")]
#[case::real_stream_completion("real_stream_completion")]
#[case::real_stream_close("real_stream_close")]
#[case::real_stream_cancellation("real_stream_cancellation")]
#[ignore = "requires the repository Python environment and LiteLLM on PYTHONPATH"]
fn component_contract(
scenario_scope: Py<PyDict>,
#[case] scenario: &str,
#[values(false, true)] retained: bool,
) -> PyResult<()> {
Python::attach(|py| {
let globals = scenario_scope.bind(py);
let source = CString::new(include_str!("fixtures/callback_components.py")).unwrap();
py.run(&source, Some(globals), None)?;
globals.get_item("run_scenario")?.unwrap().call1((
scenario,
retained,
scenario_scope.bind(py).get_item("factory")?.unwrap(),
))?;
Ok(())
})
}
#[rstest]
fn detached_release(initialized_python: &InitializedPython) -> PyResult<()> {
use litellm_python_interop::{InvocationMode, PreparedCall};
use pyo3::types::PyTuple;
let _ = initialized_python;
let (call, reference) = Python::attach(|py| -> PyResult<_> {
let globals = PyDict::new(py);
py.run(
c"import weakref\nclass Value: pass\nvalue = Value()\nreference = weakref.ref(value)",
Some(&globals),
None,
)?;
let value = globals.get_item("value")?.unwrap();
let call = PreparedCall::new(
InvocationMode::Direct,
py.None(),
PyTuple::new(py, [value])?.unbind(),
None,
);
let reference = globals.get_item("reference")?.unwrap().unbind();
globals.del_item("value")?;
Ok((call, reference))
})?;
drop(call);
Python::attach(|py| {
assert!(reference.call0(py)?.is_none(py));
Ok(())
})
}
#[rstest]
#[case::direct(false)]
#[case::awaited(true)]
fn outcome_identifies_binding(scenario_scope: Py<PyDict>, #[case] awaited: bool) -> PyResult<()> {
use litellm_python_interop::{InvocationMode, InvocationOutcome, PreparedCall};
use pyo3::types::PyTuple;
Python::attach(|py| {
let callback = py.eval(c"lambda: None", Some(scenario_scope.bind(py)), None)?;
let call = PreparedCall::new(
if awaited {
InvocationMode::Await
} else {
InvocationMode::Direct
},
callback.unbind(),
PyTuple::empty(py).unbind(),
None,
);
match call.invoke(py)? {
InvocationOutcome::Returned(value) => {
assert!(!awaited);
assert!(value.is_none(py));
}
InvocationOutcome::Awaitable(value) => {
assert!(awaited);
value.call_method0(py, "close")?;
}
}
Ok(())
})
}

View file

@ -0,0 +1,211 @@
import asyncio
from datetime import datetime
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.types.utils import Delta, ModelResponse, ModelResponseStream, StreamingChoices, Usage
def logger_for(callbacks=(), stream=False):
return Logging(
model="test",
messages=[{"role": "user", "content": "test"}],
stream=stream,
call_type="acompletion",
start_time=datetime.now(),
litellm_call_id="retained-test",
function_id="retained-test",
dynamic_async_success_callbacks=list(callbacks),
)
async def real_async_logging(owners):
observations = []
task = asyncio.current_task()
gate = asyncio.Event()
result = ModelResponse(model="test", choices=[{"message": {"role": "assistant", "content": "original"}}])
replacement = ModelResponse(model="test", choices=[{"message": {"role": "assistant", "content": "replacement"}}])
shared = {}
class Retain(CustomLogger):
async def async_logging_hook(self, kwargs, result, call_type):
observations.append(("retained", kwargs, result))
assert asyncio.current_task() is task
return kwargs, result
class MutateThenFail(CustomLogger):
async def async_logging_hook(self, kwargs, result, call_type):
asyncio.get_running_loop().call_soon(gate.set)
await gate.wait()
kwargs["retained_shared"]["changed"] = True
result.choices[0].message.content = "mutated"
raise RuntimeError("expected async callback failure")
class Replace(CustomLogger):
async def async_logging_hook(self, kwargs, result, call_type):
observations.append(("replace", kwargs, result))
return kwargs, replacement
class Observe(CustomLogger):
async def async_logging_hook(self, kwargs, result, call_type):
observations.append(("observe", kwargs, result))
assert asyncio.current_task() is task
return kwargs, result
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
observations.append(("success", kwargs, response_obj))
logger = logger_for([Retain(), MutateThenFail(), Replace(), Observe()])
logger.model_call_details["retained_shared"] = shared
owner = owners.prepare(logger.async_success_handler, (), {"result": result}, awaited=True)
await owner.invoke()
owner.close()
assert [entry[0] for entry in observations] == ["retained", "replace", "observe", "success"]
assert observations[0][2] is result and observations[1][2] is result
assert observations[2][2] is replacement and observations[3][2] is replacement
assert observations[0][1]["retained_shared"] is shared and shared["changed"]
assert result.choices[0].message.content == "mutated"
class ControlledStream:
def __init__(self):
self.originals = [
ModelResponseStream(model="test", choices=[StreamingChoices(delta=Delta(content="hello"), index=0)]),
ModelResponseStream(
model="test", choices=[StreamingChoices(delta=Delta(content=""), index=0, finish_reason="stop")]
),
ModelResponseStream(
model="test", choices=[], usage=Usage(prompt_tokens=3, completion_tokens=5, total_tokens=8)
),
]
self.chunks = iter(self.originals)
self.closed = 0
def __aiter__(self):
return self
async def __anext__(self):
try:
return next(self.chunks)
except StopIteration:
raise StopAsyncIteration
async def aclose(self):
self.closed += 1
async def drain_component_tasks():
pending = asyncio.all_tasks() - {asyncio.current_task()}
if pending:
await asyncio.gather(*pending)
async def real_stream_completion(owners):
logger = logger_for(stream=True)
completions = []
cached = []
class CacheRecorder:
async def _add_streaming_response_to_cache(self, response):
cached.append(response)
logger._llm_caching_handler = CacheRecorder()
async def complete(response, cache_hit):
completions.append(response)
logger._on_deferred_stream_complete = complete
stream = ControlledStream()
wrapper = CustomStreamWrapper(
completion_stream=stream, model="test", logging_obj=logger, custom_llm_provider="bedrock"
)
pull = owners.prepare(wrapper.__anext__, (), awaited=True)
chunks = []
retained_hidden = None
hidden_owner = None
while True:
try:
chunk = await pull.invoke()
chunks.append(chunk)
if chunk.choices and chunk.choices[0].finish_reason:
retained_hidden = chunk._hidden_params
hidden_owner = owners.prepare(lambda value: value, (retained_hidden,))
assert not completions
except StopAsyncIteration:
break
pull.close()
assert retained_hidden is not None
assert retained_hidden["usage"].total_tokens == 8
assert completions == []
response, cache_hit = logger._deferred_stream_complete_args
assert response.usage.total_tokens == 8
deferred = owners.prepare(logger._on_deferred_stream_complete, (response, cache_hit), awaited=True)
logger._on_deferred_stream_complete = None
logger._deferred_stream_complete_args = None
close = owners.prepare(wrapper.aclose, (), awaited=True)
await close.invoke()
close.close()
del wrapper, logger
await deferred.invoke()
deferred.close()
assert completions == [response]
assert retained_hidden["usage"].total_tokens == 8
assert hidden_owner.invoke() is retained_hidden
hidden_owner.close()
await drain_component_tasks()
assert len(cached) == 1 and cached[0] is not response
assert cached[0].choices[0] is not response.choices[0]
cached[0].choices[0].message.content = "cache only"
assert response.choices[0].message.content == "hello"
async def real_stream_close(owners):
source = ControlledStream()
logger = logger_for(stream=True)
wrapper = CustomStreamWrapper(
completion_stream=source, model="test", logging_obj=logger, custom_llm_provider="bedrock"
)
pull = owners.prepare(wrapper.__anext__, (), awaited=True)
chunk = await pull.invoke()
pull.close()
retained = owners.prepare(lambda value: value, (chunk,))
close = owners.prepare(wrapper.aclose, (), awaited=True)
await close.invoke()
await close.invoke()
close.close()
assert source.closed == 1
assert retained.invoke() is chunk
assert chunk.choices[0].delta.content == "hello"
retained.close()
await drain_component_tasks()
async def real_stream_cancellation(owners):
entered = asyncio.Event()
class SuspendedStream(ControlledStream):
async def __anext__(self):
entered.set()
await asyncio.Event().wait()
source = SuspendedStream()
logger = logger_for(stream=True)
wrapper = CustomStreamWrapper(
completion_stream=source, model="test", logging_obj=logger, custom_llm_provider="bedrock"
)
pull = owners.prepare(wrapper.__anext__, (), awaited=True)
task = asyncio.create_task(pull.invoke())
pull.close()
await entered.wait()
task.cancel()
try:
await task
assert False, "pull ignored cancellation"
except asyncio.CancelledError:
pass
close = owners.prepare(wrapper.aclose, (), awaited=True)
await close.invoke()
close.close()
assert source.closed == 1
await drain_component_tasks()

View file

@ -0,0 +1,519 @@
import asyncio
import contextvars
import copy
import gc
import json
import threading
import weakref
async def checkpoint():
ready = asyncio.Event()
asyncio.get_running_loop().call_soon(ready.set)
await ready.wait()
class Value:
pass
class ReferenceFactory:
def __init__(self):
self.live = 0
def prepare(self, callable, positional, keywords=None, awaited=False):
return ReferenceOwner(self, callable, positional, keywords, awaited)
class ReferenceOwner:
def __init__(self, factory, callable, positional, keywords, awaited):
self.factory = factory
self.call = (callable, positional, keywords, awaited)
factory.live += 1
def invoke(self):
if self.call is None:
raise RuntimeError("invocation owner released")
callable, positional, keywords, awaited = self.call
if not awaited:
return callable(*positional, **(keywords if keywords is not None else {}))
async def run():
return await callable(*positional, **(keywords if keywords is not None else {}))
return run()
def clone_owner(self):
return ReferenceOwner(self.factory, *self.call)
def close(self):
if self.call is not None:
released, self.call = self.call, None
self.factory.live -= 1
del released
def __del__(self):
self.close()
async def awaitable_kinds(owners):
payload = Value()
calls = []
async def coroutine(value, *, alias):
assert value is alias
calls.append("called")
return value
class CustomAwaitable:
def __await__(self):
return coroutine(payload, alias=payload).__await__()
for kind in ("async", "sync_coroutine", "custom", "future"):
future = asyncio.get_running_loop().create_future()
future.set_result(payload)
callback = {
"async": coroutine,
"sync_coroutine": lambda value, *, alias: coroutine(value, alias=alias),
"custom": lambda value, *, alias: CustomAwaitable(),
"future": lambda value, *, alias: future,
}[kind]
owner = owners.prepare(callback, (payload,), {"alias": payload}, awaited=True)
pending = owner.invoke()
before = len(calls)
owner.close()
assert await pending is payload
assert len(calls) == before + (kind != "future")
owner = owners.prepare(lambda: payload, (), awaited=True)
try:
await owner.invoke()
assert False, "non-awaitable result accepted"
except TypeError:
pass
owner.close()
inner = coroutine(payload, alias=payload)
async def returns_coroutine():
return inner
owner = owners.prepare(returns_coroutine, (), awaited=True)
assert await owner.invoke() is inner
assert inner.cr_frame is not None
inner.close()
owner.close()
direct = owners.prepare(coroutine, (payload,), {"alias": payload})
untouched = direct.invoke()
assert untouched.cr_frame is not None
assert untouched.cr_await is None
untouched.close()
direct.close()
async def identity_and_context(owners):
context = contextvars.ContextVar("retained_context", default="outside")
task = asyncio.current_task()
loop = asyncio.get_running_loop()
thread = threading.get_ident()
payload = {"nested": {}}
alias = payload["nested"]
saved = []
gate = asyncio.Event()
async def nested(value):
assert asyncio.current_task() is task
assert context.get() == "inside"
value["nested"]["nested_call"] = True
context.set("nested")
return value
async def callback(value, *, shared):
assert value is payload and shared is alias
assert asyncio.current_task() is task
assert asyncio.get_running_loop() is loop
assert threading.get_ident() == thread
assert context.get() == "at_await"
owner.close()
payload["closure_mutation"] = True
context.set("inside")
saved.append(value)
shared["before"] = True
loop.call_soon(gate.set)
await gate.wait()
inner = owners.prepare(nested, (value,), awaited=True)
try:
assert await inner.invoke() is value
finally:
inner.close()
return value
owner = owners.prepare(callback, (payload,), {"shared": alias}, awaited=True)
pending = owner.invoke()
context.set("at_await")
assert await pending is payload
assert context.get() == "nested"
assert payload["closure_mutation"] is True
assert owners.live == 0
alias["after"] = True
assert saved[0]["nested"] == {"before": True, "nested_call": True, "after": True}
async def exceptions(owners):
for error in (RuntimeError("original"), KeyboardInterrupt("original"), asyncio.CancelledError("original")):
cause = ValueError("cause")
payload = {}
async def callback():
payload["changed"] = True
raise error from cause
owner = owners.prepare(callback, (), awaited=True)
try:
await owner.invoke()
assert False, "exception lost"
except BaseException as caught:
assert caught is error and caught.__cause__ is cause
frames = []
tb = caught.__traceback__
while tb:
frames.append(tb.tb_frame.f_code.co_name)
tb = tb.tb_next
assert "callback" in frames
finally:
owner.close()
assert payload["changed"] is True
async def exception_ownership(owners):
class Callback:
def __init__(self, error):
self.error = error
async def __call__(self, value):
value.changed = True
raise self.error
value = Value()
error = RuntimeError("retained exception")
callback = Callback(error)
value_ref, callback_ref = weakref.ref(value), weakref.ref(callback)
owner = owners.prepare(callback, (value,), awaited=True)
del value, callback
try:
await owner.invoke()
assert False, "expected callback failure"
except RuntimeError as caught:
assert caught is error
owner.close()
assert value_ref().changed and callback_ref() is not None
del error
gc.collect()
assert value_ref() is None and callback_ref() is None
async def cancellation_before_start(owners):
started = []
async def callback(value):
started.append(value)
for operation in ("close", "cancel"):
value = Value()
ref = weakref.ref(value)
owner = owners.prepare(callback, (value,), awaited=True)
pending = owner.invoke()
owner.close()
del value
assert ref() is not None
if operation == "close":
pending.close()
else:
task = asyncio.create_task(pending)
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
del task
del pending
await checkpoint()
gc.collect()
assert ref() is None
assert started == []
async def cancellation_case(owners, repeated=False, suppress=False):
started, cleaning, finish = asyncio.Event(), asyncio.Event(), asyncio.Event()
value = Value()
ref = weakref.ref(value)
observed = []
async def callback(argument):
try:
started.set()
await asyncio.Event().wait()
except asyncio.CancelledError:
cleaning.set()
try:
await finish.wait()
except asyncio.CancelledError:
observed.append("second cancellation")
await finish.wait()
argument.cleaned = True
if suppress:
return argument
raise
owner = owners.prepare(callback, (value,), awaited=True)
task = asyncio.create_task(owner.invoke())
owner.close()
del value
try:
await started.wait()
task.cancel()
await cleaning.wait()
assert ref() is not None and not task.done()
if repeated:
task.cancel()
barrier = asyncio.Event()
asyncio.get_running_loop().call_soon(barrier.set)
await barrier.wait()
assert observed == ["second cancellation"] and not task.done()
finish.set()
try:
result = await task
assert suppress and result is ref() and result.cleaned
del result
except asyncio.CancelledError:
assert not suppress
assert ref().cleaned
finally:
finish.set()
if not task.done():
task.cancel()
await asyncio.gather(task, return_exceptions=True)
del task
await checkpoint()
gc.collect()
assert ref() is None
async def cancellation_unwinds(owners):
await cancellation_case(owners)
async def cancellation_during_cleanup(owners):
await cancellation_case(owners, repeated=True)
async def cancellation_suppressed(owners):
await cancellation_case(owners, suppress=True)
async def registration_and_gc(owners):
original = Value()
original_ref = weakref.ref(original)
registered = owners.prepare(lambda value: value, (original,))
active = registered.clone_owner()
registered.close()
registered = owners.prepare(lambda: "replacement", ())
del original
assert active.invoke() is original_ref()
active.close()
gc.collect()
assert original_ref() is None
assert registered.invoke() == "replacement"
registered.close()
for edge in ("callable", "positional", "keywords"):
class Callback:
def __call__(self, *args, **kwargs):
pass
value = Callback()
owner = owners.prepare(
value if edge == "callable" else lambda *a, **k: None,
(value,) if edge == "positional" else (),
{"value": value} if edge == "keywords" else None,
)
value.owner = owner
value_ref, owner_ref = weakref.ref(value), weakref.ref(owner)
del value, owner
gc.collect()
assert value_ref() is None and owner_ref() is None
assert owners.live == 0
finalized = []
class Reenter:
def __del__(self):
try:
reentrant.close()
another = owners.prepare(lambda: 42, ())
finalized.append(another.invoke())
another.close()
except BaseException as error:
finalized.append(type(error).__name__)
value = Reenter()
reentrant = owners.prepare(lambda value: None, (value,))
del value
reentrant.close()
reentrant.close()
assert finalized == [42]
assert owners.live == 0
async def background_and_session(owners):
context = contextvars.ContextVar("background_context", default="initial")
start, finish = asyncio.Event(), asyncio.Event()
payload = {"nested": {"value": "queued"}}
saved = []
async def upload(value):
assert context.get() == "submission"
start.set()
await finish.wait()
saved.append(json.dumps(value))
session = owners.prepare(lambda value: value, (payload,))
first_response, second_response = session.clone_owner(), session.clone_owner()
upload_owner = owners.prepare(upload, (payload,), awaited=True)
context.set("submission")
task = asyncio.create_task(upload_owner.invoke())
context.set("consumer")
upload_owner.close()
first_response.close()
await start.wait()
assert second_response.invoke() is payload
payload["nested"]["value"] = "later"
consumed = json.dumps(payload)
payload["nested"]["after_consumption"] = True
assert "after_consumption" not in consumed
second_response.close()
session.close()
assert owners.live == 0 and not task.done()
finish.set()
await task
assert json.loads(saved[0]) == payload
assert context.get() == "consumer"
async def stream_lifecycle(owners):
for terminal in ("exhaustion", "failure", "close"):
nested = {"usage": 0}
item = {"nested": nested}
closed = []
async def source():
try:
yield item
nested["usage"] = 12
if terminal == "failure":
raise ValueError("stream failure")
finally:
closed.append(True)
stream = source()
pull = owners.prepare(stream.__anext__, (), awaited=True)
yielded = await pull.invoke()
assert yielded is item
shallow, deep = copy.copy(yielded), copy.deepcopy(yielded)
retained = owners.prepare(lambda value: value, (yielded["nested"],))
yielded["nested"] = {"replacement": True}
if terminal == "close":
close = owners.prepare(stream.aclose, (), awaited=True)
await close.invoke()
await close.invoke()
close.close()
else:
try:
await pull.invoke()
assert False, "expected stream termination"
except (StopAsyncIteration, ValueError) as error:
assert isinstance(error, ValueError) == (terminal == "failure")
pull.close()
assert closed == [True]
assert retained.invoke() is nested
assert shallow["nested"] is nested and deep["nested"]["usage"] == 0
assert nested["usage"] == (0 if terminal == "close" else 12)
retained.close()
async def sync_stream_lifecycle(owners):
for terminal in ("exhaustion", "failure", "close"):
value = {"nested": {"usage": 0}}
closed = []
def source():
try:
yield value
value["nested"]["usage"] = 12
if terminal == "failure":
raise ValueError("stream failure")
finally:
closed.append(True)
stream = source()
pull = owners.prepare(stream.__next__, ())
assert pull.invoke() is value
saved = owners.prepare(lambda value: value, (value,))
if terminal == "close":
close = owners.prepare(stream.close, ())
close.invoke()
close.invoke()
close.close()
else:
try:
pull.invoke()
assert False, "stream did not terminate"
except (StopIteration, ValueError) as error:
assert isinstance(error, ValueError) == (terminal == "failure")
pull.close()
assert saved.invoke() is value
assert value["nested"]["usage"] == (0 if terminal == "close" else 12)
assert closed == [True]
saved.close()
async def repeated_ownership(owners):
refs = []
for batch in range(8):
gate = asyncio.Event()
async def work(value):
await gate.wait()
return None
tasks = []
for index in range(4):
value = Value()
refs.append(weakref.ref(value))
owner = owners.prepare(work, (value,), awaited=True)
tasks.append(asyncio.create_task(owner.invoke()))
owner.close()
del value
gate.set()
await asyncio.gather(*tasks)
del tasks
gc.collect()
assert owners.live == 0
assert all(ref() is None for ref in refs)
def run_scenario(name, retained, factory):
owners = factory if retained else ReferenceFactory()
async def run():
async with asyncio.timeout(15):
await globals()[name](owners)
assert owners.live == 0
pending = asyncio.all_tasks() - {asyncio.current_task()}
assert not pending, f"undrained tasks: {pending}"
asyncio.run(run())
gc.collect()
assert owners.live == 0

View file

@ -1,8 +1,9 @@
use std::ffi::CStr;
use litellm_python_interop::PreparedCall;
use litellm_python_interop::{InvocationMode, InvocationOutcome, PreparedCall};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use rstest::{fixture, rstest};
fn scope<'py>(py: Python<'py>, source: &CStr) -> PyResult<Bound<'py, PyDict>> {
let globals = PyDict::new(py);
@ -14,9 +15,11 @@ fn item<'py>(globals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> {
globals.get_item(name).unwrap().unwrap()
}
#[test]
fn retains_aliases_mutations_and_original_result() -> PyResult<()> {
Python::initialize();
#[rstest]
fn retains_aliases_mutations_and_original_result(
initialized_python: &InitializedPython,
) -> PyResult<()> {
let _ = initialized_python;
Python::attach(|py| {
let globals = scope(
py,
@ -35,11 +38,12 @@ def callback(data, *, alias):
let keywords = PyDict::new(py);
keywords.set_item("alias", item(&globals, "shared"))?;
let invocation = PreparedCall::new(
InvocationMode::Direct,
item(&globals, "callback").unbind(),
PyTuple::new(py, [&payload])?.unbind(),
Some(keywords.unbind()),
);
let result = invocation.invoke(py)?;
let result = invoke_direct(&invocation, py)?;
assert!(result.bind(py).is(&payload));
drop(invocation);
py.run(
@ -55,9 +59,11 @@ assert saved[0]['nested']['value'] == 'after'
})
}
#[test]
fn preserves_exception_identity_cause_traceback_and_prior_mutation() -> PyResult<()> {
Python::initialize();
#[rstest]
fn preserves_exception_identity_cause_traceback_and_prior_mutation(
initialized_python: &InitializedPython,
) -> PyResult<()> {
let _ = initialized_python;
Python::attach(|py| {
let globals = scope(
py,
@ -71,11 +77,12 @@ def callback(data):
",
)?;
let invocation = PreparedCall::new(
InvocationMode::Direct,
item(&globals, "callback").unbind(),
PyTuple::new(py, [item(&globals, "payload")])?.unbind(),
None,
);
let error = invocation.invoke(py).unwrap_err();
let error = invoke_direct(&invocation, py).unwrap_err();
assert!(error.value(py).is(item(&globals, "error")));
drop(invocation);
py.run(
@ -91,9 +98,9 @@ assert traceback.extract_tb(error.__traceback__)[-1].name == 'callback'
})
}
#[test]
fn returns_coroutine_without_executing_it() -> PyResult<()> {
Python::initialize();
#[rstest]
fn returns_coroutine_without_executing_it(initialized_python: &InitializedPython) -> PyResult<()> {
let _ = initialized_python;
Python::attach(|py| {
let globals = scope(
py,
@ -108,11 +115,12 @@ def callback():
",
)?;
let invocation = PreparedCall::new(
InvocationMode::Direct,
item(&globals, "callback").unbind(),
PyTuple::empty(py).unbind(),
None,
);
let result = invocation.invoke(py)?;
let result = invoke_direct(&invocation, py)?;
assert!(result.bind(py).is(item(&globals, "coroutine")));
py.run(
c"
@ -128,12 +136,22 @@ coroutine.close()
#[pyfunction]
fn reenter(py: Python<'_>, callback: Py<PyAny>, payload: Py<PyAny>) -> PyResult<Py<PyAny>> {
PreparedCall::new(callback, PyTuple::new(py, [payload])?.unbind(), None).invoke(py)
invoke_direct(
&PreparedCall::new(
InvocationMode::Direct,
callback,
PyTuple::new(py, [payload])?.unbind(),
None,
),
py,
)
}
#[test]
fn preserves_current_context_thread_and_reentry() -> PyResult<()> {
Python::initialize();
#[rstest]
fn preserves_current_context_thread_and_reentry(
initialized_python: &InitializedPython,
) -> PyResult<()> {
let _ = initialized_python;
Python::attach(|py| {
let globals = scope(
py,
@ -159,11 +177,12 @@ def outer():
)?;
globals.set_item("reenter", wrap_pyfunction!(reenter, py)?)?;
let invocation = PreparedCall::new(
InvocationMode::Direct,
item(&globals, "outer").unbind(),
PyTuple::empty(py).unbind(),
None,
);
let result = invocation.invoke(py)?;
let result = invoke_direct(&invocation, py)?;
assert!(result.bind(py).is(item(&globals, "payload")));
py.run(
c"
@ -179,9 +198,11 @@ finally:
})
}
#[test]
fn owns_arguments_until_release_and_preserves_callback_retention() -> PyResult<()> {
Python::initialize();
#[rstest]
fn owns_arguments_until_release_and_preserves_callback_retention(
initialized_python: &InitializedPython,
) -> PyResult<()> {
let _ = initialized_python;
let (invocation, globals) = Python::attach(|py| {
let globals = scope(
py,
@ -206,6 +227,7 @@ callback_ref = weakref.ref(callback)
let keywords = PyDict::new(py);
keywords.set_item("other", item(&globals, "other"))?;
let invocation = PreparedCall::new(
InvocationMode::Direct,
item(&globals, "callback").unbind(),
PyTuple::new(py, [item(&globals, "value")])?.unbind(),
Some(keywords.unbind()),
@ -220,7 +242,7 @@ callback_ref = weakref.ref(callback)
Some(globals),
None,
)?;
assert!(invocation.invoke(py)?.is_none(py));
assert!(invoke_direct(&invocation, py)?.is_none(py));
drop(invocation);
py.run(
c"
@ -249,16 +271,19 @@ fn prepare_pre_call(
keywords.set_item("api_key", py.None())?;
keywords.set_item("additional_args", view)?;
Ok(PreparedCall::new(
InvocationMode::Direct,
logger.getattr("pre_call")?.unbind(),
PyTuple::empty(py).unbind(),
Some(keywords.unbind()),
))
}
#[test]
#[rstest]
#[ignore = "requires the repository Python environment and LiteLLM on PYTHONPATH"]
fn real_ocr_logging_preserves_execution_roots_and_continues_after_error() -> PyResult<()> {
Python::initialize();
fn real_ocr_logging_preserves_execution_roots_and_continues_after_error(
initialized_python: &InitializedPython,
) -> PyResult<()> {
let _ = initialized_python;
Python::attach(|py| {
let globals = scope(
py,
@ -272,6 +297,7 @@ class Retain(CustomLogger):
self.view = kwargs['additional_args']
self.headers = self.view['headers']
self.body = self.view['complete_input_dict']
return {'ignored_replacement': True}
class MutateThenFail(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
@ -303,7 +329,7 @@ logger = Logging(
let body = item(&globals, "body").unbind();
let view = item(&globals, "view").cast_into::<PyDict>()?.unbind();
let invocation = prepare_pre_call(py, &item(&globals, "logger"), view.bind(py))?;
assert!(invocation.invoke(py)?.is_none(py));
assert!(invoke_direct(&invocation, py)?.is_none(py));
drop(invocation);
py.run(c"del headers, body, view", Some(&globals), None)?;
assert!(
@ -339,3 +365,19 @@ assert first.body['document']['value'] == 'after invocation'
)
})
}
struct InitializedPython;
#[fixture]
#[once]
fn initialized_python() -> InitializedPython {
Python::initialize();
InitializedPython
}
fn invoke_direct(call: &PreparedCall, py: Python<'_>) -> PyResult<Py<PyAny>> {
match call.invoke(py)? {
InvocationOutcome::Returned(value) => Ok(value),
InvocationOutcome::Awaitable(_) => panic!("direct binding produced an awaitable outcome"),
}
}

View file

@ -0,0 +1,114 @@
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use litellm_python_interop::{InvocationMode, InvocationOutcome, PreparedCall};
use pyo3::class::gc::{PyTraverseError, PyVisit};
use pyo3::exceptions::PyRuntimeError;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
#[pyclass]
#[derive(Default)]
pub struct OwnerFactory {
live: Arc<AtomicUsize>,
}
#[pymethods]
impl OwnerFactory {
#[getter]
fn live(&self) -> usize {
self.live.load(Ordering::SeqCst)
}
#[pyo3(signature = (callable, positional, keywords=None, awaited=false))]
fn prepare(
&self,
callable: Py<PyAny>,
positional: Py<PyTuple>,
keywords: Option<Py<PyDict>>,
awaited: bool,
) -> Owner {
self.live.fetch_add(1, Ordering::SeqCst);
Owner {
call: Some(PreparedCall::new(
if awaited {
InvocationMode::Await
} else {
InvocationMode::Direct
},
callable,
positional,
keywords,
)),
live: self.live.clone(),
}
}
}
#[pyclass(weakref)]
struct Owner {
call: Option<PreparedCall>,
live: Arc<AtomicUsize>,
}
impl Owner {
fn release(&mut self) -> Option<PreparedCall> {
let call = self.call.take();
if call.is_some() {
self.live.fetch_sub(1, Ordering::SeqCst);
}
call
}
}
#[pymethods]
impl Owner {
fn invoke(slf: &Bound<'_, Self>) -> PyResult<Py<PyAny>> {
let py = slf.py();
let call = slf
.borrow()
.call
.as_ref()
.map(|call| call.clone_ref(py))
.ok_or_else(|| PyRuntimeError::new_err("invocation owner released"))?;
match call.invoke(py)? {
InvocationOutcome::Returned(value) | InvocationOutcome::Awaitable(value) => Ok(value),
}
}
fn clone_owner(&self, py: Python<'_>) -> PyResult<Self> {
let call = self
.call
.as_ref()
.ok_or_else(|| PyRuntimeError::new_err("invocation owner released"))?;
self.live.fetch_add(1, Ordering::SeqCst);
Ok(Self {
call: Some(call.clone_ref(py)),
live: self.live.clone(),
})
}
fn close(slf: &Bound<'_, Self>) {
let released = slf.borrow_mut().release();
drop(released);
}
fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> {
if let Some(call) = &self.call {
call.traverse(visit)?;
}
Ok(())
}
fn __clear__(slf: &Bound<'_, Self>) {
Self::close(slf);
}
}
impl Drop for Owner {
fn drop(&mut self) {
drop(self.release());
}
}