4848 TrackPublication ,
4949)
5050from .transcription import Transcription
51- from .rpc import RpcError
51+ from .rpc import (
52+ RpcCallInfo ,
53+ RpcError ,
54+ RpcInterceptor ,
55+ RpcInvocationData ,
56+ _chain_incoming ,
57+ _chain_outgoing ,
58+ )
5259from ._proto .rpc_pb2 import RpcMethodInvocationResponseRequest
5360from .log import logger
5461
55- from .rpc import RpcInvocationData
5662from .data_stream import (
5763 TextStreamWriter ,
5864 TextStreamInfo ,
@@ -247,6 +253,7 @@ def __init__(
247253 self ._room_queue = room_queue
248254 self ._track_publications : dict [str , LocalTrackPublication ] = {}
249255 self ._rpc_handlers : Dict [str , RpcHandler ] = {}
256+ self ._rpc_interceptors : List [RpcInterceptor ] = []
250257 # Handles of data stream writers that have been opened but not yet
251258 # closed, so the room can drop them at disconnect. The FFI close
252259 # request consumes the handle (take_handle), so an entry is removed as
@@ -427,15 +434,27 @@ async def perform_rpc(
427434 Raises:
428435 RpcError: On failure. Details in `message`.
429436 """
437+ call = RpcCallInfo (
438+ destination_identity = destination_identity ,
439+ method = method ,
440+ payload = payload ,
441+ response_timeout = response_timeout ,
442+ max_round_trip_latency = max_round_trip_latency ,
443+ )
444+ # snapshot the interceptor list so add/remove during a call is well defined
445+ perform = _chain_outgoing (list (self ._rpc_interceptors ), self ._perform_rpc_ffi )
446+ return await perform (call )
447+
448+ async def _perform_rpc_ffi (self , call : RpcCallInfo ) -> str :
430449 req = proto_ffi .FfiRequest ()
431450 req .perform_rpc .local_participant_handle = self ._ffi_handle .handle
432- req .perform_rpc .destination_identity = destination_identity
433- req .perform_rpc .method = method
434- req .perform_rpc .payload = payload
435- if response_timeout is not None :
436- req .perform_rpc .response_timeout_ms = int (response_timeout * 1000 )
437- if max_round_trip_latency is not None :
438- req .perform_rpc .max_round_trip_latency_ms = int (max_round_trip_latency * 1000 )
451+ req .perform_rpc .destination_identity = call . destination_identity
452+ req .perform_rpc .method = call . method
453+ req .perform_rpc .payload = call . payload
454+ if call . response_timeout is not None :
455+ req .perform_rpc .response_timeout_ms = int (call . response_timeout * 1000 )
456+ if call . max_round_trip_latency is not None :
457+ req .perform_rpc .max_round_trip_latency_ms = int (call . max_round_trip_latency * 1000 )
439458
440459 queue = FfiClient .instance .queue .subscribe ()
441460 try :
@@ -449,6 +468,29 @@ async def perform_rpc(
449468
450469 return cast (str , cb .perform_rpc .payload )
451470
471+ def add_rpc_interceptor (self , interceptor : RpcInterceptor ) -> None :
472+ """
473+ Add an :class:`RpcInterceptor` that wraps every RPC this participant performs or
474+ handles. Interceptors run in the order they were added, the first being outermost.
475+ Adding the same instance twice is a no-op.
476+
477+ Args:
478+ interceptor (RpcInterceptor): The interceptor to add.
479+ """
480+ if interceptor not in self ._rpc_interceptors :
481+ self ._rpc_interceptors .append (interceptor )
482+
483+ def remove_rpc_interceptor (self , interceptor : RpcInterceptor ) -> None :
484+ """
485+ Remove a previously added :class:`RpcInterceptor`. Calls already in flight keep the
486+ chain they started with.
487+
488+ Args:
489+ interceptor (RpcInterceptor): The interceptor to remove.
490+ """
491+ if interceptor in self ._rpc_interceptors :
492+ self ._rpc_interceptors .remove (interceptor )
493+
452494 def register_rpc_method (
453495 self ,
454496 method_name : str ,
@@ -552,33 +594,21 @@ async def _handle_rpc_method_invocation(
552594 response_error : Optional [RpcError ] = None
553595 response_payload : Optional [str ] = None
554596
555- params = RpcInvocationData (request_id , caller_identity , payload , response_timeout )
556-
557- handler = self . _rpc_handlers . get ( method )
597+ params = RpcInvocationData (
598+ request_id , caller_identity , payload , response_timeout , method = method
599+ )
558600
559- if not handler :
560- response_error = RpcError ._built_in (RpcError .ErrorCode .UNSUPPORTED_METHOD )
561- else :
562- try :
563- if asyncio .iscoroutinefunction (handler ):
564- try :
565- response_payload = await asyncio .wait_for (
566- handler (params ), timeout = response_timeout
567- )
568- except asyncio .TimeoutError :
569- raise RpcError ._built_in (RpcError .ErrorCode .RESPONSE_TIMEOUT )
570- except asyncio .CancelledError :
571- raise RpcError ._built_in (RpcError .ErrorCode .RECIPIENT_DISCONNECTED )
572- else :
573- response_payload = cast (Optional [str ], handler (params ))
574- except RpcError as error :
575- response_error = error
576- except Exception :
577- logger .exception (
578- f"Uncaught error returned by RPC handler for { method } . "
579- "Returning APPLICATION_ERROR instead. "
580- )
581- response_error = RpcError ._built_in (RpcError .ErrorCode .APPLICATION_ERROR )
601+ handle = _chain_incoming (list (self ._rpc_interceptors ), self ._invoke_rpc_handler )
602+ try :
603+ response_payload = await handle (params )
604+ except RpcError as error :
605+ response_error = error
606+ except Exception :
607+ logger .exception (
608+ f"Uncaught error returned by RPC handler for { method } . "
609+ "Returning APPLICATION_ERROR instead. "
610+ )
611+ response_error = RpcError ._built_in (RpcError .ErrorCode .APPLICATION_ERROR )
582612
583613 req = proto_ffi .FfiRequest (
584614 rpc_method_invocation_response = RpcMethodInvocationResponseRequest (
@@ -595,6 +625,28 @@ async def _handle_rpc_method_invocation(
595625 err = res .rpc_method_invocation_response .error
596626 logger .error (f"error sending rpc method invocation response: { err } " )
597627
628+ async def _invoke_rpc_handler (self , invocation : RpcInvocationData ) -> Optional [str ]:
629+ """Run the registered handler for ``invocation`` (the innermost step of the chain).
630+
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.
634+ """
635+ handler = self ._rpc_handlers .get (invocation .method )
636+ if not handler :
637+ raise RpcError ._built_in (RpcError .ErrorCode .UNSUPPORTED_METHOD )
638+
639+ 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
648+ return cast (Optional [str ], handler (invocation ))
649+
598650 async def set_metadata (self , metadata : str ) -> None :
599651 """
600652 Set the metadata for the local participant.
0 commit comments