Skip to content

Commit d4e9e0f

Browse files
davidzhaoclaude
andcommitted
rpc interceptors: enforce the deadline across the chain, register by identity
Review follow-ups on the interceptor hook: - The caller's response_timeout now covers the whole incoming chain (_run_incoming_chain wraps interceptors + handler in one wait_for), so time an interceptor spends before or after next() counts against it instead of handing the handler a fresh full timeout. Deadline expiry cancels the chain and maps to RESPONSE_TIMEOUT; cancellation from outside (room disconnect) still maps to RECIPIENT_DISCONNECTED. _invoke_rpc_handler no longer wraps the handler itself. - add_rpc_interceptor / remove_rpc_interceptor compare by identity, so two distinct interceptors that compare equal coexist and only the exact instance is removed. - mypy: cast the awaited handler result; fix the non-overlapping identity check in the test. Tests: deadline burned before and after next(), outside cancellation, identity registration; the timeout test now asserts interceptors observe the cancellation. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
1 parent e55b093 commit d4e9e0f

4 files changed

Lines changed: 120 additions & 33 deletions

File tree

‎README.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -225,7 +225,7 @@ You may find it useful to adjust the `response_timeout` parameter, which indicat
225225

226226
#### Intercepting RPC calls
227227

228-
An `RpcInterceptor` wraps every RPC the local participant performs or handles, which is useful for logging, tracing, or attaching metadata to payloads. Each method receives the call and a `next` continuation; return what `next` returns. Interceptors run in the order they were added, the first being outermost, and errors from the remote side or from your handler flow through them unchanged.
228+
An `RpcInterceptor` wraps every RPC the local participant performs or handles, which is useful for logging, tracing, or attaching metadata to payloads. Each method receives the call and a `next` continuation; return what `next` returns. Interceptors run in the order they were added, the first being outermost, and errors from the remote side or from your handler flow through them unchanged. On the incoming side the caller's `response_timeout` covers the whole chain, interceptors included.
229229

230230
```python
231231
class TimingInterceptor(rtc.RpcInterceptor):

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

Lines changed: 27 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -477,7 +477,8 @@ def add_rpc_interceptor(self, interceptor: RpcInterceptor) -> None:
477477
Args:
478478
interceptor (RpcInterceptor): The interceptor to add.
479479
"""
480-
if interceptor not in self._rpc_interceptors:
480+
# identity, not equality: two distinct interceptors that compare equal must coexist
481+
if not any(existing is interceptor for existing in self._rpc_interceptors):
481482
self._rpc_interceptors.append(interceptor)
482483

483484
def remove_rpc_interceptor(self, interceptor: RpcInterceptor) -> None:
@@ -488,8 +489,9 @@ def remove_rpc_interceptor(self, interceptor: RpcInterceptor) -> None:
488489
Args:
489490
interceptor (RpcInterceptor): The interceptor to remove.
490491
"""
491-
if interceptor in self._rpc_interceptors:
492-
self._rpc_interceptors.remove(interceptor)
492+
self._rpc_interceptors = [
493+
existing for existing in self._rpc_interceptors if existing is not interceptor
494+
]
493495

494496
def register_rpc_method(
495497
self,
@@ -598,9 +600,8 @@ async def _handle_rpc_method_invocation(
598600
request_id, caller_identity, payload, response_timeout, method=method
599601
)
600602

601-
handle = _chain_incoming(list(self._rpc_interceptors), self._invoke_rpc_handler)
602603
try:
603-
response_payload = await handle(params)
604+
response_payload = await self._run_incoming_chain(params)
604605
except RpcError as error:
605606
response_error = error
606607
except Exception:
@@ -625,26 +626,36 @@ async def _handle_rpc_method_invocation(
625626
err = res.rpc_method_invocation_response.error
626627
logger.error(f"error sending rpc method invocation response: {err}")
627628

629+
async def _run_incoming_chain(self, invocation: RpcInvocationData) -> Optional[str]:
630+
"""Run the interceptor chain and the handler under the caller's response deadline.
631+
632+
The deadline covers the whole chain, so time an interceptor spends before or after
633+
``next`` counts against it; when it passes, the chain is cancelled and the caller
634+
gets ``RESPONSE_TIMEOUT``. Cancellation from outside (the room disconnecting) maps to
635+
``RECIPIENT_DISCONNECTED``, as before.
636+
"""
637+
handle = _chain_incoming(list(self._rpc_interceptors), self._invoke_rpc_handler)
638+
try:
639+
return await asyncio.wait_for(handle(invocation), timeout=invocation.response_timeout)
640+
except asyncio.TimeoutError:
641+
raise RpcError._built_in(RpcError.ErrorCode.RESPONSE_TIMEOUT) from None
642+
except asyncio.CancelledError:
643+
raise RpcError._built_in(RpcError.ErrorCode.RECIPIENT_DISCONNECTED) from None
644+
628645
async def _invoke_rpc_handler(self, invocation: RpcInvocationData) -> Optional[str]:
629646
"""Run the registered handler for ``invocation`` (the innermost step of the chain).
630647
631-
Raises ``RpcError`` for an unregistered method, a handler timeout, or a cancelled
632-
handler; any other exception from the handler propagates unchanged so interceptors
633-
can observe it before the caller's response is built.
648+
Raises ``RpcError(UNSUPPORTED_METHOD)`` when nothing is registered; any exception
649+
from the handler propagates unchanged so interceptors can observe it before the
650+
caller's response is built. The response deadline is enforced by the caller around
651+
the whole chain.
634652
"""
635653
handler = self._rpc_handlers.get(invocation.method)
636654
if not handler:
637655
raise RpcError._built_in(RpcError.ErrorCode.UNSUPPORTED_METHOD)
638656

639657
if asyncio.iscoroutinefunction(handler):
640-
try:
641-
return await asyncio.wait_for(
642-
handler(invocation), timeout=invocation.response_timeout
643-
)
644-
except asyncio.TimeoutError:
645-
raise RpcError._built_in(RpcError.ErrorCode.RESPONSE_TIMEOUT) from None
646-
except asyncio.CancelledError:
647-
raise RpcError._built_in(RpcError.ErrorCode.RECIPIENT_DISCONNECTED) from None
658+
return cast(Optional[str], await handler(invocation))
648659
return cast(Optional[str], handler(invocation))
649660

650661
async def set_metadata(self, metadata: str) -> None:

‎livekit-rtc/livekit/rtc/rpc.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -81,7 +81,9 @@ class RpcInterceptor:
8181
reaches the caller. On the incoming side, any other exception raised by the handler is
8282
also visible; the SDK converts it to ``APPLICATION_ERROR`` only after the chain returns,
8383
and a call for an unregistered method reaches the chain with ``next`` raising
84-
``UNSUPPORTED_METHOD``.
84+
``UNSUPPORTED_METHOD``. The caller's ``response_timeout`` covers the whole incoming
85+
chain: when it passes, the chain is cancelled (interceptors see ``CancelledError``) and
86+
the caller receives ``RESPONSE_TIMEOUT``.
8587
8688
Example:
8789
Time every RPC in both directions::

‎livekit-rtc/tests/test_rpc_interceptors.py‎

Lines changed: 89 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -130,8 +130,9 @@ async def greet(data: RpcInvocationData) -> str:
130130

131131
lp._rpc_handlers["greet"] = greet # register without the FFI round trip
132132

133-
handle = _incoming_chain(lp)
134-
result = await handle(RpcInvocationData("req-1", "alice", "{}", 5.0, method="greet"))
133+
result = await lp._run_incoming_chain(
134+
RpcInvocationData("req-1", "alice", "{}", 5.0, method="greet")
135+
)
135136

136137
assert result == "hello alice"
137138
assert log == ["a:in:greet>", "a:in:greet<"]
@@ -152,10 +153,11 @@ async def intercept_incoming(
152153
raise
153154

154155
lp.add_rpc_interceptor(Observe())
155-
handle = _incoming_chain(lp)
156156

157157
with pytest.raises(rtc.RpcError) as info:
158-
await handle(RpcInvocationData("req-1", "alice", "{}", 5.0, method="missing"))
158+
await lp._run_incoming_chain(
159+
RpcInvocationData("req-1", "alice", "{}", 5.0, method="missing")
160+
)
159161
assert info.value.code == rtc.RpcError.ErrorCode.UNSUPPORTED_METHOD
160162
assert seen == [info.value]
161163

@@ -187,20 +189,77 @@ def sync_ok(data: RpcInvocationData) -> str:
187189
return "sync"
188190

189191
lp._rpc_handlers.update({"boom": boom, "slow": slow, "sync_ok": sync_ok})
190-
handle = _incoming_chain(lp)
191192

192193
# the raw handler exception reaches interceptors; the SDK maps it to APPLICATION_ERROR
193194
# only when building the response, outside the chain
194195
with pytest.raises(ValueError):
195-
await handle(RpcInvocationData("r1", "alice", "{}", 5.0, method="boom"))
196+
await lp._run_incoming_chain(RpcInvocationData("r1", "alice", "{}", 5.0, method="boom"))
196197
assert isinstance(seen[-1], ValueError)
197198

199+
# the deadline cancels the chain (interceptors see the cancellation) and the caller
200+
# gets RESPONSE_TIMEOUT
201+
with pytest.raises(rtc.RpcError) as info:
202+
await lp._run_incoming_chain(RpcInvocationData("r2", "alice", "{}", 0.01, method="slow"))
203+
assert info.value.code == rtc.RpcError.ErrorCode.RESPONSE_TIMEOUT
204+
assert isinstance(seen[-1], asyncio.CancelledError)
205+
206+
result = await lp._run_incoming_chain(
207+
RpcInvocationData("r3", "alice", "{}", 5.0, method="sync_ok")
208+
)
209+
assert result == "sync"
210+
211+
212+
@pytest.mark.parametrize("delay_position", ["before_next", "after_next"])
213+
async def test_response_deadline_covers_interceptor_time(delay_position: str) -> None:
214+
"""An interceptor that burns the deadline, on either side of ``next``, times the call out
215+
instead of handing the handler a fresh full timeout."""
216+
lp = _participant()
217+
handler_ran = asyncio.Event()
218+
219+
class Slow(rtc.RpcInterceptor):
220+
async def intercept_incoming(
221+
self, invocation: RpcInvocationData, next: IncomingRpcNext
222+
) -> Optional[str]:
223+
if delay_position == "before_next":
224+
await asyncio.sleep(1.0)
225+
result = await next(invocation)
226+
if delay_position == "after_next":
227+
await asyncio.sleep(1.0)
228+
return result
229+
230+
async def fast(data: RpcInvocationData) -> str:
231+
handler_ran.set()
232+
return "fast"
233+
234+
lp.add_rpc_interceptor(Slow())
235+
lp._rpc_handlers["fast"] = fast
236+
237+
start = asyncio.get_running_loop().time()
198238
with pytest.raises(rtc.RpcError) as info:
199-
await handle(RpcInvocationData("r2", "alice", "{}", 0.01, method="slow"))
239+
await lp._run_incoming_chain(RpcInvocationData("r1", "alice", "{}", 0.05, method="fast"))
200240
assert info.value.code == rtc.RpcError.ErrorCode.RESPONSE_TIMEOUT
201-
assert seen[-1] is info.value
241+
assert asyncio.get_running_loop().time() - start < 0.5
242+
assert handler_ran.is_set() == (delay_position == "after_next")
202243

203-
assert await handle(RpcInvocationData("r3", "alice", "{}", 5.0, method="sync_ok")) == "sync"
244+
245+
async def test_outside_cancellation_maps_to_recipient_disconnected() -> None:
246+
lp = _participant()
247+
started = asyncio.Event()
248+
249+
async def hang(data: RpcInvocationData) -> str:
250+
started.set()
251+
await asyncio.sleep(10)
252+
return "never"
253+
254+
lp._rpc_handlers["hang"] = hang
255+
task = asyncio.ensure_future(
256+
lp._run_incoming_chain(RpcInvocationData("r1", "alice", "{}", 5.0, method="hang"))
257+
)
258+
await started.wait()
259+
task.cancel() # what the room does when it disconnects mid-invocation
260+
with pytest.raises(rtc.RpcError) as info:
261+
await task
262+
assert info.value.code == rtc.RpcError.ErrorCode.RECIPIENT_DISCONNECTED
204263

205264

206265
async def test_add_and_remove_interceptors() -> None:
@@ -224,13 +283,28 @@ async def fake_ffi(call: RpcCallInfo) -> str:
224283
assert log == []
225284

226285

286+
def test_registration_is_by_identity_not_equality() -> None:
287+
class Equalish(rtc.RpcInterceptor):
288+
def __eq__(self, other: object) -> bool:
289+
return isinstance(other, Equalish)
290+
291+
def __hash__(self) -> int:
292+
return 1
293+
294+
lp = _participant()
295+
first, second = Equalish(), Equalish()
296+
assert first == second and first is not second
297+
298+
lp.add_rpc_interceptor(first)
299+
lp.add_rpc_interceptor(second)
300+
assert len(lp._rpc_interceptors) == 2, "distinct interceptors that compare equal coexist"
301+
302+
lp.remove_rpc_interceptor(second)
303+
assert len(lp._rpc_interceptors) == 1
304+
assert lp._rpc_interceptors[0] is first, "only the exact instance is removed"
305+
306+
227307
def test_invocation_data_defaults_keep_positional_construction() -> None:
228308
# existing code constructs it positionally without `method`
229309
data = RpcInvocationData("req", "alice", "{}", 2.5)
230310
assert data.method == ""
231-
232-
233-
def _incoming_chain(lp: LocalParticipant) -> IncomingRpcNext:
234-
from livekit.rtc.rpc import _chain_incoming
235-
236-
return _chain_incoming(list(lp._rpc_interceptors), lp._invoke_rpc_handler)

0 commit comments

Comments
 (0)