transports(network): synchronize with time before sending data
This commit is contained in:
@@ -48,6 +48,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
|
- Websocket transports (FastAPI and Websocket) now synchronize with time before
|
||||||
|
sending data. This allows for interruptions to just work out of the box.
|
||||||
|
|
||||||
- Improved bot speaking detection for all TTS services by using actual bot
|
- Improved bot speaking detection for all TTS services by using actual bot
|
||||||
audio.
|
audio.
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import io
|
import io
|
||||||
|
import time
|
||||||
import wave
|
import wave
|
||||||
|
|
||||||
from typing import Awaitable, Callable
|
from typing import Awaitable, Callable
|
||||||
@@ -42,7 +43,6 @@ except ModuleNotFoundError as e:
|
|||||||
|
|
||||||
class FastAPIWebsocketParams(TransportParams):
|
class FastAPIWebsocketParams(TransportParams):
|
||||||
add_wav_header: bool = False
|
add_wav_header: bool = False
|
||||||
audio_frame_size: int = 6400 # 200ms
|
|
||||||
serializer: FrameSerializer
|
serializer: FrameSerializer
|
||||||
|
|
||||||
|
|
||||||
@@ -105,44 +105,52 @@ class FastAPIWebsocketOutputTransport(BaseOutputTransport):
|
|||||||
|
|
||||||
self._websocket = websocket
|
self._websocket = websocket
|
||||||
self._params = params
|
self._params = params
|
||||||
self._websocket_audio_buffer = bytes()
|
|
||||||
|
self._send_interval = (self._audio_chunk_size / self._params.audio_out_sample_rate) / 2
|
||||||
|
self._next_send_time = 0
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
await super().process_frame(frame, direction)
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
if isinstance(frame, StartInterruptionFrame):
|
if isinstance(frame, StartInterruptionFrame):
|
||||||
await self._write_frame(frame)
|
await self._write_frame(frame)
|
||||||
|
self._next_send_time = 0
|
||||||
|
|
||||||
async def write_raw_audio_frames(self, frames: bytes):
|
async def write_raw_audio_frames(self, frames: bytes):
|
||||||
self._websocket_audio_buffer += frames
|
frame = AudioRawFrame(
|
||||||
while len(self._websocket_audio_buffer):
|
audio=frames,
|
||||||
frame = AudioRawFrame(
|
sample_rate=self._params.audio_out_sample_rate,
|
||||||
audio=self._websocket_audio_buffer[: self._params.audio_frame_size],
|
num_channels=self._params.audio_out_channels,
|
||||||
sample_rate=self._params.audio_out_sample_rate,
|
)
|
||||||
num_channels=self._params.audio_out_channels,
|
|
||||||
|
if self._params.add_wav_header:
|
||||||
|
content = io.BytesIO()
|
||||||
|
ww = wave.open(content, "wb")
|
||||||
|
ww.setsampwidth(2)
|
||||||
|
ww.setnchannels(frame.num_channels)
|
||||||
|
ww.setframerate(frame.sample_rate)
|
||||||
|
ww.writeframes(frame.audio)
|
||||||
|
ww.close()
|
||||||
|
content.seek(0)
|
||||||
|
wav_frame = AudioRawFrame(
|
||||||
|
content.read(), sample_rate=frame.sample_rate, num_channels=frame.num_channels
|
||||||
)
|
)
|
||||||
|
frame = wav_frame
|
||||||
|
|
||||||
if self._params.add_wav_header:
|
payload = self._params.serializer.serialize(frame)
|
||||||
content = io.BytesIO()
|
if payload and self._websocket.client_state == WebSocketState.CONNECTED:
|
||||||
ww = wave.open(content, "wb")
|
await self._websocket.send_text(payload)
|
||||||
ww.setsampwidth(2)
|
|
||||||
ww.setnchannels(frame.num_channels)
|
|
||||||
ww.setframerate(frame.sample_rate)
|
|
||||||
ww.writeframes(frame.audio)
|
|
||||||
ww.close()
|
|
||||||
content.seek(0)
|
|
||||||
wav_frame = AudioRawFrame(
|
|
||||||
content.read(), sample_rate=frame.sample_rate, num_channels=frame.num_channels
|
|
||||||
)
|
|
||||||
frame = wav_frame
|
|
||||||
|
|
||||||
payload = self._params.serializer.serialize(frame)
|
# Simulate a clock.
|
||||||
if payload and self._websocket.client_state == WebSocketState.CONNECTED:
|
current_time = time.monotonic()
|
||||||
await self._websocket.send_text(payload)
|
sleep_duration = max(0, self._next_send_time - current_time)
|
||||||
|
await asyncio.sleep(sleep_duration)
|
||||||
|
if sleep_duration == 0:
|
||||||
|
self._next_send_time = time.monotonic() + self._send_interval
|
||||||
|
else:
|
||||||
|
self._next_send_time += self._send_interval
|
||||||
|
|
||||||
self._websocket_audio_buffer = self._websocket_audio_buffer[
|
self._websocket_audio_buffer = bytes()
|
||||||
self._params.audio_frame_size :
|
|
||||||
]
|
|
||||||
|
|
||||||
async def _write_frame(self, frame: Frame):
|
async def _write_frame(self, frame: Frame):
|
||||||
payload = self._params.serializer.serialize(frame)
|
payload = self._params.serializer.serialize(frame)
|
||||||
|
|||||||
@@ -6,6 +6,7 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import io
|
import io
|
||||||
|
import time
|
||||||
import wave
|
import wave
|
||||||
|
|
||||||
from typing import Awaitable, Callable
|
from typing import Awaitable, Callable
|
||||||
@@ -15,9 +16,12 @@ from pipecat.frames.frames import (
|
|||||||
AudioRawFrame,
|
AudioRawFrame,
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
|
Frame,
|
||||||
InputAudioRawFrame,
|
InputAudioRawFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
|
StartInterruptionFrame,
|
||||||
)
|
)
|
||||||
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.serializers.base_serializer import FrameSerializer
|
from pipecat.serializers.base_serializer import FrameSerializer
|
||||||
from pipecat.serializers.protobuf import ProtobufFrameSerializer
|
from pipecat.serializers.protobuf import ProtobufFrameSerializer
|
||||||
from pipecat.transports.base_input import BaseInputTransport
|
from pipecat.transports.base_input import BaseInputTransport
|
||||||
@@ -36,7 +40,6 @@ except ModuleNotFoundError as e:
|
|||||||
|
|
||||||
class WebsocketServerParams(TransportParams):
|
class WebsocketServerParams(TransportParams):
|
||||||
add_wav_header: bool = False
|
add_wav_header: bool = False
|
||||||
audio_frame_size: int = 6400 # 200ms
|
|
||||||
serializer: FrameSerializer = ProtobufFrameSerializer()
|
serializer: FrameSerializer = ProtobufFrameSerializer()
|
||||||
|
|
||||||
|
|
||||||
@@ -132,45 +135,59 @@ class WebsocketServerOutputTransport(BaseOutputTransport):
|
|||||||
|
|
||||||
self._websocket_audio_buffer = bytes()
|
self._websocket_audio_buffer = bytes()
|
||||||
|
|
||||||
|
self._send_interval = (self._audio_chunk_size / self._params.audio_out_sample_rate) / 2
|
||||||
|
self._next_send_time = 0
|
||||||
|
|
||||||
async def set_client_connection(self, websocket: websockets.WebSocketServerProtocol | None):
|
async def set_client_connection(self, websocket: websockets.WebSocketServerProtocol | None):
|
||||||
if self._websocket:
|
if self._websocket:
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
logger.warning("Only one client allowed, using new connection")
|
logger.warning("Only one client allowed, using new connection")
|
||||||
self._websocket = websocket
|
self._websocket = websocket
|
||||||
|
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
|
if isinstance(frame, StartInterruptionFrame):
|
||||||
|
self._next_send_time = 0
|
||||||
|
|
||||||
async def write_raw_audio_frames(self, frames: bytes):
|
async def write_raw_audio_frames(self, frames: bytes):
|
||||||
if not self._websocket:
|
if not self._websocket:
|
||||||
return
|
return
|
||||||
|
|
||||||
self._websocket_audio_buffer += frames
|
frame = AudioRawFrame(
|
||||||
while len(self._websocket_audio_buffer) >= self._params.audio_frame_size:
|
audio=frames,
|
||||||
frame = AudioRawFrame(
|
sample_rate=self._params.audio_out_sample_rate,
|
||||||
audio=self._websocket_audio_buffer[: self._params.audio_frame_size],
|
num_channels=self._params.audio_out_channels,
|
||||||
sample_rate=self._params.audio_out_sample_rate,
|
)
|
||||||
num_channels=self._params.audio_out_channels,
|
|
||||||
|
if self._params.add_wav_header:
|
||||||
|
content = io.BytesIO()
|
||||||
|
ww = wave.open(content, "wb")
|
||||||
|
ww.setsampwidth(2)
|
||||||
|
ww.setnchannels(frame.num_channels)
|
||||||
|
ww.setframerate(frame.sample_rate)
|
||||||
|
ww.writeframes(frame.audio)
|
||||||
|
ww.close()
|
||||||
|
content.seek(0)
|
||||||
|
wav_frame = AudioRawFrame(
|
||||||
|
content.read(), sample_rate=frame.sample_rate, num_channels=frame.num_channels
|
||||||
)
|
)
|
||||||
|
frame = wav_frame
|
||||||
|
|
||||||
if self._params.add_wav_header:
|
proto = self._params.serializer.serialize(frame)
|
||||||
content = io.BytesIO()
|
if proto:
|
||||||
ww = wave.open(content, "wb")
|
await self._websocket.send(proto)
|
||||||
ww.setsampwidth(2)
|
|
||||||
ww.setnchannels(frame.num_channels)
|
|
||||||
ww.setframerate(frame.sample_rate)
|
|
||||||
ww.writeframes(frame.audio)
|
|
||||||
ww.close()
|
|
||||||
content.seek(0)
|
|
||||||
wav_frame = AudioRawFrame(
|
|
||||||
content.read(), sample_rate=frame.sample_rate, num_channels=frame.num_channels
|
|
||||||
)
|
|
||||||
frame = wav_frame
|
|
||||||
|
|
||||||
proto = self._params.serializer.serialize(frame)
|
# Simulate a clock.
|
||||||
if proto:
|
current_time = time.monotonic()
|
||||||
await self._websocket.send(proto)
|
sleep_duration = max(0, self._next_send_time - current_time)
|
||||||
|
await asyncio.sleep(sleep_duration)
|
||||||
|
if sleep_duration == 0:
|
||||||
|
self._next_send_time = time.monotonic() + self._send_interval
|
||||||
|
else:
|
||||||
|
self._next_send_time += self._send_interval
|
||||||
|
|
||||||
self._websocket_audio_buffer = self._websocket_audio_buffer[
|
self._websocket_audio_buffer = bytes()
|
||||||
self._params.audio_frame_size :
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
class WebsocketServerTransport(BaseTransport):
|
class WebsocketServerTransport(BaseTransport):
|
||||||
|
|||||||
Reference in New Issue
Block a user