diff --git a/litellm-rust/crates/python-interop/src/callback.rs b/litellm-rust/crates/python-interop/src/callback.rs index fbdf7962616..3942b150daf 100644 --- a/litellm-rust/crates/python-interop/src/callback.rs +++ b/litellm-rust/crates/python-interop/src/callback.rs @@ -1,22 +1,42 @@ -use pyo3::class::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; +use pyo3::pyclass::{PyTraverseError, PyVisit}; use pyo3::sync::PyOnceLock; use pyo3::types::{PyDict, PyTuple}; -static AWAIT_CALL: PyOnceLock> = PyOnceLock::new(); +use crate::constants::{ + AWAIT_ADAPTER_FILENAME, AWAIT_ADAPTER_FUNCTION, AWAIT_ADAPTER_MODULE, AWAIT_ADAPTER_SOURCE, +}; +static AWAIT_ADAPTER: PyOnceLock> = PyOnceLock::new(); + +/// How a retained callback is bound to its caller. +/// +/// `Direct` mirrors `callable(*args, **kwargs)`: a coroutine returned by the +/// callable is handed back untouched and never awaited. `Await` mirrors +/// `await callable(*args, **kwargs)` inline in the caller's task. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum InvocationMode { Direct, Await, } +/// Result of [`PreparedCall::invoke`]. +/// +/// `Awaitable` carries an adapter coroutine that has not yet called the +/// callback. The callback runs, and any exception it raises surfaces, only +/// when Python drives that coroutine. #[derive(Debug)] pub enum InvocationOutcome { Returned(Py), Awaitable(Py), } +/// A callback plus its arguments, retained as owning Python references. +/// +/// Arguments are passed to the callback by identity, never copied, so the +/// callback observes and may mutate the caller's objects. Dropping the value +/// releases the references; a Python-visible owner must also expose them to +/// the cycle collector via [`PreparedCall::traverse`]. pub struct PreparedCall { mode: InvocationMode, callable: Py, @@ -49,29 +69,9 @@ impl PreparedCall { self.keywords.as_ref().map(|kwargs| kwargs.bind(py)), ) .map(InvocationOutcome::Returned), - InvocationMode::Await => { - if AWAIT_CALL.get(py).is_none() { - let adapter = 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")? - .unbind(); - let _ = AWAIT_CALL.set(py, adapter); - } - - let adapter = AWAIT_CALL.get(py).unwrap(); - - adapter - .call1(py, (&self.callable, &self.positional, &self.keywords)) - .map(InvocationOutcome::Awaitable) - } + InvocationMode::Await => await_adapter(py)? + .call1(py, (&self.callable, &self.positional, &self.keywords)) + .map(InvocationOutcome::Awaitable), } } @@ -93,3 +93,21 @@ impl PreparedCall { Ok(()) } } + +/// Compiling the adapter runs Python, which may re-enter this function through +/// audit hooks. `PyOnceLock` forbids re-entrant initialization, so compile +/// first and only publish a finished adapter into the cell. +fn await_adapter(py: Python<'_>) -> PyResult<&Py> { + if let Some(adapter) = AWAIT_ADAPTER.get(py) { + return Ok(adapter); + } + let compiled = PyModule::from_code( + py, + AWAIT_ADAPTER_SOURCE, + AWAIT_ADAPTER_FILENAME, + AWAIT_ADAPTER_MODULE, + )? + .getattr(AWAIT_ADAPTER_FUNCTION)? + .unbind(); + Ok(AWAIT_ADAPTER.get_or_init(py, || compiled)) +} diff --git a/litellm-rust/crates/python-interop/src/constants.rs b/litellm-rust/crates/python-interop/src/constants.rs new file mode 100644 index 00000000000..d56e5babc86 --- /dev/null +++ b/litellm-rust/crates/python-interop/src/constants.rs @@ -0,0 +1,18 @@ +use std::ffi::CStr; + +/// Python source of the coroutine adapter that awaits a retained callback +/// inline in the caller's task. It is compiled once per interpreter. +pub(crate) const AWAIT_ADAPTER_SOURCE: &CStr = + c"async def invoke_awaited(callable, positional, keywords): + if keywords is None: + return await callable(*positional) + return await callable(*positional, **keywords) +"; + +/// Filename recorded on the adapter's code object. Visible to Python +/// `compile` audit hooks and tracebacks. +pub const AWAIT_ADAPTER_FILENAME: &CStr = c"retained_callback.py"; + +pub(crate) const AWAIT_ADAPTER_MODULE: &CStr = c"_retained_callback"; + +pub(crate) const AWAIT_ADAPTER_FUNCTION: &str = "invoke_awaited"; diff --git a/litellm-rust/crates/python-interop/src/lib.rs b/litellm-rust/crates/python-interop/src/lib.rs index 7c259aa4672..28c21d5ca3a 100644 --- a/litellm-rust/crates/python-interop/src/lib.rs +++ b/litellm-rust/crates/python-interop/src/lib.rs @@ -1,7 +1,9 @@ mod callback; +mod constants; mod gil; mod marshal; pub use callback::{InvocationMode, InvocationOutcome, PreparedCall}; +pub use constants::AWAIT_ADAPTER_FILENAME; 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 index 7a39ddbd607..9a87d19c17f 100644 --- a/litellm-rust/crates/python-interop/tests/callback_lifecycle.rs +++ b/litellm-rust/crates/python-interop/tests/callback_lifecycle.rs @@ -64,6 +64,14 @@ fn scenario_scope(initialized_python: &InitializedPython) -> Py { Py::new(py, callback_owner::OwnerFactory::default()).unwrap(), ) .unwrap(); + globals + .set_item( + "AWAIT_ADAPTER_FILENAME", + litellm_python_interop::AWAIT_ADAPTER_FILENAME + .to_str() + .unwrap(), + ) + .unwrap(); run_fixture( py, &globals, @@ -135,7 +143,6 @@ fn component_contract( } #[rstest] -#[case::real_logging_queue_chain("real_logging_queue_chain")] #[case::real_logging_queue_copy_control("real_logging_queue_copy_control")] #[case::real_crowdstrike_translator_identity("real_crowdstrike_translator_identity")] #[case::real_rubrik_block_lifecycle("real_rubrik_block_lifecycle")] @@ -198,10 +205,8 @@ fn run_scenario_fixture( #[case::shallow_result("result_identity", "result_shallow")] #[case::deep_result("result_identity", "result_deep")] #[case::retained_lifetime("deferred_lifetime", "identity")] -#[case::expired_borrow("deferred_lifetime", "weak")] #[case::prepared_ownership("deferred_lifetime", "missing_handoff")] #[case::externally_owned_retained("borrowed_lifetime", "identity")] -#[case::externally_owned_borrow("borrowed_lifetime", "weak")] #[case::original_coroutine("direct_coroutine", "identity")] #[case::passthrough_coroutine("direct_coroutine", "result_passthrough")] #[serial(python_interpreter)] @@ -215,6 +220,20 @@ fn control_contract( run_control_fixture(scenario_scope, witness, control, retained, awaited) } +// The `weak` control wraps nothing: it holds only weak references and never +// calls the factory, so the `retained` axis has no effect on it. +#[rstest] +#[case::expired_borrow("deferred_lifetime")] +#[case::externally_owned_borrow("borrowed_lifetime")] +#[serial(python_interpreter)] +fn weak_control( + scenario_scope: Py, + #[case] witness: &str, + #[values(false, true)] awaited: bool, +) -> PyResult<()> { + run_control_fixture(scenario_scope, witness, "weak", false, awaited) +} + #[rstest] #[case::retained("identity")] #[case::missing_handoff("missing_handoff")] diff --git a/litellm-rust/crates/python-interop/tests/fixtures/callback_controls.py b/litellm-rust/crates/python-interop/tests/fixtures/callback_controls.py index 9463e34d0f1..6c123b64ca5 100644 --- a/litellm-rust/crates/python-interop/tests/fixtures/callback_controls.py +++ b/litellm-rust/crates/python-interop/tests/fixtures/callback_controls.py @@ -7,6 +7,8 @@ from collections.abc import Awaitable, Callable from dataclasses import dataclass from typing import Protocol, cast +from callback_lifecycle import ReferenceFactory, run_checked, settle + class PreparedInvocation(Protocol): def invoke(self) -> object: ... @@ -256,7 +258,7 @@ async def argument_identity(owners: CallFactory, awaited: bool) -> IdentityObser owner = owners.prepare(observe_async if awaited else observe, (original,), {"alias": nested}, awaited=awaited) try: pending = owner.invoke() - return await pending if awaited else pending + return await settle(pending, awaited) finally: owner.close() @@ -283,7 +285,7 @@ async def mutation_timing(owners: CallFactory, awaited: bool) -> TimingObservati original.stage = nested.stage = 1 pending = owner.invoke() original.stage = nested.stage = 2 - return await pending if awaited else pending + return await settle(pending, awaited) finally: owner.close() @@ -300,7 +302,7 @@ async def result_identity(owners: CallFactory, awaited: bool) -> IdentityObserva owner = owners.prepare(callback_async if awaited else callback, (), awaited=awaited) try: pending = owner.invoke() - result = await pending if awaited else pending + result = await settle(pending, awaited) return IdentityObservation(result is original, result.nested is original.nested, True) finally: owner.close() @@ -349,7 +351,7 @@ async def deferred_lifetime(owners: CallFactory, awaited: bool) -> LifetimeObser gc.collect() alive = tuple(reference() is not None for reference in references) pending = owner.invoke() - result = await pending if awaited else pending + result = await settle(pending, awaited) observation = LifetimeObservation(alive, result) finally: owner.close() @@ -381,7 +383,7 @@ async def borrowed_lifetime(owners: CallFactory, awaited: bool) -> BorrowedObser value.stage, alias.stage = 13, 29 pending = owner.invoke() value.stage, alias.stage = 17, 31 - return await pending if awaited else pending + return await settle(pending, awaited) finally: owner.close() @@ -467,11 +469,22 @@ def expected_control(witness: str, control: str, awaited: bool) -> object: def run_control(witness: str, control: str, retained: bool, awaited: bool, factory: LiveCallFactory) -> None: - inner = factory if retained else cast(Callable[[], LiveCallFactory], globals()["ReferenceFactory"])() + inner = factory if retained else ReferenceFactory() owners = control_factory(control, inner) async def run() -> None: - observed = await globals()[witness](owners, awaited) + observed = await WITNESSES[witness](owners, awaited) assert observed == expected_control(witness, control, awaited), (witness, control, awaited, observed) - globals()["run_checked"](inner, run()) + run_checked(inner, run()) + + +WITNESSES: dict[str, Callable[[CallFactory, bool], object]] = { + "argument_identity": argument_identity, + "mutation_timing": mutation_timing, + "result_identity": result_identity, + "deferred_lifetime": deferred_lifetime, + "borrowed_lifetime": borrowed_lifetime, + "pending_handoff": pending_handoff, + "direct_coroutine": direct_coroutine, +} diff --git a/litellm-rust/crates/python-interop/tests/fixtures/callback_integrations.py b/litellm-rust/crates/python-interop/tests/fixtures/callback_integrations.py index 210d9f848c6..430d7a656ea 100644 --- a/litellm-rust/crates/python-interop/tests/fixtures/callback_integrations.py +++ b/litellm-rust/crates/python-interop/tests/fixtures/callback_integrations.py @@ -6,7 +6,6 @@ import threading from collections import OrderedDict from dataclasses import dataclass from datetime import datetime -from functools import wraps from typing import Literal from unittest import TestCase @@ -40,23 +39,6 @@ def integration_response(url, body, status=200, headers=None): return httpx.Response(status, json=body, headers=headers, request=httpx.Request("POST", url)) -def integration_callback_scope(scenario): - @wraps(scenario) - async def run(owners): - callbacks = tuple(litellm.callbacks) - try: - return await scenario(owners) - finally: - litellm.callbacks[:] = callbacks - - return run - - -@integration_callback_scope -async def real_logging_queue_chain(owners): - return await integration_logging_queue_case(owners) - - @dataclass(frozen=True, slots=True) class QueueObservation: gcs_model_parameters: str @@ -64,7 +46,6 @@ class QueueObservation: literal_prepared_settings: str -@integration_callback_scope async def real_logging_queue_copy_control(owners): baseline = await integration_logging_queue_case(owners) copied = await integration_logging_queue_case(owners, literal_copy="payload") @@ -239,7 +220,6 @@ async def integration_logging_queue_case(owners, *, literal_copy: Literal["direc ) -@integration_callback_scope async def real_crowdstrike_translator_identity(owners): entered, release = asyncio.Event(), asyncio.Event() calls = [] @@ -304,7 +284,6 @@ async def real_crowdstrike_translator_identity(owners): assert detached[0]["content"] == user["content"] == "private text" -@integration_callback_scope async def real_rubrik_block_lifecycle(owners): for input_type, populated in (("request", False), ("response", True)): entered, release = asyncio.Event(), asyncio.Event() @@ -419,7 +398,6 @@ async def real_rubrik_block_lifecycle(owners): await asyncio.gather(rubrik._periodic_flush_task, return_exceptions=True) -@integration_callback_scope async def real_parallel_guardrail_snapshots(owners): original_mode = litellm.safe_memory_mode try: @@ -516,7 +494,6 @@ async def integration_parallel_snapshot_case(owners): await asyncio.gather(task, return_exceptions=True) -@integration_callback_scope async def real_purview_sync_background(owners): entered, release = threading.Event(), threading.Event() calls, workers = [], [] diff --git a/litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py b/litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py index d4ff68b343e..c1e133db971 100644 --- a/litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py +++ b/litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py @@ -16,6 +16,10 @@ async def checkpoint(): await ready.wait() +async def settle(pending, awaited): + return await pending if awaited else pending + + class Value: pass diff --git a/litellm-rust/crates/python-interop/tests/support/callback_owner.rs b/litellm-rust/crates/python-interop/tests/support/callback_owner.rs index 4cb38c27316..c7f39dad294 100644 --- a/litellm-rust/crates/python-interop/tests/support/callback_owner.rs +++ b/litellm-rust/crates/python-interop/tests/support/callback_owner.rs @@ -4,9 +4,9 @@ use std::sync::{ }; use litellm_python_interop::{InvocationMode, InvocationOutcome, PreparedCall}; -use pyo3::class::gc::{PyTraverseError, PyVisit}; use pyo3::exceptions::PyRuntimeError; use pyo3::prelude::*; +use pyo3::pyclass::{PyTraverseError, PyVisit}; use pyo3::types::{PyDict, PyTuple}; #[pyclass] diff --git a/litellm-rust/crates/python-interop/tests/support/python.rs b/litellm-rust/crates/python-interop/tests/support/python.rs index 5ef9c7e7b37..eb2e6c240dc 100644 --- a/litellm-rust/crates/python-interop/tests/support/python.rs +++ b/litellm-rust/crates/python-interop/tests/support/python.rs @@ -21,6 +21,12 @@ impl InitializedPython { #[once] pub fn initialized_python() -> InitializedPython { Python::initialize(); + Python::attach(|py| { + let sys = py.import("sys").unwrap(); + let path = sys.getattr("path").unwrap(); + let fixtures = concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures"); + path.call_method1("append", (fixtures,)).unwrap(); + }); InitializedPython }