transport(websocket): do not require a frame serializer
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user