refactor(tools): compact STT/TTS docstrings and comments, keep every invariant
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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: "):
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user