refactor(tools): compact STT/TTS docstrings and comments, keep every invariant

This commit is contained in:
Teknium
2026-09-02 23:11:57 -07:00
parent d34703fa86
commit 7ed48064db
7 changed files with 139 additions and 198 deletions
+4 -2
View File
@@ -64,7 +64,7 @@ _STT_M4A_ENCODE_ARGS = ("-vn", "-ac", "1", "-ar", "16000", "-c:a", "aac", "-b:a"
def _run_ffmpeg_stt_encode(ffmpeg: str, input_path: str, output_path: str, *, audio_filter: Optional[str] = None) -> None:
"""Run the shared STT m4a encode, optionally with an ``-af`` filter. Raises on failure (callers own the semantics)."""
"""Run the shared STT m4a encode, optionally with an ``-af`` filter. Raises on failure; callers own the semantics."""
command = [ffmpeg, "-y", "-i", input_path]
if audio_filter:
command += ["-af", audio_filter]
@@ -236,7 +236,9 @@ def _probe_audio_duration(file_path: str) -> Optional[float]:
ffprobe = _find_ffprobe_binary()
if not ffprobe:
return None
command = [ffprobe, "-v", "error", "-show_entries", "format=duration", "-of", "default=noprint_wrappers=1:nokey=1", file_path]
command = [
ffprobe, "-v", "error", "-show_entries", "format=duration", "-of", "default=noprint_wrappers=1:nokey=1", file_path,
]
try:
return float(_run_quiet(command, timeout=30).stdout.strip())
except Exception: # noqa: BLE001 - probe is best-effort
+29 -58
View File
@@ -1,10 +1,10 @@
"""Cloud STT providers.
OpenAI-SDK-shaped backends (groq, openai, deepinfra), Mistral Voxtral, the REST
multipart backends (xAI, ElevenLabs), and OpenAI audio credential resolution
(config > keyless local server > env > managed Nous gateway). Every name is
re-imported by ``tools/transcription_tools.py`` (patch surface), which is
imported lazily here so origin patches still intercept.
OpenAI-SDK-shaped backends (groq, openai, deepinfra), Mistral Voxtral, REST multipart
backends (xAI, ElevenLabs), and OpenAI audio credential resolution (config > keyless
local server > env > managed Nous gateway). Every name is re-imported by
``tools/transcription_tools.py`` (patch surface), imported lazily here so origin
patches still intercept.
"""
from __future__ import annotations
@@ -35,11 +35,7 @@ def _has_xai_stt_credentials() -> bool:
def _with_openai_client(api_key: str, base_url: Optional[str], file_path: str, log_label: str, body):
"""Run ``body(client)`` against a fresh OpenAI SDK client (30s timeout, no SDK retries).
Always closes the client; any exception maps to the shared envelope via
:func:`_openai_sdk_failure`.
"""
"""Run ``body(client)`` on a fresh OpenAI SDK client (30s timeout, no retries); always closed, errors -> envelope."""
try:
from openai import OpenAI
client = OpenAI(api_key=api_key, base_url=base_url, timeout=30, max_retries=0)
@@ -64,8 +60,8 @@ def _cloud_failure(exc: BaseException, file_path: str, label: str, detail: Optio
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.
APIConnectionError is checked before APITimeoutError (its subclass) so timeouts
report as connection errors, as they always have.
"""
try:
from openai import APIError, APIConnectionError, APITimeoutError
@@ -93,11 +89,7 @@ def _sdk_prompt_kwargs(language: Optional[str], prompt: Optional[str]) -> Dict[s
def _transcribe_groq(
file_path: str, model_name: str, *, language: Optional[str] = None, prompt: Optional[str] = None
) -> Dict[str, Any]:
"""Transcribe using Groq Whisper API (free tier available).
Language: hook override > ``stt.groq.language`` > ``stt.language`` > env;
otherwise Groq auto-detects.
"""
"""Transcribe via the Groq Whisper API; language: hook > ``stt.groq.language`` > ``stt.language`` > env > auto."""
from tools.transcription_tools import _HAS_OPENAI, _resolve_provider_key, _resolve_stt_language
api_key = _resolve_provider_key("GROQ_API_KEY", "groq")
if not api_key:
@@ -131,9 +123,8 @@ def _transcribe_openai(
) -> Dict[str, Any]:
"""Transcribe via the OpenAI ``audio.transcriptions.create`` SDK shape.
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``.
Shared by every OpenAI-compatible endpoint (DeepInfra etc.): explicit ``api_key``/
``base_url`` skip the OpenAI-only auth chain; ``provider_label`` names the response's provider.
"""
from tools.transcription_tools import _HAS_OPENAI, _resolve_openai_audio_client_config, _resolve_stt_language
if api_key is None:
@@ -149,8 +140,7 @@ def _transcribe_openai(
if not _HAS_OPENAI:
return _error_result("openai package not installed")
# Auto-correct a Groq-only model on the native OpenAI path only —
# 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 serve it).
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
@@ -167,14 +157,12 @@ def _transcribe_openai(
}
if language:
if model_name == "gpt-transcribe":
# gpt-transcribe replaces ``language`` with a ``languages``
# list and rejects requests sending the legacy field.
# gpt-transcribe takes a ``languages`` list and rejects 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 sent when set so the no-hook, no-config request stays byte-identical.
if prompt: # only when set so the bare request stays byte-identical
create_kwargs["prompt"] = prompt
return client.audio.transcriptions.create(**create_kwargs)
@@ -185,8 +173,7 @@ def _transcribe_openai(
message = str(exc).lower()
if not any(k in message for k in ("unsupported", "corrupted", "invalid file")):
raise
# Newer models reject some containers whisper-1 accepted
# (notably Ogg/Opus voice notes): transcode to m4a, retry once.
# Newer models reject containers whisper-1 accepted (Ogg/Opus voice notes): transcode, retry once.
converted_path, transcode_error = _transcode_audio_for_stt(file_path, work_dir)
if transcode_error:
return _error_result(transcode_error)
@@ -255,11 +242,10 @@ def _post_audio_multipart(url: str, headers: Dict[str, str], file_path: str, dat
def _rest_transcript(response, label: str, extract_detail, extract_text):
"""Turn a multipart STT response into ``(text, body, None)`` or ``(None, None, error_envelope)``.
"""Multipart STT response -> ``(text, body, None)`` or ``(None, None, error_envelope)``.
Non-200 -> ``"<label> API error (HTTP n): detail"`` (JSON detail via
*extract_detail*, else the first 300 body chars); empty text -> the
``no_speech`` envelope so callers treat silence as non-fatal.
Non-200 -> ``"<label> API error (HTTP n): detail"`` (JSON detail via *extract_detail*, else
the first 300 body chars); empty text -> the ``no_speech`` envelope (silence is non-fatal).
"""
if response.status_code != 200:
try:
@@ -299,9 +285,8 @@ def _transcribe_xai(
if prompt:
_log_prompt_unsupported("STT provider 'xai'")
# 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.
# STT is API-billed: prefer the explicit XAI_API_KEY over the xAI OAuth/Grok-subscription
# credential, which may be valid for Grok yet hit spending-limit errors on /v1/stt.
direct_api_key = str(get_env_value("XAI_API_KEY") or "").strip()
if direct_api_key:
creds = {
@@ -319,8 +304,7 @@ def _transcribe_xai(
xai_config = stt_config.get("xai") or {}
def _resolve_base_url(resolved_creds: Dict[str, str]) -> str:
# OAuth bearers are pinned to the resolver-validated xAI origin;
# config/env base URL overrides only apply to API-key credentials.
# OAuth bearers are pinned to the resolver-validated origin; overrides apply to API keys only.
if resolved_creds.get("provider") == "xai-oauth":
url = resolved_creds.get("base_url")
else:
@@ -454,12 +438,8 @@ def _transcribe_deepinfra(
def _is_local_or_private_url(url: str) -> bool:
"""True for loopback/RFC-1918/LAN-internal hosts.
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``.
"""
"""True for loopback/RFC-1918/LAN-internal hosts, where an empty ``stt.openai.api_key`` is acceptable
(local OpenAI-compatible servers ignore the auth header — no sham ``api_key: not-needed`` needed)."""
try:
from urllib.parse import urlparse
import ipaddress
@@ -476,11 +456,8 @@ def _is_local_or_private_url(url: str) -> bool:
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.
"""
"""Direct-credential ladder: config key > keyless local base_url (placeholder key so the SDK
constructs a client) > env key; None if none apply."""
from tools.transcription_tools import resolve_openai_audio_api_key
if cfg_api_key:
return cfg_api_key, (cfg_base_url or OPENAI_BASE_URL)
@@ -493,16 +470,10 @@ def _direct_openai_credentials(cfg_api_key: str, cfg_base_url: str) -> Optional[
def _resolve_openai_audio_client_config() -> tuple[str, str]:
"""Return ``(api_key, base_url)`` for the OpenAI STT client.
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.
"""
"""``(api_key, base_url)`` for the OpenAI STT client, strict on the stored ``stt`` selection:
``"nous"`` -> managed gateway ONLY (a direct OPENAI_API_KEY must NOT override it); any other
stored provider -> direct credentials ONLY (no silent managed fallback); never-configured ->
legacy ladder: direct credentials, then the managed gateway. Failures raise ValueError."""
from tools.transcription_tools import (
_load_stt_config, managed_nous_tools_enabled, nous_tool_gateway_unavailable_message,
resolve_managed_tool_gateway,
+12 -22
View File
@@ -154,11 +154,10 @@ def _dispatch_to_plugin_provider(
) -> Optional[Dict[str, Any]]:
"""Route to a plugin-registered transcription provider; None when no plugin claims the name.
Invariants re-verified here (a caller refactor can't silently break them):
built-in names never reach the registry; a same-name ``stt.providers.<name>:
type: command`` wins over a plugin. A matched plugin with ``is_available() ==
False`` returns an error envelope — not None — because the user explicitly
opted in via ``stt.provider``. Provider exceptions become the error envelope.
Invariants re-verified here so a caller refactor can't break them: built-in names never
reach the registry; a same-name command provider wins over a plugin. A matched plugin with
``is_available() == False`` returns an error envelope — not None — because the user
explicitly opted in via ``stt.provider``. Provider exceptions become the error envelope.
"""
if not provider:
return None
@@ -216,8 +215,7 @@ def _dispatch_to_plugin_provider(
return result
# Fields a pre_transcription hook may mutate. ``file_path`` is read-only —
# attempts to change it are logged and dropped.
# Fields a pre_transcription hook may mutate; ``file_path`` is read-only (logged and dropped).
_PRE_TRANSCRIPTION_MUTABLE_FIELDS = ("prompt", "language", "model")
# Whisper-family models only use the final ~224 tokens of the prompt; longer values
@@ -229,11 +227,8 @@ _WHISPER_PROMPT_CAPPED_PROVIDERS = frozenset({"local", "openai", "groq", "deepin
def _enforce_prompt_length_limit(prompt: Optional[str], provider: str) -> Optional[str]:
"""Truncate *prompt* to the whisper-family token cap, keeping the TAIL (fail-open).
Whisper conditions on the final context window, so the most recently
appended hints survive. Other providers own their own validation.
"""
"""Truncate *prompt* to the whisper-family token cap, keeping the TAIL (whisper conditions
on the final context window, so the newest hints survive). Other providers self-validate."""
if not prompt or provider not in _WHISPER_PROMPT_CAPPED_PROVIDERS:
return prompt
max_chars = _WHISPER_PROMPT_TOKEN_CAP * _PROMPT_CHARS_PER_TOKEN
@@ -251,17 +246,12 @@ def _apply_pre_transcription_hook(
*, file_path: str, provider: str, model: Optional[str], language: Optional[str],
prompt: Optional[str], source: Optional[str],
) -> tuple[Optional[str], Optional[str], Optional[str]]:
"""Fire the ``pre_transcription`` plugin hook and merge its results.
"""Fire the ``pre_transcription`` plugin hook; returns ``(model, language_override, prompt)``.
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 and are applied field-by-field, so the last
hook to write a field wins. Model values flow through the same per-backend
normalization a caller-supplied model would.
Returns ``(model, language_override, prompt)``; ``language_override`` is
None unless a hook explicitly set ``language``, so backends keep their own
config/env language resolution.
Gated on ``has_hook`` (the no-hook path never builds kwargs) and fail-open: any
plumbing error leaves the dispatch untouched. Results apply field-by-field in
registration order (last hook wins). ``language_override`` is None unless a hook
explicitly set ``language``, so backends keep their own config/env resolution.
"""
try:
from hermes_cli.plugins import has_hook, invoke_hook
+66 -101
View File
@@ -1,15 +1,13 @@
#!/usr/bin/env python3
"""Speech-to-text transcription used by the gateway for voice messages.
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.
result = transcribe_audio("/path/to/audio.ogg") # {"success", "transcript", "error"?, "provider"?}
This module owns provider resolution, the dispatcher, and the cached local model +
idle-unload state. Backends live in ``transcription_{common,audio,local,cloud,command}``
and are re-imported here so ``tools.transcription_tools.<name>`` stays the patch surface.
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. ``transcribe_audio(path)`` returns
``{"success", "transcript", "error"?, "provider"?}``. This module owns provider resolution,
the dispatcher and the cached local model + idle-unload state; backends live in
``transcription_{common,audio,local,cloud,command}`` and are re-imported here so
``tools.transcription_tools.<name>`` stays the patch surface.
"""
import logging
@@ -68,11 +66,7 @@ logger = logging.getLogger(__name__)
def get_env_value(name, default=None):
"""Read env values through the live config module.
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.
"""
"""Read env values through the live config module (resolved per call: tests monkeypatch it around import)."""
try:
from hermes_cli.config import get_env_value as _get_env_value
except ImportError:
@@ -82,9 +76,9 @@ def get_env_value(name, default=None):
def _resolve_provider_key(env_var: str, provider_id: str) -> str:
"""Resolve an STT API key via the shared voice-key resolver (config > env/.env > credential pool).
"""STT API key via the shared voice-key resolver (config > env/.env > credential pool).
Resolved at call time so tests that reload the helpers module see the live function.
Resolved per call so tests that reload the helpers module see the live function.
"""
try:
from tools.tool_backend_helpers import resolve_provider_secret
@@ -106,16 +100,14 @@ _HAS_MISTRAL = _safe_find_spec("mistralai")
_HAS_PILK = _safe_find_spec("pilk")
# 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 load.
# Local model singleton; the lock guards check-then-load against concurrent voice messages.
_local_model: Optional[object] = None
_local_model_name: Optional[str] = None
_local_model_lock = threading.Lock()
# 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 concurrent transcriptions can't spawn duplicates.
# Idle unload: one daemon thread 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 no duplicate watchers spawn.
_last_transcription_time: float = 0.0
_idle_unload_thread: Optional[threading.Thread] = None
_idle_unload_stop = threading.Event()
@@ -146,11 +138,10 @@ def is_stt_enabled(stt_config: Optional[dict] = None) -> bool:
def _resolve_stt_language(
provider_key: str, stt_config: Optional[Dict[str, Any]] = None, *, extra_keys: tuple = ()
) -> Optional[str]:
"""Resolve the language hint for an STT provider; first non-empty wins.
"""Language hint for an STT provider, first non-empty wins (never "").
Order: ``stt.<provider>.language`` (plus *extra_keys* aliases, e.g. ElevenLabs'
``language_code``) > ``stt.language`` > ``HERMES_LOCAL_STT_LANGUAGE`` env >
None (provider auto-detects). Never returns "".
``stt.<provider>.language`` (plus *extra_keys* aliases, e.g. ``language_code``) >
``stt.language`` > ``HERMES_LOCAL_STT_LANGUAGE`` env > None (provider auto-detects).
"""
if stt_config is None:
stt_config = _load_stt_config()
@@ -189,9 +180,7 @@ def _is_local_stt_provider(provider: str, stt_config: Dict[str, Any]) -> bool:
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:
if (needs_openai and not _HAS_OPENAI) or (needs_mistral and not _HAS_MISTRAL):
return False
return bool(_resolve_provider_key(env_var, provider))
return probe
@@ -208,8 +197,7 @@ 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.
# Resolved directly so a managed openai-audio gateway outage is logged with its real reason.
reason = _openai_audio_unavailable_reason()
if reason is None:
return "openai"
@@ -232,7 +220,10 @@ def _resolve_explicit_local() -> str:
backend = _detect_local_backend()
if backend:
return backend
logger.warning("STT provider 'local' configured but unavailable (install faster-whisper or set HERMES_LOCAL_STT_COMMAND)")
logger.warning(
"STT provider 'local' configured but unavailable "
"(install faster-whisper or set HERMES_LOCAL_STT_COMMAND)"
)
return "none"
@@ -253,11 +244,10 @@ _has_deepinfra_key = _has_key("DEEPINFRA_API_KEY", "deepinfra", needs_openai=Tru
# Cloud providers in AUTO-DETECT priority order:
# name -> (explicit-selection probe, auto-detect probe, explicit warning, auto-detect log)
# The two probes differ only for openai (auto-detect additionally requires the SDK;
# explicit has its own resolver that logs the real gateway reason) and xai
# (auto-detect must never raise). 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
# The probes differ only for openai (explicit has its own resolver in _EXPLICIT_RESOLVERS;
# auto-detect also requires the SDK) and xai (auto-detect must never raise). DeepInfra is
# LAST so a DEEPINFRA_API_KEY set for chat never displaces an xAI/ElevenLabs auto-selection.
# Mistral only auto-selects when the SDK is present — no lazy-install during passive
# auto-detection (explicit ``provider: mistral`` installs on first use).
_CLOUD_PROVIDER_SPECS = {
"groq": (
@@ -301,11 +291,8 @@ _EXPLICIT_RESOLVERS = {
def _resolve_explicit_provider(provider: str) -> str:
"""Resolve an explicit ``stt.provider`` to a usable provider name or ``"none"``.
Unknown names pass through untouched so the dispatcher can fail with the
provider-not-registered message.
"""
"""Explicit ``stt.provider`` -> usable name or ``"none"``; unknown names pass through untouched
so the dispatcher fails with the provider-not-registered message."""
resolver = _EXPLICIT_RESOLVERS.get(provider)
if resolver is not None:
return resolver()
@@ -320,27 +307,21 @@ def _resolve_explicit_provider(provider: str) -> str:
def _get_provider(stt_config: dict) -> str:
"""Determine which STT provider to use.
An explicit ``stt.provider`` is honoured — no silent cloud fallback. With no
provider configured, auto-detect tries local > groq > openai > mistral > xai
> elevenlabs > deepinfra.
"""
"""Which STT provider to use: an explicit ``stt.provider`` is honoured (no silent cloud
fallback); otherwise auto-detect local > groq > openai > mistral > xai > elevenlabs > deepinfra."""
if not is_stt_enabled(stt_config):
return "none"
explicit = "provider" in stt_config
provider = stt_config.get("provider", DEFAULT_PROVIDER)
# The managed "Nous Subscription" selection is serviced by the OpenAI implementation,
# routed through the managed gateway by _resolve_openai_audio_client_config.
# The managed "Nous Subscription" selection is the OpenAI backend routed via the managed gateway.
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. Only a raw config.yaml
# selection counts as explicit; otherwise autodetect (which prefers local anyway).
# Legacy DEFAULT_CONFIG seeded ``stt.provider: local`` on every install, so only a
# raw config.yaml selection counts as explicit; otherwise autodetect (local-first anyway).
try:
from tools.tool_backend_helpers import read_selection
@@ -378,12 +359,10 @@ def _unload_local_model() -> None:
def _start_idle_unload_watcher(timeout_seconds: int) -> None:
"""Ensure the single idle-unload watcher thread is running.
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.
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. The thread exits after unloading,
when the timeout becomes 0, or when the model is already gone.
"""
global _idle_unload_thread
with _idle_unload_mgmt_lock:
@@ -419,11 +398,9 @@ def _touch_transcription_time() -> None:
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.
"""Cached faster-whisper model, (re)loaded under a double-checked 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.
The returned strong reference stays valid even if the idle watcher nulls the global mid-transcription.
"""
global _local_model, _local_model_name
model = _local_model
@@ -431,8 +408,7 @@ def _get_or_load_local_model(model_name: str, local_cfg: Dict[str, Any]):
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.
# stt.local.device / compute_type pin a configuration where ``auto`` mis-detects.
_local_model = _load_local_whisper_model(
model_name, device=local_cfg.get("device", "auto"),
compute_type=local_cfg.get("compute_type", "auto"),
@@ -463,8 +439,7 @@ def _transcribe_local(
try:
stt_config = _load_stt_config()
local_cfg = stt_config.get("local") or {}
# 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.
# Reset the idle timer BEFORE loading so a long in-flight transcription isn't counted as idle.
_touch_transcription_time()
model = _get_or_load_local_model(model_name, local_cfg)
if model is None: # defensive: load failed without raising
@@ -480,9 +455,8 @@ def _transcribe_local(
try:
segments, info = model.transcribe(file_path, **transcribe_kwargs)
except Exception as exc:
# 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.
# CUDA libs can fail at dlopen-on-first-use, AFTER loading: evict the poisoned
# cached model, reload on CPU and retry once, else every later message fails.
if not _looks_like_cuda_lib_error(exc):
raise
logger.warning(
@@ -515,8 +489,8 @@ def _transcribe_local(
def _read_block_error(file_path: str) -> Optional[Dict[str, Any]]:
"""Refuse to ship a credential / secret store (auth.json, .env, OAuth tokens) to an STT
provider in plaintext. Mirrors the image-gen / video-gen read guards."""
"""Refuse to ship a credential store (auth.json, .env, OAuth tokens) to an STT provider.
Mirrors the image-gen / video-gen read guards."""
from agent.file_safety import get_read_block_error
blocked = get_read_block_error(file_path)
return _error_result(blocked) if blocked else None
@@ -527,16 +501,15 @@ def _transcribe_prepared_audio(
) -> Dict[str, Any]:
"""Transcribe a validated audio file with the configured STT provider.
``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.
``model`` overrides the config default; ``source`` is a caller-surface label
(``"gateway"``, ``"voice_mode"``) forwarded to the ``pre_transcription`` hook only.
"""
blocked = _read_block_error(file_path)
if blocked:
return blocked
# Validate before provider resolution so invalid files cannot trigger provider
# setup or lazy installation. The remote-upload size cap applies to non-local only.
# Validate before provider resolution so invalid files can't trigger provider setup
# or lazy installation; the remote-upload size cap applies to non-local only.
error = _validate_audio_file(file_path, enforce_size_limit=False)
if error:
return error
@@ -572,9 +545,8 @@ def _transcribe_prepared_audio(
shutil.rmtree(trim_cleanup_dir, ignore_errors=True)
# Built-in provider -> (stt section, config key, default, treat-empty-as-missing).
# "local_command" shares the ``stt.local`` section; xAI takes no model parameter
# (the name is logging-only); deepinfra resolves from the live catalog when empty.
# Built-in provider -> (stt section, config key, default, treat-empty-as-missing). "local_command"
# shares ``stt.local``; xAI takes no model (logging-only); deepinfra uses the live catalog when empty.
_BUILTIN_MODEL_KEYS = {
"local": ("local", "model", DEFAULT_LOCAL_MODEL, False),
"local_command": ("local", "model", DEFAULT_LOCAL_MODEL, False),
@@ -604,14 +576,12 @@ def _dispatch_stt_provider(
source: Optional[str] = None,
) -> Dict[str, Any]:
"""Route *file_path* to the handler for *provider* (built-in > command > plugin)."""
# Static ``stt.prompt`` is the base; pre_transcription hook results mutate
# on top in registration order (last hook to set a field wins).
# Static ``stt.prompt`` is the base; hook results mutate on top (last hook to set a field wins).
prompt = stt_config.get("prompt")
if not isinstance(prompt, str) or not prompt.strip():
prompt = None
# The hook fires after provider resolution and BEFORE any backend is
# invoked; ``language`` stays None unless a hook overrides it.
# Fires after provider resolution and BEFORE any backend; ``language`` stays None unless a hook sets it.
model, language, prompt = _apply_pre_transcription_hook(
file_path=file_path, provider=provider, model=model,
language=_get_stt_section(stt_config, provider).get("language"),
@@ -627,9 +597,8 @@ def _dispatch_stt_provider(
model_name = _normalize_local_model(model_name)
return handler(file_path, model_name, language=language, prompt=prompt)
# 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 providers: after built-ins (``stt.providers.openai.command`` can't override the
# real handler) and BEFORE plugins, since config is more local than a plugin install.
command_provider_config = _resolve_command_stt_provider_config(provider, stt_config)
if command_provider_config is not None:
return _transcribe_command_stt(
@@ -637,8 +606,7 @@ def _dispatch_stt_provider(
model_override=model, language_override=language, prompt=prompt,
)
# Plugin-registered backend. Plugins read per-provider config under
# ``stt.<provider>`` like built-ins; the ``model`` argument overrides it.
# Plugin backend: reads ``stt.<provider>`` 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=model or plugin_cfg.get("model"),
@@ -655,8 +623,7 @@ def _no_provider_error(provider: str, stt_config: Dict[str, Any]) -> Dict[str, A
if "provider" in stt_config and provider_key and provider_key not in BUILTIN_STT_PROVIDERS and provider_key != "none":
return _unregistered_stt_provider_error(provider_key)
# An explicit openai selection flattened to "none" carries a selection-specific
# reason (e.g. managed openai-audio gateway down); surface it with its remediation.
# An explicit openai selection flattened to "none" has a specific reason (e.g. managed gateway down).
if provider_key == "none" and str(stt_config.get("provider") or "") == "openai" and _HAS_OPENAI:
reason = _openai_audio_unavailable_reason()
if reason is not None:
@@ -675,20 +642,18 @@ def _no_provider_error(provider: str, stt_config: Dict[str, Any]) -> Dict[str, A
def transcribe_audio(
file_path: str, model: Optional[str] = None, source: Optional[str] = None
) -> Dict[str, Any]:
"""Safely validate, preprocess supported inputs, and dispatch transcription.
"""Validate, preprocess supported inputs, and dispatch transcription.
``source`` is an optional caller-surface label (``"gateway"``, ``"voice_mode"``)
forwarded to the ``pre_transcription`` hook for observability only.
``source`` is a caller-surface label (``"gateway"``, ``"voice_mode"``) forwarded to
the ``pre_transcription`` hook for observability only.
"""
# Secret-store refusal runs before ANY validation so the error names the
# real reason rather than a format error.
# Secret-store refusal runs before ANY validation so the error names the real reason.
blocked = _read_block_error(file_path)
if blocked:
return blocked
# 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.
# Cap .silk sources before the decoder runs; for other inputs the upload cap is
# provider-scoped in _transcribe_prepared_audio so local whisper can take 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:
@@ -713,8 +678,8 @@ def transcribe_audio_local_fallback(
) -> Dict[str, Any]:
"""Try an already-installed local STT backend without changing config.
For passive inbound-media recovery after the configured provider failed:
never lazy-installs or falls through to a cloud provider.
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:
+11 -5
View File
@@ -147,7 +147,9 @@ def command_failure_detail(exc: subprocess.CalledProcessError) -> str:
return "; ".join(parts) or "no command output"
def run_command_provider(command: str, timeout: float, env_passthrough: Optional[list] = None) -> subprocess.CompletedProcess:
def run_command_provider(
command: str, timeout: float, env_passthrough: Optional[list] = None,
) -> subprocess.CompletedProcess:
"""Run a command-provider shell command with process-tree idle cleanup.
``timeout`` is an IDLE timeout, reset whenever the command emits output — a
@@ -280,8 +282,10 @@ def _is_command_provider_config(config: Dict[str, Any]) -> bool:
return isinstance(command, str) and bool(command.strip())
def _resolve_command_config(provider: str, config: Dict[str, Any], reserved: FrozenSet[str]) -> Optional[Dict[str, Any]]:
"""Provider config when *provider* is a user-declared command provider; None for *reserved* names, unknown or non-command."""
def _resolve_command_config(
provider: str, config: Dict[str, Any], reserved: FrozenSet[str],
) -> Optional[Dict[str, Any]]:
"""Config of a user-declared command provider; None for *reserved* names, unknown or non-command."""
if not provider:
return None
key = provider.lower().strip()
@@ -292,7 +296,7 @@ def _resolve_command_config(provider: str, config: Dict[str, Any], reserved: Fro
def _command_timeout(config: Dict[str, Any], default: float) -> float:
"""Timeout in seconds (``timeout`` > ``timeout_seconds``); invalid or non-positive values fall back to *default*."""
"""Timeout in seconds (``timeout`` > ``timeout_seconds``); invalid or non-positive -> *default*."""
raw = config.get("timeout", config.get("timeout_seconds", default))
try:
value = float(raw)
@@ -363,7 +367,9 @@ def _configured_command_tts_output_path(path: Path, config: Dict[str, Any]) -> P
return path.with_suffix(f".{_get_command_tts_output_format(config)}")
def _generate_command_tts(text: str, output_path: str, provider_name: str, config: Dict[str, Any], tts_config: Dict[str, Any]) -> str:
def _generate_command_tts(
text: str, output_path: str, provider_name: str, config: Dict[str, Any], tts_config: Dict[str, Any],
) -> str:
"""Generate speech by running a user-configured shell command; returns the audio path it wrote.
Raises ``ValueError`` for invalid provider config and ``RuntimeError`` for
+9 -7
View File
@@ -21,8 +21,7 @@ from tools.tts_tool import _get_provider, _load_tts_config, get_env_value
logger = logging.getLogger(__name__)
# Per-sentence PCM byte cap, mirroring the 16 MiB bounded-body invariant of the
# sync providers: a buggy or hostile endpoint must not feed unbounded audio.
# Per-sentence PCM byte cap, mirroring the sync providers' 16 MiB bounded-body invariant.
_STREAM_SENTENCE_BYTE_CAP = 16 * 1024 * 1024
@@ -147,9 +146,8 @@ def _try_instantiate(name: str, tts_config: Dict) -> Optional[StreamingTTSProvid
return None
# Fallback priority for ``tts.streaming.provider: auto`` — best chunked
# latency/quality first. Deliberately hard-coded (a UX decision, not a config
# knob); edge is absent because it has no chunked-PCM API.
# Fallback priority for ``tts.streaming.provider: auto`` — best chunked latency/quality
# first. Deliberately hard-coded (a UX decision); edge is absent (no chunked-PCM API).
_PROVIDER_PRIORITY: List[str] = ["elevenlabs", "gemini", "openai", "xai"]
@@ -262,7 +260,9 @@ class GeminiStreamer(StreamingTTSProvider):
import requests
from tools.tts_tool_providers import DEFAULT_GEMINI_TTS_BASE_URL, DEFAULT_GEMINI_TTS_MODEL, DEFAULT_GEMINI_TTS_VOICE
from tools.tts_tool_providers import (
DEFAULT_GEMINI_TTS_BASE_URL, DEFAULT_GEMINI_TTS_MODEL, DEFAULT_GEMINI_TTS_VOICE,
)
api_key = _gemini_key()
model = str(self.section.get("model", DEFAULT_GEMINI_TTS_MODEL)).strip() or DEFAULT_GEMINI_TTS_MODEL
@@ -280,7 +280,9 @@ class GeminiStreamer(StreamingTTSProvider):
url = f"{base_url}/models/{model}:streamGenerateContent"
def _sse_chunks() -> Iterator[bytes]:
with requests.post(url, params={"alt": "sse", "key": api_key}, json=payload, timeout=60, stream=True) as response:
with requests.post(
url, params={"alt": "sse", "key": api_key}, json=payload, timeout=60, stream=True,
) as response:
response.raise_for_status()
for line in response.iter_lines(decode_unicode=True):
if not line or not line.startswith("data: "):
+8 -3
View File
@@ -81,10 +81,13 @@ def strip_markdown_for_tts(text: str) -> str:
def _normalize_temperature_ranges(text: str) -> str:
"""``11-17°C`` -> ``11 to 17 degrees Celsius`` (en/em dash or hyphen; unicode minus normalized)."""
number = r"([-+\u2212]?\d+(?:\.\d+)?)"
for unit, word in (("C", "Celsius"), ("F", "Fahrenheit")):
text = re.sub(
r"(?<!\w)([-+\u2212]?\d+(?:\.\d+)?)\s*[\u2013\u2014-]\s*([-+\u2212]?\d+(?:\.\d+)?)\s*°\s*" + unit + r"\b",
lambda m, w=word: f"{m.group(1).replace(chr(0x2212), '-')} to {m.group(2).replace(chr(0x2212), '-')} degrees {w}",
r"(?<!\w)" + number + r"\s*[\u2013\u2014-]\s*" + number + r"\s*°\s*" + unit + r"\b",
lambda m, w=word: (
f"{m.group(1).replace(chr(0x2212), '-')} to {m.group(2).replace(chr(0x2212), '-')} degrees {w}"
),
text,
flags=re.IGNORECASE,
)
@@ -105,7 +108,9 @@ def normalize_symbols_for_tts(text: str) -> str:
# Temperatures with a number first, then bare units ("measured in degrees C"),
# then any remaining degree symbol (angles, stray cases).
for unit, word in (("C", "Celsius"), ("F", "Fahrenheit")):
text = re.sub(r"(?<!\w)([-+]?\d+(?:\.\d+)?)\s*°\s*" + unit + r"\b", r"\1 degrees " + word, text, flags=re.IGNORECASE)
text = re.sub(
r"(?<!\w)([-+]?\d+(?:\.\d+)?)\s*°\s*" + unit + r"\b", r"\1 degrees " + word, text, flags=re.IGNORECASE,
)
for unit, word in (("C", "Celsius"), ("F", "Fahrenheit")):
text = re.sub(r"°\s*" + unit + r"\b", "degrees " + word, text, flags=re.IGNORECASE)
text = re.sub(r"(?<!\w)([-+]?\d+(?:\.\d+)?)\s*°", r"\1 degrees", text)