Add xAI HTTP TTS service
Reworks the xAI TTS integration from #4031 with consistency fixes: - Rename to XAIHttpTTSService (leaves room for future WebSocket service) - Add proper language map with all 20 supported xAI languages - Remove unnecessary deprecated InputParams/params (new service, nothing to deprecate) - Add encoding as a constructor parameter - Use Language.EN enum instead of string for default language - Linting fixes
This commit is contained in:
1
changelog/4031.added.md
Normal file
1
changelog/4031.added.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Added `XAIHttpTTSService` for text-to-speech using xAI's HTTP TTS API.
|
||||||
129
examples/foundational/07e-interruptible-xai.py
Normal file
129
examples/foundational/07e-interruptible-xai.py
Normal file
@@ -0,0 +1,129 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024-2026, Daily
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
|
#
|
||||||
|
|
||||||
|
import os
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
|
from pipecat.frames.frames import LLMRunFrame
|
||||||
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
|
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
||||||
|
from pipecat.processors.aggregators.llm_context import LLMContext
|
||||||
|
from pipecat.processors.aggregators.llm_response_universal import (
|
||||||
|
LLMContextAggregatorPair,
|
||||||
|
LLMUserAggregatorParams,
|
||||||
|
)
|
||||||
|
from pipecat.runner.types import RunnerArguments
|
||||||
|
from pipecat.runner.utils import create_transport
|
||||||
|
from pipecat.services.deepgram.stt import DeepgramSTTService
|
||||||
|
from pipecat.services.grok.llm import GrokLLMService
|
||||||
|
from pipecat.services.xai.tts import XAIHttpTTSService
|
||||||
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
|
from pipecat.transports.daily.transport import DailyParams
|
||||||
|
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
||||||
|
|
||||||
|
load_dotenv(override=True)
|
||||||
|
|
||||||
|
# We use lambdas to defer transport parameter creation until the transport
|
||||||
|
# type is selected at runtime.
|
||||||
|
transport_params = {
|
||||||
|
"daily": lambda: DailyParams(
|
||||||
|
audio_in_enabled=True,
|
||||||
|
audio_out_enabled=True,
|
||||||
|
),
|
||||||
|
"twilio": lambda: FastAPIWebsocketParams(
|
||||||
|
audio_in_enabled=True,
|
||||||
|
audio_out_enabled=True,
|
||||||
|
),
|
||||||
|
"webrtc": lambda: TransportParams(
|
||||||
|
audio_in_enabled=True,
|
||||||
|
audio_out_enabled=True,
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
||||||
|
logger.info(f"Starting bot")
|
||||||
|
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
||||||
|
|
||||||
|
tts = XAIHttpTTSService(
|
||||||
|
api_key=os.getenv("GROK_API_KEY"),
|
||||||
|
aiohttp_session=session,
|
||||||
|
settings=XAIHttpTTSService.Settings(
|
||||||
|
voice="eve",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
llm = GrokLLMService(
|
||||||
|
api_key=os.getenv("GROK_API_KEY"),
|
||||||
|
settings=GrokLLMService.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.",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
context = LLMContext()
|
||||||
|
user_aggregator, assistant_aggregator = LLMContextAggregatorPair(
|
||||||
|
context,
|
||||||
|
user_params=LLMUserAggregatorParams(vad_analyzer=SileroVADAnalyzer()),
|
||||||
|
)
|
||||||
|
|
||||||
|
pipeline = Pipeline(
|
||||||
|
[
|
||||||
|
transport.input(), # Transport user input
|
||||||
|
stt,
|
||||||
|
user_aggregator, # User responses
|
||||||
|
llm, # LLM
|
||||||
|
tts, # TTS
|
||||||
|
transport.output(), # Transport bot output
|
||||||
|
assistant_aggregator, # Assistant spoken responses
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
task = PipelineTask(
|
||||||
|
pipeline,
|
||||||
|
params=PipelineParams(
|
||||||
|
enable_metrics=True,
|
||||||
|
enable_usage_metrics=True,
|
||||||
|
audio_out_sample_rate=8000,
|
||||||
|
),
|
||||||
|
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
||||||
|
)
|
||||||
|
|
||||||
|
@transport.event_handler("on_client_connected")
|
||||||
|
async def on_client_connected(transport, client):
|
||||||
|
logger.info(f"Client connected")
|
||||||
|
# Kick off the conversation.
|
||||||
|
context.add_message(
|
||||||
|
{"role": "user", "content": "Please introduce yourself to the user."}
|
||||||
|
)
|
||||||
|
await task.queue_frames([LLMRunFrame()])
|
||||||
|
|
||||||
|
@transport.event_handler("on_client_disconnected")
|
||||||
|
async def on_client_disconnected(transport, client):
|
||||||
|
logger.info(f"Client disconnected")
|
||||||
|
await task.cancel()
|
||||||
|
|
||||||
|
runner = PipelineRunner(handle_sigint=runner_args.handle_sigint)
|
||||||
|
|
||||||
|
await runner.run(task)
|
||||||
|
|
||||||
|
|
||||||
|
async def bot(runner_args: RunnerArguments):
|
||||||
|
"""Main bot entry point compatible with Pipecat Cloud."""
|
||||||
|
transport = await create_transport(runner_args, transport_params)
|
||||||
|
await run_bot(transport, runner_args)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
from pipecat.runner.run import main
|
||||||
|
|
||||||
|
main()
|
||||||
@@ -7,6 +7,7 @@
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -24,10 +25,10 @@ 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.cartesia.tts import CartesiaTTSService
|
|
||||||
from pipecat.services.deepgram.stt import DeepgramSTTService
|
from pipecat.services.deepgram.stt import DeepgramSTTService
|
||||||
from pipecat.services.grok.llm import GrokLLMService
|
from pipecat.services.grok.llm import GrokLLMService
|
||||||
from pipecat.services.llm_service import FunctionCallParams
|
from pipecat.services.llm_service import FunctionCallParams
|
||||||
|
from pipecat.services.xai.tts import XAIHttpTTSService
|
||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
from pipecat.transports.daily.transport import DailyParams
|
from pipecat.transports.daily.transport import DailyParams
|
||||||
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
||||||
@@ -60,83 +61,88 @@ 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 = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
async with aiohttp.ClientSession() as session:
|
||||||
|
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
||||||
|
|
||||||
tts = CartesiaTTSService(
|
tts = XAIHttpTTSService(
|
||||||
api_key=os.getenv("CARTESIA_API_KEY"),
|
api_key=os.getenv("GROK_API_KEY"),
|
||||||
settings=CartesiaTTSService.Settings(
|
aiohttp_session=session,
|
||||||
voice="71a7ad14-091c-4e8e-a314-022ece01c121", # British Reading Lady
|
settings=XAIHttpTTSService.Settings(
|
||||||
),
|
voice="eve",
|
||||||
)
|
),
|
||||||
|
)
|
||||||
|
|
||||||
llm = GrokLLMService(
|
llm = GrokLLMService(
|
||||||
api_key=os.getenv("GROK_API_KEY"),
|
api_key=os.getenv("GROK_API_KEY"),
|
||||||
settings=GrokLLMService.Settings(
|
settings=GrokLLMService.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.",
|
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.",
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
# You can also register a function_name of None to get all functions
|
# You can also register a function_name of None to get all functions
|
||||||
# sent to the same callback with an additional function_name parameter.
|
# sent to the same callback with an additional function_name parameter.
|
||||||
llm.register_function("get_current_weather", fetch_weather_from_api)
|
llm.register_function("get_current_weather", fetch_weather_from_api)
|
||||||
|
|
||||||
weather_function = FunctionSchema(
|
weather_function = FunctionSchema(
|
||||||
name="get_current_weather",
|
name="get_current_weather",
|
||||||
description="Get the current weather",
|
description="Get the current weather",
|
||||||
properties={
|
properties={
|
||||||
"location": {
|
"location": {
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "The city and state, e.g. San Francisco, CA",
|
"description": "The city and state, e.g. San Francisco, CA",
|
||||||
|
},
|
||||||
|
"format": {
|
||||||
|
"type": "string",
|
||||||
|
"enum": ["celsius", "fahrenheit"],
|
||||||
|
"description": "The temperature unit to use. Infer this from the user's location.",
|
||||||
|
},
|
||||||
},
|
},
|
||||||
"format": {
|
required=["location", "format"],
|
||||||
"type": "string",
|
)
|
||||||
"enum": ["celsius", "fahrenheit"],
|
tools = ToolsSchema(standard_tools=[weather_function])
|
||||||
"description": "The temperature unit to use. Infer this from the user's location.",
|
context = LLMContext(tools=tools)
|
||||||
},
|
user_aggregator, assistant_aggregator = LLMContextAggregatorPair(
|
||||||
},
|
context,
|
||||||
required=["location", "format"],
|
user_params=LLMUserAggregatorParams(vad_analyzer=SileroVADAnalyzer()),
|
||||||
)
|
)
|
||||||
tools = ToolsSchema(standard_tools=[weather_function])
|
|
||||||
context = LLMContext(tools=tools)
|
|
||||||
user_aggregator, assistant_aggregator = LLMContextAggregatorPair(
|
|
||||||
context,
|
|
||||||
user_params=LLMUserAggregatorParams(vad_analyzer=SileroVADAnalyzer()),
|
|
||||||
)
|
|
||||||
|
|
||||||
pipeline = Pipeline(
|
pipeline = Pipeline(
|
||||||
[
|
[
|
||||||
transport.input(),
|
transport.input(),
|
||||||
stt,
|
stt,
|
||||||
user_aggregator,
|
user_aggregator,
|
||||||
llm,
|
llm,
|
||||||
tts,
|
tts,
|
||||||
transport.output(),
|
transport.output(),
|
||||||
assistant_aggregator,
|
assistant_aggregator,
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
task = PipelineTask(
|
task = PipelineTask(
|
||||||
pipeline,
|
pipeline,
|
||||||
params=PipelineParams(
|
params=PipelineParams(
|
||||||
enable_metrics=True,
|
enable_metrics=True,
|
||||||
enable_usage_metrics=True,
|
enable_usage_metrics=True,
|
||||||
),
|
),
|
||||||
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
idle_timeout_secs=runner_args.pipeline_idle_timeout_secs,
|
||||||
)
|
)
|
||||||
|
|
||||||
@transport.event_handler("on_client_connected")
|
@transport.event_handler("on_client_connected")
|
||||||
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.
|
||||||
await task.queue_frames([LLMRunFrame()])
|
context.add_message(
|
||||||
|
{"role": "user", "content": "Please introduce yourself to the user."}
|
||||||
|
)
|
||||||
|
await task.queue_frames([LLMRunFrame()])
|
||||||
|
|
||||||
@transport.event_handler("on_client_disconnected")
|
@transport.event_handler("on_client_disconnected")
|
||||||
async def on_client_disconnected(transport, client):
|
async def on_client_disconnected(transport, client):
|
||||||
logger.info(f"Client disconnected")
|
logger.info(f"Client disconnected")
|
||||||
await task.cancel()
|
await task.cancel()
|
||||||
|
|
||||||
runner = PipelineRunner(handle_sigint=runner_args.handle_sigint)
|
runner = PipelineRunner(handle_sigint=runner_args.handle_sigint)
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|
||||||
|
|
||||||
async def bot(runner_args: RunnerArguments):
|
async def bot(runner_args: RunnerArguments):
|
||||||
|
|||||||
@@ -15,23 +15,60 @@ from typing import AsyncGenerator, Optional
|
|||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from pydantic import BaseModel
|
|
||||||
|
|
||||||
from pipecat.frames.frames import ErrorFrame, Frame, TTSAudioRawFrame
|
from pipecat.frames.frames import ErrorFrame, Frame, TTSAudioRawFrame
|
||||||
from pipecat.services.settings import TTSSettings
|
from pipecat.services.settings import TTSSettings
|
||||||
from pipecat.services.tts_service import TTSService
|
from pipecat.services.tts_service import TTSService
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language, resolve_language
|
||||||
from pipecat.utils.tracing.service_decorators import traced_tts
|
from pipecat.utils.tracing.service_decorators import traced_tts
|
||||||
|
|
||||||
|
|
||||||
|
def language_to_xai_language(language: Language) -> Optional[str]:
|
||||||
|
"""Convert a Language enum to xAI language code.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
language: The Language enum value to convert.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The corresponding xAI language code, or None if not supported.
|
||||||
|
"""
|
||||||
|
LANGUAGE_MAP = {
|
||||||
|
Language.AR: "ar-EG",
|
||||||
|
Language.AR_EG: "ar-EG",
|
||||||
|
Language.AR_SA: "ar-SA",
|
||||||
|
Language.AR_AE: "ar-AE",
|
||||||
|
Language.BN: "bn",
|
||||||
|
Language.DE: "de",
|
||||||
|
Language.EN: "en",
|
||||||
|
Language.ES: "es-ES",
|
||||||
|
Language.ES_ES: "es-ES",
|
||||||
|
Language.ES_MX: "es-MX",
|
||||||
|
Language.FR: "fr",
|
||||||
|
Language.HI: "hi",
|
||||||
|
Language.ID: "id",
|
||||||
|
Language.IT: "it",
|
||||||
|
Language.JA: "ja",
|
||||||
|
Language.KO: "ko",
|
||||||
|
Language.PT: "pt-PT",
|
||||||
|
Language.PT_BR: "pt-BR",
|
||||||
|
Language.PT_PT: "pt-PT",
|
||||||
|
Language.RU: "ru",
|
||||||
|
Language.TR: "tr",
|
||||||
|
Language.VI: "vi",
|
||||||
|
Language.ZH: "zh",
|
||||||
|
}
|
||||||
|
|
||||||
|
return resolve_language(language, LANGUAGE_MAP, use_base_code=True)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class XAITTSSettings(TTSSettings):
|
class XAITTSSettings(TTSSettings):
|
||||||
"""Settings for XAITTSService."""
|
"""Settings for XAIHttpTTSService."""
|
||||||
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
class XAITTSService(TTSService):
|
class XAIHttpTTSService(TTSService):
|
||||||
"""xAI HTTP text-to-speech service.
|
"""xAI HTTP text-to-speech service.
|
||||||
|
|
||||||
The service requests raw PCM audio so emitted ``TTSAudioRawFrame`` objects
|
The service requests raw PCM audio so emitted ``TTSAudioRawFrame`` objects
|
||||||
@@ -41,31 +78,14 @@ class XAITTSService(TTSService):
|
|||||||
Settings = XAITTSSettings
|
Settings = XAITTSSettings
|
||||||
_settings: Settings
|
_settings: Settings
|
||||||
|
|
||||||
XAI_DEFAULT_SAMPLE_RATE = 24000
|
|
||||||
XAI_PCM_CODEC = "pcm"
|
|
||||||
|
|
||||||
class InputParams(BaseModel):
|
|
||||||
"""Input parameters for xAI TTS configuration.
|
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
|
||||||
Use ``settings=XAITTSService.Settings(...)`` instead.
|
|
||||||
|
|
||||||
Parameters:
|
|
||||||
language: Language for speech synthesis.
|
|
||||||
"""
|
|
||||||
|
|
||||||
language: Optional[Language] = None
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
base_url: str = "https://api.x.ai/v1/tts",
|
base_url: str = "https://api.x.ai/v1/tts",
|
||||||
voice: Optional[str] = None,
|
|
||||||
language: Optional[str | Language] = None,
|
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
|
encoding: Optional[str] = "pcm",
|
||||||
aiohttp_session: Optional[aiohttp.ClientSession] = None,
|
aiohttp_session: Optional[aiohttp.ClientSession] = None,
|
||||||
params: Optional[InputParams] = None,
|
|
||||||
settings: Optional[Settings] = None,
|
settings: Optional[Settings] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
@@ -74,54 +94,25 @@ class XAITTSService(TTSService):
|
|||||||
Args:
|
Args:
|
||||||
api_key: xAI API key for authentication.
|
api_key: xAI API key for authentication.
|
||||||
base_url: xAI TTS endpoint. Defaults to ``https://api.x.ai/v1/tts``.
|
base_url: xAI TTS endpoint. Defaults to ``https://api.x.ai/v1/tts``.
|
||||||
voice: Voice identifier. Defaults to ``"eve"``.
|
sample_rate: Audio sample rate. If None, uses default.
|
||||||
|
encoding: Output encoding format. Defaults to "pcm".
|
||||||
.. deprecated:: 0.0.105
|
|
||||||
Use ``settings=XAITTSService.Settings(voice=...)`` instead.
|
|
||||||
|
|
||||||
language: BCP-47 or base language code (for example ``"en"`` or ``"pt-BR"``).
|
|
||||||
Defaults to ``"en"``.
|
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
|
||||||
Use ``settings=XAITTSService.Settings(language=...)`` instead.
|
|
||||||
|
|
||||||
sample_rate: Output sample rate for PCM audio. Defaults to 24000 Hz.
|
|
||||||
aiohttp_session: Optional shared aiohttp session.
|
aiohttp_session: Optional shared aiohttp session.
|
||||||
params: Deprecated input parameters object.
|
settings: Runtime-updatable settings.
|
||||||
settings: Runtime-updatable settings. When provided alongside deprecated
|
|
||||||
parameters, ``settings`` values take precedence.
|
|
||||||
**kwargs: Additional keyword arguments passed to ``TTSService``.
|
**kwargs: Additional keyword arguments passed to ``TTSService``.
|
||||||
"""
|
"""
|
||||||
default_settings = self.Settings(
|
default_settings = self.Settings(
|
||||||
model=None,
|
model=None,
|
||||||
voice="eve",
|
voice="eve",
|
||||||
language="en",
|
language=Language.EN,
|
||||||
)
|
)
|
||||||
|
|
||||||
if voice is not None:
|
|
||||||
self._warn_init_param_moved_to_settings("voice", "voice")
|
|
||||||
default_settings.voice = voice
|
|
||||||
if language is not None:
|
|
||||||
self._warn_init_param_moved_to_settings("language", "language")
|
|
||||||
default_settings.language = (
|
|
||||||
self.language_to_service_language(language)
|
|
||||||
if isinstance(language, Language)
|
|
||||||
else language
|
|
||||||
)
|
|
||||||
|
|
||||||
if params is not None:
|
|
||||||
self._warn_init_param_moved_to_settings("params")
|
|
||||||
if not settings and params.language is not None:
|
|
||||||
default_settings.language = self.language_to_service_language(params.language)
|
|
||||||
|
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
default_settings.apply_update(settings)
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
pause_frame_processing=True,
|
sample_rate=sample_rate,
|
||||||
push_start_frame=True,
|
push_start_frame=True,
|
||||||
push_stop_frames=True,
|
push_stop_frames=True,
|
||||||
sample_rate=sample_rate or self.XAI_DEFAULT_SAMPLE_RATE,
|
|
||||||
settings=default_settings,
|
settings=default_settings,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
@@ -130,14 +121,22 @@ class XAITTSService(TTSService):
|
|||||||
self._base_url = base_url
|
self._base_url = base_url
|
||||||
self._session = aiohttp_session
|
self._session = aiohttp_session
|
||||||
self._session_owner = aiohttp_session is None
|
self._session_owner = aiohttp_session is None
|
||||||
|
self._encoding = encoding
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
"""Check if this service can generate processing metrics."""
|
"""Check if this service can generate processing metrics."""
|
||||||
return True
|
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 xAI's language format."""
|
"""Convert a Language enum to xAI language format.
|
||||||
return str(language)
|
|
||||||
|
Args:
|
||||||
|
language: The language to convert.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The xAI-specific language code, or None if not supported.
|
||||||
|
"""
|
||||||
|
return language_to_xai_language(language)
|
||||||
|
|
||||||
async def start(self, frame):
|
async def start(self, frame):
|
||||||
"""Start the xAI TTS service."""
|
"""Start the xAI TTS service."""
|
||||||
@@ -175,7 +174,7 @@ class XAITTSService(TTSService):
|
|||||||
"text": text,
|
"text": text,
|
||||||
"voice_id": self._settings.voice,
|
"voice_id": self._settings.voice,
|
||||||
"output_format": {
|
"output_format": {
|
||||||
"codec": self.XAI_PCM_CODEC,
|
"codec": self._encoding,
|
||||||
"sample_rate": self.sample_rate,
|
"sample_rate": self.sample_rate,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -189,12 +188,11 @@ class XAITTSService(TTSService):
|
|||||||
|
|
||||||
measuring_ttfb = True
|
measuring_ttfb = True
|
||||||
try:
|
try:
|
||||||
async with self._session.post(self._base_url, json=payload, headers=headers) as response:
|
async with self._session.post(
|
||||||
|
self._base_url, json=payload, headers=headers
|
||||||
|
) as response:
|
||||||
if response.status != 200:
|
if response.status != 200:
|
||||||
error = await response.text(errors="ignore")
|
error = await response.text(errors="ignore")
|
||||||
logger.error(
|
|
||||||
f"{self} error getting audio (status: {response.status}, error: {error})"
|
|
||||||
)
|
|
||||||
yield ErrorFrame(
|
yield ErrorFrame(
|
||||||
error=f"Error getting audio (status: {response.status}, error: {error})"
|
error=f"Error getting audio (status: {response.status}, error: {error})"
|
||||||
)
|
)
|
||||||
@@ -208,6 +206,11 @@ class XAITTSService(TTSService):
|
|||||||
if measuring_ttfb:
|
if measuring_ttfb:
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
measuring_ttfb = False
|
measuring_ttfb = False
|
||||||
yield TTSAudioRawFrame(chunk, self.sample_rate, 1, context_id=context_id)
|
yield TTSAudioRawFrame(
|
||||||
|
chunk,
|
||||||
|
self.sample_rate,
|
||||||
|
1,
|
||||||
|
context_id=context_id,
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
yield ErrorFrame(error=f"Unknown error occurred: {e}")
|
yield ErrorFrame(error=f"Unknown error occurred: {e}")
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ from pipecat.frames.frames import (
|
|||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
TTSTextFrame,
|
TTSTextFrame,
|
||||||
)
|
)
|
||||||
from pipecat.services.xai.tts import XAITTSService
|
from pipecat.services.xai.tts import XAIHttpTTSService
|
||||||
from pipecat.tests.utils import run_test
|
from pipecat.tests.utils import run_test
|
||||||
|
|
||||||
|
|
||||||
@@ -52,7 +52,7 @@ async def test_run_xai_tts_success(aiohttp_client):
|
|||||||
base_url = str(client.make_url("/v1/tts"))
|
base_url = str(client.make_url("/v1/tts"))
|
||||||
|
|
||||||
async with aiohttp.ClientSession() as session:
|
async with aiohttp.ClientSession() as session:
|
||||||
tts_service = XAITTSService(
|
tts_service = XAIHttpTTSService(
|
||||||
api_key="test-key",
|
api_key="test-key",
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
aiohttp_session=session,
|
aiohttp_session=session,
|
||||||
|
|||||||
4
uv.lock
generated
4
uv.lock
generated
@@ -4966,7 +4966,11 @@ requires-dist = [
|
|||||||
{ name = "wait-for2", marker = "python_full_version < '3.12'", specifier = ">=0.4.1,<1" },
|
{ name = "wait-for2", marker = "python_full_version < '3.12'", specifier = ">=0.4.1,<1" },
|
||||||
{ name = "websockets", marker = "extra == 'websockets-base'", specifier = ">=13.1,<16.0" },
|
{ name = "websockets", marker = "extra == 'websockets-base'", specifier = ">=13.1,<16.0" },
|
||||||
]
|
]
|
||||||
|
<<<<<<< HEAD
|
||||||
provides-extras = ["aic", "anthropic", "assemblyai", "asyncai", "aws", "aws-nova-sonic", "azure", "cartesia", "camb", "cerebras", "daily", "deepgram", "deepseek", "elevenlabs", "fal", "fireworks", "fish", "gladia", "google", "gradium", "grok", "groq", "gstreamer", "heygen", "hume", "inworld", "koala", "kokoro", "krisp", "langchain", "lemonslice", "livekit", "lmnt", "local", "local-smart-turn", "mcp", "mem0", "mistral", "mlx-whisper", "moondream", "neuphonic", "noisereduce", "novita", "nvidia", "openai", "rnnoise", "openpipe", "openrouter", "perplexity", "piper", "qwen", "remote-smart-turn", "resembleai", "rime", "riva", "runner", "sagemaker", "sambanova", "sarvam", "smallest", "sentry", "silero", "simli", "soniox", "soundfile", "speechmatics", "strands", "tavus", "together", "tracing", "ultravox", "webrtc", "websocket", "websockets-base", "whisper"]
|
provides-extras = ["aic", "anthropic", "assemblyai", "asyncai", "aws", "aws-nova-sonic", "azure", "cartesia", "camb", "cerebras", "daily", "deepgram", "deepseek", "elevenlabs", "fal", "fireworks", "fish", "gladia", "google", "gradium", "grok", "groq", "gstreamer", "heygen", "hume", "inworld", "koala", "kokoro", "krisp", "langchain", "lemonslice", "livekit", "lmnt", "local", "local-smart-turn", "mcp", "mem0", "mistral", "mlx-whisper", "moondream", "neuphonic", "noisereduce", "novita", "nvidia", "openai", "rnnoise", "openpipe", "openrouter", "perplexity", "piper", "qwen", "remote-smart-turn", "resembleai", "rime", "riva", "runner", "sagemaker", "sambanova", "sarvam", "smallest", "sentry", "silero", "simli", "soniox", "soundfile", "speechmatics", "strands", "tavus", "together", "tracing", "ultravox", "webrtc", "websocket", "websockets-base", "whisper"]
|
||||||
|
=======
|
||||||
|
provides-extras = ["aic", "anthropic", "assemblyai", "asyncai", "aws", "aws-nova-sonic", "azure", "cartesia", "camb", "cerebras", "daily", "deepgram", "deepseek", "elevenlabs", "fal", "fireworks", "fish", "gladia", "google", "gradium", "grok", "groq", "gstreamer", "heygen", "hume", "inworld", "koala", "kokoro", "krisp", "langchain", "lemonslice", "livekit", "lmnt", "local", "local-smart-turn", "mcp", "mem0", "mistral", "mlx-whisper", "moondream", "neuphonic", "noisereduce", "novita", "nvidia", "openai", "rnnoise", "openpipe", "openrouter", "perplexity", "piper", "qwen", "remote-smart-turn", "resembleai", "rime", "riva", "runner", "sagemaker", "sambanova", "sarvam", "sentry", "silero", "simli", "soniox", "soundfile", "speechmatics", "strands", "tavus", "together", "tracing", "ultravox", "webrtc", "websocket", "websockets-base", "whisper", "xai"]
|
||||||
|
>>>>>>> 43f25faca (Add xAI HTTP TTS service)
|
||||||
|
|
||||||
[package.metadata.requires-dev]
|
[package.metadata.requires-dev]
|
||||||
dev = [
|
dev = [
|
||||||
|
|||||||
Reference in New Issue
Block a user