Merge branch 'main' into hume-timestamps
This commit is contained in:
40
CHANGELOG.md
40
CHANGELOG.md
@@ -9,20 +9,51 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|
||||||
|
- Added ai-coustics integrated VAD (`AICVADAnalyzer`) with `AICFilter` factory and
|
||||||
|
example wiring; leverages the enhancement model for robust detection with no
|
||||||
|
ONNX dependency or added processing complexity.
|
||||||
|
|
||||||
|
- Added a watchdog to `DeepgramFluxSTTService` to prevent dangling tasks in case the
|
||||||
|
user was speaking and we stop receiving audio.
|
||||||
|
|
||||||
|
- Introduced a minimum confidence parameter in `DeepgramFluxSTTService` to avoid
|
||||||
|
generating transcriptions below a defined threshold.
|
||||||
|
|
||||||
- Added `ElevenLabsRealtimeSTTService` which implements the Realtime STT
|
- Added `ElevenLabsRealtimeSTTService` which implements the Realtime STT
|
||||||
service from ElevenLabs.
|
service from ElevenLabs.
|
||||||
|
|
||||||
- Added a `TTSService.includes_inter_frame_spaces` property getter, so that TTS
|
- Added ai-coustics integrated VAD (`AICVADAnalyzer`) with `AICFilter` factory and
|
||||||
services that subclass `TTSService` can indicate whether the text in the
|
example wiring; leverages the enhancement model for robust detection with no
|
||||||
`TTSTextFrame`s they push already contain any necessary inter-frame spaces.
|
ONNX dependency or added processing complexity.
|
||||||
|
|
||||||
- Added word-level timestamps support to Hume TTS service
|
- Added word-level timestamps support to Hume TTS service
|
||||||
|
|
||||||
### Changed
|
### Changed
|
||||||
|
|
||||||
|
- ⚠️ Breaking change: `LLMContext.create_image_message()` and
|
||||||
|
`LLMContext.create_audio_message()` are now async methods. This fixes and
|
||||||
|
issue where the asyncio event loop would be blocked while encoding audio or
|
||||||
|
images.
|
||||||
|
|
||||||
|
- `ConsumerProcessor` now queues frames from the producer internally instead of
|
||||||
|
pushing them directly. This allows us to subclass consumer processors and
|
||||||
|
manipulate frames before they are pushed.
|
||||||
|
|
||||||
|
- `BaseTextFilter` only require subclasses to implement the `filter()` method.
|
||||||
|
|
||||||
|
- Extracted the logic for retrying connections, and create a new `send_with_retry`
|
||||||
|
method inside `WebSocketService`.
|
||||||
|
|
||||||
|
- Refactored `DeepgramFluxSTTService` to automatically reconnect if sending a
|
||||||
|
message fails.
|
||||||
|
|
||||||
- Updated all STT and TTS services to use consistent error handling pattern with
|
- Updated all STT and TTS services to use consistent error handling pattern with
|
||||||
`push_error()` method for better pipeline error event integration.
|
`push_error()` method for better pipeline error event integration.
|
||||||
|
|
||||||
|
- Added support for `maybe_capture_participant_camera()` and
|
||||||
|
`maybe_capture_participant_screen()` for `SmallWebRTCTransport` in the runner
|
||||||
|
utils.
|
||||||
|
|
||||||
- Added Hindi support for Rime TTS services.
|
- Added Hindi support for Rime TTS services.
|
||||||
|
|
||||||
- Updated `GeminiTTSService` to use Google Cloud Text-to-Speech streaming API
|
- Updated `GeminiTTSService` to use Google Cloud Text-to-Speech streaming API
|
||||||
@@ -42,6 +73,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
|
- Fixed an issue in the `Runner` where, when using `SmallWebRTCTransport`, the
|
||||||
|
`request_data` was not being passed to the `SmallWebRTCRunnerArguments` body.
|
||||||
|
|
||||||
- Fixed subtle issue of assistant context messages ending up with double spaces
|
- Fixed subtle issue of assistant context messages ending up with double spaces
|
||||||
between words or sentences.
|
between words or sentences.
|
||||||
|
|
||||||
|
|||||||
@@ -52,7 +52,10 @@ transport_params = {
|
|||||||
async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
||||||
logger.info(f"Starting bot")
|
logger.info(f"Starting bot")
|
||||||
|
|
||||||
stt = DeepgramFluxSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
stt = DeepgramFluxSTTService(
|
||||||
|
api_key=os.getenv("DEEPGRAM_API_KEY"),
|
||||||
|
params=DeepgramFluxSTTService.InputParams(min_confidence=0.3),
|
||||||
|
)
|
||||||
|
|
||||||
tts = DeepgramTTSService(api_key=os.getenv("DEEPGRAM_API_KEY"), voice="aura-2-andromeda-en")
|
tts = DeepgramTTSService(api_key=os.getenv("DEEPGRAM_API_KEY"), voice="aura-2-andromeda-en")
|
||||||
|
|
||||||
|
|||||||
@@ -110,7 +110,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
|
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
image = Image.open(image_path)
|
image = Image.open(image_path)
|
||||||
message = LLMContext.create_image_message(
|
message = await LLMContext.create_image_message(
|
||||||
image=image.tobytes(),
|
image=image.tobytes(),
|
||||||
format="RGB",
|
format="RGB",
|
||||||
size=image.size,
|
size=image.size,
|
||||||
|
|||||||
@@ -110,7 +110,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
|
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
image = Image.open(image_path)
|
image = Image.open(image_path)
|
||||||
message = LLMContext.create_image_message(
|
message = await LLMContext.create_image_message(
|
||||||
image=image.tobytes(),
|
image=image.tobytes(),
|
||||||
format="RGB",
|
format="RGB",
|
||||||
size=image.size,
|
size=image.size,
|
||||||
|
|||||||
@@ -117,7 +117,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
|
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
image = Image.open(image_path)
|
image = Image.open(image_path)
|
||||||
message = LLMContext.create_image_message(
|
message = await LLMContext.create_image_message(
|
||||||
image=image.tobytes(),
|
image=image.tobytes(),
|
||||||
format="RGB",
|
format="RGB",
|
||||||
size=image.size,
|
size=image.size,
|
||||||
|
|||||||
@@ -110,7 +110,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
|
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
image = Image.open(image_path)
|
image = Image.open(image_path)
|
||||||
message = LLMContext.create_image_message(
|
message = await LLMContext.create_image_message(
|
||||||
image=image.tobytes(),
|
image=image.tobytes(),
|
||||||
format="RGB",
|
format="RGB",
|
||||||
size=image.size,
|
size=image.size,
|
||||||
|
|||||||
@@ -15,14 +15,21 @@ from pipecat.audio.turn.smart_turn.base_smart_turn import SmartTurnParams
|
|||||||
from pipecat.audio.turn.smart_turn.local_smart_turn_v3 import LocalSmartTurnAnalyzerV3
|
from pipecat.audio.turn.smart_turn.local_smart_turn_v3 import LocalSmartTurnAnalyzerV3
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.audio.vad.vad_analyzer import VADParams
|
from pipecat.audio.vad.vad_analyzer import VADParams
|
||||||
from pipecat.frames.frames import LLMRunFrame, UserImageRequestFrame
|
from pipecat.frames.frames import (
|
||||||
|
Frame,
|
||||||
|
LLMFullResponseEndFrame,
|
||||||
|
LLMFullResponseStartFrame,
|
||||||
|
LLMRunFrame,
|
||||||
|
TextFrame,
|
||||||
|
UserImageRequestFrame,
|
||||||
|
)
|
||||||
from pipecat.pipeline.parallel_pipeline import ParallelPipeline
|
from pipecat.pipeline.parallel_pipeline import ParallelPipeline
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
from pipecat.pipeline.task import PipelineTask
|
from pipecat.pipeline.task import PipelineTask
|
||||||
from pipecat.processors.aggregators.llm_context import LLMContext
|
from pipecat.processors.aggregators.llm_context import LLMContext
|
||||||
from pipecat.processors.aggregators.llm_response_universal import LLMContextAggregatorPair
|
from pipecat.processors.aggregators.llm_response_universal import LLMContextAggregatorPair
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import (
|
from pipecat.runner.utils import (
|
||||||
create_transport,
|
create_transport,
|
||||||
@@ -66,6 +73,27 @@ async def fetch_user_image(params: FunctionCallParams):
|
|||||||
# await params.result_callback({"result": "Image is being captured."})
|
# await params.result_callback({"result": "Image is being captured."})
|
||||||
|
|
||||||
|
|
||||||
|
class MoondreamTextFrameWrapper(FrameProcessor):
|
||||||
|
"""Wraps Moondream-provided TextFrames with LLM response start/end frames.
|
||||||
|
|
||||||
|
This processor detects TextFrames and automatically wraps them with
|
||||||
|
LLMFullResponseStartFrame and LLMFullResponseEndFrame to provide proper
|
||||||
|
response boundaries for downstream processors.
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
|
# If we receive a TextFrame, wrap it with response start/end frames
|
||||||
|
if isinstance(frame, TextFrame):
|
||||||
|
await self.push_frame(LLMFullResponseStartFrame(), direction)
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
await self.push_frame(LLMFullResponseEndFrame(), direction)
|
||||||
|
else:
|
||||||
|
# For all other frames, just pass them through
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|
||||||
# We store functions so objects (e.g. SileroVADAnalyzer) don't get
|
# We store functions so objects (e.g. SileroVADAnalyzer) don't get
|
||||||
# instantiated. The function will be called when the desired transport gets
|
# instantiated. The function will be called when the desired transport gets
|
||||||
# selected.
|
# selected.
|
||||||
@@ -130,6 +158,12 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
# If you run into weird description, try with use_cpu=True
|
# If you run into weird description, try with use_cpu=True
|
||||||
moondream = MoondreamService()
|
moondream = MoondreamService()
|
||||||
|
|
||||||
|
# Wrap TextFrames with LLM response start/end frames, which makes Moondream
|
||||||
|
# output be treated like LLM responses for the purpose of context
|
||||||
|
# aggregation. Without this, the assistant context aggregator would ignore
|
||||||
|
# Moondream output (if the TTS service is disabled).
|
||||||
|
moondream_text_wrapper = MoondreamTextFrameWrapper()
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(), # Transport user input
|
transport.input(), # Transport user input
|
||||||
@@ -137,7 +171,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
context_aggregator.user(), # User responses
|
context_aggregator.user(), # User responses
|
||||||
ParallelPipeline(
|
ParallelPipeline(
|
||||||
[llm], # LLM
|
[llm], # LLM
|
||||||
[moondream],
|
[moondream, moondream_text_wrapper],
|
||||||
),
|
),
|
||||||
tts, # TTS
|
tts, # TTS
|
||||||
transport.output(), # Transport bot output
|
transport.output(), # Transport bot output
|
||||||
|
|||||||
@@ -352,7 +352,10 @@ class TextFrame(DataFrame):
|
|||||||
class LLMTextFrame(TextFrame):
|
class LLMTextFrame(TextFrame):
|
||||||
"""Text frame generated by LLM services."""
|
"""Text frame generated by LLM services."""
|
||||||
|
|
||||||
pass
|
def __post_init__(self):
|
||||||
|
super().__post_init__()
|
||||||
|
# LLM services send text frames with all necessary spaces included
|
||||||
|
self.includes_inter_frame_spaces = True
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ translation from this universal context into whatever format it needs, using a
|
|||||||
service-specific adapter.
|
service-specific adapter.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import base64
|
import base64
|
||||||
import io
|
import io
|
||||||
import wave
|
import wave
|
||||||
@@ -137,7 +138,7 @@ class LLMContext:
|
|||||||
return {"role": role, "content": content}
|
return {"role": role, "content": content}
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def create_image_message(
|
async def create_image_message(
|
||||||
*,
|
*,
|
||||||
role: str = "user",
|
role: str = "user",
|
||||||
format: str,
|
format: str,
|
||||||
@@ -154,15 +155,21 @@ class LLMContext:
|
|||||||
image: Raw image bytes.
|
image: Raw image bytes.
|
||||||
text: Optional text to include with the image.
|
text: Optional text to include with the image.
|
||||||
"""
|
"""
|
||||||
buffer = io.BytesIO()
|
|
||||||
Image.frombytes(format, size, image).save(buffer, format="JPEG")
|
def encode_image():
|
||||||
encoded_image = base64.b64encode(buffer.getvalue()).decode("utf-8")
|
buffer = io.BytesIO()
|
||||||
|
Image.frombytes(format, size, image).save(buffer, format="JPEG")
|
||||||
|
encoded_image = base64.b64encode(buffer.getvalue()).decode("utf-8")
|
||||||
|
return encoded_image
|
||||||
|
|
||||||
|
encoded_image = await asyncio.to_thread(encode_image)
|
||||||
|
|
||||||
url = f"data:image/jpeg;base64,{encoded_image}"
|
url = f"data:image/jpeg;base64,{encoded_image}"
|
||||||
|
|
||||||
return LLMContext.create_image_url_message(role=role, url=url, text=text)
|
return LLMContext.create_image_url_message(role=role, url=url, text=text)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def create_audio_message(
|
async def create_audio_message(
|
||||||
*, role: str = "user", audio_frames: list[AudioRawFrame], text: str = "Audio follows"
|
*, role: str = "user", audio_frames: list[AudioRawFrame], text: str = "Audio follows"
|
||||||
) -> LLMContextMessage:
|
) -> LLMContextMessage:
|
||||||
"""Create a context message containing audio.
|
"""Create a context message containing audio.
|
||||||
@@ -172,21 +179,26 @@ class LLMContext:
|
|||||||
audio_frames: List of audio frame objects to include.
|
audio_frames: List of audio frame objects to include.
|
||||||
text: Optional text to include with the audio.
|
text: Optional text to include with the audio.
|
||||||
"""
|
"""
|
||||||
sample_rate = audio_frames[0].sample_rate
|
|
||||||
num_channels = audio_frames[0].num_channels
|
|
||||||
|
|
||||||
content = []
|
def encode_audio():
|
||||||
content.append({"type": "text", "text": text})
|
sample_rate = audio_frames[0].sample_rate
|
||||||
data = b"".join(frame.audio for frame in audio_frames)
|
num_channels = audio_frames[0].num_channels
|
||||||
|
|
||||||
with io.BytesIO() as buffer:
|
content = []
|
||||||
with wave.open(buffer, "wb") as wf:
|
content.append({"type": "text", "text": text})
|
||||||
wf.setsampwidth(2)
|
data = b"".join(frame.audio for frame in audio_frames)
|
||||||
wf.setnchannels(num_channels)
|
|
||||||
wf.setframerate(sample_rate)
|
|
||||||
wf.writeframes(data)
|
|
||||||
|
|
||||||
encoded_audio = base64.b64encode(buffer.getvalue()).decode("utf-8")
|
with io.BytesIO() as buffer:
|
||||||
|
with wave.open(buffer, "wb") as wf:
|
||||||
|
wf.setsampwidth(2)
|
||||||
|
wf.setnchannels(num_channels)
|
||||||
|
wf.setframerate(sample_rate)
|
||||||
|
wf.writeframes(data)
|
||||||
|
|
||||||
|
encoded_audio = base64.b64encode(buffer.getvalue()).decode("utf-8")
|
||||||
|
return encoded_audio
|
||||||
|
|
||||||
|
encoded_audio = asyncio.to_thread(encode_audio)
|
||||||
|
|
||||||
content.append(
|
content.append(
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -83,4 +83,4 @@ class ConsumerProcessor(FrameProcessor):
|
|||||||
while True:
|
while True:
|
||||||
frame = await self._queue.get()
|
frame = await self._queue.get()
|
||||||
new_frame = await self._transformer(frame)
|
new_frame = await self._transformer(frame)
|
||||||
await self.push_frame(new_frame, self._direction)
|
await self.queue_frame(new_frame, self._direction)
|
||||||
|
|||||||
@@ -264,7 +264,10 @@ def _setup_webrtc_routes(
|
|||||||
# Prepare runner arguments with the callback to run your bot
|
# Prepare runner arguments with the callback to run your bot
|
||||||
async def webrtc_connection_callback(connection):
|
async def webrtc_connection_callback(connection):
|
||||||
bot_module = _get_bot_module()
|
bot_module = _get_bot_module()
|
||||||
runner_args = SmallWebRTCRunnerArguments(webrtc_connection=connection)
|
|
||||||
|
runner_args = SmallWebRTCRunnerArguments(
|
||||||
|
webrtc_connection=connection, body=request.request_data
|
||||||
|
)
|
||||||
background_tasks.add_task(bot_module.bot, runner_args)
|
background_tasks.add_task(bot_module.bot, runner_args)
|
||||||
|
|
||||||
# Delegate handling to SmallWebRTCRequestHandler
|
# Delegate handling to SmallWebRTCRequestHandler
|
||||||
@@ -326,7 +329,8 @@ def _setup_webrtc_routes(
|
|||||||
type=request_data["type"],
|
type=request_data["type"],
|
||||||
pc_id=request_data.get("pc_id"),
|
pc_id=request_data.get("pc_id"),
|
||||||
restart_pc=request_data.get("restart_pc"),
|
restart_pc=request_data.get("restart_pc"),
|
||||||
request_data=request_data,
|
request_data=request_data.get("request_data")
|
||||||
|
or request_data.get("requestData"),
|
||||||
)
|
)
|
||||||
return await offer(webrtc_request, background_tasks)
|
return await offer(webrtc_request, background_tasks)
|
||||||
elif request.method == HTTPMethod.PATCH.value:
|
elif request.method == HTTPMethod.PATCH.value:
|
||||||
|
|||||||
@@ -281,6 +281,14 @@ async def maybe_capture_participant_camera(
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
try:
|
||||||
|
from pipecat.transports.smallwebrtc.transport import SmallWebRTCTransport
|
||||||
|
|
||||||
|
if isinstance(transport, SmallWebRTCTransport):
|
||||||
|
await transport.capture_participant_video(video_source="camera")
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
async def maybe_capture_participant_screen(
|
async def maybe_capture_participant_screen(
|
||||||
transport: BaseTransport, client: Any, framerate: int = 0
|
transport: BaseTransport, client: Any, framerate: int = 0
|
||||||
@@ -303,6 +311,14 @@ async def maybe_capture_participant_screen(
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
try:
|
||||||
|
from pipecat.transports.smallwebrtc.transport import SmallWebRTCTransport
|
||||||
|
|
||||||
|
if isinstance(transport, SmallWebRTCTransport):
|
||||||
|
await transport.capture_participant_video(video_source="screenVideo")
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
def _smallwebrtc_sdp_cleanup_ice_candidates(text: str, pattern: str) -> str:
|
def _smallwebrtc_sdp_cleanup_ice_candidates(text: str, pattern: str) -> str:
|
||||||
"""Clean up ICE candidates in SDP text for SmallWebRTC.
|
"""Clean up ICE candidates in SDP text for SmallWebRTC.
|
||||||
|
|||||||
@@ -373,9 +373,7 @@ class AnthropicLLMService(LLMService):
|
|||||||
|
|
||||||
if event.type == "content_block_delta":
|
if event.type == "content_block_delta":
|
||||||
if hasattr(event.delta, "text"):
|
if hasattr(event.delta, "text"):
|
||||||
frame = LLMTextFrame(event.delta.text)
|
await self.push_frame(LLMTextFrame(event.delta.text))
|
||||||
frame.includes_inter_frame_spaces = True
|
|
||||||
await self.push_frame(frame)
|
|
||||||
completion_tokens_estimate += self._estimate_tokens(event.delta.text)
|
completion_tokens_estimate += self._estimate_tokens(event.delta.text)
|
||||||
elif hasattr(event.delta, "partial_json") and tool_use_block:
|
elif hasattr(event.delta, "partial_json") and tool_use_block:
|
||||||
json_accumulator += event.delta.partial_json
|
json_accumulator += event.delta.partial_json
|
||||||
|
|||||||
@@ -146,15 +146,6 @@ class AsyncAITTSService(InterruptibleTTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that AsyncAI TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that AsyncAI's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Async language format.
|
"""Convert a Language enum to Async language format.
|
||||||
|
|
||||||
@@ -433,15 +424,6 @@ class AsyncAIHttpTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that AsyncAI TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that AsyncAI's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Async language format.
|
"""Convert a Language enum to Async language format.
|
||||||
|
|
||||||
|
|||||||
@@ -1078,9 +1078,7 @@ class AWSBedrockLLMService(LLMService):
|
|||||||
if "contentBlockDelta" in event:
|
if "contentBlockDelta" in event:
|
||||||
delta = event["contentBlockDelta"]["delta"]
|
delta = event["contentBlockDelta"]["delta"]
|
||||||
if "text" in delta:
|
if "text" in delta:
|
||||||
frame = LLMTextFrame(delta["text"])
|
await self.push_frame(LLMTextFrame(delta["text"]))
|
||||||
frame.includes_inter_frame_spaces = True
|
|
||||||
await self.push_frame(frame)
|
|
||||||
completion_tokens_estimate += self._estimate_tokens(delta["text"])
|
completion_tokens_estimate += self._estimate_tokens(delta["text"])
|
||||||
elif "toolUse" in delta and "input" in delta["toolUse"]:
|
elif "toolUse" in delta and "input" in delta["toolUse"]:
|
||||||
# Handle partial JSON for tool use
|
# Handle partial JSON for tool use
|
||||||
|
|||||||
@@ -209,15 +209,6 @@ class AWSPollyTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that AWS TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that AWS's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to AWS Polly language format.
|
"""Convert a Language enum to AWS Polly language format.
|
||||||
|
|
||||||
|
|||||||
@@ -151,15 +151,6 @@ class AzureBaseTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Azure TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Azure's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Azure language format.
|
"""Convert a Language enum to Azure language format.
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,9 @@
|
|||||||
|
|
||||||
"""Deepgram Flux speech-to-text service implementation."""
|
"""Deepgram Flux speech-to-text service implementation."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import json
|
import json
|
||||||
|
import time
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, AsyncGenerator, Dict, Optional
|
from typing import Any, AsyncGenerator, Dict, Optional
|
||||||
from urllib.parse import urlencode
|
from urllib.parse import urlencode
|
||||||
@@ -94,6 +96,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
mip_opt_out: Optional. Opts out requests from the Deepgram Model Improvement Program
|
mip_opt_out: Optional. Opts out requests from the Deepgram Model Improvement Program
|
||||||
(default False).
|
(default False).
|
||||||
tag: List of tags to label requests for identification during usage reporting.
|
tag: List of tags to label requests for identification during usage reporting.
|
||||||
|
min_confidence: Optional. Minimum confidence required confidence to create a TranscriptionFrame
|
||||||
"""
|
"""
|
||||||
|
|
||||||
eager_eot_threshold: Optional[float] = None
|
eager_eot_threshold: Optional[float] = None
|
||||||
@@ -102,6 +105,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
keyterm: list = []
|
keyterm: list = []
|
||||||
mip_opt_out: Optional[bool] = None
|
mip_opt_out: Optional[bool] = None
|
||||||
tag: list = []
|
tag: list = []
|
||||||
|
min_confidence: Optional[float] = None # New parameter
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -163,6 +167,13 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
self._register_event_handler("on_end_of_turn")
|
self._register_event_handler("on_end_of_turn")
|
||||||
self._register_event_handler("on_eager_end_of_turn")
|
self._register_event_handler("on_eager_end_of_turn")
|
||||||
self._register_event_handler("on_update")
|
self._register_event_handler("on_update")
|
||||||
|
self._connection_established_event = asyncio.Event()
|
||||||
|
# Watchdog task to prevent dangling tasks
|
||||||
|
# If we stop sending audio to Flux after we have received that the User has started speaking
|
||||||
|
# we never receive the user stopped speaking event unless we resume sending audio to it.
|
||||||
|
self._last_stt_time = None
|
||||||
|
self._watchdog_task = None
|
||||||
|
self._user_is_speaking = False
|
||||||
|
|
||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
"""Connect to WebSocket and start background tasks.
|
"""Connect to WebSocket and start background tasks.
|
||||||
@@ -172,9 +183,6 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
"""
|
"""
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
if self._websocket and not self._receive_task:
|
|
||||||
self._receive_task = self.create_task(self._receive_task_handler(self._report_error))
|
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
"""Disconnect from WebSocket and clean up tasks.
|
"""Disconnect from WebSocket and clean up tasks.
|
||||||
|
|
||||||
@@ -182,14 +190,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
and cleans up resources to prevent memory leaks.
|
and cleans up resources to prevent memory leaks.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
# Cancel background tasks BEFORE closing websocket
|
|
||||||
if self._receive_task:
|
|
||||||
await self.cancel_task(self._receive_task, timeout=2.0)
|
|
||||||
self._receive_task = None
|
|
||||||
|
|
||||||
# Now close the websocket
|
|
||||||
await self._disconnect_websocket()
|
await self._disconnect_websocket()
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} exception: {e}")
|
logger.error(f"{self} exception: {e}")
|
||||||
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
|
await self.push_error(ErrorFrame(error=f"{self} error: {e}"))
|
||||||
@@ -197,6 +198,25 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
# Reset state only after everything is cleaned up
|
# Reset state only after everything is cleaned up
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
|
|
||||||
|
async def _send_silence(self, duration_secs: float = 0.5):
|
||||||
|
"""Send a block of silence of the specified duration (default 500 ms)."""
|
||||||
|
sample_width = 2 # bytes per sample for 16-bit PCM
|
||||||
|
num_channels = 1 # mono
|
||||||
|
num_samples = int(self.sample_rate * duration_secs)
|
||||||
|
silence = b"\x00" * (num_samples * sample_width * num_channels)
|
||||||
|
await self._websocket.send(silence)
|
||||||
|
|
||||||
|
async def _watchdog_task_handler(self):
|
||||||
|
while self._websocket and self._websocket.state is State.OPEN:
|
||||||
|
now = time.monotonic()
|
||||||
|
# More than 500 ms without sending new audio to Flux
|
||||||
|
if self._user_is_speaking and self._last_stt_time and now - self._last_stt_time > 0.5:
|
||||||
|
logger.warning("Sending silence to Flux to prevent dangling task")
|
||||||
|
await self._send_silence()
|
||||||
|
self._last_stt_time = time.monotonic()
|
||||||
|
# check every 100ms
|
||||||
|
await asyncio.sleep(0.1)
|
||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
"""Establish WebSocket connection to API.
|
"""Establish WebSocket connection to API.
|
||||||
|
|
||||||
@@ -208,10 +228,26 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
if self._websocket and self._websocket.state is State.OPEN:
|
if self._websocket and self._websocket.state is State.OPEN:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
self._connection_established_event.clear()
|
||||||
|
self._user_is_speaking = False
|
||||||
self._websocket = await websocket_connect(
|
self._websocket = await websocket_connect(
|
||||||
self._websocket_url,
|
self._websocket_url,
|
||||||
additional_headers={"Authorization": f"Token {self._api_key}"},
|
additional_headers={"Authorization": f"Token {self._api_key}"},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Creating the receiver task
|
||||||
|
if not self._receive_task:
|
||||||
|
self._receive_task = self.create_task(
|
||||||
|
self._receive_task_handler(self._report_error)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Creating the watchdog task
|
||||||
|
if not self._watchdog_task:
|
||||||
|
self._watchdog_task = self.create_task(self._watchdog_task_handler())
|
||||||
|
|
||||||
|
# Now wait for the connection established event
|
||||||
|
logger.debug("WebSocket connected, waiting for server confirmation...")
|
||||||
|
await self._connection_established_event.wait()
|
||||||
logger.debug("Connected to Deepgram Flux Websocket")
|
logger.debug("Connected to Deepgram Flux Websocket")
|
||||||
await self._call_event_handler("on_connected")
|
await self._call_event_handler("on_connected")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -227,6 +263,16 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
metrics collection. Handles disconnection errors gracefully.
|
metrics collection. Handles disconnection errors gracefully.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
|
# Cancel background tasks BEFORE closing websocket
|
||||||
|
if self._receive_task:
|
||||||
|
await self.cancel_task(self._receive_task, timeout=2.0)
|
||||||
|
self._receive_task = None
|
||||||
|
if self._watchdog_task:
|
||||||
|
await self.cancel_task(self._watchdog_task, timeout=2.0)
|
||||||
|
self._watchdog_task = None
|
||||||
|
self._last_stt_time = None
|
||||||
|
|
||||||
|
self._connection_established_event.clear()
|
||||||
await self.stop_all_metrics()
|
await self.stop_all_metrics()
|
||||||
|
|
||||||
if self._websocket:
|
if self._websocket:
|
||||||
@@ -340,7 +386,8 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
return
|
return
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await self._websocket.send(audio)
|
self._last_stt_time = time.monotonic()
|
||||||
|
await self.send_with_retry(audio, self._report_error)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} exception: {e}")
|
logger.error(f"{self} exception: {e}")
|
||||||
yield ErrorFrame(error=f"{self} error: {e}")
|
yield ErrorFrame(error=f"{self} error: {e}")
|
||||||
@@ -463,6 +510,8 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
transcription processing.
|
transcription processing.
|
||||||
"""
|
"""
|
||||||
logger.info("Connected to Flux - ready to stream audio")
|
logger.info("Connected to Flux - ready to stream audio")
|
||||||
|
# Notify connection is established
|
||||||
|
self._connection_established_event.set()
|
||||||
|
|
||||||
async def _handle_fatal_error(self, data: Dict[str, Any]):
|
async def _handle_fatal_error(self, data: Dict[str, Any]):
|
||||||
"""Handle fatal error messages from Deepgram Flux.
|
"""Handle fatal error messages from Deepgram Flux.
|
||||||
@@ -530,6 +579,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
transcript: maybe the first few words of the turn.
|
transcript: maybe the first few words of the turn.
|
||||||
"""
|
"""
|
||||||
logger.debug("User started speaking")
|
logger.debug("User started speaking")
|
||||||
|
self._user_is_speaking = True
|
||||||
await self.push_interruption_task_frame_and_wait()
|
await self.push_interruption_task_frame_and_wait()
|
||||||
await self.broadcast_frame(UserStartedSpeakingFrame)
|
await self.broadcast_frame(UserStartedSpeakingFrame)
|
||||||
await self.start_metrics()
|
await self.start_metrics()
|
||||||
@@ -550,6 +600,22 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
logger.trace(f"Received event TurnResumed: {event}")
|
logger.trace(f"Received event TurnResumed: {event}")
|
||||||
await self._call_event_handler("on_turn_resumed")
|
await self._call_event_handler("on_turn_resumed")
|
||||||
|
|
||||||
|
def _calculate_average_confidence(self, transcript_data) -> Optional[float]:
|
||||||
|
"""Calculate the average confidence from transcript data.
|
||||||
|
|
||||||
|
Return None if the data is missing or invalid.
|
||||||
|
"""
|
||||||
|
# Example: Assume transcript_data has a list of words with confidence
|
||||||
|
words = transcript_data.get("words")
|
||||||
|
if not words or not isinstance(words, list):
|
||||||
|
return None
|
||||||
|
confidences = [
|
||||||
|
w.get("confidence") for w in words if isinstance(w.get("confidence"), (float, int))
|
||||||
|
]
|
||||||
|
if not confidences:
|
||||||
|
return None
|
||||||
|
return sum(confidences) / len(confidences)
|
||||||
|
|
||||||
async def _handle_end_of_turn(self, transcript: str, data: Dict[str, Any]):
|
async def _handle_end_of_turn(self, transcript: str, data: Dict[str, Any]):
|
||||||
"""Handle EndOfTurn events from Deepgram Flux.
|
"""Handle EndOfTurn events from Deepgram Flux.
|
||||||
|
|
||||||
@@ -569,16 +635,26 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
data: The TurnInfo message data containing event type, transcript and some extra metadata.
|
data: The TurnInfo message data containing event type, transcript and some extra metadata.
|
||||||
"""
|
"""
|
||||||
logger.debug("User stopped speaking")
|
logger.debug("User stopped speaking")
|
||||||
|
self._user_is_speaking = False
|
||||||
|
|
||||||
await self.push_frame(
|
# Compute the average confidence
|
||||||
TranscriptionFrame(
|
average_confidence = self._calculate_average_confidence(data)
|
||||||
transcript,
|
|
||||||
self._user_id,
|
if not self._params.min_confidence or average_confidence > self._params.min_confidence:
|
||||||
time_now_iso8601(),
|
await self.push_frame(
|
||||||
self._language,
|
TranscriptionFrame(
|
||||||
result=data,
|
transcript,
|
||||||
|
self._user_id,
|
||||||
|
time_now_iso8601(),
|
||||||
|
self._language,
|
||||||
|
result=data,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
)
|
else:
|
||||||
|
logger.warning(
|
||||||
|
f"Transcription confidence below min_confidence threshold: {average_confidence}"
|
||||||
|
)
|
||||||
|
|
||||||
await self._handle_transcription(transcript, True, self._language)
|
await self._handle_transcription(transcript, True, self._language)
|
||||||
await self.stop_processing_metrics()
|
await self.stop_processing_metrics()
|
||||||
await self.push_frame(UserStoppedSpeakingFrame(), FrameDirection.DOWNSTREAM)
|
await self.push_frame(UserStoppedSpeakingFrame(), FrameDirection.DOWNSTREAM)
|
||||||
|
|||||||
@@ -79,15 +79,6 @@ class DeepgramTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Deepgram TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Deepgram's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
@traced_tts
|
@traced_tts
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
"""Generate speech from text using Deepgram's TTS API.
|
"""Generate speech from text using Deepgram's TTS API.
|
||||||
@@ -177,15 +168,6 @@ class DeepgramHttpTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Deepgram TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Deepgram's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
@traced_tts
|
@traced_tts
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
"""Generate speech from text using Deepgram's TTS API.
|
"""Generate speech from text using Deepgram's TTS API.
|
||||||
|
|||||||
@@ -159,15 +159,6 @@ class FishAudioTTSService(InterruptibleTTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Fish Audio TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Fish Audio's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
"""Set the TTS model and reconnect.
|
"""Set the TTS model and reconnect.
|
||||||
|
|
||||||
|
|||||||
@@ -1452,8 +1452,6 @@ class GeminiLiveLLMService(LLMService):
|
|||||||
self._bot_text_buffer += text
|
self._bot_text_buffer += text
|
||||||
self._search_result_buffer += text # Also accumulate for grounding
|
self._search_result_buffer += text # Also accumulate for grounding
|
||||||
frame = LLMTextFrame(text=text)
|
frame = LLMTextFrame(text=text)
|
||||||
# Gemini Live text already includes any necessary inter-chunk spaces
|
|
||||||
frame.includes_inter_frame_spaces = True
|
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
# Check for grounding metadata in server content
|
# Check for grounding metadata in server content
|
||||||
|
|||||||
@@ -920,9 +920,7 @@ class GoogleLLMService(LLMService):
|
|||||||
for part in candidate.content.parts:
|
for part in candidate.content.parts:
|
||||||
if not part.thought and part.text:
|
if not part.thought and part.text:
|
||||||
search_result += part.text
|
search_result += part.text
|
||||||
frame = LLMTextFrame(part.text)
|
await self.push_frame(LLMTextFrame(part.text))
|
||||||
frame.includes_inter_frame_spaces = True
|
|
||||||
await self.push_frame(frame)
|
|
||||||
elif part.function_call:
|
elif part.function_call:
|
||||||
function_call = part.function_call
|
function_call = part.function_call
|
||||||
id = function_call.id or str(uuid.uuid4())
|
id = function_call.id or str(uuid.uuid4())
|
||||||
|
|||||||
@@ -596,15 +596,6 @@ class GoogleHttpTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Google TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Google's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Google TTS language format.
|
"""Convert a Language enum to Google TTS language format.
|
||||||
|
|
||||||
@@ -803,15 +794,6 @@ class GoogleBaseTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Google and Gemini TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Google's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Google TTS language format.
|
"""Convert a Language enum to Google TTS language format.
|
||||||
|
|
||||||
|
|||||||
@@ -111,15 +111,6 @@ class GroqTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Groq TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Groq's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
@traced_tts
|
@traced_tts
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
"""Generate speech from text using Groq's TTS API.
|
"""Generate speech from text using Groq's TTS API.
|
||||||
|
|||||||
@@ -123,15 +123,6 @@ class HumeTTSService(WordTTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Hume TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Hume's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def start(self, frame: StartFrame) -> None:
|
async def start(self, frame: StartFrame) -> None:
|
||||||
"""Start the service.
|
"""Start the service.
|
||||||
|
|
||||||
|
|||||||
@@ -250,15 +250,6 @@ class InworldTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Inworld TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Inworld's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
"""Start the Inworld TTS service.
|
"""Start the Inworld TTS service.
|
||||||
|
|
||||||
|
|||||||
@@ -124,15 +124,6 @@ class LmntTTSService(InterruptibleTTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that LMNT TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that LMNT's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to LMNT service language format.
|
"""Convert a Language enum to LMNT service language format.
|
||||||
|
|
||||||
|
|||||||
@@ -194,15 +194,6 @@ class MiniMaxHttpTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that MiniMax TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that MiniMax's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to MiniMax service language format.
|
"""Convert a Language enum to MiniMax service language format.
|
||||||
|
|
||||||
|
|||||||
@@ -151,15 +151,6 @@ class NeuphonicTTSService(InterruptibleTTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Neuphonic TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Neuphonic's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Neuphonic service language format.
|
"""Convert a Language enum to Neuphonic service language format.
|
||||||
|
|
||||||
@@ -449,15 +440,6 @@ class NeuphonicHttpTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Neuphonic TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Neuphonic's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Neuphonic service language format.
|
"""Convert a Language enum to Neuphonic service language format.
|
||||||
|
|
||||||
|
|||||||
@@ -390,9 +390,7 @@ class BaseOpenAILLMService(LLMService):
|
|||||||
# Keep iterating through the response to collect all the argument fragments
|
# Keep iterating through the response to collect all the argument fragments
|
||||||
arguments += tool_call.function.arguments
|
arguments += tool_call.function.arguments
|
||||||
elif chunk.choices[0].delta.content:
|
elif chunk.choices[0].delta.content:
|
||||||
frame = LLMTextFrame(chunk.choices[0].delta.content)
|
await self.push_frame(LLMTextFrame(chunk.choices[0].delta.content))
|
||||||
frame.includes_inter_frame_spaces = True
|
|
||||||
await self.push_frame(frame)
|
|
||||||
|
|
||||||
# When gpt-4o-audio / gpt-4o-mini-audio is used for llm or stt+llm
|
# When gpt-4o-audio / gpt-4o-mini-audio is used for llm or stt+llm
|
||||||
# we need to get LLMTextFrame for the transcript
|
# we need to get LLMTextFrame for the transcript
|
||||||
|
|||||||
@@ -678,8 +678,6 @@ class OpenAIRealtimeLLMService(LLMService):
|
|||||||
# the output modality is "text"
|
# the output modality is "text"
|
||||||
if evt.delta:
|
if evt.delta:
|
||||||
frame = LLMTextFrame(evt.delta)
|
frame = LLMTextFrame(evt.delta)
|
||||||
# OpenAI Realtime text already includes any necessary inter-chunk spaces
|
|
||||||
frame.includes_inter_frame_spaces = True
|
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
async def _handle_evt_audio_transcript_delta(self, evt):
|
async def _handle_evt_audio_transcript_delta(self, evt):
|
||||||
|
|||||||
@@ -131,15 +131,6 @@ class OpenAITTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that OpenAI TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that OpenAI's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
"""Set the TTS model to use.
|
"""Set the TTS model to use.
|
||||||
|
|
||||||
|
|||||||
@@ -66,15 +66,6 @@ class PiperTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Piper TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Piper's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
@traced_tts
|
@traced_tts
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
"""Generate speech from text using Piper's HTTP API.
|
"""Generate speech from text using Piper's HTTP API.
|
||||||
|
|||||||
@@ -501,15 +501,6 @@ class RimeHttpTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Rime TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Rime's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> str | None:
|
def language_to_service_language(self, language: Language) -> str | None:
|
||||||
"""Convert pipecat language to Rime language code.
|
"""Convert pipecat language to Rime language code.
|
||||||
|
|
||||||
|
|||||||
@@ -113,15 +113,6 @@ class RivaTTSService(TTSService):
|
|||||||
riva.client.proto.riva_tts_pb2.RivaSynthesisConfigRequest()
|
riva.client.proto.riva_tts_pb2.RivaSynthesisConfigRequest()
|
||||||
)
|
)
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Riva TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Riva's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
"""Attempt to set the TTS model.
|
"""Attempt to set the TTS model.
|
||||||
|
|
||||||
@@ -166,7 +157,6 @@ class RivaTTSService(TTSService):
|
|||||||
add_response(None)
|
add_response(None)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} exception: {e}")
|
logger.error(f"{self} exception: {e}")
|
||||||
yield ErrorFrame(error=f"{self} error: {e}")
|
|
||||||
add_response(None)
|
add_response(None)
|
||||||
|
|
||||||
await self.start_ttfb_metrics()
|
await self.start_ttfb_metrics()
|
||||||
@@ -191,6 +181,7 @@ class RivaTTSService(TTSService):
|
|||||||
resp = await asyncio.wait_for(queue.get(), timeout=RIVA_TTS_TIMEOUT_SECS)
|
resp = await asyncio.wait_for(queue.get(), timeout=RIVA_TTS_TIMEOUT_SECS)
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
logger.error(f"{self} timeout waiting for audio response")
|
logger.error(f"{self} timeout waiting for audio response")
|
||||||
|
yield ErrorFrame(error=f"{self} error: {e}")
|
||||||
|
|
||||||
await self.start_tts_usage_metrics(text)
|
await self.start_tts_usage_metrics(text)
|
||||||
yield TTSStoppedFrame()
|
yield TTSStoppedFrame()
|
||||||
|
|||||||
@@ -176,9 +176,7 @@ class SambaNovaLLMService(OpenAILLMService): # type: ignore
|
|||||||
# Keep iterating through the response to collect all the argument fragments
|
# Keep iterating through the response to collect all the argument fragments
|
||||||
arguments += tool_call.function.arguments
|
arguments += tool_call.function.arguments
|
||||||
elif chunk.choices[0].delta.content:
|
elif chunk.choices[0].delta.content:
|
||||||
frame = LLMTextFrame(chunk.choices[0].delta.content)
|
await self.push_frame(LLMTextFrame(chunk.choices[0].delta.content))
|
||||||
frame.includes_inter_frame_spaces = True
|
|
||||||
await self.push_frame(frame)
|
|
||||||
|
|
||||||
# When gpt-4o-audio / gpt-4o-mini-audio is used for llm or stt+llm
|
# When gpt-4o-audio / gpt-4o-mini-audio is used for llm or stt+llm
|
||||||
# we need to get LLMTextFrame for the transcript
|
# we need to get LLMTextFrame for the transcript
|
||||||
|
|||||||
@@ -195,15 +195,6 @@ class SarvamHttpTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Sarvam TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Sarvam's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Sarvam AI language format.
|
"""Convert a Language enum to Sarvam AI language format.
|
||||||
|
|
||||||
@@ -467,15 +458,6 @@ class SarvamTTSService(InterruptibleTTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Sarvam TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Sarvam's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
"""Convert a Language enum to Sarvam AI language format.
|
"""Convert a Language enum to Sarvam AI language format.
|
||||||
|
|
||||||
|
|||||||
@@ -105,15 +105,6 @@ class SpeechmaticsTTSService(TTSService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates that Speechmatics TTSTextFrames include necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True, indicating that Speechmatics's text frames include necessary inter-frame spaces.
|
|
||||||
"""
|
|
||||||
return True
|
|
||||||
|
|
||||||
@traced_tts
|
@traced_tts
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
"""Generate speech from text using Speechmatics' HTTP API.
|
"""Generate speech from text using Speechmatics' HTTP API.
|
||||||
|
|||||||
@@ -142,6 +142,7 @@ class TTSService(AIService):
|
|||||||
self._voice_id: str = ""
|
self._voice_id: str = ""
|
||||||
self._settings: Dict[str, Any] = {}
|
self._settings: Dict[str, Any] = {}
|
||||||
self._text_aggregator: BaseTextAggregator = text_aggregator or SimpleTextAggregator()
|
self._text_aggregator: BaseTextAggregator = text_aggregator or SimpleTextAggregator()
|
||||||
|
self._aggregated_text_includes_inter_frame_spaces: bool = False
|
||||||
self._text_filters: Sequence[BaseTextFilter] = text_filters or []
|
self._text_filters: Sequence[BaseTextFilter] = text_filters or []
|
||||||
self._transport_destination: Optional[str] = transport_destination
|
self._transport_destination: Optional[str] = transport_destination
|
||||||
self._tracing_enabled: bool = False
|
self._tracing_enabled: bool = False
|
||||||
@@ -192,23 +193,6 @@ class TTSService(AIService):
|
|||||||
CHUNK_SECONDS = 0.5
|
CHUNK_SECONDS = 0.5
|
||||||
return int(self.sample_rate * CHUNK_SECONDS * 2) # 2 bytes/sample
|
return int(self.sample_rate * CHUNK_SECONDS * 2) # 2 bytes/sample
|
||||||
|
|
||||||
@property
|
|
||||||
def includes_inter_frame_spaces(self) -> bool:
|
|
||||||
"""Indicates whether TTSTextFrames include necesary inter-frame spaces.
|
|
||||||
|
|
||||||
When True, the TTSTextFrame objects pushed by this service already
|
|
||||||
include all necessary spaces between subsequent frames. When False,
|
|
||||||
downstream processors (like the assistant context aggregator) may need
|
|
||||||
to add spacing.
|
|
||||||
|
|
||||||
Subclasses should override this property to return True if their text
|
|
||||||
generation process already includes necessary inter-frame spaces.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
False by default. Subclasses can override to return True.
|
|
||||||
"""
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
"""Set the TTS model to use.
|
"""Set the TTS model to use.
|
||||||
|
|
||||||
@@ -369,9 +353,16 @@ class TTSService(AIService):
|
|||||||
await self._maybe_pause_frame_processing()
|
await self._maybe_pause_frame_processing()
|
||||||
|
|
||||||
sentence = self._text_aggregator.text
|
sentence = self._text_aggregator.text
|
||||||
|
includes_inter_frame_spaces = self._aggregated_text_includes_inter_frame_spaces
|
||||||
|
|
||||||
|
# Reset aggregator state
|
||||||
await self._text_aggregator.reset()
|
await self._text_aggregator.reset()
|
||||||
self._processing_text = False
|
self._processing_text = False
|
||||||
await self._push_tts_frames(sentence)
|
self._aggregated_text_includes_inter_frame_spaces = False
|
||||||
|
|
||||||
|
await self._push_tts_frames(
|
||||||
|
sentence, includes_inter_frame_spaces=includes_inter_frame_spaces
|
||||||
|
)
|
||||||
if isinstance(frame, LLMFullResponseEndFrame):
|
if isinstance(frame, LLMFullResponseEndFrame):
|
||||||
if self._push_text_frames:
|
if self._push_text_frames:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
@@ -380,7 +371,8 @@ class TTSService(AIService):
|
|||||||
elif isinstance(frame, TTSSpeakFrame):
|
elif isinstance(frame, TTSSpeakFrame):
|
||||||
# Store if we were processing text or not so we can set it back.
|
# Store if we were processing text or not so we can set it back.
|
||||||
processing_text = self._processing_text
|
processing_text = self._processing_text
|
||||||
await self._push_tts_frames(frame.text)
|
# Assumption: text in TTSSpeakFrame does not include inter-frame spaces
|
||||||
|
await self._push_tts_frames(frame.text, includes_inter_frame_spaces=False)
|
||||||
# We pause processing incoming frames because we are sending data to
|
# We pause processing incoming frames because we are sending data to
|
||||||
# the TTS. We pause to avoid audio overlapping.
|
# the TTS. We pause to avoid audio overlapping.
|
||||||
await self._maybe_pause_frame_processing()
|
await self._maybe_pause_frame_processing()
|
||||||
@@ -474,11 +466,17 @@ class TTSService(AIService):
|
|||||||
text = frame.text
|
text = frame.text
|
||||||
else:
|
else:
|
||||||
text = await self._text_aggregator.aggregate(frame.text)
|
text = await self._text_aggregator.aggregate(frame.text)
|
||||||
|
# Assumption: whether inter-frame spaces are included shouldn't
|
||||||
|
# change during aggregation, so we can just use the latest frame's
|
||||||
|
# value
|
||||||
|
self._aggregated_text_includes_inter_frame_spaces = frame.includes_inter_frame_spaces
|
||||||
|
|
||||||
if text:
|
if text:
|
||||||
await self._push_tts_frames(text)
|
await self._push_tts_frames(
|
||||||
|
text, includes_inter_frame_spaces=frame.includes_inter_frame_spaces
|
||||||
|
)
|
||||||
|
|
||||||
async def _push_tts_frames(self, text: str):
|
async def _push_tts_frames(self, text: str, includes_inter_frame_spaces: bool):
|
||||||
# Remove leading newlines only
|
# Remove leading newlines only
|
||||||
text = text.lstrip("\n")
|
text = text.lstrip("\n")
|
||||||
|
|
||||||
@@ -508,7 +506,7 @@ class TTSService(AIService):
|
|||||||
# We send the original text after the audio. This way, if we are
|
# We send the original text after the audio. This way, if we are
|
||||||
# interrupted, the text is not added to the assistant context.
|
# interrupted, the text is not added to the assistant context.
|
||||||
frame = TTSTextFrame(text)
|
frame = TTSTextFrame(text)
|
||||||
frame.includes_inter_frame_spaces = self.includes_inter_frame_spaces
|
frame.includes_inter_frame_spaces = includes_inter_frame_spaces
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
async def _stop_frame_handler(self):
|
async def _stop_frame_handler(self):
|
||||||
@@ -635,6 +633,8 @@ class WordTTSService(TTSService):
|
|||||||
frame = TTSStoppedFrame()
|
frame = TTSStoppedFrame()
|
||||||
frame.pts = last_pts
|
frame.pts = last_pts
|
||||||
else:
|
else:
|
||||||
|
# Assumption: word-by-word text frames don't include spaces, so
|
||||||
|
# we can rely on the default includes_inter_frame_spaces=False
|
||||||
frame = TTSTextFrame(word)
|
frame = TTSTextFrame(word)
|
||||||
frame.pts = self._initial_word_timestamp + timestamp
|
frame.pts = self._initial_word_timestamp + timestamp
|
||||||
if frame:
|
if frame:
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ class WebsocketService(ABC):
|
|||||||
"""
|
"""
|
||||||
self._websocket: Optional[websockets.WebSocketClientProtocol] = None
|
self._websocket: Optional[websockets.WebSocketClientProtocol] = None
|
||||||
self._reconnect_on_error = reconnect_on_error
|
self._reconnect_on_error = reconnect_on_error
|
||||||
|
self._reconnect_in_progress: bool = False # Add this flag
|
||||||
|
|
||||||
async def _verify_connection(self) -> bool:
|
async def _verify_connection(self) -> bool:
|
||||||
"""Verify the websocket connection is active and responsive.
|
"""Verify the websocket connection is active and responsive.
|
||||||
@@ -66,6 +67,59 @@ class WebsocketService(ABC):
|
|||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
return await self._verify_connection()
|
return await self._verify_connection()
|
||||||
|
|
||||||
|
async def _try_reconnect(
|
||||||
|
self,
|
||||||
|
max_retries: int = 3,
|
||||||
|
report_error: Optional[Callable[[ErrorFrame], Awaitable[None]]] = None,
|
||||||
|
) -> bool:
|
||||||
|
# Prevent concurrent reconnection attempts
|
||||||
|
if self._reconnect_in_progress:
|
||||||
|
logger.warning(f"{self} reconnect attempt aborted: already in progress")
|
||||||
|
return False
|
||||||
|
|
||||||
|
self._reconnect_in_progress = True
|
||||||
|
last_exception: Optional[Exception] = None
|
||||||
|
try:
|
||||||
|
for attempt in range(1, max_retries + 1):
|
||||||
|
try:
|
||||||
|
logger.warning(f"{self} reconnecting, attempt {attempt}")
|
||||||
|
if await self._reconnect_websocket(attempt):
|
||||||
|
logger.info(f"{self} reconnected successfully on attempt {attempt}")
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
last_exception = e
|
||||||
|
logger.error(f"{self} reconnection attempt {attempt} failed: {e}")
|
||||||
|
if report_error:
|
||||||
|
await report_error(
|
||||||
|
ErrorFrame(f"{self} reconnection attempt {attempt} failed: {e}")
|
||||||
|
)
|
||||||
|
wait_time = exponential_backoff_time(attempt)
|
||||||
|
await asyncio.sleep(wait_time)
|
||||||
|
fatal_msg = f"{self} failed to reconnect after {max_retries} attempts"
|
||||||
|
if last_exception:
|
||||||
|
fatal_msg += f": {last_exception}"
|
||||||
|
logger.error(fatal_msg)
|
||||||
|
if report_error:
|
||||||
|
await report_error(ErrorFrame(fatal_msg, fatal=True))
|
||||||
|
return False
|
||||||
|
finally:
|
||||||
|
self._reconnect_in_progress = False
|
||||||
|
|
||||||
|
async def send_with_retry(self, message, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
||||||
|
"""Attempt to send a message, retrying after reconnect if necessary."""
|
||||||
|
try:
|
||||||
|
await self._websocket.send(message)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"{self} send failed: {e}, will try to reconnect")
|
||||||
|
# Try to reconnect before retrying
|
||||||
|
success = await self._try_reconnect(report_error=report_error)
|
||||||
|
if success:
|
||||||
|
logger.info(f"{self} reconnected successfully, will retry send the message")
|
||||||
|
# trying to send the message one more time
|
||||||
|
await self._websocket.send(message)
|
||||||
|
else:
|
||||||
|
logger.error(f"{self} send failed; unable to reconnect")
|
||||||
|
|
||||||
async def _receive_task_handler(self, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
async def _receive_task_handler(self, report_error: Callable[[ErrorFrame], Awaitable[None]]):
|
||||||
"""Handle websocket message receiving with automatic retry logic.
|
"""Handle websocket message receiving with automatic retry logic.
|
||||||
|
|
||||||
@@ -76,13 +130,9 @@ class WebsocketService(ABC):
|
|||||||
Args:
|
Args:
|
||||||
report_error: Callback function to report connection errors.
|
report_error: Callback function to report connection errors.
|
||||||
"""
|
"""
|
||||||
retry_count = 0
|
|
||||||
MAX_RETRIES = 3
|
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
await self._receive_messages()
|
await self._receive_messages()
|
||||||
retry_count = 0 # Reset counter on successful message receive
|
|
||||||
except ConnectionClosedOK as e:
|
except ConnectionClosedOK as e:
|
||||||
# Normal closure, don't retry
|
# Normal closure, don't retry
|
||||||
logger.debug(f"{self} connection closed normally: {e}")
|
logger.debug(f"{self} connection closed normally: {e}")
|
||||||
@@ -92,21 +142,9 @@ class WebsocketService(ABC):
|
|||||||
logger.error(message)
|
logger.error(message)
|
||||||
|
|
||||||
if self._reconnect_on_error:
|
if self._reconnect_on_error:
|
||||||
retry_count += 1
|
success = await self._try_reconnect(report_error=report_error)
|
||||||
if retry_count >= MAX_RETRIES:
|
if not success:
|
||||||
await report_error(ErrorFrame(message))
|
|
||||||
break
|
break
|
||||||
|
|
||||||
logger.warning(f"{self} connection error, will retry: {e}")
|
|
||||||
await report_error(ErrorFrame(message))
|
|
||||||
|
|
||||||
try:
|
|
||||||
if await self._reconnect_websocket(retry_count):
|
|
||||||
retry_count = 0 # Reset counter on successful reconnection
|
|
||||||
wait_time = exponential_backoff_time(retry_count)
|
|
||||||
await asyncio.sleep(wait_time)
|
|
||||||
except Exception as reconnect_error:
|
|
||||||
logger.error(f"{self} reconnection failed: {reconnect_error}")
|
|
||||||
else:
|
else:
|
||||||
await report_error(ErrorFrame(message))
|
await report_error(ErrorFrame(message))
|
||||||
break
|
break
|
||||||
|
|||||||
@@ -26,7 +26,6 @@ class BaseTextFilter(ABC):
|
|||||||
behavior, settings management, and interruption handling logic.
|
behavior, settings management, and interruption handling logic.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
async def update_settings(self, settings: Mapping[str, Any]):
|
async def update_settings(self, settings: Mapping[str, Any]):
|
||||||
"""Update the filter's configuration settings.
|
"""Update the filter's configuration settings.
|
||||||
|
|
||||||
@@ -53,7 +52,6 @@ class BaseTextFilter(ABC):
|
|||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
async def handle_interruption(self):
|
async def handle_interruption(self):
|
||||||
"""Handle interruption events in the processing pipeline.
|
"""Handle interruption events in the processing pipeline.
|
||||||
|
|
||||||
@@ -62,7 +60,6 @@ class BaseTextFilter(ABC):
|
|||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
async def reset_interruption(self):
|
async def reset_interruption(self):
|
||||||
"""Reset the filter state after an interruption has been handled.
|
"""Reset the filter state after an interruption has been handled.
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user