Update Together services to use canonical settings pattern
- STT/TTS: Use NOT_GIVEN sentinel, Settings alias, and apply_update() - TTS: Add output_format and encoding params, use audio context management, push_start_frame=True, handle binary + JSON audio - LLM: Update default model to Llama-4-Maverick - Update examples to use new settings API
This commit is contained in:
@@ -21,7 +21,6 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
)
|
)
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.openai.llm import OpenAILLMService
|
|
||||||
from pipecat.services.together.llm import TogetherLLMService
|
from pipecat.services.together.llm import TogetherLLMService
|
||||||
from pipecat.services.together.stt import TogetherSTTService
|
from pipecat.services.together.stt import TogetherSTTService
|
||||||
from pipecat.services.together.tts import TogetherTTSService
|
from pipecat.services.together.tts import TogetherTTSService
|
||||||
@@ -57,19 +56,20 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
|
|
||||||
tts = TogetherTTSService(
|
tts = TogetherTTSService(
|
||||||
api_key=os.getenv("TOGETHER_API_KEY"),
|
api_key=os.getenv("TOGETHER_API_KEY"),
|
||||||
voice="tara",
|
settings=TogetherTTSService.Settings(
|
||||||
|
model="canopylabs/orpheus-3b-0.1-ft",
|
||||||
|
voice="tara",
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
llm = TogetherLLMService(api_key=os.getenv("TOGETHER_API_KEY"))
|
llm = TogetherLLMService(
|
||||||
|
api_key=os.getenv("TOGETHER_API_KEY"),
|
||||||
|
settings=TogetherLLMService.Settings(
|
||||||
|
system_instruction="You are a helpful assistant in a voice conversation. Your responses will be spoken aloud, so avoid emojis, bullet points, or other formatting that can't be spoken. Respond to what the user said in a creative, helpful, and brief way.",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
messages = [
|
context = LLMContext()
|
||||||
{
|
|
||||||
"role": "system",
|
|
||||||
"content": "You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be spoken aloud, so avoid special characters that can't easily be spoken, such as emojis or bullet points. Respond to what the user said in a creative and helpful way.",
|
|
||||||
},
|
|
||||||
]
|
|
||||||
|
|
||||||
context = LLMContext(messages)
|
|
||||||
user_aggregator, assistant_aggregator = LLMContextAggregatorPair(
|
user_aggregator, assistant_aggregator = LLMContextAggregatorPair(
|
||||||
context,
|
context,
|
||||||
user_params=LLMUserAggregatorParams(vad_analyzer=SileroVADAnalyzer()),
|
user_params=LLMUserAggregatorParams(vad_analyzer=SileroVADAnalyzer()),
|
||||||
@@ -100,7 +100,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
async def on_client_connected(transport, client):
|
async def on_client_connected(transport, client):
|
||||||
logger.info(f"Client connected")
|
logger.info(f"Client connected")
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
messages.append({"role": "system", "content": "Please introduce yourself to the user."})
|
context.add_message({"role": "user", "content": "Please introduce yourself to the user."})
|
||||||
await task.queue_frames([LLMRunFrame()])
|
await task.queue_frames([LLMRunFrame()])
|
||||||
|
|
||||||
@transport.event_handler("on_client_disconnected")
|
@transport.event_handler("on_client_disconnected")
|
||||||
|
|||||||
@@ -72,7 +72,6 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
llm = TogetherLLMService(
|
llm = TogetherLLMService(
|
||||||
api_key=os.getenv("TOGETHER_API_KEY"),
|
api_key=os.getenv("TOGETHER_API_KEY"),
|
||||||
settings=TogetherLLMService.Settings(
|
settings=TogetherLLMService.Settings(
|
||||||
model="meta-llama/Meta-Llama-3.1-70B-Instruct-Turbo",
|
|
||||||
system_instruction="You are a helpful assistant in a voice conversation. Your responses will be spoken aloud, so avoid emojis, bullet points, or other formatting that can't be spoken. Respond to what the user said in a creative, helpful, and brief way.",
|
system_instruction="You are a helpful assistant in a voice conversation. Your responses will be spoken aloud, so avoid emojis, bullet points, or other formatting that can't be spoken. Respond to what the user said in a creative, helpful, and brief way.",
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -108,6 +108,7 @@ TESTS_07 = [
|
|||||||
("07c-interruptible-deepgram-http.py", EVAL_SIMPLE_MATH),
|
("07c-interruptible-deepgram-http.py", EVAL_SIMPLE_MATH),
|
||||||
("07d-interruptible-elevenlabs.py", EVAL_SIMPLE_MATH),
|
("07d-interruptible-elevenlabs.py", EVAL_SIMPLE_MATH),
|
||||||
("07d-interruptible-elevenlabs-http.py", EVAL_SIMPLE_MATH),
|
("07d-interruptible-elevenlabs-http.py", EVAL_SIMPLE_MATH),
|
||||||
|
("07e-interruptible-together.py", EVAL_SIMPLE_MATH),
|
||||||
("07f-interruptible-azure.py", EVAL_SIMPLE_MATH),
|
("07f-interruptible-azure.py", EVAL_SIMPLE_MATH),
|
||||||
("07f-interruptible-azure-http.py", EVAL_SIMPLE_MATH),
|
("07f-interruptible-azure-http.py", EVAL_SIMPLE_MATH),
|
||||||
("07g-interruptible-openai.py", EVAL_SIMPLE_MATH),
|
("07g-interruptible-openai.py", EVAL_SIMPLE_MATH),
|
||||||
@@ -167,6 +168,7 @@ TESTS_14 = [
|
|||||||
("14-function-calling.py", EVAL_WEATHER_AND_RESTAURANT),
|
("14-function-calling.py", EVAL_WEATHER_AND_RESTAURANT),
|
||||||
("14a-function-calling-anthropic.py", EVAL_WEATHER),
|
("14a-function-calling-anthropic.py", EVAL_WEATHER),
|
||||||
("14a-function-calling-anthropic.py", EVAL_WEATHER_AND_RESTAURANT),
|
("14a-function-calling-anthropic.py", EVAL_WEATHER_AND_RESTAURANT),
|
||||||
|
("14c-function-calling-together.py", EVAL_WEATHER),
|
||||||
("14e-function-calling-google.py", EVAL_WEATHER),
|
("14e-function-calling-google.py", EVAL_WEATHER),
|
||||||
("14e-function-calling-google.py", EVAL_WEATHER_AND_RESTAURANT),
|
("14e-function-calling-google.py", EVAL_WEATHER_AND_RESTAURANT),
|
||||||
("14f-function-calling-groq.py", EVAL_WEATHER),
|
("14f-function-calling-groq.py", EVAL_WEATHER),
|
||||||
|
|||||||
@@ -46,7 +46,7 @@ class TogetherLLMService(OpenAILLMService):
|
|||||||
Args:
|
Args:
|
||||||
api_key: The API key for accessing Together.ai's API.
|
api_key: The API key for accessing Together.ai's API.
|
||||||
base_url: The base URL for Together.ai API. Defaults to "https://api.together.xyz/v1".
|
base_url: The base URL for Together.ai API. Defaults to "https://api.together.xyz/v1".
|
||||||
model: The model identifier to use. Defaults to "meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo".
|
model: The model identifier to use. Defaults to "meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8".
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=TogetherLLMService.Settings(model=...)`` instead.
|
Use ``settings=TogetherLLMService.Settings(model=...)`` instead.
|
||||||
@@ -56,7 +56,7 @@ class TogetherLLMService(OpenAILLMService):
|
|||||||
**kwargs: Additional keyword arguments passed to OpenAILLMService.
|
**kwargs: Additional keyword arguments passed to OpenAILLMService.
|
||||||
"""
|
"""
|
||||||
# 1. Initialize default_settings with hardcoded defaults
|
# 1. Initialize default_settings with hardcoded defaults
|
||||||
default_settings = self.Settings(model="meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo")
|
default_settings = self.Settings(model="meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8")
|
||||||
|
|
||||||
# 2. Apply direct init arg overrides (deprecated)
|
# 2. Apply direct init arg overrides (deprecated)
|
||||||
if model is not None:
|
if model is not None:
|
||||||
|
|||||||
@@ -45,13 +45,15 @@ from pipecat.utils.tracing.service_decorators import traced_stt
|
|||||||
class TogetherSTTSettings(STTSettings):
|
class TogetherSTTSettings(STTSettings):
|
||||||
"""Settings for the Together AI STT service.
|
"""Settings for the Together AI STT service.
|
||||||
|
|
||||||
|
``model`` and ``language`` are inherited from ``STTSettings`` /
|
||||||
|
``ServiceSettings``.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
model: Together AI transcription model to use.
|
model: Together AI transcription model to use.
|
||||||
language: Language of the audio input.
|
language: Language of the audio input.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
model: str = "openai/whisper-large-v3"
|
pass
|
||||||
language: Language = Language.EN
|
|
||||||
|
|
||||||
|
|
||||||
class TogetherSTTService(WebsocketSTTService):
|
class TogetherSTTService(WebsocketSTTService):
|
||||||
@@ -61,38 +63,44 @@ class TogetherSTTService(WebsocketSTTService):
|
|||||||
with OpenAI-compatible speech-to-text endpoints.
|
with OpenAI-compatible speech-to-text endpoints.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
_settings: TogetherSTTSettings
|
Settings = TogetherSTTSettings
|
||||||
|
_settings: Settings
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
model: str = "openai/whisper-large-v3",
|
|
||||||
language: Language = Language.EN,
|
|
||||||
sample_rate: int = 16000,
|
sample_rate: int = 16000,
|
||||||
base_url: str = "wss://api.together.xyz/v1",
|
base_url: str = "wss://api.together.xyz/v1",
|
||||||
ttfs_p99_latency: float = TOGETHER_TTFS_P99,
|
ttfs_p99_latency: float = TOGETHER_TTFS_P99,
|
||||||
|
settings: Optional[Settings] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""Initialize the Together AI STT service.
|
"""Initialize the Together AI STT service.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
api_key: Together AI API key for authentication.
|
api_key: Together AI API key for authentication.
|
||||||
model: Together AI transcription model. Defaults to "openai/whisper-large-v3".
|
|
||||||
language: Language of the audio input. Defaults to English.
|
|
||||||
sample_rate: Audio sample rate (default: 16000). Together AI requires 16kHz input.
|
sample_rate: Audio sample rate (default: 16000). Together AI requires 16kHz input.
|
||||||
base_url: The URL of the Together AI WebSocket API.
|
base_url: The URL of the Together AI WebSocket API.
|
||||||
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
||||||
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
||||||
|
settings: Runtime-updatable settings. Allows overriding model and language.
|
||||||
**kwargs: Additional arguments passed to the parent WebsocketSTTService.
|
**kwargs: Additional arguments passed to the parent WebsocketSTTService.
|
||||||
"""
|
"""
|
||||||
|
# 1. Initialize default_settings with hardcoded defaults
|
||||||
|
default_settings = self.Settings(
|
||||||
|
model="openai/whisper-large-v3",
|
||||||
|
language=Language.EN,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Apply settings delta (canonical API, always wins)
|
||||||
|
if settings is not None:
|
||||||
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
sample_rate=sample_rate,
|
sample_rate=sample_rate,
|
||||||
ttfs_p99_latency=ttfs_p99_latency,
|
ttfs_p99_latency=ttfs_p99_latency,
|
||||||
settings=TogetherSTTSettings(
|
settings=default_settings,
|
||||||
model=model,
|
|
||||||
language=language,
|
|
||||||
),
|
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -8,12 +8,12 @@
|
|||||||
|
|
||||||
import base64
|
import base64
|
||||||
import json
|
import json
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, field
|
||||||
from typing import AsyncGenerator, Optional
|
from typing import AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from pipecat.services.settings import TTSSettings
|
from pipecat.services.settings import NOT_GIVEN, TTSSettings, _NotGiven
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from websockets.asyncio.client import connect as websocket_connect
|
from websockets.asyncio.client import connect as websocket_connect
|
||||||
@@ -26,14 +26,12 @@ except ModuleNotFoundError as e:
|
|||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
InterruptionFrame,
|
|
||||||
StartFrame,
|
StartFrame,
|
||||||
TTSAudioRawFrame,
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
|
||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
|
||||||
from pipecat.services.tts_service import WebsocketTTSService
|
from pipecat.services.tts_service import WebsocketTTSService
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.utils.tracing.service_decorators import traced_tts
|
from pipecat.utils.tracing.service_decorators import traced_tts
|
||||||
@@ -43,6 +41,9 @@ from pipecat.utils.tracing.service_decorators import traced_tts
|
|||||||
class TogetherTTSSettings(TTSSettings):
|
class TogetherTTSSettings(TTSSettings):
|
||||||
"""Settings for the Together AI TTS service.
|
"""Settings for the Together AI TTS service.
|
||||||
|
|
||||||
|
``model``, ``voice``, and ``language`` are inherited from ``TTSSettings`` /
|
||||||
|
``ServiceSettings``.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
model: Together AI TTS model to use.
|
model: Together AI TTS model to use.
|
||||||
voice: Voice to use for synthesis.
|
voice: Voice to use for synthesis.
|
||||||
@@ -50,10 +51,7 @@ class TogetherTTSSettings(TTSSettings):
|
|||||||
max_partial_length: Maximum partial text length for streaming.
|
max_partial_length: Maximum partial text length for streaming.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
model: str = "canopylabs/orpheus-3b-0.1-ft"
|
max_partial_length: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
language: Optional[Language] = Language.EN
|
|
||||||
voice: Optional[str] = "tara"
|
|
||||||
max_partial_length: Optional[int] = None
|
|
||||||
|
|
||||||
|
|
||||||
class TogetherTTSService(WebsocketTTSService):
|
class TogetherTTSService(WebsocketTTSService):
|
||||||
@@ -63,45 +61,60 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
Supports streaming synthesis with configurable voice and model options.
|
Supports streaming synthesis with configurable voice and model options.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
_settings: TogetherTTSSettings
|
Settings = TogetherTTSSettings
|
||||||
|
_settings: Settings
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
model: str = "canopylabs/orpheus-3b-0.1-ft",
|
output_format: str = "raw",
|
||||||
voice: str = "tara",
|
encoding: str = "pcm_s16le",
|
||||||
language: Optional[Language] = Language.EN,
|
|
||||||
max_partial_length: Optional[int] = None,
|
|
||||||
url: str = "wss://api.together.ai/v1/audio/speech/websocket",
|
url: str = "wss://api.together.ai/v1/audio/speech/websocket",
|
||||||
sample_rate: Optional[int] = 24000,
|
sample_rate: Optional[int] = 24000,
|
||||||
|
settings: Optional[Settings] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""Initialize the Together AI TTS service.
|
"""Initialize the Together AI TTS service.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
api_key: Together AI API key for authentication.
|
api_key: Together AI API key for authentication.
|
||||||
model: Together AI TTS model. Defaults to "canopylabs/orpheus-3b-0.1-ft".
|
output_format: Audio output container format. Supported values:
|
||||||
voice: Voice to use for synthesis. Defaults to "tara".
|
``"raw"``, ``"mp3"``, ``"wav"``, ``"opus"``, ``"aac"``,
|
||||||
language: Language of the text input. Defaults to English.
|
``"flac"``, ``"pcm"``. Defaults to ``"raw"``.
|
||||||
max_partial_length: Maximum partial text length for streaming.
|
encoding: PCM encoding when ``output_format`` is ``"raw"``.
|
||||||
|
Supported values: ``"pcm_s16le"``, ``"pcm_f32le"``.
|
||||||
|
Defaults to ``"pcm_s16le"``.
|
||||||
url: WebSocket URL for Together AI TTS API.
|
url: WebSocket URL for Together AI TTS API.
|
||||||
sample_rate: Audio sample rate (default: 24000).
|
sample_rate: Audio sample rate (default: 24000).
|
||||||
|
settings: Runtime-updatable settings. Allows overriding model, voice,
|
||||||
|
language, and max_partial_length.
|
||||||
**kwargs: Additional arguments passed to the parent service.
|
**kwargs: Additional arguments passed to the parent service.
|
||||||
"""
|
"""
|
||||||
|
# 1. Initialize default_settings with hardcoded defaults
|
||||||
|
default_settings = self.Settings(
|
||||||
|
model="canopylabs/orpheus-3b-0.1-ft",
|
||||||
|
voice="tara",
|
||||||
|
language=Language.EN,
|
||||||
|
max_partial_length=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Apply settings delta (canonical API, always wins)
|
||||||
|
if settings is not None:
|
||||||
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
|
push_text_frames=False,
|
||||||
|
push_start_frame=True,
|
||||||
sample_rate=sample_rate,
|
sample_rate=sample_rate,
|
||||||
settings=TogetherTTSSettings(
|
settings=default_settings,
|
||||||
model=model,
|
|
||||||
voice=voice,
|
|
||||||
language=language,
|
|
||||||
max_partial_length=max_partial_length,
|
|
||||||
),
|
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._url = url
|
self._url = url
|
||||||
|
self._output_format = output_format
|
||||||
|
self._encoding = encoding
|
||||||
self._session_id = None
|
self._session_id = None
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
self._context_id: Optional[str] = None
|
self._context_id: Optional[str] = None
|
||||||
@@ -116,7 +129,12 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
|
|
||||||
def _build_websocket_url(self) -> str:
|
def _build_websocket_url(self) -> str:
|
||||||
"""Build the WebSocket URL with query parameters."""
|
"""Build the WebSocket URL with query parameters."""
|
||||||
url = f"{self._url}?model={self._settings.model}&voice={self._settings.voice}"
|
url = (
|
||||||
|
f"{self._url}?model={self._settings.model}"
|
||||||
|
f"&voice={self._settings.voice}"
|
||||||
|
f"&response_format={self._output_format}"
|
||||||
|
f"&response_encoding={self._encoding}"
|
||||||
|
)
|
||||||
if self._settings.max_partial_length is not None:
|
if self._settings.max_partial_length is not None:
|
||||||
url += f"&max_partial_length={self._settings.max_partial_length}"
|
url += f"&max_partial_length={self._settings.max_partial_length}"
|
||||||
return url
|
return url
|
||||||
@@ -211,6 +229,7 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
exception=e,
|
exception=e,
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
|
await self.remove_active_audio_context()
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
self._session_id = None
|
self._session_id = None
|
||||||
await self._call_event_handler("on_disconnected")
|
await self._call_event_handler("on_disconnected")
|
||||||
@@ -236,8 +255,16 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
exception=e,
|
exception=e,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def flush_audio(self):
|
async def flush_audio(self, context_id: Optional[str] = None):
|
||||||
"""Flush any pending audio and finalize the current context."""
|
"""Flush any pending audio and finalize the current context.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
context_id: The specific context to flush. If None, falls back to the
|
||||||
|
currently active context.
|
||||||
|
"""
|
||||||
|
flush_id = context_id or self.get_active_audio_context_id()
|
||||||
|
if not flush_id or not self._websocket:
|
||||||
|
return
|
||||||
logger.trace(f"{self}: flushing audio")
|
logger.trace(f"{self}: flushing audio")
|
||||||
await self._ws_send({"type": "input_text_buffer.commit"})
|
await self._ws_send({"type": "input_text_buffer.commit"})
|
||||||
|
|
||||||
@@ -252,7 +279,10 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
this method with automatic reconnection on connection errors.
|
this method with automatic reconnection on connection errors.
|
||||||
"""
|
"""
|
||||||
async for message in self._websocket:
|
async for message in self._websocket:
|
||||||
|
# Together sends audio as binary WebSocket frames (raw PCM)
|
||||||
|
# and control/status messages as JSON text frames.
|
||||||
if not isinstance(message, str):
|
if not isinstance(message, str):
|
||||||
|
await self._handle_audio_binary(message)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -268,8 +298,7 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
elif evt_type == "session.updated":
|
elif evt_type == "session.updated":
|
||||||
await self._handle_session_updated(evt)
|
await self._handle_session_updated(evt)
|
||||||
elif evt_type == "conversation.item.input_text.received":
|
elif evt_type == "conversation.item.input_text.received":
|
||||||
text = evt.get("text", "")
|
logger.trace(f"{self} text received")
|
||||||
logger.debug(f"{self} text received: {text[:50]}{'...' if len(text) > 50 else ''}")
|
|
||||||
elif evt_type == "conversation.item.audio_output.delta":
|
elif evt_type == "conversation.item.audio_output.delta":
|
||||||
await self._handle_audio_delta(evt)
|
await self._handle_audio_delta(evt)
|
||||||
elif evt_type == "conversation.item.audio_output.done":
|
elif evt_type == "conversation.item.audio_output.done":
|
||||||
@@ -302,26 +331,48 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
updated_voice = session.get("voice")
|
updated_voice = session.get("voice")
|
||||||
logger.debug(f"{self} voice updated to: {updated_voice}")
|
logger.debug(f"{self} voice updated to: {updated_voice}")
|
||||||
|
|
||||||
|
async def _handle_audio_binary(self, data: bytes):
|
||||||
|
"""Handle a binary WebSocket frame containing raw PCM audio.
|
||||||
|
|
||||||
|
Together sends audio as binary frames and control messages as JSON text.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
data: Raw PCM audio bytes.
|
||||||
|
"""
|
||||||
|
context_id = self._context_id
|
||||||
|
if not context_id or not self.audio_context_available(context_id):
|
||||||
|
return
|
||||||
|
|
||||||
|
await self.stop_ttfb_metrics()
|
||||||
|
frame = TTSAudioRawFrame(
|
||||||
|
audio=data,
|
||||||
|
sample_rate=self.sample_rate,
|
||||||
|
num_channels=1,
|
||||||
|
context_id=context_id,
|
||||||
|
)
|
||||||
|
await self.append_to_audio_context(context_id, frame)
|
||||||
|
|
||||||
async def _handle_audio_delta(self, evt: dict):
|
async def _handle_audio_delta(self, evt: dict):
|
||||||
"""Handle an audio output delta containing a chunk of synthesized audio.
|
"""Handle a JSON audio output delta containing base64-encoded audio.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
evt: The delta event from the server.
|
evt: The delta event from the server.
|
||||||
"""
|
"""
|
||||||
|
context_id = self._context_id
|
||||||
|
if not context_id or not self.audio_context_available(context_id):
|
||||||
|
return
|
||||||
|
|
||||||
delta = evt.get("delta")
|
delta = evt.get("delta")
|
||||||
if delta:
|
if delta:
|
||||||
try:
|
await self.stop_ttfb_metrics()
|
||||||
await self.stop_ttfb_metrics()
|
audio_chunk = base64.b64decode(delta)
|
||||||
audio_chunk = base64.b64decode(delta)
|
frame = TTSAudioRawFrame(
|
||||||
frame = TTSAudioRawFrame(
|
audio=audio_chunk,
|
||||||
audio=audio_chunk,
|
sample_rate=self.sample_rate,
|
||||||
sample_rate=self.sample_rate,
|
num_channels=1,
|
||||||
num_channels=1,
|
context_id=context_id,
|
||||||
context_id=self._context_id,
|
)
|
||||||
)
|
await self.append_to_audio_context(context_id, frame)
|
||||||
await self.push_frame(frame)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{self} error processing audio delta: {e}")
|
|
||||||
|
|
||||||
async def _handle_audio_done(self, evt: dict):
|
async def _handle_audio_done(self, evt: dict):
|
||||||
"""Handle audio output completion for a speech segment.
|
"""Handle audio output completion for a speech segment.
|
||||||
@@ -329,9 +380,12 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
Args:
|
Args:
|
||||||
evt: The done event from the server.
|
evt: The done event from the server.
|
||||||
"""
|
"""
|
||||||
|
context_id = self._context_id
|
||||||
item_id = evt.get("item_id")
|
item_id = evt.get("item_id")
|
||||||
logger.debug(f"{self} audio generation complete for: {item_id}")
|
logger.debug(f"{self} audio generation complete for: {item_id}")
|
||||||
await self.push_frame(TTSStoppedFrame(context_id=self._context_id))
|
if context_id and self.audio_context_available(context_id):
|
||||||
|
await self.add_word_timestamps([("TTSStoppedFrame", 0), ("Reset", 0)], context_id)
|
||||||
|
await self.remove_audio_context(context_id)
|
||||||
|
|
||||||
async def _handle_tts_failed(self, evt: dict):
|
async def _handle_tts_failed(self, evt: dict):
|
||||||
"""Handle a TTS failure.
|
"""Handle a TTS failure.
|
||||||
@@ -339,9 +393,12 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
Args:
|
Args:
|
||||||
evt: The failed event containing error details.
|
evt: The failed event containing error details.
|
||||||
"""
|
"""
|
||||||
|
context_id = self._context_id
|
||||||
error = evt.get("error", {})
|
error = evt.get("error", {})
|
||||||
await self.push_error(error_msg=f"TTS error: {error}")
|
await self.push_error(error_msg=f"TTS error: {error}")
|
||||||
await self.push_frame(TTSStoppedFrame(context_id=self._context_id))
|
await self.push_frame(TTSStoppedFrame(context_id=context_id))
|
||||||
|
await self.stop_all_metrics()
|
||||||
|
self.reset_active_audio_context()
|
||||||
|
|
||||||
async def _handle_error(self, evt: dict):
|
async def _handle_error(self, evt: dict):
|
||||||
"""Handle a fatal error from the TTS session.
|
"""Handle a fatal error from the TTS session.
|
||||||
@@ -359,20 +416,29 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
await self.push_error(error_msg=msg)
|
await self.push_error(error_msg=msg)
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Interruption handling
|
# Audio context lifecycle
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
async def _handle_interruption(self, frame: InterruptionFrame, direction: FrameDirection):
|
async def on_audio_context_interrupted(self, context_id: str):
|
||||||
"""Handle interruption by canceling current generation.
|
"""Cancel the active Together context when the bot is interrupted.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
frame: The interruption frame.
|
context_id: The context that was interrupted.
|
||||||
direction: Frame processing direction.
|
|
||||||
"""
|
"""
|
||||||
await super()._handle_interruption(frame, direction)
|
|
||||||
await self.stop_all_metrics()
|
await self.stop_all_metrics()
|
||||||
await self._ws_send({"type": "input_text_buffer.clear"})
|
await self._ws_send({"type": "input_text_buffer.clear"})
|
||||||
|
|
||||||
|
async def on_audio_context_completed(self, context_id: str):
|
||||||
|
"""Handle context completion after all audio has been played.
|
||||||
|
|
||||||
|
No server-side cleanup is needed — the Together AI server considers
|
||||||
|
the context done once it has sent its ``audio_output.done`` message.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
context_id: The context that completed playback.
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# TTS generation
|
# TTS generation
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@@ -381,6 +447,10 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]:
|
||||||
"""Generate speech from text using Together AI's streaming API.
|
"""Generate speech from text using Together AI's streaming API.
|
||||||
|
|
||||||
|
Audio context creation and ``TTSStartedFrame`` are managed by the base
|
||||||
|
class (``push_start_frame=True``). This method only sends the text to
|
||||||
|
the WebSocket; audio frames arrive via ``_receive_messages``.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
text: The text to synthesize into speech.
|
text: The text to synthesize into speech.
|
||||||
context_id: The context ID for tracking audio frames.
|
context_id: The context ID for tracking audio frames.
|
||||||
@@ -390,25 +460,18 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
"""
|
"""
|
||||||
logger.debug(f"{self}: Generating TTS [{text}]")
|
logger.debug(f"{self}: Generating TTS [{text}]")
|
||||||
|
|
||||||
|
self._context_id = context_id
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if not self._websocket or self._websocket.state is not State.OPEN:
|
if not self._websocket or self._websocket.state is not State.OPEN:
|
||||||
await self._connect()
|
await self._connect()
|
||||||
if not self._websocket or self._websocket.state is not State.OPEN:
|
|
||||||
logger.error(f"{self} failed to connect to WebSocket")
|
|
||||||
yield TTSStoppedFrame(context_id=context_id)
|
|
||||||
return
|
|
||||||
|
|
||||||
self._context_id = context_id
|
|
||||||
|
|
||||||
await self.start_ttfb_metrics()
|
|
||||||
yield TTSStartedFrame(context_id=context_id)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await self._ws_send({"type": "input_text_buffer.append", "text": text})
|
await self._ws_send({"type": "input_text_buffer.append", "text": text})
|
||||||
await self._ws_send({"type": "input_text_buffer.commit"})
|
await self._ws_send({"type": "input_text_buffer.commit"})
|
||||||
await self.start_tts_usage_metrics(text)
|
await self.start_tts_usage_metrics(text)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} error sending message: {e}")
|
yield ErrorFrame(error=f"Error sending TTS text: {e}")
|
||||||
yield TTSStoppedFrame(context_id=context_id)
|
yield TTSStoppedFrame(context_id=context_id)
|
||||||
await self._disconnect()
|
await self._disconnect()
|
||||||
await self._connect()
|
await self._connect()
|
||||||
@@ -417,6 +480,5 @@ class TogetherTTSService(WebsocketTTSService):
|
|||||||
yield None
|
yield None
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} exception: {e}")
|
yield ErrorFrame(error=f"Error generating TTS: {e}")
|
||||||
await self.push_error(error_msg=f"Error generating TTS: {e}", exception=e)
|
|
||||||
yield TTSStoppedFrame(context_id=context_id)
|
yield TTSStoppedFrame(context_id=context_id)
|
||||||
|
|||||||
Reference in New Issue
Block a user