improve task creation and cancellation
If a FrameProcessor needs to create a task it should use FrameProcessor.create_task() and FrameProcessor.cancel_task(). This gives Pipecat more control over all the tasks that are created in Pipecat. Both functions internally use the utils module: utils.create_task() and utils.cancel_task() which should also be used outside of FrameProcessors. That is, unless strictly necessary, we should avoid using asyncio.create_task().
This commit is contained in:
@@ -28,6 +28,7 @@ from pipecat.processors.frame_processor import FrameDirection
|
||||
from pipecat.transports.base_input import BaseInputTransport
|
||||
from pipecat.transports.base_output import BaseOutputTransport
|
||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||
from pipecat.utils.utils import create_task
|
||||
|
||||
try:
|
||||
from livekit import rtc
|
||||
@@ -215,10 +216,18 @@ class LiveKitTransportClient:
|
||||
|
||||
# Wrapper methods for event handlers
|
||||
def _on_participant_connected_wrapper(self, participant: rtc.RemoteParticipant):
|
||||
asyncio.create_task(self._async_on_participant_connected(participant))
|
||||
create_task(
|
||||
self._loop,
|
||||
self._async_on_participant_connected(participant),
|
||||
"LiveKitTransportClient::_async_on_participant_connected",
|
||||
)
|
||||
|
||||
def _on_participant_disconnected_wrapper(self, participant: rtc.RemoteParticipant):
|
||||
asyncio.create_task(self._async_on_participant_disconnected(participant))
|
||||
create_task(
|
||||
self._loop,
|
||||
self._async_on_participant_disconnected(participant),
|
||||
"LiveKitTransportClient::_async_on_participant_disconnected",
|
||||
)
|
||||
|
||||
def _on_track_subscribed_wrapper(
|
||||
self,
|
||||
@@ -226,7 +235,11 @@ class LiveKitTransportClient:
|
||||
publication: rtc.RemoteTrackPublication,
|
||||
participant: rtc.RemoteParticipant,
|
||||
):
|
||||
asyncio.create_task(self._async_on_track_subscribed(track, publication, participant))
|
||||
create_task(
|
||||
self._loop,
|
||||
self._async_on_track_subscribed(track, publication, participant),
|
||||
"LiveKitTransportClient::_async_on_track_subscribed",
|
||||
)
|
||||
|
||||
def _on_track_unsubscribed_wrapper(
|
||||
self,
|
||||
@@ -234,16 +247,30 @@ class LiveKitTransportClient:
|
||||
publication: rtc.RemoteTrackPublication,
|
||||
participant: rtc.RemoteParticipant,
|
||||
):
|
||||
asyncio.create_task(self._async_on_track_unsubscribed(track, publication, participant))
|
||||
create_task(
|
||||
self._loop,
|
||||
self._async_on_track_unsubscribed(track, publication, participant),
|
||||
"LiveKitTransportClient::_async_on_track_unsubscribed",
|
||||
)
|
||||
|
||||
def _on_data_received_wrapper(self, data: rtc.DataPacket):
|
||||
asyncio.create_task(self._async_on_data_received(data))
|
||||
create_task(
|
||||
self._loop,
|
||||
self._async_on_data_received(data),
|
||||
"LiveKitTransportClient::_async_on_data_received",
|
||||
)
|
||||
|
||||
def _on_connected_wrapper(self):
|
||||
asyncio.create_task(self._async_on_connected())
|
||||
create_task(
|
||||
self._loop, self._async_on_connected(), "LiveKitTransportClient::_async_on_connected"
|
||||
)
|
||||
|
||||
def _on_disconnected_wrapper(self):
|
||||
asyncio.create_task(self._async_on_disconnected())
|
||||
create_task(
|
||||
self._loop,
|
||||
self._async_on_disconnected(),
|
||||
"LiveKitTransportClient::_async_on_disconnected",
|
||||
)
|
||||
|
||||
# Async methods for event handling
|
||||
async def _async_on_participant_connected(self, participant: rtc.RemoteParticipant):
|
||||
@@ -269,7 +296,11 @@ class LiveKitTransportClient:
|
||||
logger.info(f"Audio track subscribed: {track.sid} from participant {participant.sid}")
|
||||
self._audio_tracks[participant.sid] = track
|
||||
audio_stream = rtc.AudioStream(track)
|
||||
asyncio.create_task(self._process_audio_stream(audio_stream, participant.sid))
|
||||
create_task(
|
||||
self._loop,
|
||||
self._process_audio_stream(audio_stream, participant.sid),
|
||||
"LiveKitTransportClient::_process_audio_stream",
|
||||
)
|
||||
|
||||
async def _async_on_track_unsubscribed(
|
||||
self,
|
||||
@@ -319,23 +350,21 @@ class LiveKitInputTransport(BaseInputTransport):
|
||||
await super().start(frame)
|
||||
await self._client.connect()
|
||||
if self._params.audio_in_enabled or self._params.vad_enabled:
|
||||
self._audio_in_task = asyncio.create_task(self._audio_in_task_handler())
|
||||
self._audio_in_task = self.create_task(self._audio_in_task_handler())
|
||||
logger.info("LiveKitInputTransport started")
|
||||
|
||||
async def stop(self, frame: EndFrame):
|
||||
await super().stop(frame)
|
||||
await self._client.disconnect()
|
||||
if self._audio_in_task:
|
||||
self._audio_in_task.cancel()
|
||||
await self._audio_in_task
|
||||
await self.cancel_task(self._audio_in_task)
|
||||
logger.info("LiveKitInputTransport stopped")
|
||||
|
||||
async def cancel(self, frame: CancelFrame):
|
||||
await super().cancel(frame)
|
||||
await self._client.disconnect()
|
||||
if self._audio_in_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
||||
self._audio_in_task.cancel()
|
||||
await self._audio_in_task
|
||||
await self.cancel_task(self._audio_in_task)
|
||||
|
||||
def vad_analyzer(self) -> VADAnalyzer | None:
|
||||
return self._vad_analyzer
|
||||
@@ -347,22 +376,16 @@ class LiveKitInputTransport(BaseInputTransport):
|
||||
async def _audio_in_task_handler(self):
|
||||
logger.info("Audio input task started")
|
||||
while True:
|
||||
try:
|
||||
audio_data = await self._client.get_next_audio_frame()
|
||||
if audio_data:
|
||||
audio_frame_event, participant_id = audio_data
|
||||
pipecat_audio_frame = self._convert_livekit_audio_to_pipecat(audio_frame_event)
|
||||
input_audio_frame = InputAudioRawFrame(
|
||||
audio=pipecat_audio_frame.audio,
|
||||
sample_rate=pipecat_audio_frame.sample_rate,
|
||||
num_channels=pipecat_audio_frame.num_channels,
|
||||
)
|
||||
await self.push_audio_frame(input_audio_frame)
|
||||
except asyncio.CancelledError:
|
||||
logger.info("Audio input task cancelled")
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in audio input task: {e}")
|
||||
audio_data = await self._client.get_next_audio_frame()
|
||||
if audio_data:
|
||||
audio_frame_event, participant_id = audio_data
|
||||
pipecat_audio_frame = self._convert_livekit_audio_to_pipecat(audio_frame_event)
|
||||
input_audio_frame = InputAudioRawFrame(
|
||||
audio=pipecat_audio_frame.audio,
|
||||
sample_rate=pipecat_audio_frame.sample_rate,
|
||||
num_channels=pipecat_audio_frame.num_channels,
|
||||
)
|
||||
await self.push_audio_frame(input_audio_frame)
|
||||
|
||||
def _convert_livekit_audio_to_pipecat(
|
||||
self, audio_frame_event: rtc.AudioFrameEvent
|
||||
|
||||
Reference in New Issue
Block a user