refactor(tools/voice): extract tts delivery + wake_word engines; dedupe transcription/voice_mode helpers; compact tts providers

This commit is contained in:
Teknium
2026-09-02 13:55:33 -07:00
parent 49d15faae6
commit 2ef1e8e4e0
11 changed files with 5035 additions and 6860 deletions
+745 -1330
View File
File diff suppressed because it is too large Load Diff
+56 -86
View File
@@ -1,22 +1,16 @@
"""Provider-agnostic streaming TTS: sentence text → int16 PCM chunk iterator.
The keystone of Hermes' conversational voice UX. `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.
``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.
Two provider shapes, one contract (int16 mono PCM at ``sample_rate``):
* **True streamers** (`StreamingTTSProvider.stream`) — chunked APIs
(ElevenLabs pcm_24000, OpenAI pcm, …) that yield audio as it synthesizes.
Lowest time-to-first-audio.
* **Everyone else** — providers with no chunked API still get per-*sentence*
playback via the proven sync `text_to_speech_tool` path (handled by the
dispatcher, not here), so edge (the default) is conversational too.
Adding a streamer is `@register("name")` on a `StreamingTTSProvider` subclass;
the dispatcher, config gate (`tts.<name>.streaming`), and resolver come free.
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
@@ -32,19 +26,16 @@ from tools.tts_tool import _get_provider, _load_tts_config, get_env_value
logger = logging.getLogger(__name__)
# Upper bound on the PCM bytes accepted from one provider stream for one
# sentence. Mirrors the 16 MiB bounded-upstream-body invariant of the sync
# providers (``_read_tts_response_bytes`` in tools.tts_tool): a buggy or
# hostile endpoint must not be able to feed us unbounded audio.
# 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.
"""Provider secret lookup (config > env/.env > credential pool).
Thin, monkeypatchable seam over ``tools.tts_tool._resolve_provider_key``
(which delegates to ``resolve_provider_secret``). ALL streaming-provider
key lookups go through here — never bare ``get_env_value``.
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
@@ -54,14 +45,17 @@ def _resolve_key(env_var: str, provider_id: str) -> str:
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 (talks over it, types, hits the
# record key), 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, same as the CLI's model-switch notes). The TTL
# keeps a stale barge from annotating an unrelated message minutes later.
# 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.]"
@@ -89,11 +83,10 @@ _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 (`stream_tts_to_speaker`) 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 with the
sentence after it instead of stalling as a tiny clip.
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):
@@ -173,9 +166,8 @@ def _try_instantiate(name: str, tts_config: Dict) -> Optional[StreamingTTSProvid
# 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 — the
# dispatcher's per-sentence sync path keeps it conversational instead.
# 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"]
@@ -185,18 +177,13 @@ def resolve_streaming_provider(
) -> Optional[StreamingTTSProvider]:
"""Return a ready streamer for the *configured* provider, else ``None``.
Resolution order:
1. ``tts.streaming.provider`` (config knob) when set:
* a provider name pins that exact streamer (or ``None`` if unusable);
* ``auto`` walks the priority list (``elevenlabs → gemini → openai
→ xai``) and returns the first usable streamer — an explicit
opt-in to "give me the best chunked voice available".
2. Otherwise the *configured* TTS provider (or ``preferred`` override).
``None`` means "no chunked API for this provider" — the dispatcher
then speaks per-sentence via the sync path, preserving the user's
chosen voice. We never silently swap to a different provider just
to get streaming.
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()
@@ -213,6 +200,18 @@ def resolve_streaming_provider(
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
# ---------------------------------------------------------------------------
@@ -293,40 +292,19 @@ class OpenAIStreamer(StreamingTTSProvider):
yield from _capped(response.iter_bytes(), "OpenAI streaming TTS")
def _capped(chunks: Iterator[bytes], label: str) -> Iterator[bytes]:
"""Pass chunks through, aborting past the 16 MiB per-sentence cap.
The streaming mirror of ``_read_tts_response_bytes``'s bounded-body
invariant: one sentence of PCM should never approach the cap, so
exceeding it means a runaway/hostile upstream — stop pulling.
"""
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
@register("gemini")
class GeminiStreamer(StreamingTTSProvider):
"""Gemini ``streamGenerateContent?alt=sse`` → base64 PCM chunks (24 kHz).
Salvaged from PR #47588 (@Cdddo) and rebased onto the post-campaign
infrastructure: credentials via the provider-secret resolver, requests
(not httpx) with a bounded streamed body, and main's provider ABC.
``?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(
_resolve_key("GEMINI_API_KEY", "gemini")
or _resolve_key("GOOGLE_API_KEY", "gemini")
)
return bool(_gemini_key())
def stream(self, text: str) -> Iterator[bytes]:
import base64
@@ -340,10 +318,7 @@ class GeminiStreamer(StreamingTTSProvider):
DEFAULT_GEMINI_TTS_VOICE,
)
api_key = (
_resolve_key("GEMINI_API_KEY", "gemini")
or _resolve_key("GOOGLE_API_KEY", "gemini")
)
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(
@@ -363,8 +338,6 @@ class GeminiStreamer(StreamingTTSProvider):
},
},
}
# ``?alt=sse`` flips the response from a single JSON blob to an SSE
# feed of base64 PCM chunks — the whole point of this provider.
url = f"{base_url}/models/{model}:streamGenerateContent"
def _sse_chunks() -> Iterator[bytes]:
@@ -399,14 +372,11 @@ class GeminiStreamer(StreamingTTSProvider):
@register("xai")
class XAIStreamer(StreamingTTSProvider):
"""xAI WebSocket TTS → binary PCM frames (24 kHz mono int16).
"""xAI WebSocket TTS (``wss://api.x.ai/v1/tts``) → binary PCM frames (24 kHz mono int16).
Salvaged from PR #47588 (@Cdddo): xAI's chunked TTS API is
WebSocket-only (``wss://api.x.ai/v1/tts``). Credentials route through
``resolve_xai_http_credentials`` (OAuth or XAI_API_KEY), same as the
sync ``_generate_xai_tts`` path. The async WS loop is bridged to the
sync iterator contract via ``_collect_async`` — the seam unit tests
monkeypatch.
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
+634 -3337
View File
File diff suppressed because it is too large Load Diff
+549
View File
@@ -0,0 +1,549 @@
"""Long-form chunking, ffmpeg encoding, container repair and delivery packing.
Everything here is provider-agnostic post-processing for ``tools.tts_tool``:
split text under a per-request cap, wrap raw PCM as WAV, convert WAV/MP3 to
the target container, sniff/repair mislabelled ``.ogg`` files, and combine
final-encoded chunks under a destination platform's upload limit. Origin
module re-imports every name under its historical spelling.
"""
from __future__ import annotations
import logging
import os
import re
import shlex
import shutil
import struct
import subprocess
import tempfile
import uuid
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
from hermes_cli._subprocess_compat import windows_hide_flags
logger = logging.getLogger("tools.tts_tool")
# Final fallback when provider isn't recognised at all.
FALLBACK_MAX_TEXT_LENGTH = 4000
# PCM output specs for Gemini TTS (fixed by the API)
GEMINI_TTS_SAMPLE_RATE = 24000
GEMINI_TTS_CHANNELS = 1
GEMINI_TTS_SAMPLE_WIDTH = 2 # 16-bit PCM (L16)
# ffmpeg args producing the Ogg/Opus voice-bubble encoding Telegram & co expect.
_OPUS_VOICE_ARGS = [
"-acodec", "libopus", "-ac", "1", "-b:a", "48k", "-vbr", "on",
"-application", "voip", "-compression_level", "10",
]
# ===========================================================================
# Text chunking and delivery profiles
# ===========================================================================
@dataclass(frozen=True)
class AudioDeliveryProfile:
"""Destination-platform constraints for generated TTS audio."""
platform: str
max_file_bytes: int
safety_ratio: float = 0.85
@property
def target_file_bytes(self) -> int:
"""Conservative packing target below the platform hard limit."""
return max(1, int(self.max_file_bytes * self.safety_ratio))
_PLATFORM_AUDIO_DEFAULTS: Dict[str, Dict[str, Any]] = {
"discord": {"max_file_bytes": 10 * 1024 * 1024, "safety_ratio": 0.85},
"telegram": {"max_file_bytes": 50 * 1024 * 1024, "safety_ratio": 0.85},
"default": {"max_file_bytes": 10 * 1024 * 1024, "safety_ratio": 0.85},
}
def _resolve_audio_delivery_profile(
platform: Optional[str],
tts_config: Optional[Dict[str, Any]] = None,
) -> AudioDeliveryProfile:
"""Resolve upload constraints, including optional ``tts.delivery_profiles`` overrides."""
key = (platform or "default").lower().strip() or "default"
defaults = dict(_PLATFORM_AUDIO_DEFAULTS.get(key) or _PLATFORM_AUDIO_DEFAULTS["default"])
profiles = (tts_config or {}).get("delivery_profiles")
overrides = profiles.get(key, {}) if isinstance(profiles, dict) else {}
if isinstance(overrides, dict):
defaults.update({k: v for k, v in overrides.items() if v is not None})
max_file_bytes = defaults.get("max_file_bytes")
if isinstance(max_file_bytes, bool) or not isinstance(max_file_bytes, int) or max_file_bytes <= 0:
max_file_bytes = _PLATFORM_AUDIO_DEFAULTS["default"]["max_file_bytes"]
safety_ratio = defaults.get("safety_ratio", 0.85)
if (
isinstance(safety_ratio, bool)
or not isinstance(safety_ratio, (int, float))
or not 0 < safety_ratio <= 1
):
safety_ratio = 0.85
return AudioDeliveryProfile(platform=key, max_file_bytes=max_file_bytes, safety_ratio=float(safety_ratio))
def _pack_under_cap(pieces: List[str], max_chars: int) -> List[str]:
"""Greedily join *pieces* with single spaces, starting a new chunk past *max_chars*."""
chunks: List[str] = []
current = ""
for piece in pieces:
candidate = f"{current} {piece}".strip()
if current and len(candidate) > max_chars:
chunks.append(current)
current = piece
else:
current = candidate
if current:
chunks.append(current)
return chunks
def _split_oversized_sentence(sentence: str, max_chars: int) -> List[str]:
"""Split one over-limit sentence on word boundaries, then hard boundaries.
An over-long word flushes the running chunk and emits its slices as their
own chunks (the tail slice is not merged with following words).
"""
chunks: List[str] = []
current = ""
for word in sentence.split():
if len(word) > max_chars:
if current:
chunks.append(current)
current = ""
chunks.extend(word[i:i + max_chars] for i in range(0, len(word), max_chars))
continue
candidate = f"{current} {word}".strip()
if current and len(candidate) > max_chars:
chunks.append(current)
current = word
else:
current = candidate
if current:
chunks.append(current)
return chunks
def _split_text_for_tts(text: str, max_chars: int) -> List[str]:
"""Split text under a provider cap without dropping normalized content."""
if max_chars <= 0:
max_chars = FALLBACK_MAX_TEXT_LENGTH
normalized = " ".join((text or "").split())
if not normalized:
return []
if len(normalized) <= max_chars:
return [normalized]
expanded: List[str] = []
for sentence in re.split(r"(?<=[.!?;:,])\s+", normalized):
sentence = sentence.strip()
if not sentence:
continue
if len(sentence) <= max_chars:
expanded.append(sentence)
else:
expanded.extend(_split_oversized_sentence(sentence, max_chars))
return _pack_under_cap(expanded, max_chars)
def _pack_audio_files_for_delivery(
audio_paths: List[str],
profile: AudioDeliveryProfile,
) -> List[List[str]]:
"""Group already-final-encoded chunks under the conservative size target.
A group never mixes container suffixes (they can't be concat-copied).
"""
groups: List[List[str]] = []
current: List[str] = []
current_size = 0
current_suffix = ""
for path in audio_paths:
size = Path(path).stat().st_size
suffix = Path(path).suffix.lower()
if current and (current_size + size > profile.target_file_bytes or suffix != current_suffix):
groups.append(current)
current, current_size = [], 0
current.append(path)
current_size += size
current_suffix = suffix
if current:
groups.append(current)
return groups
# ===========================================================================
# ffmpeg encoding helpers
# ===========================================================================
def _has_ffmpeg() -> bool:
return shutil.which("ffmpeg") is not None
def _ffmpeg_run(args: List[str], *, timeout: int = 30) -> subprocess.CompletedProcess:
"""Run ``ffmpeg <args>`` headless (no stdin, hidden window on Windows)."""
return subprocess.run(
["ffmpeg", *args],
capture_output=True,
timeout=timeout,
stdin=subprocess.DEVNULL,
creationflags=windows_hide_flags(),
)
def _wav_sidecar_path(output_path: str) -> str:
"""Path a WAV-native engine writes to before conversion to *output_path*'s format."""
if output_path.endswith(".wav"):
return output_path
return output_path.rsplit(".", 1)[0] + ".wav"
def _finalize_wav_output(wav_path: str, output_path: str) -> str:
"""Move a WAV-native engine's output into the caller's requested container.
Shared by NeuTTS / Piper / KittenTTS: ffmpeg-convert when available,
otherwise rename the WAV to the expected path so the tool stays usable
(the extension is then misleading but the audio plays).
"""
if wav_path == output_path:
return output_path
ffmpeg = shutil.which("ffmpeg")
if ffmpeg:
subprocess.run(
[ffmpeg, "-i", wav_path, "-y", "-loglevel", "error", output_path],
check=True, timeout=30, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags(),
)
try:
os.remove(wav_path)
except OSError:
pass
else:
os.rename(wav_path, output_path)
return output_path
def _wrap_pcm_as_wav(
pcm_bytes: bytes,
sample_rate: int = GEMINI_TTS_SAMPLE_RATE,
channels: int = GEMINI_TTS_CHANNELS,
sample_width: int = GEMINI_TTS_SAMPLE_WIDTH,
) -> bytes:
"""Wrap raw signed-little-endian PCM (e.g. Gemini's L16) with a minimal WAV RIFF header."""
byte_rate = sample_rate * channels * sample_width
block_align = channels * sample_width
data_size = len(pcm_bytes)
fmt_chunk = struct.pack(
"<4sIHHIIHH", b"fmt ", 16, 1, channels, sample_rate, byte_rate, block_align, sample_width * 8,
)
data_chunk_header = struct.pack("<4sI", b"data", data_size)
riff_size = 4 + len(fmt_chunk) + len(data_chunk_header) + data_size
riff_header = struct.pack("<4sI4s", b"RIFF", riff_size, b"WAVE")
return riff_header + fmt_chunk + data_chunk_header + pcm_bytes
def _write_wav_bytes_as(wav_bytes: bytes, output_path: str) -> str:
"""Write in-memory WAV to *output_path*, ffmpeg-converting to its container.
``.wav`` is written directly; ``.ogg`` is forced to Opus (ffmpeg's .ogg
default is Vorbis, which voice bubbles reject); anything else is a plain
ffmpeg conversion. A failed conversion raises RuntimeError. Without
ffmpeg the raw WAV is written under the requested name (misleading
extension, but the audio still plays).
"""
if output_path.lower().endswith(".wav"):
with open(output_path, "wb") as f:
f.write(wav_bytes)
return output_path
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
tmp.write(wav_bytes)
wav_path = tmp.name
try:
ffmpeg = shutil.which("ffmpeg")
if ffmpeg:
opus = _OPUS_VOICE_ARGS if output_path.lower().endswith(".ogg") else []
cmd = [ffmpeg, "-i", wav_path, *opus, "-y", "-loglevel", "error", output_path]
result = subprocess.run(cmd, capture_output=True, timeout=30, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags())
if result.returncode != 0:
stderr = result.stderr.decode("utf-8", errors="ignore")[:300]
raise RuntimeError(f"ffmpeg conversion failed: {stderr}")
else:
logger.warning(
"ffmpeg not found; writing raw WAV to %s (extension may be misleading)",
output_path,
)
shutil.copyfile(wav_path, output_path)
finally:
try:
os.remove(wav_path)
except OSError:
pass
return output_path
def _convert_to_opus(mp3_path: str) -> Optional[str]:
"""Convert any ffmpeg-readable audio file to OGG Opus next to it; None on failure."""
if not _has_ffmpeg():
return None
return _ffmpeg_transcode_to_opus(mp3_path, mp3_path.rsplit(".", 1)[0] + ".ogg")
def _ffmpeg_transcode_to_opus(input_path: str, ogg_path: str) -> Optional[str]:
"""Transcode *input_path* to real Ogg/Opus at *ogg_path* via ffmpeg.
Safe when ``input_path == ogg_path`` (writes to a temp file, then
replaces). Returns the output path on success, None on failure.
"""
if not _has_ffmpeg():
return None
in_place = os.path.abspath(input_path) == os.path.abspath(ogg_path)
work_path = ogg_path + ".tmp.ogg" if in_place else ogg_path
try:
result = _ffmpeg_run(["-i", input_path, *_OPUS_VOICE_ARGS, "-f", "ogg", work_path, "-y"])
if result.returncode != 0:
logger.warning("ffmpeg conversion failed with return code %d: %s",
result.returncode, result.stderr.decode('utf-8', errors='ignore')[:200])
return None
if os.path.exists(work_path) and os.path.getsize(work_path) > 0:
if in_place:
os.replace(work_path, ogg_path)
return ogg_path
except subprocess.TimeoutExpired:
logger.warning("ffmpeg OGG conversion timed out after 30s")
except FileNotFoundError:
logger.warning("ffmpeg not found in PATH")
except Exception as e:
logger.warning("ffmpeg OGG conversion failed: %s", e, exc_info=True)
finally:
if in_place and os.path.exists(work_path):
try:
os.remove(work_path)
except OSError:
pass
return None
# ===========================================================================
# Container sniffing / repair
# ===========================================================================
# Several backends silently ignore the requested opus format (Edge only emits
# MP3, Piper writes WAV, xAI writes MP3, some OpenAI-compatible servers ignore
# response_format="opus"), which breaks native voice bubbles. Sniff the magic
# bytes once after synthesis and repair when they don't match the extension.
def _sniff_audio_container(path: str) -> str:
"""Return a container id ('ogg', 'wav', 'mp3', 'flac', ...) or 'unknown'."""
from tools.audio_container import sniff_container
try:
with open(path, "rb") as fh:
head = fh.read(12)
except OSError:
return "unknown"
return sniff_container(head) or "unknown"
def _repair_ogg_container(file_str: str) -> str:
"""Ensure a path claiming ``.ogg`` actually contains an Ogg container.
MP3/WAV/FLAC bytes are transcoded in place to real Ogg/Opus. On failure
the file is renamed to its sniffed real extension so platforms get an
honest file instead of a 0-second voice bubble.
"""
if not file_str.endswith(".ogg"):
return file_str
container = _sniff_audio_container(file_str)
if container in ("ogg", "unknown"):
return file_str
logger.info(
"TTS wrote %s bytes into a .ogg path (%s) — transcoding to real Ogg/Opus",
container, file_str,
)
repaired = _ffmpeg_transcode_to_opus(file_str, file_str)
if repaired:
return repaired
honest = file_str[:-4] + "." + container
try:
os.replace(file_str, honest)
logger.warning(
"Could not transcode %s to Ogg/Opus — renamed to %s so the "
"file is delivered with its real format", file_str, honest,
)
return honest
except OSError:
return file_str
# ===========================================================================
# Long-form audio combination and delivery packing
# ===========================================================================
def _concat_audio_files(
audio_paths: List[str],
output_path: str,
*,
voice_compatible: bool = False,
) -> Optional[str]:
"""Combine independently encoded chunks with ffmpeg.
OGG/Opus is always decoded and re-encoded (even without voice opt-in);
matching MP3 chunks keep their encoded frames (``-c:a copy``). Structured
containers are never byte-joined. Returns ``None`` when ffmpeg is missing
or fails so callers keep the individually valid files.
"""
if not audio_paths:
raise ValueError("No audio chunks to combine")
if len(audio_paths) == 1:
source = audio_paths[0]
if os.path.abspath(source) != os.path.abspath(output_path):
shutil.copyfile(source, output_path)
return output_path
ffmpeg = shutil.which("ffmpeg")
if not ffmpeg:
return None
destination = Path(output_path)
destination.parent.mkdir(parents=True, exist_ok=True)
concat_path = destination.with_name(f".{destination.name}.{uuid.uuid4().hex}.concat.txt")
temp_output = destination.with_name(
f".{destination.stem}.{uuid.uuid4().hex}.combining{destination.suffix}"
)
try:
with concat_path.open("w", encoding="utf-8") as concat_file:
for path in audio_paths:
concat_file.write(f"file {shlex.quote(os.path.abspath(path))}\n")
command = [
ffmpeg, "-y", "-loglevel", "error", "-f", "concat", "-safe", "0",
"-i", str(concat_path), "-vn",
]
suffix = destination.suffix.lower()
if voice_compatible or suffix in {".ogg", ".opus"}:
command.extend(["-c:a", "libopus", "-ac", "1", "-b:a", "64k", "-vbr", "off"])
elif suffix == ".mp3" and all(Path(path).suffix.lower() == ".mp3" for path in audio_paths):
command.extend(["-c:a", "copy"])
command.append(str(temp_output))
result = subprocess.run(
command,
capture_output=True,
timeout=120,
stdin=subprocess.DEVNULL,
creationflags=windows_hide_flags(),
)
if result.returncode == 0 and temp_output.exists() and temp_output.stat().st_size > 0:
os.replace(temp_output, destination)
return str(destination)
logger.warning(
"ffmpeg audio combine failed: %s",
result.stderr.decode("utf-8", errors="ignore")[:500],
)
except (OSError, subprocess.TimeoutExpired) as exc:
logger.warning("ffmpeg audio combine failed: %s", exc)
finally:
for path in (concat_path, temp_output):
try:
path.unlink()
except OSError:
pass
return None
def _build_audio_delivery_files(
audio_paths: List[str],
output_path: str,
profile: AudioDeliveryProfile,
*,
voice_compatible: bool = False,
) -> Tuple[List[str], bool]:
"""Pack final-encoded chunks and enforce the hard upload limit.
Groups are packed against the conservative target, then every combined
artifact is checked at its real post-encoding size; an over-limit group is
split in half and retried. A failed combine returns the constituent files
separately. A single chunk above the hard limit fails closed. Returns
``(final_paths, combined_any)``.
"""
if not audio_paths:
raise ValueError("No final-encoded TTS audio chunks")
for path in audio_paths:
size = Path(path).stat().st_size
if size > profile.max_file_bytes:
raise ValueError(
f"Final-encoded TTS chunk exceeds {profile.platform} delivery "
f"limit ({size} > {profile.max_file_bytes} bytes): {path}"
)
base = Path(output_path)
scratch_outputs: List[str] = []
combined_any = False
combine_index = 0
def emit(group: List[str]) -> List[str]:
nonlocal combined_any, combine_index
if len(group) == 1:
return list(group)
combine_index += 1
scratch = base.with_name(
f".{base.stem}.delivery{combine_index:03d}.{uuid.uuid4().hex}{base.suffix}"
)
combined = _concat_audio_files(group, str(scratch), voice_compatible=voice_compatible)
if not combined:
return list(group)
scratch_outputs.append(combined)
if Path(combined).stat().st_size <= profile.max_file_bytes:
combined_any = True
return [combined]
try:
Path(combined).unlink()
except OSError:
pass
midpoint = max(1, len(group) // 2)
return emit(group[:midpoint]) + emit(group[midpoint:])
packed: List[str] = []
for group in _pack_audio_files_for_delivery(audio_paths, profile):
packed.extend(emit(group))
final_paths: List[str] = []
for index, source in enumerate(packed, start=1):
if len(packed) == 1:
destination = base
else:
source_suffix = Path(source).suffix or base.suffix
destination = base.with_name(f"{base.stem}.part{index:02d}{source_suffix}")
if os.path.abspath(source) != os.path.abspath(destination):
destination.parent.mkdir(parents=True, exist_ok=True)
os.replace(source, destination)
if destination.stat().st_size > profile.max_file_bytes:
raise ValueError(
f"Final TTS deliverable exceeds {profile.platform} delivery limit: {destination}"
)
final_paths.append(str(destination))
try:
return final_paths, combined_any
finally:
for scratch in scratch_outputs:
if scratch not in final_paths:
try:
Path(scratch).unlink()
except OSError:
pass
+262
View File
@@ -0,0 +1,262 @@
"""Local on-device TTS engines for ``tools.tts_tool``: NeuTTS, Piper, KittenTTS.
All three synthesize WAV natively; :func:`_finalize_wav_output` (shared) then
converts/renames to the caller's requested container. Piper and KittenTTS keep
their loaded models in small LRU caches registered in
``_LOCAL_TTS_MODEL_CACHES`` so the origin module's warm/release lifecycle can
pre-load or drop them. ``_import_piper`` / ``_import_kittentts`` are resolved
through the origin module at call time so test monkeypatches there apply.
"""
from __future__ import annotations
import logging
import subprocess
import sys
from pathlib import Path
from typing import Any, Callable, Dict, Tuple
from tools.tts_tool_delivery import _finalize_wav_output, _wav_sidecar_path
logger = logging.getLogger("tools.tts_tool")
DEFAULT_KITTENTTS_MODEL = "KittenML/kitten-tts-nano-0.8-int8" # 25MB
DEFAULT_KITTENTTS_VOICE = "Jasper"
DEFAULT_PIPER_VOICE = "en_US-lessac-medium" # balanced size/quality
def _origin():
from tools import tts_tool
return tts_tool
# ===========================================================================
# Bounded model caches
# ===========================================================================
# Each cached entry is a whole loaded model (tens of MB); an unbounded dict
# would pin one per distinct voice for the process lifetime. Small LRU: most
# sessions use one or two voices and a cold reload is cheap.
_TTS_MODEL_CACHE_MAX = 3
# Provider name → the model cache it populates. Consulted by
# warm_tts_provider() / release_tts_provider() in the origin module; a new
# local engine adds one row here plus a loader in _local_tts_warmers().
_LOCAL_TTS_MODEL_CACHES: Dict[str, Dict[str, Any]] = {}
# Piper voices keyed on absolute .onnx path (+cuda flag); KittenTTS on model name.
_piper_voice_cache: Dict[str, Any] = {}
_kittentts_model_cache: Dict[str, Any] = {}
_LOCAL_TTS_MODEL_CACHES["piper"] = _piper_voice_cache
_LOCAL_TTS_MODEL_CACHES["kittentts"] = _kittentts_model_cache
def _tts_cache_get_or_load(cache: Dict[str, Any], key: str, load: Callable[[], Any]) -> Any:
"""Get ``key`` from ``cache`` or load it, keeping the cache LRU-bounded.
A hit refreshes recency (pop + reinsert on the insertion-ordered dict); a
miss loads then evicts LRU entries beyond ``_TTS_MODEL_CACHE_MAX``. Callers
holding an evicted reference keep it alive; only the slot is released.
"""
if key in cache:
cache[key] = cache.pop(key)
return cache[key]
value = load()
cache[key] = value
while len(cache) > _TTS_MODEL_CACHE_MAX:
cache.pop(next(iter(cache)), None)
return value
# ===========================================================================
# NeuTTS (subprocess via tools/neutts_synth.py so the ~500MB model exits after use)
# ===========================================================================
def _default_neutts_ref_audio() -> str:
return str(Path(__file__).parent / "neutts_samples" / "jo.wav")
def _default_neutts_ref_text() -> str:
return str(Path(__file__).parent / "neutts_samples" / "jo.txt")
def _generate_neutts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
neutts_config = tts_config.get("neutts") or {}
wav_path = _wav_sidecar_path(output_path)
cmd = [
sys.executable, str(Path(__file__).parent / "neutts_synth.py"),
"--text", text,
"--out", wav_path,
"--ref-audio", neutts_config.get("ref_audio", "") or _default_neutts_ref_audio(),
"--ref-text", neutts_config.get("ref_text", "") or _default_neutts_ref_text(),
"--model", neutts_config.get("model", "neuphonic/neutts-air-q4-gguf"),
"--device", neutts_config.get("device", "cpu"),
]
result = subprocess.run(cmd, capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=120, stdin=subprocess.DEVNULL)
if result.returncode != 0:
# The synth script reports success lines as "OK:" on stderr too.
error_lines = [l for l in result.stderr.strip().splitlines() if not l.startswith("OK:")]
raise RuntimeError(f"NeuTTS synthesis failed: {chr(10).join(error_lines) or 'unknown error'}")
return _finalize_wav_output(wav_path, output_path)
# ===========================================================================
# Piper (local neural VITS, 44 languages)
# ===========================================================================
def _get_piper_voices_dir() -> Path:
"""``<HERMES_HOME>/cache/piper-voices/`` so voice downloads follow profile boundaries."""
from hermes_constants import get_hermes_dir
root = Path(get_hermes_dir("cache/piper-voices", "piper_voices_cache"))
root.mkdir(parents=True, exist_ok=True)
return root
def _resolve_piper_voice_path(voice: str, download_dir: Path) -> str:
"""Resolve *voice* (an .onnx path or a voice name) to a concrete .onnx file.
Names like ``en_US-lessac-medium`` are downloaded into *download_dir* on
first use via ``python -m piper.download_voices``. Raises RuntimeError
when the model can't be located or downloaded.
"""
if not voice:
voice = DEFAULT_PIPER_VOICE
candidate = Path(voice).expanduser()
if candidate.suffix.lower() == ".onnx" and candidate.exists():
return str(candidate)
cached = download_dir / f"{voice}.onnx"
if cached.exists() and (download_dir / f"{voice}.onnx.json").exists():
return str(cached)
logger.info("[Piper] Downloading voice '%s' to %s (first use)", voice, download_dir)
try:
result = subprocess.run(
[sys.executable, "-m", "piper.download_voices", voice,
"--download-dir", str(download_dir)],
capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=300,
stdin=subprocess.DEVNULL,
)
except subprocess.TimeoutExpired as exc:
raise RuntimeError(f"Piper voice download timed out after 300s for '{voice}'") from exc
if result.returncode != 0:
stderr = (result.stderr or "").strip() or "no stderr output"
raise RuntimeError(f"Piper voice download failed for '{voice}': {stderr[:400]}")
if not cached.exists():
raise RuntimeError(
f"Piper voice download completed but {cached} is missing — "
f"check voice name (see: https://github.com/OHF-Voice/piper1-gpl/"
f"blob/main/docs/VOICES.md)"
)
return str(cached)
def _load_piper_voice_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, Dict[str, Any]]:
"""Resolve + load (or fetch from cache) the Piper voice ``tts_config`` selects.
Shared by synthesis and ``warm_tts_provider`` so a warm-up fills exactly
the cache slot the next synthesis hits. Returns ``(voice, piper_config)``.
"""
PiperVoice = _origin()._import_piper()
piper_config = tts_config.get("piper") or {} if isinstance(tts_config, dict) else {}
voice_name = piper_config.get("voice") or DEFAULT_PIPER_VOICE
download_dir = Path(piper_config.get("voices_dir") or _get_piper_voices_dir()).expanduser()
download_dir.mkdir(parents=True, exist_ok=True)
use_cuda = bool(piper_config.get("use_cuda", False))
model_path = _resolve_piper_voice_path(voice_name, download_dir)
# speaker_id is applied per call via syn_config, so one PiperVoice
# instance serves every speaker and stays out of the cache key.
cache_key = f"{model_path}::cuda={use_cuda}"
def _load_piper_voice():
logger.info("[Piper] Loading voice: %s", model_path)
v = PiperVoice.load(model_path, use_cuda=use_cuda)
logger.info("[Piper] Voice loaded")
return v
voice = _tts_cache_get_or_load(_piper_voice_cache, cache_key, _load_piper_voice)
return voice, piper_config
def _generate_piper_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
import wave
voice, piper_config = _load_piper_voice_for_config(tts_config)
# Bad speaker_id input drops to 0 (Piper's default); booleans are rejected
# outright since True/False would silently coerce to 1/0.
_raw_speaker = piper_config.get("speaker_id", 0)
speaker_id = 0 if isinstance(_raw_speaker, bool) or not isinstance(_raw_speaker, int) else _raw_speaker
# Only build a SynthesisConfig when an advanced knob is configured, so we
# don't depend on a newer piper-tts than the user's unless we must.
syn_config = None
has_advanced = any(
k in piper_config
for k in ("length_scale", "noise_scale", "noise_w_scale", "volume", "normalize_audio", "speaker_id")
)
if has_advanced:
try:
from piper import SynthesisConfig # type: ignore
syn_config = SynthesisConfig(
length_scale=float(piper_config.get("length_scale", 1.0)),
noise_scale=float(piper_config.get("noise_scale", 0.667)),
noise_w_scale=float(piper_config.get("noise_w_scale", 0.8)),
volume=float(piper_config.get("volume", 1.0)),
normalize_audio=bool(piper_config.get("normalize_audio", True)),
speaker_id=speaker_id,
)
except ImportError:
logger.warning(
"[Piper] SynthesisConfig not available in this piper-tts "
"version — advanced knobs ignored"
)
wav_path = _wav_sidecar_path(output_path)
with wave.open(wav_path, "wb") as wav_file:
if syn_config is not None:
voice.synthesize_wav(text, wav_file, syn_config=syn_config)
else:
voice.synthesize_wav(text, wav_file)
return _finalize_wav_output(wav_path, output_path)
# ===========================================================================
# KittenTTS (local ONNX, 25-80MB models, CPU only)
# ===========================================================================
def _load_kittentts_model_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, Dict[str, Any]]:
"""Load (or fetch from cache) the KittenTTS model; returns ``(model, kittentts_config)``."""
KittenTTS = _origin()._import_kittentts()
kt_config = tts_config.get("kittentts", {}) if isinstance(tts_config, dict) else {}
kt_config = kt_config or {}
model_name = kt_config.get("model", DEFAULT_KITTENTTS_MODEL)
def _load_kittentts_model():
logger.info("[KittenTTS] Loading model: %s", model_name)
m = KittenTTS(model_name)
logger.info("[KittenTTS] Model loaded successfully")
return m
model = _tts_cache_get_or_load(_kittentts_model_cache, model_name, _load_kittentts_model)
return model, kt_config
def _generate_kittentts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
model, kt_config = _load_kittentts_model_for_config(tts_config)
audio = model.generate(
text,
voice=kt_config.get("voice", DEFAULT_KITTENTTS_VOICE),
speed=kt_config.get("speed", 1.0),
clean_text=kt_config.get("clean_text", True),
) # numpy array at 24kHz
import soundfile as sf
wav_path = _wav_sidecar_path(output_path)
sf.write(wav_path, audio, 24000)
return _finalize_wav_output(wav_path, output_path)
+896
View File
@@ -0,0 +1,896 @@
"""Cloud TTS backends for ``tools.tts_tool``: Edge, ElevenLabs, xAI, MiniMax, Mistral, Gemini.
Each ``_generate_<provider>(text, output_path, tts_config) -> path`` writes one
final-encoded file. Shared here: bounded upstream response reading (16 MiB
cap so a hostile endpoint can't feed unbounded audio) and the auxiliary-model
speech-tag rewrites. OpenAI/DeepInfra stay in the origin module (they share the
managed-gateway selection logic). Seams tests monkeypatch on the origin
(``get_env_value``, ``_resolve_provider_key``, ``_import_*``) are resolved
through :func:`_origin` at call time so those patches keep applying.
"""
from __future__ import annotations
import base64
import json
import logging
import os
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Dict, Optional
from urllib.parse import urlparse
from tools.tts_tool_delivery import _wrap_pcm_as_wav, _write_wav_bytes_as
from tools.xai_http import hermes_xai_user_agent
logger = logging.getLogger("tools.tts_tool")
def _origin():
"""``tools.tts_tool``, resolved per call so monkeypatched seams there still apply."""
from tools import tts_tool
return tts_tool
# ===========================================================================
# Defaults
# ===========================================================================
DEFAULT_EDGE_VOICE = "en-US-AriaNeural"
DEFAULT_ELEVENLABS_VOICE_ID = "pNInz6obpgDQGcFmaJgB" # Adam
DEFAULT_ELEVENLABS_MODEL_ID = "eleven_multilingual_v2"
DEFAULT_ELEVENLABS_STREAMING_MODEL_ID = "eleven_flash_v2_5"
DEFAULT_MINIMAX_MODEL = "speech-02-hd"
DEFAULT_MINIMAX_VOICE_ID = "English_expressive_narrator"
DEFAULT_MINIMAX_BASE_URL = "https://api.minimax.io/v1/t2a_v2"
DEFAULT_MINIMAX_CN_BASE_URL = "https://api.minimaxi.com/v1/t2a_v2"
DEFAULT_MISTRAL_TTS_MODEL = "voxtral-mini-tts-2603"
DEFAULT_MISTRAL_TTS_VOICE_ID = "c69964a6-ab8b-4f8a-9465-ec0925096ec8" # Paul - Neutral
DEFAULT_XAI_VOICE_ID = "eve"
DEFAULT_XAI_LANGUAGE = "en"
DEFAULT_XAI_SAMPLE_RATE = 24000
DEFAULT_XAI_BIT_RATE = 128000
DEFAULT_XAI_AUTO_SPEECH_TAGS = False
DEFAULT_XAI_BASE_URL = "https://api.x.ai/v1"
# xAI `speed` accepts 0.7..1.5 (1.0 = API default, omitted from the payload).
DEFAULT_XAI_SPEED_MIN = 0.7
DEFAULT_XAI_SPEED_MAX = 1.5
DEFAULT_XAI_SPEED_DEFAULT = 1.0
# xAI `optimize_streaming_latency` is 0/1/2; >0 trades quality for time-to-first-audio.
DEFAULT_XAI_OPTIMIZE_STREAMING_LATENCY_DEFAULT = 0
# xAI `text_normalization` speaks numbers/abbreviations/symbols in written form when True.
DEFAULT_XAI_TEXT_NORMALIZATION_DEFAULT = False
DEFAULT_GEMINI_TTS_MODEL = "gemini-2.5-flash-preview-tts"
DEFAULT_GEMINI_TTS_VOICE = "Kore"
DEFAULT_GEMINI_TTS_BASE_URL = "https://generativelanguage.googleapis.com/v1beta"
DEFAULT_GEMINI_AUDIO_TAGS = False
GEMINI_AUDIO_TAG_REWRITE_TASK = "tts_audio_tags"
TTS_RESPONSE_BODY_LIMIT_BYTES = 16 * 1024 * 1024
TTS_RESPONSE_BODY_CHUNK_BYTES = 64 * 1024
def _config_bool(value: Any, default: bool = False) -> bool:
"""Coerce common YAML/env bool spellings without treating random strings as true."""
if isinstance(value, bool):
return value
if value is None:
return default
if isinstance(value, (int, float)):
return bool(value)
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in {"1", "true", "yes", "on", "enabled"}:
return True
if normalized in {"0", "false", "no", "off", "disabled"}:
return False
return default
def _tts_response_format_from_path(output_path: str) -> str:
"""Pick an OpenAI-style response format (opus/wav/flac/mp3) from the output extension."""
for ext, fmt in ((".ogg", "opus"), (".wav", "wav"), (".flac", "flac")):
if output_path.endswith(ext):
return fmt
return "mp3"
# ===========================================================================
# Bounded upstream response reading
# ===========================================================================
def _response_has_explicit_stream(response: Any) -> bool:
"""True for real ``requests`` responses (or doubles defining ``iter_content`` themselves)."""
iter_content = getattr(response, "iter_content", None)
if not callable(iter_content):
return False
response_type = type(response)
if response_type.__module__.startswith("requests."):
return True
return "iter_content" in vars(response_type)
def _close_response(response: Any) -> None:
close = getattr(response, "close", None)
if callable(close):
try:
close()
except Exception:
pass
def _read_tts_response_bytes(
response: Any,
*,
label: str,
limit: Optional[int] = None,
) -> bytes:
"""Read an upstream TTS response with a hard byte cap."""
limit = TTS_RESPONSE_BODY_LIMIT_BYTES if limit is None else limit
chunks: list[bytes] = []
total = 0
try:
if _response_has_explicit_stream(response):
iterator = response.iter_content(chunk_size=TTS_RESPONSE_BODY_CHUNK_BYTES)
else:
content = vars(response).get("content", getattr(type(response), "content", b""))
if isinstance(content, str):
content = content.encode("utf-8", errors="replace")
iterator = (content,) if isinstance(content, (bytes, bytearray)) else ()
for chunk in iterator:
if not chunk:
continue
if isinstance(chunk, str):
chunk = chunk.encode("utf-8", errors="replace")
chunk = bytes(chunk)
total += len(chunk)
if total > limit:
_close_response(response)
raise RuntimeError(f"{label} response exceeds {limit} bytes")
chunks.append(chunk)
return b"".join(chunks)
finally:
_close_response(response)
def _read_tts_response_json(
response: Any,
*,
label: str,
limit: Optional[int] = None,
) -> Dict[str, Any]:
raw = _read_tts_response_bytes(response, label=label, limit=limit)
if raw:
return json.loads(raw.decode("utf-8"))
# Unit-test doubles often only provide `.json()`. Real requests.Response
# objects took the streaming path above, so this never re-opens eager
# buffering in production.
if not _response_has_explicit_stream(response):
json_reader = getattr(response, "json", None)
if callable(json_reader):
parsed = json_reader()
return parsed if isinstance(parsed, dict) else {}
return {}
def _write_tts_response_to_file(
response: Any,
output_path: str,
*,
label: str,
limit: Optional[int] = None,
) -> None:
audio_bytes = _read_tts_response_bytes(response, label=label, limit=limit)
with open(output_path, "wb") as f:
f.write(audio_bytes)
def _extract_auxiliary_message_content(response: Any) -> str:
try:
choice = response.choices[0]
message = getattr(choice, "message", None)
if isinstance(message, dict):
return str(message.get("content") or "")
return str(getattr(message, "content", "") or "")
except Exception:
return ""
def _strip_code_fence(content: str) -> str:
"""Unwrap a ```fenced``` LLM reply; returns the stripped inner text."""
clean = (content or "").strip()
fence = re.fullmatch(r"```(?:[A-Za-z0-9_-]+)?\s*(.*?)\s*```", clean, flags=re.DOTALL)
return fence.group(1).strip() if fence else clean
# ===========================================================================
# Provider: Edge TTS (free default)
# ===========================================================================
async def _generate_edge_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
_edge_tts = _origin()._import_edge_tts()
edge_config = tts_config.get("edge") or {}
voice = edge_config.get("voice", DEFAULT_EDGE_VOICE)
speed = float(edge_config.get("speed", tts_config.get("speed", 1.0)))
kwargs = {"voice": voice}
if speed != 1.0:
pct = round((speed - 1.0) * 100)
kwargs["rate"] = f"{pct:+d}%"
communicate = _edge_tts.Communicate(text, **kwargs)
await communicate.save(output_path)
return output_path
# ===========================================================================
# Provider: ElevenLabs
# ===========================================================================
def _elevenlabs_environment_kwargs(el_config: Dict[str, Any]) -> Dict[str, Any]:
"""Client kwargs redirecting the SDK to ``tts.elevenlabs.base_url``/``wss_url``.
Empty when no base_url is set (SDK default environment). ``wss_url``
defaults to the base_url host with a ``ws(s)://`` scheme.
"""
base_url = (el_config.get("base_url") or "").rstrip("/")
if not base_url:
return {}
wss_url = (el_config.get("wss_url") or "").rstrip("/")
if not wss_url:
wss_url = re.sub(r"^http", "ws", base_url)
from elevenlabs.environment import ElevenLabsEnvironment
return {"environment": ElevenLabsEnvironment(base=base_url, wss=wss_url)}
def _generate_elevenlabs(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
origin = _origin()
api_key = (origin._resolve_provider_key("ELEVENLABS_API_KEY", "elevenlabs") or "")
if not api_key:
raise ValueError("ELEVENLABS_API_KEY not set. Get one at https://elevenlabs.io/")
el_config = tts_config.get("elevenlabs") or {}
voice_id = el_config.get("voice_id", DEFAULT_ELEVENLABS_VOICE_ID)
model_id = el_config.get("model_id", DEFAULT_ELEVENLABS_MODEL_ID)
output_format = "opus_48000_64" if output_path.endswith(".ogg") else "mp3_44100_128"
ElevenLabs = origin._import_elevenlabs()
client = ElevenLabs(api_key=api_key, **_elevenlabs_environment_kwargs(el_config))
audio_generator = client.text_to_speech.convert(
text=text,
voice_id=voice_id,
model_id=model_id,
output_format=output_format,
)
with open(output_path, "wb") as f:
for chunk in audio_generator:
f.write(chunk)
return output_path
# ===========================================================================
# Provider: xAI TTS (dedicated /v1/tts endpoint, not the OpenAI audio shape)
# ===========================================================================
_XAI_INLINE_SPEECH_TAGS = (
"pause", "long-pause", "hum-tune", "laugh", "chuckle", "giggle", "cry", "tsk",
"tongue-click", "lip-smack", "breath", "inhale", "exhale", "sigh",
)
_XAI_WRAPPING_SPEECH_TAGS = (
"soft", "whisper", "loud", "build-intensity", "decrease-intensity", "higher-pitch",
"lower-pitch", "slow", "fast", "sing-song", "singing", "laugh-speak", "emphasis",
)
_XAI_SPEECH_TAG_RE = re.compile(
r"(\[(?:" + "|".join(_XAI_INLINE_SPEECH_TAGS) + r")\]|</?(?:" + "|".join(_XAI_WRAPPING_SPEECH_TAGS) + r")>)",
flags=re.IGNORECASE,
)
_XAI_FIRST_SENTENCE_RE = re.compile(r"^(.{12,120}?[.!?…])\s+(?=\S)", flags=re.DOTALL)
def _xai_bool_config(value: Any, default: bool = False) -> bool:
return _config_bool(value, default=default)
def _apply_xai_auto_speech_tags(text: str) -> str:
"""Add xAI speech tags for more natural voice-mode replies.
Local conservative pass first ([pause] between paragraphs and after the
first sentence). If the text carried no explicit speech tags already, the
auxiliary model then rewrites it with the richer xAI tag set; any failure
falls back to the locally tagged text.
"""
clean = text.strip()
if not clean:
return text
local = re.sub(r"\n\s*\n+", " [pause] ", clean)
local = re.sub(r"\s*\n\s*", " ", local)
if not _XAI_SPEECH_TAG_RE.search(local):
local = _XAI_FIRST_SENTENCE_RE.sub(r"\1 [pause] ", local, count=1)
local = re.sub(r"\s{2,}", " ", local).strip()
# Explicit user/model tags are trusted as-is.
if _XAI_SPEECH_TAG_RE.search(clean):
return local
inline = ", ".join(_XAI_INLINE_SPEECH_TAGS)
wrapping = ", ".join(_XAI_WRAPPING_SPEECH_TAGS)
system_prompt = (
"You rewrite transcripts for the xAI /v1/tts endpoint by inserting "
"expressive speech tags.\n\n"
"Valid inline tags (use as `[tag]`): " + inline + ".\n"
"Valid wrapping tags (use as `[tag]...[/tag]`): " + wrapping + ".\n\n"
"Rules:\n"
"- Preserve the spoken words, order, and meaning.\n"
"- Do not add new spoken sentences or remove existing spoken words.\n"
"- Use inline `[tag]` for short modifiers (laughs, sighs, pause, etc.).\n"
"- Use wrapping `[tag]...[/tag]` for sustained effects (whisper, soft, slow, fast, loud, etc.).\n"
"- Do not use angle-bracket tags like `<tag>...</tag>` — xAI uses BBCode-style closing tags with `[/tag]`.\n"
"- Do not use SSML.\n"
"- Do not explain or comment.\n"
"- Return only the tagged TTS script."
)
try:
from agent.auxiliary_client import call_llm
response = call_llm(
task="tts_audio_tags",
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": f"TRANSCRIPT TO TAG:\n{local}"},
],
temperature=0.7,
)
tagged = _strip_code_fence(_extract_auxiliary_message_content(response))
return tagged or local
except Exception as exc:
logger.debug("xAI TTS audio tag rewrite failed; using locally-tagged text: %s", exc)
return local
def _clamped_number(raw: Any, cast, lo, hi):
"""Parse an optional numeric knob and clamp into [lo, hi]; ``None``/unparseable -> None.
Mirrors the historical inline logic exactly, including that an empty
string is passed to the clamp unconverted (a TypeError the caller's
generic handler reports as a TTS failure).
"""
if raw is None:
return None
if raw != "":
try:
raw = cast(raw)
except (TypeError, ValueError):
return None
return max(lo, min(hi, raw))
def _generate_xai_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
import requests
from tools.xai_http import resolve_xai_http_credentials
# TTS is API-billed: a subscription OAuth bearer can authorize chat while
# returning 403 for /v1/tts, so prefer an explicit XAI_API_KEY with OAuth
# as the fallback.
creds = resolve_xai_http_credentials(prefer_api_key=True)
api_key = str(creds.get("api_key") or "").strip()
if not api_key:
raise ValueError("No xAI credentials found. Configure xAI OAuth in `hermes model` or set XAI_API_KEY.")
xai_config = tts_config.get("xai") or {}
voice_id = str(xai_config.get("voice_id", DEFAULT_XAI_VOICE_ID)).strip() or DEFAULT_XAI_VOICE_ID
language = str(xai_config.get("language", DEFAULT_XAI_LANGUAGE)).strip() or DEFAULT_XAI_LANGUAGE
sample_rate = int(xai_config.get("sample_rate", DEFAULT_XAI_SAMPLE_RATE))
bit_rate = int(xai_config.get("bit_rate", DEFAULT_XAI_BIT_RATE))
auto_speech_tags = _xai_bool_config(
xai_config.get("auto_speech_tags", xai_config.get("speech_tags")),
DEFAULT_XAI_AUTO_SPEECH_TAGS,
)
# ``tts.xai.speed`` overrides global ``tts.speed``; out-of-range values are
# clamped into the API's 0.7..1.5 band rather than 400ing the request.
speed = _clamped_number(
xai_config.get("speed", tts_config.get("speed")),
float, DEFAULT_XAI_SPEED_MIN, DEFAULT_XAI_SPEED_MAX,
)
optimize_streaming_latency = _clamped_number(
xai_config.get("optimize_streaming_latency", tts_config.get("optimize_streaming_latency")),
int, 0, 2,
)
text_normalization = _xai_bool_config(
xai_config.get("text_normalization"),
DEFAULT_XAI_TEXT_NORMALIZATION_DEFAULT,
)
if auto_speech_tags:
text = _apply_xai_auto_speech_tags(text)
if creds.get("provider") == "xai-oauth":
base_url = str(creds.get("base_url") or DEFAULT_XAI_BASE_URL).strip().rstrip("/")
else:
base_url = str(
xai_config.get("base_url")
or creds.get("base_url")
or _origin().get_env_value("XAI_BASE_URL")
or DEFAULT_XAI_BASE_URL
).strip().rstrip("/")
# Send the documented minimal POST /v1/tts shape; optional fields are
# attached only when they differ from the API defaults.
codec = "wav" if output_path.endswith(".wav") else "mp3"
payload: Dict[str, Any] = {
"text": text,
"voice_id": voice_id,
"language": language,
}
if (
codec != "mp3"
or sample_rate != DEFAULT_XAI_SAMPLE_RATE
or (codec == "mp3" and bit_rate != DEFAULT_XAI_BIT_RATE)
):
output_format: Dict[str, Any] = {"codec": codec}
if sample_rate:
output_format["sample_rate"] = sample_rate
if codec == "mp3" and bit_rate:
output_format["bit_rate"] = bit_rate
payload["output_format"] = output_format
if speed is not None and speed != DEFAULT_XAI_SPEED_DEFAULT:
payload["speed"] = speed
if (
optimize_streaming_latency is not None
and optimize_streaming_latency != DEFAULT_XAI_OPTIMIZE_STREAMING_LATENCY_DEFAULT
):
payload["optimize_streaming_latency"] = optimize_streaming_latency
if text_normalization:
payload["text_normalization"] = True
response = requests.post(
f"{base_url}/tts",
headers={
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
"User-Agent": hermes_xai_user_agent(),
},
json=payload,
timeout=60,
stream=True,
)
response.raise_for_status()
_write_tts_response_to_file(response, output_path, label="xAI TTS")
return output_path
# ===========================================================================
# Provider: MiniMax TTS
# ===========================================================================
@dataclass(frozen=True)
class _MiniMaxTTSRuntime:
"""A region-bound MiniMax endpoint and credential (key excluded from ``repr``)."""
region: str
endpoint: str
credential_source: str
api_key: str = field(repr=False)
def _resolve_minimax_tts_runtime(
tts_config: Dict[str, Any],
) -> _MiniMaxTTSRuntime:
"""Select MiniMax TTS region, endpoint, and credential atomically.
An explicit ``tts.minimax.region`` wins. Without one, the legacy global
credential wins when present; a China credential is selected only when it
is the sole configured MiniMax credential.
"""
mm_config = tts_config.get("minimax", {})
if not isinstance(mm_config, dict):
mm_config = {}
resolve_key = _origin()._resolve_provider_key
credentials = {
"global": ("MINIMAX_API_KEY", str(resolve_key("MINIMAX_API_KEY", "minimax") or "").strip()),
"cn": ("MINIMAX_CN_API_KEY", str(resolve_key("MINIMAX_CN_API_KEY", "minimax") or "").strip()),
}
endpoints = {"global": DEFAULT_MINIMAX_BASE_URL, "cn": DEFAULT_MINIMAX_CN_BASE_URL}
configured_region = str(mm_config.get("region") or "").strip().lower()
if configured_region and configured_region not in endpoints:
raise ValueError("tts.minimax.region must be 'global' or 'cn'")
if configured_region:
region = configured_region
elif credentials["global"][1]:
region = "global"
elif credentials["cn"][1]:
region = "cn"
else:
region = "global"
credential_source, api_key = credentials[region]
if not api_key:
raise ValueError(f"{credential_source} not set for MiniMax TTS region {region!r}")
endpoint = str(mm_config.get("base_url") or endpoints[region]).strip()
endpoint_host = (urlparse(endpoint).hostname or "").lower()
official_region_hosts = {
"global": frozenset({"api.minimax.io", "api.minimax.chat"}),
"cn": frozenset({"api.minimaxi.com"}),
}
other_region = "cn" if region == "global" else "global"
if endpoint_host in official_region_hosts[other_region]:
raise ValueError(
f"tts.minimax.base_url points to the {other_region!r} MiniMax endpoint "
f"but region is {region!r}"
)
return _MiniMaxTTSRuntime(
region=region,
endpoint=endpoint,
credential_source=credential_source,
api_key=api_key,
)
def _raise_minimax_api_error(result: Dict[str, Any]) -> None:
base_resp = result.get("base_resp", {})
status_code = base_resp.get("status_code", -1)
if status_code != 0:
status_msg = base_resp.get("status_msg", "unknown error")
raise RuntimeError(f"MiniMax TTS API error (code {status_code}): {status_msg}")
def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
"""Generate audio via MiniMax.
Two endpoints, detected from the URL: ``t2a_v2`` (nested payload, JSON
reply with hex-encoded audio) and legacy ``text_to_speech`` (flat payload,
raw ``audio/*`` body).
"""
import requests
runtime = _resolve_minimax_tts_runtime(tts_config)
mm_config = tts_config.get("minimax", {})
if not isinstance(mm_config, dict):
mm_config = {}
model = mm_config.get("model", DEFAULT_MINIMAX_MODEL)
voice_id = mm_config.get("voice_id", DEFAULT_MINIMAX_VOICE_ID)
base_url = runtime.endpoint
# MiniMax accounts scope TTS requests by GroupId (``?GroupId=<id>`` on the
# t2a_v2 URL). Config or MINIMAX_GROUP_ID; only attach when absent from the URL.
group_id = (
str(mm_config.get("group_id") or "").strip()
or (_origin().get_env_value("MINIMAX_GROUP_ID") or "").strip()
)
if group_id and "GroupId=" not in base_url:
sep = "&" if "?" in base_url else "?"
base_url = f"{base_url}{sep}GroupId={group_id}"
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {runtime.api_key}",
}
is_t2a_v2 = "t2a_v2" in base_url
if is_t2a_v2:
payload = {
"model": model,
"text": text,
"voice_setting": {
"voice_id": voice_id,
"speed": mm_config.get("speed", 1.0),
"vol": mm_config.get("vol", 1.0),
"pitch": mm_config.get("pitch", 0),
"emotion": mm_config.get("emotion", "neutral"),
},
"audio_setting": {
"sample_rate": mm_config.get("sample_rate", 32000),
"bitrate": mm_config.get("bitrate", 128000),
"format": "mp3",
"channel": 1,
},
}
else:
payload = {"model": model, "text": text, "voice_id": voice_id}
response = requests.post(base_url, json=payload, headers=headers, timeout=60, stream=True)
if is_t2a_v2:
response.raise_for_status()
result = _read_tts_response_json(response, label="MiniMax TTS")
_raise_minimax_api_error(result)
hex_audio = result.get("data", {}).get("audio", "")
if not hex_audio:
raise RuntimeError("MiniMax TTS returned empty audio data")
with open(output_path, "wb") as f:
f.write(bytes.fromhex(hex_audio))
return output_path
content_type = response.headers.get("Content-Type", "")
if "audio/" in content_type:
_write_tts_response_to_file(response, output_path, label="MiniMax TTS")
return output_path
# Non-audio reply: surface the API error if the body is JSON.
raw_body = b""
try:
raw_body = _read_tts_response_bytes(response, label="MiniMax TTS")
result = json.loads(raw_body.decode("utf-8")) if raw_body else {}
_raise_minimax_api_error(result)
except (json.JSONDecodeError, UnicodeDecodeError, TypeError):
response.raise_for_status()
raise RuntimeError(
f"MiniMax TTS returned unexpected Content-Type '{content_type}' "
f"({len(raw_body)} bytes)"
)
raise RuntimeError("MiniMax TTS returned no audio data")
# ===========================================================================
# Provider: Mistral (Voxtral TTS) — base64 audio, native Opus for voice bubbles
# ===========================================================================
def _generate_mistral_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
origin = _origin()
api_key = (origin._resolve_provider_key("MISTRAL_API_KEY", "mistral") or "")
if not api_key:
raise ValueError("MISTRAL_API_KEY not set. Get one at https://console.mistral.ai/")
mi_config = tts_config.get("mistral") or {}
model = mi_config.get("model", DEFAULT_MISTRAL_TTS_MODEL)
voice_id = mi_config.get("voice_id") or DEFAULT_MISTRAL_TTS_VOICE_ID
base_url = mi_config.get("base_url") # the Mistral SDK calls it server_url
Mistral = origin._import_mistral_client()
client_kwargs: Dict[str, Any] = {"api_key": api_key}
if base_url:
client_kwargs["server_url"] = base_url
try:
with Mistral(**client_kwargs) as client:
response = client.audio.speech.complete(
model=model,
input=text,
voice_id=voice_id,
response_format=_tts_response_format_from_path(output_path),
)
audio_bytes = base64.b64decode(response.audio_data)
except ValueError:
raise
except Exception as e:
logger.error("Mistral TTS failed: %s", e, exc_info=True)
raise RuntimeError(f"Mistral TTS failed: {type(e).__name__}") from e
with open(output_path, "wb") as f:
f.write(audio_bytes)
return output_path
# ===========================================================================
# Provider: Google Gemini TTS
# ===========================================================================
def _resolve_gemini_persona_prompt_path(gemini_config: Dict[str, Any]) -> Optional[Path]:
"""``tts.gemini.persona_prompt_file`` as a Path (relative -> under HERMES_HOME), or None."""
raw = gemini_config.get("persona_prompt_file")
if not isinstance(raw, str) or not raw.strip():
return None
path = Path(os.path.expandvars(raw.strip())).expanduser()
if not path.is_absolute():
try:
from hermes_constants import get_hermes_home
path = get_hermes_home() / path
except Exception:
path = Path.cwd() / path
return path
def _read_gemini_persona_prompt(gemini_config: Dict[str, Any]) -> str:
"""Read the Gemini persona prompt file, failing soft on config mistakes."""
path = _resolve_gemini_persona_prompt_path(gemini_config)
if path is None:
return ""
try:
return path.read_text(encoding="utf-8").strip()
except (OSError, UnicodeDecodeError) as exc:
logger.warning("Gemini TTS persona prompt file unavailable at %s: %s", path, exc)
return ""
def _gemini_model_supports_audio_tags(model: str) -> bool:
"""Only Gemini 3.1 TTS models are known to honor expressive audio tags."""
normalized = (model or "").strip().lower().rsplit("/", 1)[-1]
return "gemini-3.1" in normalized and "tts" in normalized
def _gemini_audio_tags_enabled(gemini_config: Dict[str, Any], model: str) -> bool:
raw = gemini_config.get("audio_tags")
if isinstance(raw, dict):
raw = raw.get("enabled")
if not _config_bool(raw, default=DEFAULT_GEMINI_AUDIO_TAGS):
return False
if not _gemini_model_supports_audio_tags(model):
logger.warning(
"Gemini TTS audio_tags enabled, but model %s is not known to support "
"Gemini audio tags; skipping hidden tag rewrite",
model,
)
return False
return True
def _rewrite_gemini_tts_audio_tags(text: str, persona_prompt: str = "") -> str:
"""Use the configured auxiliary model to insert Gemini audio tags (falls back to *text*)."""
transcript = text.strip()
if not transcript:
return text
system_prompt = (
"You rewrite transcripts for Gemini 3.1 Flash TTS by inserting expressive "
"audio tags.\n\n"
"Audio tags are inline square-bracket modifiers such as [whispers], "
"[excitedly], [very slow], [sarcastically], [laughs], [sighs], or [gasp]. "
"There is no fixed allowlist. Use creative freeform tags generously but "
"naturally to control tone, pace, emotional vibe, emphasis, section-level "
"delivery, and non-verbal sounds. Use English audio tags even when the "
"spoken transcript is not English.\n\n"
"Rules:\n"
"- Preserve the spoken words, order, and meaning.\n"
"- Do not add new spoken sentences or remove existing spoken words.\n"
"- Use square brackets for every audio tag.\n"
"- Do not use SSML or XML tags.\n"
"- Do not explain or comment.\n"
"- Return only the tagged TTS script."
)
context = persona_prompt.strip() or "(none)"
user_prompt = f"PERSONA AND DIRECTOR CONTEXT:\n{context}\n\nTRANSCRIPT TO TAG:\n{transcript}"
try:
from agent.auxiliary_client import call_llm
response = call_llm(
task=GEMINI_AUDIO_TAG_REWRITE_TASK,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
temperature=0.7,
)
tagged = _strip_code_fence(_extract_auxiliary_message_content(response))
return tagged or text
except Exception as exc:
logger.warning("Gemini TTS audio tag rewrite failed; using untagged text: %s", exc)
return text
def _compose_gemini_tts_prompt(
text: str,
gemini_config: Dict[str, Any],
persona_prompt: Optional[str] = None,
) -> str:
"""Build the Gemini prompt from persona direction plus the live transcript.
A ``{transcript}`` / ``{{transcript}}`` placeholder in the persona prompt is
substituted in place; otherwise the transcript is appended under a heading.
"""
transcript = text.strip()
if persona_prompt is None:
persona_prompt = _read_gemini_persona_prompt(gemini_config)
if not persona_prompt:
return transcript
preamble = (
"Synthesize speech from the TRANSCRIPT only. Treat AUDIO PROFILE, "
"SCENE, DIRECTOR'S NOTES, and SAMPLE CONTEXT as performance direction; "
"do not speak those sections aloud."
)
for pattern in (r"\{\{\s*transcript\s*\}\}", r"\{\s*transcript\s*\}"):
compiled = re.compile(pattern, flags=re.IGNORECASE)
if compiled.search(persona_prompt):
return f"{preamble}\n\n{compiled.sub(transcript, persona_prompt)}".strip()
return f"{preamble}\n\n{persona_prompt}\n\n#### TRANSCRIPT\n{transcript}".strip()
def _generate_gemini_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
"""Generate audio via Gemini ``generateContent`` with ``responseModalities=["AUDIO"]``.
The API returns raw 24kHz mono 16-bit PCM as base64; it is wrapped as WAV
and ffmpeg-converted to MP3/Opus when the caller asked for those (no
ffmpeg -> the WAV is written under the requested name, same as NeuTTS).
"""
import requests
origin = _origin()
api_key = (
origin._resolve_provider_key("GEMINI_API_KEY", "gemini")
or origin._resolve_provider_key("GOOGLE_API_KEY", "gemini")
)
if not api_key:
raise ValueError(
"GEMINI_API_KEY not set. Get one at https://aistudio.google.com/app/apikey"
)
raw_gemini_config = tts_config.get("gemini") or {}
gemini_config = raw_gemini_config if isinstance(raw_gemini_config, dict) else {}
model = str(gemini_config.get("model", DEFAULT_GEMINI_TTS_MODEL)).strip() or DEFAULT_GEMINI_TTS_MODEL
voice = str(gemini_config.get("voice", DEFAULT_GEMINI_TTS_VOICE)).strip() or DEFAULT_GEMINI_TTS_VOICE
base_url = str(
gemini_config.get("base_url")
or origin.get_env_value("GEMINI_BASE_URL")
or DEFAULT_GEMINI_TTS_BASE_URL
).strip().rstrip("/")
persona_prompt = _read_gemini_persona_prompt(gemini_config)
tts_script = text
if _gemini_audio_tags_enabled(gemini_config, model):
tts_script = _rewrite_gemini_tts_audio_tags(text, persona_prompt=persona_prompt)
prompt_text = _compose_gemini_tts_prompt(tts_script, gemini_config, persona_prompt=persona_prompt)
max_len = origin._resolve_max_text_length("gemini", tts_config)
if len(prompt_text) > max_len:
raise ValueError(
"Gemini TTS composed prompt exceeds the provider request limit "
f"({len(prompt_text)} > {max_len} chars). Reduce the persona/audio-tag "
"prompt or lower tts.gemini.max_text_length so long-form text is "
"split with enough prompt headroom."
)
payload: Dict[str, Any] = {
"contents": [{"parts": [{"text": prompt_text}]}],
"generationConfig": {
"responseModalities": ["AUDIO"],
"speechConfig": {
"voiceConfig": {
"prebuiltVoiceConfig": {"voiceName": voice},
},
},
},
}
headers = {"Content-Type": "application/json"}
if urlparse(base_url).hostname == "generativelanguage.googleapis.com":
try:
import hermes_cli as _hermes_cli
_hermes_version = str(_hermes_cli.__version__)
except Exception:
_hermes_version = "0.0.0"
# Gemini partner-integration guidance: identify the client.
headers["X-Goog-Api-Client"] = f"hermes-agent/{_hermes_version}"
response = requests.post(
f"{base_url}/models/{model}:generateContent",
params={"key": api_key},
headers=headers,
json=payload,
timeout=60,
stream=True,
)
if response.status_code != 200:
raw_body = _read_tts_response_bytes(response, label="Gemini TTS")
try:
if raw_body:
err = json.loads(raw_body.decode("utf-8")).get("error", {})
elif not _response_has_explicit_stream(response) and callable(getattr(response, "json", None)):
err = response.json().get("error", {})
else:
err = {}
detail = err.get("message") or raw_body.decode("utf-8", errors="replace")[:300]
except Exception:
detail = raw_body.decode("utf-8", errors="replace")[:300]
raise RuntimeError(f"Gemini TTS API error (HTTP {response.status_code}): {detail}")
try:
data = _read_tts_response_json(response, label="Gemini TTS")
parts = data["candidates"][0]["content"]["parts"]
audio_part = next((p for p in parts if "inlineData" in p or "inline_data" in p), None)
if audio_part is None:
raise RuntimeError("Gemini TTS response contained no audio data")
inline = audio_part.get("inlineData") or audio_part.get("inline_data") or {}
audio_b64 = inline.get("data", "")
except (KeyError, IndexError, TypeError) as e:
raise RuntimeError(f"Gemini TTS response was malformed: {e}") from e
if not audio_b64:
raise RuntimeError("Gemini TTS returned empty audio data")
return _write_wav_bytes_as(_wrap_pcm_as_wav(base64.b64decode(audio_b64)), output_path)
+484
View File
@@ -0,0 +1,484 @@
"""Speaker-side streaming pipeline for ``tools.tts_tool.stream_tts_to_speaker``.
Turns a queue of LLM text deltas into audio the moment each sentence is
complete. Two paths share the sentence cutter (``tools.tts_streaming``):
* :class:`_StreamerPlayback` — a registered chunked streamer (ElevenLabs,
OpenAI, …). Every sentence gets a prefetch thread that fires the HTTP
request immediately and buffers PCM into a per-sentence queue; one playback
worker drains those queues in FIFO order through a sounddevice OutputStream
(or a temp WAV + system player when PortAudio is unavailable).
* :class:`_SyncSentencePipeline` — every other provider (edge, piper,
plugins). Per-sentence ``text_to_speech_tool`` synthesis on a single-thread
executor, overlapped with playback so sentence n+1 synthesizes while n plays.
Seams tests monkeypatch on the origin module (``_load_tts_config``,
``_import_sounddevice``, ``text_to_speech_tool``, ``_strip_markdown_for_tts``)
are resolved through :func:`_origin` at call time.
"""
from __future__ import annotations
import logging
import os
import platform
import queue
import tempfile
import threading
from concurrent.futures import Future, ThreadPoolExecutor
from typing import Callable, Iterable, Iterator, List, Optional
logger = logging.getLogger("tools.tts_tool")
def _origin():
from tools import tts_tool
return tts_tool
def _unlink_quietly(path: Optional[str]) -> None:
if path:
try:
os.unlink(path)
except OSError:
pass
def _align_int16_chunks(chunks: Iterable[bytes], stop_evt: threading.Event) -> Iterator[bytes]:
"""Yield int16-aligned byte chunks; a dangling odd byte is padded at the end."""
leftover = b""
for chunk in chunks:
if stop_evt.is_set():
break
buf = leftover + chunk
aligned_len = len(buf) - (len(buf) % 2)
if aligned_len >= 2:
yield buf[:aligned_len]
leftover = buf[aligned_len:] if aligned_len < len(buf) else b""
if leftover:
yield b"\x00"
def _play_via_tempfile(audio_iter: Iterable[bytes], stop_evt: threading.Event, sample_rate: int = 24000) -> None:
"""Write PCM chunks to a temp WAV file and play it with the system player."""
tmp = None
tmp_path = None
try:
import wave
tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)
tmp_path = tmp.name
with wave.open(tmp, "wb") as wf:
wf.setnchannels(1)
wf.setsampwidth(2) # 16-bit
wf.setframerate(sample_rate)
for aligned in _align_int16_chunks(audio_iter, stop_evt):
wf.writeframes(aligned)
# wave.open() on a file object does NOT close it. On Windows the open
# write handle blocks the player and the unlink below (WinError 32),
# so release it before playback.
tmp.close()
from tools.voice_mode import play_audio_file
play_audio_file(tmp_path)
except Exception as exc:
logger.warning("Temp-file TTS fallback failed: %s", exc)
finally:
if tmp is not None:
try:
tmp.close() # idempotent; ensures close on early error
except Exception:
pass
_unlink_quietly(tmp_path)
def _drain_chunks(chunk_queue: "queue.Queue[Optional[bytes]]") -> List[bytes]:
"""Collect one sentence's PCM chunks up to the ``None`` sentinel."""
chunks: List[bytes] = []
while True:
chunk = chunk_queue.get()
if chunk is None:
return chunks
chunks.append(chunk)
class _SyncSentencePipeline:
"""Overlap per-sentence synthesis with playback for non-streaming providers.
Serial synthesize-then-play added a full synthesis-time of dead air per
sentence — for a local model at real-time-factor ~1, as long silent as
speaking. One single-thread synthesis executor (sentences FIFO; providers
never see concurrent calls) feeds one playback worker through a small
bounded queue: while sentence n plays, n+1 is already synthesizing. The
bound keeps lookahead/temp files small and gives the caller backpressure.
``text_to_speech_tool`` / ``play_audio_file`` are resolved late so tests
that monkeypatch them keep working.
"""
def __init__(self, stop_event: threading.Event, *, lookahead: int = 2):
self._stop = stop_event
self._queue: "queue.Queue[Optional[tuple[str, Future]]]" = queue.Queue(maxsize=max(1, lookahead))
self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="tts-sync-synth")
self._player = threading.Thread(target=self._drain, name="tts-sync-play", daemon=True)
self._player.start()
def speak(self, cleaned: str) -> None:
"""Queue one sentence. Blocks only when the lookahead bound is full."""
if self._stop.is_set():
return
future = self._executor.submit(self._synthesize_to_tmp, cleaned)
self._queue.put((cleaned, future))
def close(self) -> None:
"""Flush queued sentences in order (skipped if stopped), then join."""
self._queue.put(None)
self._player.join()
self._executor.shutdown(wait=True)
def _synthesize_to_tmp(self, cleaned: str) -> Optional[str]:
if self._stop.is_set():
return None
tmp_path = None
try:
fd, tmp_path = tempfile.mkstemp(suffix=".mp3")
os.close(fd)
_origin().text_to_speech_tool(text=cleaned, output_path=tmp_path)
return tmp_path
except Exception as exc:
logger.warning("Sync per-sentence TTS synthesis failed: %s", exc)
_unlink_quietly(tmp_path)
return None
def _drain(self) -> None:
while True:
item = self._queue.get()
if item is None:
return
_sentence, future = item
tmp_path = None
try:
tmp_path = future.result()
if (tmp_path and not self._stop.is_set()
and os.path.isfile(tmp_path)
and os.path.getsize(tmp_path) > 0):
from tools.voice_mode import play_audio_file
play_audio_file(tmp_path)
except Exception as exc:
logger.warning("Sync per-sentence TTS failed: %s", exc)
finally:
_unlink_quietly(tmp_path)
class _StreamerPlayback:
"""Prefetch + FIFO playback for a chunked :class:`StreamingTTSProvider`.
``speak(text)`` calls ``streamer.stream()`` right away and hands the
iterator to a prefetch thread (at most 3 in flight) that buffers chunks
into a bounded per-sentence queue; the single playback worker plays those
queues in order, so sentence N+1 is already arriving while N plays.
Output goes to a PortAudio stream when one could be opened, otherwise via
temp WAV files. A failing PortAudio write is retried on a reinitialized
stream up to ``_MAX_REINIT`` times before falling back to temp files.
"""
_MAX_REINIT = 3
_CHUNK_QUEUE_MAX = 64
def __init__(self, streamer, stop_event: threading.Event):
self.streamer = streamer
self.stop_event = stop_event
self.output_stream = self._open_output_stream()
self._audio_queue: "queue.Queue[Optional[queue.Queue[Optional[bytes]]]]" = queue.Queue()
self._prefetch_threads: List[threading.Thread] = []
self._prefetch_sem = threading.Semaphore(3)
self._worker = threading.Thread(target=self._playback_worker, daemon=True)
self._worker.start()
# -- PortAudio stream management ---------------------------------------
def _create_output_stream(self):
sd = _origin()._import_sounddevice()
stream = sd.OutputStream(
samplerate=self.streamer.sample_rate,
channels=self.streamer.channels,
dtype="int16",
)
stream.start()
return stream
def _open_output_stream(self):
# On macOS skip sounddevice entirely: PortAudio/CoreAudio init triggers
# a kTCCServiceMediaLibrary permission prompt even though output needs
# no media-library access. None routes every sentence through the
# tempfile -> play_audio_file -> afplay path.
if platform.system() == "Darwin":
return None
try:
return self._create_output_stream()
except (ImportError, OSError) as exc:
logger.debug("sounddevice not available, streamer→tempfile: %s", exc)
except Exception as exc:
logger.warning("sounddevice OutputStream failed: %s", exc)
return None
def _reinit_output_stream(self):
"""Close the broken PortAudio stream and try to create a fresh one."""
if self.output_stream is not None:
try:
self.output_stream.stop()
self.output_stream.close()
except Exception:
pass
try:
self.output_stream = self._create_output_stream()
logger.info("TTS: PortAudio output stream reinitialized after error")
except Exception as exc:
logger.warning("TTS: PortAudio stream reinit failed: %s", exc)
self.output_stream = None
return self.output_stream
def close_output_stream(self) -> None:
"""Always release the device so a later stream can open it."""
if self.output_stream is not None:
try:
self.output_stream.stop()
self.output_stream.close()
except Exception:
pass
# -- prefetch ----------------------------------------------------------
def speak(self, text: str) -> None:
"""Start ``streamer.stream(text)`` and prefetch its chunks immediately."""
try:
audio_iter = self.streamer.stream(text)
except Exception as exc:
logger.warning("Streaming TTS synthesis failed: %s", exc)
return
self._prefetch_sem.acquire()
chunk_queue: "queue.Queue[Optional[bytes]]" = queue.Queue(maxsize=self._CHUNK_QUEUE_MAX)
self._audio_queue.put(chunk_queue)
t = threading.Thread(target=self._consume_to_queue, args=(audio_iter, chunk_queue), daemon=True)
self._prefetch_threads.append(t)
t.start()
def _consume_to_queue(self, audio_iter: Iterator[bytes], chunk_queue: "queue.Queue[Optional[bytes]]") -> None:
try:
for chunk in audio_iter:
if self.stop_event.is_set():
logger.info(
"TTS CUT: prefetch cancelled (stop_event set "
"mid-sentence) — partial audio only"
)
break
chunk_queue.put(chunk, timeout=30.0)
except Exception as exc:
logger.warning(
"TTS CUT: streaming TTS prefetch failed mid-sentence "
"(partial audio only): %s",
exc,
)
finally:
chunk_queue.put(None) # sentinel: no more chunks
self._prefetch_sem.release()
# -- playback ----------------------------------------------------------
def _play_sentence_via_tempfile(self, chunk_queue) -> None:
_play_via_tempfile(iter(_drain_chunks(chunk_queue)), self.stop_event, self.streamer.sample_rate)
def _playback_worker(self) -> None:
"""Single consumer: play audio segments from the queue in order."""
if self.output_stream is None:
while True:
chunk_queue = self._audio_queue.get()
if chunk_queue is None:
break
if self.stop_event.is_set():
continue
self._play_sentence_via_tempfile(chunk_queue)
return
import numpy as _np
try:
from tools.voice_mode import mark_audio_output_active
except Exception:
def mark_audio_output_active(_active):
return None
def write_pcm(stream, buf: bytes) -> None:
stream.write(_np.frombuffer(buf, dtype="<i2").reshape(-1, 1))
mark_audio_output_active(True)
try:
reinit_count = 0
current_stream = self.output_stream
while True:
chunk_queue = self._audio_queue.get()
if chunk_queue is None:
break
if self.stop_event.is_set():
continue
if current_stream is None:
self._play_sentence_via_tempfile(chunk_queue)
continue
pcm_leftover = b""
while True:
chunk = chunk_queue.get()
if chunk is None or self.stop_event.is_set():
break
buf = pcm_leftover + chunk
aligned_len = len(buf) - (len(buf) % 2)
if aligned_len >= 2:
try:
write_pcm(current_stream, buf[:aligned_len])
except Exception as write_exc:
logger.warning(
"PortAudio write failed, attempting "
"stream reinit: %s",
write_exc,
)
if reinit_count < self._MAX_REINIT:
reinit_count += 1
current_stream = self._reinit_output_stream()
if current_stream is not None:
try:
write_pcm(current_stream, buf[:aligned_len])
except Exception:
pass
pcm_leftover = buf[aligned_len:] if aligned_len < len(buf) else b""
continue
else:
logger.warning(
"TTS: PortAudio reinit exhausted "
"after %d attempts, falling back "
"to tempfile for remaining "
"sentences",
self._MAX_REINIT,
)
current_stream = None
break
pcm_leftover = buf[aligned_len:] if aligned_len < len(buf) else b""
finally:
mark_audio_output_active(False)
def finish(self) -> None:
"""Send the end sentinel, then wait for playback and prefetch threads."""
self._audio_queue.put(None)
self._worker.join(timeout=300.0)
for t in self._prefetch_threads:
t.join(timeout=10.0)
self.close_output_stream()
def stream_tts_to_speaker(
text_queue: queue.Queue,
stop_event: threading.Event,
tts_done_event: threading.Event,
display_callback: Optional[Callable[[str], None]] = None,
provider: Optional[str] = None,
):
"""Consume text deltas from *text_queue*, cut them into sentences, and speak
each one the moment it's ready — the conversational path.
A registered streaming provider plays chunked PCM for the lowest latency;
every other provider (edge, the default) is spoken per-sentence via the
sync ``text_to_speech_tool`` path, so audio still starts on sentence one.
Protocol:
* The producer puts ``str`` deltas onto *text_queue*.
* A ``None`` sentinel signals end-of-text (flush remaining buffer).
* *stop_event* aborts early (barge-in / user interrupt).
* *tts_done_event* is **set** in the ``finally`` block so callers
waiting on it (continuous voice mode) know playback is finished.
"""
tts_done_event.clear()
origin = _origin()
sync_pipeline: Optional[_SyncSentencePipeline] = None
playback: Optional[_StreamerPlayback] = None
try:
tts_config = origin._load_tts_config()
# Prefer a chunked streamer for low time-to-first-audio; otherwise
# per-sentence sync synthesis (universal — edge + every non-streamer).
from tools.tts_streaming import SentenceChunker, resolve_streaming_provider
streamer = resolve_streaming_provider(tts_config, preferred=provider)
stream_max_len = 0
if streamer is None:
sync_pipeline = _SyncSentencePipeline(stop_event)
else:
try:
stream_max_len = origin._resolve_max_text_length(
provider or origin._get_provider(tts_config), tts_config
)
except Exception:
stream_max_len = 0
playback = _StreamerPlayback(streamer, stop_event)
chunker = SentenceChunker()
long_flush_len = 100
queue_timeout = 0.5
spoken_sentences: list[str] = [] # skip duplicate/near-duplicate sentences (LLM repetition)
def _speak_sentence(sentence: str) -> None:
if stop_event.is_set():
return
cleaned = origin._strip_markdown_for_tts(sentence).strip()
if not cleaned:
return
cleaned_lower = cleaned.lower().rstrip(".!,")
if any(prev.lower().rstrip(".!,") == cleaned_lower for prev in spoken_sentences):
return
spoken_sentences.append(cleaned)
if display_callback is not None:
display_callback(sentence) # raw sentence on screen before TTS processing
if sync_pipeline is not None:
sync_pipeline.speak(cleaned)
return
if stream_max_len and len(cleaned) > stream_max_len:
cleaned = cleaned[:stream_max_len]
playback.speak(cleaned)
while not stop_event.is_set():
try:
delta = text_queue.get(timeout=queue_timeout)
except queue.Empty:
# Idle producer: flush a long buffer instead of sitting on it
if len(chunker.buf) > long_flush_len:
for sentence in chunker.flush():
_speak_sentence(sentence)
continue
if delta is None:
for sentence in chunker.flush():
_speak_sentence(sentence)
break
for sentence in chunker.feed(delta):
_speak_sentence(sentence)
while True:
try:
text_queue.get_nowait()
except queue.Empty:
break
except Exception as exc:
logger.warning("Streaming TTS pipeline error: %s", exc)
finally:
# Flush the sync pipeline first: queued sentences finish playing (or
# are skipped when stop_event is set) BEFORE tts_done_event fires, so
# continuous voice mode never reopens the mic over its own voice.
if sync_pipeline is not None:
try:
sync_pipeline.close()
except Exception:
pass
# The end sentinel lives in finally: so an exception in the text pump
# still lets the playback worker exit.
if playback is not None:
playback.finish()
tts_done_event.set()
+89 -155
View File
@@ -1,33 +1,26 @@
"""Resolve the active profile's STT/TTS config for CLIENT-DIRECT voice.
The desktop app can cut the audio relay hop (mic → gateway → provider and
provider → gateway → speaker) by calling the voice providers directly with
the profile's own credentials, fetched over the authenticated REST channel
at voice-session start. This module is the single resolver behind
``GET /api/audio/voice-config``: it reuses the exact provider/key/model/
language resolution chains ``tools.transcription_tools`` and
``tools.tts_tool`` use, so what the client receives is byte-for-byte what
the gateway itself would use for the same request.
The desktop can skip the audio relay hop (mic → gateway → provider) by calling
voice providers directly with the profile's own credentials, fetched over the
authenticated REST channel at voice-session start. This is the single resolver
behind ``GET /api/audio/voice-config``; it reuses the exact provider/key/model/
language chains of ``tools.transcription_tools`` and ``tools.tts_tool`` so the
client receives byte-for-byte what the gateway itself would use.
Design rules:
* **Same-trust boundary.** The endpoint is profile-scoped and rides the
same auth as every other REST route. A client that can reach it can
already drive the agent (terminal included), so handing it the voice
key is not a privilege escalation — but keys still never touch client
disk (the desktop holds them in renderer memory only) and are never
logged here.
* **Relay is the floor, not an error.** Providers that can only run on
the gateway host (local whisper, edge-tts, command providers, plugins)
resolve to ``{"mode": "relay"}`` and the desktop falls back to the
existing ``/api/audio/*`` relay endpoints. A resolution failure also
degrades to relay — the relay endpoint will surface the real error.
* **No new key stores.** Everything is read through the live resolvers;
nothing is persisted anywhere new.
* **Same-trust boundary.** The endpoint is profile-scoped and rides the same
auth as every REST route — a client that can reach it can already drive the
agent, so handing it the voice key is no escalation. Keys still never touch
client disk (renderer memory only) and are never logged here.
* **Relay is the floor, not an error.** Server-host-only providers (local
whisper, edge-tts, command providers, plugins) resolve to ``{"mode": "relay"}``
and the desktop falls back to ``/api/audio/*``. A resolution failure also
degrades to relay — the relay endpoint surfaces the real error.
* **No new key stores.** Everything is read through the live resolvers.
Config gate: ``voice.client_direct`` (config.yaml, default ``true``).
When false every provider reports relay and the desktop behaves exactly
as before this feature.
Config gate: ``voice.client_direct`` (config.yaml, default ``true``). When
false every provider reports relay and the desktop behaves as before.
"""
from __future__ import annotations
@@ -74,10 +67,39 @@ def _relay(reason: str) -> Dict[str, Any]:
return {"mode": "relay", "reason": reason}
def _section(config: Any, provider: str) -> Dict[str, Any]:
"""The provider's own sub-dict of an STT/TTS config, shape-guarded."""
section = config.get(provider) if isinstance(config, dict) else None
return section if isinstance(section, dict) else {}
def _direct(wire: str, provider: str, base_url: Any, api_key: str, model: Any, **extra: Any) -> Dict[str, Any]:
return {"mode": "direct", "wire": wire, "provider": provider, "base_url": base_url,
"api_key": api_key, "model": model, **extra}
def _deepinfra_model(section: Dict[str, Any], kind: str) -> Optional[str]:
"""Configured model, else the first catalog model of ``kind`` (stt/tts)."""
from hermes_cli.models import deepinfra_model_ids
model = section.get("model")
if not model:
candidates = deepinfra_model_ids(kind)
model = candidates[0] if candidates else None
return model
# ---------------------------------------------------------------------------
# STT
# ---------------------------------------------------------------------------
# provider -> (env var, default-model attr on transcription_tools, base_url).
# ``base_url`` is a transcription_tools attr name or a literal URL.
_STT_KEYED: Dict[str, tuple[str, str, str]] = {
"groq": ("GROQ_API_KEY", "DEFAULT_GROQ_STT_MODEL", "GROQ_BASE_URL"),
"mistral": ("MISTRAL_API_KEY", "DEFAULT_MISTRAL_STT_MODEL", "https://api.mistral.ai/v1"),
}
def _resolve_stt_client_config() -> Dict[str, Any]:
from tools import transcription_tools as tt
@@ -99,22 +121,21 @@ def _resolve_stt_client_config() -> Dict[str, Any]:
provider, stt_config,
extra_keys=("language_code",) if provider == "elevenlabs" else (),
)
section = stt_config.get(provider) if isinstance(stt_config, dict) else None
section = section if isinstance(section, dict) else {}
section = _section(stt_config, provider)
if provider == "groq":
api_key = tt._resolve_provider_key("GROQ_API_KEY", "groq")
def direct(wire: str, base_url: Any, api_key: str, model: Any) -> Dict[str, Any]:
return _direct(wire, provider, base_url, api_key, model, language=language)
def env_base_url(env_var: str, default: str) -> str:
return str(section.get("base_url") or tt.get_env_value(env_var) or default).strip().rstrip("/")
if provider in _STT_KEYED:
env_var, default_model, base = _STT_KEYED[provider]
api_key = tt._resolve_provider_key(env_var, provider)
if not api_key:
return _relay("no credentials")
return {
"mode": "direct",
"wire": STT_WIRE_OPENAI,
"provider": "groq",
"base_url": tt.GROQ_BASE_URL,
"api_key": api_key,
"model": section.get("model") or tt.DEFAULT_GROQ_STT_MODEL,
"language": language,
}
return direct(STT_WIRE_OPENAI, getattr(tt, base, base), api_key,
section.get("model") or getattr(tt, default_model))
if provider == "openai":
# Handles the Nous-managed selection too: the resolver returns the
@@ -124,29 +145,7 @@ def _resolve_stt_client_config() -> Dict[str, Any]:
api_key, base_url = tt._resolve_openai_audio_client_config()
except ValueError as exc:
return _relay(f"openai resolution failed: {exc}")
return {
"mode": "direct",
"wire": STT_WIRE_OPENAI,
"provider": "openai",
"base_url": base_url,
"api_key": api_key,
"model": section.get("model") or tt.DEFAULT_STT_MODEL,
"language": language,
}
if provider == "mistral":
api_key = tt._resolve_provider_key("MISTRAL_API_KEY", "mistral")
if not api_key:
return _relay("no credentials")
return {
"mode": "direct",
"wire": STT_WIRE_OPENAI,
"provider": "mistral",
"base_url": "https://api.mistral.ai/v1",
"api_key": api_key,
"model": section.get("model") or tt.DEFAULT_MISTRAL_STT_MODEL,
"language": language,
}
return direct(STT_WIRE_OPENAI, base_url, api_key, section.get("model") or tt.DEFAULT_STT_MODEL)
if provider == "xai":
# API key only. An xAI OAuth bearer refreshes server-side mid-session;
@@ -154,61 +153,26 @@ def _resolve_stt_client_config() -> Dict[str, Any]:
api_key = str(tt.get_env_value("XAI_API_KEY") or "").strip()
if not api_key:
return _relay("xai oauth (server-managed) or no credentials")
base_url = str(
section.get("base_url")
or tt.get_env_value("XAI_STT_BASE_URL")
or tt.XAI_STT_BASE_URL
).strip().rstrip("/")
return {
"mode": "direct",
"wire": STT_WIRE_XAI,
"provider": "xai",
"base_url": base_url,
"api_key": api_key,
"model": None,
"language": language,
}
return direct(STT_WIRE_XAI, env_base_url("XAI_STT_BASE_URL", tt.XAI_STT_BASE_URL), api_key, None)
if provider == "elevenlabs":
api_key = tt._resolve_provider_key("ELEVENLABS_API_KEY", "elevenlabs")
if not api_key:
return _relay("no credentials")
base_url = str(
section.get("base_url")
or tt.get_env_value("ELEVENLABS_STT_BASE_URL")
or tt.ELEVENLABS_STT_BASE_URL
).strip().rstrip("/")
return {
"mode": "direct",
"wire": STT_WIRE_ELEVENLABS,
"provider": "elevenlabs",
"base_url": base_url,
"api_key": api_key,
"model": section.get("model") or tt.DEFAULT_ELEVENLABS_STT_MODEL,
"language": language,
}
base_url = env_base_url("ELEVENLABS_STT_BASE_URL", tt.ELEVENLABS_STT_BASE_URL)
return direct(STT_WIRE_ELEVENLABS, base_url, api_key,
section.get("model") or tt.DEFAULT_ELEVENLABS_STT_MODEL)
if provider == "deepinfra":
api_key = tt._resolve_provider_key("DEEPINFRA_API_KEY", "deepinfra")
if not api_key:
return _relay("no credentials")
from hermes_cli.models import deepinfra_base_url, deepinfra_model_ids
from hermes_cli.models import deepinfra_base_url
model = section.get("model")
if not model:
candidates = deepinfra_model_ids("stt")
model = candidates[0] if candidates else None
model = _deepinfra_model(section, "stt")
if not model:
return _relay("no deepinfra stt model")
return {
"mode": "direct",
"wire": STT_WIRE_OPENAI,
"provider": "deepinfra",
"base_url": deepinfra_base_url(section),
"api_key": api_key,
"model": model,
"language": language,
}
return direct(STT_WIRE_OPENAI, deepinfra_base_url(section), api_key, model)
return _relay(f"provider {provider!r} has no client wire")
@@ -233,8 +197,7 @@ def _resolve_tts_client_config() -> Dict[str, Any]:
api_key, base_url, is_managed = tts._resolve_openai_audio_client_config()
except ValueError as exc:
return _relay(f"openai resolution failed: {exc}")
oai = tts_config.get("openai") if isinstance(tts_config, dict) else None
oai = oai if isinstance(oai, dict) else {}
oai = _section(tts_config, "openai")
model = oai.get("model") or tts.DEFAULT_OPENAI_MODEL
config_base = oai.get("base_url")
if config_base:
@@ -248,58 +211,33 @@ def _resolve_tts_client_config() -> Dict[str, Any]:
speed = float(oai.get("speed", speed_default))
except (TypeError, ValueError):
speed = 1.0
return {
"mode": "direct",
"wire": TTS_WIRE_OPENAI,
"provider": "openai",
"base_url": base_url,
"api_key": api_key,
"model": model,
"voice": oai.get("voice") or tts.DEFAULT_OPENAI_VOICE,
"speed": speed,
}
return _direct(TTS_WIRE_OPENAI, "openai", base_url, api_key, model,
voice=oai.get("voice") or tts.DEFAULT_OPENAI_VOICE, speed=speed)
if provider == "elevenlabs":
api_key = tts._resolve_provider_key("ELEVENLABS_API_KEY", "elevenlabs")
if not api_key:
return _relay("no credentials")
el = tts_config.get("elevenlabs") if isinstance(tts_config, dict) else None
el = el if isinstance(el, dict) else {}
return {
"mode": "direct",
"wire": TTS_WIRE_ELEVENLABS,
"provider": "elevenlabs",
"base_url": str(el.get("base_url") or "https://api.elevenlabs.io/v1").rstrip("/"),
"api_key": api_key,
"model": el.get("model_id") or tts.DEFAULT_ELEVENLABS_MODEL_ID,
"voice": el.get("voice_id") or tts.DEFAULT_ELEVENLABS_VOICE_ID,
"speed": None,
}
el = _section(tts_config, "elevenlabs")
return _direct(
TTS_WIRE_ELEVENLABS, "elevenlabs",
str(el.get("base_url") or "https://api.elevenlabs.io/v1").rstrip("/"),
api_key, el.get("model_id") or tts.DEFAULT_ELEVENLABS_MODEL_ID,
voice=el.get("voice_id") or tts.DEFAULT_ELEVENLABS_VOICE_ID, speed=None,
)
if provider == "deepinfra":
api_key = tts._resolve_provider_key("DEEPINFRA_API_KEY", "deepinfra")
if not api_key:
return _relay("no credentials")
from hermes_cli.models import deepinfra_base_url, deepinfra_model_ids
from hermes_cli.models import deepinfra_base_url
di = tts_config.get("deepinfra") if isinstance(tts_config, dict) else None
di = di if isinstance(di, dict) else {}
model = di.get("model")
if not model:
candidates = deepinfra_model_ids("tts")
model = candidates[0] if candidates else None
di = _section(tts_config, "deepinfra")
model = _deepinfra_model(di, "tts")
if not model:
return _relay("no deepinfra tts model")
return {
"mode": "direct",
"wire": TTS_WIRE_OPENAI,
"provider": "deepinfra",
"base_url": deepinfra_base_url(di),
"api_key": api_key,
"model": model,
"voice": di.get("voice") or "af_bella",
"speed": None,
}
return _direct(TTS_WIRE_OPENAI, "deepinfra", deepinfra_base_url(di), api_key, model,
voice=di.get("voice") or "af_bella", speed=None)
# edge / minimax / xai / mistral / gemini / neutts / kittentts / piper:
# either server-host-only engines or wire shapes the desktop doesn't
@@ -323,15 +261,11 @@ def resolve_client_voice_config() -> Dict[str, Any]:
disabled = _relay("voice.client_direct disabled")
return {"stt": disabled, "tts": disabled}
try:
stt = _resolve_stt_client_config()
except Exception:
logger.exception("client voice-config STT resolution failed")
stt = _relay("resolution error")
try:
tts = _resolve_tts_client_config()
except Exception:
logger.exception("client voice-config TTS resolution failed")
tts = _relay("resolution error")
return {"stt": stt, "tts": tts}
out: Dict[str, Any] = {}
for key, resolver in (("stt", _resolve_stt_client_config), ("tts", _resolve_tts_client_config)):
try:
out[key] = resolver()
except Exception:
logger.exception("client voice-config %s resolution failed", key.upper())
out[key] = _relay("resolution error")
return out
+648 -1211
View File
File diff suppressed because it is too large Load Diff
+341 -741
View File
File diff suppressed because it is too large Load Diff
+331
View File
@@ -0,0 +1,331 @@
"""Wake-word hotword engines (openWakeWord / sherpa-onnx KWS / Porcupine).
All three run fully on-device. Config, platform probes and sensitivity
accessors live in :mod:`tools.wake_word`; engines read them lazily through
that module so test seams (``patch("tools.wake_word.<name>")``) keep working.
"""
from __future__ import annotations
import logging
import os
from pathlib import Path
from typing import Any, Dict, Optional
logger = logging.getLogger("tools.wake_word")
def _ww():
from tools import wake_word
return wake_word
class _Engine:
"""Minimal hotword-engine contract: feed int16 frames, get a bool."""
frame_length: int = 1280 # 80 ms at 16 kHz
#: (matched phrase, profile name) of the most recent fire. Multi-phrase
#: engines (sherpa) set this for profile routing; single-phrase engines
#: leave it None (callers fall back to configured phrase / active profile).
last_match: Optional[tuple[str, str]] = None
def process(self, frame) -> bool: # frame: 1-D int16 ndarray
raise NotImplementedError
def reset(self) -> None:
"""Clear any internal audio/feature buffer (called on every (re)start)."""
def close(self) -> None:
pass
def _looks_like_path(value: str) -> bool:
return os.sep in value or value.endswith((".onnx", ".tflite", ".ppn")) or os.path.exists(value)
def _sub(cfg: Dict[str, Any], key: str) -> Dict[str, Any]:
sub = cfg.get(key)
return sub if isinstance(sub, dict) else {}
class _OpenWakeWordEngine(_Engine):
"""openWakeWord — free, local ONNX/tflite hotword detection.
Scores one ~80 ms frame at a time; ``sensitivity`` IS the raw 0..1
threshold (higher = stricter). A real utterance holds the score high
across frames while a stray ambient phoneme spikes one, so we require
``confirmation_frames`` consecutive over-threshold frames before firing.
"""
frame_length = 1280 # openWakeWord recommends 80 ms frames.
def __init__(self, cfg: Dict[str, Any]):
from tools import lazy_deps
lazy_deps.ensure("wake.openwakeword", prompt=False)
import openwakeword
from openwakeword.model import Model
ww = _ww()
model_ref = str(_sub(cfg, "openwakeword").get("model") or ww._BUNDLED_MODEL_NAME).strip()
framework = self._usable_framework(ww.resolve_inference_framework(cfg))
self._threshold = ww._sensitivity(cfg)
self._confirm_needed = ww._confirmation_frames(cfg)
self._confirm_streak = 0
# Default (or explicit "hey_hermes") → the bundled model; a built-in
# name or custom path is used as-is.
if model_ref.lower() in ww._BUNDLED_MODEL_ALIASES:
model_ref = ww._bundled_wakeword_path(framework)
# download_models() also fetches the shared feature models (melspectrogram
# + embedding) needed for ANY model, so a custom path must call it too or a
# fresh install crashes on a missing melspectrogram.onnx.
try:
openwakeword.utils.download_models([model_ref])
except Exception as e: # pragma: no cover - network/path dependent
logger.debug("openwakeword model download skipped: %s", e)
self._model = Model(wakeword_models=[model_ref], inference_framework=framework)
self._labels = list(self._model.models.keys())
@staticmethod
def _usable_framework(framework: str) -> str:
"""Refuse openWakeWord's silent tflite→onnx downgrade.
Without a tflite runtime openWakeWord falls back to onnx, which on
macOS ARM64 is the backend whose embedding model never fires — the
listener would arm and stay deaf. Install + bridge the runtime first
(the platform gate lives here because dep specs can't carry PEP 508
markers); on that Mac raise instead of downgrading.
"""
ww = _ww()
if framework != "tflite" or ww.ensure_tflite_runtime():
return framework
try:
from tools import lazy_deps
lazy_deps.ensure("wake.openwakeword.tflite", prompt=False)
except Exception as e:
logger.debug("wake word: tflite runtime install failed: %s", e)
if ww.ensure_tflite_runtime():
return framework
if ww._is_macos_arm64():
raise RuntimeError(
"The wake word needs the tflite backend on this Mac, but its "
"runtime is missing. Install it with: pip install ai-edge-litert"
)
logger.warning("wake word: no tflite runtime available — falling back to onnx")
return "onnx"
def process(self, frame) -> bool:
scores = self._model.predict(frame)
if not any(score >= self._threshold for score in scores.values()):
self._confirm_streak = 0
return False
self._confirm_streak += 1
if self._confirm_streak < self._confirm_needed:
return False
self._confirm_streak = 0
return True
def reset(self) -> None:
# Clears openWakeWord's rolling feature/prediction buffer so stale audio
# captured before a pause can't re-fire the moment we resume.
self._confirm_streak = 0
try:
self._model.reset()
except Exception:
pass
def close(self) -> None:
self.reset()
# sherpa-onnx open-vocabulary KWS model: a small streaming zipformer transducer
# (English, GigaSpeech); one-time download cached under HERMES_HOME. Keywords
# are typed phrases tokenized at RUNTIME — no training step.
_SHERPA_KWS_MODEL_URL = (
"https://github.com/k2-fsa/sherpa-onnx/releases/download/kws-models/"
"sherpa-onnx-kws-zipformer-gigaspeech-3.3M-2024-01-01.tar.bz2"
)
_SHERPA_KWS_MODEL_DIR = "sherpa-onnx-kws-zipformer-gigaspeech-3.3M-2024-01-01"
def _sherpa_model_root() -> Path:
from hermes_constants import get_hermes_home
return get_hermes_home() / "cache" / "wakewords"
def _ensure_sherpa_model(root: Optional[Path] = None) -> Path:
"""Download + unpack the sherpa KWS model once; return its directory."""
root = root or _sherpa_model_root()
target = root / _SHERPA_KWS_MODEL_DIR
if (target / "tokens.txt").exists():
return target
import tarfile
import urllib.request
root.mkdir(parents=True, exist_ok=True)
archive = root / f"{_SHERPA_KWS_MODEL_DIR}.tar.bz2"
logger.info("wake word: downloading sherpa KWS model (one-time, ~13 MB)")
urllib.request.urlretrieve(_SHERPA_KWS_MODEL_URL, archive) # noqa: S310
with tarfile.open(archive, "r:bz2") as tf:
tf.extractall(root, filter="data")
archive.unlink(missing_ok=True)
if not (target / "tokens.txt").exists():
raise RuntimeError(f"sherpa KWS model unpack failed: {target}")
return target
class _SherpaKwsEngine(_Engine):
"""sherpa-onnx open-vocabulary keyword spotting — any typed phrase, zero training.
``wake_word.phrase`` is BPE-tokenized at runtime against the model's
vocabulary, so here ``phrase`` is DETECTION config, not a cosmetic label.
"""
frame_length = 1280 # streaming zipformer accepts any chunk; match capture path.
def __init__(self, cfg: Dict[str, Any]):
from tools import lazy_deps
lazy_deps.ensure("wake.sherpa", prompt=False)
import sherpa_onnx
from sherpa_onnx import text2token
ww = _ww()
model_dir = str(_sub(cfg, "sherpa").get("model_dir") or "").strip()
d = Path(model_dir) if model_dir else _ensure_sherpa_model()
if not (d / "tokens.txt").exists():
raise RuntimeError(f"sherpa KWS model not found at {d}")
# Phrase set: this profile's own phrase plus — when profile routing is
# on — every other wake-enabled profile's phrase, so ONE listener can
# wake any profile. phrase → profile is kept for routing the match back.
phrase = str(ww._get(cfg, "phrase") or "hey hermes").strip()
phrase_map: Dict[str, str] = {phrase: ww._active_profile_name()}
if bool(cfg.get("profile_routing", True)):
for prof, p in ww.enrolled_profile_phrases().items():
phrase_map.setdefault(p.strip(), prof)
phrases = list(phrase_map)
tokens = text2token(
[p.upper() for p in phrases],
tokens=str(d / "tokens.txt"),
tokens_type="bpe",
bpe_model=str(d / "bpe.model"),
)
import tempfile
# sherpa keyword entries reject spaces in the @display-name; underscore
# them and map display → profile for match routing.
self._display_to_profile: Dict[str, str] = {}
kw = tempfile.NamedTemporaryFile(
mode="w", suffix=".txt", prefix="hermes-kws-", delete=False, encoding="utf-8"
)
for p, toks in zip(phrases, tokens):
display = p.upper().replace(" ", "_")
self._display_to_profile[display] = phrase_map[p]
kw.write(" ".join(toks) + f" @{display}\n")
kw.close()
self._keywords_file = kw.name
self.last_match: Optional[tuple[str, str]] = None
# Shared 0..1 sensitivity → sherpa keywords_threshold. 0.5 lands on
# sherpa's recommended 0.25; a stricter 0.35 missed ~12% of true
# positives in live TTS matrix tests while 0.25 held zero false fires.
threshold = 0.05 + 0.4 * ww._sensitivity(cfg)
def _model_file(pattern: str) -> str:
hits = sorted(d.glob(pattern))
if not hits:
raise RuntimeError(f"sherpa KWS model file missing: {d}/{pattern}")
return str(hits[0])
self._spotter = sherpa_onnx.KeywordSpotter(
tokens=str(d / "tokens.txt"),
encoder=_model_file("encoder-*[!8].onnx"),
decoder=_model_file("decoder-*[!8].onnx"),
joiner=_model_file("joiner-*[!8].onnx"),
keywords_file=self._keywords_file,
keywords_threshold=threshold,
num_threads=1,
)
self._stream = self._spotter.create_stream()
def process(self, frame) -> bool:
import numpy as np
samples = np.asarray(frame, dtype=np.float32) / 32768.0
self._stream.accept_waveform(_ww().SAMPLE_RATE, samples)
fired = False
while self._spotter.is_ready(self._stream):
self._spotter.decode_stream(self._stream)
result = self._spotter.get_result(self._stream)
if result:
fired = True
display = str(result)
self.last_match = (
display.replace("_", " ").lower(),
self._display_to_profile.get(display, ""),
)
# Reset decoder state so one utterance can't fire repeatedly.
self._spotter.reset_stream(self._stream)
return fired
def reset(self) -> None:
# Fresh stream drops buffered audio/decoder state (pause → resume must
# not re-fire on stale audio).
try:
self._stream = self._spotter.create_stream()
except Exception:
pass
def close(self) -> None:
try:
os.unlink(self._keywords_file)
except OSError:
pass
class _PorcupineEngine(_Engine):
"""Picovoice Porcupine — premium, on-device, needs an access key."""
def __init__(self, cfg: Dict[str, Any]):
from tools import lazy_deps
lazy_deps.ensure("wake.porcupine", prompt=False)
import pvporcupine
access_key = (os.getenv("PORCUPINE_ACCESS_KEY") or "").strip()
if not access_key:
raise RuntimeError(
"Porcupine wake word requires PORCUPINE_ACCESS_KEY "
"(get a free key at https://console.picovoice.ai)."
)
keyword = str(_sub(cfg, "porcupine").get("keyword") or "jarvis").strip()
# Porcupine's `sensitivities` runs the OPPOSITE way to our shared knob
# (higher = looser); invert so "higher = stricter" holds for every engine.
kwargs: Dict[str, Any] = {"access_key": access_key, "sensitivities": [1.0 - _ww()._sensitivity(cfg)]}
kwargs["keyword_paths" if _looks_like_path(keyword) else "keywords"] = [keyword]
self._porcupine = pvporcupine.create(**kwargs)
self.frame_length = self._porcupine.frame_length
def process(self, frame) -> bool:
# pvporcupine wants a plain list/sequence of int16 samples.
return self._porcupine.process(frame) >= 0
def close(self) -> None:
try:
self._porcupine.delete()
except Exception:
pass