transports(websockets): cancel or wait for tasks to finish
This commit is contained in:
@@ -16,6 +16,8 @@ from loguru import logger
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
|
CancelFrame,
|
||||||
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
InputAudioRawFrame,
|
InputAudioRawFrame,
|
||||||
OutputAudioRawFrame,
|
OutputAudioRawFrame,
|
||||||
@@ -27,6 +29,7 @@ from pipecat.serializers.base_serializer import FrameSerializer, FrameSerializer
|
|||||||
from pipecat.transports.base_input import BaseInputTransport
|
from pipecat.transports.base_input import BaseInputTransport
|
||||||
from pipecat.transports.base_output import BaseOutputTransport
|
from pipecat.transports.base_output import BaseOutputTransport
|
||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
|
from pipecat.utils.asyncio import cancel_task
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from fastapi import WebSocket
|
from fastapi import WebSocket
|
||||||
@@ -72,6 +75,14 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
|
|||||||
await self._callbacks.on_client_connected(self._websocket)
|
await self._callbacks.on_client_connected(self._websocket)
|
||||||
self._receive_task = self.create_task(self._receive_messages())
|
self._receive_task = self.create_task(self._receive_messages())
|
||||||
|
|
||||||
|
async def stop(self, frame: EndFrame):
|
||||||
|
await super().stop(frame)
|
||||||
|
await cancel_task(self._receive_task)
|
||||||
|
|
||||||
|
async def cancel(self, frame: CancelFrame):
|
||||||
|
await super().cancel(frame)
|
||||||
|
await cancel_task(self._receive_task)
|
||||||
|
|
||||||
def _iter_data(self) -> typing.AsyncIterator[bytes | str]:
|
def _iter_data(self) -> typing.AsyncIterator[bytes | str]:
|
||||||
if self._params.serializer.type == FrameSerializerType.BINARY:
|
if self._params.serializer.type == FrameSerializerType.BINARY:
|
||||||
return self._websocket.iter_bytes()
|
return self._websocket.iter_bytes()
|
||||||
|
|||||||
@@ -76,12 +76,11 @@ class WebsocketServerInputTransport(BaseInputTransport):
|
|||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
await super().stop(frame)
|
await super().stop(frame)
|
||||||
self._stop_server_event.set()
|
self._stop_server_event.set()
|
||||||
await self._server_task
|
await self.wait_for_task(self._server_task)
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
await super().cancel(frame)
|
await super().cancel(frame)
|
||||||
self._stop_server_event.set()
|
await self.cancel_task(self._server_task)
|
||||||
await self._server_task
|
|
||||||
|
|
||||||
async def _server_task_handler(self):
|
async def _server_task_handler(self):
|
||||||
logger.info(f"Starting websocket server on {self._host}:{self._port}")
|
logger.info(f"Starting websocket server on {self._host}:{self._port}")
|
||||||
|
|||||||
Reference in New Issue
Block a user