adding session_timeout in fastapi

This commit is contained in:
Vaibhav159
2024-11-21 14:56:42 +05:30
parent 7dfa886669
commit 6e8e7fa19a

View File

@@ -44,11 +44,13 @@ except ModuleNotFoundError as e:
class FastAPIWebsocketParams(TransportParams): class FastAPIWebsocketParams(TransportParams):
add_wav_header: bool = False add_wav_header: bool = False
serializer: FrameSerializer serializer: FrameSerializer
session_timeout: int | None = None
class FastAPIWebsocketCallbacks(BaseModel): class FastAPIWebsocketCallbacks(BaseModel):
on_client_connected: Callable[[WebSocket], Awaitable[None]] on_client_connected: Callable[[WebSocket], Awaitable[None]]
on_client_disconnected: Callable[[WebSocket], Awaitable[None]] on_client_disconnected: Callable[[WebSocket], Awaitable[None]]
on_session_timeout: Callable[[WebSocket], Awaitable[None]]
class FastAPIWebsocketInputTransport(BaseInputTransport): class FastAPIWebsocketInputTransport(BaseInputTransport):
@@ -67,6 +69,7 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
self._monitor_websocket_task = self.get_event_loop().create_task(self._monitor_websocket())
await self._callbacks.on_client_connected(self._websocket) await self._callbacks.on_client_connected(self._websocket)
self._receive_task = self.get_event_loop().create_task(self._receive_messages()) self._receive_task = self.get_event_loop().create_task(self._receive_messages())
@@ -88,6 +91,16 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
await self._callbacks.on_client_disconnected(self._websocket) await self._callbacks.on_client_disconnected(self._websocket)
async def _monitor_websocket(self):
"""
Wait for self._params.session_timeout seconds, if the websocket is still open, trigger timeout event.
"""
try:
await asyncio.sleep(self._params.session_timeout)
await self._callbacks.on_session_timeout(self._websocket)
except asyncio.CancelledError:
logger.info(f"Monitoring task cancelled for: {self._websocket}")
class FastAPIWebsocketOutputTransport(BaseOutputTransport): class FastAPIWebsocketOutputTransport(BaseOutputTransport):
def __init__(self, websocket: WebSocket, params: FastAPIWebsocketParams, **kwargs): def __init__(self, websocket: WebSocket, params: FastAPIWebsocketParams, **kwargs):
@@ -163,6 +176,7 @@ class FastAPIWebsocketTransport(BaseTransport):
self._callbacks = FastAPIWebsocketCallbacks( self._callbacks = FastAPIWebsocketCallbacks(
on_client_connected=self._on_client_connected, on_client_connected=self._on_client_connected,
on_client_disconnected=self._on_client_disconnected, on_client_disconnected=self._on_client_disconnected,
on_session_timeout=self._on_session_timeout,
) )
self._input = FastAPIWebsocketInputTransport( self._input = FastAPIWebsocketInputTransport(
@@ -176,6 +190,7 @@ class FastAPIWebsocketTransport(BaseTransport):
# these handlers. # these handlers.
self._register_event_handler("on_client_connected") self._register_event_handler("on_client_connected")
self._register_event_handler("on_client_disconnected") self._register_event_handler("on_client_disconnected")
self._register_event_handler("on_session_timeout")
def input(self) -> FastAPIWebsocketInputTransport: def input(self) -> FastAPIWebsocketInputTransport:
return self._input return self._input
@@ -188,3 +203,6 @@ class FastAPIWebsocketTransport(BaseTransport):
async def _on_client_disconnected(self, websocket): async def _on_client_disconnected(self, websocket):
await self._call_event_handler("on_client_disconnected", websocket) await self._call_event_handler("on_client_disconnected", websocket)
async def _on_session_timeout(self, websocket):
await self._call_event_handler("on_session_timeout", websocket)