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:
Mark Backman
2026-03-24 10:30:27 -04:00
parent 79dafb9ac9
commit 9af5c482c6
6 changed files with 277 additions and 134 deletions

1
changelog/4031.added.md Normal file
View File

@@ -0,0 +1 @@
- Added `XAIHttpTTSService` for text-to-speech using xAI's HTTP TTS API.

View 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()

View File

@@ -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):

View File

@@ -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}")

View File

@@ -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
View File

@@ -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 = [