transport(websocket): do not require a frame serializer

This commit is contained in:
Aleix Conchillo Flaqué
2025-05-24 23:15:13 -07:00
parent 2cdfaa0a82
commit 071a9307c9
3 changed files with 36 additions and 11 deletions

View File

@@ -26,7 +26,7 @@ from pipecat.frames.frames import (
TransportMessageFrame, TransportMessageFrame,
TransportMessageUrgentFrame, TransportMessageUrgentFrame,
) )
from pipecat.processors.frame_processor import FrameDirection from pipecat.processors.frame_processor import FrameDirection, FrameProcessorSetup
from pipecat.serializers.base_serializer import FrameSerializer, FrameSerializerType from pipecat.serializers.base_serializer import FrameSerializer, FrameSerializerType
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
@@ -45,7 +45,7 @@ except ModuleNotFoundError as e:
class FastAPIWebsocketParams(TransportParams): class FastAPIWebsocketParams(TransportParams):
add_wav_header: bool = False add_wav_header: bool = False
serializer: FrameSerializer serializer: Optional[FrameSerializer] = None
session_timeout: Optional[int] = None session_timeout: Optional[int] = None
@@ -125,7 +125,8 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._client.setup(frame) await self._client.setup(frame)
await self._params.serializer.setup(frame) if self._params.serializer:
await self._params.serializer.setup(frame)
if not self._monitor_websocket_task and self._params.session_timeout: if not self._monitor_websocket_task and self._params.session_timeout:
self._monitor_websocket_task = self.create_task(self._monitor_websocket()) self._monitor_websocket_task = self.create_task(self._monitor_websocket())
await self._client.trigger_client_connected() await self._client.trigger_client_connected()
@@ -158,6 +159,9 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
async def _receive_messages(self): async def _receive_messages(self):
try: try:
async for message in self._client.receive(): async for message in self._client.receive():
if not self._params.serializer:
continue
frame = await self._params.serializer.deserialize(message) frame = await self._params.serializer.deserialize(message)
if not frame: if not frame:
@@ -203,7 +207,8 @@ class FastAPIWebsocketOutputTransport(BaseOutputTransport):
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._client.setup(frame) await self._client.setup(frame)
await self._params.serializer.setup(frame) if self._params.serializer:
await self._params.serializer.setup(frame)
self._send_interval = (self.audio_chunk_size / self.sample_rate) / 2 self._send_interval = (self.audio_chunk_size / self.sample_rate) / 2
await self.set_transport_ready(frame) await self.set_transport_ready(frame)
@@ -266,6 +271,9 @@ class FastAPIWebsocketOutputTransport(BaseOutputTransport):
await self._write_audio_sleep() await self._write_audio_sleep()
async def _write_frame(self, frame: Frame): async def _write_frame(self, frame: Frame):
if not self._params.serializer:
return
try: try:
payload = await self._params.serializer.serialize(frame) payload = await self._params.serializer.serialize(frame)
if payload: if payload:
@@ -302,7 +310,9 @@ class FastAPIWebsocketTransport(BaseTransport):
on_session_timeout=self._on_session_timeout, on_session_timeout=self._on_session_timeout,
) )
is_binary = self._params.serializer.type == FrameSerializerType.BINARY is_binary = False
if self._params.serializer:
is_binary = self._params.serializer.type == FrameSerializerType.BINARY
self._client = FastAPIWebsocketClient(websocket, is_binary, self._callbacks) self._client = FastAPIWebsocketClient(websocket, is_binary, self._callbacks)
self._input = FastAPIWebsocketInputTransport( self._input = FastAPIWebsocketInputTransport(

View File

@@ -34,7 +34,7 @@ from pipecat.utils.asyncio import BaseTaskManager
class WebsocketClientParams(TransportParams): class WebsocketClientParams(TransportParams):
add_wav_header: bool = True add_wav_header: bool = True
serializer: FrameSerializer = ProtobufFrameSerializer() serializer: Optional[FrameSerializer] = None
class WebsocketClientCallbacks(BaseModel): class WebsocketClientCallbacks(BaseModel):
@@ -133,7 +133,8 @@ class WebsocketClientInputTransport(BaseInputTransport):
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._params.serializer.setup(frame) if self._params.serializer:
await self._params.serializer.setup(frame)
await self._session.setup(frame) await self._session.setup(frame)
await self._session.connect() await self._session.connect()
await self.set_transport_ready(frame) await self.set_transport_ready(frame)
@@ -151,6 +152,8 @@ class WebsocketClientInputTransport(BaseInputTransport):
await self._transport.cleanup() await self._transport.cleanup()
async def on_message(self, websocket, message): async def on_message(self, websocket, message):
if not self._params.serializer:
return
frame = await self._params.serializer.deserialize(message) frame = await self._params.serializer.deserialize(message)
if not frame: if not frame:
return return
@@ -184,7 +187,8 @@ class WebsocketClientOutputTransport(BaseOutputTransport):
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
self._send_interval = (self.audio_chunk_size / self.sample_rate) / 2 self._send_interval = (self.audio_chunk_size / self.sample_rate) / 2
await self._params.serializer.setup(frame) if self._params.serializer:
await self._params.serializer.setup(frame)
await self._session.setup(frame) await self._session.setup(frame)
await self._session.connect() await self._session.connect()
await self.set_transport_ready(frame) await self.set_transport_ready(frame)
@@ -231,6 +235,8 @@ class WebsocketClientOutputTransport(BaseOutputTransport):
await self._write_audio_sleep() await self._write_audio_sleep()
async def _write_frame(self, frame: Frame): async def _write_frame(self, frame: Frame):
if not self._params.serializer:
return
payload = await self._params.serializer.serialize(frame) payload = await self._params.serializer.serialize(frame)
if payload: if payload:
await self._session.send(payload) await self._session.send(payload)
@@ -255,6 +261,7 @@ class WebsocketClientTransport(BaseTransport):
super().__init__() super().__init__()
self._params = params or WebsocketClientParams() self._params = params or WebsocketClientParams()
self._params.serializer = self._params.serializer or ProtobufFrameSerializer()
callbacks = WebsocketClientCallbacks( callbacks = WebsocketClientCallbacks(
on_connected=self._on_connected, on_connected=self._on_connected,

View File

@@ -40,7 +40,7 @@ except ModuleNotFoundError as e:
class WebsocketServerParams(TransportParams): class WebsocketServerParams(TransportParams):
add_wav_header: bool = False add_wav_header: bool = False
serializer: FrameSerializer serializer: Optional[FrameSerializer] = None
session_timeout: Optional[int] = None session_timeout: Optional[int] = None
@@ -80,7 +80,8 @@ class WebsocketServerInputTransport(BaseInputTransport):
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._params.serializer.setup(frame) if self._params.serializer:
await self._params.serializer.setup(frame)
if not self._server_task: if not self._server_task:
self._server_task = self.create_task(self._server_task_handler()) self._server_task = self.create_task(self._server_task_handler())
await self.set_transport_ready(frame) await self.set_transport_ready(frame)
@@ -134,6 +135,9 @@ class WebsocketServerInputTransport(BaseInputTransport):
# Handle incoming messages # Handle incoming messages
try: try:
async for message in websocket: async for message in websocket:
if not self._params.serializer:
continue
frame = await self._params.serializer.deserialize(message) frame = await self._params.serializer.deserialize(message)
if not frame: if not frame:
@@ -194,7 +198,8 @@ class WebsocketServerOutputTransport(BaseOutputTransport):
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._params.serializer.setup(frame) if self._params.serializer:
await self._params.serializer.setup(frame)
self._send_interval = (self.audio_chunk_size / self.sample_rate) / 2 self._send_interval = (self.audio_chunk_size / self.sample_rate) / 2
await self.set_transport_ready(frame) await self.set_transport_ready(frame)
@@ -252,6 +257,9 @@ class WebsocketServerOutputTransport(BaseOutputTransport):
await self._write_audio_sleep() await self._write_audio_sleep()
async def _write_frame(self, frame: Frame): async def _write_frame(self, frame: Frame):
if not self._params.serializer:
return
try: try:
payload = await self._params.serializer.serialize(frame) payload = await self._params.serializer.serialize(frame)
if payload and self._websocket: if payload and self._websocket: