introduce TaskManager and PipelineRunner event loop
This commit is contained in:
@@ -6,7 +6,7 @@
|
||||
|
||||
import asyncio
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Awaitable, Callable, List
|
||||
from typing import Any, Awaitable, Callable, List, Optional
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import BaseModel
|
||||
@@ -17,7 +17,6 @@ from pipecat.frames.frames import (
|
||||
AudioRawFrame,
|
||||
CancelFrame,
|
||||
EndFrame,
|
||||
Frame,
|
||||
InputAudioRawFrame,
|
||||
OutputAudioRawFrame,
|
||||
StartFrame,
|
||||
@@ -28,7 +27,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.asyncio import create_task
|
||||
from pipecat.utils.asyncio import TaskManager
|
||||
|
||||
try:
|
||||
from livekit import rtc
|
||||
@@ -72,7 +71,6 @@ class LiveKitTransportClient:
|
||||
room_name: str,
|
||||
params: LiveKitParams,
|
||||
callbacks: LiveKitCallbacks,
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
transport_name: str,
|
||||
):
|
||||
self._url = url
|
||||
@@ -80,9 +78,8 @@ class LiveKitTransportClient:
|
||||
self._room_name = room_name
|
||||
self._params = params
|
||||
self._callbacks = callbacks
|
||||
self._loop = loop
|
||||
self._transport_name = transport_name
|
||||
self._room = rtc.Room(loop=loop)
|
||||
self._room: rtc.Room | None = None
|
||||
self._participant_id: str = ""
|
||||
self._connected = False
|
||||
self._disconnect_counter = 0
|
||||
@@ -91,20 +88,32 @@ class LiveKitTransportClient:
|
||||
self._audio_tracks = {}
|
||||
self._audio_queue = asyncio.Queue()
|
||||
self._other_participant_has_joined = False
|
||||
|
||||
# Set up room event handlers
|
||||
self._room.on("participant_connected")(self._on_participant_connected_wrapper)
|
||||
self._room.on("participant_disconnected")(self._on_participant_disconnected_wrapper)
|
||||
self._room.on("track_subscribed")(self._on_track_subscribed_wrapper)
|
||||
self._room.on("track_unsubscribed")(self._on_track_unsubscribed_wrapper)
|
||||
self._room.on("data_received")(self._on_data_received_wrapper)
|
||||
self._room.on("connected")(self._on_connected_wrapper)
|
||||
self._room.on("disconnected")(self._on_disconnected_wrapper)
|
||||
self._task_manager: Optional[TaskManager] = None
|
||||
|
||||
@property
|
||||
def participant_id(self) -> str:
|
||||
return self._participant_id
|
||||
|
||||
@property
|
||||
def room(self) -> rtc.Room:
|
||||
if not self._room:
|
||||
raise Exception(f"{self}: missing room object (pipeline not started?)")
|
||||
return self._room
|
||||
|
||||
async def setup(self, frame: StartFrame):
|
||||
if not self._task_manager:
|
||||
self._task_manager = frame.task_manager
|
||||
self._room = rtc.Room(loop=self._task_manager.get_event_loop())
|
||||
|
||||
# Set up room event handlers
|
||||
self.room.on("participant_connected")(self._on_participant_connected_wrapper)
|
||||
self.room.on("participant_disconnected")(self._on_participant_disconnected_wrapper)
|
||||
self.room.on("track_subscribed")(self._on_track_subscribed_wrapper)
|
||||
self.room.on("track_unsubscribed")(self._on_track_unsubscribed_wrapper)
|
||||
self.room.on("data_received")(self._on_data_received_wrapper)
|
||||
self.room.on("connected")(self._on_connected_wrapper)
|
||||
self.room.on("disconnected")(self._on_disconnected_wrapper)
|
||||
|
||||
@retry(stop=stop_after_attempt(3), wait=wait_exponential(multiplier=1, min=4, max=10))
|
||||
async def connect(self):
|
||||
if self._connected:
|
||||
@@ -115,7 +124,7 @@ class LiveKitTransportClient:
|
||||
logger.info(f"Connecting to {self._room_name}")
|
||||
|
||||
try:
|
||||
await self._room.connect(
|
||||
await self.room.connect(
|
||||
self._url,
|
||||
self._token,
|
||||
options=rtc.RoomOptions(auto_subscribe=True),
|
||||
@@ -124,7 +133,7 @@ class LiveKitTransportClient:
|
||||
# Increment disconnect counter if we successfully connected.
|
||||
self._disconnect_counter += 1
|
||||
|
||||
self._participant_id = self._room.local_participant.sid
|
||||
self._participant_id = self.room.local_participant.sid
|
||||
logger.info(f"Connected to {self._room_name}")
|
||||
|
||||
# Set up audio source and track
|
||||
@@ -136,7 +145,7 @@ class LiveKitTransportClient:
|
||||
)
|
||||
options = rtc.TrackPublishOptions()
|
||||
options.source = rtc.TrackSource.SOURCE_MICROPHONE
|
||||
await self._room.local_participant.publish_track(self._audio_track, options)
|
||||
await self.room.local_participant.publish_track(self._audio_track, options)
|
||||
|
||||
await self._callbacks.on_connected()
|
||||
|
||||
@@ -157,7 +166,7 @@ class LiveKitTransportClient:
|
||||
return
|
||||
|
||||
logger.info(f"Disconnecting from {self._room_name}")
|
||||
await self._room.disconnect()
|
||||
await self.room.disconnect()
|
||||
self._connected = False
|
||||
logger.info(f"Disconnected from {self._room_name}")
|
||||
await self._callbacks.on_disconnected()
|
||||
@@ -168,11 +177,11 @@ class LiveKitTransportClient:
|
||||
|
||||
try:
|
||||
if participant_id:
|
||||
await self._room.local_participant.publish_data(
|
||||
await self.room.local_participant.publish_data(
|
||||
data, reliable=True, destination_identities=[participant_id]
|
||||
)
|
||||
else:
|
||||
await self._room.local_participant.publish_data(data, reliable=True)
|
||||
await self.room.local_participant.publish_data(data, reliable=True)
|
||||
except Exception as e:
|
||||
logger.error(f"Error sending data: {e}")
|
||||
|
||||
@@ -186,10 +195,10 @@ class LiveKitTransportClient:
|
||||
logger.error(f"Error publishing audio: {e}")
|
||||
|
||||
def get_participants(self) -> List[str]:
|
||||
return [p.sid for p in self._room.remote_participants.values()]
|
||||
return [p.sid for p in self.room.remote_participants.values()]
|
||||
|
||||
async def get_participant_metadata(self, participant_id: str) -> dict:
|
||||
participant = self._room.remote_participants.get(participant_id)
|
||||
participant = self.room.remote_participants.get(participant_id)
|
||||
if participant:
|
||||
return {
|
||||
"id": participant.sid,
|
||||
@@ -200,17 +209,17 @@ class LiveKitTransportClient:
|
||||
return {}
|
||||
|
||||
async def set_participant_metadata(self, metadata: str):
|
||||
await self._room.local_participant.set_metadata(metadata)
|
||||
await self.room.local_participant.set_metadata(metadata)
|
||||
|
||||
async def mute_participant(self, participant_id: str):
|
||||
participant = self._room.remote_participants.get(participant_id)
|
||||
participant = self.room.remote_participants.get(participant_id)
|
||||
if participant:
|
||||
for track in participant.tracks.values():
|
||||
if track.kind == "audio":
|
||||
await track.set_enabled(False)
|
||||
|
||||
async def unmute_participant(self, participant_id: str):
|
||||
participant = self._room.remote_participants.get(participant_id)
|
||||
participant = self.room.remote_participants.get(participant_id)
|
||||
if participant:
|
||||
for track in participant.tracks.values():
|
||||
if track.kind == "audio":
|
||||
@@ -218,17 +227,15 @@ class LiveKitTransportClient:
|
||||
|
||||
# Wrapper methods for event handlers
|
||||
def _on_participant_connected_wrapper(self, participant: rtc.RemoteParticipant):
|
||||
create_task(
|
||||
self._loop,
|
||||
self._task_manager.create_task(
|
||||
self._async_on_participant_connected(participant),
|
||||
f"{self._transport_name}::LiveKitTransportClient::_async_on_participant_connected",
|
||||
f"{self}::_async_on_participant_connected",
|
||||
)
|
||||
|
||||
def _on_participant_disconnected_wrapper(self, participant: rtc.RemoteParticipant):
|
||||
create_task(
|
||||
self._loop,
|
||||
self._task_manager.create_task(
|
||||
self._async_on_participant_disconnected(participant),
|
||||
f"{self._transport_name}::LiveKitTransportClient::_async_on_participant_disconnected",
|
||||
f"{self}::_async_on_participant_disconnected",
|
||||
)
|
||||
|
||||
def _on_track_subscribed_wrapper(
|
||||
@@ -237,10 +244,9 @@ class LiveKitTransportClient:
|
||||
publication: rtc.RemoteTrackPublication,
|
||||
participant: rtc.RemoteParticipant,
|
||||
):
|
||||
create_task(
|
||||
self._loop,
|
||||
self._task_manager.create_task(
|
||||
self._async_on_track_subscribed(track, publication, participant),
|
||||
f"{self._transport_name}::LiveKitTransportClient::_async_on_track_subscribed",
|
||||
f"{self}::_async_on_track_subscribed",
|
||||
)
|
||||
|
||||
def _on_track_unsubscribed_wrapper(
|
||||
@@ -249,31 +255,23 @@ class LiveKitTransportClient:
|
||||
publication: rtc.RemoteTrackPublication,
|
||||
participant: rtc.RemoteParticipant,
|
||||
):
|
||||
create_task(
|
||||
self._loop,
|
||||
self._task_manager.create_task(
|
||||
self._async_on_track_unsubscribed(track, publication, participant),
|
||||
f"{self._transport_name}::LiveKitTransportClient::_async_on_track_unsubscribed",
|
||||
f"{self}::_async_on_track_unsubscribed",
|
||||
)
|
||||
|
||||
def _on_data_received_wrapper(self, data: rtc.DataPacket):
|
||||
create_task(
|
||||
self._loop,
|
||||
self._task_manager.create_task(
|
||||
self._async_on_data_received(data),
|
||||
f"{self._transport_name}::LiveKitTransportClient::_async_on_data_received",
|
||||
f"{self}::_async_on_data_received",
|
||||
)
|
||||
|
||||
def _on_connected_wrapper(self):
|
||||
create_task(
|
||||
self._loop,
|
||||
self._async_on_connected(),
|
||||
f"{self._transport_name}::LiveKitTransportClient::_async_on_connected",
|
||||
)
|
||||
self._task_manager.create_task(self._async_on_connected(), f"{self}::_async_on_connected")
|
||||
|
||||
def _on_disconnected_wrapper(self):
|
||||
create_task(
|
||||
self._loop,
|
||||
self._async_on_disconnected(),
|
||||
f"{self._transport_name}::LiveKitTransportClient::_async_on_disconnected",
|
||||
self._task_manager.create_task(
|
||||
self._async_on_disconnected(), f"{self}::_async_on_disconnected"
|
||||
)
|
||||
|
||||
# Async methods for event handling
|
||||
@@ -300,10 +298,9 @@ 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)
|
||||
create_task(
|
||||
self._loop,
|
||||
self._task_manager.create_task(
|
||||
self._process_audio_stream(audio_stream, participant.sid),
|
||||
f"{self._transport_name}::LiveKitTransportClient::_process_audio_stream",
|
||||
f"{self}::_process_audio_stream",
|
||||
)
|
||||
|
||||
async def _async_on_track_unsubscribed(
|
||||
@@ -342,6 +339,9 @@ class LiveKitTransportClient:
|
||||
frame, participant_id = await self._audio_queue.get()
|
||||
return frame, participant_id
|
||||
|
||||
def __str__(self):
|
||||
return f"{self._transport_name}::LiveKitTransportClient"
|
||||
|
||||
|
||||
class LiveKitInputTransport(BaseInputTransport):
|
||||
def __init__(self, client: LiveKitTransportClient, params: LiveKitParams, **kwargs):
|
||||
@@ -352,6 +352,7 @@ class LiveKitInputTransport(BaseInputTransport):
|
||||
|
||||
async def start(self, frame: StartFrame):
|
||||
await super().start(frame)
|
||||
await self._client.setup(frame)
|
||||
await self._client.connect()
|
||||
if self._params.audio_in_enabled or self._params.vad_enabled:
|
||||
self._audio_in_task = self.create_task(self._audio_in_task_handler())
|
||||
@@ -395,13 +396,10 @@ class LiveKitInputTransport(BaseInputTransport):
|
||||
self, audio_frame_event: rtc.AudioFrameEvent
|
||||
) -> AudioRawFrame:
|
||||
audio_frame = audio_frame_event.frame
|
||||
audio_data = audio_frame.data
|
||||
original_sample_rate = audio_frame.sample_rate
|
||||
|
||||
if original_sample_rate != self._params.audio_in_sample_rate:
|
||||
audio_data = resample_audio(
|
||||
audio_data, original_sample_rate, self._params.audio_in_sample_rate
|
||||
)
|
||||
audio_data = resample_audio(
|
||||
audio_frame.data.tobytes(), audio_frame.sample_rate, self._params.audio_in_sample_rate
|
||||
)
|
||||
|
||||
return AudioRawFrame(
|
||||
audio=audio_data,
|
||||
@@ -417,6 +415,7 @@ class LiveKitOutputTransport(BaseOutputTransport):
|
||||
|
||||
async def start(self, frame: StartFrame):
|
||||
await super().start(frame)
|
||||
await self._client.setup(frame)
|
||||
await self._client.connect()
|
||||
logger.info("LiveKitOutputTransport started")
|
||||
|
||||
@@ -461,9 +460,8 @@ class LiveKitTransport(BaseTransport):
|
||||
params: LiveKitParams = LiveKitParams(),
|
||||
input_name: str | None = None,
|
||||
output_name: str | None = None,
|
||||
loop: asyncio.AbstractEventLoop | None = None,
|
||||
):
|
||||
super().__init__(input_name=input_name, output_name=output_name, loop=loop)
|
||||
super().__init__(input_name=input_name, output_name=output_name)
|
||||
|
||||
callbacks = LiveKitCallbacks(
|
||||
on_connected=self._on_connected,
|
||||
@@ -478,7 +476,7 @@ class LiveKitTransport(BaseTransport):
|
||||
self._params = params
|
||||
|
||||
self._client = LiveKitTransportClient(
|
||||
url, token, room_name, self._params, callbacks, self._loop, self.name
|
||||
url, token, room_name, self._params, callbacks, self.name
|
||||
)
|
||||
self._input: LiveKitInputTransport | None = None
|
||||
self._output: LiveKitOutputTransport | None = None
|
||||
@@ -544,7 +542,7 @@ class LiveKitTransport(BaseTransport):
|
||||
|
||||
async def _on_audio_track_subscribed(self, participant_id: str):
|
||||
await self._call_event_handler("on_audio_track_subscribed", participant_id)
|
||||
participant = self._client._room.remote_participants.get(participant_id)
|
||||
participant = self._client.room.remote_participants.get(participant_id)
|
||||
if participant:
|
||||
for publication in participant.audio_tracks.values():
|
||||
self._client._on_track_subscribed_wrapper(
|
||||
|
||||
Reference in New Issue
Block a user