adding session_timeout in fastapi
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user