diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index add38eb66e7..31b23d5a5c9 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -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 diff --git a/litellm-rust/crates/python-interop/src/callback.rs b/litellm-rust/crates/python-interop/src/callback.rs index 951a79b41e0..6c73f7c9a5d 100644 --- a/litellm-rust/crates/python-interop/src/callback.rs +++ b/litellm-rust/crates/python-interop/src/callback.rs @@ -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> = PyOnceLock::new(); + +#[derive(Clone, Copy)] +pub enum InvocationMode { + Direct, + Await, +} + +#[derive(Debug)] +pub enum InvocationOutcome { + Returned(Py), + Awaitable(Py), +} + pub struct PreparedCall { + mode: InvocationMode, callable: Py, positional: Py, keywords: Option>, } impl PreparedCall { - pub fn new(callable: Py, positional: Py, keywords: Option>) -> Self { + pub fn new( + mode: InvocationMode, + callable: Py, + positional: Py, + keywords: Option>, + ) -> Self { Self { + mode, callable, positional, keywords, } } - pub fn invoke(&self, py: Python<'_>) -> PyResult> { - self.callable.call( - py, - self.positional.bind(py), - self.keywords.as_ref().map(|kwargs| kwargs.bind(py)), - ) + pub fn invoke(&self, py: Python<'_>) -> PyResult { + 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(()) } } diff --git a/litellm-rust/crates/python-interop/src/lib.rs b/litellm-rust/crates/python-interop/src/lib.rs index a34d30e2a3a..7c259aa4672 100644 --- a/litellm-rust/crates/python-interop/src/lib.rs +++ b/litellm-rust/crates/python-interop/src/lib.rs @@ -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}; diff --git a/litellm-rust/crates/python-interop/tests/callback_lifecycle.rs b/litellm-rust/crates/python-interop/tests/callback_lifecycle.rs new file mode 100644 index 00000000000..69f0705515f --- /dev/null +++ b/litellm-rust/crates/python-interop/tests/callback_lifecycle.rs @@ -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 { + 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, + #[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, + #[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, #[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(()) + }) +} diff --git a/litellm-rust/crates/python-interop/tests/fixtures/callback_components.py b/litellm-rust/crates/python-interop/tests/fixtures/callback_components.py new file mode 100644 index 00000000000..38324fb6a66 --- /dev/null +++ b/litellm-rust/crates/python-interop/tests/fixtures/callback_components.py @@ -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() diff --git a/litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py b/litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py new file mode 100644 index 00000000000..0b4bbe2f1fc --- /dev/null +++ b/litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py @@ -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 diff --git a/litellm-rust/crates/python-interop/tests/prepared_call.rs b/litellm-rust/crates/python-interop/tests/prepared_call.rs index e3219601285..100cb6383e5 100644 --- a/litellm-rust/crates/python-interop/tests/prepared_call.rs +++ b/litellm-rust/crates/python-interop/tests/prepared_call.rs @@ -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> { 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, payload: Py) -> PyResult> { - 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::()?.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> { + match call.invoke(py)? { + InvocationOutcome::Returned(value) => Ok(value), + InvocationOutcome::Awaitable(_) => panic!("direct binding produced an awaitable outcome"), + } +} diff --git a/litellm-rust/crates/python-interop/tests/support/callback_owner.rs b/litellm-rust/crates/python-interop/tests/support/callback_owner.rs new file mode 100644 index 00000000000..4cb38c27316 --- /dev/null +++ b/litellm-rust/crates/python-interop/tests/support/callback_owner.rs @@ -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, +} + +#[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, + positional: Py, + keywords: Option>, + 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, + live: Arc, +} + +impl Owner { + fn release(&mut self) -> Option { + 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> { + 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 { + 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()); + } +}