Files
hermes-agent/tools/tts_streaming.py

459 lines
17 KiB
Python

"""Provider-agnostic streaming TTS: sentence text → int16 PCM chunk iterator.
``stream_tts_to_speaker`` (``tools.tts_tool``) owns the sentence buffer,
sounddevice output and stop/queue protocol; this module owns the *provider*
half — turning one sentence into audio the moment it's ready so playback starts
on sentence one instead of after the whole reply.
One contract (int16 mono PCM at ``sample_rate``): **true streamers**
(`StreamingTTSProvider.stream`) wrap chunked APIs (ElevenLabs pcm_24000, OpenAI
pcm, …); providers with no chunked API (edge, the default) still get per-
*sentence* playback via the sync ``text_to_speech_tool`` path in the dispatcher.
Adding a streamer is ``@register("name")`` on a subclass; the dispatcher, config
gate (``tts.<name>.streaming``) and resolver come free.
"""
from __future__ import annotations
import logging
import re
import time
from abc import ABC, abstractmethod
from typing import Callable, Dict, Iterator, List, Optional
from tools.tool_backend_helpers import resolve_openai_audio_api_key
from tools.tts_tool import _get_provider, _load_tts_config, get_env_value
logger = logging.getLogger(__name__)
# Per-sentence PCM byte cap, mirroring the 16 MiB bounded-body invariant of the
# sync providers: a buggy or hostile endpoint must not feed unbounded audio.
_STREAM_SENTENCE_BYTE_CAP = 16 * 1024 * 1024
def _resolve_key(env_var: str, provider_id: str) -> str:
"""Provider secret lookup (config > env/.env > credential pool).
Monkeypatchable seam over ``tools.tts_tool._resolve_provider_key``. ALL
streaming-provider key lookups go through here — never bare ``get_env_value``.
"""
try:
from tools.tts_tool import _resolve_provider_key
return _resolve_provider_key(env_var, provider_id) or ""
except Exception:
return get_env_value(env_var) or ""
def _gemini_key() -> str:
return _resolve_key("GEMINI_API_KEY", "gemini") or _resolve_key("GOOGLE_API_KEY", "gemini")
# ---------------------------------------------------------------------------
# Interruption latch — lets the model know it was cut off mid-speech
# ---------------------------------------------------------------------------
# When the user barges in on a spoken reply, the surface marks the latch; the
# next turn's submit path takes it and prepends SPEECH_INTERRUPTED_NOTE to the
# model-bound message (API-call local, never persisted). The TTL keeps a stale
# barge from annotating an unrelated message minutes later.
SPEECH_INTERRUPTED_NOTE = (
"[Note: the user interrupted your previous spoken reply before it finished.]"
)
_INTERRUPT_TTL_S = 120.0
_interrupted_at: Optional[float] = None
def mark_speech_interrupted() -> None:
global _interrupted_at
_interrupted_at = time.monotonic()
def take_speech_interrupted() -> bool:
"""Pop the latch; True when a barge happened within the TTL."""
global _interrupted_at
at, _interrupted_at = _interrupted_at, None
return at is not None and time.monotonic() - at < _INTERRUPT_TTL_S
# Sentence boundary: after .!? followed by whitespace, or a blank line.
SENTENCE_BOUNDARY_RE = re.compile(r"(?<=[.!?])(?:\s|\n)|(?:\n\n)")
_THINK_BLOCK_RE = re.compile(r"<think[\s>].*?</think>", flags=re.DOTALL)
class SentenceChunker:
"""Incremental sentence cutter for LLM token deltas.
Shared by the speaker pipeline and the speak-stream WebSocket so every
surface cuts speech identically. Strips ``<think>`` blocks (even split
across deltas) and merges fragments shorter than *min_len* into the
following sentence, so "Ha!" rides along instead of stalling as a tiny clip.
"""
def __init__(self, min_len: int = 20):
self.min_len = min_len
self.buf = ""
def feed(self, delta: str) -> List[str]:
"""Absorb *delta*; return every complete sentence now ready to speak."""
self.buf = _THINK_BLOCK_RE.sub("", self.buf + delta)
if "<think" in self.buf and "</think>" not in self.buf:
return [] # open think tag — the closing tag may arrive next delta
out: List[str] = []
start = 0 # skip boundaries that would leave the head too short
while m := SENTENCE_BOUNDARY_RE.search(self.buf, start):
head = self.buf[: m.end()]
if len(head.strip()) < self.min_len:
start = m.end()
continue
out.append(head)
self.buf = self.buf[m.end():]
start = 0
return out
def flush(self) -> List[str]:
"""Drain the tail (end-of-text or long-idle flush)."""
tail = _THINK_BLOCK_RE.sub("", self.buf).strip()
self.buf = ""
return [tail] if tail else []
# ---------------------------------------------------------------------------
# ABC + registry
# ---------------------------------------------------------------------------
class StreamingTTSProvider(ABC):
"""Yields raw int16, little-endian, mono PCM chunks at ``sample_rate``."""
sample_rate: int = 24000
channels: int = 1
sample_width: int = 2 # bytes/sample (int16)
def __init__(self, tts_config: Dict, section: Dict):
self.tts_config = tts_config
self.section = section
@staticmethod
@abstractmethod
def available() -> bool:
"""True when this provider's credentials/SDK are usable right now."""
@abstractmethod
def stream(self, text: str) -> Iterator[bytes]:
"""Yield PCM chunks for ``text``. Raise on failure (caller logs)."""
_REGISTRY: Dict[str, type[StreamingTTSProvider]] = {}
def register(name: str) -> Callable[[type[StreamingTTSProvider]], type[StreamingTTSProvider]]:
def _wrap(cls: type[StreamingTTSProvider]) -> type[StreamingTTSProvider]:
_REGISTRY[name] = cls
return cls
return _wrap
def _try_instantiate(name: str, tts_config: Dict) -> Optional[StreamingTTSProvider]:
"""Construct the registered streamer *name* if it's usable, else None."""
cls = _REGISTRY.get(name)
if cls is None or not cls.available():
return None
try:
return cls(tts_config, tts_config.get(name) or {})
except Exception as exc: # pragma: no cover - defensive
logger.debug("streaming provider %s init failed: %s", name, exc)
return None
# Fallback priority for ``tts.streaming.provider: auto`` — best chunked
# latency/quality first. Deliberately hard-coded (a UX decision, not a config
# knob); edge is absent because it has no chunked-PCM API.
_PROVIDER_PRIORITY: List[str] = ["elevenlabs", "gemini", "openai", "xai"]
def resolve_streaming_provider(
tts_config: Dict,
preferred: Optional[str] = None,
) -> Optional[StreamingTTSProvider]:
"""Return a ready streamer for the *configured* provider, else ``None``.
1. ``tts.streaming.provider`` when set: a name pins that exact streamer
(or ``None`` if unusable); ``auto`` walks ``_PROVIDER_PRIORITY`` and
returns the first usable one.
2. Otherwise the configured TTS provider (or ``preferred``). ``None`` means
"no chunked API" — the dispatcher speaks per-sentence via the sync path,
preserving the user's chosen voice. We never silently swap providers
just to get streaming.
"""
streaming_cfg = tts_config.get("streaming") or {}
pinned = str(streaming_cfg.get("provider") or "").lower().strip()
if pinned == "auto":
for name in _PROVIDER_PRIORITY:
inst = _try_instantiate(name, tts_config)
if inst is not None:
return inst
return None
if pinned:
return _try_instantiate(pinned, tts_config)
name = (preferred or _get_provider(tts_config)).lower().strip()
return _try_instantiate(name, tts_config)
def _capped(chunks: Iterator[bytes], label: str) -> Iterator[bytes]:
"""Pass chunks through, aborting past the per-sentence byte cap (runaway/hostile upstream)."""
total = 0
for chunk in chunks:
total += len(chunk)
if total > _STREAM_SENTENCE_BYTE_CAP:
logger.warning("%s exceeded %d bytes for one sentence; truncating",
label, _STREAM_SENTENCE_BYTE_CAP)
return
yield chunk
# ---------------------------------------------------------------------------
# Providers
# ---------------------------------------------------------------------------
@register("elevenlabs")
class ElevenLabsStreamer(StreamingTTSProvider):
"""ElevenLabs chunked HTTP → pcm_24000 (the original reference path)."""
sample_rate = 24000
@staticmethod
def available() -> bool:
return bool(_resolve_key("ELEVENLABS_API_KEY", "elevenlabs"))
def stream(self, text: str) -> Iterator[bytes]:
from tools.tts_tool import (
DEFAULT_ELEVENLABS_STREAMING_MODEL_ID,
DEFAULT_ELEVENLABS_VOICE_ID,
_elevenlabs_environment_kwargs,
_import_elevenlabs,
)
client = _import_elevenlabs()(
api_key=_resolve_key("ELEVENLABS_API_KEY", "elevenlabs"),
**_elevenlabs_environment_kwargs(self.section),
)
voice_id = self.section.get("voice_id", DEFAULT_ELEVENLABS_VOICE_ID)
model_id = self.section.get(
"streaming_model_id",
self.section.get("model_id", DEFAULT_ELEVENLABS_STREAMING_MODEL_ID),
)
yield from client.text_to_speech.convert(
text=text,
voice_id=voice_id,
model_id=model_id,
output_format="pcm_24000",
)
def _openai_config_api_key() -> str:
"""Return ``tts.openai.api_key`` from config.yaml, or empty string."""
try:
openai_cfg = (_load_tts_config().get("openai") or {})
except Exception:
return ""
return openai_cfg.get("api_key") or ""
@register("openai")
class OpenAIStreamer(StreamingTTSProvider):
"""OpenAI speech with ``response_format=pcm`` (24 kHz mono int16)."""
sample_rate = 24000
@staticmethod
def available() -> bool:
return bool(_openai_config_api_key() or resolve_openai_audio_api_key())
def stream(self, text: str) -> Iterator[bytes]:
from openai import OpenAI
client = OpenAI(
api_key=(self.section.get("api_key") or resolve_openai_audio_api_key()),
base_url=(
self.section.get("base_url")
or get_env_value("OPENAI_BASE_URL")
or None
),
)
model = self.section.get("model", "gpt-4o-mini-tts")
voice = self.section.get("voice", "alloy")
with client.audio.speech.with_streaming_response.create(
model=model,
voice=voice,
input=text,
response_format="pcm",
) as response:
yield from _capped(response.iter_bytes(), "OpenAI streaming TTS")
@register("gemini")
class GeminiStreamer(StreamingTTSProvider):
"""Gemini ``streamGenerateContent?alt=sse`` → base64 PCM chunks (24 kHz).
``?alt=sse`` flips the response from one JSON blob to an SSE feed of
base64 PCM chunks. Uses requests with a bounded streamed body.
"""
sample_rate = 24000
@staticmethod
def available() -> bool:
return bool(_gemini_key())
def stream(self, text: str) -> Iterator[bytes]:
import base64
import json as _json
import requests
from tools.tts_tool import (
DEFAULT_GEMINI_TTS_BASE_URL,
DEFAULT_GEMINI_TTS_MODEL,
DEFAULT_GEMINI_TTS_VOICE,
)
api_key = _gemini_key()
model = str(self.section.get("model", DEFAULT_GEMINI_TTS_MODEL)).strip() or DEFAULT_GEMINI_TTS_MODEL
voice = str(self.section.get("voice", DEFAULT_GEMINI_TTS_VOICE)).strip() or DEFAULT_GEMINI_TTS_VOICE
base_url = str(
self.section.get("base_url")
or get_env_value("GEMINI_BASE_URL")
or DEFAULT_GEMINI_TTS_BASE_URL
).strip().rstrip("/")
payload = {
"contents": [{"parts": [{"text": text}]}],
"generationConfig": {
"responseModalities": ["AUDIO"],
"speechConfig": {
"voiceConfig": {
"prebuiltVoiceConfig": {"voiceName": voice},
},
},
},
}
url = f"{base_url}/models/{model}:streamGenerateContent"
def _sse_chunks() -> Iterator[bytes]:
with requests.post(
url,
params={"alt": "sse", "key": api_key},
json=payload,
timeout=60,
stream=True,
) as response:
response.raise_for_status()
for line in response.iter_lines(decode_unicode=True):
if not line or not line.startswith("data: "):
continue
try:
event = _json.loads(line[len("data: "):])
parts = event["candidates"][0]["content"]["parts"]
except (ValueError, KeyError, IndexError, TypeError):
continue
for part in parts:
inline = part.get("inlineData") or part.get("inline_data") or {}
b64 = inline.get("data", "")
if not b64:
continue
try:
yield base64.b64decode(b64)
except (ValueError, TypeError) as exc:
logger.warning("Gemini SSE: bad base64 audio: %s", exc)
yield from _capped(_sse_chunks(), "Gemini streaming TTS")
@register("xai")
class XAIStreamer(StreamingTTSProvider):
"""xAI WebSocket TTS (``wss://api.x.ai/v1/tts``) → binary PCM frames (24 kHz mono int16).
Credentials route through ``resolve_xai_http_credentials`` (OAuth or
XAI_API_KEY), same as the sync path. The async WS loop is bridged to the
sync iterator contract via ``_collect_async`` — the seam unit tests patch.
"""
sample_rate = 24000
@staticmethod
def available() -> bool:
try:
from tools.xai_http import resolve_xai_http_credentials
creds = resolve_xai_http_credentials()
return bool(str(creds.get("api_key") or "").strip())
except Exception:
return False
def stream(self, text: str) -> Iterator[bytes]:
yield from _capped(iter(self._collect_async(text)), "xAI streaming TTS")
# -- async→sync bridge (test seam) ------------------------------------
def _collect_async(self, text: str) -> List[bytes]:
import asyncio
return asyncio.run(self._drain_async(text))
async def _drain_async(self, text: str) -> List[bytes]:
frames: List[bytes] = []
async for frame in self._async_frames(text):
frames.append(frame)
return frames
async def _async_frames(self, text: str):
import json as _json
import websockets
from tools.tts_tool import DEFAULT_XAI_VOICE_ID
from tools.xai_http import resolve_xai_http_credentials
creds = resolve_xai_http_credentials()
api_key = str(creds.get("api_key") or "").strip()
if not api_key:
raise RuntimeError("No xAI credentials for streaming TTS")
voice = str(self.section.get("voice_id", DEFAULT_XAI_VOICE_ID)).strip() or DEFAULT_XAI_VOICE_ID
ws_url = str(
self.section.get("streaming_url") or "wss://api.x.ai/v1/tts"
).strip()
async with websockets.connect(
ws_url, extra_headers={"Authorization": f"Bearer {api_key}"}
) as ws:
await ws.send(_json.dumps({
"text": text,
"voice_id": voice,
"response_format": "pcm",
}))
try:
while True:
message = await ws.recv()
if isinstance(message, (bytes, bytearray, memoryview)):
yield bytes(message)
continue
try:
envelope = _json.loads(message)
except (ValueError, TypeError):
if message == "done":
return
continue
etype = envelope.get("type")
if etype == "done":
return
if etype == "error":
logger.warning("xAI WS error envelope: %s",
envelope.get("error") or envelope.get("message") or envelope)
return
except Exception as exc:
if exc.__class__.__name__ == "ConnectionClosed":
return
logger.warning("xAI WS receive failed: %s", exc)
return