mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
675 lines
21 KiB
Python
675 lines
21 KiB
Python
import asyncio
|
|
import contextvars
|
|
import copy
|
|
import gc
|
|
import inspect
|
|
import json
|
|
import sys
|
|
import threading
|
|
import weakref
|
|
from unittest import TestCase
|
|
|
|
|
|
async def checkpoint():
|
|
ready = asyncio.Event()
|
|
asyncio.get_running_loop().call_soon(ready.set)
|
|
await ready.wait()
|
|
|
|
|
|
class Value:
|
|
pass
|
|
|
|
|
|
def cold_awaited_adapter_reentry(factory):
|
|
events = []
|
|
compilations = []
|
|
|
|
def invoke_failure():
|
|
error = LookupError("cold-cache callback error")
|
|
|
|
async def failing():
|
|
raise error
|
|
|
|
owner = factory.prepare(failing, (), awaited=True)
|
|
pending = owner.invoke()
|
|
try:
|
|
assert inspect.getcoroutinestate(pending) == inspect.CORO_CREATED
|
|
with TestCase().assertRaises(LookupError) as caught:
|
|
pending.send(None)
|
|
assert caught.exception is error
|
|
finally:
|
|
pending.close()
|
|
owner.close()
|
|
error.__traceback__ = None
|
|
|
|
def audit(event, args):
|
|
if event != "compile" or args[1] != "retained_callback.py":
|
|
return
|
|
compilations.append(args[1])
|
|
if len(compilations) == 1:
|
|
events.append("entered")
|
|
invoke_failure()
|
|
events.append("nested completed")
|
|
|
|
sys.addaudithook(audit)
|
|
invoke_failure()
|
|
events.append("outer completed")
|
|
assert events == ["entered", "nested completed", "outer completed"]
|
|
assert len(compilations) == 2
|
|
assert factory.live == 0
|
|
|
|
result = object()
|
|
direct = factory.prepare(lambda value: value, (result,))
|
|
try:
|
|
assert direct.invoke() is result
|
|
finally:
|
|
direct.close()
|
|
|
|
async def successful(value, *, alias):
|
|
assert value is alias
|
|
await checkpoint()
|
|
return value
|
|
|
|
owner = factory.prepare(successful, (result,), {"alias": result}, awaited=True)
|
|
try:
|
|
assert asyncio.run(owner.invoke()) is result
|
|
finally:
|
|
owner.close()
|
|
assert len(compilations) == 2
|
|
assert factory.live == 0
|
|
|
|
|
|
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=future: 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)
|
|
with TestCase().assertRaises(TypeError):
|
|
await owner.invoke()
|
|
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=payload, error=error, cause=cause):
|
|
payload["changed"] = True
|
|
raise error from cause
|
|
|
|
owner = owners.prepare(callback, (), awaited=True)
|
|
caught_error = None
|
|
try:
|
|
await owner.invoke()
|
|
except BaseException as caught:
|
|
caught_error = caught
|
|
finally:
|
|
owner.close()
|
|
assert caught_error is error and caught_error.__cause__ is cause
|
|
frames = []
|
|
tb = caught_error.__traceback__
|
|
while tb:
|
|
frames.append(tb.tb_frame.f_code.co_name)
|
|
tb = tb.tb_next
|
|
assert "callback" in frames
|
|
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
|
|
caught_error = None
|
|
try:
|
|
await owner.invoke()
|
|
except RuntimeError as caught:
|
|
caught_error = caught
|
|
assert caught_error is error
|
|
owner.close()
|
|
assert value_ref().changed and callback_ref() is not None
|
|
del error, caught_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(item=item, nested=nested, terminal=terminal, closed=closed):
|
|
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:
|
|
with TestCase().assertRaises(ValueError if terminal == "failure" else StopAsyncIteration):
|
|
await pull.invoke()
|
|
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(value=value, terminal=terminal, closed=closed):
|
|
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:
|
|
with TestCase().assertRaises(ValueError if terminal == "failure" else StopIteration):
|
|
pull.invoke()
|
|
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, gate=gate):
|
|
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)
|
|
|
|
|
|
async def retained_field_replacement(owners):
|
|
original = {"messages": [{"content": "original"}]}
|
|
replacement = {"messages": [{"content": "replacement"}]}
|
|
event = {"payload": original, "alias": original}
|
|
saved = []
|
|
|
|
def retain(value):
|
|
saved.append(value["payload"])
|
|
|
|
def replace(value):
|
|
value["payload"] = replacement
|
|
value["alias"]["messages"][0]["content"] = "mutated original"
|
|
|
|
for callback in (retain, replace):
|
|
owner = owners.prepare(callback, (event,))
|
|
try:
|
|
assert owner.invoke() is None
|
|
finally:
|
|
owner.close()
|
|
assert saved[0] is original is event["alias"]
|
|
assert event["payload"] is replacement
|
|
assert saved[0]["messages"][0]["content"] == "mutated original"
|
|
replacement["messages"][0]["content"] = "mutated replacement"
|
|
assert event["payload"]["messages"][0]["content"] == "mutated replacement"
|
|
assert original["messages"][0]["content"] == "mutated original"
|
|
|
|
|
|
async def queued_graph_ownership(owners):
|
|
queue = asyncio.Queue()
|
|
sentinel = Value()
|
|
reference = weakref.ref(sentinel)
|
|
payload = {"sentinel": sentinel, "nested": {"status": "queued"}}
|
|
snapshot = json.dumps(payload["nested"])
|
|
enqueue = owners.prepare(queue.put_nowait, (payload,))
|
|
try:
|
|
enqueue.invoke()
|
|
finally:
|
|
enqueue.close()
|
|
del sentinel, payload
|
|
gc.collect()
|
|
assert owners.live == 0 and reference() is not None
|
|
queued = queue.get_nowait()
|
|
queued["nested"]["status"] = "changed before flush"
|
|
assert json.loads(json.dumps(queued["nested"])) == {"status": "changed before flush"}
|
|
assert json.loads(snapshot) == {"status": "queued"}
|
|
assert queued["sentinel"] is reference()
|
|
queue.task_done()
|
|
del queued
|
|
gc.collect()
|
|
assert reference() is None
|
|
|
|
|
|
async def detached_work_after_error(owners):
|
|
for raises in (False, True):
|
|
entered, release = asyncio.Event(), asyncio.Event()
|
|
tasks, observed = [], []
|
|
value = Value()
|
|
value.status = "before return"
|
|
reference = weakref.ref(value)
|
|
|
|
async def consume(argument, entered=entered, release=release, observed=observed):
|
|
entered.set()
|
|
await release.wait()
|
|
observed.append(argument.status)
|
|
|
|
def callback(argument, tasks=tasks, consume=consume, raises=raises):
|
|
tasks.append(asyncio.create_task(consume(argument)))
|
|
if raises:
|
|
raise ValueError("after task creation")
|
|
|
|
owner = owners.prepare(callback, (value,))
|
|
try:
|
|
if raises:
|
|
with TestCase().assertRaisesRegex(ValueError, "after task creation"):
|
|
owner.invoke()
|
|
else:
|
|
assert owner.invoke() is None
|
|
finally:
|
|
owner.close()
|
|
del value
|
|
try:
|
|
await entered.wait()
|
|
assert owners.live == 0 and reference() is not None
|
|
reference().status = "after invocation"
|
|
release.set()
|
|
await tasks[0]
|
|
assert observed == ["after invocation"]
|
|
finally:
|
|
release.set()
|
|
await asyncio.gather(*tasks, return_exceptions=True)
|
|
tasks.clear()
|
|
await checkpoint()
|
|
gc.collect()
|
|
assert reference() is None
|
|
|
|
|
|
def run_checked(owners, scenario):
|
|
baseline = owners.live
|
|
|
|
async def run():
|
|
await asyncio.wait_for(scenario, timeout=15)
|
|
assert owners.live == baseline
|
|
pending = asyncio.all_tasks() - {asyncio.current_task()}
|
|
assert not pending, f"undrained tasks: {pending}"
|
|
|
|
asyncio.run(run())
|
|
gc.collect()
|
|
assert owners.live == baseline
|
|
|
|
|
|
def run_scenario(name, retained, factory):
|
|
owners = factory if retained else ReferenceFactory()
|
|
assert owners.live == 0
|
|
run_checked(owners, globals()[name](owners))
|