feat: Add support for pipecat video stream; fix the bug of duplicate participants when connecting; fix the bug of RTVI messages sent via livekit messages;

This commit is contained in:
Alex Zhou
2025-09-06 10:41:33 +08:00
parent b9c96fd623
commit 13b73d4406

View File

@@ -13,6 +13,7 @@ event handling for conversational AI applications.
import asyncio import asyncio
from dataclasses import dataclass from dataclasses import dataclass
import json
from typing import Any, Awaitable, Callable, List, Optional from typing import Any, Awaitable, Callable, List, Optional
from loguru import logger from loguru import logger
@@ -22,6 +23,7 @@ from pipecat.audio.utils import create_stream_resampler
from pipecat.audio.vad.vad_analyzer import VADAnalyzer from pipecat.audio.vad.vad_analyzer import VADAnalyzer
from pipecat.frames.frames import ( from pipecat.frames.frames import (
AudioRawFrame, AudioRawFrame,
ImageRawFrame,
CancelFrame, CancelFrame,
EndFrame, EndFrame,
OutputAudioRawFrame, OutputAudioRawFrame,
@@ -31,6 +33,7 @@ from pipecat.frames.frames import (
TransportMessageFrame, TransportMessageFrame,
TransportMessageUrgentFrame, TransportMessageUrgentFrame,
UserAudioRawFrame, UserAudioRawFrame,
UserImageRawFrame,
) )
from pipecat.processors.frame_processor import FrameDirection, FrameProcessorSetup from pipecat.processors.frame_processor import FrameDirection, FrameProcessorSetup
from pipecat.transports.base_input import BaseInputTransport from pipecat.transports.base_input import BaseInputTransport
@@ -40,10 +43,13 @@ from pipecat.utils.asyncio.task_manager import BaseTaskManager
try: try:
from livekit import rtc from livekit import rtc
from livekit.rtc._proto import video_frame_pb2 as proto_video_frame
from tenacity import retry, stop_after_attempt, wait_exponential from tenacity import retry, stop_after_attempt, wait_exponential
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
logger.error(f"Exception: {e}") logger.error(f"Exception: {e}")
logger.error("In order to use LiveKit, you need to `pip install pipecat-ai[livekit]`.") logger.error(
"In order to use LiveKit, you need to `pip install pipecat-ai[livekit]`."
)
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
# DTMF mapping according to RFC 4733 # DTMF mapping according to RFC 4733
@@ -114,6 +120,8 @@ class LiveKitCallbacks(BaseModel):
on_participant_disconnected: Callable[[str], Awaitable[None]] on_participant_disconnected: Callable[[str], Awaitable[None]]
on_audio_track_subscribed: Callable[[str], Awaitable[None]] on_audio_track_subscribed: Callable[[str], Awaitable[None]]
on_audio_track_unsubscribed: Callable[[str], Awaitable[None]] on_audio_track_unsubscribed: Callable[[str], Awaitable[None]]
on_video_track_subscribed: Callable[[str], Awaitable[None]]
on_video_track_unsubscribed: Callable[[str], Awaitable[None]]
on_data_received: Callable[[bytes, str], Awaitable[None]] on_data_received: Callable[[bytes, str], Awaitable[None]]
on_first_participant_joined: Callable[[str], Awaitable[None]] on_first_participant_joined: Callable[[str], Awaitable[None]]
@@ -158,8 +166,11 @@ class LiveKitTransportClient:
self._audio_track: Optional[rtc.LocalAudioTrack] = None self._audio_track: Optional[rtc.LocalAudioTrack] = None
self._audio_tracks = {} self._audio_tracks = {}
self._audio_queue = asyncio.Queue() self._audio_queue = asyncio.Queue()
self._video_tracks = {}
self._video_queue = asyncio.Queue()
self._other_participant_has_joined = False self._other_participant_has_joined = False
self._task_manager: Optional[BaseTaskManager] = None self._task_manager: Optional[BaseTaskManager] = None
self._async_lock = asyncio.Lock()
@property @property
def participant_id(self) -> str: def participant_id(self) -> str:
@@ -198,7 +209,9 @@ class LiveKitTransportClient:
# Set up room event handlers # Set up room event handlers
self.room.on("participant_connected")(self._on_participant_connected_wrapper) self.room.on("participant_connected")(self._on_participant_connected_wrapper)
self.room.on("participant_disconnected")(self._on_participant_disconnected_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_subscribed")(self._on_track_subscribed_wrapper)
self.room.on("track_unsubscribed")(self._on_track_unsubscribed_wrapper) self.room.on("track_unsubscribed")(self._on_track_unsubscribed_wrapper)
self.room.on("data_received")(self._on_data_received_wrapper) self.room.on("data_received")(self._on_data_received_wrapper)
@@ -215,66 +228,74 @@ class LiveKitTransportClient:
Args: Args:
frame: The start frame containing initialization parameters. frame: The start frame containing initialization parameters.
""" """
self._out_sample_rate = self._params.audio_out_sample_rate or frame.audio_out_sample_rate self._out_sample_rate = (
self._params.audio_out_sample_rate or frame.audio_out_sample_rate
)
@retry(stop=stop_after_attempt(3), wait=wait_exponential(multiplier=1, min=4, max=10)) @retry(
stop=stop_after_attempt(3), wait=wait_exponential(multiplier=1, min=4, max=10)
)
async def connect(self): async def connect(self):
"""Connect to the LiveKit room with retry logic.""" """Connect to the LiveKit room with retry logic."""
if self._connected: async with self._async_lock:
# Increment disconnect counter if already connected. if self._connected:
self._disconnect_counter += 1 # Increment disconnect counter if already connected.
return self._disconnect_counter += 1
return
logger.info(f"Connecting to {self._room_name}") logger.info(f"Connecting to {self._room_name}")
try: try:
await self.room.connect( await self.room.connect(
self._url, self._url,
self._token, self._token,
options=rtc.RoomOptions(auto_subscribe=True), options=rtc.RoomOptions(auto_subscribe=True),
) )
self._connected = True self._connected = True
# Increment disconnect counter if we successfully connected. # Increment disconnect counter if we successfully connected.
self._disconnect_counter += 1 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}") logger.info(f"Connected to {self._room_name}")
# Set up audio source and track # Set up audio source and track
self._audio_source = rtc.AudioSource( self._audio_source = rtc.AudioSource(
self._out_sample_rate, self._params.audio_out_channels self._out_sample_rate, self._params.audio_out_channels
) )
self._audio_track = rtc.LocalAudioTrack.create_audio_track( self._audio_track = rtc.LocalAudioTrack.create_audio_track(
"pipecat-audio", self._audio_source "pipecat-audio", self._audio_source
) )
options = rtc.TrackPublishOptions() options = rtc.TrackPublishOptions()
options.source = rtc.TrackSource.SOURCE_MICROPHONE 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() await self._callbacks.on_connected()
# Check if there are already participants in the room # Check if there are already participants in the room
participants = self.get_participants() participants = self.get_participants()
if participants and not self._other_participant_has_joined: if participants and not self._other_participant_has_joined:
self._other_participant_has_joined = True self._other_participant_has_joined = True
await self._callbacks.on_first_participant_joined(participants[0]) await self._callbacks.on_first_participant_joined(participants[0])
except Exception as e: except Exception as e:
logger.error(f"Error connecting to {self._room_name}: {e}") logger.error(f"Error connecting to {self._room_name}: {e}")
raise raise
async def disconnect(self): async def disconnect(self):
"""Disconnect from the LiveKit room.""" """Disconnect from the LiveKit room."""
# Decrement leave counter when leaving. async with self._async_lock:
self._disconnect_counter -= 1 # Decrement leave counter when leaving.
self._disconnect_counter -= 1
if not self._connected or self._disconnect_counter > 0: if not self._connected or self._disconnect_counter > 0:
return return
logger.info(f"Disconnecting from {self._room_name}") logger.info(f"Disconnecting from {self._room_name}")
await self.room.disconnect() await self.room.disconnect()
self._connected = False self._connected = False
logger.info(f"Disconnected from {self._room_name}") logger.info(f"Disconnected from {self._room_name}")
await self._callbacks.on_disconnected() await self._callbacks.on_disconnected()
async def send_data(self, data: bytes, participant_id: Optional[str] = None): async def send_data(self, data: bytes, participant_id: Optional[str] = None):
"""Send data to participants in the room. """Send data to participants in the room.
@@ -437,7 +458,9 @@ class LiveKitTransportClient:
def _on_connected_wrapper(self): def _on_connected_wrapper(self):
"""Wrapper for connected events.""" """Wrapper for connected events."""
self._task_manager.create_task(self._async_on_connected(), f"{self}::_async_on_connected") self._task_manager.create_task(
self._async_on_connected(), f"{self}::_async_on_connected"
)
def _on_disconnected_wrapper(self): def _on_disconnected_wrapper(self):
"""Wrapper for disconnected events.""" """Wrapper for disconnected events."""
@@ -454,7 +477,9 @@ class LiveKitTransportClient:
self._other_participant_has_joined = True self._other_participant_has_joined = True
await self._callbacks.on_first_participant_joined(participant.sid) await self._callbacks.on_first_participant_joined(participant.sid)
async def _async_on_participant_disconnected(self, participant: rtc.RemoteParticipant): async def _async_on_participant_disconnected(
self, participant: rtc.RemoteParticipant
):
"""Handle participant disconnected events.""" """Handle participant disconnected events."""
logger.info(f"Participant disconnected: {participant.identity}") logger.info(f"Participant disconnected: {participant.identity}")
await self._callbacks.on_participant_disconnected(participant.sid) await self._callbacks.on_participant_disconnected(participant.sid)
@@ -469,7 +494,9 @@ class LiveKitTransportClient:
): ):
"""Handle track subscribed events.""" """Handle track subscribed events."""
if track.kind == rtc.TrackKind.KIND_AUDIO: if track.kind == rtc.TrackKind.KIND_AUDIO:
logger.info(f"Audio track subscribed: {track.sid} from participant {participant.sid}") logger.info(
f"Audio track subscribed: {track.sid} from participant {participant.sid}"
)
self._audio_tracks[participant.sid] = track self._audio_tracks[participant.sid] = track
audio_stream = rtc.AudioStream(track) audio_stream = rtc.AudioStream(track)
self._task_manager.create_task( self._task_manager.create_task(
@@ -477,6 +504,17 @@ class LiveKitTransportClient:
f"{self}::_process_audio_stream", f"{self}::_process_audio_stream",
) )
await self._callbacks.on_audio_track_subscribed(participant.sid) await self._callbacks.on_audio_track_subscribed(participant.sid)
elif track.kind == rtc.TrackKind.KIND_VIDEO:
logger.info(
f"Video track subscribed: {track.sid} from participant {participant.sid}"
)
self._video_tracks[participant.sid] = track
video_stream = rtc.VideoStream(track)
self._task_manager.create_task(
self._process_video_stream(video_stream, participant.sid),
f"{self}::_process_video_stream",
)
await self._callbacks.on_video_track_subscribed(participant.sid)
async def _async_on_track_unsubscribed( async def _async_on_track_unsubscribed(
self, self,
@@ -485,9 +523,13 @@ class LiveKitTransportClient:
participant: rtc.RemoteParticipant, participant: rtc.RemoteParticipant,
): ):
"""Handle track unsubscribed events.""" """Handle track unsubscribed events."""
logger.info(f"Track unsubscribed: {publication.sid} from {participant.identity}") logger.info(
f"Track unsubscribed: {publication.sid} from {participant.identity}"
)
if track.kind == rtc.TrackKind.KIND_AUDIO: if track.kind == rtc.TrackKind.KIND_AUDIO:
await self._callbacks.on_audio_track_unsubscribed(participant.sid) await self._callbacks.on_audio_track_unsubscribed(participant.sid)
elif track.kind == rtc.TrackKind.KIND_VIDEO:
await self._callbacks.on_video_track_unsubscribed(participant.sid)
async def _async_on_data_received(self, data: rtc.DataPacket): async def _async_on_data_received(self, data: rtc.DataPacket):
"""Handle data received events.""" """Handle data received events."""
@@ -503,7 +545,9 @@ class LiveKitTransportClient:
logger.info(f"Disconnected from {self._room_name}. Reason: {reason}") logger.info(f"Disconnected from {self._room_name}. Reason: {reason}")
await self._callbacks.on_disconnected() await self._callbacks.on_disconnected()
async def _process_audio_stream(self, audio_stream: rtc.AudioStream, participant_id: str): async def _process_audio_stream(
self, audio_stream: rtc.AudioStream, participant_id: str
):
"""Process incoming audio stream from a participant.""" """Process incoming audio stream from a participant."""
logger.info(f"Started processing audio stream for participant {participant_id}") logger.info(f"Started processing audio stream for participant {participant_id}")
async for event in audio_stream: async for event in audio_stream:
@@ -518,6 +562,23 @@ class LiveKitTransportClient:
frame, participant_id = await self._audio_queue.get() frame, participant_id = await self._audio_queue.get()
yield frame, participant_id yield frame, participant_id
async def _process_video_stream(
self, video_stream: rtc.VideoStream, participant_id: str
):
"""Process incoming video stream from a participant."""
logger.info(f"Started processing video stream for participant {participant_id}")
async for event in video_stream:
if isinstance(event, rtc.VideoFrameEvent):
await self._video_queue.put((event, participant_id))
else:
logger.warning(f"Received unexpected event type: {type(event)}")
async def get_next_video_frame(self):
"""Get the next video frame from the queue."""
while True:
frame, participant_id = await self._video_queue.get()
yield frame, participant_id
def __str__(self): def __str__(self):
"""String representation of the LiveKit transport client.""" """String representation of the LiveKit transport client."""
return f"{self._transport_name}::LiveKitTransportClient" return f"{self._transport_name}::LiveKitTransportClient"
@@ -550,6 +611,7 @@ class LiveKitInputTransport(BaseInputTransport):
self._client = client self._client = client
self._audio_in_task = None self._audio_in_task = None
self._video_in_task = None
self._vad_analyzer: Optional[VADAnalyzer] = params.vad_analyzer self._vad_analyzer: Optional[VADAnalyzer] = params.vad_analyzer
self._resampler = create_stream_resampler() self._resampler = create_stream_resampler()
@@ -582,6 +644,8 @@ class LiveKitInputTransport(BaseInputTransport):
await self._client.connect() await self._client.connect()
if not self._audio_in_task and self._params.audio_in_enabled: if not self._audio_in_task and self._params.audio_in_enabled:
self._audio_in_task = self.create_task(self._audio_in_task_handler()) self._audio_in_task = self.create_task(self._audio_in_task_handler())
if not self._video_in_task and self._params.video_in_enabled:
self._video_in_task = self.create_task(self._video_in_task_handler())
await self.set_transport_ready(frame) await self.set_transport_ready(frame)
logger.info("LiveKitInputTransport started") logger.info("LiveKitInputTransport started")
@@ -595,6 +659,8 @@ class LiveKitInputTransport(BaseInputTransport):
await self._client.disconnect() await self._client.disconnect()
if self._audio_in_task: if self._audio_in_task:
await self.cancel_task(self._audio_in_task) await self.cancel_task(self._audio_in_task)
if self._video_in_task:
await self.cancel_task(self._video_in_task)
logger.info("LiveKitInputTransport stopped") logger.info("LiveKitInputTransport stopped")
async def cancel(self, frame: CancelFrame): async def cancel(self, frame: CancelFrame):
@@ -607,6 +673,8 @@ class LiveKitInputTransport(BaseInputTransport):
await self._client.disconnect() await self._client.disconnect()
if self._audio_in_task and self._params.audio_in_enabled: if self._audio_in_task and self._params.audio_in_enabled:
await self.cancel_task(self._audio_in_task) await self.cancel_task(self._audio_in_task)
if self._video_in_task and self._params.video_in_enabled:
await self.cancel_task(self._video_in_task)
async def setup(self, setup: FrameProcessorSetup): async def setup(self, setup: FrameProcessorSetup):
"""Setup the input transport with shared client setup. """Setup the input transport with shared client setup.
@@ -629,7 +697,9 @@ class LiveKitInputTransport(BaseInputTransport):
message: The message data to send. message: The message data to send.
sender: ID of the message sender. sender: ID of the message sender.
""" """
frame = LiveKitTransportMessageUrgentFrame(message=message, participant_id=sender) frame = LiveKitTransportMessageUrgentFrame(
message=message, participant_id=sender
)
await self.push_frame(frame) await self.push_frame(frame)
async def _audio_in_task_handler(self): async def _audio_in_task_handler(self):
@@ -655,6 +725,29 @@ class LiveKitInputTransport(BaseInputTransport):
) )
await self.push_audio_frame(input_audio_frame) await self.push_audio_frame(input_audio_frame)
async def _video_in_task_handler(self):
"""Handle incoming video frames from participants."""
logger.info("Video input task started")
video_iterator = self._client.get_next_video_frame()
async for video_data in video_iterator:
if video_data:
video_frame_event, participant_id = video_data
pipecat_video_frame = await self._convert_livekit_video_to_pipecat(
video_frame_event=video_frame_event
)
# Skip frames with no video data
if len(pipecat_video_frame.image) == 0:
continue
input_video_frame = UserImageRawFrame(
user_id=participant_id,
image=pipecat_video_frame.image,
size=pipecat_video_frame.size,
format=pipecat_video_frame.format,
)
await self.push_video_frame(input_video_frame)
async def _convert_livekit_audio_to_pipecat( async def _convert_livekit_audio_to_pipecat(
self, audio_frame_event: rtc.AudioFrameEvent self, audio_frame_event: rtc.AudioFrameEvent
) -> AudioRawFrame: ) -> AudioRawFrame:
@@ -671,6 +764,21 @@ class LiveKitInputTransport(BaseInputTransport):
num_channels=audio_frame.num_channels, num_channels=audio_frame.num_channels,
) )
async def _convert_livekit_video_to_pipecat(
self,
video_frame_event: rtc.VideoFrameEvent,
) -> ImageRawFrame:
"""Convert LiveKit video frame to Pipecat video frame."""
rgb_frame = video_frame_event.frame.convert(
proto_video_frame.VideoBufferType.RGB24
)
image_frame = ImageRawFrame(
image=rgb_frame.data,
size=(rgb_frame.width, rgb_frame.height),
format="RGB",
)
return image_frame
class LiveKitOutputTransport(BaseOutputTransport): class LiveKitOutputTransport(BaseOutputTransport):
"""Handles outgoing media streams and events to LiveKit rooms. """Handles outgoing media streams and events to LiveKit rooms.
@@ -752,16 +860,24 @@ class LiveKitOutputTransport(BaseOutputTransport):
await super().cleanup() await super().cleanup()
await self._transport.cleanup() await self._transport.cleanup()
async def send_message(self, frame: TransportMessageFrame | TransportMessageUrgentFrame): async def send_message(
self, frame: TransportMessageFrame | TransportMessageUrgentFrame
):
"""Send a transport message to participants. """Send a transport message to participants.
Args: Args:
frame: The transport message frame to send. frame: The transport message frame to send.
""" """
if isinstance(frame, (LiveKitTransportMessageFrame, LiveKitTransportMessageUrgentFrame)): message = frame.message
await self._client.send_data(frame.message.encode(), frame.participant_id) if isinstance(message, dict):
# fix message encoding for dict-like messages, e.g. RTVI messages.
message = json.dumps(message, ensure_ascii=False)
if isinstance(
frame, (LiveKitTransportMessageFrame, LiveKitTransportMessageUrgentFrame)
):
await self._client.send_data(message.encode(), frame.participant_id)
else: else:
await self._client.send_data(frame.message.encode()) await self._client.send_data(message.encode())
async def write_audio_frame(self, frame: OutputAudioRawFrame): async def write_audio_frame(self, frame: OutputAudioRawFrame):
"""Write an audio frame to the LiveKit room. """Write an audio frame to the LiveKit room.
@@ -838,6 +954,8 @@ class LiveKitTransport(BaseTransport):
on_participant_disconnected=self._on_participant_disconnected, on_participant_disconnected=self._on_participant_disconnected,
on_audio_track_subscribed=self._on_audio_track_subscribed, on_audio_track_subscribed=self._on_audio_track_subscribed,
on_audio_track_unsubscribed=self._on_audio_track_unsubscribed, on_audio_track_unsubscribed=self._on_audio_track_unsubscribed,
on_video_track_subscribed=self._on_video_track_subscribed,
on_video_track_unsubscribed=self._on_video_track_unsubscribed,
on_data_received=self._on_data_received, on_data_received=self._on_data_received,
on_first_participant_joined=self._on_first_participant_joined, on_first_participant_joined=self._on_first_participant_joined,
) )
@@ -855,6 +973,8 @@ class LiveKitTransport(BaseTransport):
self._register_event_handler("on_participant_disconnected") self._register_event_handler("on_participant_disconnected")
self._register_event_handler("on_audio_track_subscribed") self._register_event_handler("on_audio_track_subscribed")
self._register_event_handler("on_audio_track_unsubscribed") self._register_event_handler("on_audio_track_unsubscribed")
self._register_event_handler("on_video_track_subscribed")
self._register_event_handler("on_video_track_unsubscribed")
self._register_event_handler("on_data_received") self._register_event_handler("on_data_received")
self._register_event_handler("on_first_participant_joined") self._register_event_handler("on_first_participant_joined")
self._register_event_handler("on_participant_left") self._register_event_handler("on_participant_left")
@@ -960,7 +1080,9 @@ class LiveKitTransport(BaseTransport):
async def _on_participant_disconnected(self, participant_id: str): async def _on_participant_disconnected(self, participant_id: str):
"""Handle participant disconnected events.""" """Handle participant disconnected events."""
await self._call_event_handler("on_participant_disconnected", participant_id) await self._call_event_handler("on_participant_disconnected", participant_id)
await self._call_event_handler("on_participant_left", participant_id, "disconnected") await self._call_event_handler(
"on_participant_left", participant_id, "disconnected"
)
async def _on_audio_track_subscribed(self, participant_id: str): async def _on_audio_track_subscribed(self, participant_id: str):
"""Handle audio track subscribed events.""" """Handle audio track subscribed events."""
@@ -976,6 +1098,20 @@ class LiveKitTransport(BaseTransport):
"""Handle audio track unsubscribed events.""" """Handle audio track unsubscribed events."""
await self._call_event_handler("on_audio_track_unsubscribed", participant_id) await self._call_event_handler("on_audio_track_unsubscribed", participant_id)
async def _on_video_track_subscribed(self, participant_id: str):
"""Handle video track subscribed events."""
await self._call_event_handler("on_video_track_subscribed", participant_id)
participant = self._client.room.remote_participants.get(participant_id)
if participant:
for publication in participant.video_tracks.values():
self._client._on_track_subscribed_wrapper(
publication.track, publication, participant
)
async def _on_video_track_unsubscribed(self, participant_id: str):
"""Handle video track unsubscribed events."""
await self._call_event_handler("on_video_track_unsubscribed", participant_id)
async def _on_data_received(self, data: bytes, participant_id: str): async def _on_data_received(self, data: bytes, participant_id: str):
"""Handle data received events.""" """Handle data received events."""
if self._input: if self._input:
@@ -990,10 +1126,14 @@ class LiveKitTransport(BaseTransport):
participant_id: Optional specific participant to send to. participant_id: Optional specific participant to send to.
""" """
if self._output: if self._output:
frame = LiveKitTransportMessageFrame(message=message, participant_id=participant_id) frame = LiveKitTransportMessageFrame(
message=message, participant_id=participant_id
)
await self._output.send_message(frame) await self._output.send_message(frame)
async def send_message_urgent(self, message: str, participant_id: Optional[str] = None): async def send_message_urgent(
self, message: str, participant_id: Optional[str] = None
):
"""Send an urgent message to participants in the room. """Send an urgent message to participants in the room.
Args: Args: