Skip to content

Commit 91a23d8

Browse files
davidzhaoclaude
andcommitted
rpc interceptors: bound the unwind after outside cancellation; await any awaitable a handler returns
Awaiting the chain directly forwarded an outside cancel (the room disconnecting) into the chain and kept the invocation task parked until the chain finished, so a handler that swallowed cancellation held room.disconnect() up to the caller's deadline. The chain is now awaited through a shield: the cancel is raised here at once, the chain is cancelled explicitly and given a bounded time to unwind, and a handler that does not stop is named in a warning. Handlers are typed as any callable returning a payload or an awaitable of one; dispatch now checks the result with inspect.isawaitable instead of iscoroutinefunction, which missed callable objects with an async __call__ and sync wrappers returning a coroutine. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
1 parent 1db942a commit 91a23d8

2 files changed

Lines changed: 95 additions & 8 deletions

File tree

‎livekit-rtc/livekit/rtc/participant.py‎

Lines changed: 28 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616

1717
import ctypes
1818
import asyncio
19+
import inspect
1920
import datetime
2021
import enum
2122
import os
@@ -237,6 +238,10 @@ def disconnect_reason(
237238

238239
RpcHandler = Callable[["RpcInvocationData"], Union[Awaitable[Optional[str]], Optional[str]]]
239240

241+
# how long a cancelled incoming RPC chain gets to unwind before the caller is answered
242+
# without it (the room's disconnect waits on these invocations)
243+
_RPC_CANCEL_UNWIND_TIMEOUT = 2.0
244+
240245

241246
F = TypeVar(
242247
"F", bound=Callable[[RpcInvocationData], Union[Awaitable[Optional[str]], Optional[str]]]
@@ -658,14 +663,25 @@ def _on_deadline() -> None:
658663

659664
deadline = loop.call_later(invocation.response_timeout, _on_deadline)
660665
try:
661-
return await chain_task
666+
# shielded: a cancel from outside (the room disconnecting) is raised here at
667+
# once. Awaiting the chain directly would instead forward the cancel to it and
668+
# keep this task parked until the chain finished, so a handler that ignored
669+
# cancellation held up room.disconnect() for as long as the caller's deadline.
670+
return await asyncio.shield(chain_task)
662671
except asyncio.CancelledError:
663672
if deadline_fired:
664673
raise RpcError._built_in(RpcError.ErrorCode.RESPONSE_TIMEOUT) from None
665-
# cancelled from outside: awaiting propagated the cancel into the chain; let it
666-
# finish unwinding before answering the caller
667-
if not chain_task.done():
668-
await asyncio.wait([chain_task])
674+
# cancelled from outside: stop the chain and let it unwind before answering the
675+
# caller, but not for long; this is the path room.disconnect() waits on
676+
chain_task.cancel()
677+
_, pending = await asyncio.wait([chain_task], timeout=_RPC_CANCEL_UNWIND_TIMEOUT)
678+
if pending:
679+
logger.warning(
680+
"RPC handler for %s did not stop within %.1fs of being cancelled; "
681+
"answering the caller without it",
682+
invocation.method,
683+
_RPC_CANCEL_UNWIND_TIMEOUT,
684+
)
669685
raise RpcError._built_in(RpcError.ErrorCode.RECIPIENT_DISCONNECTED) from None
670686
except Exception:
671687
if deadline_fired:
@@ -686,9 +702,13 @@ async def _invoke_rpc_handler(self, invocation: RpcInvocationData) -> Optional[s
686702
if not handler:
687703
raise RpcError._built_in(RpcError.ErrorCode.UNSUPPORTED_METHOD)
688704

689-
if asyncio.iscoroutinefunction(handler):
690-
return cast(Optional[str], await handler(invocation))
691-
return cast(Optional[str], handler(invocation))
705+
# RpcHandler admits any callable returning a payload or an awaitable of one: a
706+
# coroutine function, but also a callable object with an async __call__ or a sync
707+
# wrapper handing back a coroutine, which iscoroutinefunction would not recognize
708+
result = handler(invocation)
709+
if inspect.isawaitable(result):
710+
result = await result
711+
return result
692712

693713
async def set_metadata(self, metadata: str) -> None:
694714
"""

‎livekit-rtc/tests/test_rpc_interceptors.py‎

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -321,6 +321,73 @@ async def hang(data: RpcInvocationData) -> str:
321321
assert info.value.code == rtc.RpcError.ErrorCode.RECIPIENT_DISCONNECTED
322322

323323

324+
async def test_unwind_after_outside_cancellation_is_bounded(
325+
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
326+
) -> None:
327+
"""A handler that ignores cancellation must not hang the disconnect that cancelled it:
328+
the caller is answered once the unwind bound passes, and the offender is named.
329+
330+
Awaiting the chain directly could never bound this: asyncio forwards the cancel to the
331+
chain and keeps the awaiting task parked until the chain finishes, so the wait would
332+
only begin once the stubborn handler had already stopped (at the caller's deadline)."""
333+
from livekit.rtc import participant as participant_mod
334+
335+
monkeypatch.setattr(participant_mod, "_RPC_CANCEL_UNWIND_TIMEOUT", 0.05)
336+
lp = _participant()
337+
started = asyncio.Event()
338+
release = asyncio.Event()
339+
340+
async def ignores_cancel(data: RpcInvocationData) -> str:
341+
started.set()
342+
try:
343+
await asyncio.sleep(10)
344+
except asyncio.CancelledError:
345+
pass
346+
await release.wait() # keeps running long after it was told to stop
347+
return "late"
348+
349+
lp._rpc_handlers["stubborn"] = ignores_cancel
350+
task = asyncio.ensure_future(
351+
lp._run_incoming_chain(RpcInvocationData("r1", "alice", "{}", 5.0, method="stubborn"))
352+
)
353+
await started.wait()
354+
task.cancel()
355+
loop = asyncio.get_running_loop()
356+
t0 = loop.time()
357+
with pytest.raises(rtc.RpcError) as info:
358+
await task
359+
assert info.value.code == rtc.RpcError.ErrorCode.RECIPIENT_DISCONNECTED
360+
assert loop.time() - t0 < 1.0
361+
assert any("stubborn" in r.getMessage() for r in caplog.records)
362+
363+
release.set() # let the orphaned handler finish so the loop closes clean
364+
await asyncio.sleep(0)
365+
366+
367+
async def test_handlers_returning_an_awaitable_are_awaited() -> None:
368+
"""RpcHandler admits any callable returning a payload or an awaitable of one, not only
369+
coroutine functions: a callable object with an async __call__, a sync wrapper handing
370+
back a coroutine, and a plain sync handler all work."""
371+
lp = _participant()
372+
373+
class Handler:
374+
async def __call__(self, data: RpcInvocationData) -> str:
375+
return f"obj:{data.payload}"
376+
377+
async def _inner(data: RpcInvocationData) -> str:
378+
return f"wrapped:{data.payload}"
379+
380+
lp._rpc_handlers["obj"] = Handler()
381+
lp._rpc_handlers["wrapped"] = lambda data: _inner(data) # sync, returns a coroutine
382+
lp._rpc_handlers["sync"] = lambda data: f"sync:{data.payload}"
383+
384+
for method, expected in (("obj", "obj:x"), ("wrapped", "wrapped:x"), ("sync", "sync:x")):
385+
result = await lp._run_incoming_chain(
386+
RpcInvocationData("r1", "alice", "x", 1.0, method=method)
387+
)
388+
assert result == expected, method
389+
390+
324391
async def test_add_and_remove_interceptors() -> None:
325392
lp = _participant()
326393
log: list[str] = []

0 commit comments

Comments
 (0)