refactor(tools/voice): extract tts delivery + wake_word engines; dedupe transcription/voice_mode helpers; compact tts providers
This commit is contained in:
+745
-1330
File diff suppressed because it is too large
Load Diff
+56
-86
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
+341
-741
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
Reference in New Issue
Block a user