From 2ef1e8e4e00884d2b6ae4001e52ca050a8918f06 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 13:55:33 -0700 Subject: [PATCH] refactor(tools/voice): extract tts delivery + wake_word engines; dedupe transcription/voice_mode helpers; compact tts providers --- tools/transcription_tools.py | 2075 +++++++----------- tools/tts_streaming.py | 142 +- tools/tts_tool.py | 3971 ++++++---------------------------- tools/tts_tool_delivery.py | 549 +++++ tools/tts_tool_local.py | 262 +++ tools/tts_tool_providers.py | 896 ++++++++ tools/tts_tool_speaker.py | 484 +++++ tools/voice_client_config.py | 244 +-- tools/voice_mode.py | 1859 ++++++---------- tools/wake_word.py | 1082 +++------ tools/wake_word_engines.py | 331 +++ 11 files changed, 5035 insertions(+), 6860 deletions(-) create mode 100644 tools/tts_tool_delivery.py create mode 100644 tools/tts_tool_local.py create mode 100644 tools/tts_tool_providers.py create mode 100644 tools/tts_tool_speaker.py create mode 100644 tools/wake_word_engines.py diff --git a/tools/transcription_tools.py b/tools/transcription_tools.py index c568e7bce7..9fd9637478 100644 --- a/tools/transcription_tools.py +++ b/tools/transcription_tools.py @@ -1,36 +1,16 @@ #!/usr/bin/env python3 -""" -Transcription Tools Module +"""Speech-to-text transcription used by the gateway for voice messages. -Provides speech-to-text transcription with six providers: +Built-in providers: local (faster-whisper, default/free), local_command, groq, +openai (also serves the managed ``nous`` selection), mistral, xai, elevenlabs, +deepinfra; plus user-declared command providers and plugin providers. - - **local** (default, free) — faster-whisper running locally, no API key needed. - Auto-downloads the model (~150 MB for ``base``) on first use. - - **groq** (free tier) — Groq Whisper API, requires ``GROQ_API_KEY``. - - **openai** (paid) — OpenAI Whisper API, requires ``VOICE_TOOLS_OPENAI_KEY``. - - **mistral** — Mistral Voxtral Transcribe API, requires ``MISTRAL_API_KEY``. - - **xai** — xAI Grok STT API, requires ``XAI_API_KEY``. High accuracy, - Inverse Text Normalization, diarization, 21 languages. - - **elevenlabs** — ElevenLabs Scribe API, requires ``ELEVENLABS_API_KEY``. - -Used by the messaging gateway to automatically transcribe voice messages -sent by users on Telegram, Discord, WhatsApp, Slack, and Signal. - -Supported input formats: mp3, mp4, mpeg, mpga, m4a, wav, webm, ogg, aac - -Usage:: - - from tools.transcription_tools import transcribe_audio - - result = transcribe_audio("/path/to/audio.ogg") - if result["success"]: - print(result["transcript"]) + result = transcribe_audio("/path/to/audio.ogg") # {"success", "transcript", "error"?, "provider"?} """ import logging import os import platform -import queue import re import shlex import shutil @@ -45,7 +25,7 @@ from urllib.parse import urljoin from hermes_cli._subprocess_compat import windows_hide_flags from utils import is_truthy_value from tools.managed_tool_gateway import resolve_managed_tool_gateway -from tools.tts_command_provider import ( +from tools.tts_command_provider import ( # noqa: F401 — aliases are patched by tests command_env_passthrough as _command_stt_env_passthrough, quote_command_placeholder as _quote_command_stt_placeholder, render_command_template as _render_command_stt_template, @@ -61,12 +41,12 @@ from tools.tool_backend_helpers import ( logger = logging.getLogger(__name__) + def get_env_value(name, default=None): """Read env values through the live config module. - Tests may monkeypatch and later restore ``hermes_cli.config.get_env_value`` - before this module is imported. Resolve the helper at call time so STT does - not keep a stale imported function for the rest of the test process. + Resolved at call time: tests monkeypatch/restore ``hermes_cli.config.get_env_value`` + around this module's import, so a cached import would go stale. """ try: from hermes_cli.config import get_env_value as _get_env_value @@ -77,13 +57,9 @@ def get_env_value(name, default=None): def _resolve_provider_key(env_var: str, provider_id: str) -> str: - """Resolve an STT provider API key via the shared voice-key resolver. + """Resolve an STT API key via the shared voice-key resolver (config > env/.env > credential pool). - Delegates to ``tools.tool_backend_helpers.resolve_provider_secret`` — - the single owner of STT/TTS key resolution (config > env/.env > the - credential pool populated by ``hermes auth add ``). - Resolved at call time so tests that reload the helpers module see the - live function. + Resolved at call time so tests that reload the helpers module see the live function. """ try: from tools.tool_backend_helpers import resolve_provider_secret @@ -91,6 +67,7 @@ def _resolve_provider_key(env_var: str, provider_id: str) -> str: return str(get_env_value(env_var) or "").strip() return resolve_provider_secret(env_var, provider_id, env_getter=get_env_value) + # --------------------------------------------------------------------------- # Optional imports — graceful degradation # --------------------------------------------------------------------------- @@ -129,7 +106,7 @@ GROQ_BASE_URL = os.getenv("GROQ_BASE_URL", "https://api.groq.com/openai/v1") OPENAI_BASE_URL = os.getenv("STT_OPENAI_BASE_URL", "https://api.openai.com/v1") XAI_STT_BASE_URL = os.getenv("XAI_STT_BASE_URL", "https://api.x.ai/v1") ELEVENLABS_STT_BASE_URL = os.getenv("ELEVENLABS_STT_BASE_URL", "https://api.elevenlabs.io/v1") -# DeepInfra STT base URL now resolved via hermes_cli.models.deepinfra_base_url (shared). +# DeepInfra STT base URL is resolved via hermes_cli.models.deepinfra_base_url (shared). SUPPORTED_FORMATS = {".mp3", ".mp4", ".mpeg", ".mpga", ".m4a", ".wav", ".webm", ".ogg", ".oga", ".opus", ".aac", ".flac", ".caf"} LOCAL_NATIVE_AUDIO_FORMATS = {".wav", ".aiff", ".aif"} @@ -139,30 +116,35 @@ MAX_FILE_SIZE = 25 * 1024 * 1024 # 25 MB OPENAI_MODELS = {"whisper-1", "gpt-4o-mini-transcribe", "gpt-4o-transcribe", "gpt-transcribe"} GROQ_MODELS = {"whisper-large-v3", "whisper-large-v3-turbo", "distil-whisper-large-v3-en"} -# Singleton for the local model — loaded once, reused across calls +# Singleton for the local model — loaded once, reused across calls. The lock +# guards the check-then-load so two concurrent voice messages can't both +# download/load the model. _local_model: Optional[object] = None _local_model_name: Optional[str] = None -# Guards the check-then-load of the module-global model cache above. -# Without it, two concurrent voice messages can both see `_local_model is -# None` and download/load the whisper model twice (#24767). _local_model_lock = threading.Lock() -# --- Idle unload --------------------------------------------------------------- -# The model singleton above is loaded once and never released — hundreds of MB -# of RAM/VRAM sit idle between voice messages. On long-running gateway -# processes (especially with local LLMs competing for the same GPU) this is -# wasteful. A single long-lived daemon thread checks _last_transcription_time -# and unloads the model after a configurable idle period, then exits. The next -# voice message reloads the model and restarts the watcher transparently. +# Idle unload: a single daemon thread checks _last_transcription_time and +# releases the model (hundreds of MB of RAM/VRAM) after a configurable idle +# period, then exits; the next voice message reloads and restarts it. +# _idle_unload_mgmt_lock serializes the start check so two concurrent +# transcriptions can't both observe "no watcher alive" and spawn duplicates. _last_transcription_time: float = 0.0 _idle_unload_thread: Optional[threading.Thread] = None _idle_unload_stop = threading.Event() -# Serializes watcher start checks so two concurrent transcriptions can't -# both observe "no watcher alive" and spawn duplicates. _idle_unload_mgmt_lock = threading.Lock() _IDLE_UNLOAD_CHECK_INTERVAL = 30 # seconds between idle checks + +def _error_result(error: str, **extra: Any) -> Dict[str, Any]: + """Standard failure envelope shared by every provider and validator.""" + return {"success": False, "transcript": "", "error": error, **extra} + + +def _ok_result(transcript: str, provider: str) -> Dict[str, Any]: + return {"success": True, "transcript": transcript, "provider": provider} + + # --------------------------------------------------------------------------- # Config helpers # --------------------------------------------------------------------------- @@ -181,8 +163,15 @@ def is_stt_enabled(stt_config: Optional[dict] = None) -> bool: """Return whether STT is enabled in config.""" if stt_config is None: stt_config = _load_stt_config() - enabled = stt_config.get("enabled", True) - return is_truthy_value(enabled, default=True) + return is_truthy_value(stt_config.get("enabled", True), default=True) + + +def _get_stt_section(stt_config: Dict[str, Any], name: str) -> Dict[str, Any]: + """Return an stt sub-section if it's a dict, else an empty dict.""" + if not isinstance(stt_config, dict): + return {} + section = stt_config.get(name) + return section if isinstance(section, dict) else {} def _resolve_stt_language( @@ -191,23 +180,16 @@ def _resolve_stt_language( *, extra_keys: tuple = (), ) -> Optional[str]: - """Resolve the language hint for an STT provider (class-level, all providers). + """Resolve the language hint for an STT provider; first non-empty wins. - Resolution order (first non-empty wins): - 1. ``stt..language`` (plus any *extra_keys* aliases, e.g. - ElevenLabs' historical ``language_code``) - 2. ``stt.language`` — global default for every provider - 3. ``HERMES_LOCAL_STT_LANGUAGE`` env var (legacy escape hatch) - 4. ``None`` — let the provider auto-detect - - Returns a stripped ISO-639-1-ish code or None. Never returns "". + Order: ``stt..language`` (plus *extra_keys* aliases, e.g. ElevenLabs' + ``language_code``) > ``stt.language`` > ``HERMES_LOCAL_STT_LANGUAGE`` env > + None (provider auto-detects). Never returns "". """ if stt_config is None: stt_config = _load_stt_config() provider_cfg = _get_stt_section(stt_config, provider_key) - candidates = [provider_cfg.get("language")] - for key in extra_keys: - candidates.append(provider_cfg.get(key)) + candidates = [provider_cfg.get(key) for key in ("language", *extra_keys)] if isinstance(stt_config, dict): candidates.append(stt_config.get("language")) candidates.append(os.getenv(LOCAL_STT_LANGUAGE_ENV)) @@ -239,9 +221,25 @@ def _find_ffmpeg_binary() -> Optional[str]: return _find_binary("ffmpeg") -# Shared encode profile for every STT-bound m4a we produce (transcode and -# silence-trim): 16 kHz mono 32 kbps AAC, faststart. One owner — codec or -# bitrate changes must not drift between the two paths. +def _find_ffprobe_binary() -> Optional[str]: + return _find_binary("ffprobe") + + +def _find_whisper_binary() -> Optional[str]: + return _find_binary("whisper") + + +def _run_quiet(command: list, *, timeout: float, env: Optional[dict] = None) -> subprocess.CompletedProcess: + """``subprocess.run`` for STT helper binaries: checked, captured, utf-8 text, no stdin, hidden window.""" + return subprocess.run( + command, check=True, capture_output=True, text=True, + encoding="utf-8", errors="replace", timeout=timeout, + stdin=subprocess.DEVNULL, env=env, creationflags=windows_hide_flags(), + ) + + +# Shared encode profile for every STT-bound m4a (transcode and silence-trim): +# 16 kHz mono 32 kbps AAC, faststart. One owner so codec/bitrate never drift. _STT_M4A_ENCODE_ARGS = ( "-vn", "-ac", "1", "-ar", "16000", "-c:a", "aac", "-b:a", "32k", "-movflags", "+faststart", @@ -253,29 +251,21 @@ def _run_ffmpeg_stt_encode( ) -> None: """Run the shared STT m4a encode, optionally with an ``-af`` filter. - Raises on failure (CalledProcessError / TimeoutExpired) — callers own - the error semantics (transcode reports, trim swallows). + Raises on failure — callers own the error semantics (transcode reports, trim swallows). """ command = [ffmpeg, "-y", "-i", input_path] if audio_filter: command += ["-af", audio_filter] command += [*_STT_M4A_ENCODE_ARGS, output_path] - subprocess.run( - command, check=True, capture_output=True, text=True, - encoding="utf-8", errors="replace", timeout=120, - stdin=subprocess.DEVNULL, creationflags=windows_hide_flags(), - ) + _run_quiet(command, timeout=120) def _transcode_audio_for_stt(file_path: str, work_dir: str) -> tuple[Optional[str], Optional[str]]: - """Transcode ``file_path`` to a compact, broadly-accepted .m4a for STT upload. + """Transcode to a compact 16 kHz mono AAC/m4a for STT upload. - Newer OpenAI transcription models (``gpt-4o-transcribe``, - ``gpt-4o-mini-transcribe``) reject some containers the legacy ``whisper-1`` - endpoint accepted -- notably the Ogg/Opus voice notes messaging apps send -- - and gateway downloads occasionally arrive with a misleading extension. - Normalizing to 16 kHz mono AAC/m4a produces a small file the endpoints - accept. Returns ``(converted_path, None)`` on success or ``(None, error)``. + Newer OpenAI models reject containers ``whisper-1`` accepted (notably Ogg/Opus + voice notes) and gateway downloads may carry a misleading extension. + Returns ``(converted_path, None)`` or ``(None, error)``. """ ffmpeg = _find_ffmpeg_binary() if not ffmpeg: @@ -293,20 +283,14 @@ def _transcode_audio_for_stt(file_path: str, work_dir: str) -> tuple[Optional[st return None, f"failed to transcode audio for the STT API: {exc}" -def _find_whisper_binary() -> Optional[str]: - return _find_binary("whisper") - - def _get_local_command_template() -> Optional[str]: configured = os.getenv(LOCAL_STT_COMMAND_ENV, "").strip() if configured: return configured - whisper_binary = _find_whisper_binary() if whisper_binary: - quoted_binary = shlex.quote(whisper_binary) return ( - f"{quoted_binary} {{input_path}} --model {{model}} --output_format txt " + f"{shlex.quote(whisper_binary)} {{input_path}} --model {{model}} --output_format txt " "--output_dir {output_dir} --language {language}" ) return None @@ -317,48 +301,33 @@ def _has_local_command() -> bool: def _normalize_local_model(model_name: Optional[str]) -> str: - """Return a valid faster-whisper model size, mapping cloud-only names to the default. - - Cloud providers like OpenAI use names such as ``whisper-1`` which are not - valid for faster-whisper (which expects ``tiny``, ``base``, ``small``, - ``medium``, or ``large-v*``). When such a name is detected we fall back to - the default local model and emit a warning so the user knows what happened. - """ - if not model_name or model_name in OPENAI_MODELS or model_name in GROQ_MODELS: - if model_name and (model_name in OPENAI_MODELS or model_name in GROQ_MODELS): - logger.warning( - "STT model '%s' is a cloud-only name and cannot be used with the local " - "provider. Falling back to '%s'. Set stt.local.model to a valid " - "faster-whisper size (tiny, base, small, medium, large-v3).", - model_name, - DEFAULT_LOCAL_MODEL, - ) + """Return a valid faster-whisper size; cloud-only names (``whisper-1`` …) fall back to the default with a warning.""" + if not model_name: + return DEFAULT_LOCAL_MODEL + if model_name in OPENAI_MODELS or model_name in GROQ_MODELS: + logger.warning( + "STT model '%s' is a cloud-only name and cannot be used with the local " + "provider. Falling back to '%s'. Set stt.local.model to a valid " + "faster-whisper size (tiny, base, small, medium, large-v3).", + model_name, + DEFAULT_LOCAL_MODEL, + ) return DEFAULT_LOCAL_MODEL return model_name -def _normalize_local_command_model(model_name: Optional[str]) -> str: - return _normalize_local_model(model_name) +_normalize_local_command_model = _normalize_local_model def _try_lazy_install_stt() -> bool: - """Attempt to lazy-install faster-whisper and return True on success. - - The module-level ``_HAS_FASTER_WHISPER`` flag is set at import time and - cached. If the package wasn't installed at startup, calling ``ensure()`` - installs it. This function re-checks dynamically after installation so - the provider can use it immediately without a process restart. - """ + """Lazy-install faster-whisper and re-check dynamically so it's usable without a restart.""" try: from tools.lazy_deps import ensure - # prompt=False: never raise a blocking input() prompt mid-session. - # Under the interactive CLI prompt_toolkit owns stdin, so a bare - # input() deadlocks the terminal (#40490). The install is already - # gated by security.allow_lazy_installs, so reaching here is opt-in. + # prompt=False: a bare input() deadlocks under the interactive CLI where + # prompt_toolkit owns stdin; the install is already gated by + # security.allow_lazy_installs, so reaching here is opt-in. ensure("stt.faster_whisper", prompt=False) - # Re-check dynamically after install - import importlib.util as _iu - if _iu.find_spec("faster_whisper"): + if _ilu.find_spec("faster_whisper"): return True logger.warning( "faster-whisper was installed but importlib still cannot find it " @@ -377,85 +346,52 @@ def _try_lazy_install_stt() -> bool: return False -# Names of the STT providers with native handlers in this module. -# Kept in sync with ``agent.transcription_registry._BUILTIN_NAMES`` — -# a regression test fails if they drift. The plugin hook from -# issue #30398-style follow-up rejects plugins registering under any -# of these names; the dispatcher in ``transcribe_audio`` short-circuits -# them defensively as well. +# Providers with native handlers here. Kept in sync with +# ``agent.transcription_registry._BUILTIN_NAMES`` (a regression test fails on +# drift); plugins may not register under these names and the dispatcher +# short-circuits them before command/plugin lookup. BUILTIN_STT_PROVIDERS = frozenset({ - "local", - "local_command", - "groq", - "openai", - "mistral", - "xai", - "elevenlabs", - "deepinfra", + "local", "local_command", "groq", "openai", "mistral", "xai", "elevenlabs", "deepinfra", }) +# Built-in providers that upload audio to a remote API. +CLOUD_STT_PROVIDERS = frozenset(BUILTIN_STT_PROVIDERS - {"local", "local_command"}) + # --------------------------------------------------------------------------- # Command-provider registry (``stt.providers.: type: command``) # --------------------------------------------------------------------------- # -# Mirrors the TTS command-provider registry shipped in PR #17843 — same -# placeholder grammar, same shell-quote-aware rendering, same process-tree -# termination on timeout. Lets any whisper CLI / ASR CLI / curl pipeline -# become an STT backend with zero Python. -# -# Resolution order: -# 1. Built-in (``local``, ``local_command``, ``groq``, ``openai``, -# ``mistral``, ``xai``) → native handler. **Always wins.** -# 2. ``stt.providers.: type: command`` → command-provider runner. -# 3. Plugin-registered TranscriptionProvider → plugin dispatch. -# 4. No match → "No STT provider available". -# -# The single-env-var ``HERMES_LOCAL_STT_COMMAND`` escape hatch is preserved -# untouched via the built-in ``local_command`` path. Use the command-provider -# registry when you want MULTIPLE shell-driven STT engines, or you want a -# named provider you can pick via ``stt.provider`` in config.yaml. +# Mirrors the TTS command-provider registry: same placeholder grammar, +# shell-quote-aware rendering and process-tree termination on timeout. +# Resolution order: built-in name (always wins) > stt.providers. command +# > plugin-registered TranscriptionProvider > "No STT provider available". +# The single-env-var HERMES_LOCAL_STT_COMMAND escape hatch stays untouched via +# the built-in ``local_command`` path. DEFAULT_COMMAND_STT_TIMEOUT_SECONDS = 300 DEFAULT_COMMAND_STT_LANGUAGE = "en" DEFAULT_COMMAND_STT_OUTPUT_FORMAT = "txt" COMMAND_STT_OUTPUT_FORMATS = frozenset({"txt", "json", "srt", "vtt"}) -def _get_stt_section(stt_config: Dict[str, Any], name: str) -> Dict[str, Any]: - """Return an stt sub-section if it's a dict, else an empty dict.""" - if not isinstance(stt_config, dict): - return {} - section = stt_config.get(name) - return section if isinstance(section, dict) else {} - - def _get_named_stt_provider_config( stt_config: Dict[str, Any], name: str, ) -> Dict[str, Any]: - """Return the config dict for a user-declared STT command provider. + """Return the config for a user-declared STT provider, or {}. - Looks up ``stt.providers.`` first (the canonical location), and - falls back to ``stt.`` so users who followed the built-in layout - still work. Returns an empty dict when the provider is not declared. - - Built-in names are NOT special-cased here — the caller short-circuits - them before this is consulted, AND ``_is_command_stt_provider_config`` - requires an explicit ``command:`` value, so a built-in section like - ``stt.openai`` (which has ``model``/``language`` but no ``command``) - can't accidentally be treated as a command provider. + ``stt.providers.`` is canonical; ``stt.`` is accepted for + back-compat only when *name* is not a built-in, so a user's ``stt.openai`` + block still means the OpenAI provider. Built-in sections can't be mistaken + for command providers anyway: ``_is_command_stt_provider_config`` requires + an explicit ``command:``. """ providers = _get_stt_section(stt_config, "providers") - section = providers.get(name) if isinstance(providers, dict) else None + section = providers.get(name) if isinstance(section, dict): return section - # Back-compat: allow ``stt.`` for user-declared providers too, - # but only when the name is not a built-in (so a user's ``stt.openai`` - # block still means the OpenAI provider, not a custom command). if name.lower() not in BUILTIN_STT_PROVIDERS: - legacy = _get_stt_section(stt_config, name) - if legacy: - return legacy + return _get_stt_section(stt_config, name) return {} @@ -474,49 +410,19 @@ def _resolve_command_stt_provider_config( provider: str, stt_config: Dict[str, Any], ) -> Optional[Dict[str, Any]]: - """Return the provider config if *provider* resolves to a command type. - - Built-in provider names are rejected (they have native handlers). - Returns None when the name is a built-in, ``"none"``, unknown, or not - a command type. - """ + """Return the provider config if *provider* is a command type; None for built-ins, ``none``, unknown.""" if not provider: return None key = provider.lower().strip() if key in BUILTIN_STT_PROVIDERS or key == "none": return None config = _get_named_stt_provider_config(stt_config, key) - if _is_command_stt_provider_config(config): - return config - return None + return config if _is_command_stt_provider_config(config) else None def _is_local_stt_provider(provider: str, stt_config: Dict[str, Any]) -> bool: """Return whether *provider* is exempt from Hermes's remote upload cap.""" - key = (provider or "").lower().strip() - if key in {"local", "local_command"}: - return True - return False - - -def _iter_command_stt_providers(stt_config: Dict[str, Any]): - """Yield (name, config) pairs for every declared command-type STT provider.""" - if not isinstance(stt_config, dict): - return - providers = _get_stt_section(stt_config, "providers") - for name, cfg in (providers or {}).items(): - if isinstance(name, str) and name.lower() not in BUILTIN_STT_PROVIDERS: - if _is_command_stt_provider_config(cfg): - yield name, cfg - - -def _has_any_command_stt_provider(stt_config: Optional[Dict[str, Any]] = None) -> bool: - """Return True when any command-type STT provider is configured.""" - if stt_config is None: - stt_config = _load_stt_config() - for _name, _cfg in _iter_command_stt_providers(stt_config): - return True - return False + return (provider or "").lower().strip() in {"local", "local_command"} def _get_command_stt_timeout(config: Dict[str, Any]) -> float: @@ -526,36 +432,20 @@ def _get_command_stt_timeout(config: Dict[str, Any]) -> float: value = float(raw) except (TypeError, ValueError): return float(DEFAULT_COMMAND_STT_TIMEOUT_SECONDS) - if value <= 0: - return float(DEFAULT_COMMAND_STT_TIMEOUT_SECONDS) - return value + return value if value > 0 else float(DEFAULT_COMMAND_STT_TIMEOUT_SECONDS) def _get_command_stt_output_format(config: Dict[str, Any]) -> str: """Return the validated output format (txt/json/srt/vtt).""" - raw = ( - config.get("format") - or config.get("output_format") - or DEFAULT_COMMAND_STT_OUTPUT_FORMAT - ) + raw = config.get("format") or config.get("output_format") or DEFAULT_COMMAND_STT_OUTPUT_FORMAT fmt = str(raw).lower().strip().lstrip(".") return fmt if fmt in COMMAND_STT_OUTPUT_FORMATS else DEFAULT_COMMAND_STT_OUTPUT_FORMAT def _read_command_stt_output(output_path: Path, stdout: str, fmt: str) -> str: - """Return the transcript text from a command-provider invocation. + """Return the transcript: non-empty output file > non-empty stdout (curl one-liners) > RuntimeError. - Resolution: - 1. If ``output_path`` exists and is non-empty → read it (raw text). - 2. Else if ``stdout`` is non-empty → use stdout (lets users write - curl-style one-liners that emit transcript to stdout instead of - writing a file). - 3. Else → raise RuntimeError (no usable output produced). - - For JSON format, we still return the raw bytes — extracting a - ``text`` field is out of scope; users either configure ``format: txt`` - or post-process JSON downstream. (Same trade-off as TTS: the runner - doesn't try to be clever about output shape.) + JSON output is returned raw — users configure ``format: txt`` or post-process. """ if output_path.exists(): try: @@ -572,6 +462,10 @@ def _read_command_stt_output(output_path: Path, stdout: str, fmt: str) -> str: ) +def _log_prompt_unsupported(label: str) -> None: + logger.debug("%s does not support transcription prompts — proceeding without the prompt.", label) + + def _transcribe_command_stt( file_path: str, provider_name: str, @@ -583,46 +477,24 @@ def _transcribe_command_stt( ) -> Dict[str, Any]: """Transcribe via a user-declared ``stt.providers.: type: command``. - Placeholder grammar: - - | Placeholder | Substituted with | - |-------------------|-----------------------------------------------------------| - | ``{input_path}`` | absolute path to the audio file (original location) | - | ``{output_path}`` | absolute path the provider should write its transcript to | - | ``{output_dir}`` | parent dir of ``{output_path}`` | - | ``{format}`` | configured output format (``txt`` / ``json`` / ``srt`` / ``vtt``) | - | ``{language}`` | configured language code (default ``en``) | - | ``{model}`` | configured model id (empty when not set) | - - All placeholders are shell-quote-aware (see ``_render_command_stt_template``). - Doubled braces ``{{`` and ``}}`` are preserved as literal braces. - - Returns the standard transcribe-response envelope (``success``, - ``transcript``, ``provider``, ``error``). + Placeholders (all shell-quote-aware; ``{{``/``}}`` stay literal): + ``{input_path}`` original audio path, ``{output_path}`` file to write the + transcript to, ``{output_dir}`` its parent, ``{format}`` txt/json/srt/vtt, + ``{language}`` (default ``en``), ``{model}`` (empty when unset). """ if prompt: - logger.debug( - "Command STT provider '%s' does not support transcription " - "prompts — proceeding without the prompt.", provider_name, - ) + _log_prompt_unsupported(f"Command STT provider '{provider_name}'") + + def fail(error: str) -> Dict[str, Any]: + return _error_result(error, provider=provider_name) command_template = str(config.get("command") or "").strip() if not command_template: - return { - "success": False, - "transcript": "", - "provider": provider_name, - "error": f"stt.providers.{provider_name}.command is not configured", - } + return fail(f"stt.providers.{provider_name}.command is not configured") audio = Path(file_path).expanduser() if not audio.exists(): - return { - "success": False, - "transcript": "", - "provider": provider_name, - "error": f"Audio file not found: {file_path}", - } + return fail(f"Audio file not found: {file_path}") timeout = _get_command_stt_timeout(config) output_format = _get_command_stt_output_format(config) @@ -657,15 +529,7 @@ def _transcribe_command_stt( env_passthrough=_command_stt_env_passthrough(config), ) except subprocess.TimeoutExpired: - return { - "success": False, - "transcript": "", - "provider": provider_name, - "error": ( - f"STT command provider '{provider_name}' timed out after " - f"{timeout:g}s" - ), - } + return fail(f"STT command provider '{provider_name}' timed out after {timeout:g}s") except subprocess.CalledProcessError as exc: detail_parts = [] if exc.stderr: @@ -673,53 +537,154 @@ def _transcribe_command_stt( if exc.stdout: detail_parts.append(f"stdout: {exc.stdout.strip()}") detail = "; ".join(detail_parts) or "no command output" - return { - "success": False, - "transcript": "", - "provider": provider_name, - "error": ( - f"STT command provider '{provider_name}' exited with code " - f"{exc.returncode}: {detail}" - ), - } + return fail( + f"STT command provider '{provider_name}' exited with code " + f"{exc.returncode}: {detail}" + ) try: transcript_text = _read_command_stt_output( output_path, result.stdout or "", output_format, ) except RuntimeError as exc: - return { - "success": False, - "transcript": "", - "provider": provider_name, - "error": str(exc), - } + return fail(str(exc)) except OSError as exc: - return { - "success": False, - "transcript": "", - "provider": provider_name, - "error": f"STT command provider '{provider_name}' failed: {exc}", - } + return fail(f"STT command provider '{provider_name}' failed: {exc}") logger.info( "Transcribed %s via command STT provider '%s' (%d chars)", audio.name, provider_name, len(transcript_text), ) - return { - "success": True, - "transcript": transcript_text, - "provider": provider_name, - } + return _ok_result(transcript_text, provider_name) + + +# --------------------------------------------------------------------------- +# Provider resolution +# --------------------------------------------------------------------------- + + +def _has_xai_stt_credentials() -> bool: + from tools.xai_http import resolve_xai_http_credentials + + return bool(resolve_xai_http_credentials().get("api_key")) + + +def _has_xai_stt_credentials_quietly() -> bool: + try: + return _has_xai_stt_credentials() + except Exception: + return False + + +def _has_key(env_var: str, provider: str, *, needs_openai: bool = False, needs_mistral: bool = False): + """Availability probe factory: optional SDK flag AND a resolvable API key.""" + def probe() -> bool: + if needs_openai and not _HAS_OPENAI: + return False + if needs_mistral and not _HAS_MISTRAL: + return False + return bool(_resolve_provider_key(env_var, provider)) + return probe + + +_has_groq_key = _has_key("GROQ_API_KEY", "groq", needs_openai=True) +_has_mistral_key = _has_key("MISTRAL_API_KEY", "mistral", needs_mistral=True) +_has_elevenlabs_key = _has_key("ELEVENLABS_API_KEY", "elevenlabs") +_has_deepinfra_key = _has_key("DEEPINFRA_API_KEY", "deepinfra", needs_openai=True) + + +def _resolve_explicit_openai() -> str: + if not _HAS_OPENAI: + logger.warning("STT provider 'openai' configured but no API key available") + return "none" + # Resolve directly rather than via the boolean probe so a managed + # openai-audio gateway outage is logged with its real reason, not a + # generic "no API key" hint. + try: + _resolve_openai_audio_client_config() + return "openai" + except ValueError as exc: + logger.warning("STT provider 'openai' configured but unavailable: %s", exc) + return "none" + + +def _resolve_explicit_local() -> str: + if _HAS_FASTER_WHISPER: + return "local" + if _has_local_command(): + return "local_command" + if _try_lazy_install_stt(): + return "local" + logger.warning( + "STT provider 'local' configured but unavailable " + "(install faster-whisper or set HERMES_LOCAL_STT_COMMAND)" + ) + return "none" + + +def _resolve_explicit_local_command() -> str: + if _has_local_command(): + return "local_command" + if _HAS_FASTER_WHISPER: + logger.info("Local STT command unavailable, using local faster-whisper") + return "local" + logger.warning("STT provider 'local_command' configured but unavailable") + return "none" + + +def _explicit_cloud_resolver(name: str, probe, warning: str): + def resolve() -> str: + if probe(): + return name + logger.warning(warning) + return "none" + return resolve + + +# Explicit ``stt.provider`` selections -> resolver returning the provider or "none". +_EXPLICIT_PROVIDER_RESOLVERS = { + "local": _resolve_explicit_local, + "local_command": _resolve_explicit_local_command, + "openai": _resolve_explicit_openai, + "groq": _explicit_cloud_resolver( + "groq", _has_groq_key, "STT provider 'groq' configured but GROQ_API_KEY not set"), + "mistral": _explicit_cloud_resolver( + "mistral", _has_mistral_key, + "STT provider 'mistral' configured but mistralai package " + "not installed or MISTRAL_API_KEY not set"), + "xai": _explicit_cloud_resolver( + "xai", _has_xai_stt_credentials, "STT provider 'xai' configured but no xAI credentials are available"), + "elevenlabs": _explicit_cloud_resolver( + "elevenlabs", _has_elevenlabs_key, "STT provider 'elevenlabs' configured but ELEVENLABS_API_KEY not set"), + "deepinfra": _explicit_cloud_resolver( + "deepinfra", _has_deepinfra_key, + "STT provider 'deepinfra' configured but DEEPINFRA_API_KEY not set " + "(or openai package missing)"), +} + +# Auto-detect ladder for cloud providers, in priority order: (check, name, log). +# DeepInfra is LAST so a DEEPINFRA_API_KEY set for the chat surface never +# displaces an existing xAI/ElevenLabs auto-selection. Mistral only +# auto-selects when the SDK is already present — no lazy-install during +# passive auto-detection (explicit ``provider: mistral`` installs on first use). +_AUTO_DETECT_CLOUD = ( + (_has_groq_key, "groq", "No local STT available, using Groq Whisper API"), + (lambda: _HAS_OPENAI and _has_openai_audio_backend(), + "openai", "No local STT available, using OpenAI Whisper API"), + (_has_mistral_key, "mistral", "No local STT available, using Mistral Voxtral Transcribe API"), + (_has_xai_stt_credentials_quietly, "xai", "No local STT available, using xAI Grok STT API"), + (_has_elevenlabs_key, "elevenlabs", "No local STT available, using ElevenLabs Scribe STT API"), + (_has_deepinfra_key, "deepinfra", "No local STT available, using DeepInfra Whisper API"), +) def _get_provider(stt_config: dict) -> str: """Determine which STT provider to use. - When ``stt.provider`` is explicitly set in config, that choice is - honoured — no silent cloud fallback. When no provider is configured, - auto-detect tries: local > groq (free) > openai (paid). + An explicit ``stt.provider`` is honoured — no silent cloud fallback. With no + provider configured, auto-detect tries local > groq > openai > mistral > xai + > elevenlabs > deepinfra. """ if not is_stt_enabled(stt_config): return "none" @@ -727,20 +692,17 @@ def _get_provider(stt_config: dict) -> str: explicit = "provider" in stt_config provider = stt_config.get("provider", DEFAULT_PROVIDER) - # The managed "Nous Subscription" selection (stt.provider: nous) is - # serviced by the OpenAI provider implementation, routed through the - # managed openai-audio gateway by _resolve_openai_audio_client_config. + # The managed "Nous Subscription" selection is serviced by the OpenAI + # implementation, routed through the managed gateway by + # _resolve_openai_audio_client_config. if isinstance(provider, str) and provider.strip().lower() == "nous": provider = "openai" if explicit and provider == "local": - # Legacy DEFAULT_CONFIG seeded ``stt.provider: local`` on every - # install, so a merged-config "local" is not proof of a user pick. - # ``read_selection`` reads the raw config.yaml: when the raw file - # holds an stt selection (picker- or hand-written ``local``) it is - # honored; when the merged "local" came only from a legacy default - # merge, take the autodetect branch (which prefers local first - # anyway, so a genuine local user is unaffected when it's available). + # Legacy DEFAULT_CONFIG seeded ``stt.provider: local`` on every install, + # so a merged-config "local" is not proof of a user pick. Only a raw + # config.yaml selection counts as explicit; otherwise autodetect (which + # prefers local first anyway). try: from tools.tool_backend_helpers import read_selection @@ -749,162 +711,37 @@ def _get_provider(stt_config: dict) -> str: except Exception: # pragma: no cover — helpers are in-repo pass - # --- Explicit provider: respect the user's choice ---------------------- - if explicit: - if provider == "local": - if _HAS_FASTER_WHISPER: - return "local" - if _has_local_command(): - return "local_command" - # Try lazy-install before giving up - if _try_lazy_install_stt(): - return "local" - logger.warning( - "STT provider 'local' configured but unavailable " - "(install faster-whisper or set HERMES_LOCAL_STT_COMMAND)" - ) - return "none" - - if provider == "local_command": - if _has_local_command(): - return "local_command" - if _HAS_FASTER_WHISPER: - logger.info("Local STT command unavailable, using local faster-whisper") - return "local" - logger.warning( - "STT provider 'local_command' configured but unavailable" - ) - return "none" - - if provider == "groq": - if _HAS_OPENAI and _resolve_provider_key("GROQ_API_KEY", "groq"): - return "groq" - logger.warning( - "STT provider 'groq' configured but GROQ_API_KEY not set" - ) - return "none" - - if provider == "openai": - if _HAS_OPENAI: - # Resolve directly instead of via the boolean probe: the - # probe flattens _resolve_openai_audio_client_config's - # selection-specific ValueError into False, so a managed - # openai-audio gateway outage would be logged as a generic - # "no API key" hint (#93045). - try: - _resolve_openai_audio_client_config() - return "openai" - except ValueError as exc: - logger.warning( - "STT provider 'openai' configured but unavailable: %s", exc - ) - return "none" - logger.warning( - "STT provider 'openai' configured but no API key available" - ) - return "none" - - if provider == "mistral": - if _HAS_MISTRAL and _resolve_provider_key("MISTRAL_API_KEY", "mistral"): - return "mistral" - logger.warning( - "STT provider 'mistral' configured but mistralai package " - "not installed or MISTRAL_API_KEY not set" - ) - return "none" - - if provider == "xai": - from tools.xai_http import resolve_xai_http_credentials - - if resolve_xai_http_credentials().get("api_key"): - return "xai" - logger.warning( - "STT provider 'xai' configured but no xAI credentials are available" - ) - return "none" - - if provider == "elevenlabs": - if _resolve_provider_key("ELEVENLABS_API_KEY", "elevenlabs"): - return "elevenlabs" - logger.warning( - "STT provider 'elevenlabs' configured but ELEVENLABS_API_KEY not set" - ) - return "none" - - if provider == "deepinfra": - if _HAS_OPENAI and _resolve_provider_key("DEEPINFRA_API_KEY", "deepinfra"): - return "deepinfra" - logger.warning( - "STT provider 'deepinfra' configured but DEEPINFRA_API_KEY not set " - "(or openai package missing)" - ) - return "none" - - return provider # Unknown — let it fail downstream - - # --- Auto-detect (no explicit provider): - # local > groq > openai > mistral > xai > elevenlabs > deepinfra --- - # DeepInfra is tried LAST so adding DEEPINFRA_API_KEY (commonly set for the - # chat surface) never silently displaces an existing xAI/ElevenLabs STT - # auto-selection; a DeepInfra-only box still resolves to it. mistral is - # intentionally skipped while `mistralai` is quarantined on PyPI (malicious - # 2.4.6 release on 2026-05-12). + resolver = _EXPLICIT_PROVIDER_RESOLVERS.get(provider) + return resolver() if resolver else provider # Unknown — let it fail downstream if _HAS_FASTER_WHISPER: return "local" if _has_local_command(): return "local_command" - # Try lazy-install before falling through to cloud providers if _try_lazy_install_stt(): return "local" - if _HAS_OPENAI and _resolve_provider_key("GROQ_API_KEY", "groq"): - logger.info("No local STT available, using Groq Whisper API") - return "groq" - if _HAS_OPENAI and _has_openai_audio_backend(): - logger.info("No local STT available, using OpenAI Whisper API") - return "openai" - # Only auto-select Mistral if the SDK is already present — don't trigger a - # lazy-install during passive auto-detection. Explicit `provider: mistral` - # (above) does lazy-install on first transcription call. - if _HAS_MISTRAL and _resolve_provider_key("MISTRAL_API_KEY", "mistral"): - logger.info("No local STT available, using Mistral Voxtral Transcribe API") - return "mistral" - try: - from tools.xai_http import resolve_xai_http_credentials - - if resolve_xai_http_credentials().get("api_key"): - logger.info("No local STT available, using xAI Grok STT API") - return "xai" - except Exception: - pass - if _resolve_provider_key("ELEVENLABS_API_KEY", "elevenlabs"): - logger.info("No local STT available, using ElevenLabs Scribe STT API") - return "elevenlabs" - if _HAS_OPENAI and _resolve_provider_key("DEEPINFRA_API_KEY", "deepinfra"): - logger.info("No local STT available, using DeepInfra Whisper API") - return "deepinfra" + for available, name, message in _AUTO_DETECT_CLOUD: + if available(): + logger.info(message) + return name return "none" def _unregistered_stt_provider_error(provider: str) -> Dict[str, Any]: key = str(provider or "").strip() - return { - "success": False, - "transcript": "", - "provider": key, - "error_type": "provider_not_registered", - "error": ( - f"stt.provider='{key}' is set but no built-in, command, or plugin " - "provider registered that name. Run `hermes plugins list` to see " - "installed STT plugins, or configure a command provider under " - f"`stt.providers.{key}.command`." - ), - } + return _error_result( + f"stt.provider='{key}' is set but no built-in, command, or plugin " + "provider registered that name. Run `hermes plugins list` to see " + "installed STT plugins, or configure a command provider under " + f"`stt.providers.{key}.command`.", + provider=key, + error_type="provider_not_registered", + ) # --------------------------------------------------------------------------- -# Plugin provider dispatch (issue follow-up to #30398 — STT pluggability) +# Plugin provider dispatch # --------------------------------------------------------------------------- @@ -917,50 +754,21 @@ def _dispatch_to_plugin_provider( language: Optional[str] = None, prompt: Optional[str] = None, ) -> Optional[Dict[str, Any]]: - """Route the call to a plugin-registered transcription provider, or - return None. + """Route to a plugin-registered transcription provider; None when no plugin claims the name. - Returns the transcribe-response dict on dispatch, or ``None`` when no - plugin claimed the provider name. - - Resolution invariants enforced here: - - 1. Built-in provider names short-circuit — never reach the plugin - registry. The caller (``transcribe_audio``) handles ``local``, - ``groq``, ``openai``, etc. via its existing elif chain; this - function defensively rejects those names so a plugin can't be - silently dispatched under a built-in name even if it somehow - slipped past the registry's built-in shadow guard. - 2. Same-name command-type provider declared under - ``stt.providers.: type: command`` wins over a plugin. The - caller short-circuits to the command runner before reaching us, - but we re-verify here so a refactor of the caller can't silently - break the invariant (matches TTS PR #17843 precedence rule). - 3. Plugin dispatch fires only when ``provider`` matches a - registered :class:`TranscriptionProvider` whose ``name`` equals - the configured value. Unknown names with no plugin registered - return None (caller surfaces the configured-provider error when - the name came from ``stt.provider``). - 4. Availability gating: when the matched plugin reports - ``is_available() == False`` (missing API key, missing optional - SDK, etc.) this returns an error envelope identifying the - plugin as unavailable — **not** ``None`` — because the user - explicitly opted into this plugin via ``stt.provider`` and the - generic fallthrough message would be misleading. - - Provider exceptions are caught and converted into the standard - error envelope (matches the legacy built-in error shapes — the - gateway/CLI caller already expects ``{success: False, error: - "...", transcript: ""}`` on failure). + Invariants (re-verified here even though the caller short-circuits first, + so a caller refactor can't silently break them): built-in names never reach + the registry; a same-name ``stt.providers.: type: command`` wins over + a plugin. A matched plugin reporting ``is_available() == False`` returns an + error envelope — not None — because the user explicitly opted in via + ``stt.provider`` and the generic fall-through message would mislead. + Provider exceptions become the standard error envelope. """ if not provider: return None key = provider.lower().strip() if key in BUILTIN_STT_PROVIDERS or key == "none": return None - # Defense in depth: command-provider check should already have - # short-circuited the caller. If a same-name command config exists, - # bail so the command path wins. if stt_config is not None and _is_command_stt_provider_config( _get_named_stt_provider_config(stt_config, key) ): @@ -972,11 +780,8 @@ def _dispatch_to_plugin_provider( _ensure_plugins_discovered() plugin_provider = get_provider(key) if plugin_provider is None: - # Long-lived sessions may have discovered plugins before a - # bundled backend was patched in or before config changed. - # Retry once with a forced refresh before surfacing fall- - # through. Mirrors the image_gen / browser dispatcher - # recovery pattern. + # Long-lived sessions may have discovered plugins before a backend + # was patched in or config changed — retry once with a forced refresh. _ensure_plugins_discovered(force=True) plugin_provider = get_provider(key) except Exception as exc: # noqa: BLE001 — discovery failure is non-fatal @@ -985,17 +790,8 @@ def _dispatch_to_plugin_provider( if plugin_provider is None: return None - # Availability gate: when a plugin reports it's not configured - # (missing API key, missing optional SDK, etc.) surface a clean - # error envelope **instead of** falling through to the generic - # "No STT provider" message. The user explicitly set - # ``stt.provider: `` in config — surfacing the plugin's - # own availability failure is more actionable than the generic - # auto-detect-failure error, and avoids routing the call into a - # plugin that's about to crash messily. - # - # ``is_available()`` MUST NOT raise per the ABC contract; defend - # anyway so a buggy plugin can't break dispatch for everyone. + # ``is_available()`` MUST NOT raise per the ABC contract; defend anyway so + # a buggy plugin can't break dispatch for everyone. try: available = plugin_provider.is_available() except Exception as exc: # noqa: BLE001 @@ -1009,21 +805,15 @@ def _dispatch_to_plugin_provider( "STT plugin provider '%s' reports not available; returning " "unavailability envelope.", key, ) - return { - "success": False, - "transcript": "", - "error": ( - f"STT plugin '{key}' is not available — check that its " - "required credentials / dependencies are configured." - ), - "provider": key, - } + return _error_result( + f"STT plugin '{key}' is not available — check that its " + "required credentials / dependencies are configured.", + provider=key, + ) logger.info("Transcribing with plugin STT provider '%s'...", key) - # Plugin providers receive the transcription prompt via the ABC's - # existing ``**extra`` kwargs — no signature change needed. The key is - # only sent when a prompt is actually set so providers that predate it - # see byte-identical calls on the no-prompt path. + # The prompt travels via the ABC's ``**extra`` kwargs and is only sent when + # set, so pre-prompt providers see byte-identical calls on the no-prompt path. extra_kwargs: Dict[str, Any] = {} if prompt is not None: extra_kwargs["prompt"] = prompt @@ -1038,45 +828,28 @@ def _dispatch_to_plugin_provider( logger.warning( "STT plugin provider '%s' raised: %s", key, exc, exc_info=True, ) - return { - "success": False, - "transcript": "", - "error": f"STT plugin '{key}' raised: {exc}", - "provider": key, - } + return _error_result(f"STT plugin '{key}' raised: {exc}", provider=key) - # Defensive: plugins should return a dict matching the contract. If - # they don't, surface a clear error envelope rather than leaking a - # weird object back to the gateway. if not isinstance(result, dict): - return { - "success": False, - "transcript": "", - "error": f"STT plugin '{key}' returned a non-dict result", - "provider": key, - } - # Stamp provider if the plugin forgot to. + return _error_result(f"STT plugin '{key}' returned a non-dict result", provider=key) result.setdefault("provider", key) return result # --------------------------------------------------------------------------- -# pre_transcription plugin hook (issue #64168 — STT prompt/vocab threading) +# pre_transcription plugin hook (STT prompt/vocab threading) # --------------------------------------------------------------------------- -# Fields a pre_transcription hook may mutate. ``file_path`` is deliberately -# absent — it is read-only; attempts to change it are logged and dropped. +# Fields a pre_transcription hook may mutate. ``file_path`` is read-only — +# attempts to change it are logged and dropped. _PRE_TRANSCRIPTION_MUTABLE_FIELDS = ("prompt", "language", "model") -# Whisper-family models silently use only the final ~224 tokens of the -# prompt/initial_prompt; longer values waste upload bytes and can trip -# stricter OpenAI-compatible servers. Enforce the cap client-side for the -# whisper-family backends: truncate with a warning, never error. -# Approximation: ~4 characters per token (no tokenizer dependency). +# Whisper-family models only use the final ~224 tokens of the prompt; longer +# values waste upload bytes and can trip stricter OpenAI-compatible servers. +# Enforced client-side (truncate with a warning, never error), ~4 chars/token. _WHISPER_PROMPT_TOKEN_CAP = 224 _PROMPT_CHARS_PER_TOKEN = 4 -# Providers whose prompt parameter feeds a whisper-family model. _WHISPER_PROMPT_CAPPED_PROVIDERS = frozenset( {"local", "openai", "groq", "deepinfra"} ) @@ -1085,12 +858,10 @@ _WHISPER_PROMPT_CAPPED_PROVIDERS = frozenset( def _enforce_prompt_length_limit( prompt: Optional[str], provider: str ) -> Optional[str]: - """Truncate *prompt* to the provider's known token cap (fail-open). + """Truncate *prompt* to the whisper-family token cap, keeping the TAIL (fail-open). - Only whisper-family backends have a documented ~224-token prompt window; - other providers (mistral, plugin providers) own their own validation. - Truncation keeps the TAIL of the prompt because whisper conditions on - the final context window — the most recently appended hints survive. + Whisper conditions on the final context window, so the most recently + appended hints survive. Other providers own their own validation. """ if not prompt or provider not in _WHISPER_PROMPT_CAPPED_PROVIDERS: return prompt @@ -1119,32 +890,20 @@ def _apply_pre_transcription_hook( ) -> tuple[Optional[str], Optional[str], Optional[str]]: """Fire the ``pre_transcription`` plugin hook and merge its results. - Mirrors the ``transform_*`` hook mechanics (``transform_tool_result``): - gated on ``has_hook`` so the no-hook dispatch path never builds hook - kwargs, and fail-open — any hook-plumbing error leaves the dispatch - untouched. ``invoke_hook`` returns results in registration order, and - plugin discovery scans plugin directories in sorted order, so multiple - plugins' hints compose deterministically (sorted by plugin id, then - each plugin's own registration order). Each dict result is applied - field-by-field on top of the previous ones, so the last hook to write - a field wins (last-writer-wins per field). + Gated on ``has_hook`` so the no-hook path never builds hook kwargs, and + fail-open: any hook-plumbing error leaves the dispatch untouched. Results + arrive in registration order (plugins discovered in sorted order) and are + applied field-by-field, so the last hook to write a field wins. Model + values are accepted as-is and flow through the same per-backend + normalization a caller-supplied model would. - Model values are accepted as-is: the dispatcher has no catalog-level - validation today, so a hook-set model flows through the exact same - per-backend normalization/auto-correction (``_normalize_local_model``, - the Groq/OpenAI cross-corrections) a caller-supplied model would, and - otherwise errors at the backend as it would today. - - Returns ``(model, language_override, prompt)``. ``language_override`` - is ``None`` unless a hook explicitly set ``language`` — backends keep - their existing config/env language resolution when no hook overrides - it. + Returns ``(model, language_override, prompt)``; ``language_override`` is + None unless a hook explicitly set ``language``, so backends keep their own + config/env language resolution. """ try: from hermes_cli.plugins import has_hook, invoke_hook - # No-hook short-circuit: keep the no-plugin dispatch path - # byte-identical (no kwargs built, no invoke_hook call). if not has_hook("pre_transcription"): return model, None, prompt @@ -1163,7 +922,6 @@ def _apply_pre_transcription_hook( continue for key, value in hook_result.items(): if key == "file_path": - # file_path is read-only for hooks — log and drop. logger.warning( "pre_transcription hook attempted to change " "file_path (read-only) — ignoring the attempt." @@ -1186,9 +944,7 @@ def _apply_pre_transcription_hook( if "model" in overrides: model = overrides["model"] if "prompt" in overrides: - # Hook results win over the static ``stt.prompt`` config value — - # config is the base, hooks mutate on top. An empty string - # clears the config prompt. + # Hooks win over the static ``stt.prompt`` config; "" clears it. prompt = overrides["prompt"] or None return model, overrides.get("language") or None, prompt except Exception as _hook_err: # noqa: BLE001 — hook plumbing is fail-open @@ -1206,13 +962,11 @@ def _validate_audio_file_size(audio_path: Path) -> Optional[Dict[str, Any]]: try: file_size = audio_path.stat().st_size except OSError as e: - return {"success": False, "transcript": "", "error": f"Failed to access file: {e}"} + return _error_result(f"Failed to access file: {e}") if file_size > MAX_FILE_SIZE: - return { - "success": False, - "transcript": "", - "error": f"File too large: {file_size / (1024*1024):.1f}MB (max {MAX_FILE_SIZE / (1024*1024):.0f}MB)", - } + return _error_result( + f"File too large: {file_size / (1024*1024):.1f}MB (max {MAX_FILE_SIZE / (1024*1024):.0f}MB)" + ) return None @@ -1225,17 +979,17 @@ def _validate_audio_source_file( audio_path = Path(file_path) if os.path.islink(audio_path): - return {"success": False, "transcript": "", "error": f"Path is a symbolic link: {file_path}"} + return _error_result(f"Path is a symbolic link: {file_path}") if not audio_path.exists(): - return {"success": False, "transcript": "", "error": f"Audio file not found: {file_path}"} + return _error_result(f"Audio file not found: {file_path}") if not audio_path.is_file(): - return {"success": False, "transcript": "", "error": f"Path is not a file: {file_path}"} + return _error_result(f"Path is not a file: {file_path}") if enforce_size_limit: return _validate_audio_file_size(audio_path) try: audio_path.stat() except OSError as e: - return {"success": False, "transcript": "", "error": f"Failed to access file: {e}"} + return _error_result(f"Failed to access file: {e}") return None @@ -1251,13 +1005,11 @@ def _validate_audio_file( if source_error: return source_error - audio_path = Path(file_path) - if audio_path.suffix.lower() not in SUPPORTED_FORMATS: - return { - "success": False, - "transcript": "", - "error": f"Unsupported format: {audio_path.suffix}. Supported: {', '.join(sorted(SUPPORTED_FORMATS))}", - } + suffix = Path(file_path).suffix + if suffix.lower() not in SUPPORTED_FORMATS: + return _error_result( + f"Unsupported format: {suffix}. Supported: {', '.join(sorted(SUPPORTED_FORMATS))}" + ) return None @@ -1269,19 +1021,17 @@ def _prepare_audio_for_transcription( if audio_path.suffix.lower() != ".silk": return file_path, None, None if not _HAS_PILK: - # pilk is a tiny silk-v3 codec binding — lazy-install it on first - # .silk voice note instead of bloating the base install. + # pilk is a tiny silk-v3 codec binding — lazy-install on first .silk + # voice note instead of bloating the base install. try: from tools.lazy_deps import ensure as _lazy_ensure _lazy_ensure("stt.silk", prompt=False) except Exception: pass if not _safe_find_spec("pilk"): - return None, None, { - "success": False, - "transcript": "", - "error": "Unsupported format: .silk. Install the optional 'pilk' dependency to enable WeChat voice transcription.", - } + return None, None, _error_result( + "Unsupported format: .silk. Install the optional 'pilk' dependency to enable WeChat voice transcription." + ) temp_dir = tempfile.mkdtemp(prefix="hermes-silk-") converted_path = os.path.join(temp_dir, f"{audio_path.stem}.wav") @@ -1295,47 +1045,28 @@ def _prepare_audio_for_transcription( except Exception as exc: shutil.rmtree(temp_dir, ignore_errors=True) logger.error("Failed to convert .silk audio %s: %s", file_path, exc, exc_info=True) - return None, None, { - "success": False, - "transcript": "", - "error": f"Failed to convert .silk audio for transcription: {exc}", - } + return None, None, _error_result(f"Failed to convert .silk audio for transcription: {exc}") + # --------------------------------------------------------------------------- # Provider: local (faster-whisper) # --------------------------------------------------------------------------- -# Substrings that identify a missing/unloadable CUDA runtime library. When -# ctranslate2 (the backend for faster-whisper) cannot dlopen one of these, the -# "auto" device picker has already committed to CUDA and the model can no -# longer be used — we fall back to CPU and reload. -# -# Deliberately narrow: we match on library-name tokens and dlopen phrasing so -# we DO NOT accidentally catch legitimate runtime failures like "CUDA out of -# memory" — those should surface to the user, not silently fall back to CPU -# (a 32GB audio clip on CPU at int8 isn't useful either). +# Substrings identifying a missing/unloadable CUDA runtime library: when +# ctranslate2 can't dlopen one of these the "auto" device picker has already +# committed to CUDA, so we fall back to CPU and reload. Deliberately narrow +# (library names + dlopen phrasing) so legitimate runtime failures like "CUDA +# out of memory" surface to the user instead of silently running on CPU. _CUDA_LIB_ERROR_MARKERS = ( - "libcublas", - "libcudnn", - "libcudart", - "cannot be loaded", - "cannot open shared object", - "no kernel image is available", - "CUBLAS_STATUS_NOT_SUPPORTED", - "no CUDA-capable device", + "libcublas", "libcudnn", "libcudart", "cannot be loaded", "cannot open shared object", + "no kernel image is available", "CUBLAS_STATUS_NOT_SUPPORTED", "no CUDA-capable device", "CUDA driver version is insufficient", ) def _looks_like_cuda_lib_error(exc: BaseException) -> bool: - """Heuristic: is this exception a missing/broken CUDA runtime library? - - ctranslate2 raises plain RuntimeError with messages like - ``Library libcublas.so.12 is not found or cannot be loaded``. We want to - catch missing/unloadable shared libs and driver-mismatch errors, NOT - legitimate runtime failures ("CUDA out of memory", model bugs, etc.). - """ + """Heuristic: is this a missing/broken CUDA runtime library (not a legitimate runtime failure)?""" msg = str(exc) return any(marker in msg for marker in _CUDA_LIB_ERROR_MARKERS) @@ -1354,33 +1085,21 @@ def _sysctl_value(name: str) -> str: def _should_force_faster_whisper_cpu() -> bool: - """Avoid faster-whisper device autodetection paths known to hard-abort. - - On Apple Silicon, especially when Python is running as x86_64 under - Rosetta, ctranslate2's ``device=\"auto\"`` path can abort inside native - code before Python can catch an exception. Force CPU so local STT remains - reliable for gateway voice messages. - """ + """Force CPU on Apple Silicon (incl. x86_64 under Rosetta), where ctranslate2's + ``device="auto"`` can abort inside native code before Python can catch it.""" if platform.system() != "Darwin": return False - - machine = platform.machine().lower() - if machine in {"arm64", "aarch64"}: + if platform.machine().lower() in {"arm64", "aarch64"}: return True - - # Under Rosetta, platform.machine() reports x86_64. sysctl.proc_translated - # tells us this process is translated, while hw.optional.arm64 distinguishes - # Apple Silicon hosts from Intel Macs. + # Under Rosetta platform.machine() reports x86_64; sysctl.proc_translated + # flags translation and hw.optional.arm64 distinguishes Apple Silicon hosts. if _sysctl_value("sysctl.proc_translated") == "1": return True return _sysctl_value("hw.optional.arm64") == "1" def _get_idle_unload_seconds(local_cfg: Dict[str, Any]) -> int: - """Resolve the idle unload timeout from config. - - 0 = never unload (default). Negative values are treated as 0. - """ + """Resolve the idle unload timeout from config; 0 = never (default), negatives clamp to 0.""" try: val = int(local_cfg.get("unload_after_idle_seconds", 0)) except (TypeError, ValueError): @@ -1389,11 +1108,7 @@ def _get_idle_unload_seconds(local_cfg: Dict[str, Any]) -> int: def _unload_local_model() -> None: - """Release the cached local whisper model and free its memory. - - Safe to call from any thread. The model lock prevents races with a - concurrent transcription that is mid-load. - """ + """Release the cached local whisper model. Thread-safe via the model lock.""" global _local_model, _local_model_name with _local_model_lock: if _local_model is not None: @@ -1406,19 +1121,14 @@ def _unload_local_model() -> None: def _start_idle_unload_watcher(timeout_seconds: int) -> None: - """Ensure the idle-unload watcher thread is running. + """Ensure the single idle-unload watcher thread is running. - A single long-lived watcher: started only when none is alive, so the - per-transcription cost is one lock + one ``is_alive()`` check — no - stop/join/restart churn on the response path. The loop re-reads the - configured timeout from config every cycle, so changing - ``stt.local.unload_after_idle_seconds`` takes effect within one check - interval without a restart. After unloading (or when the timeout is set - to 0/never, or the model is already gone) the thread exits; the next - transcription restarts it. - - ``timeout_seconds`` seeds the first cycle so a just-written config is - honored even if a concurrent config read would race. + Started only when none is alive (one lock + one ``is_alive()`` per + transcription). The loop re-reads ``stt.local.unload_after_idle_seconds`` + every cycle so config edits apply within one interval; ``timeout_seconds`` + seeds the first cycle so a just-written config is honored even if a + concurrent read races. After unloading, when the timeout becomes 0, or when + the model is already gone, the thread exits; the next transcription restarts it. """ global _idle_unload_thread with _idle_unload_mgmt_lock: @@ -1432,8 +1142,6 @@ def _start_idle_unload_watcher(timeout_seconds: int) -> None: break if _local_model is None: break - # Re-read the timeout each cycle: config edits apply without - # waiting for the next voice message. try: timeout = _get_idle_unload_seconds( _load_stt_config().get("local") or {} @@ -1442,8 +1150,7 @@ def _start_idle_unload_watcher(timeout_seconds: int) -> None: timeout = initial_timeout if timeout <= 0: break # unload disabled mid-flight — stand down - idle_for = time.monotonic() - _last_transcription_time - if idle_for >= timeout: + if time.monotonic() - _last_transcription_time >= timeout: _unload_local_model() break @@ -1463,26 +1170,15 @@ def _touch_transcription_time() -> None: def _load_local_whisper_model(model_name: str, device: str = "auto", compute_type: str = "auto"): """Load faster-whisper with graceful CUDA → CPU fallback. - faster-whisper's ``device="auto"`` picks CUDA when the ctranslate2 wheel - ships CUDA shared libs, even on hosts where the NVIDIA runtime - (``libcublas.so.12`` / ``libcudnn*``) isn't installed — common on WSL2 - without CUDA-on-WSL, headless servers, and CPU-only developer machines. - On those hosts the load itself sometimes succeeds and the dlopen failure - only surfaces at first ``transcribe()`` call. - - ``device`` / ``compute_type`` default to ``"auto"`` so the historical - behaviour is unchanged; pass explicit values from ``stt.local.device`` / - ``stt.local.compute_type`` to pin a configuration (#9088). - - We try the requested config first (fast CUDA path when it works), and on - any CUDA library load failure fall back to CPU + int8. + ``device="auto"`` picks CUDA whenever the ctranslate2 wheel ships CUDA libs, + even on hosts without the NVIDIA runtime (WSL2, headless servers, CPU-only + dev boxes). Try the requested config first; on a CUDA library load failure + fall back to CPU + int8. Pass ``stt.local.device`` / ``compute_type`` to pin. """ force_cpu = _should_force_faster_whisper_cpu() if force_cpu: - # Importing ctranslate2/faster-whisper itself can abort on some - # Apple Silicon/Rosetta installs because multiple Intel OpenMP runtimes - # are already loaded. Set this before importing faster_whisper so the - # gateway survives, then keep inference on CPU to avoid device probing. + # Importing ctranslate2 can itself abort on Apple Silicon/Rosetta when + # multiple Intel OpenMP runtimes are loaded — set before the import. os.environ.setdefault("KMP_DUPLICATE_LIB_OK", "TRUE") from faster_whisper import WhisperModel @@ -1506,35 +1202,36 @@ def _load_local_whisper_model(model_name: str, device: str = "auto", compute_typ return WhisperModel(model_name, device="cpu", compute_type="int8") -# Silence-hallucination hardening defaults for local faster-whisper. -# Whisper decodes SOMETHING even from pure silence/noise — often short junk -# tokens ("You", "Thank you.", other-language phrases). Three layers kill the -# class at the source (all tunable under ``stt.local``): -# 1. vad_filter (Silero VAD, bundled with faster-whisper): silence never -# reaches the model. ``stt.local.vad: false`` restores raw behavior -# (e.g. transcribing music/ambient audio). -# 2. condition_on_previous_text=False: one hallucinated token can't seed a -# run of them; negligible quality cost for voice-note-length audio. -# 3. Segment confidence gate (see _is_hallucinated_segment): drops segments -# the model itself flags as probably-not-speech AND low-confidence. +# Silence-hallucination hardening for local faster-whisper (whisper decodes +# junk like "You"/"Thank you." from pure silence). Three layers, all tunable +# under ``stt.local``: Silero VAD so silence never reaches the model +# (``vad: false`` restores raw behaviour for music/ambient audio); +# condition_on_previous_text=False so one hallucinated token can't seed a run; +# and the segment confidence gate in _is_hallucinated_segment. _VAD_MIN_SILENCE_MS_DEFAULT = 500 _NO_SPEECH_PROB_THRESHOLD_DEFAULT = 0.6 _LOGPROB_THRESHOLD_DEFAULT = -1.0 +def _config_number(cfg: Dict[str, Any], key: str, default, cast=float): + """Read ``cfg[key]`` through *cast*, falling back to *default* on bad values.""" + try: + return cast(cfg.get(key, default)) + except (TypeError, ValueError): + return default + + def build_local_transcribe_kwargs(stt_config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: """Build the kwargs for EVERY local faster-whisper ``model.transcribe`` call. - Single owner for the anti-hallucination hardening — any new local-whisper - call site must go through this helper instead of hand-rolling kwargs. + Single owner for the anti-hallucination hardening — new local-whisper call + sites must go through here instead of hand-rolling kwargs. """ stt_config = stt_config if isinstance(stt_config, dict) else _load_stt_config() local_cfg = stt_config.get("local") or {} kwargs: Dict[str, Any] = { "beam_size": 5, - # Don't feed the previous window's text back as a prompt: a single - # hallucinated token otherwise seeds a self-reinforcing run of them. "condition_on_previous_text": False, } @@ -1543,25 +1240,19 @@ def build_local_transcribe_kwargs(stt_config: Optional[Dict[str, Any]] = None) - vad_enabled = True if bool(vad_enabled): kwargs["vad_filter"] = True - try: - min_silence_ms = int( - local_cfg.get("vad_min_silence_ms", _VAD_MIN_SILENCE_MS_DEFAULT) + kwargs["vad_parameters"] = { + "min_silence_duration_ms": _config_number( + local_cfg, "vad_min_silence_ms", _VAD_MIN_SILENCE_MS_DEFAULT, int ) - except (TypeError, ValueError): - min_silence_ms = _VAD_MIN_SILENCE_MS_DEFAULT - kwargs["vad_parameters"] = {"min_silence_duration_ms": min_silence_ms} + } else: kwargs["vad_filter"] = False - # Push the confidence gate down into faster-whisper itself. Without this the - # library's own internal defaults (no_speech_threshold=0.6, log_prob_ - # threshold=-1.0) drop low-confidence segments BEFORE they reach our - # _is_hallucinated_segment post-filter, so the ``stt.local`` threshold knobs - # were dead for that first gate. Non-English speech decodes at a lower - # avg_logprob, so the English-tuned defaults silently discard whole - # utterances. Mapping the same config values through keeps both gates in - # sync and makes the knobs actually usable. Defaults are unchanged, so - # behavior is identical unless a user tunes them. + # Push the confidence gate into faster-whisper itself: its internal + # defaults drop low-confidence segments BEFORE our post-filter sees them, + # so without this the ``stt.local`` threshold knobs were dead for that + # first gate (non-English speech decodes at lower avg_logprob and was + # silently discarded). Same values feed both gates; defaults unchanged. no_speech_threshold, log_prob_threshold = _confidence_thresholds(local_cfg) kwargs["no_speech_threshold"] = no_speech_threshold kwargs["log_prob_threshold"] = log_prob_threshold @@ -1579,26 +1270,18 @@ def build_local_transcribe_kwargs(stt_config: Optional[Dict[str, Any]] = None) - def _confidence_thresholds(local_cfg: Dict[str, Any]) -> tuple[float, float]: """Resolve (no_speech_prob, avg_logprob) gate thresholds from config.""" - try: - no_speech = float( - local_cfg.get("no_speech_prob_threshold", _NO_SPEECH_PROB_THRESHOLD_DEFAULT) - ) - except (TypeError, ValueError): - no_speech = _NO_SPEECH_PROB_THRESHOLD_DEFAULT - try: - logprob = float(local_cfg.get("logprob_threshold", _LOGPROB_THRESHOLD_DEFAULT)) - except (TypeError, ValueError): - logprob = _LOGPROB_THRESHOLD_DEFAULT - return no_speech, logprob + return ( + _config_number(local_cfg, "no_speech_prob_threshold", _NO_SPEECH_PROB_THRESHOLD_DEFAULT), + _config_number(local_cfg, "logprob_threshold", _LOGPROB_THRESHOLD_DEFAULT), + ) def _is_hallucinated_segment(segment: Any, no_speech_threshold: float, logprob_threshold: float) -> bool: """True when a segment is very likely a silence hallucination. - Conservative AND gate (matches openai-whisper's own heuristic): the model - must BOTH think the window is non-speech (high no_speech_prob) AND have - decoded it with low confidence (low avg_logprob). Quiet-but-real speech - fails one of the two conditions and survives. + Conservative AND gate (openai-whisper's own heuristic): the model must BOTH + think the window is non-speech AND have decoded it with low confidence, so + quiet-but-real speech survives. Unknown segment shapes are never dropped. """ no_speech_prob = getattr(segment, "no_speech_prob", None) avg_logprob = getattr(segment, "avg_logprob", None) @@ -1608,7 +1291,6 @@ def _is_hallucinated_segment(segment: Any, no_speech_threshold: float, logprob_t no_speech_prob = float(no_speech_prob) avg_logprob = float(avg_logprob) except (TypeError, ValueError): - # Unknown segment shape (plugin/test doubles) — never drop. return False return no_speech_prob > no_speech_threshold and avg_logprob < logprob_threshold @@ -1630,6 +1312,31 @@ def _join_confident_segments(segments: Any, local_cfg: Dict[str, Any]) -> str: return " ".join(kept).strip() +def _get_or_load_local_model(model_name: str, local_cfg: Dict[str, Any]): + """Return the cached faster-whisper model, (re)loading under the lock when needed. + + Double-checked lock: concurrent voice messages must not both download/load. + The returned strong reference stays valid even if the idle watcher nulls the + module global mid-transcription. + """ + global _local_model, _local_model_name + model = _local_model + if model is None or _local_model_name != model_name: + with _local_model_lock: + if _local_model is None or _local_model_name != model_name: + logger.info("Loading faster-whisper model '%s' (first load downloads the model)...", model_name) + # stt.local.device / compute_type let users pin a configuration + # where ``auto`` mis-detects; the loader keeps the CUDA→CPU fallback. + _local_model = _load_local_whisper_model( + model_name, + device=local_cfg.get("device", "auto"), + compute_type=local_cfg.get("compute_type", "auto"), + ) + _local_model_name = model_name + model = _local_model + return model + + def _transcribe_local( file_path: str, model_name: str, @@ -1640,65 +1347,34 @@ def _transcribe_local( """Transcribe using faster-whisper (local, free).""" global _local_model, _local_model_name - if not _HAS_FASTER_WHISPER: - if not _try_lazy_install_stt(): - return {"success": False, "transcript": "", "error": "faster-whisper not installed"} + if not _HAS_FASTER_WHISPER and not _try_lazy_install_stt(): + return _error_result("faster-whisper not installed") try: local_cfg = _load_stt_config().get("local") or {} - # Reset the idle timer BEFORE loading/transcribing so the idle-unload - # watcher can't count a long in-flight transcription as idle time and - # unload mid-use. + # Reset the idle timer BEFORE loading/transcribing so the watcher can't + # count a long in-flight transcription as idle time and unload mid-use. _touch_transcription_time() - # Lazy-load the model (downloads on first use, ~150 MB for 'base'). - # Double-checked lock: concurrent voice messages must not both - # download/load the model (#24767). - # ``model`` is a strong local reference bound under the lock: the idle - # watcher may null the module global at any time, but this - # transcription keeps using the instance it grabbed. - model = _local_model - if model is None or _local_model_name != model_name: - with _local_model_lock: - if _local_model is None or _local_model_name != model_name: - logger.info("Loading faster-whisper model '%s' (first load downloads the model)...", model_name) - # Honour stt.local.device / stt.local.compute_type from config so - # users on hosts where ``auto`` mis-detects (NVIDIA libs present but - # not usable, etc.) can pin a working configuration (#9088). - # _load_local_whisper_model retains the CUDA→CPU fallback for the - # auto/CUDA paths. - _local_model = _load_local_whisper_model( - model_name, - device=local_cfg.get("device", "auto"), - compute_type=local_cfg.get("compute_type", "auto"), - ) - _local_model_name = model_name - model = _local_model - + model = _get_or_load_local_model(model_name, local_cfg) if model is None: # defensive: load failed without raising - return {"success": False, "transcript": "", "error": "Local whisper model failed to load"} - # Shared hardened kwargs: VAD filter (default on), no cross-window - # conditioning, language/initial_prompt resolution — one owner for - # every local faster-whisper call site. + return _error_result("Local whisper model failed to load") + stt_config = _load_stt_config() local_config = stt_config.get("local") or {} transcribe_kwargs = build_local_transcribe_kwargs(stt_config) - # pre_transcription hook overrides win over the config-resolved - # values from build_local_transcribe_kwargs. + # pre_transcription hook overrides win over config-resolved values. if language: transcribe_kwargs["language"] = language if prompt: - # faster-whisper's vocabulary/context hint parameter. transcribe_kwargs["initial_prompt"] = prompt try: segments, info = model.transcribe(file_path, **transcribe_kwargs) transcript = _join_confident_segments(segments, local_config) except Exception as exc: - # CUDA runtime libs sometimes only fail at dlopen-on-first-use, - # AFTER the model loaded successfully. Evict the broken cached - # model, reload on CPU, retry once. Without this the module- - # global `_local_model` is poisoned and every subsequent voice - # message on this process fails identically until restart. + # CUDA libs sometimes only fail at dlopen-on-first-use, AFTER the + # model loaded. Evict the poisoned cached model, reload on CPU and + # retry once — otherwise every later voice message fails until restart. if not _looks_like_cuda_lib_error(exc): raise logger.warning( @@ -1724,11 +1400,11 @@ def _transcribe_local( if idle_timeout > 0: _start_idle_unload_watcher(idle_timeout) - return {"success": True, "transcript": transcript, "provider": "local"} + return _ok_result(transcript, "local") except Exception as e: logger.error("Local transcription failed: %s", e, exc_info=True) - return {"success": False, "transcript": "", "error": f"Local transcription failed: {e}"} + return _error_result(f"Local transcription failed: {e}") def _prepare_local_audio(file_path: str, work_dir: str) -> tuple[Optional[str], Optional[str]]: @@ -1742,10 +1418,8 @@ def _prepare_local_audio(file_path: str, work_dir: str) -> tuple[Optional[str], return None, "Local STT fallback requires ffmpeg for non-WAV inputs, but ffmpeg was not found" converted_path = os.path.join(work_dir, f"{audio_path.stem}.wav") - command = [ffmpeg, "-y", "-i", file_path, converted_path] - try: - subprocess.run(command, check=True, capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=300, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags()) + _run_quiet([ffmpeg, "-y", "-i", file_path, converted_path], timeout=300) return converted_path, None except subprocess.TimeoutExpired: logger.error("ffmpeg conversion timed out for %s", file_path) @@ -1791,33 +1465,23 @@ def _transcribe_local_command( ) -> Dict[str, Any]: """Run the configured local STT command template and read back a .txt transcript.""" if prompt: - logger.debug( - "STT provider 'local_command' does not support transcription " - "prompts — proceeding without the prompt." - ) + _log_prompt_unsupported("STT provider 'local_command'") command_template = _get_local_command_template() if not command_template: - return { - "success": False, - "transcript": "", - "error": ( - f"{LOCAL_STT_COMMAND_ENV} not configured and no local whisper binary was found" - ), - } + return _error_result( + f"{LOCAL_STT_COMMAND_ENV} not configured and no local whisper binary was found" + ) - # Language: hook override > stt.local.language > stt.language > env var - # > "en" default. - language = ( - language or _resolve_stt_language("local") or DEFAULT_LOCAL_STT_LANGUAGE - ) + # Language: hook override > stt.local.language > stt.language > env > "en". + language = language or _resolve_stt_language("local") or DEFAULT_LOCAL_STT_LANGUAGE normalized_model = _normalize_local_command_model(model_name) try: with tempfile.TemporaryDirectory(prefix="hermes-local-stt-") as output_dir: prepared_input, prep_error = _prepare_local_audio(file_path, output_dir) if prep_error: - return {"success": False, "transcript": "", "error": prep_error} + return _error_result(prep_error) command = command_template.format( input_path=shlex.quote(prepared_input), @@ -1825,32 +1489,17 @@ def _transcribe_local_command( language=shlex.quote(language), model=shlex.quote(normalized_model), ) - # Scrub Hermes secrets from the child env (sibling path to #56332 / - # _run_command_stt — this local-whisper path previously inherited - # the full process environment). + # Scrub Hermes secrets from the child env (same policy as _run_command_stt). from tools.environments.local import hermes_subprocess_env - child_env = hermes_subprocess_env(inherit_credentials=False) - subprocess.run( - shlex.split(command), - check=True, - capture_output=True, - text=True, - encoding="utf-8", - errors="replace", - timeout=300, - stdin=subprocess.DEVNULL, - env=child_env, - creationflags=windows_hide_flags(), + _run_quiet( + shlex.split(command), timeout=300, + env=hermes_subprocess_env(inherit_credentials=False), ) txt_files = sorted(Path(output_dir).glob("*.txt")) if not txt_files: - return { - "success": False, - "transcript": "", - "error": "Local STT command completed but did not produce a .txt transcript", - } + return _error_result("Local STT command completed but did not produce a .txt transcript") transcript_text = txt_files[0].read_text(encoding="utf-8").strip() logger.info( @@ -1859,27 +1508,52 @@ def _transcribe_local_command( normalized_model, len(transcript_text), ) - return {"success": True, "transcript": transcript_text, "provider": "local_command"} + return _ok_result(transcript_text, "local_command") except KeyError as e: - return { - "success": False, - "transcript": "", - "error": f"Invalid {LOCAL_STT_COMMAND_ENV} template, missing placeholder: {e}", - } + return _error_result(f"Invalid {LOCAL_STT_COMMAND_ENV} template, missing placeholder: {e}") except subprocess.CalledProcessError as e: details = e.stderr.strip() or e.stdout.strip() or str(e) logger.error("Local STT command failed for %s: %s", file_path, details) - return {"success": False, "transcript": "", "error": f"Local STT failed: {details}"} + return _error_result(f"Local STT failed: {details}") except Exception as e: logger.error("Unexpected error during local command transcription: %s", e, exc_info=True) - return {"success": False, "transcript": "", "error": f"Local transcription failed: {e}"} + return _error_result(f"Local transcription failed: {e}") + # --------------------------------------------------------------------------- -# Provider: groq (Whisper API — free tier) +# OpenAI-SDK-shaped providers: groq, openai (+ deepinfra via openai) # --------------------------------------------------------------------------- +def _close_client(client: Any) -> None: + close = getattr(client, "close", None) + if callable(close): + close() + + +def _openai_sdk_failure(exc: BaseException, file_path: str, log_label: str) -> Dict[str, Any]: + """Map an OpenAI-SDK-shaped exception to the shared error envelope. + + Order matters: APIConnectionError is checked before APITimeoutError (its + subclass) so timeouts report as connection errors, as they always have. + """ + try: + from openai import APIError, APIConnectionError, APITimeoutError + except ImportError: # pragma: no cover — callers gate on _HAS_OPENAI + APIError = APIConnectionError = APITimeoutError = () + if isinstance(exc, PermissionError): + return _error_result(f"Permission denied: {file_path}") + if isinstance(exc, APIConnectionError): + return _error_result(f"Connection error: {exc}") + if isinstance(exc, APITimeoutError): + return _error_result(f"Request timeout: {exc}") + if isinstance(exc, APIError): + return _error_result(f"API error: {exc}") + logger.error("%s transcription failed: %s", log_label, exc, exc_info=True) + return _error_result(f"Transcription failed: {exc}") + + def _transcribe_groq( file_path: str, model_name: str, @@ -1889,28 +1563,25 @@ def _transcribe_groq( ) -> Dict[str, Any]: """Transcribe using Groq Whisper API (free tier available). - Honours an optional ISO-639-1 language hint resolved from a - ``pre_transcription`` hook override > ``stt.groq.language`` > - ``stt.language`` (config.yaml) > ``HERMES_LOCAL_STT_LANGUAGE`` (env). - When none is set, Groq Whisper auto-detects. + Language: hook override > ``stt.groq.language`` > ``stt.language`` > env; + otherwise Groq auto-detects. """ api_key = _resolve_provider_key("GROQ_API_KEY", "groq") if not api_key: - return {"success": False, "transcript": "", "error": "GROQ_API_KEY not set"} + return _error_result("GROQ_API_KEY not set") if not _HAS_OPENAI: - return {"success": False, "transcript": "", "error": "openai package not installed"} + return _error_result("openai package not installed") # Auto-correct model if caller passed an OpenAI-only model if model_name in OPENAI_MODELS: logger.info("Model %s not available on Groq, using %s", model_name, DEFAULT_GROQ_STT_MODEL) model_name = DEFAULT_GROQ_STT_MODEL - # Language: hook override > stt.groq.language > stt.language > env. language = language or _resolve_stt_language("groq") try: - from openai import OpenAI, APIError, APIConnectionError, APITimeoutError + from openai import OpenAI client = OpenAI(api_key=api_key, base_url=GROQ_BASE_URL, timeout=30, max_retries=0) try: create_kwargs = { @@ -1920,8 +1591,7 @@ def _transcribe_groq( if language: create_kwargs["language"] = language if prompt: - # Only send the prompt when set so the no-hook, no-config - # request stays byte-identical to today's. + # Only sent when set so the no-hook, no-config request stays byte-identical. create_kwargs["prompt"] = prompt with open(file_path, "rb") as audio_file: transcription = client.audio.transcriptions.create( @@ -1933,27 +1603,12 @@ def _transcribe_groq( logger.info("Transcribed %s via Groq API (%s, lang=%s, %d chars)", Path(file_path).name, model_name, language or "auto", len(transcript_text)) - return {"success": True, "transcript": transcript_text, "provider": "groq"} + return _ok_result(transcript_text, "groq") finally: - close = getattr(client, "close", None) - if callable(close): - close() + _close_client(client) - except PermissionError: - return {"success": False, "transcript": "", "error": f"Permission denied: {file_path}"} - except APIConnectionError as e: - return {"success": False, "transcript": "", "error": f"Connection error: {e}"} - except APITimeoutError as e: - return {"success": False, "transcript": "", "error": f"Request timeout: {e}"} - except APIError as e: - return {"success": False, "transcript": "", "error": f"API error: {e}"} except Exception as e: - logger.error("Groq transcription failed: %s", e, exc_info=True) - return {"success": False, "transcript": "", "error": f"Transcription failed: {e}"} - -# --------------------------------------------------------------------------- -# Provider: openai (Whisper API) -# --------------------------------------------------------------------------- + return _openai_sdk_failure(e, file_path, "Groq") def _transcribe_openai( @@ -1968,42 +1623,31 @@ def _transcribe_openai( ) -> Dict[str, Any]: """Transcribe via the OpenAI ``audio.transcriptions.create`` SDK shape. - Also serves as the shared backend for every OpenAI-compatible STT - endpoint (DeepInfra etc.) — callers pass an explicit ``api_key`` / - ``base_url`` to skip the OpenAI-only auth chain, and a - ``provider_label`` so the response carries the right ``provider`` - name. + Shared backend for every OpenAI-compatible STT endpoint (DeepInfra etc.): + callers pass explicit ``api_key``/``base_url`` to skip the OpenAI-only auth + chain and a ``provider_label`` for the response's ``provider``. """ if api_key is None: try: api_key, fallback_base = _resolve_openai_audio_client_config() except ValueError as exc: - return {"success": False, "transcript": "", "error": str(exc)} + return _error_result(str(exc)) base_url = base_url or fallback_base - # Language: hook override > stt..language > stt.language > - # env > auto-detect. Explicit language hint improves accuracy for - # non-English languages. + # Language: hook override > stt..language > stt.language > env > auto. language = language or _resolve_stt_language(provider_label) if not _HAS_OPENAI: - return {"success": False, "transcript": "", "error": "openai package not installed"} + return _error_result("openai package not installed") - # Auto-correct model if caller passed a Groq-only model. Only applies - # to the native OpenAI path — third-party endpoints may legitimately - # serve a whisper-large-v3 variant. + # Auto-correct a Groq-only model on the native OpenAI path only — + # third-party endpoints may legitimately serve a whisper-large-v3 variant. if provider_label == "openai" and model_name in GROQ_MODELS: logger.info("Model %s not available on OpenAI, using %s", model_name, DEFAULT_STT_MODEL) model_name = DEFAULT_STT_MODEL try: - from openai import ( - OpenAI, - APIError, - APIConnectionError, - APITimeoutError, - BadRequestError, - ) + from openai import OpenAI, BadRequestError client = OpenAI(api_key=api_key, base_url=base_url, timeout=30, max_retries=0) def _create_transcription(path: str): @@ -2015,16 +1659,14 @@ def _transcribe_openai( } if language: if model_name == "gpt-transcribe": - # gpt-transcribe replaces the singular ``language`` - # field with a ``languages`` list; the API rejects - # requests that send the legacy field. + # gpt-transcribe replaces ``language`` with a ``languages`` + # list and rejects requests sending the legacy field. create_kwargs["extra_body"] = {"languages": [language]} else: create_kwargs["language"] = language logger.debug("Using language hint '%s' for OpenAI STT", language) if prompt: - # Only send the prompt when set so the no-hook, no-config - # request stays byte-identical to today's. + # Only sent when set so the no-hook, no-config request stays byte-identical. create_kwargs["prompt"] = prompt return client.audio.transcriptions.create(**create_kwargs) @@ -2036,12 +1678,11 @@ def _transcribe_openai( message = str(exc).lower() if not any(k in message for k in ("unsupported", "corrupted", "invalid file")): raise - # Newer models (e.g. gpt-4o-transcribe) reject some containers - # whisper-1 accepted (notably Ogg/Opus voice notes). Transcode - # to a compact .m4a and retry once. + # Newer models reject some containers whisper-1 accepted + # (notably Ogg/Opus voice notes): transcode to m4a, retry once. converted_path, transcode_error = _transcode_audio_for_stt(file_path, work_dir) if transcode_error: - return {"success": False, "transcript": "", "error": transcode_error} + return _error_result(transcode_error) logger.info( "Retrying %s STT after transcoding %s to m4a (API rejected the original container)", provider_label, Path(file_path).name, @@ -2054,23 +1695,13 @@ def _transcribe_openai( Path(file_path).name, provider_label, model_name, len(transcript_text), ) - return {"success": True, "transcript": transcript_text, "provider": provider_label} + return _ok_result(transcript_text, provider_label) finally: - close = getattr(client, "close", None) - if callable(close): - close() + _close_client(client) - except PermissionError: - return {"success": False, "transcript": "", "error": f"Permission denied: {file_path}"} - except APIConnectionError as e: - return {"success": False, "transcript": "", "error": f"Connection error: {e}"} - except APITimeoutError as e: - return {"success": False, "transcript": "", "error": f"Request timeout: {e}"} - except APIError as e: - return {"success": False, "transcript": "", "error": f"API error: {e}"} except Exception as e: - logger.error("%s transcription failed: %s", provider_label, e, exc_info=True) - return {"success": False, "transcript": "", "error": f"Transcription failed: {e}"} + return _openai_sdk_failure(e, file_path, provider_label) + # --------------------------------------------------------------------------- # Provider: mistral (Voxtral Transcribe API) @@ -2084,14 +1715,10 @@ def _transcribe_mistral( language: Optional[str] = None, prompt: Optional[str] = None, ) -> Dict[str, Any]: - """Transcribe using Mistral Voxtral Transcribe API. - - Uses the ``mistralai`` Python SDK to call ``/v1/audio/transcriptions``. - Requires ``MISTRAL_API_KEY`` environment variable. - """ + """Transcribe with the ``mistralai`` SDK (``/v1/audio/transcriptions``); requires ``MISTRAL_API_KEY``.""" api_key = _resolve_provider_key("MISTRAL_API_KEY", "mistral") if not api_key: - return {"success": False, "transcript": "", "error": "MISTRAL_API_KEY not set"} + return _error_result("MISTRAL_API_KEY not set") try: try: @@ -2107,14 +1734,12 @@ def _transcribe_mistral( "model": model_name, "file": {"content": audio_file, "file_name": Path(file_path).name}, } - # Language: hook override > stt.mistral.language > - # stt.language > env > auto. + # Language: hook override > stt.mistral.language > stt.language > env > auto. language = language or _resolve_stt_language("mistral") if language: complete_kwargs["language"] = language if prompt: - # Only send the prompt when set so the no-hook, no-config - # request stays byte-identical to today's. + # Only sent when set so the no-hook, no-config request stays byte-identical. complete_kwargs["prompt"] = prompt result = client.audio.transcriptions.complete(**complete_kwargs) @@ -2123,20 +1748,45 @@ def _transcribe_mistral( "Transcribed %s via Mistral API (%s, %d chars)", Path(file_path).name, model_name, len(transcript_text), ) - return {"success": True, "transcript": transcript_text, "provider": "mistral"} + return _ok_result(transcript_text, "mistral") except PermissionError: - return {"success": False, "transcript": "", "error": f"Permission denied: {file_path}"} + return _error_result(f"Permission denied: {file_path}") except Exception as e: logger.error("Mistral transcription failed: %s", e, exc_info=True) - return {"success": False, "transcript": "", "error": f"Mistral transcription failed: {type(e).__name__}"} + return _error_result(f"Mistral transcription failed: {type(e).__name__}") # --------------------------------------------------------------------------- -# Provider: xAI (Grok STT API) +# REST multipart providers: xAI, ElevenLabs # --------------------------------------------------------------------------- +def _post_audio_multipart(url: str, headers: Dict[str, str], file_path: str, data: Dict[str, str]): + import requests + + with open(file_path, "rb") as audio_file: + return requests.post( + url, headers=headers, files={"file": (Path(file_path).name, audio_file)}, + data=data, timeout=120, + ) + + +def _http_error_detail(response, extract) -> str: + """``extract(json_body)`` -> detail string, falling back to the first 300 chars of the body.""" + try: + return extract(response.json()) or response.text[:300] + except Exception: + return response.text[:300] + + +def _elevenlabs_error_detail(err_body: Dict[str, Any]) -> str: + error_value = err_body.get("detail") or err_body.get("error") + if isinstance(error_value, dict): + return str(error_value.get("message") or error_value) + return str(error_value) if error_value else "" + + def _transcribe_xai( file_path: str, model_name: str, @@ -2144,24 +1794,15 @@ def _transcribe_xai( language: Optional[str] = None, prompt: Optional[str] = None, ) -> Dict[str, Any]: - """Transcribe using xAI Grok STT API. - - Uses the ``POST /v1/stt`` REST endpoint with multipart/form-data. - Supports Inverse Text Normalization, diarization, and word-level timestamps. - Requires ``XAI_API_KEY`` environment variable. - """ + """Transcribe via xAI ``POST /v1/stt`` (multipart). Supports ITN, diarization, word timestamps.""" from tools.xai_http import resolve_xai_http_credentials if prompt: - logger.debug( - "STT provider 'xai' does not support transcription prompts — " - "proceeding without the prompt." - ) + _log_prompt_unsupported("STT provider 'xai'") - # STT is an API-billed endpoint. Prefer the explicit XAI_API_KEY over the - # general xAI OAuth/Grok-subscription credential; subscription OAuth may be - # valid for Grok while returning personal-team spending-limit errors for - # /v1/stt. Other xAI integrations keep their existing resolver precedence. + # STT is API-billed: prefer the explicit XAI_API_KEY over the general xAI + # OAuth/Grok-subscription credential, which may be valid for Grok yet hit + # personal-team spending-limit errors on /v1/stt. direct_api_key = str(get_env_value("XAI_API_KEY") or "").strip() if direct_api_key: creds = { @@ -2175,11 +1816,9 @@ def _transcribe_xai( creds = resolve_xai_http_credentials() api_key = str(creds.get("api_key") or "").strip() if not api_key: - return { - "success": False, - "transcript": "", - "error": "No xAI credentials found. Configure xAI OAuth in `hermes model` or set XAI_API_KEY", - } + return _error_result( + "No xAI credentials found. Configure xAI OAuth in `hermes model` or set XAI_API_KEY" + ) stt_config = _load_stt_config() xai_config = stt_config.get("xai") or {} @@ -2201,13 +1840,10 @@ def _transcribe_xai( base_url = _resolve_base_url(creds) # Language: hook override > stt.xai.language > stt.language > env. language = language or _resolve_stt_language("xai", stt_config) or "" - # .get("format", True) already defaults to True when the key is absent; - # is_truthy_value only normalizes truthy/falsy strings from config. use_format = is_truthy_value(xai_config.get("format", True)) use_diarize = is_truthy_value(xai_config.get("diarize", False)) try: - import requests from tools.xai_http import hermes_xai_user_agent data: Dict[str, str] = {} @@ -2219,19 +1855,11 @@ def _transcribe_xai( data["diarize"] = "true" def _post_transcription(bearer: str, endpoint_base_url: str): - with open(file_path, "rb") as audio_file: - return requests.post( - f"{endpoint_base_url}/stt", - headers={ - "Authorization": f"Bearer {bearer}", - "User-Agent": hermes_xai_user_agent(), - }, - files={ - "file": (Path(file_path).name, audio_file), - }, - data=data, - timeout=120, - ) + return _post_audio_multipart( + f"{endpoint_base_url}/stt", + {"Authorization": f"Bearer {bearer}", "User-Agent": hermes_xai_user_agent()}, + file_path, data, + ) response = _post_transcription(api_key, base_url) @@ -2262,28 +1890,14 @@ def _transcribe_xai( ) if response.status_code != 200: - detail = "" - try: - err_body = response.json() - detail = err_body.get("error", {}).get("message", "") or response.text[:300] - except Exception: - detail = response.text[:300] - return { - "success": False, - "transcript": "", - "error": f"xAI STT API error (HTTP {response.status_code}): {detail}", - } + detail = _http_error_detail(response, lambda body: body.get("error", {}).get("message", "")) + return _error_result(f"xAI STT API error (HTTP {response.status_code}): {detail}") result = response.json() transcript_text = result.get("text", "").strip() if not transcript_text: - return { - "success": False, - "transcript": "", - "error": "xAI STT returned empty transcript", - "no_speech": True, - } + return _error_result("xAI STT returned empty transcript", no_speech=True) logger.info( "Transcribed %s via xAI Grok STT (lang=%s, %.1fs audio, %d chars)", @@ -2293,13 +1907,13 @@ def _transcribe_xai( len(transcript_text), ) - return {"success": True, "transcript": transcript_text, "provider": "xai"} + return _ok_result(transcript_text, "xai") except PermissionError: - return {"success": False, "transcript": "", "error": f"Permission denied: {file_path}"} + return _error_result(f"Permission denied: {file_path}") except Exception as e: logger.error("xAI STT transcription failed: %s", e, exc_info=True) - return {"success": False, "transcript": "", "error": f"xAI STT transcription failed: {e}"} + return _error_result(f"xAI STT transcription failed: {e}") # --------------------------------------------------------------------------- @@ -2316,14 +1930,11 @@ def _transcribe_elevenlabs( ) -> Dict[str, Any]: """Transcribe using ElevenLabs Scribe STT API.""" if prompt: - logger.debug( - "STT provider 'elevenlabs' does not support transcription " - "prompts — proceeding without the prompt." - ) + _log_prompt_unsupported("STT provider 'elevenlabs'") api_key = _resolve_provider_key("ELEVENLABS_API_KEY", "elevenlabs") if not api_key: - return {"success": False, "transcript": "", "error": "ELEVENLABS_API_KEY not set"} + return _error_result("ELEVENLABS_API_KEY not set") stt_config = _load_stt_config() elevenlabs_config = stt_config.get("elevenlabs") or {} @@ -2340,8 +1951,6 @@ def _transcribe_elevenlabs( diarize = is_truthy_value(elevenlabs_config.get("diarize", False)) try: - import requests - data: Dict[str, str] = { "model_id": model_name, "tag_audio_events": "true" if tag_audio_events else "false", @@ -2350,43 +1959,17 @@ def _transcribe_elevenlabs( if language_code: data["language_code"] = language_code - with open(file_path, "rb") as audio_file: - response = requests.post( - f"{base_url}/speech-to-text", - headers={"xi-api-key": api_key}, - files={"file": (Path(file_path).name, audio_file)}, - data=data, - timeout=120, - ) + response = _post_audio_multipart( + f"{base_url}/speech-to-text", {"xi-api-key": api_key}, file_path, data, + ) if response.status_code != 200: - detail = "" - try: - err_body = response.json() - error_value = err_body.get("detail") or err_body.get("error") - if isinstance(error_value, dict): - detail = str(error_value.get("message") or error_value) - elif error_value: - detail = str(error_value) - else: - detail = response.text[:300] - except Exception: - detail = response.text[:300] - return { - "success": False, - "transcript": "", - "error": f"ElevenLabs STT API error (HTTP {response.status_code}): {detail}", - } + detail = _http_error_detail(response, _elevenlabs_error_detail) + return _error_result(f"ElevenLabs STT API error (HTTP {response.status_code}): {detail}") - result = response.json() - transcript_text = _extract_transcript_text(result) + transcript_text = _extract_transcript_text(response.json()) if not transcript_text: - return { - "success": False, - "transcript": "", - "error": "ElevenLabs STT returned empty transcript", - "no_speech": True, - } + return _error_result("ElevenLabs STT returned empty transcript", no_speech=True) logger.info( "Transcribed %s via ElevenLabs Scribe (%s, %d chars)", @@ -2395,13 +1978,13 @@ def _transcribe_elevenlabs( len(transcript_text), ) - return {"success": True, "transcript": transcript_text, "provider": "elevenlabs"} + return _ok_result(transcript_text, "elevenlabs") except PermissionError: - return {"success": False, "transcript": "", "error": f"Permission denied: {file_path}"} + return _error_result(f"Permission denied: {file_path}") except Exception as e: logger.error("ElevenLabs STT transcription failed: %s", e, exc_info=True) - return {"success": False, "transcript": "", "error": f"ElevenLabs STT transcription failed: {e}"} + return _error_result(f"ElevenLabs STT transcription failed: {e}") # --------------------------------------------------------------------------- @@ -2416,42 +1999,26 @@ def _transcribe_deepinfra( language: Optional[str] = None, prompt: Optional[str] = None, ) -> Dict[str, Any]: - """Resolve DeepInfra credentials/model, then delegate to the OpenAI handler. - - DeepInfra's STT endpoint is OpenAI-compatible, so the actual SDK - call lives in :func:`_transcribe_openai` — this wrapper only owns - DeepInfra-specific credential and model resolution, using the shared - ``hermes_cli.models`` helpers so every DeepInfra surface resolves the - base URL and model ids identically. - """ + """Resolve DeepInfra credentials/model (via the shared ``hermes_cli.models`` + helpers), then delegate to :func:`_transcribe_openai`.""" api_key = _resolve_provider_key("DEEPINFRA_API_KEY", "deepinfra") if not api_key: - return {"success": False, "transcript": "", "error": "DEEPINFRA_API_KEY not set"} + return _error_result("DEEPINFRA_API_KEY not set") from hermes_cli.models import deepinfra_base_url, deepinfra_model_ids - stt_config = _load_stt_config() - # ``stt.deepinfra: null`` in YAML yields None, not {} — coalesce so the - # ``.get`` calls don't raise (no stt.deepinfra block in DEFAULT_CONFIG to - # deep-merge over the null). - di_config = stt_config.get("deepinfra") if isinstance(stt_config, dict) else None - if not isinstance(di_config, dict): - di_config = {} - base_url = deepinfra_base_url(di_config) + # ``stt.deepinfra: null`` in YAML yields None, not {} — coalesce. + base_url = deepinfra_base_url(_get_stt_section(_load_stt_config(), "deepinfra")) if not model_name: candidates = deepinfra_model_ids("stt") if not candidates: - return { - "success": False, - "transcript": "", - "error": ( - "No DeepInfra STT model available. Pin one in " - "config.yaml under stt.deepinfra.model, or check " - "connectivity to api.deepinfra.com so the live catalog " - "can be fetched." - ), - } + return _error_result( + "No DeepInfra STT model available. Pin one in " + "config.yaml under stt.deepinfra.model, or check " + "connectivity to api.deepinfra.com so the live catalog " + "can be fetched." + ) model_name = candidates[0] return _transcribe_openai( @@ -2469,52 +2036,31 @@ def _transcribe_deepinfra( # Cloud pre-upload silence trim # --------------------------------------------------------------------------- # -# Local faster-whisper gets Silero VAD (build_local_transcribe_kwargs) so -# silence never reaches the model. Cloud providers get no such protection: -# the raw file is uploaded, so every second of silence is paid for twice — -# once in upload time and once in per-audio-minute billing — and cloud -# Whisper hallucinates junk tokens on silent stretches exactly like local -# Whisper did before the VAD hardening. -# -# Before uploading to a built-in cloud provider we collapse long pauses with -# ffmpeg's silenceremove filter, keeping ``stt.cloud_trim_keep_ms`` of every -# pause so word boundaries and natural pacing survive. The trim is purely -# best-effort — ANY of these falls back to uploading the original untouched: -# - ``stt.cloud_trim_silence: false`` -# - ffmpeg or ffprobe not installed -# - the trim command fails or times out -# - the trimmed result is suspiciously empty (mostly-silence clip — the -# provider, not a client-side heuristic, decides whether it has speech) -# - the trim saves less than ~10% (re-encoding for nothing) -# -# Command-type and plugin providers are deliberately NOT trimmed: they may -# wrap local CLIs that want the original bytes (and may run their own VAD). +# Local faster-whisper gets Silero VAD; cloud providers get the raw file, so +# every second of silence is paid for twice (upload + per-minute billing) and +# cloud Whisper hallucinates on it. Before uploading to a built-in cloud +# provider we collapse long pauses with ffmpeg's silenceremove, keeping +# ``stt.cloud_trim_keep_ms`` of each pause so word boundaries survive. +# Purely best-effort — ANY of these uploads the original untouched: +# ``stt.cloud_trim_silence: false``, ffmpeg/ffprobe missing, trim failure or +# timeout, a ~empty result (the provider, not a dB heuristic, decides "no +# speech"), or <10% saving. Command-type and plugin providers are NOT trimmed: +# they may wrap local CLIs that want the original bytes. _CLOUD_TRIM_THRESHOLD_DB_DEFAULT = -40 # audio below this level counts as silence _CLOUD_TRIM_KEEP_MS_DEFAULT = 300 # how much of each pause survives the trim _CLOUD_TRIM_MIN_SAVING = 0.10 # use the trimmed file only when >=10% shorter _CLOUD_TRIM_MIN_RESULT_SECONDS = 0.3 # all-silence guard floor: never upload ~empty audio -# Below this duration the trim can't pay for itself: a >=10% saving on a short -# clip is ~a second of audio, several providers bill a per-request minimum -# anyway (Groq: 10s), and the encode would sit on the synchronous voice-note -# response path. Skip the whole pipeline. +# Below this the trim can't pay for itself (several providers bill a 10s +# minimum per request) and the encode would sit on the synchronous voice-note path. _CLOUD_TRIM_MIN_INPUT_SECONDS = 12.0 -# Built-in providers that upload audio to a remote API. -CLOUD_STT_PROVIDERS = frozenset(BUILTIN_STT_PROVIDERS - {"local", "local_command"}) - - -def _find_ffprobe_binary() -> Optional[str]: - return _find_binary("ffprobe") - def _probe_audio_duration(file_path: str) -> Optional[float]: """Return the audio duration in seconds via ffprobe, or None. - Canonical sync seconds-probe. ``gateway/run.py._probe_audio_duration`` - (async, returns a display string) and the Telegram adapter's - ``_probe_voice_duration_seconds`` carry local variants of the same - ffprobe invocation — keep the command shape in sync. + Canonical sync probe; ``gateway/run.py._probe_audio_duration`` and the + Telegram adapter carry local variants — keep the command shape in sync. """ ffprobe = _find_ffprobe_binary() if not ffprobe: @@ -2526,12 +2072,7 @@ def _probe_audio_duration(file_path: str) -> Optional[float]: file_path, ] try: - result = subprocess.run( - command, check=True, capture_output=True, text=True, - encoding="utf-8", errors="replace", timeout=30, - stdin=subprocess.DEVNULL, creationflags=windows_hide_flags(), - ) - return float(result.stdout.strip()) + return float(_run_quiet(command, timeout=30).stdout.strip()) except Exception: # noqa: BLE001 - probe is best-effort return None @@ -2539,17 +2080,10 @@ def _probe_audio_duration(file_path: str) -> Optional[float]: def _cloud_trim_settings(stt_config: Dict[str, Any]) -> tuple[bool, int, int]: """Resolve (enabled, threshold_db, keep_ms) for the cloud silence trim.""" cfg = stt_config if isinstance(stt_config, dict) else {} - # is_truthy_value: the module's established config-boolean normalizer — - # a YAML string "false" must disable, exactly like is_stt_enabled. + # is_truthy_value: a YAML string "false" must disable, exactly like is_stt_enabled. enabled = is_truthy_value(cfg.get("cloud_trim_silence", True), default=True) - try: - threshold_db = int(cfg.get("cloud_trim_threshold_db", _CLOUD_TRIM_THRESHOLD_DB_DEFAULT)) - except (TypeError, ValueError): - threshold_db = _CLOUD_TRIM_THRESHOLD_DB_DEFAULT - try: - keep_ms = int(cfg.get("cloud_trim_keep_ms", _CLOUD_TRIM_KEEP_MS_DEFAULT)) - except (TypeError, ValueError): - keep_ms = _CLOUD_TRIM_KEEP_MS_DEFAULT + threshold_db = _config_number(cfg, "cloud_trim_threshold_db", _CLOUD_TRIM_THRESHOLD_DB_DEFAULT, int) + keep_ms = _config_number(cfg, "cloud_trim_keep_ms", _CLOUD_TRIM_KEEP_MS_DEFAULT, int) return enabled, threshold_db, max(keep_ms, 0) @@ -2558,11 +2092,9 @@ def _trim_silence_for_cloud_stt( ) -> Optional[str]: """Return a silence-trimmed copy of *file_path* for cloud upload, or None. - ``None`` always means "upload the original file": the trim is disabled, - the tools are missing, the clip is too short for a trim to pay for - itself, the trim failed, the clip is mostly silence, or trimming would - not save enough to justify the re-encode. On success the caller owns - deleting the returned file's parent directory. + ``None`` always means "upload the original" (disabled, tools missing, clip + too short, trim failed, mostly silence, or not enough saving). On success + the caller owns deleting the returned file's parent directory. """ enabled, threshold_db, keep_ms = _cloud_trim_settings(stt_config) if not enabled: @@ -2576,8 +2108,6 @@ def _trim_silence_for_cloud_stt( logger.debug("Cloud STT silence trim skipped: could not probe %s", file_path) return None if original_duration < _CLOUD_TRIM_MIN_INPUT_SECONDS: - # Short clip: savings can't matter (some providers bill a 10s - # minimum per request anyway) — skip the encode entirely. logger.debug( "Cloud STT silence trim skipped for %s: %.1fs is below the %.0fs gate", Path(file_path).name, original_duration, _CLOUD_TRIM_MIN_INPUT_SECONDS, @@ -2602,8 +2132,6 @@ def _trim_silence_for_cloud_stt( _run_ffmpeg_stt_encode(ffmpeg, file_path, trimmed_path, audio_filter=filter_expr) trimmed_duration = _probe_audio_duration(trimmed_path) if not trimmed_duration or trimmed_duration < min_result_seconds: - # Mostly/all silence. Deciding "no speech" belongs to the - # provider, not a client-side dB heuristic — upload the original. logger.debug( "Cloud STT silence trim discarded for %s: trimmed result ~empty (%.2fs)", Path(file_path).name, trimmed_duration or 0.0, @@ -2641,51 +2169,30 @@ def _transcribe_prepared_audio( model: Optional[str] = None, source: Optional[str] = None, ) -> Dict[str, Any]: - """ - Transcribe an audio file using the configured STT provider. + """Transcribe a validated audio file with the configured STT provider. - Provider priority: - 1. User config (``stt.provider`` in config.yaml) - 2. Auto-detect: local > Groq > OpenAI > Mistral > xAI > ElevenLabs - - Args: - file_path: Absolute path to the audio file to transcribe. - model: Override the model. If None, uses config or provider default. - source: Optional caller-surface label (e.g. ``"gateway"``, - ``"voice_mode"``) forwarded to the ``pre_transcription`` - plugin hook for observability. Not used for dispatch. - - Returns: - dict with keys: - - "success" (bool): Whether transcription succeeded - - "transcript" (str): The transcribed text (empty on failure) - - "error" (str, optional): Error message if success is False - - "provider" (str, optional): Which provider was used + ``model`` overrides the config/provider default; ``source`` is a caller-surface + label (``"gateway"``, ``"voice_mode"``) forwarded to the ``pre_transcription`` + hook for observability only. Returns the standard result envelope. """ # Refuse to feed a credential / secret store (auth.json, .env, OAuth - # tokens, mcp-tokens/, ...) to an STT provider: an external provider would - # ship its plaintext contents to a third-party API. Mirrors the local-input - # read guard added to image-gen (587be5b5b) and xAI video-gen (104232979). + # tokens, ...) to an STT provider, which would ship its plaintext to a + # third-party API. Mirrors the image-gen / video-gen read guards. from agent.file_safety import get_read_block_error blocked = get_read_block_error(file_path) if blocked: - return {"success": False, "transcript": "", "error": blocked} + return _error_result(blocked) - # Apply common path validation before provider resolution so invalid files - # cannot trigger provider setup or lazy installation. The remote-upload - # size cap is enforced separately below, only for non-local providers. + # Validate before provider resolution so invalid files cannot trigger + # provider setup or lazy installation. The remote-upload size cap is + # enforced below, only for non-local providers. error = _validate_audio_file(file_path, enforce_size_limit=False) if error: return error - # Load config and determine provider stt_config = _load_stt_config() if not is_stt_enabled(stt_config): - return { - "success": False, - "transcript": "", - "error": "STT is disabled in config.yaml (stt.enabled: false).", - } + return _error_result("STT is disabled in config.yaml (stt.enabled: false).") provider = _get_provider(stt_config) if not _is_local_stt_provider(provider, stt_config): @@ -2696,16 +2203,11 @@ def _transcribe_prepared_audio( # Convert CAF (iMessage voice notes) to WAV for cloud STT providers. if Path(file_path).suffix.lower() == ".caf" and provider not in ("local", "local_command"): converted = _convert_caf_to_wav(file_path) - if converted: - file_path = converted - else: - return {"success": False, "transcript": "", - "error": "CAF audio could not be converted to WAV."} + if not converted: + return _error_result("CAF audio could not be converted to WAV.") + file_path = converted - # Pre-upload silence trim for built-in cloud providers: local whisper gets - # Silero VAD, cloud endpoints get the raw file — collapse long pauses - # client-side so silence isn't uploaded, billed, or hallucinated on. - # Best-effort: any failure uploads the original untouched. + # Best-effort pre-upload silence trim for built-in cloud providers. trim_cleanup_dir: Optional[str] = None if provider in CLOUD_STT_PROVIDERS: trimmed = _trim_silence_for_cloud_stt(file_path, stt_config) @@ -2720,6 +2222,40 @@ def _transcribe_prepared_audio( shutil.rmtree(trim_cleanup_dir, ignore_errors=True) +def _builtin_model_name(provider: str, stt_config: Dict[str, Any], model: Optional[str]) -> str: + """Resolve the model for a built-in provider: caller override > ``stt.`` config > default.""" + if model: + return model + if provider in ("local", "local_command"): + return _get_stt_section(stt_config, "local").get("model", DEFAULT_LOCAL_MODEL) + if provider == "xai": + return "grok-stt" # xAI STT takes no model parameter — logging only + cfg = _get_stt_section(stt_config, provider) + if provider == "groq": + return cfg.get("model") or DEFAULT_GROQ_STT_MODEL + if provider == "openai": + return cfg.get("model", DEFAULT_STT_MODEL) + if provider == "mistral": + return cfg.get("model", DEFAULT_MISTRAL_STT_MODEL) + if provider == "elevenlabs": + return cfg.get("model_id", DEFAULT_ELEVENLABS_STT_MODEL) + return cfg.get("model") or "" # deepinfra: resolved from the live catalog when empty + + +# Built-in provider -> handler. Looked up at call time so tests may patch the +# module-level ``_transcribe_*`` functions. +_BUILTIN_STT_HANDLERS = { + "local": lambda *a, **kw: _transcribe_local(*a, **kw), + "local_command": lambda *a, **kw: _transcribe_local_command(*a, **kw), + "groq": lambda *a, **kw: _transcribe_groq(*a, **kw), + "openai": lambda *a, **kw: _transcribe_openai(*a, **kw), + "mistral": lambda *a, **kw: _transcribe_mistral(*a, **kw), + "xai": lambda *a, **kw: _transcribe_xai(*a, **kw), + "elevenlabs": lambda *a, **kw: _transcribe_elevenlabs(*a, **kw), + "deepinfra": lambda *a, **kw: _transcribe_deepinfra(*a, **kw), +} + + def _dispatch_stt_provider( file_path: str, provider: str, @@ -2728,20 +2264,14 @@ def _dispatch_stt_provider( source: Optional[str] = None, ) -> Dict[str, Any]: """Route *file_path* to the handler for *provider* (built-in > command > plugin).""" - # Optional static transcription prompt (``stt.prompt`` in config.yaml): - # vocabulary/context hints threaded to prompt-capable backends. - # Ordering: config is the base; pre_transcription hook results mutate on - # top, in registration order, so the last hook to set a field wins. + # Static ``stt.prompt`` is the base; pre_transcription hook results mutate + # on top in registration order (last hook to set a field wins). prompt = stt_config.get("prompt") if not isinstance(prompt, str) or not prompt.strip(): prompt = None - # pre_transcription plugin hook — fires after provider resolution and - # BEFORE any backend (built-in, command-type, or plugin-registered) is - # invoked. Hooks may mutate prompt/language/model; file_path is - # read-only. The helper short-circuits on has_hook() so the no-hook - # dispatch path stays byte-identical. ``language`` stays None unless a - # hook overrides it — backends keep their own config/env resolution. + # The hook fires after provider resolution and BEFORE any backend is + # invoked; ``language`` stays None unless a hook overrides it. model, language, prompt = _apply_pre_transcription_hook( file_path=file_path, provider=provider, @@ -2750,79 +2280,18 @@ def _dispatch_stt_provider( prompt=prompt, source=source, ) - - # Whisper-family prompt windows top out around 224 tokens — truncate - # (keeping the tail) with a warning rather than erroring or letting a - # strict server reject the request. prompt = _enforce_prompt_length_limit(prompt, provider) - if provider == "local": - local_cfg = stt_config.get("local") or {} - model_name = _normalize_local_model( - model or local_cfg.get("model", DEFAULT_LOCAL_MODEL) - ) - return _transcribe_local( - file_path, model_name, language=language, prompt=prompt, - ) + handler = _BUILTIN_STT_HANDLERS.get(provider) + if handler is not None: + model_name = _builtin_model_name(provider, stt_config, model) + if provider in ("local", "local_command"): + model_name = _normalize_local_model(model_name) + return handler(file_path, model_name, language=language, prompt=prompt) - if provider == "local_command": - local_cfg = stt_config.get("local") or {} - model_name = _normalize_local_command_model( - model or local_cfg.get("model", DEFAULT_LOCAL_MODEL) - ) - return _transcribe_local_command( - file_path, model_name, language=language, prompt=prompt, - ) - - if provider == "groq": - groq_cfg = stt_config.get("groq") or {} - model_name = model or groq_cfg.get("model") or DEFAULT_GROQ_STT_MODEL - return _transcribe_groq( - file_path, model_name, language=language, prompt=prompt, - ) - - if provider == "openai": - openai_cfg = stt_config.get("openai") or {} - model_name = model or openai_cfg.get("model", DEFAULT_STT_MODEL) - return _transcribe_openai( - file_path, model_name, language=language, prompt=prompt, - ) - - if provider == "mistral": - mistral_cfg = stt_config.get("mistral") or {} - model_name = model or mistral_cfg.get("model", DEFAULT_MISTRAL_STT_MODEL) - return _transcribe_mistral( - file_path, model_name, language=language, prompt=prompt, - ) - - if provider == "xai": - # xAI Grok STT doesn't use a model parameter — pass through for logging - model_name = model or "grok-stt" - return _transcribe_xai( - file_path, model_name, language=language, prompt=prompt, - ) - - if provider == "elevenlabs": - elevenlabs_cfg = stt_config.get("elevenlabs") or {} - model_name = model or elevenlabs_cfg.get("model_id", DEFAULT_ELEVENLABS_STT_MODEL) - return _transcribe_elevenlabs( - file_path, model_name, language=language, prompt=prompt, - ) - - if provider == "deepinfra": - di_config = stt_config.get("deepinfra") # may be None (YAML null) - di_config = di_config if isinstance(di_config, dict) else {} - model_name = model or di_config.get("model") or "" - return _transcribe_deepinfra( - file_path, model_name, language=language, prompt=prompt, - ) - - # User-declared command-type provider - # (``stt.providers.: type: command``). Fires after the built-in - # elif chain — built-in names short-circuit upstream so a user's - # ``stt.providers.openai.command`` can't override the real OpenAI - # handler — and BEFORE the plugin dispatcher, because config is more - # local than a plugin install (same precedence rule as TTS PR #17843). + # User-declared command provider: after built-ins (so ``stt.providers.openai + # .command`` can't override the real handler) and BEFORE plugins, because + # config is more local than a plugin install (same precedence as TTS). command_provider_config = _resolve_command_stt_provider_config(provider, stt_config) if command_provider_config is not None: return _transcribe_command_stt( @@ -2835,28 +2304,15 @@ def _dispatch_stt_provider( prompt=prompt, ) - # Plugin-registered STT backend (e.g. OpenRouter, SenseAudio, - # Gemini-STT). Fires only when ``provider`` is neither a built-in - # nor ``"none"`` AND there is no same-name command provider. The - # dispatcher enforces built-ins-always-win + command-wins-over-plugin - # defensively. Returns None when no plugin is registered for the - # configured name; explicit configured names get a provider-specific - # error before the generic auto-detect fallback below. - # - # Plugin-scoped config namespace mirrors the built-in pattern - # (``stt.openai.model``, ``stt.mistral.model``): plugins read their - # per-provider config under ``stt.`` and the dispatcher - # forwards ``language`` from there. Top-level ``model`` argument - # overrides any config-set model. - plugin_cfg = stt_config.get(provider, {}) if isinstance(stt_config.get(provider), dict) else {} - plugin_language = language or _resolve_stt_language(provider, stt_config) - plugin_model = model or plugin_cfg.get("model") + # Plugin-registered backend. Plugins read per-provider config under + # ``stt.`` like built-ins; the ``model`` argument overrides it. + plugin_cfg = _get_stt_section(stt_config, provider) plugin_result = _dispatch_to_plugin_provider( file_path, provider, stt_config, - model=plugin_model, - language=plugin_language, + model=model or plugin_cfg.get("model"), + language=language or _resolve_stt_language(provider, stt_config), prompt=prompt, ) if plugin_result is not None: @@ -2872,28 +2328,22 @@ def _dispatch_stt_provider( return _unregistered_stt_provider_error(provider_key) # An explicit openai selection flattened to "none" carries a - # selection-specific reason (e.g. the managed openai-audio gateway is - # unavailable). Surface it — with its `hermes tools` remediation — - # instead of the all-provider setup hint (#93045). + # selection-specific reason (e.g. managed openai-audio gateway down); + # surface it with its remediation instead of the all-provider hint. if provider_key == "none" and str(stt_config.get("provider") or "") == "openai" and _HAS_OPENAI: try: _resolve_openai_audio_client_config() except ValueError as exc: - return {"success": False, "transcript": "", "error": str(exc)} + return _error_result(str(exc)) - # No provider available - return { - "success": False, - "transcript": "", - "error": ( - "No STT provider available. Install faster-whisper for free local " - f"transcription, configure {LOCAL_STT_COMMAND_ENV} or install a local whisper CLI, " - "set GROQ_API_KEY for free Groq Whisper, set MISTRAL_API_KEY for Mistral " - "Voxtral Transcribe, configure xAI OAuth or set XAI_API_KEY for xAI Grok STT, " - "set ELEVENLABS_API_KEY for ElevenLabs Scribe, or set VOICE_TOOLS_OPENAI_KEY " - "or OPENAI_API_KEY for the OpenAI Whisper API." - ), - } + return _error_result( + "No STT provider available. Install faster-whisper for free local " + f"transcription, configure {LOCAL_STT_COMMAND_ENV} or install a local whisper CLI, " + "set GROQ_API_KEY for free Groq Whisper, set MISTRAL_API_KEY for Mistral " + "Voxtral Transcribe, configure xAI OAuth or set XAI_API_KEY for xAI Grok STT, " + "set ELEVENLABS_API_KEY for ElevenLabs Scribe, or set VOICE_TOOLS_OPENAI_KEY " + "or OPENAI_API_KEY for the OpenAI Whisper API." + ) def transcribe_audio( @@ -2903,22 +2353,19 @@ def transcribe_audio( ) -> Dict[str, Any]: """Safely validate, preprocess supported inputs, and dispatch transcription. - ``source`` is an optional caller-surface label (e.g. ``"gateway"``, - ``"voice_mode"``) forwarded to the ``pre_transcription`` plugin hook for - observability. Not used for dispatch. + ``source`` is an optional caller-surface label (``"gateway"``, ``"voice_mode"``) + forwarded to the ``pre_transcription`` hook for observability only. """ - # Refuse to feed a credential / secret store (auth.json, .env, OAuth - # tokens, mcp-tokens/, ...) to an STT provider — before ANY validation or - # preprocessing, so the refusal names the real reason rather than a - # format error. Mirrors the image-gen / video-gen read guards. + # Secret-store refusal runs before ANY validation so the error names the + # real reason rather than a format error. from agent.file_safety import get_read_block_error blocked = get_read_block_error(file_path) if blocked: - return {"success": False, "transcript": "", "error": blocked} + return _error_result(blocked) - # Cap .silk sources before the decoder runs (decoder safety). For all - # other inputs the remote-upload size cap is provider-scoped and enforced - # in _transcribe_prepared_audio, so local whisper can handle big files. + # Cap .silk sources before the decoder runs (decoder safety); for all other + # inputs the upload cap is provider-scoped in _transcribe_prepared_audio, + # so local whisper can handle big files. is_silk = Path(file_path).suffix.lower() == ".silk" source_error = _validate_audio_source_file(file_path, enforce_size_limit=is_silk) if source_error: @@ -2928,11 +2375,7 @@ def transcribe_audio( if prep_error: return prep_error if prepared_path is None: - return { - "success": False, - "transcript": "", - "error": "Audio preprocessing did not produce a file for transcription.", - } + return _error_result("Audio preprocessing did not produce a file for transcription.") try: prepared_error = _validate_audio_file(prepared_path, enforce_size_limit=False) @@ -2945,12 +2388,11 @@ def transcribe_audio( def _is_local_or_private_url(url: str) -> bool: - """True when *url* points at a loopback/RFC-1918/LAN-internal host. + """True for loopback/RFC-1918/LAN-internal hosts. - Used to decide whether an empty ``stt.openai.api_key`` is acceptable: - local OpenAI-compatible STT servers (faster-whisper-server, speaches, - vLLM whisper variants...) ignore the auth header, so users shouldn't - have to write a sham ``api_key: not-needed`` in config.yaml. + Decides whether an empty ``stt.openai.api_key`` is acceptable: local + OpenAI-compatible STT servers ignore the auth header, so users shouldn't + need a sham ``api_key: not-needed``. """ try: from urllib.parse import urlparse @@ -2975,48 +2417,49 @@ def transcribe_audio_local_fallback( ) -> Dict[str, Any]: """Try an already-installed local STT backend without changing config. - This is intended for passive inbound-media recovery after the configured - provider has failed. It deliberately does not lazy-install dependencies or - fall through to another cloud provider. + For passive inbound-media recovery after the configured provider failed: + never lazy-installs or falls through to a cloud provider. """ error = _validate_audio_file(file_path) if error: return error - stt_config = _load_stt_config() - local_cfg = stt_config.get("local") or {} + local_cfg = _load_stt_config().get("local") or {} local_model = model or local_cfg.get("model", DEFAULT_LOCAL_MODEL) if _HAS_FASTER_WHISPER: - return _transcribe_local( - file_path, - _normalize_local_model(local_model), - ) + return _transcribe_local(file_path, _normalize_local_model(local_model)) if _has_local_command(): - return _transcribe_local_command( - file_path, - _normalize_local_command_model(local_model), - ) - return { - "success": False, - "transcript": "", - "error": "No installed local STT backend is available.", - "provider": "local", - } + return _transcribe_local_command(file_path, _normalize_local_command_model(local_model)) + return _error_result("No installed local STT backend is available.", provider="local") + + +def _direct_openai_credentials(cfg_api_key: str, cfg_base_url: str) -> Optional[tuple[str, str]]: + """Direct-credential ladder: config key > keyless local base_url > env key; None if none apply. + + A local OpenAI-compatible server needs no key — send a placeholder so the + SDK doesn't refuse to construct a client. + """ + if cfg_api_key: + return cfg_api_key, (cfg_base_url or OPENAI_BASE_URL) + if cfg_base_url and _is_local_or_private_url(cfg_base_url): + return "not-needed", cfg_base_url + direct_api_key = resolve_openai_audio_api_key() + if direct_api_key: + return direct_api_key, OPENAI_BASE_URL + return None def _resolve_openai_audio_client_config() -> tuple[str, str]: """Return ``(api_key, base_url)`` for the OpenAI STT client. - Strict selection semantics (switch on the stored ``stt`` provider - string; previously this resolver never read the stored gateway intent): - - ``"nous"`` (or legacy ``use_gateway: true``) → managed gateway ONLY; - unentitled/unreachable is a selection-naming error (a direct - OPENAI_API_KEY must NOT override it). - - any other stored stt provider → direct credentials ONLY; missing - credentials is a selection-naming error — no silent managed fallback. - - never-configured stt section → legacy ladder: config key → local - base_url → env key → managed gateway. + Strict selection semantics on the stored ``stt`` provider string: + - ``"nous"`` → managed gateway ONLY; unentitled/unreachable is a + selection-naming error (a direct OPENAI_API_KEY must NOT override it). + - any other stored provider → direct credentials ONLY; missing credentials + is a selection-naming error — no silent managed fallback. + - never-configured stt section → legacy ladder: direct credentials, then + the managed gateway. """ from tools.tool_backend_helpers import ( NOUS_MANAGED_PROVIDER, @@ -3024,8 +2467,7 @@ def _resolve_openai_audio_client_config() -> tuple[str, str]: selection_error, ) - stt_config = _load_stt_config() - openai_cfg = stt_config.get("openai") or {} + openai_cfg = _load_stt_config().get("openai") or {} cfg_api_key = openai_cfg.get("api_key", "") cfg_base_url = openai_cfg.get("base_url", "") @@ -3044,15 +2486,11 @@ def _resolve_openai_audio_client_config() -> tuple[str, str]: f"{managed_gateway.gateway_origin.rstrip('/')}/", "v1" ) + direct = _direct_openai_credentials(cfg_api_key, cfg_base_url) + if direct is not None: + return direct + if selected is not None: - # Stored vendor selection: direct credentials only. - if cfg_api_key: - return cfg_api_key, (cfg_base_url or OPENAI_BASE_URL) - if cfg_base_url and _is_local_or_private_url(cfg_base_url): - return "not-needed", cfg_base_url - direct_api_key = resolve_openai_audio_api_key() - if direct_api_key: - return direct_api_key, OPENAI_BASE_URL raise ValueError(selection_error( "stt", selected, @@ -3060,19 +2498,6 @@ def _resolve_openai_audio_client_config() -> tuple[str, str]: "VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set", )) - # Never-configured stt section: legacy credential ladder. - if cfg_api_key: - return cfg_api_key, (cfg_base_url or OPENAI_BASE_URL) - - # A local OpenAI-compatible server needs no key — send a placeholder so - # the SDK doesn't refuse to construct a client (#25193, credit @nnnet). - if cfg_base_url and _is_local_or_private_url(cfg_base_url): - return "not-needed", cfg_base_url - - direct_api_key = resolve_openai_audio_api_key() - if direct_api_key: - return direct_api_key, OPENAI_BASE_URL - managed_gateway = resolve_managed_tool_gateway("openai-audio") if managed_gateway is None: message = "Neither stt.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set" @@ -3091,24 +2516,14 @@ def _resolve_openai_audio_client_config() -> tuple[str, str]: def _extract_transcript_text(transcription: Any) -> str: - """Normalize text and JSON transcription responses to a plain string.""" - text: Optional[str] = None - + """Normalize text / object / dict transcription responses to a plain string.""" if isinstance(transcription, str): text = transcription.strip() - - if text is None and hasattr(transcription, "text"): - value = getattr(transcription, "text") - if isinstance(value, str): - text = value.strip() - - if text is None and isinstance(transcription, dict): - value = transcription.get("text") - if isinstance(value, str): - text = value.strip() - - if text is None: - text = str(transcription).strip() + else: + value = getattr(transcription, "text", None) + if not isinstance(value, str) and isinstance(transcription, dict): + value = transcription.get("text") + text = value.strip() if isinstance(value, str) else str(transcription).strip() match = re.match( r"\s*language\s+[\w.-]+(?:\s*[^<]*)?\s*\s*(?P.*)", diff --git a/tools/tts_streaming.py b/tools/tts_streaming.py index aa0f1e4540..b62d7747e5 100644 --- a/tools/tts_streaming.py +++ b/tools/tts_streaming.py @@ -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..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..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"].*?", 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 - ```` 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 ```` 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 diff --git a/tools/tts_tool.py b/tools/tts_tool.py index 89dfa3938d..ffed37c0e1 100644 --- a/tools/tts_tool.py +++ b/tools/tts_tool.py @@ -14,19 +14,22 @@ Built-in TTS providers: - KittenTTS (local, free, no API key): On-device 25MB model - Piper (local, free, no API key): OHF-Voice/piper1-gpl neural VITS, 44 languages -Custom command providers: -- Users can declare any number of named providers with ``type: command`` - under ``tts.providers.`` in ``~/.hermes/config.yaml``. Hermes - writes the input text to a temp file and runs the configured shell - command, which must produce the audio file at the expected path. - See the Local Command section of ``website/docs/user-guide/features/tts.md``. +Custom command providers: any number of named ``type: command`` providers under +``tts.providers.`` in ``~/.hermes/config.yaml``; Hermes writes the text to +a temp file and runs the shell template (see the Local Command section of +``website/docs/user-guide/features/tts.md``). -Output formats: -- Opus (.ogg) for Telegram voice bubbles (requires ffmpeg for Edge TTS) -- MP3 (.mp3) for everything else (CLI, Discord, WhatsApp) +Output: Opus (.ogg) for voice-bubble platforms (Telegram etc.), MP3 elsewhere. +Configuration lives under the ``tts:`` key; the user chooses provider/voice, +the model just sends text. -Configuration is loaded from ~/.hermes/config.yaml under the 'tts:' key. -The user chooses the provider and voice; the model just sends text. +Module layout: this file owns config resolution, the command/plugin provider +layers, the OpenAI/DeepInfra backends (managed-gateway aware), provider +dispatch, the lifecycle leases and the tool registration. Sibling modules: +``tts_tool_providers`` (cloud backends), ``tts_tool_local`` (on-device engines ++ model caches), ``tts_tool_delivery`` (chunking / ffmpeg / packing), +``tts_tool_speaker`` (streaming speaker pipeline). Their names are re-imported +here so ``tools.tts_tool.`` keeps resolving. Usage: from tools.tts_tool import text_to_speech_tool, check_tts_requirements @@ -35,38 +38,32 @@ Usage: """ import asyncio -import base64 import datetime import importlib.util import json import logging import os -import queue -import platform import re -import shlex -import shutil import subprocess import tempfile import threading import time import uuid -from concurrent.futures import Future, ThreadPoolExecutor -from dataclasses import dataclass, field from pathlib import Path -from typing import Callable, Dict, Any, Iterator, List, Optional, Tuple -from urllib.parse import urljoin, urlparse +from typing import Callable, Dict, Any, List, Optional +from urllib.parse import urljoin -from hermes_cli._subprocess_compat import windows_hide_flags from hermes_constants import display_hermes_home logger = logging.getLogger(__name__) + + def get_env_value(name, default=None): """Read env values through the live config module. - Tests may monkeypatch and later restore ``hermes_cli.config.get_env_value`` - before this module is imported. Resolve the helper at call time so TTS does - not keep a stale imported function for the rest of the test process. + Resolved at call time: tests monkeypatch/restore + ``hermes_cli.config.get_env_value`` and must not leave TTS holding a stale + function for the rest of the process. """ try: from hermes_cli.config import get_env_value as _get_env_value @@ -79,11 +76,9 @@ def get_env_value(name, default=None): def _resolve_provider_key(env_var: str, provider_id: str) -> str: """Resolve a TTS provider API key via the shared voice-key resolver. - Delegates to ``tools.tool_backend_helpers.resolve_provider_secret`` — - the single owner of STT/TTS key resolution (config > env/.env > the - credential pool populated by ``hermes auth add ``). - Resolved at call time so tests that reload the helpers module see the - live function. + ``tools.tool_backend_helpers.resolve_provider_secret`` is the single owner + of STT/TTS key resolution (config > env/.env > credential pool). Resolved + at call time so tests that reload the helpers module see the live function. """ try: from tools.tool_backend_helpers import resolve_provider_secret @@ -91,14 +86,13 @@ def _resolve_provider_key(env_var: str, provider_id: str) -> str: return str(get_env_value(env_var) or "").strip() return resolve_provider_secret(env_var, provider_id, env_getter=get_env_value) + from tools.managed_tool_gateway import resolve_managed_tool_gateway from tools.tts_command_provider import ( command_env_passthrough as _command_provider_env_passthrough, - quote_command_placeholder as _quote_command_tts_placeholder, render_command_template as _render_command_tts_template, run_command_provider as _run_command_tts, - shell_quote_context as _shell_quote_context, - terminate_command_process_tree as _terminate_command_tts_process_tree, + shell_quote_context as _shell_quote_context, # noqa: F401 — tests import via this module ) from tools.tool_backend_helpers import ( NOUS_MANAGED_PROVIDER, @@ -108,207 +102,194 @@ from tools.tool_backend_helpers import ( resolve_openai_audio_api_key, selection_error, ) -from tools.xai_http import hermes_xai_user_agent +from tools.tts_tool_delivery import ( # noqa: F401 — historical names re-exported + FALLBACK_MAX_TEXT_LENGTH, + AudioDeliveryProfile, + _build_audio_delivery_files, + _concat_audio_files, + _convert_to_opus, + _has_ffmpeg, + _pack_audio_files_for_delivery, + _repair_ogg_container, + _resolve_audio_delivery_profile, + _sniff_audio_container, + _split_oversized_sentence, + _split_text_for_tts, + _wrap_pcm_as_wav, +) +from tools.tts_tool_providers import ( # noqa: F401 — historical names re-exported + DEFAULT_ELEVENLABS_MODEL_ID, + DEFAULT_ELEVENLABS_STREAMING_MODEL_ID, + DEFAULT_ELEVENLABS_VOICE_ID, + DEFAULT_GEMINI_TTS_BASE_URL, + DEFAULT_GEMINI_TTS_MODEL, + DEFAULT_GEMINI_TTS_VOICE, + DEFAULT_MINIMAX_BASE_URL, + DEFAULT_MINIMAX_CN_BASE_URL, + DEFAULT_XAI_BASE_URL, + DEFAULT_XAI_VOICE_ID, + TTS_RESPONSE_BODY_LIMIT_BYTES, + _XAI_FIRST_SENTENCE_RE, + _XAI_INLINE_SPEECH_TAGS, + _XAI_WRAPPING_SPEECH_TAGS, + _apply_xai_auto_speech_tags, + _elevenlabs_environment_kwargs, + _generate_edge_tts, + _generate_elevenlabs, + _generate_gemini_tts, + _generate_minimax_tts, + _generate_mistral_tts, + _generate_xai_tts, + _read_tts_response_bytes, + _resolve_minimax_tts_runtime, + _tts_response_format_from_path, +) +from tools.tts_tool_local import ( # noqa: F401 — historical names re-exported + DEFAULT_PIPER_VOICE, + _LOCAL_TTS_MODEL_CACHES, + _TTS_MODEL_CACHE_MAX, + _generate_kittentts, + _generate_neutts, + _generate_piper_tts, + _kittentts_model_cache, + _load_kittentts_model_for_config, + _load_piper_voice_for_config, + _piper_voice_cache, + _resolve_piper_voice_path, + _tts_cache_get_or_load, +) +from tools.tts_tool_speaker import ( # noqa: F401 — historical names re-exported + stream_tts_to_speaker, +) # --------------------------------------------------------------------------- # Lazy imports -- providers are imported only when actually used to avoid # crashing in headless environments (SSH, Docker, WSL, no PortAudio). # --------------------------------------------------------------------------- -def _import_edge_tts(): - """Lazy import edge_tts. Returns the module or raises ImportError.""" +def _lazy_ensure(feature: str) -> None: + """Best-effort ``tools.lazy_deps.ensure`` so an SDK installs on first use. + + Users who enabled a provider by editing config.yaml never ran the + post-setup hook. Any failure (lazy_deps missing, install refused) falls + through so the raw import below still raises a clean ImportError. + """ try: - from tools.lazy_deps import ensure as _lazy_ensure - _lazy_ensure("tts.edge", prompt=False) - except ImportError: - pass + from tools.lazy_deps import ensure + ensure(feature, prompt=False) except Exception: pass + + +def _import_edge_tts(): + """Lazy import edge_tts. Returns the module or raises ImportError.""" + _lazy_ensure("tts.edge") import edge_tts return edge_tts -def _import_elevenlabs(): - """Lazy import ElevenLabs client. Returns the class or raises ImportError. - Calls :func:`tools.lazy_deps.ensure` first so the SDK gets installed on - demand if the user picked ElevenLabs as their TTS provider but never ran - the post-setup hook (e.g. enabled it by editing config.yaml directly). - Raises ``ImportError`` on lazy-install failure so existing callers' - error-handling paths keep working. - """ - try: - from tools.lazy_deps import FeatureUnavailable, ensure - ensure("tts.elevenlabs", prompt=False) - except ImportError: - # lazy_deps module itself missing — fall through to the raw import - # so older code paths still get a clean ImportError. - pass - except Exception: - pass +def _import_elevenlabs(): + """Lazy import the ElevenLabs client class or raise ImportError.""" + _lazy_ensure("tts.elevenlabs") from elevenlabs.client import ElevenLabs return ElevenLabs -def _elevenlabs_environment_kwargs(el_config: Dict[str, Any]) -> Dict[str, Any]: - """Build ElevenLabs client kwargs honoring config base_url/wss_url. - - ``tts.elevenlabs.base_url`` (and optionally ``wss_url``) redirect the SDK - to a self-hosted / proxy endpoint via an ``ElevenLabsEnvironment``. When - neither is set the SDK default environment is used. ``wss_url`` defaults - to the ``base_url`` host with a ``wss://`` scheme when omitted. - """ - 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 _import_openai_client(): - """Lazy import OpenAI client. Returns the class or raises ImportError.""" from openai import OpenAI as OpenAIClient return OpenAIClient -def _import_mistral_client(): - """Lazy import Mistral client. Returns the class or raises ImportError. - Calls :func:`tools.lazy_deps.ensure` first so the ``mistralai`` SDK gets - installed on demand if the user picked Mistral as their STT/TTS provider - but never ran the post-setup hook (e.g. enabled it by editing config.yaml - directly). Mirrors the ElevenLabs lazy-import path. - """ - try: - from tools.lazy_deps import ensure - ensure("tts.mistral", prompt=False) - except ImportError: - pass - except Exception: - pass +def _import_mistral_client(): + """Lazy import the Mistral client class or raise ImportError.""" + _lazy_ensure("tts.mistral") from mistralai.client import Mistral return Mistral + def _import_sounddevice(): - """Lazy import sounddevice. Returns the module or raises ImportError/OSError.""" + """Raises ImportError/OSError when PortAudio is unavailable.""" import sounddevice as sd return sd def _import_kittentts(): - """Lazy import KittenTTS. Returns the class or raises ImportError.""" from kittentts import KittenTTS return KittenTTS def _import_piper(): - """Lazy import Piper. Returns the PiperVoice class or raises ImportError. - - Piper is an optional, fully-local neural TTS engine (Home Assistant / - Open Home Foundation). ``pip install piper-tts`` provides cross-platform - wheels (Linux / macOS / Windows, x86_64 + ARM64) with embedded espeak-ng. - Voice models (.onnx + .onnx.json) are downloaded on first use. - """ + """``pip install piper-tts`` ships cross-platform wheels with embedded espeak-ng.""" from piper import PiperVoice return PiperVoice +def _package_installed(name: str) -> bool: + try: + return importlib.util.find_spec(name) is not None + except Exception: + return False + + +def _check_neutts_available() -> bool: + return _package_installed("neutts") + + +def _check_kittentts_available() -> bool: + return _package_installed("kittentts") + + +def _check_piper_available() -> bool: + return _package_installed("piper") + + # =========================================================================== # Defaults # =========================================================================== DEFAULT_PROVIDER = "edge" -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_OPENAI_MODEL = "gpt-4o-mini-tts" -# The managed OpenAI audio gateway (Nous portal proxy) only proxies these speech -# models. A user's tts.openai.model set for *direct* OpenAI (e.g. "tts-1-hd") -# is rejected with a 400 "Unsupported managed OpenAI speech model", so it must be -# coerced to a supported model when routing through the gateway. +# The managed OpenAI audio gateway (Nous portal proxy) only proxies these +# speech models; anything else is 400 "Unsupported managed OpenAI speech model". MANAGED_OPENAI_TTS_MODELS = frozenset({"gpt-4o-mini-tts"}) -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 DEFAULT_OPENAI_VOICE = "alloy" DEFAULT_OPENAI_BASE_URL = "https://api.openai.com/v1" -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 TTS `speed` accepts 0.7..1.5; 1.0 is the API default (omitted => default). -DEFAULT_XAI_SPEED_MIN = 0.7 -DEFAULT_XAI_SPEED_MAX = 1.5 -DEFAULT_XAI_SPEED_DEFAULT = 1.0 -# xAI TTS `optimize_streaming_latency` accepts 0, 1, or 2; 0 (best quality) is -# the API default (omitted => default). Values >0 trade quality for time-to-first-audio. -DEFAULT_XAI_OPTIMIZE_STREAMING_LATENCY_DEFAULT = 0 -# xAI TTS `text_normalization` is a boolean (default False). When enabled, -# the model normalizes written-form text (numbers, abbreviations, symbols) -# into spoken-form before generating audio. -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" -# Base URL now resolved via hermes_cli.models.deepinfra_base_url (shared). +# DeepInfra base URL is resolved via hermes_cli.models.deepinfra_base_url (shared). DEFAULT_DEEPINFRA_TTS_VOICE = "default" -# 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) -TTS_RESPONSE_BODY_LIMIT_BYTES = 16 * 1024 * 1024 -TTS_RESPONSE_BODY_CHUNK_BYTES = 64 * 1024 + def _get_default_output_dir() -> str: from hermes_constants import get_hermes_dir return str(get_hermes_dir("cache/audio", "audio_cache")) + DEFAULT_OUTPUT_DIR = _get_default_output_dir() _DEFAULT_OUTPUT_DIR_AT_IMPORT = DEFAULT_OUTPUT_DIR + def _default_output_dir() -> str: """Return the active profile's audio output dir at call time. - Same bug class as skills_tool (f8723c478) and skills_sync (#65828): - long-lived multi-profile runtimes (dashboard console, TUI/Desktop backend, - cron, kanban workers) import this module once under the launch - HERMES_HOME and later scope requests to a different profile via - ``hermes_constants.set_hermes_home_override()`` — a frozen module - constant keeps writing synthesized audio into the launch profile's - cache instead of the active profile's (#98749). Keep the legacy - ``DEFAULT_OUTPUT_DIR`` module attribute for tests and external patchers; - when it has not been patched, re-resolve from the live profile-scoped - HERMES_HOME on every call. + Long-lived multi-profile runtimes (dashboard, TUI/Desktop backend, cron) + import this module once and later switch profiles via + ``set_hermes_home_override()``; a frozen constant would keep writing into + the launch profile's cache. ``DEFAULT_OUTPUT_DIR`` stays as a module + attribute for tests/patchers and wins whenever it has been patched. """ configured = DEFAULT_OUTPUT_DIR if configured != _DEFAULT_OUTPUT_DIR_AT_IMPORT: return configured return _get_default_output_dir() -# --------------------------------------------------------------------------- -# Per-provider input-character limits (from official provider docs). -# A single global cap was wrong: OpenAI is 4096, xAI is 15k, MiniMax is 10k, -# ElevenLabs is model-dependent (5k / 10k / 30k / 40k), Gemini has a 32k-token -# context window. Users can override any of these via -# ``tts..max_text_length`` in config.yaml. -# --------------------------------------------------------------------------- + +# Per-provider input-character caps (from official provider docs); override +# via ``tts..max_text_length``. PROVIDER_MAX_TEXT_LENGTH: Dict[str, int] = { "edge": 5000, # edge-tts practical sync limit "openai": 4096, # https://platform.openai.com/docs/guides/text-to-speech "xai": 15000, # https://docs.x.ai/developers/model-capabilities/audio/text-to-speech "minimax": 10000, # https://platform.minimax.io/docs/api-reference/speech-t2a-http (sync) "mistral": 4000, # conservative; no published per-request cap - "gemini": 32000, # Gemini TTS has a 32k-token context window; char cap is conservative + "gemini": 32000, # 32k-token context window; char cap is conservative "elevenlabs": 10000, # fallback when model-aware lookup can't resolve (multilingual_v2) "neutts": 2000, # local model, quality falls off on long text "kittentts": 2000, # local 25MB model @@ -327,148 +308,37 @@ ELEVENLABS_MODEL_MAX_TEXT_LENGTH: Dict[str, int] = { "eleven_flash_v2_5": 40000, } - -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 _response_has_explicit_stream(response: Any) -> bool: - 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 use the streaming path above, so this fallback does not re-open - # the production eager-buffering behavior. - 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) - -# Final fallback when provider isn't recognised at all. -FALLBACK_MAX_TEXT_LENGTH = 4000 - # Back-compat alias. Prefer ``_resolve_max_text_length()`` for new code. MAX_TEXT_LENGTH = FALLBACK_MAX_TEXT_LENGTH +def _positive_int_override(value: Any) -> Optional[int]: + """A user ``max_text_length`` override, or None when absent/bool/non-positive.""" + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + return None + return value + + def _resolve_max_text_length( provider: Optional[str], tts_config: Optional[Dict[str, Any]] = None, ) -> int: """Return the input-character cap for *provider*. - Resolution order: - 1. ``tts..max_text_length`` (user override in config.yaml) - 2. ``tts.providers..max_text_length`` for user-declared - command providers - 3. ElevenLabs model-aware table (keyed on configured ``model_id``) - 4. ``PROVIDER_MAX_TEXT_LENGTH`` default - 5. ``DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH`` when the provider is a - command-type user provider without an explicit cap - 6. ``FALLBACK_MAX_TEXT_LENGTH`` (4000) - - Non-positive or non-integer overrides fall through to the default so a - broken config can't accidentally disable truncation entirely. + Order: ``tts..max_text_length`` > ElevenLabs model table > + ``PROVIDER_MAX_TEXT_LENGTH`` > command provider's own ``max_text_length`` + (else ``DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH``) > ``FALLBACK_MAX_TEXT_LENGTH``. + Non-positive / non-int overrides fall through so a broken config can't + disable truncation. """ if not provider: return FALLBACK_MAX_TEXT_LENGTH key = provider.lower().strip() cfg = tts_config or {} - # Built-in-style override at tts..max_text_length wins first, - # matching historical behavior. prov_cfg = cfg.get(key) if isinstance(cfg.get(key), dict) else {} - override = prov_cfg.get("max_text_length") if prov_cfg else None - if isinstance(override, bool): - override = None - if isinstance(override, int) and override > 0: + override = _positive_int_override(prov_cfg.get("max_text_length") if prov_cfg else None) + if override: return override if key == "elevenlabs": @@ -480,189 +350,19 @@ def _resolve_max_text_length( if key in PROVIDER_MAX_TEXT_LENGTH: return PROVIDER_MAX_TEXT_LENGTH[key] - # User-declared command provider (under tts.providers.) if key not in BUILTIN_TTS_PROVIDERS: named = _get_named_provider_config(cfg, key) if _is_command_provider_config(named): - named_override = named.get("max_text_length") - if isinstance(named_override, bool): - named_override = None - if isinstance(named_override, int) and named_override > 0: - return named_override - return DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH + return _positive_int_override(named.get("max_text_length")) or DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH return FALLBACK_MAX_TEXT_LENGTH -# =========================================================================== -# Long-form chunking and delivery packing -# =========================================================================== - -@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 per-platform overrides.""" - key = (platform or "default").lower().strip() or "default" - defaults = dict( - _PLATFORM_AUDIO_DEFAULTS.get(key) or _PLATFORM_AUDIO_DEFAULTS["default"] - ) - cfg = tts_config or {} - profiles = cfg.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 _split_oversized_sentence(sentence: str, max_chars: int) -> List[str]: - """Split one over-limit sentence on word boundaries, then hard boundaries.""" - words = sentence.split() - chunks: List[str] = [] - current = "" - for word in words: - 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] - - sentences = [ - sentence.strip() - for sentence in re.split(r"(?<=[.!?;:,])\s+", normalized) - if sentence.strip() - ] - expanded: List[str] = [] - for sentence in sentences: - if len(sentence) <= max_chars: - expanded.append(sentence) - else: - expanded.extend(_split_oversized_sentence(sentence, max_chars)) - - chunks: List[str] = [] - current = "" - for sentence in expanded: - candidate = f"{current} {sentence}".strip() - if current and len(candidate) > max_chars: - chunks.append(current) - current = sentence - else: - current = candidate - if current: - chunks.append(current) - return chunks - - -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.""" - 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_suffix = "" - current.append(path) - current_size += size - current_suffix = suffix - if current: - groups.append(current) - return groups - - # =========================================================================== # Config loader -- reads tts: section from ~/.hermes/config.yaml # =========================================================================== def _load_tts_config() -> Dict[str, Any]: - """ - Load TTS configuration from ~/.hermes/config.yaml. - - Returns a dict with provider settings. Falls back to defaults - for any missing fields. - """ + """Return the ``tts`` config section ({} when unavailable).""" try: from hermes_cli.config import load_config config = load_config() @@ -676,15 +376,13 @@ def _load_tts_config() -> Dict[str, Any]: def _get_provider(tts_config: Dict[str, Any]) -> str: - """Get the explicitly configured TTS provider or the free default. + """The explicitly configured TTS provider, or the free default. - Inference credentials do not imply consent to paid speech generation. - Users opt into cloud TTS by setting ``tts.provider`` (normally through - ``hermes tools``); otherwise the historical Edge backend remains active. - - The managed "Nous Subscription" selection (``tts.provider: nous``) is - serviced by the OpenAI provider implementation, routed through the - managed openai-audio gateway by ``_resolve_openai_audio_client_config``. + Inference credentials do not imply consent to paid speech generation: + cloud TTS is opt-in via ``tts.provider``. The managed selection + (``tts.provider: nous``) is serviced by the OpenAI implementation, routed + through the managed openai-audio gateway by + ``_resolve_openai_audio_client_config``. """ provider = (tts_config.get("provider") or DEFAULT_PROVIDER).lower().strip() if provider == NOUS_MANAGED_PROVIDER: @@ -692,96 +390,11 @@ def _get_provider(tts_config: Dict[str, Any]) -> str: return provider -@dataclass(frozen=True) -class _MiniMaxTTSRuntime: - """A region-bound MiniMax endpoint and credential. - - The credential is excluded from ``repr`` so diagnostics cannot expose it - accidentally. - """ - - 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 = {} - - credentials = { - "global": ( - "MINIMAX_API_KEY", - str(_resolve_provider_key("MINIMAX_API_KEY", "minimax") or "").strip(), - ), - "cn": ( - "MINIMAX_CN_API_KEY", - str(_resolve_provider_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, - ) - - # =========================================================================== # Custom command providers (type: command under tts.providers.) # =========================================================================== # -# Users can declare any number of command-type providers alongside the -# built-ins so they can plug any local CLI (Piper, VoxCPM, Kokoro CLIs, -# custom voice-cloning scripts, etc.) into Hermes without any Python code -# changes. The config shape is:: +# Config shape:: # # tts: # provider: piper-en @@ -791,20 +404,12 @@ def _resolve_minimax_tts_runtime( # command: "piper -m ~/model.onnx -f {output_path} < {input_path}" # output_format: wav # -# Hermes writes the input text to a temp UTF-8 file, runs the command with -# placeholder substitution, and reads the audio file the command wrote to -# ``{output_path}``. Supported placeholders: ``{input_path}``, -# ``{text_path}`` (alias for input_path), ``{output_path}``, ``{format}``, -# ``{voice}``, ``{model}``, ``{speed}``. Use ``{{`` / ``}}`` for literal braces. -# -# Built-in provider names always win over an entry with the same name under -# ``tts.providers``, so user config can't silently shadow ``edge`` etc. -# -# Placeholder values are shell-quoted for their surrounding context -# (bare / single / double quote), so paths with spaces work transparently. +# Placeholders: ``{input_path}``, ``{text_path}`` (alias), ``{output_path}``, +# ``{format}``, ``{voice}``, ``{model}``, ``{speed}``; ``{{``/``}}`` for literal +# braces. Values are shell-quoted for their surrounding quote context. Built-in +# provider names always win over a same-named entry under ``tts.providers``. -# Built-in provider names. Any ``tts.provider`` value NOT in this set is -# interpreted as a reference to ``tts.providers.``. +# Any ``tts.provider`` value NOT in this set refers to ``tts.providers.``. BUILTIN_TTS_PROVIDERS = frozenset({ "edge", "elevenlabs", @@ -826,10 +431,8 @@ COMMAND_TTS_OUTPUT_FORMATS = frozenset( ) DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH = 5000 -# Platforms whose native voice-bubble delivery requires Ogg/Opus audio. -# Previously only Telegram was recognized, so Matrix/Feishu/WhatsApp/Signal -# voice replies were synthesized as MP3 and rendered as broken attachments -# (#14841, #45557 and siblings). +# Platforms whose native voice-bubble delivery requires Ogg/Opus audio +# (MP3 renders as a broken attachment there). OPUS_VOICE_PLATFORMS = frozenset({ "telegram", "matrix", @@ -838,6 +441,11 @@ OPUS_VOICE_PLATFORMS = frozenset({ "signal", }) +# Built-ins that emit Opus natively when asked for .ogg (no ffmpeg needed). +_NATIVE_OPUS_PROVIDERS = frozenset({"openai", "elevenlabs", "mistral", "gemini"}) +# Built-ins whose native output (MP3/WAV) needs ffmpeg for voice-bubble delivery. +_FFMPEG_OPUS_PROVIDERS = frozenset({"edge", "neutts", "minimax", "xai", "kittentts", "piper"}) + def _get_provider_section(tts_config: Dict[str, Any], name: str) -> Dict[str, Any]: """Return a provider config block if it's a dict, else an empty dict.""" @@ -851,19 +459,16 @@ def _get_named_provider_config( tts_config: Dict[str, Any], name: str, ) -> Dict[str, Any]: - """Return the config dict for a user-declared provider. + """Config dict for a user-declared provider, or {}. - Looks up ``tts.providers.`` first (the canonical location), and - falls back to ``tts.`` so users who followed the built-in layout - still work. Returns an empty dict when the provider is not declared. + ``tts.providers.`` is canonical; ``tts.`` is accepted as + back-compat only for non-built-in names (so a user's ``tts.openai`` block + still means the OpenAI provider, not a custom command). """ providers = _get_provider_section(tts_config, "providers") section = providers.get(name) if isinstance(providers, dict) else None if isinstance(section, dict): return section - # Back-compat: allow ``tts.`` for user-declared providers too, - # but only when the name is not a built-in (so a user's ``tts.openai`` - # block still means the OpenAI provider, not a custom command). if name.lower() not in BUILTIN_TTS_PROVIDERS: legacy = _get_provider_section(tts_config, name) if legacy: @@ -872,7 +477,7 @@ def _get_named_provider_config( def _is_command_provider_config(config: Dict[str, Any]) -> bool: - """Return True when *config* declares a command-type provider.""" + """True when *config* declares a command-type provider (has a non-empty ``command``).""" if not isinstance(config, dict): return False ptype = str(config.get("type") or "").strip().lower() @@ -886,11 +491,10 @@ def _resolve_command_provider_config( provider: str, tts_config: Dict[str, Any], ) -> Optional[Dict[str, Any]]: - """Return the provider config if *provider* resolves to a command type. + """The provider config when *provider* is a user-declared command provider. - Built-in provider names are rejected (they have native handlers). - Returns None when the name is a built-in, unknown, or not a command - type. + None for built-in names (native handlers win), unknown names, or + non-command types. """ if not provider: return None @@ -909,39 +513,24 @@ def _dispatch_to_plugin_provider( provider: str, tts_config: Dict[str, Any], ) -> Optional[str]: - """Route the call to a plugin-registered TTS provider, or return None. + """Route to a plugin-registered TTS provider; None means "fall through". - Returns the path to the written audio file on dispatch, or ``None`` - to fall through to the next resolution layer (built-in dispatch or - Edge TTS default). + Invariants enforced here even though the caller checks them too, so a + caller refactor can't silently break them: - Resolution invariants enforced here (matches issue #30398): + 1. Built-in names never reach the plugin registry. + 2. A same-named ``type: command`` provider wins over a plugin. + 3. Dispatch fires only for a registered :class:`TTSProvider` whose name + equals the configured value; unknown names return None. - 1. Built-in provider names short-circuit — never reach the plugin - registry. The caller is responsible for the elif chain that - handles ``edge``/``openai``/etc.; this function explicitly - rejects those names defensively. - 2. Command-type providers declared under - ``tts.providers.: type: command`` (PR #17843) win over a - plugin with the same name. The caller passes us only when its - own command-provider check returned None — we re-verify here so - a refactor of the caller can't silently break the invariant. - 3. Plugin dispatch fires only when ``provider`` matches a registered - :class:`TTSProvider` whose ``name`` equals the configured value. - Unknown names return None (caller falls through to Edge default). - - Plugin exceptions are caught and re-raised — the outer - ``text_to_speech_tool`` try/except converts them to the standard - error envelope, matching how command-provider failures surface. + Plugin exceptions propagate — the outer ``text_to_speech_tool`` converts + them to the standard error envelope. """ if not provider: return None key = provider.lower().strip() if key in BUILTIN_TTS_PROVIDERS: return None - # Defense in depth: command-provider check should already have - # short-circuited the caller. If a same-name command config exists, - # bail so the command path wins. if _is_command_provider_config(_get_named_provider_config(tts_config, key)): return None try: @@ -951,11 +540,8 @@ def _dispatch_to_plugin_provider( _ensure_plugins_discovered() plugin_provider = get_provider(key) if plugin_provider is None: - # Long-lived sessions may have discovered plugins before the - # bundled backend was patched in or before config changed. - # Retry once with a forced refresh before surfacing fall- - # through. Mirrors the image_gen / browser dispatcher - # recovery pattern. + # Long-lived sessions may have discovered plugins before this one + # was installed/enabled; retry once with a forced refresh. _ensure_plugins_discovered(force=True) plugin_provider = get_provider(key) except Exception as exc: # noqa: BLE001 — discovery failure is non-fatal @@ -964,22 +550,15 @@ def _dispatch_to_plugin_provider( if plugin_provider is None: return None - # Resolve voice / model / format from tts_config — providers should - # treat all of these as optional and fall back to their own defaults - # when None is passed (matches the ABC contract documented on - # ``TTSProvider.synthesize``). - voice = tts_config.get("voice") if isinstance(tts_config, dict) else None - model = tts_config.get("model") if isinstance(tts_config, dict) else None - speed = tts_config.get("speed") if isinstance(tts_config, dict) else None - fmt = ( - tts_config.get("output_format", DEFAULT_COMMAND_TTS_OUTPUT_FORMAT) - if isinstance(tts_config, dict) - else DEFAULT_COMMAND_TTS_OUTPUT_FORMAT - ) + # voice/model/speed/format are optional per the TTSProvider.synthesize + # contract; providers fall back to their own defaults on None. + cfg = tts_config if isinstance(tts_config, dict) else {} + voice = cfg.get("voice") + model = cfg.get("model") + speed = cfg.get("speed") + fmt = cfg.get("output_format", DEFAULT_COMMAND_TTS_OUTPUT_FORMAT) - logger.info( - "Generating speech with plugin TTS provider '%s'...", key, - ) + logger.info("Generating speech with plugin TTS provider '%s'...", key) written = plugin_provider.synthesize( text, output_path, @@ -988,18 +567,14 @@ def _dispatch_to_plugin_provider( speed=float(speed) if isinstance(speed, (int, float)) else None, format=str(fmt).lower() if fmt else "mp3", ) - # Provider contract: returns the (possibly rewritten) output path. - # Defensive against a provider returning None or a non-string — - # fall back to the caller's expected output_path. + # Contract: returns the (possibly rewritten) output path; tolerate None. return written if isinstance(written, str) and written else output_path def _plugin_provider_is_voice_compatible(provider: str) -> bool: - """Return True when the registered plugin provider opts into voice - bubble delivery via its ``voice_compatible`` property. + """True when the registered plugin provider opts into voice-bubble delivery. - Defensive: any registry or property access failure means False - (matches the safe default for the command-provider path). + Any registry/property failure means False (safe default, like command providers). """ if not provider: return False @@ -1014,9 +589,7 @@ def _plugin_provider_is_voice_compatible(provider: str) -> bool: return False return bool(plugin_provider.voice_compatible) except Exception as exc: # noqa: BLE001 - logger.debug( - "tts plugin voice_compatible check failed for '%s': %s", key, exc, - ) + logger.debug("tts plugin voice_compatible check failed for '%s': %s", key, exc) return False @@ -1026,13 +599,16 @@ def _iter_command_providers(tts_config: Dict[str, Any]): return providers = _get_provider_section(tts_config, "providers") for name, cfg in (providers or {}).items(): - if isinstance(name, str) and name.lower() not in BUILTIN_TTS_PROVIDERS: - if _is_command_provider_config(cfg): - yield name, cfg + if ( + isinstance(name, str) + and name.lower() not in BUILTIN_TTS_PROVIDERS + and _is_command_provider_config(cfg) + ): + yield name, cfg def _get_command_tts_timeout(config: Dict[str, Any]) -> float: - """Return timeout in seconds, falling back when invalid.""" + """Timeout in seconds; invalid or non-positive values fall back to the default.""" raw = config.get("timeout", config.get("timeout_seconds", DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS)) try: value = float(raw) @@ -1047,7 +623,7 @@ def _get_command_tts_output_format( config: Dict[str, Any], output_path: Optional[str] = None, ) -> str: - """Return the validated output format (mp3/wav/ogg/flac).""" + """Validated output format: the output path's suffix wins, then ``format``/``output_format``.""" if output_path: suffix = Path(output_path).suffix.lower().strip().lstrip(".") if suffix in COMMAND_TTS_OUTPUT_FORMATS: @@ -1062,7 +638,7 @@ def _get_command_tts_output_format( def _is_command_tts_voice_compatible(config: Dict[str, Any]) -> bool: - """Return True only when the user explicitly opted in to voice delivery.""" + """True only when the user explicitly opted in to voice delivery.""" value = config.get("voice_compatible", False) if isinstance(value, str): return value.strip().lower() in {"1", "true", "yes", "on"} @@ -1084,9 +660,9 @@ def _generate_command_tts( ) -> str: """Generate speech by running a user-configured shell command. - Returns the absolute path of the audio file the command wrote. - Raises ``ValueError`` when the provider config is invalid, and - ``RuntimeError`` for timeouts / non-zero exits / empty output. + Returns the absolute path of the audio file the command wrote. Raises + ``ValueError`` for invalid provider config and ``RuntimeError`` for + timeouts / non-zero exits / empty output. """ command_template = str(config.get("command") or "").strip() if not command_template: @@ -1157,415 +733,9 @@ def _has_any_command_tts_provider(tts_config: Optional[Dict[str, Any]] = None) - # =========================================================================== -# ffmpeg Opus conversion (Edge TTS MP3 -> OGG Opus for Telegram) -# =========================================================================== -def _has_ffmpeg() -> bool: - """Check if ffmpeg is available on the system.""" - return shutil.which("ffmpeg") is not None - - -def _convert_to_opus(mp3_path: str) -> Optional[str]: - """ - Convert an audio file (MP3/WAV/anything ffmpeg reads) to OGG Opus - format for Telegram voice bubbles. - - Args: - mp3_path: Path to the input audio file. - - Returns: - Path to the .ogg file, or None if conversion fails. - """ - if not _has_ffmpeg(): - return None - - ogg_path = mp3_path.rsplit(".", 1)[0] + ".ogg" - return _ffmpeg_transcode_to_opus(mp3_path, ogg_path) - - -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 = subprocess.run( - ["ffmpeg", "-i", input_path, "-acodec", "libopus", - "-ac", "1", "-b:a", "48k", "-vbr", "on", - "-application", "voip", "-compression_level", "10", "-f", "ogg", - work_path, "-y"], - capture_output=True, timeout=30, - stdin=subprocess.DEVNULL, - creationflags=windows_hide_flags(), - ) - 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 — class-level guard against "MP3/WAV bytes in a .ogg -# file". Several TTS backends silently ignore the requested opus format -# (Edge only emits MP3, Piper writes WAV, xAI writes MP3, some -# OpenAI-compatible servers reject/ignore response_format="opus"), which -# breaks native voice bubbles on Telegram/Matrix/Feishu/WhatsApp. Rather -# than special-casing every provider, sniff the magic bytes once after -# synthesis and repair the container when it doesn't match the extension. -# --------------------------------------------------------------------------- - -def _sniff_audio_container(path: str) -> str: - """Return a container id ('ogg', 'wav', 'mp3', 'flac', ...) or 'unknown'. - - Delegates to the shared magic-byte sniffer in ``tools.audio_container`` - (one module owns container detection for both this outbound repair and - the inbound gateway audio cache). - """ - 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. - - When the bytes are MP3/WAV/FLAC (a backend ignored the opus request), - transcode in place to real Ogg/Opus. On any failure, rename to the - sniffed real extension so downstream players/platforms at least get an - honest file instead of a 0-second voice bubble. Returns the (possibly - updated) path. - """ - 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 - - # ffmpeg unavailable/failed: rename to the honest extension. - 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 when a custom provider did - not opt in to voice-message presentation. Matching MP3 chunks preserve their - encoded frames. A failed or unavailable combine returns ``None`` so callers - can preserve the original, individually valid files. Structured audio - containers are never byte-joined. - """ - 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 - ): - # Matching MP3 provider chunks already share one output codec/config. - # Preserve those encoded frames instead of imposing a second lossy pass. - 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. - - Packing uses the conservative target. Every combined artifact is then - checked at its actual post-encoding size; an over-limit group is split and - retried. If combining fails, the valid constituent files are returned - separately. A single final-encoded chunk above the hard limit fails closed - rather than returning an upload that the destination will reject. - """ - 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) - combined_size = Path(combined).stat().st_size - if combined_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: " - f"{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 - - -# =========================================================================== -# Provider: Edge TTS (free) -# =========================================================================== -async def _generate_edge_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str: - """ - Generate audio using Edge TTS. - - Args: - text: Text to convert. - output_path: Where to save the MP3 file. - tts_config: TTS config dict. - - Returns: - Path to the saved audio file. - """ - _edge_tts = _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 (premium) -# =========================================================================== -def _generate_elevenlabs(text: str, output_path: str, tts_config: Dict[str, Any]) -> str: - """ - Generate audio using ElevenLabs. - - Args: - text: Text to convert. - output_path: Where to save the audio file. - tts_config: TTS config dict. - - Returns: - Path to the saved audio file. - """ - api_key = (_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) - - # Determine output format based on file extension - if output_path.endswith(".ogg"): - output_format = "opus_48000_64" - else: - output_format = "mp3_44100_128" - - ElevenLabs = _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, - ) - - # audio_generator yields chunks -- write them all - with open(output_path, "wb") as f: - for chunk in audio_generator: - f.write(chunk) - - return output_path - - -def _tts_response_format_from_path(output_path: str) -> str: - """Pick an OpenAI-compatible TTS response format from the output extension.""" - if output_path.endswith(".ogg"): - return "opus" - if output_path.endswith(".wav"): - return "wav" - if output_path.endswith(".flac"): - return "flac" - return "mp3" - - -# =========================================================================== -# Provider: OpenAI TTS (also used by every OpenAI-compatible TTS endpoint — -# DeepInfra delegates here via _generate_deepinfra_tts). +# Provider: OpenAI TTS (also every OpenAI-compatible endpoint — DeepInfra +# delegates here). Kept in the origin module: it shares the managed-gateway +# selection logic below. # =========================================================================== def _generate_openai_tts( text: str, @@ -1581,35 +751,15 @@ def _generate_openai_tts( ) -> str: """Generate audio via the OpenAI ``audio.speech.create`` SDK shape. - Optional kwargs let OpenAI-compatible backends (DeepInfra etc.) reuse - this function — they resolve credentials/model themselves and pass - them through, skipping the OpenAI-only ``_resolve_openai_audio_client_config``. - - Args: - text: Text to convert. - output_path: Where to save the audio file. - tts_config: TTS config dict (used for ``tts.openai`` sub-block - and the global ``speed`` default). - api_key: Bearer token. When None, resolved from the OpenAI auth - chain (config → env → managed gateway). - base_url: API base URL. When None, falls back to - ``tts.openai.base_url`` then the OpenAI default. - model: Model id. When None, reads ``tts.openai.model``. - voice: Voice id. When None, reads ``tts.openai.voice``. - speed: Playback speed. When None, reads ``tts.openai.speed`` / - ``tts.speed``. - instructions: Optional voice-design guidance (tone, emotion, pacing, - accent, whispering). Forwarded to `audio.speech.create` when - truthy; omitted otherwise so ``tts-1``/``tts-1-hd`` and strict - OpenAI-compatible servers that reject unknown kwargs are - unaffected. - - Returns: - Path to the saved audio file. + Explicit kwargs let OpenAI-compatible backends (DeepInfra) pass their own + credentials/model/voice and skip ``_resolve_openai_audio_client_config`` + (the managed-gateway path). When None: ``api_key`` comes from the OpenAI + auth chain, ``base_url`` from ``tts.openai.base_url`` then the auth-chain + fallback then the OpenAI default, model/voice/speed from ``tts.openai`` + (speed falling back to global ``tts.speed``). ``instructions`` is + forwarded only when truthy so ``tts-1`` and strict OpenAI-compatible + servers that reject unknown kwargs are unaffected. """ - # Only resolve the OpenAI auth chain when the caller didn't pass explicit - # credentials. OpenAI-compatible backends (DeepInfra) pass api_key / - # base_url / model / voice through and never hit the managed-gateway path. fallback_base: Optional[str] = None is_managed = False explicit_base_url = base_url is not None @@ -1624,21 +774,17 @@ def _generate_openai_tts( voice = oai_config.get("voice", DEFAULT_OPENAI_VOICE) config_base_url = oai_config.get("base_url") if base_url is None: - # Config override wins over the auth-chain fallback (restores the - # pre-refactor precedence, where tts.openai.base_url beat the resolved - # default); the auth-chain value is the last-resort default. An - # explicit base_url arg from an OpenAI-compatible caller (DeepInfra) - # skips this block entirely and always wins. + # Config override beats the auth-chain fallback; an explicit arg + # (DeepInfra) skipped this block and always wins. base_url = config_base_url or fallback_base or DEFAULT_OPENAI_BASE_URL if speed is None: speed_default = tts_config.get("speed", 1.0) if isinstance(tts_config, dict) else 1.0 speed = float(oai_config.get("speed", speed_default)) language = oai_config.get("language") - # The managed OpenAI audio gateway only proxies MANAGED_OPENAI_TTS_MODELS. - # A model set for direct OpenAI (e.g. "tts-1-hd") 400s there with - # "Unsupported managed OpenAI speech model", so coerce it — unless the user - # redirected base_url to their own endpoint, in which case respect it. + # The managed gateway only proxies MANAGED_OPENAI_TTS_MODELS; coerce a + # direct-OpenAI model (e.g. "tts-1-hd") unless the user redirected + # base_url to their own endpoint. if ( is_managed and not explicit_base_url @@ -1681,25 +827,12 @@ def _generate_openai_tts( close() -# =========================================================================== -# Provider: DeepInfra TTS -# =========================================================================== -# -# DeepInfra serves TTS over an OpenAI-compatible /v1/openai/audio/speech -# endpoint. Models are discovered live via the shared catalog helper -# (filtered by the ``tts`` surface tag) — no hardcoded model ids in this -# file, so retired models disappear from hermes the next time the -# catalog is fetched without a patch. - - def _generate_deepinfra_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str: """Resolve DeepInfra credentials/model, then delegate to the OpenAI handler. - DeepInfra's audio endpoint is OpenAI-compatible, so there's no need - to duplicate the SDK call — we just pass an explicit api_key / - base_url / model / voice through. Model ids and the base URL come from - the shared ``hermes_cli.models`` helpers so every DeepInfra surface - resolves them identically. + DeepInfra's audio endpoint is OpenAI-compatible. Model ids come live from + the shared ``hermes_cli.models`` catalog helpers (no hardcoded ids, so + retired models disappear without a patch). """ api_key = _resolve_provider_key("DEEPINFRA_API_KEY", "deepinfra") if not api_key: @@ -1708,9 +841,7 @@ def _generate_deepinfra_tts(text: str, output_path: str, tts_config: Dict[str, A "or set the env var directly." ) - # ``tts.deepinfra: null`` in YAML yields None, not {} — coalesce so the - # ``.get`` calls below don't raise AttributeError (there is no - # tts.deepinfra block in DEFAULT_CONFIG to deep-merge over the null). + # ``tts.deepinfra: null`` yields None (no DEFAULT_CONFIG block to merge over). di_config = tts_config.get("deepinfra") if isinstance(tts_config, dict) else None if not isinstance(di_config, dict): di_config = {} @@ -1739,947 +870,22 @@ def _generate_deepinfra_tts(text: str, output_path: str, tts_config: Dict[str, A ) -# =========================================================================== -# Provider: xAI TTS -# =========================================================================== -_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")\]|)", - 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. - - First applies a conservative local transform (inserts [pause] between - paragraphs and after the first sentence). Then, if the result contains - no explicit user/model speech tags, asks the configured auxiliary model - to rewrite the transcript with a richer set of xAI-supported tags - (laughs, sighs, whispers, soft/loud, slow/fast, etc.) so the voice - output sounds more expressive. Falls back to the local result on any - auxiliary-model failure. - """ - clean = text.strip() - if not clean: - return text - - # Local conservative pass: pauses only. - local = clean - local = re.sub(r"\n\s*\n+", " [pause] ", local) - 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() - - # If the user/model already supplied explicit speech tags, trust them - # and don't re-rewrite. - if _XAI_SPEECH_TAG_RE.search(clean): - return local - - # Auxiliary rewrite for richer emotion tags (mirrors the Gemini path). - 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 `...` — 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 = _extract_auxiliary_message_content(response).strip() - # Strip markdown fences if the LLM wrapped the response. - fence = re.fullmatch(r"```(?:[A-Za-z0-9_-]+)?\s*(.*?)\s*```", tagged, flags=re.DOTALL) - if fence: - tagged = fence.group(1).strip() - 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 _generate_xai_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str: - """ - Generate audio using xAI TTS. - - xAI exposes a dedicated /v1/tts endpoint instead of the OpenAI audio.speech - API shape, so this is implemented as a separate backend. - """ - 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 (#87045, same root cause as x_search #88040), - # 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``; the xAI TTS API - # accepts 0.7..1.5 (1.0 = normal). Out-of-range values are clamped so a - # misconfigured agent can't 400 the request — the API would reject - # anything outside the band. - speed = xai_config.get("speed", tts_config.get("speed")) - if speed is not None and speed != "": - try: - speed = float(speed) - except (TypeError, ValueError): - speed = None - if speed is not None: - speed = max(DEFAULT_XAI_SPEED_MIN, min(DEFAULT_XAI_SPEED_MAX, speed)) - # ``tts.xai.optimize_streaming_latency`` is 0, 1, or 2 (xAI-specific; - # trades chunk-boundary quality for time-to-first-audio). - optimize_streaming_latency = xai_config.get( - "optimize_streaming_latency", - tts_config.get("optimize_streaming_latency"), - ) - if optimize_streaming_latency is not None and optimize_streaming_latency != "": - try: - optimize_streaming_latency = int(optimize_streaming_latency) - except (TypeError, ValueError): - optimize_streaming_latency = None - if optimize_streaming_latency is not None: - optimize_streaming_latency = max(0, min(2, optimize_streaming_latency)) - # ``tts.xai.text_normalization`` enables spoken-form normalization - # (numbers, abbreviations, symbols → words). Defaults to False. - 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 get_env_value("XAI_BASE_URL") - or DEFAULT_XAI_BASE_URL - ).strip().rstrip("/") - - # Match the documented minimal POST /v1/tts shape by default. Only send - # output_format when Hermes actually needs a non-default format/override. - 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 - # Only attach `speed` when the caller asked for something other than the - # API default (1.0). Keeps the existing minimal-payload contract for - # users who never touch the knob. - if speed is not None and speed != DEFAULT_XAI_SPEED_DEFAULT: - payload["speed"] = speed - # Only attach `optimize_streaming_latency` when the caller explicitly - # opts in to a non-default value (anything other than 0). - if ( - optimize_streaming_latency is not None - and optimize_streaming_latency != DEFAULT_XAI_OPTIMIZE_STREAMING_LATENCY_DEFAULT - ): - payload["optimize_streaming_latency"] = optimize_streaming_latency - # Only attach `text_normalization` when explicitly enabled (default is False). - 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 -# =========================================================================== -def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str: - """ - Generate audio using MiniMax TTS API. - - Supports two endpoints: - - v1/text_to_speech: simple payload, returns raw audio (Content-Type: audio/mpeg) - - v1/t2a_v2: nested voice_setting/audio_setting, returns JSON with hex-encoded audio - - Args: - text: Text to convert (max 10,000 characters). - output_path: Where to save the audio file. - tts_config: TTS config dict. - - Returns: - Path to the saved audio file. - """ - 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 - 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") - sample_rate = mm_config.get("sample_rate", 32000) - bitrate = mm_config.get("bitrate", 128000) - - # MiniMax accounts scope TTS requests by GroupId. When present, the docs - # show it as a ?GroupId= query param on the t2a_v2 URL. Accept it - # from config or from the MINIMAX_GROUP_ID env var; only attach when the - # URL doesn't already carry one. - group_id = ( - str(mm_config.get("group_id") or "").strip() - or (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}", - } - - # Detect endpoint from URL - is_t2a_v2 = "t2a_v2" in base_url - - if is_t2a_v2: - # t2a_v2 endpoint: nested voice_setting/audio_setting structure - payload = { - "model": model, - "text": text, - "voice_setting": { - "voice_id": voice_id, - "speed": speed, - "vol": vol, - "pitch": pitch, - "emotion": emotion, - }, - "audio_setting": { - "sample_rate": sample_rate, - "bitrate": bitrate, - "format": "mp3", - "channel": 1, - }, - } - else: - # text_to_speech endpoint: flat payload - 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: - # t2a_v2 returns JSON with hex-encoded audio - response.raise_for_status() - result = _read_tts_response_json(response, label="MiniMax TTS") - 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}") - - hex_audio = result.get("data", {}).get("audio", "") - if not hex_audio: - raise RuntimeError("MiniMax TTS returned empty audio data") - - audio_bytes = bytes.fromhex(hex_audio) - with open(output_path, "wb") as f: - f.write(audio_bytes) - return output_path - - else: - # text_to_speech returns raw audio directly - 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 - - # Fallback: try parsing as JSON - try: - raw_body = _read_tts_response_bytes(response, label="MiniMax TTS") - result = json.loads(raw_body.decode("utf-8")) if raw_body else {} - 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}") - except (json.JSONDecodeError, UnicodeDecodeError, TypeError): - response.raise_for_status() - raise RuntimeError( - f"MiniMax TTS returned unexpected Content-Type '{content_type}' " - f"({len(raw_body) if 'raw_body' in locals() else 0} bytes)" - ) - - raise RuntimeError("MiniMax TTS returned no audio data") - - -# =========================================================================== -# Provider: Mistral (Voxtral TTS) -# =========================================================================== -def _generate_mistral_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str: - """Generate audio using Mistral Voxtral TTS API. - - The API returns base64-encoded audio; this function decodes it - and writes the raw bytes to *output_path*. - Supports native Opus output for Telegram voice bubbles. - """ - api_key = (_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 - # Class-level base_url parity: every cloud TTS provider section supports - # base_url. The Mistral SDK calls it server_url. - base_url = mi_config.get("base_url") - - if output_path.endswith(".ogg"): - response_format = "opus" - elif output_path.endswith(".wav"): - response_format = "wav" - elif output_path.endswith(".flac"): - response_format = "flac" - else: - response_format = "mp3" - - Mistral = _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=response_format, - ) - 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 _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 with a standard WAV RIFF header. - - Gemini TTS returns audio/L16;codec=pcm;rate=24000 -- raw PCM samples with - no container. We add a minimal WAV header so the file is playable and - ffmpeg can re-encode it to MP3/Opus downstream. - """ - import struct - - 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, # fmt chunk size (PCM) - 1, # audio format (PCM) - 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 _resolve_gemini_persona_prompt_path(gemini_config: Dict[str, Any]) -> Optional[Path]: - """Return the configured persona prompt file path, if any.""" - raw = gemini_config.get("persona_prompt_file") - if not isinstance(raw, str) or not raw.strip(): - return None - - expanded = os.path.expandvars(raw.strip()) - path = Path(expanded).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: - """Return True for Gemini TTS models known to support 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") - enabled = _config_bool(raw, default=DEFAULT_GEMINI_AUDIO_TAGS) - if not enabled: - 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 _clean_gemini_audio_tag_rewrite(content: str) -> str: - clean = (content or "").strip() - fence = re.fullmatch(r"```(?:[A-Za-z0-9_-]+)?\s*(.*?)\s*```", clean, flags=re.DOTALL) - if fence: - clean = fence.group(1).strip() - return clean - - -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 _rewrite_gemini_tts_audio_tags(text: str, persona_prompt: str = "") -> str: - """Use the configured auxiliary model to insert Gemini audio tags.""" - 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 = ( - "PERSONA AND DIRECTOR CONTEXT:\n" - f"{context}\n\n" - "TRANSCRIPT TO TAG:\n" - f"{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 = _clean_gemini_audio_tag_rewrite(_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.""" - 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." - ) - - placeholder_patterns = ( - re.compile(r"\{\{\s*transcript\s*\}\}", flags=re.IGNORECASE), - re.compile(r"\{\s*transcript\s*\}", flags=re.IGNORECASE), - ) - prompt = persona_prompt - for pattern in placeholder_patterns: - if pattern.search(prompt): - prompt = pattern.sub(transcript, prompt) - return f"{preamble}\n\n{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 using Google Gemini TTS. - - Gemini's generateContent endpoint with responseModalities=["AUDIO"] returns - raw 24kHz mono 16-bit PCM (L16) as base64. We wrap it with a WAV RIFF - header to produce a playable file, then ffmpeg-convert to MP3 / Opus if - the caller requested those formats (same pattern as NeuTTS). - - Args: - text: Text to convert (prompt-style; supports inline direction like - "Say cheerfully:" and audio tags like [whispers]). - output_path: Where to save the audio file (.wav, .mp3, or .ogg). - tts_config: TTS config dict. - - Returns: - Path to the saved audio file. - """ - import requests - - api_key = ( - _resolve_provider_key("GEMINI_API_KEY", "gemini") - or _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 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 = _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" - # Include Hermes client context following Gemini's partner - # integration guidance: - # https://ai.google.dev/gemini-api/docs/partner-integration - headers["X-Goog-Api-Client"] = f"hermes-agent/{_hermes_version}" - - endpoint = f"{base_url}/models/{model}:generateContent" - response = requests.post( - endpoint, - params={"key": api_key}, - headers=headers, - json=payload, - timeout=60, - stream=True, - ) - if response.status_code != 200: - # Surface the API error message when present - 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") - - pcm_bytes = base64.b64decode(audio_b64) - wav_bytes = _wrap_pcm_as_wav(pcm_bytes) - - # Fast path: caller wants WAV directly, just write. - if output_path.lower().endswith(".wav"): - with open(output_path, "wb") as f: - f.write(wav_bytes) - return output_path - - # Otherwise write WAV to a temp file and ffmpeg-convert to the target - # format (.mp3 or .ogg). If ffmpeg is missing, fall back to renaming the - # WAV -- this matches the NeuTTS behavior and keeps the tool usable on - # systems without ffmpeg (audio still plays, just with a misleading - # extension). - with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: - tmp.write(wav_bytes) - wav_path = tmp.name - - try: - ffmpeg = shutil.which("ffmpeg") - if ffmpeg: - # For .ogg output, force libopus encoding (Telegram voice bubbles - # require Opus specifically; ffmpeg's default for .ogg is Vorbis). - if output_path.lower().endswith(".ogg"): - cmd = [ - ffmpeg, "-i", wav_path, - "-acodec", "libopus", "-ac", "1", - "-b:a", "48k", "-vbr", "on", - "-application", "voip", "-compression_level", "10", - "-y", "-loglevel", "error", - output_path, - ] - else: - cmd = [ffmpeg, "-i", wav_path, "-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 - - -# =========================================================================== -# NeuTTS (local, on-device TTS via neutts_cli) -# =========================================================================== - -def _check_neutts_available() -> bool: - """Check if the neutts engine is importable (installed locally).""" - try: - import importlib.util - return importlib.util.find_spec("neutts") is not None - except Exception: - return False - - -def _check_kittentts_available() -> bool: - """Check if the kittentts engine is importable (installed locally).""" - try: - import importlib.util - return importlib.util.find_spec("kittentts") is not None - except Exception: - return False - - -def _default_neutts_ref_audio() -> str: - """Return path to the bundled default voice reference audio.""" - return str(Path(__file__).parent / "neutts_samples" / "jo.wav") - - -def _default_neutts_ref_text() -> str: - """Return path to the bundled default voice reference transcript.""" - return str(Path(__file__).parent / "neutts_samples" / "jo.txt") - - -def _generate_neutts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str: - """Generate speech using the local NeuTTS engine. - - Runs synthesis in a subprocess via tools/neutts_synth.py to keep the - ~500MB model in a separate process that exits after synthesis. - Outputs WAV; the caller handles conversion for Telegram if needed. - """ - import sys - - neutts_config = tts_config.get("neutts") or {} - 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") - - # NeuTTS outputs WAV natively — use a .wav path for generation, - # let the caller convert to the final format afterward. - wav_path = output_path - if not output_path.endswith(".wav"): - wav_path = output_path.rsplit(".", 1)[0] + ".wav" - - synth_script = str(Path(__file__).parent / "neutts_synth.py") - cmd = [ - sys.executable, synth_script, - "--text", text, - "--out", wav_path, - "--ref-audio", ref_audio, - "--ref-text", ref_text, - "--model", model, - "--device", device, - ] - - result = subprocess.run(cmd, capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=120, stdin=subprocess.DEVNULL) - if result.returncode != 0: - stderr = result.stderr.strip() - # Filter out the "OK:" line from stderr - error_lines = [l for l in stderr.splitlines() if not l.startswith("OK:")] - raise RuntimeError(f"NeuTTS synthesis failed: {chr(10).join(error_lines) or 'unknown error'}") - - # If the caller wanted .mp3 or .ogg, convert from WAV - if wav_path != output_path: - ffmpeg = shutil.which("ffmpeg") - if ffmpeg: - conv_cmd = [ffmpeg, "-i", wav_path, "-y", "-loglevel", "error", output_path] - subprocess.run(conv_cmd, check=True, timeout=30, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags()) - os.remove(wav_path) - else: - # No ffmpeg — just rename the WAV to the expected path - os.rename(wav_path, output_path) - - return output_path - - -# =========================================================================== -# Provider: Piper (local, neural VITS, 44 languages) -# =========================================================================== - -# Each cached entry below is a whole loaded TTS model (tens of MB). An -# unbounded dict pins one model per distinct voice/model for the process -# lifetime, so a surface that sweeps voices grows memory with no ceiling. Cap -# each cache with a small LRU — most sessions use one or two voices, and a -# reload on a cold miss is cheap next to keeping every model resident. -_TTS_MODEL_CACHE_MAX = 3 - - -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. - - Refreshes recency on a hit (insertion-ordered dict: pop + reinsert), loads - on a miss, then evicts least-recently-used entries beyond the cap. An entry - evicted while a caller still holds its returned reference stays alive for - that caller; only the cache 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 - - # =========================================================================== # Local-engine lifecycle: warm-up / release driven by TTS-output toggles # =========================================================================== # -# Local engines (Piper, KittenTTS) load their model lazily on the first -# synthesis call, so the first spoken reply after a user turns on "read -# replies aloud" / a voice conversation pays the whole load (plus a voice -# download on a fresh install) as dead air before the first word. And once -# loaded, the model stays resident for the process lifetime even after every -# TTS-output toggle is off again. -# -# The toggles ARE the intent signal. Every surface that flips speech output -# on holds a *lease* here (warming the configured engine as a side effect); -# flipping it off releases the lease, and when the last lease is gone the -# local model caches are dropped. Lease-counting instead of a bare -# on/off keeps one surface's "off" from unloading a model another surface -# (TUI /voice tts, desktop read-aloud, desktop conversation) still needs — -# they share this process's caches. -# -# Cloud providers have no resident model; warming them is a no-op beyond -# making sure the lazily-installed SDK is importable (edge-tts), which is -# also first-use latency users see as silence. - -# Provider name → local model cache it populates. The single registry both -# warm_tts_provider() and the release path consult — a new local engine adds -# one row here (at its cache declaration) plus a loader in -# _local_tts_warmers() and gets warm/release for free. -_LOCAL_TTS_MODEL_CACHES: Dict[str, Dict[str, Any]] = {} - +# Local engines load their model lazily on first synthesis, so the first spoken +# reply after a user turns speech output on pays the whole load as dead air, +# and the model then stays resident forever. The toggles ARE the intent +# signal: every surface that flips speech output on holds a *lease* here +# (warming the configured engine); when the last lease is released the local +# model caches are dropped. Lease-counting keeps one surface's "off" from +# unloading a model another surface in this process still needs. Cloud +# providers have nothing resident; warming them only ensures the lazily +# installed SDK is importable. def _local_tts_warmers() -> Dict[str, Callable[[Dict[str, Any]], Any]]: - # Resolved lazily: the loader functions are defined later in this module. + """Provider name → loader populating that engine's cache slot (same key synthesis uses).""" return { "piper": lambda cfg: _load_piper_voice_for_config(cfg)[0], "kittentts": lambda cfg: _load_kittentts_model_for_config(cfg)[0], @@ -2751,18 +957,12 @@ def warm_tts_provider( ) -> Dict[str, Any]: """Pre-load the configured TTS provider so the next synthesis starts hot. - * Local engines (Piper, KittenTTS): resolve the configured voice/model - exactly as synthesis would (including first-use voice download) and - load it into the same LRU cache slot synthesis reads. - * Lazily-installed cloud SDKs (edge-tts, ElevenLabs, Mistral): make sure - the SDK is importable, installing it if lazy installs are allowed. - * User-declared providers: command providers run ``warm_command`` when - set; plugin providers get :meth:`TTSProvider.warm`. - * Everything else: nothing to warm — reported as ``action: "noop"``. - - Never raises; the result dict carries ``warmed`` / ``action`` / ``error`` - so callers on a toggle path can log and move on. Blocking — callers on a - UI thread should run it in the background. + Local engines load their voice/model into the same LRU slot synthesis + reads (including first-use download); lazily-installed cloud SDKs are + made importable; user-declared providers get ``warm_command`` / + :meth:`TTSProvider.warm`; everything else is ``action: "noop"``. + Never raises — the result dict carries ``warmed`` / ``action`` / + ``error``. Blocking; UI threads should run it in the background. """ if tts_config is None: tts_config = _load_tts_config() @@ -2813,12 +1013,10 @@ def warm_tts_provider( def release_tts_provider(provider: Optional[str] = None) -> Dict[str, Any]: """Drop resident local TTS models so their memory is returned. - With ``provider`` given, only that engine's cache is cleared; otherwise - every local engine cache is and the configured user-declared provider - (plugin ``release()`` / command ``release_command``) is signalled. - Cloud providers hold nothing to release. - Returns ``{"released": }``. The next - synthesis simply reloads (or a warm-up does it ahead of time). + With ``provider`` given only that engine's cache is cleared; otherwise + every local cache is, and the configured user-declared provider is + signalled (plugin ``release()`` / command ``release_command``). Returns + ``{"released": }``. """ name = (provider or "").lower().strip() if not name: @@ -2836,11 +1034,10 @@ def release_tts_provider(provider: Optional[str] = None) -> Dict[str, Any]: def acquire_tts_lease(lease: str, tts_config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: - """Register ``lease`` as a live TTS-output consumer and warm the provider. + """Register ``lease`` (e.g. ``"desktop:read-aloud"``) as a live consumer and warm the provider. - ``lease`` names the surface/toggle (e.g. ``"desktop:read-aloud"``, - ``"tui:voice-tts"``). Re-acquiring an existing lease is idempotent (still - re-warms — cheap on a cache hit, and heals a cache cleared elsewhere). + Re-acquiring is idempotent but still re-warms (cheap on a cache hit, and + heals a cache cleared elsewhere). """ with _tts_lease_lock: _tts_leases.add(lease) @@ -2853,9 +1050,8 @@ def acquire_tts_lease(lease: str, tts_config: Optional[Dict[str, Any]] = None) - def release_tts_lease(lease: str) -> Dict[str, Any]: """Drop ``lease``; when it was the last one, unload resident local models. - Releasing a lease that was never acquired is a no-op (still reports the - live holder count) so surfaces can call it unconditionally on their - "off" path. + Releasing a never-acquired lease is a no-op (still reports the holder + count) so surfaces can call it unconditionally on their "off" path. """ with _tts_lease_lock: _tts_leases.discard(lease) @@ -2877,278 +1073,228 @@ def _reset_tts_leases_for_tests() -> None: _tts_leases.clear() -# Module-level cache for Piper voice instances. Voices are keyed on their -# absolute .onnx model path so switching voices doesn't invalidate older -# cached voices. -_piper_voice_cache: Dict[str, Any] = {} -_LOCAL_TTS_MODEL_CACHES["piper"] = _piper_voice_cache +# =========================================================================== +# Built-in provider dispatch +# =========================================================================== +# provider -> (importer-name or None, "package missing" error, log line, +# generator-name). Names are looked up in module globals at call time so +# tests that monkeypatch ``tools.tts_tool._import_x`` / ``_generate_x`` apply. +_BUILTIN_DISPATCH: Dict[str, tuple] = { + "elevenlabs": ( + "_import_elevenlabs", + "ElevenLabs provider selected but 'elevenlabs' package not installed. Run: pip install elevenlabs", + "Generating speech with ElevenLabs...", + "_generate_elevenlabs", + ), + "openai": ( + "_import_openai_client", + "OpenAI provider selected but 'openai' package not installed.", + "Generating speech with OpenAI TTS...", + "_generate_openai_tts", + ), + "deepinfra": ( + "_import_openai_client", + "DeepInfra TTS uses the 'openai' SDK but it isn't installed.", + "Generating speech with DeepInfra TTS...", + "_generate_deepinfra_tts", + ), + "minimax": (None, None, "Generating speech with MiniMax TTS...", "_generate_minimax_tts"), + "xai": (None, None, "Generating speech with xAI TTS...", "_generate_xai_tts"), + "mistral": ( + "_import_mistral_client", + "Mistral provider selected but 'mistralai' package not installed. " + "Run `hermes setup` to install Mistral support.", + "Generating speech with Mistral Voxtral TTS...", + "_generate_mistral_tts", + ), + "gemini": (None, None, "Generating speech with Google Gemini TTS...", "_generate_gemini_tts"), + "kittentts": ( + "_import_kittentts", + "KittenTTS provider selected but 'kittentts' package not installed. " + "Run 'hermes setup tts' and choose KittenTTS, or install manually: " + "pip install https://github.com/KittenML/KittenTTS/releases/download/0.8.1/kittentts-0.8.1-py3-none-any.whl", + "Generating speech with KittenTTS (local, ~25MB)...", + "_generate_kittentts", + ), + "piper": ( + "_import_piper", + "Piper provider selected but 'piper-tts' package not installed. " + "Run 'hermes tools' and select Piper under TTS, or install manually: " + "pip install piper-tts", + "Generating speech with Piper (local)...", + "_generate_piper_tts", + ), +} +_NEUTTS_MISSING_ERROR = ( + "NeuTTS provider selected but neutts is not installed. " + "Run hermes setup and choose NeuTTS, or install espeak-ng and run python -m pip install -U neutts[all]." +) -def _check_piper_available() -> bool: - """Check whether the piper-tts package is importable.""" +def _error_json(message: str) -> str: + return json.dumps({"success": False, "error": message}, ensure_ascii=False) + + +def _run_edge_tts(text: str, file_str: str, tts_config: Dict[str, Any]) -> None: + """Run the async Edge generator from sync code (worker thread; direct run if that fails).""" try: - import importlib.util - return importlib.util.find_spec("piper") is not None - except Exception: - return False + import concurrent.futures + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: + pool.submit( + lambda: asyncio.run(_generate_edge_tts(text, file_str, tts_config)) + ).result(timeout=60) + except RuntimeError: + asyncio.run(_generate_edge_tts(text, file_str, tts_config)) -def _get_piper_voices_dir() -> Path: - """Return the directory where Hermes caches Piper voice models. +def _select_builtin_engine(provider: str) -> tuple: + """Check a built-in provider's SDK. Returns ``(engine, None)`` or ``(provider, error_json)``. - Resolves to ``~/.hermes/cache/piper-voices/`` under the active - HERMES_HOME so voice downloads follow profile boundaries. + Unknown names take the Edge default; when edge-tts is missing, NeuTTS is + the local fallback (``engine`` then differs from ``provider``). """ - 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* (a model name or path) to a concrete .onnx file path. - - Accepts any of: - - Absolute / expanded path to an .onnx file the user already has - - A voice *name* like ``en_US-lessac-medium`` (downloads to - ``download_dir`` on first use via ``python -m piper.download_voices``) - - Raises RuntimeError if the model can't be located or downloaded. - """ - if not voice: - voice = DEFAULT_PIPER_VOICE - - # Case 1: user gave a direct file path. - candidate = Path(voice).expanduser() - if candidate.suffix.lower() == ".onnx" and candidate.exists(): - return str(candidate) - - # Case 2: user gave a voice *name*. See if it's already downloaded. - cached = download_dir / f"{voice}.onnx" - if cached.exists() and (download_dir / f"{voice}.onnx.json").exists(): - return str(cached) - - # Case 3: download the voice. piper ships a download helper module. - import sys as _sys - 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 :func:`warm_tts_provider` so a warm-up populates - exactly the cache slot the next synthesis call will hit — same voice - resolution, same download-on-first-use, same cache key. - - Returns ``(voice, piper_config)``. - """ - PiperVoice = _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.speaker_id — the same - # PiperVoice instance serves all speakers, so it stays out of the cache - # key. Multi-speaker workflows share one model load. - 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: - """Generate speech using the local Piper engine. - - Loads the voice model once per process (cached by absolute path) and - writes a WAV file. Caller is responsible for converting to MP3/Opus - via ffmpeg when a different output format is required. - """ - import wave - - voice, piper_config = _load_piper_voice_for_config(tts_config) - - # Tolerant speaker_id parse: drop bad input (non-int strings, lists, dicts) - # to 0 (Piper's own default). Booleans are rejected outright — True/False - # would silently coerce to 1/0 and hide a config mistake. - _raw_speaker = piper_config.get("speaker_id", 0) - if isinstance(_raw_speaker, bool) or not isinstance(_raw_speaker, int): - speaker_id = 0 - else: - speaker_id = _raw_speaker - - # Optional synthesis knobs — only pass a SynthesisConfig when at least - # one advanced knob is configured, so we don't depend on a newer Piper - # version than the user's installed one unless we need to. - 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", - ) + entry = _BUILTIN_DISPATCH.get(provider) + if entry is not None: + importer_name, missing_error = entry[0], entry[1] + if importer_name is not None and not _importable(globals()[importer_name]): + return provider, _error_json(missing_error) + return provider, None + if provider == "neutts": + if not _check_neutts_available(): + return provider, _error_json(_NEUTTS_MISSING_ERROR) + logger.info("Generating speech with NeuTTS (local)...") + return provider, None + if _importable(_import_edge_tts): + return provider, None # Edge default; the reported provider stays as configured + if _check_neutts_available(): + logger.info("Edge TTS not available, falling back to NeuTTS (local)...") + return "neutts", None + return provider, _error_json( + "No TTS provider available. Install edge-tts (pip install edge-tts) " + "or set up NeuTTS for local synthesis." ) - 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" - ) - # Piper outputs WAV. Caller handles downstream MP3/Opus conversion. - wav_path = output_path - if not output_path.endswith(".wav"): - wav_path = output_path.rsplit(".", 1)[0] + ".wav" - 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) +def _synthesize_builtin(engine: str, text: str, file_str: str, tts_config: Dict[str, Any], instructions: Optional[str]) -> None: + """Run the already-selected built-in *engine* (the caller logs the engine-selection line).""" + entry = _BUILTIN_DISPATCH.get(engine) + if entry is not None: + logger.info(entry[2]) + if engine == "openai": + _generate_openai_tts(text, file_str, tts_config, instructions=instructions) else: - voice.synthesize_wav(text, wav_file) - - # Convert to desired format if caller requested mp3/ogg - if wav_path != output_path: - ffmpeg = shutil.which("ffmpeg") - if ffmpeg: - conv_cmd = [ffmpeg, "-i", wav_path, "-y", "-loglevel", "error", output_path] - subprocess.run(conv_cmd, check=True, timeout=30, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags()) - try: - os.remove(wav_path) - except OSError: - pass - else: - # No ffmpeg — keep WAV and return that path - os.rename(wav_path, output_path) - - return output_path + globals()[entry[3]](text, file_str, tts_config) + elif engine == "neutts": + _generate_neutts(text, file_str, tts_config) + else: + logger.info("Generating speech with Edge TTS...") + _run_edge_tts(text, file_str, tts_config) -# =========================================================================== -# Provider: KittenTTS (local, lightweight) -# =========================================================================== +def _finalize_voice_delivery( + file_str: str, + provider: str, + command_provider_config: Optional[Dict[str, Any]], + want_opus: bool, +) -> tuple: + """Decide voice-bubble eligibility and Opus-convert when needed. -# Module-level cache for KittenTTS model instance -_kittentts_model_cache: Dict[str, Any] = {} -_LOCAL_TTS_MODEL_CACHES["kittentts"] = _kittentts_model_cache - - -def _load_kittentts_model_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, Dict[str, Any]]: - """Load (or fetch from cache) the KittenTTS model ``tts_config`` selects. - - Shared by synthesis and :func:`warm_tts_provider` — same model name, - same cache key. Returns ``(model, kittentts_config)``. + Command and plugin providers are documents by default and opt in via + ``voice_compatible``; native-Opus built-ins are voice-compatible when the + platform wants Opus and they wrote .ogg; MP3/WAV built-ins are converted + with ffmpeg only when the platform needs Opus. Returns ``(path, voice_compatible)``. """ - KittenTTS = _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) + voice_compatible = False + if command_provider_config is not None: + opted_in = _is_command_tts_voice_compatible(command_provider_config) + elif provider not in BUILTIN_TTS_PROVIDERS: + opted_in = _plugin_provider_is_voice_compatible(provider) + elif want_opus and provider in _FFMPEG_OPUS_PROVIDERS and not file_str.endswith(".ogg"): + opus_path = _convert_to_opus(file_str) + if opus_path: + return opus_path, True + return file_str, False + elif provider in _NATIVE_OPUS_PROVIDERS: + return file_str, want_opus and file_str.endswith(".ogg") + else: + return file_str, False - 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: - """Generate speech using KittenTTS local ONNX model. - - KittenTTS is a lightweight TTS engine (25-80MB models) that runs - entirely on CPU without requiring a GPU or API key. - - Args: - text: Text to convert to speech. - output_path: Where to save the audio file. - tts_config: TTS config dict. - - Returns: - Path to the saved audio file. - """ - model, kt_config = _load_kittentts_model_for_config(tts_config) - voice = kt_config.get("voice", DEFAULT_KITTENTTS_VOICE) - speed = kt_config.get("speed", 1.0) - clean_text = kt_config.get("clean_text", True) - - # Generate audio (returns numpy array at 24kHz) - audio = model.generate(text, voice=voice, speed=speed, clean_text=clean_text) - - # Save as WAV - import soundfile as sf - wav_path = output_path - if not output_path.endswith(".wav"): - wav_path = output_path.rsplit(".", 1)[0] + ".wav" - - sf.write(wav_path, audio, 24000) - - # Convert to desired format if needed - if wav_path != output_path: - ffmpeg = shutil.which("ffmpeg") - if ffmpeg: - conv_cmd = [ffmpeg, "-i", wav_path, "-y", "-loglevel", "error", output_path] - subprocess.run(conv_cmd, check=True, timeout=30, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags()) - os.remove(wav_path) - else: - # No ffmpeg — rename the WAV to the expected path - os.rename(wav_path, output_path) - - return output_path + if opted_in: + if not file_str.endswith(".ogg"): + opus_path = _convert_to_opus(file_str) + if opus_path: + file_str = opus_path + voice_compatible = file_str.endswith(".ogg") + return file_str, voice_compatible # =========================================================================== # Main tool function # =========================================================================== + +def _apply_call_overrides(tts_config: Dict[str, Any], speed: Optional[float], provider: Optional[str]): + """Apply per-call ``speed`` (clamped, on a shallow copy) and resolve the provider name.""" + if speed is not None: + clamped = max(0.25, min(4.0, float(speed))) + tts_config = dict(tts_config) # shallow copy to avoid mutating the cache + tts_config["speed"] = clamped + provider = provider.lower().strip() if provider else _get_provider(tts_config) + return tts_config, provider + + +def _session_platform() -> tuple: + """``(platform, wants_opus)`` — platforms delivering voice bubbles only as Ogg/Opus want Opus.""" + from gateway.session_context import get_session_env + platform = get_session_env("HERMES_SESSION_PLATFORM", "").lower() + return platform, platform in OPUS_VOICE_PLATFORMS + + +def _resolve_output_base( + output_path: Optional[str], + provider: str, + command_provider_config: Optional[Dict[str, Any]], + want_opus: bool, +) -> tuple: + """Pick the output file. Returns ``(Path, None)`` or ``(None, error_json)``. + + A caller-supplied path is rejected on ``..`` traversal (bug or + prompt-injection; an absolute path is fine) and on protected credential/ + system locations. Command providers get their configured extension. + Default: ``', flags=re.DOTALL), ' '), + (re.compile(r'```[\s\S]*?```'), ' '), + (re.compile(r'\[([^\]]+)\]\([^)]+\)'), r'\1'), + (re.compile(r'https?://\S+'), ''), + (re.compile(r'\*\*(.+?)\*\*'), r'\1'), + (re.compile(r'\*(.+?)\*'), r'\1'), + (re.compile(r'`(.+?)`'), r'\1'), + (re.compile(r'^#+\s*', flags=re.MULTILINE), ''), + (re.compile(r'^\s*[-*]\s+', flags=re.MULTILINE), ''), + (re.compile(r'---+'), ''), + # Emoji + variation selectors/ZWJ: providers speak them as awkward labels. + (re.compile('[\U0001F000-\U0001FAFF\u2600-\u27BF\uFE0F\u200D\U000E0020-\U000E007F]+'), ' '), + (re.compile(r'\n{3,}'), '\n\n'), ) -# Strip ... reasoning blocks before TTS — models with -# /reasoning show enabled produce think blocks that shouldn't be spoken. -_THINK_BLOCK = re.compile(r'].*?', flags=re.DOTALL) - def _strip_markdown_for_tts(text: str) -> str: """Prepare text for speech via the shared cleaner in tts_text_normalize. One cleaner for every TTS path (tool, gateway auto-TTS, voice-mode - streaming, web dashboard): strips reasoning blocks, the - file-mutation verifier footer, markdown, and emoji; expands units and - symbols; and flattens newlines to sentence breaks so newline-sensitive - providers (Kokoro) speak the whole script. Falls back to the legacy - regex pipeline if the normalizer ever fails. + streaming, web dashboard): strips blocks, the verifier footer, + markdown and emoji; expands units; flattens newlines so newline-sensitive + providers (Kokoro) speak the whole script. Falls back to the legacy regex + pipeline if the normalizer ever fails. """ try: from tools.tts_text_normalize import prepare_spoken_text return prepare_spoken_text(text, max_chars=None) except Exception: pass - text = _THINK_BLOCK.sub(' ', text) - text = _MD_CODE_BLOCK.sub(' ', text) - text = _MD_LINK.sub(r'\1', text) - text = _MD_URL.sub('', text) - text = _MD_BOLD.sub(r'\1', text) - text = _MD_ITALIC.sub(r'\1', text) - text = _MD_INLINE_CODE.sub(r'\1', text) - text = _MD_HEADER.sub('', text) - text = _MD_LIST_ITEM.sub('', text) - text = _MD_HR.sub('', text) - text = _EMOJI.sub(' ', text) - text = _MD_EXCESS_NL.sub('\n\n', text) + for pattern, repl in _LEGACY_TTS_STRIP_STEPS: + text = pattern.sub(repl, text) return text.strip() -class _SyncSentencePipeline: - """Overlap per-sentence synthesis with playback for non-streaming providers. - - The universal sync fallback used to run strictly serially per sentence — - synthesize, play, and only then start synthesizing the next sentence — so - every sentence boundary added a full synthesis-time of dead air. For local - model providers that cost dominates the conversation: a provider at - real-time-factor ~1 spends as long silent between sentences as it does - speaking. Chunked streamers already avoid this; this closes the same gap - for everyone else (edge, piper, plugin providers, …) without touching the - provider contract. - - Shape: one synthesis worker (single-threaded executor, so sentences are - synthesized FIFO and providers never see concurrent calls from this loop — - same effective concurrency as the serial path) feeding one playback worker - through a small bounded queue. While sentence *n* plays, sentence *n+1* is - already synthesizing. The bound keeps lookahead — and the temp files it - implies — small, and gives natural backpressure to the caller. - - ``synthesize``/``play`` are resolved late (module global / import inside - the worker) so tests that monkeypatch ``text_to_speech_tool`` or - ``tools.voice_mode`` keep working unchanged. - """ - - 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) - 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) - if tmp_path: - try: - os.unlink(tmp_path) - except OSError: - pass - 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: - if tmp_path: - try: - os.unlink(tmp_path) - except OSError: - pass - - -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*, buffer them into sentences, and - speak each sentence the moment it's ready — the conversational path. - - Provider-agnostic. A registered streaming provider (ElevenLabs, OpenAI, …) - plays chunked PCM through one sounddevice stream 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 instead - of after the whole reply. - - Protocol: - * The producer puts ``str`` deltas onto *text_queue*. - * A ``None`` sentinel signals end-of-text (flush remaining buffer). - * *stop_event* can be set to abort 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() - sync_pipeline: Optional[_SyncSentencePipeline] = None - - try: - output_stream = None - streamer = None # type: ignore[assignment] - _worker_thread = None - _audio_queue = None # type: ignore[assignment] - _prefetch_threads = [] - tts_config = _load_tts_config() - - # Prefer a chunked streamer for low time-to-first-audio; fall back to - # 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) - - # No chunked streamer: per-sentence sync synthesis, pipelined so the - # next sentence synthesizes while the current one plays (closed in the - # finally block, which flushes anything still queued). - sync_pipeline = _SyncSentencePipeline(stop_event) if streamer is None else None - - stream_max_len = 0 - if streamer is not None: - try: - stream_max_len = _resolve_max_text_length( - provider or _get_provider(tts_config), tts_config - ) - except Exception: - stream_max_len = 0 - # On macOS, skip the sounddevice OutputStream entirely: PortAudio/ - # CoreAudio init triggers a kTCCServiceMediaLibrary permission - # prompt even though output needs no media-library access. Leaving - # output_stream=None routes each sentence through the tempfile - # -> play_audio_file -> afplay path. See PR #62601 / #13291. - if platform.system() == "Darwin": - output_stream = None - else: - try: - sd = _import_sounddevice() - output_stream = sd.OutputStream( - samplerate=streamer.sample_rate, - channels=streamer.channels, - dtype="int16", - ) - output_stream.start() - except (ImportError, OSError) as exc: - logger.debug("sounddevice not available, streamer→tempfile: %s", exc) - output_stream = None - except Exception as exc: - logger.warning("sounddevice OutputStream failed: %s", exc) - output_stream = None - - chunker = SentenceChunker() - long_flush_len = 100 - queue_timeout = 0.5 - _spoken_sentences: list[str] = [] # track spoken sentences to skip duplicates - - # --- Per-sentence prefetch pipeline --- - # Every sentence gets its own streamer.stream() call the moment it's - # complete. A background prefetch thread fires the HTTP request - # immediately, buffering PCM chunks into a per-segment queue. The - # single playback worker drains these queues in FIFO order. This - # means sentence N+1's HTTP request fires WHILE sentence N is still - # playing, so by the time the worker reaches it, audio is already - # arriving — no inter-sentence gap. - _audio_queue: queue.Queue[Optional[queue.Queue[Optional[bytes]]]] = queue.Queue() - _prefetch_threads: list[threading.Thread] = [] - _prefetch_sem = threading.Semaphore(3) - _CHUNK_QUEUE_MAX = 64 - - def _create_output_stream(): - """Create and start a fresh PortAudio OutputStream.""" - sd = _import_sounddevice() - new_stream = sd.OutputStream( - samplerate=streamer.sample_rate, - channels=streamer.channels, - dtype="int16", - ) - new_stream.start() - return new_stream - - def _consume_to_queue( - audio_iter: Iterator[bytes], - chunk_queue: "queue.Queue[Optional[bytes]]", - ) -> None: - """Consume a generator into a thread-safe queue.""" - try: - for chunk in audio_iter: - if 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 - _prefetch_sem.release() # free a prefetch slot - - def _reinit_output_stream(): - """Close the broken PortAudio stream and try to create a fresh one.""" - nonlocal output_stream - if output_stream is not None: - try: - output_stream.stop() - output_stream.close() - except Exception: - pass - try: - new_stream = _create_output_stream() - output_stream = new_stream - logger.info( - "TTS: PortAudio output stream reinitialized after error" - ) - return new_stream - except Exception as exc: - logger.warning( - "TTS: PortAudio stream reinit failed: %s", exc - ) - output_stream = None - return None - - def _playback_worker() -> None: - """Single consumer: play audio segments from the queue in order.""" - assert streamer is not None - if output_stream is not None: - import numpy as _np - - try: - from tools.voice_mode import mark_audio_output_active - except Exception: - def mark_audio_output_active(_active): - return None - - mark_audio_output_active(True) - try: - _max_reinit = 3 - _reinit_count = 0 - _current_stream = output_stream - while True: - chunk_queue = _audio_queue.get() - if chunk_queue is None: - break - if stop_event.is_set(): - continue - if _current_stream is None: - _chunks = [] - while True: - chunk = chunk_queue.get() - if chunk is None: - break - _chunks.append(chunk) - _play_via_tempfile( - iter(_chunks), stop_event, streamer.sample_rate - ) - continue - _pcm_leftover = b"" - while True: - chunk = chunk_queue.get() - if chunk is None: - break - if stop_event.is_set(): - break - _buf = _pcm_leftover + chunk - _aligned_len = len(_buf) - (len(_buf) % 2) - if _aligned_len >= 2: - try: - _current_stream.write( - _np.frombuffer( - _buf[:_aligned_len], dtype=" None: - """Synthesize *text_to_speak* and start prefetching immediately.""" - assert streamer is not None - try: - audio_iter = streamer.stream(text_to_speak) - except Exception as exc: - logger.warning("Streaming TTS synthesis failed: %s", exc) - return - _prefetch_sem.acquire() - chunk_queue: "queue.Queue[Optional[bytes]]" = queue.Queue(maxsize=_CHUNK_QUEUE_MAX) - _audio_queue.put(chunk_queue) - t = threading.Thread( - target=_consume_to_queue, - args=(audio_iter, chunk_queue), - daemon=True, - ) - _prefetch_threads.append(t) - t.start() - - _worker_thread: Optional[threading.Thread] = None - if streamer is not None: - _worker_thread = threading.Thread(target=_playback_worker, daemon=True) - _worker_thread.start() - - def _speak_sentence(sentence: str): - """Display sentence and route to the appropriate audio path.""" - if stop_event.is_set(): - return - cleaned = _strip_markdown_for_tts(sentence).strip() - if not cleaned: - return - # Skip duplicate/near-duplicate sentences (LLM repetition) - cleaned_lower = cleaned.lower().rstrip(".!,") - for prev in _spoken_sentences: - if prev.lower().rstrip(".!,") == cleaned_lower: - return - _spoken_sentences.append(cleaned) - # Display raw sentence on screen before TTS processing - if display_callback is not None: - display_callback(sentence) - # No chunked streamer → per-sentence sync synthesis (universal), - # pipelined: this enqueues and returns, so sentence n+1 is already - # synthesizing while sentence n is still playing. - if sync_pipeline is not None: - sync_pipeline.speak(cleaned) - return - # Truncate very long sentences to the provider's per-request cap. - if stream_max_len and len(cleaned) > stream_max_len: - cleaned = cleaned[:stream_max_len] - # Every sentence gets its own prefetch thread — the HTTP request - # fires the moment the sentence boundary is detected, so audio for - # sentence N+1 is already buffering while sentence N plays. - _enqueue_audio(cleaned) - - def _align_int16_chunks(chunks, stop_evt): - """Yield int16-aligned byte chunks from an iterable.""" - 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, stop_evt, sample_rate=24000): - """Write PCM chunks to a temp WAV file and play it.""" - 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() given a file object flushes but does NOT close it - # (it only closes files it opened itself, by name), so the OS - # handle to tmp stays open. On Windows an open write handle - # blocks the system player from reading the file and blocks the - # os.unlink() below (WinError 32, swallowed → temp .wav files - # pile up). Release the handle before playback and cleanup. - 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 - if tmp_path: - try: - os.unlink(tmp_path) - except OSError: - pass - - while not stop_event.is_set(): - # Read next delta from queue - 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: - # End-of-text sentinel: flush whatever remains - for sentence in chunker.flush(): - _speak_sentence(sentence) - break - - for sentence in chunker.feed(delta): - _speak_sentence(sentence) - - # Drain any remaining items from the queue - while True: - try: - text_queue.get_nowait() - except queue.Empty: - break - - # output_stream is closed in the finally block below - - 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 - # Signal the playback worker that no more audio is coming. This lives - # in finally: so an exception in the text pump still sends the sentinel. - if streamer is not None and _worker_thread is not None: - _audio_queue.put(None) - _worker_thread.join(timeout=300.0) - for t in _prefetch_threads: - t.join(timeout=10.0) - # Always close the audio output stream to avoid locking the device - if output_stream is not None: - try: - output_stream.stop() - output_stream.close() - except Exception: - pass - tts_done_event.set() - - # =========================================================================== # Main -- quick diagnostics # =========================================================================== @@ -4466,18 +1770,11 @@ if __name__ == "__main__": print("🔊 Text-to-Speech Tool Module") print("=" * 50) - def _check(importer, label): - try: - importer() - return True - except ImportError: - return False - print("\nProvider availability:") - print(f" Edge TTS: {'installed' if _check(_import_edge_tts, 'edge') else 'not installed (pip install edge-tts)'}") - print(f" ElevenLabs: {'installed' if _check(_import_elevenlabs, 'el') else 'not installed (pip install elevenlabs)'}") + print(f" Edge TTS: {'installed' if _importable(_import_edge_tts) else 'not installed (pip install edge-tts)'}") + print(f" ElevenLabs: {'installed' if _importable(_import_elevenlabs) else 'not installed (pip install elevenlabs)'}") print(f" API Key: {'set' if _resolve_provider_key('ELEVENLABS_API_KEY', 'elevenlabs') else 'not set'}") - print(f" OpenAI: {'installed' if _check(_import_openai_client, 'oai') else 'not installed'}") + print(f" OpenAI: {'installed' if _importable(_import_openai_client) else 'not installed'}") print( " API Key: " f"{'set' if resolve_openai_audio_api_key() else 'not set (VOICE_TOOLS_OPENAI_KEY or OPENAI_API_KEY)'}" diff --git a/tools/tts_tool_delivery.py b/tools/tts_tool_delivery.py new file mode 100644 index 0000000000..ad4bd48fef --- /dev/null +++ b/tools/tts_tool_delivery.py @@ -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 `` 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 diff --git a/tools/tts_tool_local.py b/tools/tts_tool_local.py new file mode 100644 index 0000000000..0786ca3d09 --- /dev/null +++ b/tools/tts_tool_local.py @@ -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: + """``/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) diff --git a/tools/tts_tool_providers.py b/tools/tts_tool_providers.py new file mode 100644 index 0000000000..611cdca01a --- /dev/null +++ b/tools/tts_tool_providers.py @@ -0,0 +1,896 @@ +"""Cloud TTS backends for ``tools.tts_tool``: Edge, ElevenLabs, xAI, MiniMax, Mistral, Gemini. + +Each ``_generate_(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")\]|)", + 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 `...` — 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=`` 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) diff --git a/tools/tts_tool_speaker.py b/tools/tts_tool_speaker.py new file mode 100644 index 0000000000..a995adf805 --- /dev/null +++ b/tools/tts_tool_speaker.py @@ -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="= 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() diff --git a/tools/voice_client_config.py b/tools/voice_client_config.py index 713aa9f1e5..3be1cf1494 100644 --- a/tools/voice_client_config.py +++ b/tools/voice_client_config.py @@ -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 diff --git a/tools/voice_mode.py b/tools/voice_mode.py index 80652bdb0f..33bff5adf7 100644 --- a/tools/voice_mode.py +++ b/tools/voice_mode.py @@ -9,12 +9,10 @@ Dependencies (optional): or: uv sync --extra voice """ -import difflib import logging import math import os import platform -import re import shlex import shutil import subprocess @@ -28,28 +26,36 @@ from typing import Any, Callable, Dict, List, Optional logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Lazy audio imports -- never imported at module level to avoid crashing -# in headless environments (SSH, Docker, WSL, no PortAudio). -# --------------------------------------------------------------------------- +from tools.voice_mode_transcript import ( # noqa: F401 - re-exported; tests patch tools.voice_mode. + _voice_config, + WHISPER_HALLUCINATIONS, + _HALLUCINATION_REPEAT_RE, + is_whisper_hallucination, + DEFAULT_VOICE_STOP_PHRASES, + _load_voice_stop_phrases, + is_voice_stop_phrase, + DEFAULT_TTS_ECHO_SIMILARITY_THRESHOLD, + MIN_FRAGMENT_LENGTH_FOR_ECHO, + _normalize_for_echo_compare, + is_tts_echo, + voice_stop_hint, +) + +# ── Lazy audio imports ── +# Never imported at module level: crashes headless environments (SSH, Docker, +# WSL, no PortAudio). def _import_audio(): - """Lazy-import sounddevice and numpy. Returns (sd, np). - - Raises ImportError or OSError if the libraries are not available - (e.g. PortAudio missing on headless servers). - """ + """Lazy-import sounddevice and numpy; returns (sd, np). Raises ImportError + or OSError when unavailable (e.g. PortAudio missing on headless servers).""" import sounddevice as sd import numpy as np return sd, np def _import_numpy(): - """Lazy-import numpy only (no sounddevice). Returns the module. - - Used where we need to synthesize/convert audio samples but must NOT - import sounddevice — see _sounddevice_output_allowed. - """ + """Lazy-import numpy only — for synthesizing samples where sounddevice + must NOT be imported (see _sounddevice_output_allowed).""" import numpy as np return np @@ -57,12 +63,10 @@ def _import_numpy(): def _sounddevice_output_allowed() -> bool: """Whether sounddevice may be used for audio OUTPUT. - Returns False on macOS: importing/initializing sounddevice - (PortAudio/CoreAudio) for output triggers a kTCCServiceMediaLibrary - permission prompt, even though playback needs no media-library access. - On macOS all output is routed through ``afplay`` instead. This does NOT - affect audio *input* (recording), which legitimately needs microphone - permission. See PR #62601 / #13291. + False on macOS: initializing PortAudio/CoreAudio for output triggers a + kTCCServiceMediaLibrary prompt even though playback needs no media-library + access, so all output goes through ``afplay`` there. Does NOT affect + *input* (recording), which legitimately needs microphone permission. """ return platform.system() != "Darwin" @@ -77,20 +81,32 @@ def _play_int16_via_tempfile(audio, sample_rate: int) -> None: try: 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) - wf.writeframes(audio.tobytes()) + _write_wav_frames(tmp, audio.tobytes(), sample_rate) play_audio_file(tmp_path) except Exception as e: logger.debug("Tone tempfile playback failed: %s", e) finally: if tmp_path: - try: - os.unlink(tmp_path) - except OSError: - pass + _unlink_quietly(tmp_path) + + +def _write_wav_frames(dest, frames: bytes, sample_rate: int) -> None: + """Write raw 16-bit mono PCM *frames* as a WAV to *dest* (path or file object).""" + with wave.open(dest, "wb") as wf: + wf.setnchannels(CHANNELS) + wf.setsampwidth(SAMPLE_WIDTH) + wf.setframerate(sample_rate) + wf.writeframes(frames) + + +def _unlink_quietly(path: Optional[str]) -> None: + """Best-effort unlink; missing/undeletable files are ignored.""" + if not path: + return + try: + os.unlink(path) + except OSError: + pass def _audio_available() -> bool: @@ -102,12 +118,14 @@ def _audio_available() -> bool: return False -def _default_input_samplerate(sd) -> int: - """Return the preferred capture rate for the default input device. +def _rms(np, data) -> float: + """Root-mean-square level of an int16 block.""" + return float(np.sqrt(np.mean(data.astype(np.float64) ** 2))) - Falls back to the Whisper-friendly 16 kHz constant when the backend does - not expose a numeric default rate. - """ + +def _default_input_samplerate(sd) -> int: + """Preferred capture rate for the default input device; falls back to the + Whisper-friendly 16 kHz constant when the backend exposes no numeric rate.""" try: info = sd.query_devices(None, "input") rate = info.get("default_samplerate") if isinstance(info, dict) else getattr(info, "default_samplerate", None) @@ -124,12 +142,9 @@ from hermes_constants import is_termux as _is_termux_environment def _voice_capture_install_hint() -> str: if _is_termux_environment(): return "pkg install python-numpy portaudio && python -m pip install sounddevice" - # If we're running inside a venv (e.g. the bundled Hermes venv at - # ~/.hermes/profiles//hermes-agent/venv/), `pip install` on the - # user's PATH won't reach the right site-packages — the bare hint sends - # them off to whichever Python their shell resolves first, which on macOS - # is often a system Python under Rosetta with a totally separate wheel - # index. Point them at the actual interpreter pip is sitting next to. + # Inside a venv (e.g. the bundled Hermes venv) a bare `pip install` may hit + # whichever Python the shell resolves first (on macOS often a Rosetta + # system Python) — point at the venv's own pip instead. try: if sys.prefix != getattr(sys, "base_prefix", sys.prefix): pip_in_venv = Path(sys.prefix) / "bin" / "pip" @@ -140,19 +155,42 @@ def _voice_capture_install_hint() -> str: return "pip install sounddevice numpy" +def _portaudio_missing_message() -> str: + """Error text for "sounddevice imports but PortAudio's shared library is + missing" — a pip install can't fix that, so point at the system package.""" + if _is_termux_environment(): + hint = " Termux: pkg install portaudio" + else: + hint = ( + " Linux: sudo apt-get install libportaudio2\n" + " macOS: brew install portaudio" + ) + return f"PortAudio system library not found -- install it first:\n{hint}\nThen retry /voice on." + + +_TERMUX_APP_MISSING_WARNING = ( + "Termux:API Android app is not installed. Install/update the Termux:API app to use termux-microphone-record." +) + + def _termux_microphone_command() -> Optional[str]: if not _is_termux_environment(): return None return shutil.which("termux-microphone-record") +def _run_quiet(cmd: List[str], *, timeout: float, check: bool) -> subprocess.CompletedProcess: + """subprocess.run with captured, utf-8-decoded output and no stdin.""" + return subprocess.run( + cmd, capture_output=True, text=True, encoding='utf-8', errors='replace', + timeout=timeout, check=check, stdin=subprocess.DEVNULL, + ) -# Probes used to detect whether the Termux:API Android app is installed. -# `pm list packages` is the canonical Android lookup but is unreliable on -# some devices: on certain ROMs / Android API levels `pm` itself isn't on -# Termux's PATH while `cmd package` is, on others `pm` returns nothing for -# the calling user even when the app is present. We try both before -# concluding that the app is genuinely missing (issue #31015). + +# Probes for the Termux:API Android app. `pm list packages` is the canonical +# lookup but on some ROMs `pm` isn't on Termux's PATH while `cmd package` is, +# and on others `pm` returns nothing for the calling user even when the app is +# present — so both are tried before concluding the app is missing. _TERMUX_API_PACKAGE_PROBES = ( ("pm", "list", "packages", "com.termux.api"), ("cmd", "package", "list", "packages", "com.termux.api"), @@ -162,25 +200,13 @@ _TERMUX_API_PACKAGE_PROBES = ( def _termux_api_app_installed() -> bool: """Return True iff the Termux:API Android app is installed. - Strategy (issue #31015): - - 1. Try each probe in ``_TERMUX_API_PACKAGE_PROBES`` and look for - ``package:com.termux.api`` in stdout. Any positive hit is - authoritative — return True. - 2. If every probe is *inconclusive* (binary missing, permission - denied, timeout, non-zero exit) we cannot honestly say the app - is missing; fall back to trusting the ``termux-microphone-record`` - binary on PATH. The binary ships with the ``termux-api`` package - and is only useful when the Android app is installed; users who - installed the package deliberately almost always have the app - too. A false negative on this gate blocks ``/voice on`` - outright (the symptom reported in #31015), while a false - positive only surfaces a precise runtime error from the binary - itself — strictly more actionable. - 3. If at least one probe ran cleanly and definitively did not - mention the package, treat the app as missing and return False - — that's the genuine "Termux:API CLI installed without the app" - case the existing warning was written for. + Any probe reporting ``package:com.termux.api`` is authoritative. If EVERY + probe is inconclusive (binary missing, permission denied, timeout, non-zero + exit) we cannot honestly say the app is missing, so trust the + ``termux-microphone-record`` binary on PATH instead: a false negative here + blocks ``/voice on`` outright, while a false positive only surfaces a + precise runtime error from the binary. If at least one probe ran cleanly + without mentioning the package, the app is genuinely missing. """ if not _is_termux_environment(): return False @@ -188,18 +214,8 @@ def _termux_api_app_installed() -> bool: inconclusive = False for cmd in _TERMUX_API_PACKAGE_PROBES: try: - result = subprocess.run( - list(cmd), - capture_output=True, - text=True, encoding='utf-8', errors='replace', - timeout=5, - check=False, - stdin=subprocess.DEVNULL, - ) - except (FileNotFoundError, PermissionError, OSError): - inconclusive = True - continue - except subprocess.TimeoutExpired: + result = _run_quiet(list(cmd), timeout=5, check=False) + except (OSError, subprocess.TimeoutExpired): inconclusive = True continue if result.returncode != 0: @@ -221,47 +237,40 @@ def _termux_voice_capture_available() -> bool: return _termux_microphone_command() is not None and _termux_api_app_installed() -def _pulse_socket_reachable() -> bool: - """Return True if a PulseAudio/PipeWire socket is reachable on disk. - - Covers the common case where a sound server runs locally (e.g. on a - remote SSH host) without ``PULSE_SERVER``/``PIPEWIRE_REMOTE`` being set -- - the client just connects to the default socket under the runtime dir. - We look at ``PULSE_SERVER`` unix paths, ``PULSE_RUNTIME_PATH``, and - ``XDG_RUNTIME_DIR`` for a ``pulse/native`` or ``pipewire-0`` socket - (issue #35622). - """ - import socket - import stat - +def _pulse_socket_candidates() -> List[str]: + """Socket paths a PulseAudio/PipeWire client would try by default.""" candidates: List[str] = [] - - pulse_server = os.environ.get('PULSE_SERVER', '') # PULSE_SERVER may be "unix:/path", "unix:/path;..." or a bare path. - for part in pulse_server.split(';'): + for part in os.environ.get('PULSE_SERVER', '').split(';'): part = part.strip() if part.startswith('unix:'): candidates.append(part[len('unix:'):]) - pulse_runtime = os.environ.get('PULSE_RUNTIME_PATH') if pulse_runtime: candidates.append(os.path.join(pulse_runtime, 'native')) - xdg_runtime = os.environ.get('XDG_RUNTIME_DIR') if xdg_runtime: candidates.append(os.path.join(xdg_runtime, 'pulse', 'native')) candidates.append(os.path.join(xdg_runtime, 'pipewire-0')) + return [c for c in candidates if c] - for path in candidates: - if not path: - continue + +def _pulse_socket_reachable() -> bool: + """True if a PulseAudio/PipeWire socket is reachable on disk. + + Covers a sound server running locally (e.g. on a remote SSH host) without + ``PULSE_SERVER``/``PIPEWIRE_REMOTE`` set. A socket file must also accept a + connection — a stale socket left by a dead server does not count. + """ + import socket + import stat + + for path in _pulse_socket_candidates(): try: if not stat.S_ISSOCK(os.stat(path).st_mode): continue except OSError: continue - # Confirm the socket actually accepts a connection -- a stale socket - # file left by a dead server should not count as reachable. sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) try: sock.settimeout(0.5) @@ -274,27 +283,72 @@ def _pulse_socket_reachable() -> bool: return False +def _probe_audio_libraries( + warnings: List[str], notices: List[str], *, + has_forwarded_audio: bool, termux_mic_cmd: Optional[str], termux_app_installed: bool, +) -> None: + """Import sounddevice and query devices; append the outcome to warnings/notices. + + Host audio forwarding or Termux:API capture downgrade "no devices" / + "query failed" to notices — in WSL with PulseAudio device queries can fail + even though recording/playback works fine. + """ + termux_capture = bool(termux_mic_cmd and termux_app_installed) + try: + sd, _ = _import_audio() + except ImportError: + if termux_capture: + notices.append("Termux:API microphone recording available (sounddevice not required)") + elif termux_mic_cmd and not termux_app_installed: + warnings.append(_TERMUX_APP_MISSING_WARNING) + else: + warnings.append(f"Audio libraries not installed ({_voice_capture_install_hint()})") + return + except OSError: + if termux_capture: + notices.append("Termux:API microphone recording available (PortAudio not required)") + elif termux_mic_cmd and not termux_app_installed: + warnings.append(_TERMUX_APP_MISSING_WARNING) + else: + warnings.append(_portaudio_missing_message()) + return + + try: + if sd.query_devices(): + return + if has_forwarded_audio: + notices.append("No PortAudio devices detected but host audio forwarding is configured -- continuing") + elif termux_capture: + notices.append("No PortAudio devices detected, but Termux:API microphone capture is available") + else: + warnings.append("No audio input/output devices detected") + except Exception: + if has_forwarded_audio: + notices.append("Audio device query failed but host audio forwarding is configured -- continuing") + elif termux_capture: + notices.append("PortAudio device query failed, but Termux:API microphone capture is available") + else: + warnings.append("Audio subsystem error (PortAudio cannot query devices)") + + def detect_audio_environment() -> dict: """Detect if the current environment supports audio I/O. - Returns dict with 'available' (bool), 'warnings' (list of hard-fail - reasons that block voice mode), and 'notices' (list of informational - messages that do NOT block voice mode). + Returns dict with 'available' (bool), 'warnings' (hard-fail reasons that + block voice mode), and 'notices' (informational, do NOT block). SSH, + containers and WSL normally have no audio devices, but a reachable sound + server (PulseAudio/PipeWire socket or forwarding env vars) is honored. """ - warnings = [] # hard-fail: these block voice mode - notices = [] # informational: logged but don't block + warnings: List[str] = [] + notices: List[str] = [] termux_mic_cmd = _termux_microphone_command() termux_app_installed = _termux_api_app_installed() - termux_capture = bool(termux_mic_cmd and termux_app_installed) has_forwarded_audio = bool( os.environ.get('PULSE_SERVER') or os.environ.get('PIPEWIRE_REMOTE') or _pulse_socket_reachable() ) - # SSH detection -- normally no audio devices, but honor a reachable - # sound server (PulseAudio/PipeWire socket or forwarding env vars), which - # works fine over SSH (issue #35622). if any(os.environ.get(v) for v in ('SSH_CLIENT', 'SSH_TTY', 'SSH_CONNECTION')): if has_forwarded_audio: notices.append("Running over SSH with a reachable PulseAudio/PipeWire sound server") @@ -307,10 +361,6 @@ def detect_audio_environment() -> dict: " # or: export PULSE_SERVER=unix:$XDG_RUNTIME_DIR/pulse/native" ) - # Docker/Podman container detection — honor host audio forwarding. - # When the user mounts a PulseAudio/PipeWire socket into the container - # and points PULSE_SERVER / PIPEWIRE_REMOTE at it, audio works fine - # (issue #21203). Only block when no forwarding is configured. from hermes_constants import is_container if is_container(): if has_forwarded_audio: @@ -325,95 +375,35 @@ def detect_audio_environment() -> dict: " PipeWire: -e PIPEWIRE_REMOTE=$XDG_RUNTIME_DIR/pipewire-0" ) - # WSL detection — a reachable sound server makes audio work in WSL. - # Honor any forwarding (PulseAudio bridge OR a forwarded PipeWire/Pulse - # socket), mirroring the SSH and container blocks above. When no - # forwarding is configured, only hard-block if the WSL2 PowerShell TTS - # fallback (Media.SoundPlayer via powershell.exe, see play_audio_file) - # isn't available either. The PowerShell path only covers OUTPUT (TTS - # playback) -- microphone recording genuinely still needs the - # PulseAudio bridge -- so when it's the only thing available we - # downgrade to a notice (keeps the same recording guidance visible, - # but doesn't block /voice on for TTS-only usage). - try: - with open('/proc/version', 'r', encoding="utf-8") as f: - if 'microsoft' in f.read().lower(): - if has_forwarded_audio: - notices.append("Running in WSL with a reachable PulseAudio/PipeWire sound server") - elif _wsl_powershell_tts_available(): - notices.append( - "Running in WSL without a PulseAudio bridge -- TTS playback " - "will use the PowerShell/Media.SoundPlayer fallback. " - "Voice INPUT (recording) still requires a PulseAudio bridge:\n" - " 1. Set PULSE_SERVER=unix:/mnt/wslg/PulseServer\n" - " 2. Create ~/.asoundrc pointing ALSA at PulseAudio\n" - " 3. Verify with: arecord -d 3 /tmp/test.wav && aplay /tmp/test.wav" - ) - else: - warnings.append( - "Running in WSL -- audio requires a forwarded sound server.\n" - " PulseAudio: export PULSE_SERVER=unix:/mnt/wslg/PulseServer\n" - " PipeWire: export PIPEWIRE_REMOTE=$XDG_RUNTIME_DIR/pipewire-0\n" - " Then verify: arecord -d 3 /tmp/test.wav && aplay /tmp/test.wav" - ) - except (FileNotFoundError, PermissionError, OSError): - pass + # WSL: the PowerShell/Media.SoundPlayer fallback only covers OUTPUT, so + # when it is the only thing available we downgrade to a notice (recording + # guidance stays visible, but TTS-only usage isn't blocked). + if _is_wsl2_env(): + if has_forwarded_audio: + notices.append("Running in WSL with a reachable PulseAudio/PipeWire sound server") + elif _wsl_powershell_tts_available(): + notices.append( + "Running in WSL without a PulseAudio bridge -- TTS playback " + "will use the PowerShell/Media.SoundPlayer fallback. " + "Voice INPUT (recording) still requires a PulseAudio bridge:\n" + " 1. Set PULSE_SERVER=unix:/mnt/wslg/PulseServer\n" + " 2. Create ~/.asoundrc pointing ALSA at PulseAudio\n" + " 3. Verify with: arecord -d 3 /tmp/test.wav && aplay /tmp/test.wav" + ) + else: + warnings.append( + "Running in WSL -- audio requires a forwarded sound server.\n" + " PulseAudio: export PULSE_SERVER=unix:/mnt/wslg/PulseServer\n" + " PipeWire: export PIPEWIRE_REMOTE=$XDG_RUNTIME_DIR/pipewire-0\n" + " Then verify: arecord -d 3 /tmp/test.wav && aplay /tmp/test.wav" + ) - # Check audio libraries - try: - sd, _ = _import_audio() - try: - devices = sd.query_devices() - if not devices: - if has_forwarded_audio: - notices.append( - "No PortAudio devices detected but host audio forwarding is configured -- continuing" - ) - elif termux_capture: - notices.append("No PortAudio devices detected, but Termux:API microphone capture is available") - else: - warnings.append("No audio input/output devices detected") - except Exception: - # In WSL with PulseAudio, device queries can fail even though - # recording/playback works fine. Don't block if host audio - # forwarding is configured. - if has_forwarded_audio: - notices.append( - "Audio device query failed but host audio forwarding is configured -- continuing" - ) - elif termux_capture: - notices.append("PortAudio device query failed, but Termux:API microphone capture is available") - else: - warnings.append("Audio subsystem error (PortAudio cannot query devices)") - except ImportError: - if termux_capture: - notices.append("Termux:API microphone recording available (sounddevice not required)") - elif termux_mic_cmd and not termux_app_installed: - warnings.append( - "Termux:API Android app is not installed. Install/update the Termux:API app to use termux-microphone-record." - ) - else: - warnings.append(f"Audio libraries not installed ({_voice_capture_install_hint()})") - except OSError: - if termux_capture: - notices.append("Termux:API microphone recording available (PortAudio not required)") - elif termux_mic_cmd and not termux_app_installed: - warnings.append( - "Termux:API Android app is not installed. Install/update the Termux:API app to use termux-microphone-record." - ) - elif _is_termux_environment(): - warnings.append( - "PortAudio system library not found -- install it first:\n" - " Termux: pkg install portaudio\n" - "Then retry /voice on." - ) - else: - warnings.append( - "PortAudio system library not found -- install it first:\n" - " Linux: sudo apt-get install libportaudio2\n" - " macOS: brew install portaudio\n" - "Then retry /voice on." - ) + _probe_audio_libraries( + warnings, notices, + has_forwarded_audio=has_forwarded_audio, + termux_mic_cmd=termux_mic_cmd, + termux_app_installed=termux_app_installed, + ) return { "available": not warnings, @@ -421,9 +411,7 @@ def detect_audio_environment() -> dict: "notices": notices, } -# --------------------------------------------------------------------------- -# Recording parameters -# --------------------------------------------------------------------------- +# ── Recording parameters ── SAMPLE_RATE = 16000 # Whisper native rate CHANNELS = 1 # Mono DTYPE = "int16" # 16-bit PCM @@ -437,54 +425,42 @@ SILENCE_DURATION_SECONDS = 3.0 # Seconds of continuous silence before auto-stop _TEMP_DIR = os.path.join(tempfile.gettempdir(), "hermes_voice") -# ============================================================================ -# Audio cues (beep tones) -# ============================================================================ +# ── Audio cues (beep tones) ── _DEFAULT_BEEP_VOLUME = 0.3 # Backward-compatible default (matches prior hardcoded value) def _get_beep_volume() -> float: - """Read ``voice.beep_volume`` from config.yaml; clamps to 0.0-1.0. - - Defaults to 0.3 when the key is missing, invalid, or when the config - system can't be imported (e.g. broken ~/.hermes/config.yaml during a - partial install). Failures fall back silently so the audio cue never - breaks the voice loop on a degenerate config. - """ - try: - from hermes_cli.config import load_config - voice_cfg = load_config().get("voice", {}) - if not isinstance(voice_cfg, dict): - return _DEFAULT_BEEP_VOLUME - raw = voice_cfg.get("beep_volume", _DEFAULT_BEEP_VOLUME) - except Exception: - return _DEFAULT_BEEP_VOLUME + """``voice.beep_volume`` clamped to 0.0-1.0; 0.3 when missing/invalid so the + audio cue never breaks the voice loop on a degenerate config.""" + raw = _voice_config().get("beep_volume", _DEFAULT_BEEP_VOLUME) try: volume = float(raw) except (TypeError, ValueError): return _DEFAULT_BEEP_VOLUME - if isinstance(raw, bool) or volume < 0.0 or volume > 1.0 or _is_nan(volume): + if isinstance(raw, bool) or volume < 0.0 or volume > 1.0 or math.isnan(volume): return _DEFAULT_BEEP_VOLUME return volume -def _is_nan(value: float) -> bool: - try: - return math.isnan(value) - except Exception: - return False +def _sd_play_blocking(sd, audio, sample_rate: int, *, timeout: float, blocksize: int = 0) -> None: + """``sd.play`` then poll until the stream goes idle or *timeout* passes. + + ``sd.wait()`` calls ``Event.wait()`` without a timeout and hangs forever if + the audio device stalls, so poll with a ceiling and force-stop instead. + """ + sd.play(audio, samplerate=sample_rate, blocksize=blocksize) + deadline = time.monotonic() + timeout + while sd.get_stream() and sd.get_stream().active and time.monotonic() < deadline: + time.sleep(0.01) + sd.stop() def play_beep(frequency: int = 880, duration: float = 0.12, count: int = 1) -> None: - """Play a short beep tone using numpy + sounddevice. + """Play *count* short beeps of *frequency* Hz (default 880 = A5), *duration* s each. - Args: - frequency: Tone frequency in Hz (default 880 = A5). - duration: Duration of each beep in seconds. - count: Number of beeps to play (with short gap between). + The tone is synthesized with numpy only, so the macOS TCC prompt is not + triggered on the synthesis step; on macOS output goes through afplay. """ - # Synthesize the tone with numpy only (no sounddevice import yet, so the - # macOS TCC prompt is not triggered on the synthesis step). try: np = _import_numpy() except ImportError: @@ -498,8 +474,8 @@ def play_beep(frequency: int = 880, duration: float = 0.12, count: int = 1) -> N parts = [] for i in range(count): t = np.linspace(0, duration, samples_per_beep, endpoint=False) - # Apply fade in/out to avoid click artifacts tone = np.sin(2 * np.pi * frequency * t) + # Fade in/out to avoid click artifacts. fade_len = min(int(SAMPLE_RATE * 0.01), samples_per_beep // 4) tone[:fade_len] *= np.linspace(0, 1, fade_len) tone[-fade_len:] *= np.linspace(1, 0, fade_len) @@ -509,7 +485,6 @@ def play_beep(frequency: int = 880, duration: float = 0.12, count: int = 1) -> N audio = np.concatenate(parts) - # On macOS, route the tone through afplay instead of sounddevice. if not _sounddevice_output_allowed(): _play_int16_via_tempfile(audio, SAMPLE_RATE) return @@ -518,28 +493,18 @@ def play_beep(frequency: int = 880, duration: float = 0.12, count: int = 1) -> N sd, _ = _import_audio() except (ImportError, OSError): return - sd.play(audio, samplerate=SAMPLE_RATE) - # sd.wait() calls Event.wait() without timeout — hangs forever if the - # audio device stalls. Poll with a 2s ceiling and force-stop. - deadline = time.monotonic() + 2.0 - while sd.get_stream() and sd.get_stream().active and time.monotonic() < deadline: - time.sleep(0.01) - sd.stop() + _sd_play_blocking(sd, audio, SAMPLE_RATE, timeout=2.0) except Exception as e: logger.debug("Beep playback failed: %s", e) -# ============================================================================ -# Thinking sound — calm ambient "blub blub" while the agent works -# ============================================================================ -# During a voice conversation the agent can think / run tools for minutes with -# zero audio, which reads as "it died". A quiet, repeating pair of soft water- -# bubble blips fills that gap. Fully synthesized with numpy (no binary asset), -# volume-scaled by voice.beep_volume, gated by voice.thinking_sound (default -# on), and macOS-TCC-safe: sounddevice OUTPUT is gated there -# (_sounddevice_output_allowed), and spawning afplay every second would churn -# subprocesses, so on macOS the thinking sound is skipped silently. - +# ── Thinking sound — calm ambient "blub blub" while the agent works ── +# The agent can think / run tools for minutes with zero audio, which reads as +# "it died"; a quiet repeating pair of soft water-bubble blips fills the gap. +# Synthesized with numpy, scaled by voice.beep_volume, gated by +# voice.thinking_sound (default on). macOS: sounddevice OUTPUT is TCC-gated +# and spawning afplay every second would churn subprocesses, so it is skipped. +# # The host's *should_play* callback decides when blips are allowed; the # module-level output ref-count below tracks when real audio (TTS sentences, # file playback) is actually flowing so hosts have an accurate signal. @@ -551,11 +516,10 @@ _audio_output_lock = threading.Lock() def mark_audio_output_active(active: bool) -> None: """Reference-count real audio output (TTS/file playback). - Playback paths bracket their work with ``mark_audio_output_active(True)`` - / ``(False)`` so ``is_audio_output_active()`` reflects whether speech - audio is leaving the speakers RIGHT NOW — unlike the per-turn TTS-done - events, which stay 'busy' for a whole turn even while the pipeline is - silently waiting for text. + Playback paths bracket their work with ``(True)`` / ``(False)`` so + ``is_audio_output_active()`` reflects whether speech audio is leaving the + speakers RIGHT NOW — unlike per-turn TTS-done events, which stay 'busy' + for a whole turn even while the pipeline silently waits for text. """ global _audio_output_active_count with _audio_output_lock: @@ -577,17 +541,10 @@ _thinking_stop: Optional[threading.Event] = None def thinking_sound_enabled() -> bool: """Config gate: ``voice.thinking_sound`` (default True).""" try: - from hermes_cli.config import load_config from utils import is_truthy_value - - voice_cfg = load_config().get("voice", {}) - if isinstance(voice_cfg, dict): - return is_truthy_value( - voice_cfg.get("thinking_sound", True), default=True - ) + return is_truthy_value(_voice_config().get("thinking_sound", True), default=True) except Exception: - pass - return True + return True def _synth_thinking_blip(np, frequency: float) -> "Any": @@ -678,9 +635,14 @@ def stop_thinking_sound() -> None: stop.set() -# ============================================================================ -# Termux Audio Recorder -# ============================================================================ +# ── Termux Audio Recorder ── +def _new_recording_path(ext: str) -> str: + """Timestamped ``recording_*.`` path under _TEMP_DIR (created on demand).""" + os.makedirs(_TEMP_DIR, exist_ok=True) + timestamp = time.strftime("%Y%m%d_%H%M%S") + return os.path.join(_TEMP_DIR, f"recording_{timestamp}.{ext}") + + class TermuxAudioRecorder: """Recorder backend that uses Termux:API microphone capture commands.""" @@ -725,9 +687,7 @@ class TermuxAudioRecorder: with self._lock: if self._recording: return - os.makedirs(_TEMP_DIR, exist_ok=True) - timestamp = time.strftime("%Y%m%d_%H%M%S") - self._recording_path = os.path.join(_TEMP_DIR, f"recording_{timestamp}.aac") + self._recording_path = _new_recording_path("aac") command = [ mic_cmd, @@ -738,7 +698,7 @@ class TermuxAudioRecorder: "-c", str(CHANNELS), ] try: - subprocess.run(command, capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=15, check=True, stdin=subprocess.DEVNULL) + _run_quiet(command, timeout=15, check=True) except subprocess.CalledProcessError as e: details = (e.stderr or e.stdout or str(e)).strip() raise RuntimeError(f"Termux microphone start failed: {details}") from e @@ -755,60 +715,47 @@ class TermuxAudioRecorder: mic_cmd = _termux_microphone_command() if not mic_cmd: return - subprocess.run([mic_cmd, "-q"], capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=15, check=False, stdin=subprocess.DEVNULL) + _run_quiet([mic_cmd, "-q"], timeout=15, check=False) + + def _reset_state(self) -> tuple: + """Clear recording state under the lock; return (was_recording, path, started_at).""" + with self._lock: + was_recording, path, started_at = self._recording, self._recording_path, self._start_time + self._recording = False + self._recording_path = None + self._current_rms = 0 + return was_recording, path, started_at def stop(self) -> Optional[str]: - with self._lock: - if not self._recording: - return None - self._recording = False - path = self._recording_path - self._recording_path = None - started_at = self._start_time - self._current_rms = 0 + was_recording, path, started_at = self._reset_state() + if not was_recording: + return None self._stop_termux_recording() if not path or not os.path.isfile(path): return None - if time.monotonic() - started_at < 0.3: - try: - os.unlink(path) - except OSError: - pass - return None - if os.path.getsize(path) <= 0: - try: - os.unlink(path) - except OSError: - pass + # Discard sub-0.3s taps and empty files. + if time.monotonic() - started_at < 0.3 or os.path.getsize(path) <= 0: + _unlink_quietly(path) return None logger.info("Termux voice recording stopped: %s", path) return path def cancel(self) -> None: - with self._lock: - path = self._recording_path - self._recording = False - self._recording_path = None - self._current_rms = 0 + _, path, _ = self._reset_state() try: self._stop_termux_recording() except Exception: pass if path and os.path.isfile(path): - try: - os.unlink(path) - except OSError: - pass + _unlink_quietly(path) logger.info("Termux voice recording cancelled") def shutdown(self) -> None: self.cancel() -# ============================================================================ -# AudioRecorder -# ============================================================================ +# ── AudioRecorder ── class AudioRecorder: """Thread-safe audio recorder using sounddevice.InputStream. @@ -834,34 +781,34 @@ class AudioRecorder: self._recording = False self._start_time: float = 0.0 self._sample_rate: int = SAMPLE_RATE - # Silence detection state - self._has_spoken = False - self._speech_start: float = 0.0 # When speech attempt began - self._dip_start: float = 0.0 # When current below-threshold dip began - self._min_speech_duration: float = 0.3 # Seconds of speech needed to confirm - self._max_dip_tolerance: float = 0.3 # Max dip duration before resetting speech - self._silence_start: float = 0.0 - self._resume_start: float = 0.0 # Tracks sustained speech after silence starts - self._resume_dip_start: float = 0.0 # Dip tolerance tracker for resume detection self._on_silence_stop = None self._silence_threshold: int = SILENCE_RMS_THRESHOLD self._silence_duration: float = SILENCE_DURATION_SECONDS + self._min_speech_duration: float = 0.3 # Seconds of speech needed to confirm + self._max_dip_tolerance: float = 0.3 # Max dip duration before resetting speech self._max_wait: float = 15.0 # Max seconds to wait for speech before auto-stop # Hard cap on total recording length, wired from voice.max_recording_seconds - # by the CLI before each recording. 0 (or unset) = no cap (previous behaviour). + # by the CLI before each recording. 0 (or unset) = no cap. self._max_recording_seconds: float = 0.0 - # Peak RMS seen during recording (for speech presence check in stop()) - self._peak_rms: int = 0 - # Live audio level (read by UI for visual feedback) - self._current_rms: int = 0 + self._peak_rms: int = 0 # for the speech-presence check in stop() + self._current_rms: int = 0 # live level, read by the UI + self._reset_detection_state() + + def _reset_detection_state(self) -> None: + """Reset per-recording silence-detection trackers.""" + self._has_spoken = False + self._speech_start: float = 0.0 # When speech attempt began + self._dip_start: float = 0.0 # When current below-threshold dip began + self._silence_start: float = 0.0 + self._resume_start: float = 0.0 # Tracks sustained speech after silence starts + self._resume_dip_start: float = 0.0 # Dip tolerance tracker for resume detection def _max_duration_reached(self, elapsed: float) -> bool: """Whether the configured hard recording-length cap has elapsed. ``voice.max_recording_seconds`` is applied by the CLI before each recording (see ``HermesCLI._voice_start_recording``). A value <= 0 - (or unset) disables the cap, preserving the previous unbounded - behaviour. + (or unset) disables the cap. """ cap = self._max_recording_seconds return bool(cap and cap > 0 and elapsed >= cap) @@ -884,15 +831,111 @@ class AudioRecorder: """Whether audio recording is currently active.""" return self._recording + # -- silence detection --------------------------------------------------- + + def _track_speech(self, rms: int, now: float) -> None: + """Advance the speech/dip trackers for one audio block. + + Speech is confirmed after ``_min_speech_duration`` above threshold, + tolerating dips shorter than ``_max_dip_tolerance`` (micro-pauses + between syllables). After confirmation only SUSTAINED resumed speech + resets the silence timer — brief ambient spikes must not. + """ + if rms > self._silence_threshold: + self._dip_start = 0.0 + if self._speech_start == 0.0: + self._speech_start = now + elif not self._has_spoken and now - self._speech_start >= self._min_speech_duration: + self._has_spoken = True + logger.debug("Speech confirmed (%.2fs above threshold)", + now - self._speech_start) + if not self._has_spoken: + self._silence_start = 0.0 + else: + # Resumed speech mirrors initial detection: track, tolerate + # short dips, confirm after _min_speech_duration. + self._resume_dip_start = 0.0 + if self._resume_start == 0.0: + self._resume_start = now + elif now - self._resume_start >= self._min_speech_duration: + self._silence_start = 0.0 + self._resume_start = 0.0 + elif self._has_spoken: + # Below threshold after confirmed speech: dip-tolerant resume reset. + if self._resume_start > 0: + if self._resume_dip_start == 0.0: + self._resume_dip_start = now + elif now - self._resume_dip_start >= self._max_dip_tolerance: + self._resume_start = 0.0 + self._resume_dip_start = 0.0 + elif self._speech_start > 0: + # Speech attempt dipped; a long enough dip is genuine silence. + if self._dip_start == 0.0: + self._dip_start = now + elif now - self._dip_start >= self._max_dip_tolerance: + logger.debug("Speech attempt reset (dip lasted %.2fs)", + now - self._dip_start) + self._speech_start = 0.0 + self._dip_start = 0.0 + + def _should_auto_stop(self, rms: int, now: float) -> bool: + """Auto-stop when: the user spoke then stayed silent for + ``_silence_duration``; no speech at all for ``_max_wait``; or the + hard ``voice.max_recording_seconds`` cap elapsed (independent of + speech, so a continuous speaker still stops).""" + elapsed = now - self._start_time + if self._has_spoken and rms <= self._silence_threshold: + if self._silence_start == 0.0: + self._silence_start = now + elif now - self._silence_start >= self._silence_duration: + logger.info("Silence detected (%.1fs), auto-stopping", self._silence_duration) + return True + elif not self._has_spoken and elapsed >= self._max_wait: + logger.info("No speech within %.0fs, auto-stopping", self._max_wait) + return True + if self._max_duration_reached(elapsed): + logger.info("Max recording length reached (%.0fs), auto-stopping", + self._max_recording_seconds) + return True + return False + + def _fire_silence_callback(self) -> None: + """Invoke ``on_silence_stop`` once, in a daemon thread.""" + with self._lock: + cb = self._on_silence_stop + self._on_silence_stop = None # fire only once + if not cb: + return + + def _safe_cb(): + try: + cb() + except Exception as e: + logger.error("Silence callback failed: %s", e, exc_info=True) + threading.Thread(target=_safe_cb, daemon=True).start() + + def _on_audio_block(self, np, indata) -> None: + """Per-block work for the InputStream callback while recording.""" + self._frames.append(indata.copy()) + rms = int(_rms(np, indata)) + self._current_rms = rms + self._peak_rms = max(self._peak_rms, rms) + if self._on_silence_stop is None: + return + now = time.monotonic() + self._track_speech(rms, now) + if self._should_auto_stop(rms, now): + self._fire_silence_callback() + # -- public methods ------------------------------------------------------ def _ensure_stream(self) -> None: """Create the audio InputStream once and keep it alive. - The stream stays open for the lifetime of the recorder. Between - recordings the callback simply discards audio chunks (``_recording`` - is ``False``). This avoids the CoreAudio bug where closing and - re-opening an ``InputStream`` hangs indefinitely on macOS. + The stream stays open for the lifetime of the recorder; between + recordings the callback simply discards chunks (``_recording`` is + False). This avoids the CoreAudio bug where closing and re-opening an + ``InputStream`` hangs indefinitely on macOS. """ if self._stream is not None: return # already alive @@ -902,105 +945,8 @@ class AudioRecorder: def _callback(indata, frames, time_info, status): # noqa: ARG001 if status: logger.debug("sounddevice status: %s", status) - # When not recording the stream is idle — discard audio. - if not self._recording: - return - self._frames.append(indata.copy()) - - # Compute RMS for level display and silence detection - rms = int(np.sqrt(np.mean(indata.astype(np.float64) ** 2))) - self._current_rms = rms - self._peak_rms = max(self._peak_rms, rms) - - # Silence detection - if self._on_silence_stop is not None: - now = time.monotonic() - elapsed = now - self._start_time - - if rms > self._silence_threshold: - # Audio is above threshold -- this is speech (or noise). - self._dip_start = 0.0 # Reset dip tracker - if self._speech_start == 0.0: - self._speech_start = now - elif not self._has_spoken and now - self._speech_start >= self._min_speech_duration: - self._has_spoken = True - logger.debug("Speech confirmed (%.2fs above threshold)", - now - self._speech_start) - # After speech is confirmed, only reset silence timer if - # speech is sustained (>0.3s above threshold). Brief - # spikes from ambient noise should NOT reset the timer. - if not self._has_spoken: - self._silence_start = 0.0 - else: - # Track resumed speech with dip tolerance. - # Brief dips below threshold are normal during speech, - # so we mirror the initial speech detection pattern: - # start tracking, tolerate short dips, confirm after 0.3s. - self._resume_dip_start = 0.0 # Above threshold — no dip - if self._resume_start == 0.0: - self._resume_start = now - elif now - self._resume_start >= self._min_speech_duration: - self._silence_start = 0.0 - self._resume_start = 0.0 - elif self._has_spoken: - # Below threshold after speech confirmed. - # Use dip tolerance before resetting resume tracker — - # natural speech has brief dips below threshold. - if self._resume_start > 0: - if self._resume_dip_start == 0.0: - self._resume_dip_start = now - elif now - self._resume_dip_start >= self._max_dip_tolerance: - # Sustained dip — user actually stopped speaking - self._resume_start = 0.0 - self._resume_dip_start = 0.0 - elif self._speech_start > 0: - # We were in a speech attempt but RMS dipped. - # Tolerate brief dips (micro-pauses between syllables). - if self._dip_start == 0.0: - self._dip_start = now - elif now - self._dip_start >= self._max_dip_tolerance: - # Dip lasted too long -- genuine silence, reset - logger.debug("Speech attempt reset (dip lasted %.2fs)", - now - self._dip_start) - self._speech_start = 0.0 - self._dip_start = 0.0 - - # Fire silence callback when: - # 1. User spoke then went silent for silence_duration, OR - # 2. No speech detected at all for max_wait seconds - should_fire = False - if self._has_spoken and rms <= self._silence_threshold: - # User was speaking and now is silent - if self._silence_start == 0.0: - self._silence_start = now - elif now - self._silence_start >= self._silence_duration: - logger.info("Silence detected (%.1fs), auto-stopping", - self._silence_duration) - should_fire = True - elif not self._has_spoken and elapsed >= self._max_wait: - logger.info("No speech within %.0fs, auto-stopping", - self._max_wait) - should_fire = True - - # 3. Hard cap on total recording length (voice.max_recording_seconds). - # Independent of speech/silence so a continuous speaker past the - # configured limit still auto-stops instead of recording forever. - if not should_fire and self._max_duration_reached(elapsed): - logger.info("Max recording length reached (%.0fs), auto-stopping", - self._max_recording_seconds) - should_fire = True - - if should_fire: - with self._lock: - cb = self._on_silence_stop - self._on_silence_stop = None # fire only once - if cb: - def _safe_cb(): - try: - cb() - except Exception as e: - logger.error("Silence callback failed: %s", e, exc_info=True) - threading.Thread(target=_safe_cb, daemon=True).start() + if self._recording: + self._on_audio_block(np, indata) # Create stream — may block on CoreAudio (first call only). stream = None @@ -1027,36 +973,16 @@ class AudioRecorder: def start(self, on_silence_stop=None) -> None: """Start capturing audio from the default input device. - The underlying InputStream is created once and kept alive across - recordings. Subsequent calls simply reset detection state and - toggle frame collection via ``_recording``. - - Args: - on_silence_stop: Optional callback invoked (in a daemon thread) when - silence is detected after speech. The callback receives no arguments. - Use this to auto-stop recording and trigger transcription. - - Raises ``RuntimeError`` if sounddevice/numpy are not installed - or if a recording is already in progress. + The InputStream is created once and kept alive across recordings; + later calls reset detection state and toggle frame collection. + *on_silence_stop* is invoked (in a daemon thread, no arguments) when + silence follows speech — use it to auto-stop and transcribe. + Raises ``RuntimeError`` if sounddevice/numpy are not installed. """ try: sd, _ = _import_audio() except OSError as e: - # sounddevice imports but PortAudio's shared library is missing — - # a pip install can't fix that; point at the system package - # instead of misreporting missing Python packages (#18432). - if _is_termux_environment(): - portaudio_hint = " Termux: pkg install portaudio" - else: - portaudio_hint = ( - " Linux: sudo apt-get install libportaudio2\n" - " macOS: brew install portaudio" - ) - raise RuntimeError( - "PortAudio system library not found -- install it first:\n" - f"{portaudio_hint}\n" - "Then retry /voice on." - ) from e + raise RuntimeError(_portaudio_missing_message()) from e except ImportError as e: raise RuntimeError( "Voice mode requires sounddevice and numpy.\n" @@ -1069,16 +995,10 @@ class AudioRecorder: self._frames = [] self._start_time = time.monotonic() - self._has_spoken = False - self._speech_start = 0.0 - self._dip_start = 0.0 - self._silence_start = 0.0 - self._resume_start = 0.0 - self._resume_dip_start = 0.0 + self._reset_detection_state() self._peak_rms = 0 self._current_rms = 0 self._on_silence_stop = on_silence_stop - # Ensure the persistent stream is alive (no-op after first call). self._sample_rate = _default_input_samplerate(sd) self._ensure_stream() @@ -1113,11 +1033,8 @@ class AudioRecorder: def stop(self) -> Optional[str]: """Stop recording and write captured audio to a WAV file. - The underlying stream is kept alive for reuse — only frame - collection is stopped. - - Returns: - Path to the WAV file, or ``None`` if no audio was captured. + The stream stays alive for reuse — only frame collection stops. + Returns the WAV path, or ``None`` if no usable audio was captured. """ with self._lock: if not self._recording: @@ -1125,12 +1042,10 @@ class AudioRecorder: self._recording = False self._current_rms = 0 - # Stream stays alive — no close needed. if not self._frames: return None - # Concatenate frames and write WAV _, np = _import_audio() audio_data = np.concatenate(self._frames, axis=0) self._frames = [] @@ -1153,24 +1068,21 @@ class AudioRecorder: return self._write_wav(audio_data, sample_rate=self._sample_rate) - def cancel(self) -> None: - """Stop recording and discard all captured audio. - - The underlying stream is kept alive for reuse. - """ + def _discard(self) -> None: with self._lock: self._recording = False self._frames = [] self._on_silence_stop = None self._current_rms = 0 + + def cancel(self) -> None: + """Stop recording and discard all captured audio (stream stays alive).""" + self._discard() logger.info("Voice recording cancelled") def shutdown(self) -> None: """Release the audio stream. Call when voice mode is disabled.""" - with self._lock: - self._recording = False - self._frames = [] - self._on_silence_stop = None + self._discard() # Close stream OUTSIDE the lock to avoid deadlock with audio callback self._close_stream_with_timeout() logger.info("AudioRecorder shut down") @@ -1179,22 +1091,10 @@ class AudioRecorder: @staticmethod def _write_wav(audio_data, *, sample_rate: int = SAMPLE_RATE) -> str: - """Write numpy int16 audio data to a WAV file. - - Returns the file path. - """ - os.makedirs(_TEMP_DIR, exist_ok=True) - timestamp = time.strftime("%Y%m%d_%H%M%S") - wav_path = os.path.join(_TEMP_DIR, f"recording_{timestamp}.wav") - - with wave.open(wav_path, "wb") as wf: - wf.setnchannels(CHANNELS) - wf.setsampwidth(SAMPLE_WIDTH) - wf.setframerate(sample_rate) - wf.writeframes(audio_data.tobytes()) - - file_size = os.path.getsize(wav_path) - logger.info("WAV written: %s (%d bytes)", wav_path, file_size) + """Write numpy int16 audio data to a WAV file; returns the path.""" + wav_path = _new_recording_path("wav") + _write_wav_frames(wav_path, audio_data.tobytes(), sample_rate) + logger.info("WAV written: %s (%d bytes)", wav_path, os.path.getsize(wav_path)) return wav_path @@ -1205,255 +1105,40 @@ def create_audio_recorder() -> AudioRecorder | TermuxAudioRecorder: return AudioRecorder() -# ============================================================================ -# Whisper hallucination filter -# ============================================================================ -# Whisper commonly hallucinates these phrases on silent/near-silent audio. -WHISPER_HALLUCINATIONS = { - "thank you.", - "thank you", - "thanks for watching.", - "thanks for watching", - "subscribe to my channel.", - "subscribe to my channel", - "like and subscribe.", - "like and subscribe", - "please subscribe.", - "please subscribe", - "thank you for watching.", - "thank you for watching", - "bye.", - "bye", - "you", - "the end.", - "the end", - # Non-English hallucinations (common on silence) - "продолжение следует", - "продолжение следует...", - "sous-titres", - "sous-titres réalisés par la communauté d'amara.org", - "sottotitoli creati dalla comunità amara.org", - "untertitel von stephanie geiges", - "amara.org", - "www.mooji.org", - "ご視聴ありがとうございました", -} - -# Regex patterns for repetitive hallucinations (e.g. "Thank you. Thank you. Thank you.") -_HALLUCINATION_REPEAT_RE = re.compile( - r'^(?:thank you|thanks|bye|you|ok|okay|the end|\.|\s|,|!)+$', - flags=re.IGNORECASE, -) - - -def is_whisper_hallucination(transcript: str) -> bool: - """Check if a transcript is a known Whisper hallucination on silence.""" - cleaned = transcript.strip().lower() - if not cleaned: - return True - # Exact match against known phrases - if cleaned.rstrip('.!') in WHISPER_HALLUCINATIONS or cleaned in WHISPER_HALLUCINATIONS: - return True - # Repetitive patterns (e.g. "Thank you. Thank you. Thank you. you") - if _HALLUCINATION_REPEAT_RE.match(cleaned): - return True - return False - - -# ============================================================================ -# Voice-chat stop phrases -# ============================================================================ - -DEFAULT_VOICE_STOP_PHRASES = ("stop",) - - -def _load_voice_stop_phrases() -> tuple: - """Return the configured ``voice.stop_phrases`` list (default: ("stop",)). - - Malformed config (scalar, dict, list of non-strings) falls back to the - default rather than crashing the voice loop. - """ - try: - from hermes_cli.config import load_config - voice_cfg = load_config().get("voice", {}) - if isinstance(voice_cfg, dict): - raw = voice_cfg.get("stop_phrases", DEFAULT_VOICE_STOP_PHRASES) - if isinstance(raw, str): - raw = [raw] - if isinstance(raw, (list, tuple)): - phrases = tuple( - str(p).strip().lower() for p in raw - if isinstance(p, (str, int, float)) and str(p).strip() - ) - return phrases # empty tuple = feature disabled - except Exception: - pass - return DEFAULT_VOICE_STOP_PHRASES - - -def is_voice_stop_phrase(transcript: str, stop_phrases: Optional[tuple] = None) -> bool: - """Return True when *transcript* is EXACTLY a configured stop phrase. - - Ends the voice conversation when the user says "stop" (or another - configured phrase) and nothing else. Deliberately strict: the whole - utterance — after lowercasing and stripping surrounding punctuation — - must equal a phrase, so "stop doing that and try again" still reaches - the agent. Configure via ``voice.stop_phrases`` in config.yaml - (set ``[]`` to disable). - """ - if not transcript: - return False - cleaned = transcript.strip().lower().strip(".,!?;: \t\n\"'") - if not cleaned: - return False - if stop_phrases is None: - stop_phrases = _load_voice_stop_phrases() - return cleaned in stop_phrases - - -# Similarity ratio (difflib.SequenceMatcher, 0..1) above which a -# playback-phase barge transcript is treated as a self-capture of Hermes' -# own just-spoken TTS rather than genuine user speech. See #75780: the -# full-duplex listener has no acoustic echo cancellation, so speaker bleed -# on the mic can trip the barge trigger and get transcribed nearly -# verbatim from the TTS text, creating a TTS -> STT -> TTS feedback loop. -DEFAULT_TTS_ECHO_SIMILARITY_THRESHOLD = 0.6 - -# Minimum normalized-transcript length (in characters) required before the -# fragment sliding-window fallback runs. Below this, any same-length window -# of `spoken_text` that happens to contain the transcript verbatim (e.g. a -# genuine one-word barge-in like "yes" landing inside a longer reply that -# also says "yes") scores a trivial 1.0 ratio and would otherwise be -# misread as a self-capture. A real self-capture fragment spans at least -# the pre-roll buffer plus time-to-silence, so it is normally well above -# this length; a short genuine interjection is not (#75792 review). -MIN_FRAGMENT_LENGTH_FOR_ECHO = 10 - - -def _normalize_for_echo_compare(text: str) -> str: - return re.sub(r"\s+", " ", text).strip().lower() - - -def is_tts_echo( - transcript: str, - spoken_text: str, - threshold: float = DEFAULT_TTS_ECHO_SIMILARITY_THRESHOLD, -) -> bool: - """Return True when *transcript* looks like a self-capture of *spoken_text*. - - Compares a playback-phase barge-in transcript against the TTS text - Hermes just spoke using a character-level similarity ratio, which works - across languages without word-tokenization. A genuine user interjection - is very unlikely to closely match Hermes' own words, so a high ratio is - a strong signal of speaker-bleed self-capture (fail-closed guard for the - playback-phase full-duplex listener, which has no acoustic echo - cancellation; see #75780). - - The playback-phase capture is cut immediately when the barge trigger - fires and only spans the pre-roll buffer plus time-to-silence, so for - any spoken reply longer than a clause the transcript is a short - FRAGMENT of `spoken_text`, not a near-verbatim repeat of the whole - thing. A whole-string ratio dilutes towards 0 as `spoken_text` grows - past the fragment's length, so when the whole-string check misses, we - also slide a window sized to the transcript's character length across - `spoken_text` and compare against each window, catching a short - fragment echoed from within a much longer multi-sentence reply. This - windowing is character-based (not word-split), so it also works for - languages without whitespace between words. Transcripts shorter than - `MIN_FRAGMENT_LENGTH_FOR_ECHO` skip this fallback entirely, since a - short genuine interjection can trivially match an equally short window - of unrelated spoken text. - """ - if not transcript or not spoken_text: - return False - a = _normalize_for_echo_compare(transcript) - b = _normalize_for_echo_compare(spoken_text) - if not a or not b: - return False - if difflib.SequenceMatcher(None, a, b).ratio() >= threshold: - return True - if len(a) < MIN_FRAGMENT_LENGTH_FOR_ECHO or len(a) >= len(b): - return False - for start in range(0, len(b) - len(a) + 1): - window = b[start : start + len(a)] - if difflib.SequenceMatcher(None, a, window).ratio() >= threshold: - return True - return False - - -def voice_stop_hint() -> str: - """One-line 'Say "stop" to end the voice chat.' hint for voice-mode start. - - Sources the phrase from ``voice.stop_phrases`` (first entry) so a custom - phrase renders correctly; returns "" when stop phrases are disabled - (``stop_phrases: []``) so surfaces show no hint at all. Every surface - that announces voice-mode start (CLI /voice on, TUI, desktop) uses this - one owner instead of hardcoding the wording. - """ - phrases = _load_voice_stop_phrases() - if not phrases: - return "" - return f'Say "{phrases[0]}" to end the voice chat.' - - -# ============================================================================ -# STT dispatch -# ============================================================================ +# ── STT dispatch ── def transcribe_recording(wav_path: str, model: Optional[str] = None) -> Dict[str, Any]: - """Transcribe a WAV recording using the existing Whisper pipeline. + """Transcribe a WAV via ``tools.transcription_tools.transcribe_audio()``, + filtering Whisper hallucinations on silent audio. - Delegates to ``tools.transcription_tools.transcribe_audio()``. - Filters out known Whisper hallucinations on silent audio. - - Args: - wav_path: Path to the WAV file. - model: Whisper model name (default: from config or ``whisper-1``). - - Returns: - Dict with ``success``, ``transcript``, and optionally ``error``. + Returns dict with ``success``, ``transcript``, and optionally ``error``. """ from tools.transcription_tools import MAX_FILE_SIZE, transcribe_audio result = transcribe_audio(wav_path, model=model, source="voice_mode") - # Only chunk when the provider itself reports "File too large" — - # local providers (faster-whisper, whisper.cpp, etc.) have no upload - # cap so ``transcribe_audio`` will never return this error for them. + # Only chunk when the provider itself reports "File too large" — local + # providers have no upload cap and never return this error. if not result.get("success") and "File too large" in result.get("error", ""): result = _transcribe_wav_in_chunks(wav_path, model=model, max_file_size=MAX_FILE_SIZE) - # Filter out Whisper hallucinations (common on silent/near-silent audio). # A configured voice-chat stop phrase is checked FIRST and always survives: # phrases like "bye" or "okay" overlap the hallucination blocklist/repeat - # regex, and swallowing them here would make saying "bye" (when configured - # as a stop phrase) silently fail to end the voice chat. + # regex, and swallowing them would make saying "bye" fail to end the chat. if result.get("success"): raw_transcript = result.get("transcript", "") - if is_whisper_hallucination(raw_transcript) and not is_voice_stop_phrase( - raw_transcript - ): + if is_whisper_hallucination(raw_transcript) and not is_voice_stop_phrase(raw_transcript): logger.info("Filtered Whisper hallucination: %r", result["transcript"]) return {"success": True, "transcript": "", "filtered": True} - # Providers that flag no_speech (empty transcript) failed to hear words, - # not to transcribe — treat like silence so the voice loop re-listens - # quietly instead of surfacing "Transcription failed". + # Providers that flag no_speech failed to hear words, not to transcribe — + # treat like silence so the voice loop re-listens quietly instead of + # surfacing "Transcription failed". if result.get("no_speech"): return {"success": True, "transcript": "", "no_speech": True} return result -def _should_chunk_for_transcription(file_path: str, max_file_size: int) -> bool: - """Return whether a CLI WAV recording needs to be split before STT.""" - if not file_path.lower().endswith(".wav"): - return False - try: - return os.path.getsize(file_path) > max_file_size - except OSError: - return False - - def _transcribe_wav_in_chunks( wav_path: str, *, @@ -1497,11 +1182,7 @@ def _transcribe_wav_in_chunks( return {"success": False, "transcript": "", "error": f"Chunked transcription failed: {e}"} finally: for chunk_path in chunk_paths: - try: - if os.path.isfile(chunk_path): - os.unlink(chunk_path) - except OSError: - pass + _unlink_quietly(chunk_path) def _split_wav_for_transcription(wav_path: str, *, max_file_size: int) -> List[str]: @@ -1536,25 +1217,17 @@ def _split_wav_for_transcription(wav_path: str, *, max_file_size: int) -> List[s try: with wave.open(chunk_path, "wb") as chunk: - chunk.setnchannels(params.nchannels) - chunk.setsampwidth(params.sampwidth) - chunk.setframerate(params.framerate) - chunk.setcomptype(params.comptype, params.compname) + chunk.setparams(params._replace(nframes=0)) chunk.writeframes(frames) chunk_paths.append(chunk_path) except Exception: - try: - os.unlink(chunk_path) - except OSError: - pass + _unlink_quietly(chunk_path) raise return chunk_paths -# ============================================================================ -# Audio playback (interruptable) -# ============================================================================ +# ── Audio playback (interruptable) ── # Global reference to the active playback process so it can be interrupted. _active_playback: Optional[subprocess.Popen] = None @@ -1581,22 +1254,11 @@ def stop_playback() -> None: pass -def _is_wsl() -> bool: - """True when running inside Windows Subsystem for Linux.""" - try: - with open("/proc/version", "r", encoding="utf-8", errors="replace") as f: - return "microsoft" in f.read().lower() - except Exception: - return False - - def _is_wsl2_env() -> bool: - """Return True when running inside WSL2 (Windows Subsystem for Linux 2). + """True when running inside WSL (Microsoft kernel signature in /proc/version). - Reads /proc/version and checks for the Microsoft kernel signature. - Returns False on any error (non-WSL Linux, Docker, SSH, etc.). - Extracted as a module-level function so tests can patch it directly - without fighting builtins.open patching complexity. + Returns False on any error (non-WSL Linux, Docker, SSH, etc.). A + module-level function so tests can patch it instead of ``builtins.open``. """ try: with open("/proc/version", encoding="utf-8", errors="replace") as _fv: @@ -1605,13 +1267,15 @@ def _is_wsl2_env() -> bool: return False -def _wsl_powershell_tts_available() -> bool: - """Return True when the WSL2 PowerShell TTS playback fallback can be used. +_is_wsl = _is_wsl2_env - This only covers OUTPUT (TTS playback via Media.SoundPlayer on the - Windows host) -- it does NOT make microphone recording work. A caller - using this to relax the audio-environment gate must still surface the - existing PulseAudio-bridge guidance for recording/STT. + +def _wsl_powershell_tts_available() -> bool: + """True when the WSL2 PowerShell TTS playback fallback can be used. + + Only covers OUTPUT (Media.SoundPlayer on the Windows host) — it does NOT + make microphone recording work, so callers relaxing the audio-environment + gate must still surface the PulseAudio-bridge guidance for recording. """ return bool( _is_wsl2_env() @@ -1623,15 +1287,10 @@ def _wsl_powershell_tts_available() -> bool: def play_audio_file(file_path: str) -> bool: """Play an audio file through the default output device. - Strategy: - 1. WAV files via ``sounddevice.play()`` when available. - 2. System commands: ``afplay`` (macOS), ``ffplay`` (cross-platform), - ``aplay`` (Linux ALSA). - - Playback can be interrupted by calling ``stop_playback()``. - - Returns: - ``True`` if playback succeeded, ``False`` otherwise. + WAV files go through ``sounddevice.play()`` when allowed; otherwise system + players: ``afplay`` (macOS), the WSL2 PowerShell bridge, ``ffplay``, + ``aplay`` (Linux). Interruptible via ``stop_playback()``. Returns True on + success. """ # Ref-count real speaker output for the whole call so the thinking-sound # loop (and any other ambient cue) knows audio is flowing right now. @@ -1642,352 +1301,173 @@ def play_audio_file(file_path: str) -> bool: mark_audio_output_active(False) -def _play_audio_file_impl(file_path: str) -> bool: - global _active_playback +def _play_wav_via_sounddevice(file_path: str) -> bool: + """Play a WAV through sounddevice; False if the audio libs are unavailable + or playback failed (caller falls through to system players).""" + try: + sd, np = _import_audio() + with wave.open(file_path, "rb") as wf: + frames = wf.readframes(wf.getnframes()) + audio_data = np.frombuffer(frames, dtype=np.int16) + sample_rate = wf.getframerate() + # WSLg RDP audio needs a warmup to avoid crackling at the start: the + # RDP virtual channel takes ~100 ms to stabilise, and the small default + # blocksize exasperates clock-adjustment jitter (microsoft/wslg#1257). + if _is_wsl(): + silence_samples = int(0.1 * sample_rate) + fade_samples = int(0.1 * sample_rate) + fade = np.linspace(0.0, 1.0, fade_samples, dtype=np.float64) + audio_float = audio_data.astype(np.float64) + audio_float[:fade_samples] *= fade + tail = np.zeros(int(0.05 * sample_rate), dtype=np.int16) + audio_data = np.concatenate([ + np.zeros(silence_samples, dtype=np.int16), + audio_float.astype(np.int16), + tail, + ]) + blocksize = 4096 + else: + blocksize = 0 # default (auto) + + _sd_play_blocking( + sd, audio_data, sample_rate, + timeout=len(audio_data) / sample_rate + 2.0, blocksize=blocksize, + ) + return True + except (ImportError, OSError): + return False # audio libs not available, fall through to system players + except Exception as e: + logger.debug("sounddevice playback failed: %s", e) + return False + + +def _wsl_powershell_player_cmd(file_path: str) -> Optional[List[str]]: + """Build the WSL2 PowerShell fallback player command, or None. + + In WSL without a PulseAudio bridge ffplay/aplay have no device, but + Media.SoundPlayer on the Windows host always does: convert to a + uniquely-named WAV in Windows %TEMP% (so concurrent TTS calls don't + collide) and play it. The WAV is deleted unconditionally, and the ORIGINAL + ffmpeg/powershell exit status is re-raised past that cleanup (rm -f always + exits 0) so the player loop can fall through to the next player. + """ + if not (shutil.which("powershell.exe") and shutil.which("ffmpeg") and _is_wsl2_env()): + return None + try: + import uuid + + def _out(cmd): + return subprocess.check_output(cmd, stderr=subprocess.DEVNULL, timeout=3).decode(errors="replace").strip() + + win_tmp_wsl = _out(["wslpath", "-u", _out(["cmd.exe", "/c", "echo %TEMP%"])]) + if not win_tmp_wsl: + return None + wsl_wav = os.path.join(win_tmp_wsl, f"hermes-tts-{uuid.uuid4().hex[:8]}.wav") + win_wav = _out(["wslpath", "-w", wsl_wav]) + if not win_wav: + return None + win_wav_safe = win_wav.replace("'", "''") + ps_script = f"(New-Object Media.SoundPlayer '{win_wav_safe}').PlaySync()" + ps_cmd = " && ".join([ + shlex.join(["ffmpeg", "-i", file_path, "-f", "wav", wsl_wav, "-loglevel", "quiet", "-y"]), + shlex.join(["powershell.exe", "-NoProfile", "-Command", ps_script]), + ]) + cleanup = shlex.join(["rm", "-f", wsl_wav]) + # Full path so the which(cmd[0]) check in the player loop passes. + return ["/bin/sh", "-c", f"( {ps_cmd} ); rc=$?; {cleanup}; exit $rc"] + except Exception: + return None # WSL path resolution failed; fall through to ffplay/aplay + + +def _system_player_candidates(file_path: str) -> List[List[str]]: + """Ordered system-player commands for this platform.""" + system = platform.system() + players: List[List[str]] = [] + if system == "Darwin": + players.append(["afplay", file_path]) + if system == "Linux": + ps_cmd = _wsl_powershell_player_cmd(file_path) + if ps_cmd: + players.append(ps_cmd) + players.append(["ffplay", "-nodisp", "-autoexit", "-loglevel", "quiet", file_path]) + if system == "Linux": + players.append(["aplay", "-q", file_path]) + return players + + +def _set_active_playback(proc) -> None: + global _active_playback + with _playback_lock: + _active_playback = proc + + +def _run_system_player(cmd: List[str]) -> bool: + """Run one player to completion (interruptible via stop_playback).""" + proc = None + try: + # Sibling of the TTS/STT credential scrub: system audio players must + # not inherit gateway tokens / API keys. + from tools.environments.local import hermes_subprocess_env + + proc = subprocess.Popen( + cmd, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + stdin=subprocess.DEVNULL, + env=hermes_subprocess_env(inherit_credentials=False), + ) + _set_active_playback(proc) + proc.wait(timeout=300) + rc = proc.returncode + _set_active_playback(None) + if rc == 0: + return True + # Non-zero exit: e.g. WSL ffplay/aplay with no audio device, or the + # PowerShell fallback failing. Fall through to the next player. + logger.debug("System player %s exited with code %d, trying next", cmd[0], rc) + except subprocess.TimeoutExpired: + logger.warning("System player %s timed out, killing process", cmd[0]) + if proc is not None: + proc.kill() + proc.wait() + _set_active_playback(None) + except Exception as e: + logger.debug("System player %s failed: %s", cmd[0], e) + _set_active_playback(None) + return False + + +def _play_audio_file_impl(file_path: str) -> bool: if not os.path.isfile(file_path): logger.warning("Audio file not found: %s", file_path) return False - # Skip sounddevice for output where it is not allowed (macOS): PortAudio/ - # CoreAudio init triggers a kTCCServiceMediaLibrary permission prompt even - # though playback needs no media-library access. afplay (added to the - # system-player list below) handles all formats natively instead. - if file_path.endswith(".wav") and _sounddevice_output_allowed(): - try: - sd, np = _import_audio() - with wave.open(file_path, "rb") as wf: - frames = wf.readframes(wf.getnframes()) - audio_data = np.frombuffer(frames, dtype=np.int16) - sample_rate = wf.getframerate() + # sounddevice output is skipped on macOS (PortAudio/CoreAudio init triggers + # a TCC media-library prompt); afplay handles all formats there instead. + if file_path.endswith(".wav") and _sounddevice_output_allowed() and _play_wav_via_sounddevice(file_path): + return True - # WSLg RDP audio needs a warmup to avoid crackling at the start. - # The RDP virtual-channel connection takes ~100 ms to stabilise, - # and small default blocksize exasperates timing jitter caused by - # systemd-timesyncd clock adjustments (microsoft/wslg#1257). - if _is_wsl(): - silence_samples = int(0.1 * sample_rate) - fade_samples = int(0.1 * sample_rate) - fade = np.linspace(0.0, 1.0, fade_samples, dtype=np.float64) - audio_float = audio_data.astype(np.float64) - audio_float[:fade_samples] *= fade - tail = np.zeros(int(0.05 * sample_rate), dtype=np.int16) - audio_data = np.concatenate([ - np.zeros(silence_samples, dtype=np.int16), - audio_float.astype(np.int16), - tail, - ]) - blocksize = 4096 - else: - blocksize = 0 # default (auto) - - sd.play(audio_data, samplerate=sample_rate, blocksize=blocksize) - # sd.wait() calls Event.wait() without timeout — hangs forever if - # the audio device stalls. Poll with a ceiling and force-stop. - duration_secs = len(audio_data) / sample_rate - deadline = time.monotonic() + duration_secs + 2.0 - while sd.get_stream() and sd.get_stream().active and time.monotonic() < deadline: - time.sleep(0.01) - sd.stop() + for cmd in _system_player_candidates(file_path): + if shutil.which(cmd[0]) and _run_system_player(cmd): return True - except (ImportError, OSError): - pass # audio libs not available, fall through to system players - except Exception as e: - logger.debug("sounddevice playback failed: %s", e) - - # Fall back to system audio players (using Popen for interruptability) - system = platform.system() - players = [] - - if system == "Darwin": - players.append(["afplay", file_path]) - - # WSL2 PowerShell fallback: when running in WSL without a PulseAudio - # bridge, ffplay and aplay have no audio device. If powershell.exe and - # ffmpeg are available, convert the audio to a uniquely-named WAV in the - # Windows %TEMP% directory and play it via Media.SoundPlayer -- which - # always has a working audio device on the Windows host (#17608). - # A unique suffix prevents concurrent Hermes TTS calls from colliding on - # the same filename. The WAV is deleted in the shell pipeline - # unconditionally (success or failure), and the ORIGINAL ffmpeg/ - # powershell exit status is preserved past that cleanup so the player - # loop below can correctly fall through to ffplay/aplay on failure. - if system == "Linux" and shutil.which("powershell.exe") and shutil.which("ffmpeg"): - if _is_wsl2_env(): - try: - import uuid - _win_tmp_raw = subprocess.check_output( - ["cmd.exe", "/c", "echo %TEMP%"], - stderr=subprocess.DEVNULL, timeout=3, - ).decode(errors="replace").strip() - _win_tmp_wsl = subprocess.check_output( - ["wslpath", "-u", _win_tmp_raw], - stderr=subprocess.DEVNULL, timeout=3, - ).decode(errors="replace").strip() - if _win_tmp_wsl: - # Unique suffix prevents concurrent TTS playback collision. - _unique = uuid.uuid4().hex[:8] - _wsl_wav = os.path.join(_win_tmp_wsl, f"hermes-tts-{_unique}.wav") - _win_wav = subprocess.check_output( - ["wslpath", "-w", _wsl_wav], - stderr=subprocess.DEVNULL, timeout=3, - ).decode(errors="replace").strip() - if _win_wav: - _win_wav_safe = _win_wav.replace("'", "''") - _ps_script = ( - f"(New-Object Media.SoundPlayer '{_win_wav_safe}').PlaySync()" - ) - _ps_cmd = " && ".join([ - shlex.join(["ffmpeg", "-i", file_path, "-f", "wav", - _wsl_wav, "-loglevel", "quiet", "-y"]), - shlex.join(["powershell.exe", "-NoProfile", "-Command", - _ps_script]), - ]) - _cleanup = shlex.join(["rm", "-f", _wsl_wav]) - # Capture the (ffmpeg && powershell) exit status into - # $rc BEFORE cleanup runs, then exit with that status - # instead of rm -f's (rm -f always exits 0, which - # would otherwise mask a conversion/playback failure - # and prevent falling through to the next player). - _full_cmd = f"( {_ps_cmd} ); rc=$?; {_cleanup}; exit $rc" - # Use full path so the which(cmd[0]) check in the player loop passes. - players.insert(0, ["/bin/sh", "-c", _full_cmd]) - except Exception: - pass # WSL path resolution failed; fall through to ffplay/aplay - - players.append(["ffplay", "-nodisp", "-autoexit", "-loglevel", "quiet", file_path]) - if system == "Linux": - players.append(["aplay", "-q", file_path]) - - for cmd in players: - exe = shutil.which(cmd[0]) - if exe: - try: - # Sibling of TTS/STT credential scrub (#70342 / #56332): system - # audio players must not inherit gateway tokens / API keys. - from tools.environments.local import hermes_subprocess_env - - proc = subprocess.Popen( - cmd, - stdout=subprocess.DEVNULL, - stderr=subprocess.DEVNULL, - stdin=subprocess.DEVNULL, - env=hermes_subprocess_env(inherit_credentials=False), - ) - with _playback_lock: - _active_playback = proc - proc.wait(timeout=300) - rc = proc.returncode - with _playback_lock: - _active_playback = None - if rc == 0: - return True - # Non-zero exit: player failed (e.g. WSL ffplay/aplay with no - # audio device, or the PowerShell fallback's ffmpeg/playback - # step failing). Fall through to the next player in the list. - logger.debug("System player %s exited with code %d, trying next", cmd[0], rc) - except subprocess.TimeoutExpired: - logger.warning("System player %s timed out, killing process", cmd[0]) - proc.kill() - proc.wait() - with _playback_lock: - _active_playback = None - except Exception as e: - logger.debug("System player %s failed: %s", cmd[0], e) - with _playback_lock: - _active_playback = None logger.warning("No audio player available for %s", file_path) return False -# ============================================================================ -# Barge-in — detect the user speaking over TTS playback -# ============================================================================ -def listen_for_speech( - should_stop: Callable[[], bool], - threshold: Optional[int] = None, - sustained_ms: int = 300, - calibration_ms: int = 400, - capture: bool = False, - on_trigger: Optional[Callable[[], None]] = None, - pre_roll_ms: int = 1200, - endpoint_silence_ms: int = 1250, - max_utterance_ms: int = 30_000, -): - """Block until sustained speech is heard on the mic, or *should_stop*. - - Barge-in monitor: run in a side thread while TTS is playing. Without - *capture* it returns ``True`` when the user started talking (cut playback). - With ``capture=True`` it ALSO records the interruption — a rolling - *pre_roll_ms* buffer means the utterance is kept from its first syllable, - not from the moment detection tripped — and keeps rolling until the user - goes quiet for *endpoint_silence_ms*, then returns the WAV path (or - ``None`` if speech never tripped). *on_trigger* fires at the moment of - detection so the caller can stop playback while capture continues. - - The noise floor is calibrated from the first *calibration_ms* of input — - playback is already audible then, so speaker bleed is baked into the - floor and only louder-than-playback speech trips the trigger. Requiring - *sustained_ms* of consecutive above-threshold blocks filters out coughs, - keyboard thumps, and playback transients. - """ - try: - sd, np = _import_audio() - except (ImportError, OSError): - return None if capture else False - - from collections import deque - - block = int(SAMPLE_RATE * 0.03) # 30ms blocks - calib_blocks = max(1, calibration_ms // 30) - trip_blocks = max(1, sustained_ms // 30) - endpoint_blocks = max(1, endpoint_silence_ms // 30) - max_blocks = max(1, max_utterance_ms // 30) - - # Rolling floor window: continuously tracks TTS speaker-bleed volume - # throughout playback, not just the first calibration_ms. This is the - # key fix for false barge-in — a one-shot calibration freezes a floor - # from the opening TTS passage, but later louder passages exceed the - # stale floor and false-trigger. The rolling window keeps the floor - # current so only genuinely louder-than-playback speech trips the VAD. - floor_window: "deque[float]" = deque(maxlen=max(calib_blocks, 100)) # ~3s rolling - pre_roll: deque = deque(maxlen=max(1, pre_roll_ms // 30)) - consecutive = 0 - min_floor = 0.0 # baseline from initial calibration; floor never drops below this - block_idx = 0 # block counter for diagnostic logging - - try: - with sd.InputStream(samplerate=SAMPLE_RATE, channels=1, dtype="int16", blocksize=block) as stream: - while not should_stop(): - data, _ = stream.read(block) - rms = float(np.sqrt(np.mean(data.astype(np.float64) ** 2))) - if capture: - pre_roll.append(data.copy()) - block_idx += 1 - - # Wait for at least calib_blocks before evaluating. During - # the initial warmup we always feed the window so calibration - # has data to work with. - if len(floor_window) < calib_blocks: - floor_window.append(rms) - continue - - # Lock a minimum floor from the initial calibration samples. - # During inter-sentence pauses the rolling window can flush - # with near-silence, collapsing the 90th-percentile floor - # toward zero and false-triggering on the next rising - # sentence. min_floor keeps the trigger from ever dropping - # below the baseline TTS playback level established during - # the initial calibration_ms window. - # - # If the grace period ended during an inter-sentence gap the - # calibration samples near-silence. Locking a near-zero - # floor sets the trigger so low that TTS blocks exceed it, - # are excluded from the rolling window (rms >= trigger), and - # the floor freezes — guaranteeing a false trigger the moment - # TTS resumes. Clamp min_floor to SILENCE_RMS_THRESHOLD * 2 - # (400 RMS) so the 8x multiplier yields a trigger of at least - # (500-2000 RMS) stays below it and feeds the rolling window, - # while genuine speech (3000-8000 RMS) can still trip it. - if min_floor == 0.0 and len(floor_window) >= calib_blocks: - _pct90 = float(np.percentile(list(floor_window), 90)) - min_floor = max(_pct90, SILENCE_RMS_THRESHOLD * 2) - else: - _pct90 = float(np.percentile(list(floor_window), 90)) - - # Use the 90th percentile of the ROLLING window for the - # noise floor so the trigger reflects the loudest parts of - # recent playback — not a frozen snapshot from TTS onset. - _floor = max(_pct90, min_floor) - # 8.0x multiplier: TTS speaker bleed has wide - # volume variation between sentences and within sentences. - # At 5x, louder TTS passages exceed the trigger, get - # excluded from the floor window, and create a low-stale - # floor that false-triggers on the next loud passage. - # 8x gives enough headroom for TTS dynamics to stay below - # the trigger and get absorbed into the rolling floor. - trigger = max(float(threshold or SILENCE_RMS_THRESHOLD * 2), _floor * 8.0) - # Ceiling: never let the trigger exceed 4000 RMS, otherwise - # a very loud TTS passage would push the trigger so high - # that genuine speech (which is typically 3000–8000 RMS) - # couldn't trip it. - trigger = min(trigger, 4000.0) - - # Only feed the floor window with blocks that are NOT above - # the current trigger — speech blocks would inflate the floor - # and make the trigger unreachable. - if rms < trigger: - floor_window.append(rms) - - consecutive = consecutive + 1 if rms >= trigger else 0 - if consecutive > 0: - logger.debug( - "VAD above-trigger: block=%d rms=%.0f floor=%.0f trigger=%.0f " - "consec=%d/%d min_floor=%.0f window_len=%d", - block_idx, rms, _floor, trigger, consecutive, - trip_blocks, min_floor, len(floor_window), - ) - if consecutive < trip_blocks: - continue - - # Tripped — the user is talking over playback. - logger.info( - "VAD TRIPPED: block=%d rms=%.0f floor=%.0f trigger=%.0f " - "consec=%d min_floor=%.0f — cutting TTS playback", - block_idx, rms, _floor, trigger, consecutive, min_floor, - ) - if on_trigger: - try: - on_trigger() - except Exception as e: - logger.debug("Barge-in trigger callback failed: %s", e) - if not capture: - return True - - # Keep rolling until the user goes quiet. Playback is stopped - # now, so plain silence endpointing (recorder threshold) works. - frames: List[Any] = list(pre_roll) - quiet = 0 - for _ in range(max_blocks): - data, _ = stream.read(block) - frames.append(data.copy()) - rms = float(np.sqrt(np.mean(data.astype(np.float64) ** 2))) - quiet = quiet + 1 if rms < SILENCE_RMS_THRESHOLD else 0 - if quiet >= endpoint_blocks: - break - return AudioRecorder._write_wav(np.concatenate(frames, axis=0)) - except Exception as e: - logger.debug("Barge-in listener failed: %s", e) - return None if capture else False - - -# ============================================================================ -# Full-duplex agent-turn listener -# ============================================================================ -# -# One listener for the WHOLE agent turn in continuous voice mode: armed the -# moment an utterance is submitted, disarmed when the turn is fully done -# (response + TTS finished). It replaces the scattered per-playback barge -# monitors, which had two class-level failures: -# -# 1. HALF-DUPLEX GAP: the monitor only spawned when TTS playback started, -# so during LLM generation (seconds to minutes) there was NO microphone -# listener at all — the user could not interject by voice. -# 2. PLAYBACK DEAFNESS: the monitor calibrated its noise floor WHILE the -# speaker was already blasting TTS, baking speaker bleed into the floor; -# with an 8x multiplier the trigger became unreachable for normal speech, -# and requiring a full second of strictly CONSECUTIVE 30ms blocks above -# trigger meant any intra-word dip reset the counter. -# -# This listener calibrates against the QUIET room at turn start (before any -# playback exists), freezes that baseline through playback (never calibrating -# against its own speaker bleed), and trips on a windowed majority of blocks -# instead of a strict consecutive run. +# ── Full-duplex agent-turn listener ── +# One listener for the WHOLE agent turn in continuous voice mode: armed when an +# utterance is submitted, disarmed when the turn (response + TTS) is done. It +# replaced per-playback barge monitors that (1) only listened during TTS, so +# the user could not interject during LLM generation, and (2) calibrated their +# noise floor against active speaker bleed with a strict consecutive-block +# counter, making the trigger unreachable for normal speech. This listener +# calibrates against the QUIET room at turn start, freezes that baseline +# through playback, and trips on a windowed majority of blocks. # Minimum trigger while TTS audio is flowing. Speaker bleed reaching the mic -# through air at conversational volume is typically well under this (a few -# hundred RMS at arm's length; ~1000-1400 with loud speakers close to the +# is typically a few hundred RMS (~1000-1400 with loud speakers close to the # mic), while direct speech at normal distance measures 3000-8000 RMS. PLAYBACK_MIN_TRIGGER = 1500.0 @@ -1995,29 +1475,38 @@ PLAYBACK_MIN_TRIGGER = 1500.0 # what normal speech (3000-8000 RMS) can reach. TRIGGER_CEILING = 4000.0 -# Default trigger multiplier over the quiet-room floor. The old 8x default -# only made sense when the floor was (wrongly) calibrated against active TTS -# bleed; against a genuine quiet-room baseline (typically 50-300 RMS) the -# synthetic-frame tests show 3x separates speech from ambient cleanly while -# staying reachable (300 RMS floor * 3 = 900 trigger vs 3000+ RMS speech). +# Trigger multiplier over the quiet-room floor (typically 50-300 RMS): 3x +# separates speech from ambient cleanly while staying reachable +# (300 RMS floor * 3 = 900 trigger vs 3000+ RMS speech). DEFAULT_BARGE_MULTIPLIER = 3.0 -def _voice_debug_enabled() -> bool: - return os.environ.get("HERMES_VOICE_DEBUG", "").strip() == "1" - - def _vad_log(msg: str) -> None: """VAD decision-point diagnostic — always logger.debug, plus stderr when HERMES_VOICE_DEBUG=1 so live hardware tuning doesn't need a log tail.""" logger.debug(msg) - if _voice_debug_enabled(): + if os.environ.get("HERMES_VOICE_DEBUG", "").strip() == "1": try: print(f"[voice-vad] {msg}", file=sys.stderr, flush=True) except Exception: pass +def _capture_until_quiet(stream, np, block: int, pre_roll, *, endpoint_blocks: int, max_blocks: int) -> str: + """Keep reading after a trip until *endpoint_blocks* of quiet (or + *max_blocks*), then write pre-roll + capture to a WAV and return its path. + Playback was cut by the trigger, so plain silence endpointing works.""" + frames: List[Any] = list(pre_roll) + quiet = 0 + for _ in range(max_blocks): + data, _ = stream.read(block) + frames.append(data.copy()) + quiet = quiet + 1 if _rms(np, data) < SILENCE_RMS_THRESHOLD else 0 + if quiet >= endpoint_blocks: + break + return AudioRecorder._write_wav(np.concatenate(frames, axis=0)) + + def full_duplex_listen( should_stop: Callable[[], bool], is_playing: Optional[Callable[[], bool]] = None, @@ -2032,26 +1521,18 @@ def full_duplex_listen( ) -> Optional[str]: """Listen across an ENTIRE agent turn; return the captured interruption. - Runs from utterance-submit to turn-complete. Two phases, decided per - 30ms block by *is_playing* (usually ``is_audio_output_active``): + Two phases, decided per 30ms block by *is_playing* (usually + ``is_audio_output_active``): ``generation`` (no TTS) — the first + *calibration_ms* of the quiet room set the noise floor and the trigger is + quiet_floor x *multiplier*; ``playback`` (TTS flowing) — the quiet + baseline is HELD (never recalibrated against speaker bleed), the trigger + is clamped up to ``PLAYBACK_MIN_TRIGGER`` so bleed alone can't trip it, + and a *grace_ms* window after playback starts suppresses onset transients. - * ``generation`` — no TTS audio flowing. The room is quiet; the first - *calibration_ms* establish the noise floor (pre-playback calibration). - Trigger = quiet_floor x *multiplier* (clamped to a sane minimum), i.e. - ordinary speech detection. - * ``playback`` — TTS audio flowing. The quiet baseline is HELD (never - recalibrated against speaker bleed); the trigger is additionally - clamped up to ``PLAYBACK_MIN_TRIGGER`` so bleed alone can't trip it, - and a *grace_ms* window after playback first starts suppresses trips - from the playback onset transient. - - Detection is a windowed majority — >=80% of the last *sustained_ms* - worth of blocks above trigger (with the current block above) — so - intra-word energy dips don't reset progress the way the old strictly- - consecutive counter did. - - On detection ``on_trigger(phase)`` fires (cut TTS / interrupt the turn), - then capture continues from the rolling *pre_roll_ms* buffer until + Detection is a windowed majority — >=80% of the last *sustained_ms* of + blocks above trigger (current block included) — so intra-word dips don't + reset progress. On detection ``on_trigger(phase)`` fires, capture + continues from the rolling *pre_roll_ms* buffer until *endpoint_silence_ms* of quiet, and the WAV path is returned. Returns ``None`` when *should_stop* ends the turn without speech. """ @@ -2082,13 +1563,18 @@ def full_duplex_listen( blocks_since_playback = 10_000 block_idx = 0 + def _floor(seq) -> tuple: + """(pct90, floor): 90th percentile of the quiet-phase RMS window, floor never below the silence threshold.""" + pct90 = float(np.percentile(list(seq), 90)) if seq else float(SILENCE_RMS_THRESHOLD) + return pct90, max(pct90, float(SILENCE_RMS_THRESHOLD)) + try: with sd.InputStream( samplerate=SAMPLE_RATE, channels=1, dtype="int16", blocksize=block ) as stream: while not should_stop(): data, _ = stream.read(block) - rms = float(np.sqrt(np.mean(data.astype(np.float64) ** 2))) + rms = _rms(np, data) pre_roll.append(data.copy()) block_idx += 1 @@ -2101,12 +1587,7 @@ def full_duplex_listen( if not playing: ambient.append(rms) if len(ambient) >= calib_blocks or playing: - pct90 = ( - float(np.percentile(list(ambient), 90)) - if ambient - else float(SILENCE_RMS_THRESHOLD) - ) - quiet_floor = max(pct90, float(SILENCE_RMS_THRESHOLD)) + pct90, quiet_floor = _floor(ambient) floor_locked = True _vad_log( f"calibrated quiet floor={quiet_floor:.0f} " @@ -2133,22 +1614,14 @@ def full_duplex_listen( # Trigger: quiet baseline x multiplier, phase-clamped. trigger = quiet_floor * mult - if playing: - trigger = max(trigger, PLAYBACK_MIN_TRIGGER) - else: - trigger = max(trigger, float(SILENCE_RMS_THRESHOLD) * 2) + trigger = max(trigger, PLAYBACK_MIN_TRIGGER if playing else float(SILENCE_RMS_THRESHOLD) * 2) trigger = min(trigger, TRIGGER_CEILING) - # Keep the quiet floor current with ambient drift — but ONLY - # while nothing is playing (never absorb speaker bleed) and - # the block isn't speech. + # Track ambient drift — ONLY while nothing is playing (never + # absorb speaker bleed) and the block isn't speech. if not playing and rms < trigger: ambient.append(rms) - if ambient: - quiet_floor = max( - float(np.percentile(list(ambient), 90)), - float(SILENCE_RMS_THRESHOLD), - ) + _, quiet_floor = _floor(ambient) above = rms >= trigger if above and grace_remaining > 0: @@ -2184,32 +1657,20 @@ def full_duplex_listen( except Exception as e: logger.debug("full-duplex trigger callback failed: %s", e) - # Capture until the user goes quiet. Playback was cut by - # on_trigger, so plain silence endpointing works. - frames: List[Any] = list(pre_roll) - quiet = 0 - for _ in range(max_blocks): - data, _ = stream.read(block) - frames.append(data.copy()) - rms = float(np.sqrt(np.mean(data.astype(np.float64) ** 2))) - quiet = quiet + 1 if rms < SILENCE_RMS_THRESHOLD else 0 - if quiet >= endpoint_blocks: - break - return AudioRecorder._write_wav(np.concatenate(frames, axis=0)) + return _capture_until_quiet( + stream, np, block, pre_roll, + endpoint_blocks=endpoint_blocks, max_blocks=max_blocks, + ) except Exception as e: logger.debug("Full-duplex listener failed: %s", e) return None -# ============================================================================ -# Requirements check -# ============================================================================ +# ── Requirements check ── def _check_plugin_stt_provider(provider: str) -> bool: """Return True when *provider* resolves to an available STT plugin.""" - if not provider: - return False - key = provider.lower().strip() - if key == "none": + key = (provider or "").lower().strip() + if not key or key == "none": return False try: from agent.transcription_registry import get_provider @@ -2223,9 +1684,7 @@ def _check_plugin_stt_provider(provider: str) -> bool: _ensure_plugins_discovered(force=True) plugin_provider = get_provider(key) except Exception as exc: # noqa: BLE001 - discovery failure is non-fatal - logger.debug( - "STT plugin requirements check skipped for '%s': %s", key, exc, - ) + logger.debug("STT plugin requirements check skipped for '%s': %s", key, exc) return False if plugin_provider is None: @@ -2237,21 +1696,29 @@ def _check_plugin_stt_provider(provider: str) -> bool: logger.warning( "STT plugin provider '%s' is_available() raised during requirements " "check: %s - treating as unavailable", - key, - exc, - exc_info=True, + key, exc, exc_info=True, ) return False +# STT providers handled natively by tools.transcription_tools -> status label. +_NATIVE_STT_LABELS = { + "local": "local faster-whisper", + "local_command": "local command", + "groq": "Groq", + "openai": "OpenAI", + "mistral": "Mistral Voxtral", + "xai": "xAI Grok STT", + "elevenlabs": "ElevenLabs Scribe", +} + + def check_voice_requirements() -> Dict[str, Any]: """Check if all voice mode requirements are met. - Returns: - Dict with ``available``, ``audio_available``, ``stt_available``, - ``missing_packages``, and ``details``. + Returns dict with ``available``, ``audio_available``, ``stt_available``, + ``missing_packages``, ``details`` and ``environment``. """ - # Determine STT provider availability from tools.transcription_tools import ( _get_provider, _load_stt_config, @@ -2261,15 +1728,7 @@ def check_voice_requirements() -> Dict[str, Any]: stt_config = _load_stt_config() stt_enabled = is_stt_enabled(stt_config) stt_provider = _get_provider(stt_config) - native_stt_available = stt_provider in { - "local", - "local_command", - "groq", - "openai", - "mistral", - "xai", - "elevenlabs", - } + native_stt_available = stt_provider in _NATIVE_STT_LABELS command_stt_config = None plugin_stt_available = False if stt_enabled and not native_stt_available: @@ -2291,7 +1750,6 @@ def check_voice_requirements() -> Dict[str, Any]: if not has_audio: missing.extend(["sounddevice", "numpy"]) - # Environment detection env_check = detect_audio_environment() available = has_audio and stt_available and env_check["available"] @@ -2306,20 +1764,8 @@ def check_voice_requirements() -> Dict[str, Any]: if not stt_enabled: details_parts.append("STT provider: DISABLED in config (stt.enabled: false)") - elif stt_provider == "local": - details_parts.append("STT provider: OK (local faster-whisper)") - elif stt_provider == "local_command": - details_parts.append("STT provider: OK (local command)") - elif stt_provider == "groq": - details_parts.append("STT provider: OK (Groq)") - elif stt_provider == "openai": - details_parts.append("STT provider: OK (OpenAI)") - elif stt_provider == "mistral": - details_parts.append("STT provider: OK (Mistral Voxtral)") - elif stt_provider == "xai": - details_parts.append("STT provider: OK (xAI Grok STT)") - elif stt_provider == "elevenlabs": - details_parts.append("STT provider: OK (ElevenLabs Scribe)") + elif stt_provider in _NATIVE_STT_LABELS: + details_parts.append(f"STT provider: OK ({_NATIVE_STT_LABELS[stt_provider]})") elif command_stt_config is not None: details_parts.append(f"STT provider: OK (command: {stt_provider})") elif plugin_stt_available: @@ -2346,18 +1792,10 @@ def check_voice_requirements() -> Dict[str, Any]: } -# ============================================================================ -# Temp file cleanup -# ============================================================================ +# ── Temp file cleanup ── def cleanup_temp_recordings(max_age_seconds: int = 3600) -> int: - """Remove old temporary voice recording files. - - Args: - max_age_seconds: Delete files older than this (default: 1 hour). - - Returns: - Number of files deleted. - """ + """Remove ``recording_*.wav`` temp files older than *max_age_seconds* + (default 1 hour); returns the number deleted.""" if not os.path.isdir(_TEMP_DIR): return 0 @@ -2367,8 +1805,7 @@ def cleanup_temp_recordings(max_age_seconds: int = 3600) -> int: for entry in os.scandir(_TEMP_DIR): if entry.is_file() and entry.name.startswith("recording_") and entry.name.endswith(".wav"): try: - age = now - entry.stat().st_mtime - if age > max_age_seconds: + if now - entry.stat().st_mtime > max_age_seconds: os.unlink(entry.path) deleted += 1 except OSError: diff --git a/tools/wake_word.py b/tools/wake_word.py index f433dc6602..ac59bab889 100644 --- a/tools/wake_word.py +++ b/tools/wake_word.py @@ -1,32 +1,17 @@ """Wake-word ("Hey Hermes") detection — hands-free session trigger. -A lightweight, always-on hotword listener that fires a callback when a wake -phrase is spoken — the "Hey Siri" / "Alexa" pattern. Shared by the CLI, TUI, and -desktop GUI (one of them owns it, gated by ``wake_surface_enabled``): say the -wake word, Hermes opens a fresh session and captures voice via the existing -pipeline, then answers. +An always-on hotword listener shared by CLI, TUI and desktop GUI (one owns it, +gated by ``wake_surface_enabled``): on wake Hermes opens a fresh session and +captures voice via the existing pipeline. Engines (openwakeword default, +sherpa open-vocabulary, porcupine premium) are all on-device and live in +:mod:`tools.wake_word_engines`; this module owns config, the capture loop and +the process-wide listener singleton. -Three engines, all fully on-device (no audio leaves the machine for detection): - -* **openwakeword** (default, free, no API key) — loads an ONNX model. Defaults - to the bundled "hey hermes" model (``tools/wakewords/``) so the wake word - works out of the box; or point ``wake_word.openwakeword.model`` at a built-in - name (``hey_jarvis``, ``alexa``, …) or a custom ``.onnx`` for another phrase. -* **sherpa** (free, no API key, open vocabulary) — sherpa-onnx keyword - spotting. Detects ANY typed phrase with no training: set - ``wake_word.phrase`` and the phrase is tokenized at runtime against a small - streaming zipformer model (~13 MB English model, one-time download). -* **porcupine** (premium) — Picovoice's engine. Needs ``PORCUPINE_ACCESS_KEY``; - supports built-in keywords and custom ``.ppn`` files from the Picovoice - Console. - -Audio capture reuses the same 16 kHz mono int16 ``sounddevice`` path as voice -mode. The detector runs on its own daemon thread; callers ``pause()`` it while a -voice turn holds the microphone and ``resume()`` it once the system is idle -again (two input streams on one device is unreliable cross-platform). - -Nothing here mutates agent context or the prompt cache — on wake we hand a plain -string to the caller, exactly like a voice transcript. +Capture reuses voice mode's 16 kHz mono int16 ``sounddevice`` path on a daemon +thread; callers ``pause()`` while a voice turn holds the mic and ``resume()`` +once idle (two input streams on one device is unreliable cross-platform). +Nothing here touches agent context or the prompt cache — on wake the caller +gets a plain string, like a transcript. """ from __future__ import annotations @@ -36,50 +21,58 @@ import os import sys import threading import time +from dataclasses import dataclass from pathlib import Path from typing import Any, Callable, Dict, Optional +from tools.wake_word_engines import ( # noqa: F401 (re-exported for callers/tests) + _SHERPA_KWS_MODEL_DIR, _SHERPA_KWS_MODEL_URL, _Engine, _OpenWakeWordEngine, _PorcupineEngine, + _SherpaKwsEngine, _ensure_sherpa_model, _looks_like_path, _sherpa_model_root, +) + logger = logging.getLogger(__name__) -# 16 kHz mono int16 — Whisper-native and what both engines expect. +# 16 kHz mono int16 — Whisper-native and what every engine expects. SAMPLE_RATE = 16000 -# Minimum gap between two consecutive wake fires, so one "hey hermes" can't -# retrigger across several frames while the caller is still reacting. +# Minimum gap between two wake fires, so one "hey hermes" can't retrigger +# across several frames while the caller is still reacting. _FIRE_COOLDOWN_SECONDS = 2.0 _START_TIMEOUT_SECONDS = 5.0 -# Ambient-speech rejection: openWakeWord scores one ~80ms frame at a time, and a -# stray phoneme in background conversation can spike a single frame over the -# threshold. A real utterance of the phrase holds the score high across several -# consecutive frames, so we require N-in-a-row above threshold before firing. -# This is the primary lever against unintended triggers on ambient talk. +# Ambient-speech rejection: require N consecutive over-threshold frames before +# firing (a stray phoneme spikes one frame; a real phrase holds several). _DEFAULT_CONFIRMATION_FRAMES = 3 -# Dead-mic detection: an int16 stream whose peak stays at/below this for this -# many consecutive seconds is flagged as silent. Desktop push-to-talk and the -# backend listener use different capture paths, so one can work while the -# backend-selected stream is all zeros. +# Dead-mic detection: an int16 stream whose peak stays at/below _SILENCE_PEAK +# for this many consecutive seconds is flagged silent (desktop push-to-talk and +# the backend listener use different capture paths, so one can work while the +# backend-selected stream is all zeros). _SILENCE_PEAK = 10 _SILENCE_ALERT_SECONDS = 10 +# provider alias -> (engine class name on this module, lazy_deps feature). +# Unknown providers probe as openwakeword but fail to build. +_PROVIDERS: Dict[str, tuple[str, str]] = { + "porcupine": ("_PorcupineEngine", "wake.porcupine"), + **{k: ("_SherpaKwsEngine", "wake.sherpa") for k in ("sherpa", "sherpa-onnx", "kws", "open")}, + **{k: ("_OpenWakeWordEngine", "wake.openwakeword") for k in ("openwakeword", "oww", "local")}, +} + class WakeWordInUse(RuntimeError): """Raised when another surface or process owns the wake-word listener.""" -# --------------------------------------------------------------------------- -# Config -# --------------------------------------------------------------------------- +# ── Config ── _DEFAULTS: Dict[str, Any] = { "enabled": False, "surface": "auto", "input_device": None, - # Where PCM is captured: - # "local" — PortAudio on the backend host (historic default) - # "client" — desktop/TUI streams int16 frames via wake.feed - # "auto" — local when a device exists, else client capture + # Where PCM is captured: "local" (PortAudio on the backend host), + # "client" (desktop/TUI streams int16 frames via wake.feed), or + # "auto" (local when a device exists, else client capture). "capture": "auto", "provider": "openwakeword", "phrase": "hey hermes", @@ -88,8 +81,8 @@ _DEFAULTS: Dict[str, Any] = { "start_new_session": True, } -# Bundled "hey hermes" model (tools/wakewords/) — the default, so the wake word -# works out of the box. Config names in _ALIASES resolve to it, not a built-in. +# Bundled "hey hermes" model (tools/wakewords/) — the default. Config names in +# _ALIASES resolve to it, not to an openWakeWord built-in. _BUNDLED_MODEL_NAME = "hey_hermes" _BUNDLED_MODEL_ALIASES = frozenset({"", "hey_hermes", "hey hermes", "hermes"}) @@ -107,16 +100,9 @@ def _is_macos_arm64() -> bool: def default_inference_framework() -> str: - """The openWakeWord backend to use on this platform. - - openWakeWord's ONNX backend produces near-zero scores on macOS ARM64 — its - shared *embedding* model is the broken stage (the melspectrogram front-end - and the wake classifier both match tflite exactly). The detector arms, the - microphone works, and no phrase can ever cross the threshold. Prefer the - tflite backend there; ONNX stays the default everywhere else. - - Upstream: https://github.com/dscripka/openWakeWord/issues/336 - """ + """tflite on macOS ARM64, onnx elsewhere: openWakeWord's ONNX *embedding* + model scores near-zero on Apple Silicon (upstream #336) — the detector arms + but no phrase ever crosses threshold.""" return "tflite" if _is_macos_arm64() else "onnx" @@ -124,16 +110,10 @@ _warned_onnx_coerced = False def resolve_inference_framework(cfg: Dict[str, Any]) -> str: - """Resolve the effective openWakeWord backend from config. - - Honors an explicit ``openwakeword.inference_framework`` — EXCEPT the one - combination that is provably dead: an explicit ``onnx`` on macOS ARM64, - where ONNX's embedding model never lets a phrase cross threshold (upstream - #336). Existing macOS users who pinned ``onnx`` before the tflite fix landed - would otherwise keep a wake word that arms but never fires. Coerce that one - case to tflite (with a one-time warning) instead of silently shipping a dead - ear. Every other explicit value is respected as-is; empty falls back to the - platform default. + """Effective openWakeWord backend: explicit ``openwakeword.inference_framework`` + or the platform default. The one provably dead combination — explicit + ``onnx`` on macOS ARM64 (upstream #336) — is coerced to tflite with a + one-time warning so a pre-fix pin doesn't keep a wake word that never fires. """ global _warned_onnx_coerced @@ -157,14 +137,12 @@ def resolve_inference_framework(cfg: Dict[str, Any]) -> str: return framework - def ensure_tflite_runtime() -> bool: """Make ``import tflite_runtime.interpreter`` resolve, returning success. - openWakeWord hardcodes that import but only declares ``tflite-runtime`` for - ``platform_system == "Linux"``; on macOS the equivalent wheel is - ``ai-edge-litert``. Alias the module so the upstream import succeeds. The - alias is process-local — nothing is written to site-packages. + openWakeWord hardcodes that import but only declares ``tflite-runtime`` on + Linux; on macOS the equivalent wheel is ``ai-edge-litert``. Alias the + module in-process (nothing is written to site-packages). """ try: import tflite_runtime.interpreter # noqa: F401 @@ -204,6 +182,15 @@ def _get(cfg: Dict[str, Any], key: str) -> Any: return _DEFAULTS.get(key) if val is None else val +def _clamped(cfg: Dict[str, Any], key: str, cast, lo, hi): + """Numeric config value via ``cast``, defaulting on junk, clamped to lo..hi.""" + try: + n = cast(_get(cfg, key)) + except (TypeError, ValueError): + n = cast(_DEFAULTS[key]) + return min(max(n, lo), hi) + + def _provider(cfg: Dict[str, Any]) -> str: return str(_get(cfg, "provider")).strip().lower() or "openwakeword" @@ -215,32 +202,20 @@ def _input_device(cfg: Dict[str, Any]) -> int | str | None: return None if isinstance(raw, int): return raw - value = str(raw).strip() - return value or None + return str(raw).strip() or None def _sensitivity(cfg: Dict[str, Any]) -> float: - raw = _get(cfg, "sensitivity") - try: - s = float(raw) - except (TypeError, ValueError): - s = float(_DEFAULTS["sensitivity"]) - return min(max(s, 0.0), 1.0) + return _clamped(cfg, "sensitivity", float, 0.0, 1.0) def _confirmation_frames(cfg: Dict[str, Any]) -> int: - """How many consecutive over-threshold frames are required to fire. + """Consecutive over-threshold frames required to fire, clamped 1..10. - ``1`` restores the old single-frame behaviour; higher values reject - ambient-speech blips at the cost of a few tens of ms of extra latency. - Clamped to a sane 1..10. + ``1`` restores single-frame behaviour; higher rejects ambient blips at the + cost of a few tens of ms of latency. """ - raw = _get(cfg, "confirmation_frames") - try: - n = int(raw) - except (TypeError, ValueError): - n = _DEFAULT_CONFIRMATION_FRAMES - return min(max(n, 1), 10) + return _clamped(cfg, "confirmation_frames", int, 1, 10) def wake_phrase(cfg: Optional[Dict[str, Any]] = None) -> str: @@ -257,9 +232,11 @@ def resolve_capture_mode( ) -> str: """Return ``local`` or ``client`` capture mode for this arm. - ``prefer_client`` is set by remote desktop (Mac mic, headless backend). - ``force_local`` keeps CLI/TUI on the process mic. Config ``capture`` is - ``auto`` | ``local`` | ``client``. + ``prefer_client`` is set by remote desktop; ``force_local`` keeps CLI/TUI on + the process mic. Under ``auto`` a working backend input always wins (local + desktops keep PortAudio + ``input_device``); client is the fallback only for + a preferring surface with no usable backend mic — CLI/TUI stay local so + status reports the real requirement rather than a path nothing will feed. """ cfg = cfg if cfg is not None else load_wake_word_config() if force_local: @@ -269,20 +246,18 @@ def resolve_capture_mode( return "client" if raw == "local": return "local" - # auto: a working backend input always wins so local desktops keep - # PortAudio and the configured ``input_device`` selection. Client capture - # is the fallback for a preferring surface (desktop remote) on a backend - # with no usable mic — the headless VPS/Cloud case. - if _local_input_device_ready(): - return "local" - if prefer_client: + if prefer_client and not _local_input_device_ready(): return "client" - # No local mic and no client preference (CLI/TUI): stay local so status - # reports the real requirement instead of advertising a capture path - # nothing will feed. return "local" +def _input_channels(info: Any) -> int: + channels = info.get("max_input_channels") if isinstance(info, dict) else None + if channels is None: + channels = getattr(info, "max_input_channels", 0) + return int(channels or 0) + + def _local_input_device_ready() -> bool: """True when PortAudio is importable and at least one input device exists.""" try: @@ -291,24 +266,12 @@ def _local_input_device_ready() -> bool: return False try: devices = sd.query_devices() - except Exception: - return False - if isinstance(devices, dict): - return int(devices.get("max_input_channels") or 0) > 0 - try: - for dev in devices: - channels = dev.get("max_input_channels") if isinstance(dev, dict) else None - if channels is None: - channels = getattr(dev, "max_input_channels", 0) - if int(channels or 0) > 0: - return True - except Exception: - return False - # Also accept a resolvable default input (some hosts list devices oddly). - try: - info = sd.query_devices(None, "input") - channels = info.get("max_input_channels") if isinstance(info, dict) else 0 - return int(channels or 0) > 0 + if isinstance(devices, dict): + return _input_channels(devices) > 0 + if any(_input_channels(dev) > 0 for dev in devices): + return True + # Also accept a resolvable default input (some hosts list devices oddly). + return _input_channels(sd.query_devices(None, "input")) > 0 except Exception: return False @@ -316,9 +279,9 @@ def _local_input_device_ready() -> bool: def wake_surface_enabled(surface: str, cfg: Optional[Dict[str, Any]] = None) -> bool: """Should ``surface`` (``cli`` / ``tui`` / ``gui``) host the listener? - True when the wake word is enabled and the configured ``surface`` is either - ``auto`` or this exact surface. ``auto`` makes a surface eligible; the - process/machine ownership lock still permits only the first claimant. + True when enabled and the configured ``surface`` is ``auto`` or this exact + surface. ``auto`` only makes a surface eligible; the process/machine + ownership lock still permits a single claimant. """ cfg = cfg if cfg is not None else load_wake_word_config() if not cfg.get("enabled"): @@ -327,9 +290,7 @@ def wake_surface_enabled(surface: str, cfg: Optional[Dict[str, Any]] = None) -> return want == "auto" or want == surface.strip().lower() -# --------------------------------------------------------------------------- -# Multi-profile phrase enrollment (open-vocabulary routing) -# --------------------------------------------------------------------------- +# ── Multi-profile phrase enrollment (open-vocabulary routing) ── def _active_profile_name() -> str: try: @@ -343,11 +304,10 @@ def _active_profile_name() -> str: def enrolled_profile_phrases() -> Dict[str, str]: """Map ``profile name -> wake phrase`` for every wake-enabled profile. - Reads each profile's own ``config.yaml`` raw (cheap, no full config merge). - A profile is enrolled when its ``wake_word.enabled`` is truthy; its phrase - defaults to ``"hey "`` when unset. Used by the sherpa engine to - listen for every enrolled profile's phrase at once and route the wake to - the matching profile. Best-effort: unreadable profiles are skipped. + Reads each profile's own ``config.yaml`` raw (``load_config()`` targets only + the ACTIVE profile). Enrolled = ``wake_word.enabled`` truthy; phrase defaults + to ``"hey "``. The sherpa engine listens for all of them at once and + routes the wake to the matching profile. Best-effort: unreadable skipped. """ phrases: Dict[str, str] = {} try: @@ -357,10 +317,7 @@ def enrolled_profile_phrases() -> Dict[str, str]: for info in list_profiles(): name = getattr(info, "name", None) or str(info) try: - # Multi-profile read: load_config() targets the ACTIVE - # profile's home, so read each profile's file directly. - cfg_path = Path(get_profile_dir(name)) / "config.yaml" - raw = read_user_config_raw(cfg_path) + raw = read_user_config_raw(Path(get_profile_dir(name)) / "config.yaml") wc = raw.get("wake_word") or {} if not isinstance(wc, dict) or not wc.get("enabled"): continue @@ -374,9 +331,7 @@ def enrolled_profile_phrases() -> Dict[str, str]: return phrases -# --------------------------------------------------------------------------- -# Audio capture (lazy — never import sounddevice at module load) -# --------------------------------------------------------------------------- +# ── Audio capture (lazy — never import sounddevice at module load) ── def _import_audio(): import numpy as np @@ -396,8 +351,8 @@ def _audio_available() -> bool: def _describe_input_device(sd, selector: int | str | None) -> Dict[str, Any]: """Resolve a PortAudio selector into JSON-safe diagnostics. - Device discovery is diagnostic only. ``InputStream`` remains the authority - on whether the selected device can actually open at the requested format. + Diagnostic only: ``InputStream`` remains the authority on whether the + device can actually open at the requested format. """ details: Dict[str, Any] = {"selector": selector} try: @@ -405,28 +360,26 @@ def _describe_input_device(sd, selector: int | str | None) -> Dict[str, Any]: except Exception as e: details["error"] = str(e) return details + if not isinstance(info, dict): + return details - if isinstance(info, dict): - name = info.get("name") - if name: - details["name"] = str(name) - channels = info.get("max_input_channels") - if isinstance(channels, (int, float)): - details["max_input_channels"] = int(channels) - rate = info.get("default_samplerate") - if isinstance(rate, (int, float)): - details["default_samplerate"] = float(rate) - hostapi_index = info.get("hostapi") - if isinstance(hostapi_index, (int, float)): - details["hostapi_index"] = int(hostapi_index) - try: - hostapi = sd.query_hostapis(int(hostapi_index)) - hostapi_name = hostapi.get("name") if isinstance(hostapi, dict) else None - if hostapi_name: - details["hostapi"] = str(hostapi_name) - except Exception: - pass - + if info.get("name"): + details["name"] = str(info["name"]) + for key, out_key, cast in ( + ("max_input_channels", "max_input_channels", int), + ("default_samplerate", "default_samplerate", float), + ("hostapi", "hostapi_index", int), + ): + if isinstance(info.get(key), (int, float)): + details[out_key] = cast(info[key]) + if "hostapi_index" in details: + try: + hostapi = sd.query_hostapis(details["hostapi_index"]) + hostapi_name = hostapi.get("name") if isinstance(hostapi, dict) else None + if hostapi_name: + details["hostapi"] = str(hostapi_name) + except Exception: + pass return details @@ -458,13 +411,12 @@ def _resample_audio_frame(np, frame, output_length: int): return np.zeros(output_length, dtype=np.int16) if source.size > output_length: - # Match the desktop wake capture path: average each source window when - # reducing the rate so speech energy is retained instead of decimating. + # Average each source window when reducing (matches the desktop wake + # capture path) so speech energy is retained instead of decimated. edges = np.linspace(0, source.size, output_length + 1, dtype=np.int64) values = np.add.reduceat(source, edges[:-1]) / np.diff(edges) else: - # Unusual low-rate devices need interpolation to reach the 16 kHz - # frame size expected by every wake-word engine. + # Unusual low-rate devices: interpolate up to the 16 kHz frame size. source_positions = np.arange(source.size, dtype=np.float64) target_positions = np.linspace(0, source.size - 1, output_length) values = np.interp(target_positions, source_positions, source) @@ -492,356 +444,22 @@ def silent_audio_hint(details: Dict[str, Any]) -> str: ) -# --------------------------------------------------------------------------- -# Engines -# --------------------------------------------------------------------------- - -class _Engine: - """Minimal hotword-engine contract: feed int16 frames, get a bool.""" - - frame_length: int = 1280 # 80 ms at 16 kHz - - #: Optional (matched phrase, profile name) of the most recent fire. - #: Multi-phrase engines (sherpa) set this for profile routing; the - #: single-phrase engines leave it None (callers fall back to the - #: 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).""" - pass - - 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) - ) - - -class _OpenWakeWordEngine(_Engine): - """openWakeWord — free, local ONNX hotword detection.""" - - # openWakeWord recommends 80 ms frames (1280 samples) for efficiency. - frame_length = 1280 - - 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 - - sub = cfg.get("openwakeword") if isinstance(cfg.get("openwakeword"), dict) else {} - model_ref = str(sub.get("model") or _BUNDLED_MODEL_NAME).strip() - framework = resolve_inference_framework(cfg) - # openWakeWord returns a 0..1 score per frame; sensitivity IS the raw - # threshold a score must clear. Higher = stricter (fewer false fires). - # Default 0.6 sits above openWakeWord's permissive 0.5 baseline, which - # let near-misses like "hey hor" through. - self._threshold = _sensitivity(cfg) - self._confirm_needed = _confirmation_frames(cfg) - self._confirm_streak = 0 - - # openWakeWord silently downgrades tflite -> onnx when no tflite runtime - # imports (model.py). On macOS ARM64 that lands on the backend whose - # embedding model is broken, so the listener would arm and never fire. - # Install + bridge the runtime first, and refuse the downgrade rather - # than ship a dead ear. - if framework == "tflite" and not ensure_tflite_runtime(): - # Same lazy-install contract as every other backend; the platform - # gate lives here because dep specs can't carry PEP 508 markers. - try: - lazy_deps.ensure("wake.openwakeword.tflite", prompt=False) - except Exception as e: - logger.debug("wake word: tflite runtime install failed: %s", e) - if not ensure_tflite_runtime(): - if _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") - framework = "onnx" - - # Default (or explicit "hey_hermes") → the bundled model; a built-in name - # or custom path is used as-is. - if model_ref.lower() in _BUNDLED_MODEL_ALIASES: - model_ref = _bundled_wakeword_path(framework) - - # openWakeWord needs its shared feature models (melspectrogram + embedding) - # for ANY model — download_models() fetches those first on every call, so a - # custom path must call it too, else a fresh install crashes on a missing - # melspectrogram.onnx. A built-in name additionally pulls that pretrained - # model; a path matches nothing in the catalog and is a no-op beyond base. - 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) - models = [model_ref] - - self._model = Model(wakeword_models=models, inference_framework=framework) - self._labels = list(self._model.models.keys()) - - def process(self, frame) -> bool: - scores = self._model.predict(frame) - over = any(score >= self._threshold for score in scores.values()) - # Require N consecutive over-threshold frames: a real phrase holds the - # score high across frames, a stray ambient phoneme spikes just one. - if over: - self._confirm_streak += 1 - if self._confirm_streak >= self._confirm_needed: - self._confirm_streak = 0 - return True - return False - self._confirm_streak = 0 - return False - - 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, unlike openWakeWord/Porcupine custom models. -_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. - - The configured ``wake_word.phrase`` is BPE-tokenized at runtime against the - model's vocabulary, so "hey hermes", "hey coder", or any other phrase works - immediately. Here ``phrase`` is DETECTION config, not a cosmetic label. - """ - - # sherpa's streaming zipformer consumes arbitrary chunk sizes; 1280 - # samples (80 ms) matches the shared capture path. - frame_length = 1280 - - 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 - - sub = cfg.get("sherpa") if isinstance(cfg.get("sherpa"), dict) else {} - model_dir = str(sub.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 ("hey hermes" / "hey coder" / ...). display-name → - # profile is kept for routing the match back. - phrase = str(_get(cfg, "phrase") or "hey hermes").strip() - own_profile = _active_profile_name() - phrase_map: Dict[str, str] = {phrase: own_profile} - if bool(cfg.get("profile_routing", True)): - for prof, p in enrolled_profile_phrases().items(): - phrase_map.setdefault(p.strip(), prof) - - phrases = list(phrase_map) - # Runtime tokenization of the arbitrary phrases — the open-vocab core. - 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 - #: (phrase display name, profile) of the most recent fire, for routing. - self.last_match: Optional[tuple[str, str]] = None - - # Map the shared 0..1 sensitivity onto sherpa's keywords_threshold. - # 0.5 lands exactly on sherpa's recommended default (0.25); live TTS - # matrix testing showed our previous stricter mapping (0.35) missed - # ~12% of true positives while 0.25 held zero false fires. - threshold = 0.05 + 0.4 * _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(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 all 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)." - ) - - sub = cfg.get("porcupine") if isinstance(cfg.get("porcupine"), dict) else {} - keyword = str(sub.get("keyword") or "jarvis").strip() - # Porcupine's `sensitivities` runs the OPPOSITE way to our shared knob: - # per Picovoice, higher = more true positives AND more false alarms - # (looser). Our config contract is "higher = stricter" everywhere, so - # invert it here to keep one consistent meaning across all engines. - porcupine_sensitivity = 1.0 - _sensitivity(cfg) - - kwargs: Dict[str, Any] = {"access_key": access_key, "sensitivities": [porcupine_sensitivity]} - if _looks_like_path(keyword): - kwargs["keyword_paths"] = [keyword] - else: - kwargs["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 - +# ── Engines (implementations live in tools.wake_word_engines) ── def _build_engine(cfg: Dict[str, Any]) -> _Engine: provider = _provider(cfg) - if provider == "porcupine": - return _PorcupineEngine(cfg) - if provider in ("sherpa", "sherpa-onnx", "kws", "open"): - return _SherpaKwsEngine(cfg) - if provider in ("openwakeword", "oww", "local"): - return _OpenWakeWordEngine(cfg) - raise ValueError(f"Unknown wake_word provider: {provider!r}") + if provider not in _PROVIDERS: + raise ValueError(f"Unknown wake_word provider: {provider!r}") + return globals()[_PROVIDERS[provider][0]](cfg) -# --------------------------------------------------------------------------- -# Requirements probe (for /wake status + enable path) -# --------------------------------------------------------------------------- +# ── Requirements probe (for /wake status + enable path) ── def _stt_ready() -> bool: """Is a speech-to-text provider configured and enabled? - A wake without STT arms the mic but every captured utterance dies at - transcription — a useless (and confusing) experience. Same standard as - voice mode's ``check_voice_requirements``: enabled + a real provider. + A wake without STT arms the mic but every utterance dies at transcription. + Same standard as voice mode's ``check_voice_requirements``. """ try: from tools.transcription_tools import _get_provider, _load_stt_config, is_stt_enabled @@ -852,18 +470,16 @@ def _stt_ready() -> bool: return False -def _tts_ready() -> bool: - """Can the configured text-to-speech provider run (or install at first use)? +_LAZY_TTS_FEATURES = {"edge": "tts.edge", "elevenlabs": "tts.elevenlabs", "mistral": "tts.mistral"} - The wake flow is fully hands-free (wake → speak → hear the reply); without - TTS the reply is silent and the loop is pointless. + +def _tts_ready() -> bool: + """Can the configured TTS provider run (or install at first use)? PROBE, not an installer: ``check_tts_requirements`` lazily pip-installs the - provider SDK via ``_import_*`` → ``lazy_deps.ensure`` — running that inside - a status poll froze wake.status for the length of a pip install (and a - failed install marked the wake word unavailable, unmounting the desktop - ear). When the provider's deps aren't installed yet, "installable at first - use" counts as ready and we never touch pip from here. + provider SDK, which froze wake.status polls for a whole pip run (a failed + install unmounted the desktop ear). Uninstalled deps count as ready iff + lazy installs are allowed; pip is never touched from here. """ try: from tools.tts_tool import _get_provider, _load_tts_config @@ -872,18 +488,12 @@ def _tts_ready() -> bool: except Exception: return False - _LAZY_TTS_FEATURES = { - "edge": "tts.edge", - "elevenlabs": "tts.elevenlabs", - "mistral": "tts.mistral", - } feature = _LAZY_TTS_FEATURES.get(provider) if feature is not None: try: from tools import lazy_deps if not lazy_deps.is_available(feature): - # Not installed: ready iff it can install at first speak. return lazy_deps._allow_lazy_installs() except Exception: return False @@ -902,37 +512,26 @@ def check_wake_word_requirements(cfg: Optional[Dict[str, Any]] = None) -> Dict[s provider = _provider(cfg) from tools import lazy_deps - if provider == "porcupine": - feature = "wake.porcupine" - elif provider in ("sherpa", "sherpa-onnx", "kws", "open"): - feature = "wake.sherpa" - else: - feature = "wake.openwakeword" + feature = _PROVIDERS.get(provider, ("", "wake.openwakeword"))[1] deps_ok = lazy_deps.is_available(feature) lazy_ok = lazy_deps._allow_lazy_installs() - # The audio probe imports sounddevice + numpy — two of the very packages - # the lazy installer would fetch — so it can only be trusted once the - # feature's deps are installed. On a fresh install (deps missing, lazy - # installs allowed) we defer the mic check: the engine constructors call - # ``lazy_deps.ensure()`` and the stream-open surfaces any real audio - # problem. Gating ``available`` on the probe here made the lazy-install - # path unreachable (the probe always failed before ensure() could run). + # The audio probe imports sounddevice + numpy — packages the lazy installer + # would fetch — so only trust it once deps are installed; on a fresh install + # the engine constructors' ``lazy_deps.ensure()`` + stream-open surface any + # real audio problem (gating on the probe made lazy install unreachable). audio_ok = _audio_available() if deps_ok else False key_ok = True - # The full wake loop is wake → record → STT → agent → TTS. Arming without - # either end configured gives a mic that hears you and then does nothing - # the user can perceive — refuse with a pointer instead. + # Loop is wake → record → STT → agent → TTS; without either end the mic + # hears you and nothing perceptible happens — refuse with a hint. stt_ok = _stt_ready() tts_ok = _tts_ready() hint = "" - # The tflite backend needs a runtime openWakeWord doesn't declare off Linux. - # Report it as a real remediation instead of arming a detector that can't fire. + # tflite needs a runtime openWakeWord doesn't declare off Linux; report it + # as a remediation instead of arming a detector that can't fire. tflite_ok = True - if provider not in ("porcupine", "sherpa", "sherpa-onnx", "kws", "open"): - framework = resolve_inference_framework(cfg) - if framework == "tflite": - tflite_ok = ensure_tflite_runtime() or lazy_deps.is_available("wake.openwakeword.tflite") or lazy_ok + if feature == "wake.openwakeword" and resolve_inference_framework(cfg) == "tflite": + tflite_ok = ensure_tflite_runtime() or lazy_deps.is_available("wake.openwakeword.tflite") or lazy_ok if provider == "porcupine" and not (os.getenv("PORCUPINE_ACCESS_KEY") or "").strip(): key_ok = False @@ -951,14 +550,9 @@ def check_wake_word_requirements(cfg: Optional[Dict[str, Any]] = None) -> Dict[s f"(Voice section) or see the voice-mode docs.") capture_mode = resolve_capture_mode(cfg) - local_input_ok = _local_input_device_ready() if deps_ok else False # Client capture needs deps (engine) but not a server-side PortAudio device. if capture_mode == "client": - mic_ok = deps_ok or (not deps_ok and lazy_ok) - if deps_ok and not hint: - # No server mic required; clear the local-device hint if that was set. - if hint.startswith("Microphone capture needs"): - hint = "" + mic_ok = deps_ok or lazy_ok else: mic_ok = (deps_ok and audio_ok) or (not deps_ok and lazy_ok) if deps_ok and not audio_ok and not hint: @@ -973,7 +567,7 @@ def check_wake_word_requirements(cfg: Optional[Dict[str, Any]] = None) -> Dict[s "provider": provider, "deps_available": deps_ok, "audio_available": audio_ok, - "local_input_available": local_input_ok, + "local_input_available": _local_input_device_ready() if deps_ok else False, "capture": capture_mode, "access_key_set": key_ok, "stt_available": stt_ok, @@ -983,9 +577,36 @@ def check_wake_word_requirements(cfg: Optional[Dict[str, Any]] = None) -> Dict[s } -# --------------------------------------------------------------------------- -# Detector -# --------------------------------------------------------------------------- +# ── Detector ── + +@dataclass +class _Capture: + """One armed audio source: a PortAudio stream (local) or the feed queue (client).""" + + stream: Any = None # sounddevice.InputStream, None in client mode + queue: Any = None # client-capture frame queue, None in local mode + np: Any = None + rate: int = SAMPLE_RATE + frame_length: int = 1280 # samples per read at ``rate`` + + def read(self): + """One raw block; None when no client frame arrived within 250 ms. + Stream errors propagate.""" + if self.stream is not None: + return self.stream.read(self.frame_length)[0] + try: + return self.queue.get(timeout=0.25) + except Exception: + return None + + def close(self) -> None: + try: + if self.stream is not None: + self.stream.stop() + self.stream.close() + except Exception: + pass + class WakeWordDetector: """Background hotword listener. Fires ``on_wake()`` when the phrase is heard. @@ -1019,9 +640,8 @@ class WakeWordDetector: import queue as _queue self._audio_q: "_queue.Queue[Any]" = _queue.Queue(maxsize=64) - # True when the stream is open but every frame is (near-)silence. - # Surfaced via wake.status / /wake status so users can tell "armed" - # from "deaf". + # True when the stream is open but every frame is (near-)silence, so + # status surfaces can tell "armed" from "deaf". self.audio_silent = False self._silent_frames = 0 @@ -1033,8 +653,8 @@ class WakeWordDetector: def feed(self, pcm_int16) -> None: """Enqueue one int16 mono frame (or raw bytes) for client capture. - Frame length should match ``engine.frame_length`` (typically 1280 samples - at 16 kHz). Short frames are zero-padded; long frames are split. + Short frames are zero-padded to ``engine.frame_length``; long frames are + split. On queue overflow the oldest frame is dropped to stay real-time. """ if not self.external_audio: return @@ -1049,12 +669,8 @@ class WakeWordDetector: fl = int(self.engine.frame_length) if fl <= 0: return - # Split / pad into engine frames - offset = 0 - n = int(arr.shape[0]) - while offset < n: + for offset in range(0, int(arr.shape[0]), fl): chunk = arr[offset : offset + fl] - offset += fl if chunk.shape[0] < fl: pad = np.zeros(fl, dtype=np.int16) pad[: chunk.shape[0]] = chunk @@ -1062,12 +678,8 @@ class WakeWordDetector: try: self._audio_q.put_nowait(chunk) except Exception: - # Drop oldest on overflow so we stay real-time - try: + try: # full: drop the oldest frame, then retry once self._audio_q.get_nowait() - except Exception: - pass - try: self._audio_q.put_nowait(chunk) except Exception: pass @@ -1122,14 +734,8 @@ class WakeWordDetector: finally: self._callback_inflight.clear() - def _run(self, ready: threading.Event, - startup_errors: list[BaseException]) -> None: - frame_length = self.engine.frame_length - capture_frame_length = frame_length - capture_rate = SAMPLE_RATE - np = None - stream = None - + def _open_capture(self, frame_length: int) -> _Capture: + """Open the audio source; raises on any local-mic failure.""" if self.external_audio: # Drain any stale frames from a previous arm. try: @@ -1141,48 +747,87 @@ class WakeWordDetector: "wake word: client-capture mode (frame=%d, rate=%d) — waiting for wake.feed", frame_length, SAMPLE_RATE, ) - else: - try: - sd, np = _import_audio() - except (ImportError, OSError) as e: - logger.error("wake word: audio libraries unavailable: %s", e) - startup_errors.append(e) - ready.set() - return + return _Capture(queue=self._audio_q, frame_length=frame_length) - self.input_device_details = _describe_input_device(sd, self.input_device) - capture_rate = _capture_sample_rate(self.input_device_details) - capture_frame_length = max( - 1, int(round(frame_length * capture_rate / SAMPLE_RATE)) + try: + sd, np = _import_audio() + except (ImportError, OSError) as e: + logger.error("wake word: audio libraries unavailable: %s", e) + raise + + self.input_device_details = _describe_input_device(sd, self.input_device) + cap = _Capture(np=np, rate=_capture_sample_rate(self.input_device_details)) + cap.frame_length = max(1, int(round(frame_length * cap.rate / SAMPLE_RATE))) + logger.info( + "wake word: opening microphone device=%s selector=%r hostapi=%s " + "default_rate=%s capture_rate=%d engine_rate=%d", + self.input_device_details.get("name") or "system default", + self.input_device, + self.input_device_details.get("hostapi") or "unknown", + self.input_device_details.get("default_samplerate") or "unknown", + cap.rate, + SAMPLE_RATE, + ) + try: + cap.stream = sd.InputStream( + device=self.input_device, + samplerate=cap.rate, + channels=1, + dtype="int16", + blocksize=cap.frame_length, ) - logger.info( - "wake word: opening microphone device=%s selector=%r hostapi=%s " - "default_rate=%s capture_rate=%d engine_rate=%d", - self.input_device_details.get("name") or "system default", - self.input_device, - self.input_device_details.get("hostapi") or "unknown", - self.input_device_details.get("default_samplerate") or "unknown", - capture_rate, - SAMPLE_RATE, - ) - try: - stream = sd.InputStream( - device=self.input_device, - samplerate=capture_rate, - channels=1, - dtype="int16", - blocksize=capture_frame_length, + cap.stream.start() + except Exception as e: + logger.error("wake word: failed to open microphone: %s", e) + raise + return cap + + def _note_silence(self, frame, silent_alert_frames: int) -> None: + """Track consecutive near-zero frames; flag/unflag ``audio_silent``.""" + try: + peak = int(abs(frame).max()) if len(frame) else 0 + except Exception: + peak = _SILENCE_PEAK + 1 + if peak <= _SILENCE_PEAK: + self._silent_frames += 1 + if self._silent_frames == silent_alert_frames: + self.audio_silent = True + logger.warning( + "wake word: mic delivers only silence (peak<=%d for %ds); %s", + _SILENCE_PEAK, _SILENCE_ALERT_SECONDS, + silent_audio_hint(self.input_device_details), ) - stream.start() - except Exception as e: - logger.error("wake word: failed to open microphone: %s", e) - startup_errors.append(e) - ready.set() - return + elif self._silent_frames: + if self.audio_silent: + logger.info("wake word: mic audio detected — stream healthy") + self._silent_frames = 0 + self.audio_silent = False - # Drop any buffered audio/feature state so a resume right after a voice - # turn can't immediately re-fire on audio captured before the pause (the - # wake → voice → resume → wake runaway loop). + def _fire(self) -> None: + """Honor the cooldown, then run ``on_wake`` on its own thread (once).""" + now = time.monotonic() + if now - self._last_fire < self.cooldown: + logger.debug("wake word: detection within cooldown — ignored") + return + self._last_fire = now + logger.info("wake word: phrase detected — firing callback") + if not self._callback_inflight.is_set(): + self._callback_inflight.set() + threading.Thread(target=self._dispatch_wake, daemon=True, name="wake-word-callback").start() + + def _run(self, ready: threading.Event, + startup_errors: list[BaseException]) -> None: + frame_length = self.engine.frame_length + try: + cap = self._open_capture(frame_length) + except Exception as e: + startup_errors.append(e) + ready.set() + return + + # Drop buffered audio/feature state so a resume right after a voice turn + # can't re-fire on audio captured before the pause (the wake → voice → + # resume → wake runaway loop). try: self.engine.reset() except Exception: @@ -1192,83 +837,40 @@ class WakeWordDetector: frame_length, SAMPLE_RATE, self.external_audio) ready.set() failed = False - # ~seconds of consecutive near-zero frames before we flag the stream - # as silent. silent_alert_frames = max(1, int(_SILENCE_ALERT_SECONDS * SAMPLE_RATE / max(1, frame_length))) try: while not self._stop.is_set(): try: - if self.external_audio: - try: - frame = self._audio_q.get(timeout=0.25) - except Exception: - # No client frames yet — count as silence for status. - self._silent_frames += 1 - if self._silent_frames == silent_alert_frames: - self.audio_silent = True - continue - data = frame - else: - data, _overflow = stream.read(capture_frame_length) + data = cap.read() except Exception as e: logger.warning("wake word: stream read error: %s", e) failed = not self._stop.is_set() break - frame = data[:, 0] if getattr(data, "ndim", 1) == 2 else data - if capture_rate != SAMPLE_RATE: - frame = _resample_audio_frame(np, frame, frame_length) - try: - peak = int(abs(frame).max()) if len(frame) else 0 - except Exception: - peak = _SILENCE_PEAK + 1 - if peak <= _SILENCE_PEAK: + if data is None: + # No client frames yet — count as silence for status. self._silent_frames += 1 if self._silent_frames == silent_alert_frames: self.audio_silent = True - logger.warning( - "wake word: mic delivers only silence (peak<=%d for %ds); %s", - _SILENCE_PEAK, _SILENCE_ALERT_SECONDS, - silent_audio_hint(self.input_device_details), - ) - elif self._silent_frames: - if self.audio_silent: - logger.info("wake word: mic audio detected — stream healthy") - self._silent_frames = 0 - self.audio_silent = False + continue + frame = data[:, 0] if getattr(data, "ndim", 1) == 2 else data + if cap.rate != SAMPLE_RATE: + frame = _resample_audio_frame(cap.np, frame, frame_length) + self._note_silence(frame, silent_alert_frames) try: fired = self.engine.process(frame) except Exception as e: logger.debug("wake word: engine error: %s", e) continue if fired: - now = time.monotonic() - if now - self._last_fire >= self.cooldown: - self._last_fire = now - logger.info("wake word: phrase detected — firing callback") - if not self._callback_inflight.is_set(): - self._callback_inflight.set() - threading.Thread( - target=self._dispatch_wake, - daemon=True, - name="wake-word-callback", - ).start() - else: - logger.debug("wake word: detection within cooldown — ignored") + self._fire() finally: - if stream is not None: - try: - stream.stop() - stream.close() - except Exception: - pass + cap.close() logger.info("wake word: stream closed") if failed and self.on_failure is not None: self.on_failure(self) -# --------------------------------------------------------------------------- -# Process-wide singleton (mirrors hermes_cli.voice's continuous API) -# --------------------------------------------------------------------------- +# ── Process-wide singleton (mirrors hermes_cli.voice's continuous API) ── _detector: Optional[WakeWordDetector] = None _detector_owner: object | None = None @@ -1282,25 +884,31 @@ def _lock_path() -> Path: return get_default_hermes_root() / "runtime" / "wake-word.lock" +def _flock(handle, acquire: bool) -> None: + """Non-blocking exclusive lock (or unlock) of one byte / whole file, per OS.""" + if os.name == "nt": + import msvcrt + + if acquire: # msvcrt needs at least one byte to lock + handle.seek(0, os.SEEK_END) + if handle.tell() == 0: + handle.write(b"\0") + handle.flush() + handle.seek(0) + msvcrt.locking(handle.fileno(), msvcrt.LK_NBLCK if acquire else msvcrt.LK_UNLCK, 1) + else: + import fcntl + + fcntl.flock(handle.fileno(), (fcntl.LOCK_EX | fcntl.LOCK_NB) if acquire else fcntl.LOCK_UN) + + def _acquire_machine_lock(path: Optional[Path] = None): """Acquire the cross-process microphone lease, or raise WakeWordInUse.""" lock_path = path or _lock_path() lock_path.parent.mkdir(parents=True, exist_ok=True) handle = open(lock_path, "a+b") try: - if os.name == "nt": - import msvcrt - - handle.seek(0, os.SEEK_END) - if handle.tell() == 0: - handle.write(b"\0") - handle.flush() - handle.seek(0) - msvcrt.locking(handle.fileno(), msvcrt.LK_NBLCK, 1) - else: - import fcntl - - fcntl.flock(handle.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB) + _flock(handle, True) except (OSError, BlockingIOError) as e: handle.close() raise WakeWordInUse("Wake-word microphone is already owned.") from e @@ -1311,31 +919,34 @@ def _release_machine_lock(handle) -> None: if handle is None: return try: - if os.name == "nt": - import msvcrt - - handle.seek(0) - msvcrt.locking(handle.fileno(), msvcrt.LK_UNLCK, 1) - else: - import fcntl - - fcntl.flock(handle.fileno(), fcntl.LOCK_UN) + _flock(handle, False) except OSError: pass finally: handle.close() +def _clear_singleton_locked() -> tuple[Optional[WakeWordDetector], Any]: + """Forget the armed detector (caller holds ``_detector_lock``); returns (detector, lock handle).""" + global _detector, _detector_owner, _detector_file_lock + det, handle = _detector, _detector_file_lock + _detector = None + _detector_owner = None + _detector_file_lock = None + return det, handle + + +def _owned_detector(owner: object) -> Optional[WakeWordDetector]: + """The armed detector iff ``owner`` holds the lease (caller holds the lock).""" + return _detector if _detector is not None and _detector_owner is owner else None + + def _detector_failed(detector: WakeWordDetector) -> None: """Release ownership if the active microphone stream dies unexpectedly.""" - global _detector, _detector_owner, _detector_file_lock with _detector_lock: if _detector is not detector: return - lock_handle = _detector_file_lock - _detector = None - _detector_owner = None - _detector_file_lock = None + _, lock_handle = _clear_singleton_locked() try: detector.engine.close() finally: @@ -1383,52 +994,46 @@ def start_listening( detector.start() return detector except Exception: - if _detector is not None: - try: - _detector.stop() - except Exception: - pass - _detector = None - _detector_owner = None - _detector_file_lock = None + det, _ = _clear_singleton_locked() + try: + if det is not None: + det.stop() + except Exception: + pass _release_machine_lock(lock_handle) raise def owns_listener(owner: object) -> bool: with _detector_lock: - return _detector is not None and _detector_owner is owner + return _owned_detector(owner) is not None + + +def _owned_call(owner: object, method: str) -> bool: + with _detector_lock: + det = _owned_detector(owner) + if det is None: + return False + getattr(det, method)() + return True def pause_listening(*, owner: object) -> bool: """Release the microphone only when ``owner`` holds the lease.""" - with _detector_lock: - if _detector is None or _detector_owner is not owner: - return False - _detector.pause() - return True + return _owned_call(owner, "pause") def resume_listening(*, owner: object) -> bool: """Re-open the microphone only when ``owner`` holds the lease.""" - with _detector_lock: - if _detector is None or _detector_owner is not owner: - return False - _detector.resume() - return True + return _owned_call(owner, "resume") def stop_listening(*, owner: object) -> bool: """Fully stop the detector only when ``owner`` holds the lease.""" - global _detector, _detector_owner, _detector_file_lock with _detector_lock: - if _detector is None or _detector_owner is not owner: + if _owned_detector(owner) is None: return False - det = _detector - lock_handle = _detector_file_lock - _detector = None - _detector_owner = None - _detector_file_lock = None + det, lock_handle = _clear_singleton_locked() try: det.stop() finally: @@ -1436,9 +1041,13 @@ def stop_listening(*, owner: object) -> bool: return True -def is_listening() -> bool: +def _current_detector() -> Optional[WakeWordDetector]: with _detector_lock: - det = _detector + return _detector + + +def is_listening() -> bool: + det = _current_detector() return det is not None and det.running @@ -1446,18 +1055,15 @@ def audio_is_silent() -> bool: """True when the armed stream has delivered only silence (dead mic). The stream opens fine but every frame is zeros, so detection can never - fire. Lets status surfaces show "listening but the microphone appears - silent" instead of a healthy state. + fire; status surfaces show "listening but the microphone appears silent". """ - with _detector_lock: - det = _detector + det = _current_detector() return det is not None and det.audio_silent def get_input_device_status(cfg: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: """Return configured/active PortAudio input diagnostics for status UIs.""" - with _detector_lock: - det = _detector + det = _current_detector() if det is not None: return dict(det.input_device_details) @@ -1473,11 +1079,8 @@ def get_input_device_status(cfg: Optional[Dict[str, Any]] = None) -> Dict[str, A def get_last_match() -> Optional[tuple[str, str]]: """(matched phrase, profile) of the most recent wake fire, if the engine reports per-phrase matches (sherpa multi-profile routing). None otherwise.""" - with _detector_lock: - det = _detector - if det is None: - return None - return getattr(det.engine, "last_match", None) + det = _current_detector() + return None if det is None else getattr(det.engine, "last_match", None) def feed_audio(*, owner: object, pcm_int16) -> bool: @@ -1486,19 +1089,16 @@ def feed_audio(*, owner: object, pcm_int16) -> bool: Returns True when the frame was accepted for ``owner``'s armed detector. """ with _detector_lock: - if _detector is None or _detector_owner is not owner: + det = _owned_detector(owner) + if det is None or not det.external_audio: return False - if not _detector.external_audio: - return False - det = _detector det.feed(pcm_int16) return True def detector_frame_info() -> Dict[str, Any]: """Sample rate + frame length for client capture streamers.""" - with _detector_lock: - det = _detector + det = _current_detector() if det is None: return {"sample_rate": SAMPLE_RATE, "frame_length": 1280} return { diff --git a/tools/wake_word_engines.py b/tools/wake_word_engines.py new file mode 100644 index 0000000000..adca319077 --- /dev/null +++ b/tools/wake_word_engines.py @@ -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.")``) 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