New approach to reconnect STT services after updating settings.
This commit is contained in:
@@ -28,6 +28,7 @@ from pipecat.frames.frames import (
|
|||||||
STTMuteFrame,
|
STTMuteFrame,
|
||||||
STTUpdateSettingsFrame,
|
STTUpdateSettingsFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
|
UserStoppedSpeakingFrame,
|
||||||
VADUserStartedSpeakingFrame,
|
VADUserStartedSpeakingFrame,
|
||||||
VADUserStoppedSpeakingFrame,
|
VADUserStoppedSpeakingFrame,
|
||||||
)
|
)
|
||||||
@@ -163,6 +164,16 @@ class STTService(AIService):
|
|||||||
self._keepalive_task: Optional[asyncio.Task] = None
|
self._keepalive_task: Optional[asyncio.Task] = None
|
||||||
self._last_audio_time: float = 0
|
self._last_audio_time: float = 0
|
||||||
|
|
||||||
|
# VAD-aware reconnect state
|
||||||
|
# Whether it is safe to reconnect right now (False while the user is speaking).
|
||||||
|
self._can_reconnect: bool = True
|
||||||
|
# Whether a reconnect has been requested but deferred until speaking ends.
|
||||||
|
self._need_reconnect: bool = False
|
||||||
|
# Whether a reconnect cycle is currently in progress.
|
||||||
|
self._reconnecting: bool = False
|
||||||
|
# Audio frames received while _reconnecting is True, replayed after reconnect.
|
||||||
|
self._reconnect_audio_buffer: list[tuple[AudioRawFrame, FrameDirection]] = []
|
||||||
|
|
||||||
self._register_event_handler("on_connected")
|
self._register_event_handler("on_connected")
|
||||||
self._register_event_handler("on_disconnected")
|
self._register_event_handler("on_disconnected")
|
||||||
self._register_event_handler("on_connection_error")
|
self._register_event_handler("on_connection_error")
|
||||||
@@ -290,6 +301,7 @@ class STTService(AIService):
|
|||||||
await super().cleanup()
|
await super().cleanup()
|
||||||
await self._cancel_ttfb_timeout()
|
await self._cancel_ttfb_timeout()
|
||||||
await self._cancel_keepalive_task()
|
await self._cancel_keepalive_task()
|
||||||
|
self._reconnect_audio_buffer.clear()
|
||||||
|
|
||||||
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
||||||
"""Apply an STT settings delta.
|
"""Apply an STT settings delta.
|
||||||
@@ -331,15 +343,19 @@ class STTService(AIService):
|
|||||||
async def process_audio_frame(self, frame: AudioRawFrame, direction: FrameDirection):
|
async def process_audio_frame(self, frame: AudioRawFrame, direction: FrameDirection):
|
||||||
"""Process an audio frame for speech recognition.
|
"""Process an audio frame for speech recognition.
|
||||||
|
|
||||||
If the service is muted, this method does nothing. Otherwise, it
|
If a reconnect is in progress, the frame is buffered and replayed
|
||||||
processes the audio frame and runs speech-to-text on it, yielding
|
once the connection is restored. If the service is muted, the frame
|
||||||
transcription results. If the frame has a user_id, it is stored
|
is dropped. Otherwise the frame is sent to the STT service and, if
|
||||||
for later use in transcription.
|
a user_id is present, it is stored for use in transcription results.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
frame: The audio frame to process.
|
frame: The audio frame to process.
|
||||||
direction: The direction of frame processing.
|
direction: The direction of frame processing.
|
||||||
"""
|
"""
|
||||||
|
if self._reconnecting:
|
||||||
|
self._reconnect_audio_buffer.append((frame, direction))
|
||||||
|
return
|
||||||
|
|
||||||
if self._muted:
|
if self._muted:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -390,6 +406,9 @@ class STTService(AIService):
|
|||||||
elif isinstance(frame, VADUserStoppedSpeakingFrame):
|
elif isinstance(frame, VADUserStoppedSpeakingFrame):
|
||||||
await self._handle_vad_user_stopped_speaking(frame)
|
await self._handle_vad_user_stopped_speaking(frame)
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
elif isinstance(frame, UserStoppedSpeakingFrame):
|
||||||
|
await self._handle_user_stopped_speaking(frame)
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
elif isinstance(frame, STTUpdateSettingsFrame):
|
elif isinstance(frame, STTUpdateSettingsFrame):
|
||||||
if frame.service is not None and frame.service is not self:
|
if frame.service is not None and frame.service is not self:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
@@ -483,10 +502,25 @@ class STTService(AIService):
|
|||||||
"""
|
"""
|
||||||
await self._reset_stt_ttfb_state()
|
await self._reset_stt_ttfb_state()
|
||||||
self._user_speaking = True
|
self._user_speaking = True
|
||||||
|
self._can_reconnect = False
|
||||||
self._finalize_requested = False
|
self._finalize_requested = False
|
||||||
self._finalize_pending = False
|
self._finalize_pending = False
|
||||||
self._last_transcript_time = 0
|
self._last_transcript_time = 0
|
||||||
|
|
||||||
|
async def _handle_user_stopped_speaking(self, frame: UserStoppedSpeakingFrame):
|
||||||
|
"""Handle user stopped speaking frame.
|
||||||
|
|
||||||
|
Called when the user's full turn has ended and the transcription has been
|
||||||
|
received. Re-enables reconnection and triggers any deferred reconnect that
|
||||||
|
was requested while the user was speaking.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
frame: The user stopped speaking frame.
|
||||||
|
"""
|
||||||
|
self._can_reconnect = True
|
||||||
|
if self._need_reconnect:
|
||||||
|
await self._reconnect()
|
||||||
|
|
||||||
async def _handle_vad_user_stopped_speaking(self, frame: VADUserStoppedSpeakingFrame):
|
async def _handle_vad_user_stopped_speaking(self, frame: VADUserStoppedSpeakingFrame):
|
||||||
"""Handle VAD user stopped speaking frame.
|
"""Handle VAD user stopped speaking frame.
|
||||||
|
|
||||||
@@ -546,6 +580,57 @@ class STTService(AIService):
|
|||||||
await self.cancel_task(self._keepalive_task)
|
await self.cancel_task(self._keepalive_task)
|
||||||
self._keepalive_task = None
|
self._keepalive_task = None
|
||||||
|
|
||||||
|
async def _reconnect(self):
|
||||||
|
"""Perform a full reconnect cycle with audio buffering.
|
||||||
|
|
||||||
|
Sets ``_reconnecting`` so incoming audio frames are buffered rather than
|
||||||
|
sent to a dead connection. Delegates the actual connection reset to
|
||||||
|
``_do_reconnect()``. After the new connection is established all buffered
|
||||||
|
frames are replayed. On failure the error is reported via ``push_error``
|
||||||
|
and the ``on_connection_error`` event handler.
|
||||||
|
"""
|
||||||
|
logger.info(f"{self} reconnecting...")
|
||||||
|
self._reconnect_audio_buffer.clear()
|
||||||
|
self._reconnecting = True
|
||||||
|
self._need_reconnect = False
|
||||||
|
try:
|
||||||
|
await self._do_reconnect()
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"{self} reconnect failed: {e}")
|
||||||
|
await self._call_event_handler("on_connection_error", str(e))
|
||||||
|
await self.push_error(f"{self} reconnect failed: {e}", exception=e)
|
||||||
|
return
|
||||||
|
finally:
|
||||||
|
self._reconnecting = False
|
||||||
|
|
||||||
|
# Replay audio frames that arrived while the connection was down.
|
||||||
|
for buffered_frame, buffered_direction in self._reconnect_audio_buffer:
|
||||||
|
await self.process_audio_frame(buffered_frame, buffered_direction)
|
||||||
|
self._reconnect_audio_buffer.clear()
|
||||||
|
|
||||||
|
async def _do_reconnect(self):
|
||||||
|
"""Perform the service-specific connection reset.
|
||||||
|
|
||||||
|
Called by ``_reconnect()`` inside the reconnecting guard. The default
|
||||||
|
implementation is a no-op. Subclasses that support explicit reconnection
|
||||||
|
(e.g. ``WebsocketSTTService``) should override this to tear down and
|
||||||
|
re-establish their connection.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def _request_reconnect(self):
|
||||||
|
"""Reconnect immediately if safe, or defer until after the current user turn.
|
||||||
|
|
||||||
|
Reconnection is unsafe while the user is speaking because the service is
|
||||||
|
actively receiving audio. Calling this method while the user is speaking
|
||||||
|
schedules a reconnect that fires as soon as ``UserStoppedSpeakingFrame``
|
||||||
|
is received.
|
||||||
|
"""
|
||||||
|
if self._can_reconnect:
|
||||||
|
await self._reconnect()
|
||||||
|
else:
|
||||||
|
self._need_reconnect = True
|
||||||
|
|
||||||
async def _keepalive_task_handler(self):
|
async def _keepalive_task_handler(self):
|
||||||
"""Send periodic silent audio to prevent the server from closing the connection.
|
"""Send periodic silent audio to prevent the server from closing the connection.
|
||||||
|
|
||||||
@@ -737,6 +822,16 @@ class WebsocketSTTService(STTService, WebsocketService):
|
|||||||
await super()._disconnect()
|
await super()._disconnect()
|
||||||
await self._cancel_keepalive_task()
|
await self._cancel_keepalive_task()
|
||||||
|
|
||||||
|
async def _do_reconnect(self):
|
||||||
|
"""Disconnect and reconnect the websocket.
|
||||||
|
|
||||||
|
Called by ``STTService._reconnect()`` inside the reconnecting guard.
|
||||||
|
Tears down the current websocket connection and re-establishes it.
|
||||||
|
Keepalive management is handled by ``_connect`` / ``_disconnect``.
|
||||||
|
"""
|
||||||
|
await self._disconnect()
|
||||||
|
await self._connect()
|
||||||
|
|
||||||
async def _reconnect_websocket(self, attempt_number: int) -> bool:
|
async def _reconnect_websocket(self, attempt_number: int) -> bool:
|
||||||
"""Reconnect and restart keepalive task.
|
"""Reconnect and restart keepalive task.
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user