mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(rust): complete retained callback lifecycle foundation
This commit is contained in:
parent
e399ec52f1
commit
fc1877718d
8 changed files with 1139 additions and 35 deletions
2
.github/workflows/test-rust.yml
vendored
2
.github/workflows/test-rust.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
153
litellm-rust/crates/python-interop/tests/callback_lifecycle.rs
Normal file
153
litellm-rust/crates/python-interop/tests/callback_lifecycle.rs
Normal 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(())
|
||||
})
|
||||
}
|
||||
211
litellm-rust/crates/python-interop/tests/fixtures/callback_components.py
vendored
Normal file
211
litellm-rust/crates/python-interop/tests/fixtures/callback_components.py
vendored
Normal 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()
|
||||
519
litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py
vendored
Normal file
519
litellm-rust/crates/python-interop/tests/fixtures/callback_lifecycle.py
vendored
Normal 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
|
||||
|
|
@ -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"),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue