refactor(tools): unify STT/TTS command-provider config helpers, REST STT flow, dispatch tables

This commit is contained in:
Teknium
2026-09-02 22:41:57 -07:00
parent 113f04616b
commit 4d37eddc2f
9 changed files with 476 additions and 831 deletions

View File

@@ -2,10 +2,9 @@
Binary discovery, the shared ffmpeg m4a encode (transcode + silence trim),
source/format validation, WeChat .silk decoding, CAF conversion and the
best-effort cloud pre-upload silence trim.
Split out of ``tools/transcription_tools.py``, which re-imports every name (patch
surface) and is imported lazily here so origin patches still intercept.
best-effort cloud pre-upload silence trim. Every name is re-imported by
``tools/transcription_tools.py`` (patch surface), which is imported lazily here
so origin patches still intercept.
"""
from __future__ import annotations
@@ -61,17 +60,11 @@ def _run_quiet(command: list, *, timeout: float, env: Optional[dict] = None) ->
# 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",
)
_STT_M4A_ENCODE_ARGS = ("-vn", "-ac", "1", "-ar", "16000", "-c:a", "aac", "-b:a", "32k", "-movflags", "+faststart")
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 error semantics (transcode reports, trim swallows).
"""
"""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]
@@ -80,11 +73,10 @@ def _run_ffmpeg_stt_encode(ffmpeg: str, input_path: str, output_path: str, *, au
def _transcode_audio_for_stt(file_path: str, work_dir: str) -> tuple[Optional[str], Optional[str]]:
"""Transcode to a compact 16 kHz mono AAC/m4a for STT upload.
"""Transcode to a compact 16 kHz mono AAC/m4a for STT upload; ``(converted_path, None)`` 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)``.
"""
from tools.transcription_tools import _find_ffmpeg_binary, _run_ffmpeg_stt_encode
ffmpeg = _find_ffmpeg_binary()
@@ -110,16 +102,13 @@ def _validate_audio_file_size(audio_path: Path, *, enforce_size_limit: bool = Tr
except OSError as e:
return _error_result(f"Failed to access file: {e}")
if enforce_size_limit and file_size > MAX_FILE_SIZE:
return _error_result(
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
def _validate_audio_source_file(file_path: str, *, enforce_size_limit: bool = True) -> Optional[Dict[str, Any]]:
"""Validate source path safety (and optionally size) before any decoder runs."""
audio_path = Path(file_path)
if os.path.islink(audio_path):
return _error_result(f"Path is a symbolic link: {file_path}")
if not audio_path.exists():
@@ -137,9 +126,7 @@ def _validate_audio_file(file_path: str, *, enforce_size_limit: bool = True) ->
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 _error_result(f"Unsupported format: {suffix}. Supported: {', '.join(sorted(SUPPORTED_FORMATS))}")
return None
@@ -150,8 +137,7 @@ def _prepare_audio_for_transcription(file_path: str) -> tuple[Optional[str], Opt
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 on first .silk
# voice note instead of bloating the base install.
# pilk is a tiny silk-v3 codec binding — lazy-installed on first .silk voice note.
_lazy_ensure_quietly("stt.silk")
if not _safe_find_spec("pilk"):
return None, None, _error_result(
@@ -221,23 +207,22 @@ def _convert_caf_to_wav(file_path: str) -> Optional[str]:
# ---- Cloud pre-upload silence trim --------------------------------------
#
# 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.
# Local faster-whisper gets Silero VAD; cloud providers get the raw file, so 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:
# ``stt.cloud_trim_silence: false``, ffmpeg/ffprobe missing, trim failure/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 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.
# Below this the trim can't pay for itself (several providers bill a 10s minimum)
# and the encode would sit on the synchronous voice-note path.
_CLOUD_TRIM_MIN_INPUT_SECONDS = 12.0
@@ -251,12 +236,7 @@ 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
@@ -274,11 +254,9 @@ def _cloud_trim_settings(stt_config: Dict[str, Any]) -> tuple[bool, int, int]:
def _trim_silence_for_cloud_stt(file_path: str, stt_config: Dict[str, Any]) -> Optional[str]:
"""Return a silence-trimmed copy of *file_path* for cloud upload, or None.
"""Return a silence-trimmed copy of *file_path* for cloud upload, or None (= upload the original).
``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.
On success the caller owns deleting the returned file's parent directory.
"""
from tools.transcription_tools import _find_ffmpeg_binary, _probe_audio_duration, _run_ffmpeg_stt_encode
enabled, threshold_db, keep_ms = _cloud_trim_settings(stt_config)
@@ -301,8 +279,7 @@ def _trim_silence_for_cloud_stt(file_path: str, stt_config: Dict[str, Any]) -> O
return None
keep_seconds = keep_ms / 1000.0
# start_periods=1 strips leading silence; stop_periods=-1 collapses every
# interior/trailing silence, keeping ``keep_seconds`` of each pause.
# start_periods=1 strips leading silence; stop_periods=-1 collapses every interior/trailing silence.
filter_expr = (
f"silenceremove="
f"start_periods=1:start_threshold={threshold_db}dB:start_silence={keep_seconds}:"
@@ -310,8 +287,7 @@ def _trim_silence_for_cloud_stt(file_path: str, stt_config: Dict[str, Any]) -> O
)
work_dir = tempfile.mkdtemp(prefix="hermes-stt-trim-")
trimmed_path = os.path.join(work_dir, f"{Path(file_path).stem or 'audio'}-trimmed.m4a")
# Scale the all-silence guard with keep_ms: an output consisting solely
# of kept pause must never be uploaded as "speech".
# Scale the all-silence guard with keep_ms: output that is solely kept pause must never upload as "speech".
min_result_seconds = max(_CLOUD_TRIM_MIN_RESULT_SECONDS, 2 * keep_seconds)
keep_result = False
try:

View File

@@ -2,10 +2,9 @@
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).
Split out of ``tools/transcription_tools.py``, which re-imports every name (patch
surface) and is imported lazily here so origin patches still intercept.
(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.
"""
from __future__ import annotations
@@ -14,7 +13,7 @@ import logging
import re
import tempfile
from pathlib import Path
from typing import Any, Dict, Optional
from typing import Any, Callable, Dict, Optional
from urllib.parse import urljoin
from utils import is_truthy_value
@@ -35,12 +34,6 @@ def _has_xai_stt_credentials() -> bool:
return bool(resolve_xai_http_credentials().get("api_key"))
def _close_client(client: Any) -> None:
close = getattr(client, "close", None)
if callable(close):
close()
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).
@@ -53,7 +46,9 @@ def _with_openai_client(api_key: str, base_url: Optional[str], file_path: str, l
try:
return body(client)
finally:
_close_client(client)
close = getattr(client, "close", None)
if callable(close):
close()
except Exception as e:
return _openai_sdk_failure(e, file_path, log_label)
@@ -78,16 +73,23 @@ def _openai_sdk_failure(exc: BaseException, file_path: str, log_label: str) -> D
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}")
for cls, label in ((APIConnectionError, "Connection error"), (APITimeoutError, "Request timeout"), (APIError, "API error")):
if isinstance(exc, cls):
return _error_result(f"{label}: {exc}")
logger.error("%s transcription failed: %s", log_label, exc, exc_info=True)
return _error_result(f"Transcription failed: {exc}")
def _sdk_prompt_kwargs(language: Optional[str], prompt: Optional[str]) -> Dict[str, Any]:
"""``language``/``prompt`` create-kwargs, each only when set so the bare request stays byte-identical."""
kwargs: Dict[str, Any] = {}
if language:
kwargs["language"] = language
if prompt:
kwargs["prompt"] = prompt
return kwargs
def _transcribe_groq(
file_path: str,
model_name: str,
@@ -104,7 +106,6 @@ def _transcribe_groq(
api_key = _resolve_provider_key("GROQ_API_KEY", "groq")
if not api_key:
return _error_result("GROQ_API_KEY not set")
if not _HAS_OPENAI:
return _error_result("openai package not installed")
@@ -116,25 +117,12 @@ def _transcribe_groq(
language = language or _resolve_stt_language("groq")
def _run(client):
create_kwargs = {
"model": model_name,
"response_format": "text",
}
if language:
create_kwargs["language"] = language
if prompt:
# Only sent when set so the no-hook, no-config request stays byte-identical.
create_kwargs["prompt"] = prompt
create_kwargs = {"model": model_name, "response_format": "text", **_sdk_prompt_kwargs(language, prompt)}
with open(file_path, "rb") as audio_file:
transcription = client.audio.transcriptions.create(
file=audio_file,
**create_kwargs,
)
transcription = client.audio.transcriptions.create(file=audio_file, **create_kwargs)
transcript_text = str(transcription).strip()
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 _ok_result(transcript_text, "groq")
return _with_openai_client(api_key, GROQ_BASE_URL, file_path, "Groq", _run)
@@ -222,7 +210,6 @@ def _transcribe_openai(
"Transcribed %s via %s (%s, %d chars)",
Path(file_path).name, provider_label, model_name, len(transcript_text),
)
return _ok_result(transcript_text, provider_label)
return _with_openai_client(api_key, base_url, file_path, provider_label, _run)
@@ -247,17 +234,13 @@ def _transcribe_mistral(
with Mistral(api_key=api_key) as client:
with open(file_path, "rb") as audio_file:
# Language: hook override > stt.mistral.language > stt.language > env > auto.
language = language or _resolve_stt_language("mistral")
complete_kwargs: Dict[str, Any] = {
"model": model_name,
"file": {"content": audio_file, "file_name": Path(file_path).name},
**_sdk_prompt_kwargs(language, prompt),
}
# 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 sent when set so the no-hook, no-config request stays byte-identical.
complete_kwargs["prompt"] = prompt
result = client.audio.transcriptions.complete(**complete_kwargs)
transcript_text = _extract_transcript_text(result)
@@ -271,6 +254,9 @@ def _transcribe_mistral(
return _cloud_failure(e, file_path, "Mistral transcription", type(e).__name__)
# ---- REST multipart backends (xAI, ElevenLabs) ----------------------------
def _post_audio_multipart(url: str, headers: Dict[str, str], file_path: str, data: Dict[str, str]):
import requests
@@ -281,22 +267,18 @@ def _post_audio_multipart(url: str, headers: Dict[str, str], file_path: str, dat
)
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 _rest_transcript(response, label: str, extract_detail, extract_text):
"""Turn a multipart STT response into ``(text, body, None)`` or ``(None, None, error_envelope)``.
Non-200 -> ``"<label> API error (HTTP n): detail"``; empty text -> the
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.
"""
if response.status_code != 200:
detail = _http_error_detail(response, extract_detail)
try:
detail = extract_detail(response.json()) or response.text[:300]
except Exception:
detail = response.text[:300]
return None, None, _error_result(f"{label} API error (HTTP {response.status_code}): {detail}")
body = response.json()
text = extract_text(body)
@@ -305,11 +287,24 @@ def _rest_transcript(response, label: str, extract_detail, extract_text):
return text, body, None
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 _rest_provider(
file_path: str,
provider: str,
label: str,
post: Callable[[], Any],
extract_detail,
extract_text,
log: Callable[[str, Dict[str, Any]], None],
) -> Dict[str, Any]:
"""Shared REST flow: ``post()`` -> envelope/transcript -> ``log(text, body)`` -> ok; exceptions -> ``_cloud_failure``."""
try:
transcript_text, body, error = _rest_transcript(post(), label, extract_detail, extract_text)
if error:
return error
log(transcript_text, body)
return _ok_result(transcript_text, provider)
except Exception as e:
return _cloud_failure(e, file_path, f"{label} transcription")
def _transcribe_xai(
@@ -357,7 +352,7 @@ def _transcribe_xai(
# Language: hook override > stt.xai.language > stt.language > env.
language = language or _resolve_stt_language("xai", stt_config) or ""
try:
def _post() -> Any:
from tools.xai_http import hermes_xai_user_agent
data: Dict[str, str] = {}
@@ -375,7 +370,6 @@ def _transcribe_xai(
)
response = _post_transcription(api_key, _resolve_base_url(creds))
if response.status_code in {401, 403} and creds.get("provider") == "xai-oauth":
logger.info("xAI STT got HTTP %d; refreshing OAuth credentials and retrying once", response.status_code)
try:
@@ -387,22 +381,27 @@ def _transcribe_xai(
logger.warning(
"xAI STT OAuth refresh-and-retry after HTTP %d failed: %s", response.status_code, retry_exc,
)
return response
transcript_text, result, error = _rest_transcript(
response, "xAI STT",
lambda body: body.get("error", {}).get("message", ""),
lambda body: body.get("text", "").strip(),
)
if error:
return error
def _log(transcript_text: str, result: Dict[str, Any]) -> None:
logger.info(
"Transcribed %s via xAI Grok STT (lang=%s, %.1fs audio, %d chars)",
Path(file_path).name, result.get("language", language), result.get("duration", 0), len(transcript_text),
)
return _ok_result(transcript_text, "xai")
except Exception as e:
return _cloud_failure(e, file_path, "xAI STT transcription")
return _rest_provider(
file_path, "xai", "xAI STT", _post,
lambda body: body.get("error", {}).get("message", ""),
lambda body: body.get("text", "").strip(),
_log,
)
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_elevenlabs(
@@ -428,7 +427,8 @@ def _transcribe_elevenlabs(
).strip().rstrip("/")
# Language: hook override > stt.elevenlabs.language(_code) > stt.language.
language_code = language or _resolve_stt_language("elevenlabs", stt_config, extra_keys=("language_code",)) or ""
try:
def _post() -> Any:
data: Dict[str, str] = {
"model_id": model_name,
"tag_audio_events": str(is_truthy_value(elevenlabs_config.get("tag_audio_events", False))).lower(),
@@ -436,22 +436,17 @@ def _transcribe_elevenlabs(
}
if language_code:
data["language_code"] = language_code
return _post_audio_multipart(f"{base_url}/speech-to-text", {"xi-api-key": api_key}, file_path, data)
response = _post_audio_multipart(f"{base_url}/speech-to-text", {"xi-api-key": api_key}, file_path, data)
transcript_text, _body, error = _rest_transcript(
response, "ElevenLabs STT", _elevenlabs_error_detail, _extract_transcript_text,
)
if error:
return error
def _log(transcript_text: str, _body: Dict[str, Any]) -> None:
logger.info(
"Transcribed %s via ElevenLabs Scribe (%s, %d chars)",
Path(file_path).name, model_name, len(transcript_text),
)
return _ok_result(transcript_text, "elevenlabs")
except Exception as e:
return _cloud_failure(e, file_path, "ElevenLabs STT transcription")
return _rest_provider(
file_path, "elevenlabs", "ElevenLabs STT", _post, _elevenlabs_error_detail, _extract_transcript_text, _log,
)
def _transcribe_deepinfra(
@@ -461,8 +456,7 @@ def _transcribe_deepinfra(
language: Optional[str] = None,
prompt: Optional[str] = None,
) -> Dict[str, Any]:
"""Resolve DeepInfra credentials/model (via the shared ``hermes_cli.models``
helpers), then delegate to :func:`_transcribe_openai`."""
"""Resolve DeepInfra credentials/model (shared ``hermes_cli.models`` helpers), then delegate to :func:`_transcribe_openai`."""
from tools.transcription_tools import _load_stt_config, _resolve_provider_key, _transcribe_openai
api_key = _resolve_provider_key("DEEPINFRA_API_KEY", "deepinfra")
if not api_key:
@@ -488,6 +482,9 @@ def _transcribe_deepinfra(
)
# ---- OpenAI audio credential resolution -----------------------------------
def _is_local_or_private_url(url: str) -> bool:
"""True for loopback/RFC-1918/LAN-internal hosts.

View File

@@ -2,10 +2,9 @@
``stt.providers.<name>: type: command`` registry, plugin-registered
``TranscriptionProvider`` dispatch, and the ``pre_transcription`` hook that
threads prompt/language/model overrides into every backend.
Split out of ``tools/transcription_tools.py``, which re-imports every name (patch
surface) and is imported lazily here so origin patches still intercept.
threads prompt/language/model overrides into every backend. Every name is
re-imported by ``tools/transcription_tools.py`` (patch surface), which is
imported lazily here so origin patches still intercept.
"""
from __future__ import annotations
@@ -17,12 +16,13 @@ from pathlib import Path
from typing import Any, Dict, Optional
from tools.tts_command_provider import (
command_env_passthrough as _command_stt_env_passthrough,
render_command_template as _render_command_stt_template,
_command_output_format, _command_timeout, _is_command_provider_config as _is_command_stt_provider_config,
_named_provider_config, _resolve_command_config, command_env_passthrough as _command_stt_env_passthrough,
command_failure_detail, render_command_template as _render_command_stt_template,
run_command_provider as _run_command_stt,
)
from tools.transcription_common import (
BUILTIN_STT_PROVIDERS, _error_result, _get_stt_section, _log_prompt_unsupported, _ok_result,
BUILTIN_STT_PROVIDERS, _error_result, _log_prompt_unsupported, _ok_result,
)
# Log-record parity with the origin module.
@@ -31,86 +31,38 @@ logger = logging.getLogger("tools.transcription_tools")
# ---- Command-provider registry (``stt.providers.<name>: type: command``) ---
#
# 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.<name> 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.
# Mirrors the TTS command-provider registry (same placeholder grammar, quote-aware
# rendering, process-tree termination on timeout). Resolution order: built-in name
# (always wins) > stt.providers.<name> command > plugin 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"})
_NON_COMMAND_STT_NAMES = frozenset(BUILTIN_STT_PROVIDERS | {"none"})
def _get_named_stt_provider_config(
stt_config: Dict[str, Any],
name: str,
) -> Dict[str, Any]:
"""Return the config for a user-declared STT provider, or {}.
``stt.providers.<name>`` is canonical; ``stt.<name>`` 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(section, dict):
return section
if name.lower() not in BUILTIN_STT_PROVIDERS:
return _get_stt_section(stt_config, name)
return {}
def _get_named_stt_provider_config(stt_config: Dict[str, Any], name: str) -> Dict[str, Any]:
"""``stt.providers.<name>`` (canonical), else ``stt.<name>`` for non-built-in names only."""
return _named_provider_config(stt_config, name, BUILTIN_STT_PROVIDERS)
def _is_command_stt_provider_config(config: Dict[str, Any]) -> bool:
"""Return True when *config* declares a command-type STT provider."""
if not isinstance(config, dict):
return False
ptype = str(config.get("type") or "").strip().lower()
if ptype and ptype != "command":
return False
command = config.get("command")
return isinstance(command, str) and bool(command.strip())
def _resolve_command_stt_provider_config(
provider: str,
stt_config: Dict[str, Any],
) -> Optional[Dict[str, Any]]:
"""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)
return config if _is_command_stt_provider_config(config) else None
def _resolve_command_stt_provider_config(provider: str, stt_config: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""The provider config if *provider* is a command type; None for built-ins, ``none``, unknown."""
return _resolve_command_config(provider, stt_config, _NON_COMMAND_STT_NAMES)
def _get_command_stt_timeout(config: Dict[str, Any]) -> float:
"""Return timeout in seconds (``timeout`` > ``timeout_seconds``), falling back when invalid or <= 0."""
raw = config.get("timeout", config.get("timeout_seconds", DEFAULT_COMMAND_STT_TIMEOUT_SECONDS))
try:
value = float(raw)
except (TypeError, ValueError):
value = 0.0
return value if value > 0 else float(DEFAULT_COMMAND_STT_TIMEOUT_SECONDS)
return _command_timeout(config, 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
fmt = str(raw).lower().strip().lstrip(".")
return fmt if fmt in COMMAND_STT_OUTPUT_FORMATS else DEFAULT_COMMAND_STT_OUTPUT_FORMAT
return _command_output_format(config, COMMAND_STT_OUTPUT_FORMATS, DEFAULT_COMMAND_STT_OUTPUT_FORMAT)
def _read_command_stt_output(output_path: Path, stdout: str, fmt: str) -> str:
"""Return the transcript: non-empty output file > non-empty stdout (curl one-liners) > RuntimeError.
JSON output is returned raw — users configure ``format: txt`` or post-process.
"""
"""Transcript: non-empty output file > non-empty stdout (curl one-liners) > RuntimeError. JSON is returned raw."""
if output_path.exists():
try:
content = output_path.read_text(encoding="utf-8").strip()
@@ -120,10 +72,7 @@ def _read_command_stt_output(output_path: Path, stdout: str, fmt: str) -> str:
return content
if stdout and stdout.strip():
return stdout.strip()
raise RuntimeError(
f"Command STT provider wrote no output file at {output_path} "
f"and produced no stdout"
)
raise RuntimeError(f"Command STT provider wrote no output file at {output_path} and produced no stdout")
def _transcribe_command_stt(
@@ -137,9 +86,8 @@ def _transcribe_command_stt(
) -> Dict[str, Any]:
"""Transcribe via a user-declared ``stt.providers.<name>: type: command``.
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,
Placeholders (shell-quote-aware; ``{{``/``}}`` stay literal): ``{input_path}``,
``{output_path}`` (transcript file), ``{output_dir}``, ``{format}`` txt/json/srt/vtt,
``{language}`` (default ``en``), ``{model}`` (empty when unset).
"""
from tools.transcription_tools import _resolve_stt_language
@@ -160,10 +108,8 @@ def _transcribe_command_stt(
timeout = _get_command_stt_timeout(config)
output_format = _get_command_stt_output_format(config)
language = (
language_override
or config.get("language")
or _resolve_stt_language(provider_name, stt_config)
or DEFAULT_COMMAND_STT_LANGUAGE
language_override or config.get("language")
or _resolve_stt_language(provider_name, stt_config) or DEFAULT_COMMAND_STT_LANGUAGE
)
model = model_override or config.get("model") or ""
@@ -171,12 +117,9 @@ def _transcribe_command_stt(
with tempfile.TemporaryDirectory(prefix=f"hermes-cmd-stt-{provider_name}-") as tmpdir:
output_path = Path(tmpdir) / f"transcript.{output_format}"
placeholders = {
"input_path": str(audio.resolve()),
"output_path": str(output_path),
"output_dir": str(output_path.parent),
"format": output_format,
"language": str(language),
"model": str(model),
"input_path": str(audio.resolve()), "output_path": str(output_path),
"output_dir": str(output_path.parent), "format": output_format,
"language": str(language), "model": str(model),
}
command = _render_command_stt_template(command_template, placeholders)
logger.info("Transcribing %s via command STT provider '%s'...", audio.name, provider_name)
@@ -185,11 +128,9 @@ def _transcribe_command_stt(
except subprocess.TimeoutExpired:
return fail(f"STT command provider '{provider_name}' timed out after {timeout:g}s")
except subprocess.CalledProcessError as exc:
detail_parts = [
f"{stream}: {text.strip()}" for stream, text in (("stderr", exc.stderr), ("stdout", exc.stdout)) if text
]
detail = "; ".join(detail_parts) or "no command output"
return fail(f"STT command provider '{provider_name}' exited with code {exc.returncode}: {detail}")
return fail(
f"STT command provider '{provider_name}' exited with code {exc.returncode}: {command_failure_detail(exc)}"
)
except RuntimeError as exc:
return fail(str(exc))
except OSError as exc:
@@ -222,22 +163,18 @@ 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 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.<name>: 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.
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.
"""
if not provider:
return None
key = provider.lower().strip()
if key in BUILTIN_STT_PROVIDERS or key == "none":
if key in _NON_COMMAND_STT_NAMES:
return None
if stt_config is not None and _is_command_stt_provider_config(
_get_named_stt_provider_config(stt_config, key)
):
if stt_config is not None and _is_command_stt_provider_config(_get_named_stt_provider_config(stt_config, key)):
return None
try:
from agent.transcription_registry import get_provider
@@ -292,9 +229,9 @@ def _dispatch_to_plugin_provider(
# attempts to change it are logged and dropped.
_PRE_TRANSCRIPTION_MUTABLE_FIELDS = ("prompt", "language", "model")
# 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-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
_WHISPER_PROMPT_CAPPED_PROVIDERS = frozenset({"local", "openai", "groq", "deepinfra"})
@@ -314,10 +251,7 @@ def _enforce_prompt_length_limit(prompt: Optional[str], provider: str) -> Option
logger.warning(
"Transcription prompt is ~%d tokens; whisper-family provider '%s' "
"only uses the final ~%d — truncating to the last %d characters.",
len(prompt) // _PROMPT_CHARS_PER_TOKEN,
provider,
_WHISPER_PROMPT_TOKEN_CAP,
max_chars,
len(prompt) // _PROMPT_CHARS_PER_TOKEN, provider, _WHISPER_PROMPT_TOKEN_CAP, max_chars,
)
return prompt[-max_chars:]
@@ -335,9 +269,8 @@ def _apply_pre_transcription_hook(
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
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
@@ -351,13 +284,8 @@ def _apply_pre_transcription_hook(
return model, None, prompt
hook_results = invoke_hook(
"pre_transcription",
file_path=file_path,
provider=provider,
model=model,
language=language,
prompt=prompt,
source=source,
"pre_transcription", file_path=file_path, provider=provider,
model=model, language=language, prompt=prompt, source=source,
)
overrides: Dict[str, Any] = {}
for hook_result in hook_results:

View File

@@ -1,7 +1,7 @@
"""Constants, result envelopes and tiny config readers shared by every STT module.
Split out of ``tools/transcription_tools.py``; every name is re-imported there, so
``tools.transcription_tools.<name>`` keeps resolving (and monkeypatching) as before.
Every name is re-imported by ``tools/transcription_tools.py``, so
``tools.transcription_tools.<name>`` keeps resolving (and monkeypatching).
"""
from __future__ import annotations
@@ -11,6 +11,8 @@ import os
import subprocess # noqa: F401 (type annotation only)
from typing import Any, Dict
from tools.tts_command_provider import _get_provider_section as _get_stt_section # noqa: F401 (re-exported)
# Log-record parity with the origin module.
logger = logging.getLogger("tools.transcription_tools")
@@ -39,10 +41,9 @@ 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"}
# Providers with native handlers. 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.
# Providers with native handlers. 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",
})
@@ -59,14 +60,6 @@ def _ok_result(transcript: str, provider: str) -> Dict[str, Any]:
return {"success": True, "transcript": transcript, "provider": provider}
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 _lazy_ensure_quietly(dep: str) -> None:
"""Best-effort ``tools.lazy_deps.ensure(dep, prompt=False)``; failures are swallowed.

View File

@@ -2,11 +2,9 @@
faster-whisper loading (CUDA->CPU fallback, Apple Silicon pinning), the
anti-hallucination transcribe kwargs and segment gate, and the local whisper CLI
(``local_command``) provider. The cached-model singleton and its idle-unload
watcher stay in ``transcription_tools`` (they own the module state).
Split out of ``tools/transcription_tools.py``, which re-imports every name (patch
surface) and is imported lazily here so origin patches still intercept.
(``local_command``) provider. The cached-model singleton and idle-unload watcher
stay in ``transcription_tools`` (module state), which re-imports every name here
(patch surface) and is imported lazily so origin patches still intercept.
"""
from __future__ import annotations
@@ -69,9 +67,8 @@ def _try_lazy_install_stt() -> bool:
"""Lazy-install faster-whisper and re-check dynamically so it's usable without a restart."""
try:
from tools.lazy_deps import ensure
# 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.
# 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.
ensure("stt.faster_whisper", prompt=False)
if _ilu.find_spec("faster_whisper"):
return True
@@ -89,10 +86,9 @@ def _try_lazy_install_stt() -> bool:
return False
# 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
# Substrings identifying a missing/unloadable CUDA runtime library: 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",
@@ -111,10 +107,7 @@ def _sysctl_value(name: str) -> str:
"""Return a sysctl value, or an empty string when unavailable."""
try:
return subprocess.check_output(
["/usr/sbin/sysctl", "-n", name],
stderr=subprocess.DEVNULL,
text=True,
timeout=2,
["/usr/sbin/sysctl", "-n", name], stderr=subprocess.DEVNULL, text=True, timeout=2,
).strip()
except Exception:
return ""
@@ -140,10 +133,10 @@ def _get_idle_unload_seconds(local_cfg: Dict[str, Any]) -> int:
def _load_local_whisper_model(model_name: str, device: str = "auto", compute_type: str = "auto"):
"""Load faster-whisper with graceful CUDA → CPU fallback.
``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.
``device="auto"`` picks CUDA whenever the ctranslate2 wheel ships CUDA libs, even
on hosts without the NVIDIA runtime (WSL2, headless servers). 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:
@@ -172,23 +165,18 @@ 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 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.
# 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 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 — new local-whisper call
sites must go through here instead of hand-rolling kwargs.
"""
"""Kwargs for EVERY local faster-whisper ``model.transcribe`` call — single owner of the anti-hallucination hardening."""
from tools.transcription_tools import _load_stt_config, _resolve_stt_language
stt_config = stt_config if isinstance(stt_config, dict) else _load_stt_config()
local_cfg = stt_config.get("local") or {}
@@ -205,11 +193,10 @@ def build_local_transcribe_kwargs(stt_config: Optional[Dict[str, Any]] = None) -
"min_silence_duration_ms": _config_number(local_cfg, "vad_min_silence_ms", _VAD_MIN_SILENCE_MS_DEFAULT, int)
}
# 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.
# Push the confidence gate into faster-whisper itself: its internal defaults drop
# low-confidence segments BEFORE our post-filter sees them, so 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.
kwargs["no_speech_threshold"], kwargs["log_prob_threshold"] = _confidence_thresholds(local_cfg)
forced_lang = _resolve_stt_language("local", stt_config)
@@ -234,9 +221,8 @@ def _confidence_thresholds(local_cfg: Dict[str, Any]) -> tuple[float, float]:
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 (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.
Conservative AND gate (openai-whisper's own heuristic): non-speech AND low decode
confidence, so quiet-but-real speech survives. Unknown segment shapes are never dropped.
"""
try:
no_speech_prob = float(getattr(segment, "no_speech_prob"))
@@ -254,8 +240,7 @@ def _join_confident_segments(segments: Any, local_cfg: Dict[str, Any]) -> str:
if _is_hallucinated_segment(segment, no_speech_threshold, logprob_threshold):
logger.debug(
"Dropping probable hallucinated segment %r (no_speech_prob=%.3f, avg_logprob=%.3f)",
getattr(segment, "text", ""),
getattr(segment, "no_speech_prob", float("nan")),
getattr(segment, "text", ""), getattr(segment, "no_speech_prob", float("nan")),
getattr(segment, "avg_logprob", float("nan")),
)
continue
@@ -290,10 +275,8 @@ def _transcribe_local_command(
return _error_result(prep_error)
command = command_template.format(
input_path=shlex.quote(prepared_input),
output_dir=shlex.quote(output_dir),
language=shlex.quote(language),
model=shlex.quote(normalized_model),
input_path=shlex.quote(prepared_input), output_dir=shlex.quote(output_dir),
language=shlex.quote(language), model=shlex.quote(normalized_model),
)
# Scrub Hermes secrets from the child env (same policy as _run_command_stt).
from tools.environments.local import hermes_subprocess_env

View File

@@ -7,10 +7,9 @@ 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 sibling modules
(``transcription_{common,audio,local,cloud,command}``) and are re-imported
here so ``tools.transcription_tools.<name>`` stays the patch/import surface.
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
@@ -108,17 +107,15 @@ _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
# download/load the model.
# guards the check-then-load so two concurrent voice messages can't both load.
_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 two concurrent
# transcriptions can't both observe "no watcher alive" and spawn duplicates.
# 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.
_last_transcription_time: float = 0.0
_idle_unload_thread: Optional[threading.Thread] = None
_idle_unload_stop = threading.Event()
@@ -192,13 +189,6 @@ def _is_local_stt_provider(provider: str, stt_config: Dict[str, Any]) -> bool:
# ---- Provider resolution ------------------------------------------------
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:
@@ -210,19 +200,19 @@ def _has_key(env_var: str, provider: str, *, needs_openai: bool = False, needs_m
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 _has_xai_stt_credentials_quietly() -> bool:
try:
return _has_xai_stt_credentials()
except Exception:
return False
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.
# 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.
reason = _openai_audio_unavailable_reason()
if reason is None:
return "openai"
@@ -245,10 +235,7 @@ 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"
@@ -262,14 +249,19 @@ def _resolve_explicit_local_command() -> str:
return "none"
_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)
# 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) 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 auto-detection (explicit
# ``provider: mistral`` installs on first use).
# 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
# auto-detection (explicit ``provider: mistral`` installs on first use).
_CLOUD_PROVIDER_SPECS = {
"groq": (
_has_groq_key, _has_groq_key,
@@ -278,13 +270,12 @@ _CLOUD_PROVIDER_SPECS = {
),
"openai": (
None, lambda: _HAS_OPENAI and _has_openai_audio_backend(),
None, # explicit openai has its own resolver (logs the real gateway reason)
None,
"No local STT available, using OpenAI Whisper API",
),
"mistral": (
_has_mistral_key, _has_mistral_key,
"STT provider 'mistral' configured but mistralai package "
"not installed or MISTRAL_API_KEY not set",
"STT provider 'mistral' configured but mistralai package not installed or MISTRAL_API_KEY not set",
"No local STT available, using Mistral Voxtral Transcribe API",
),
"xai": (
@@ -299,12 +290,18 @@ _CLOUD_PROVIDER_SPECS = {
),
"deepinfra": (
_has_deepinfra_key, _has_deepinfra_key,
"STT provider 'deepinfra' configured but DEEPINFRA_API_KEY not set "
"(or openai package missing)",
"STT provider 'deepinfra' configured but DEEPINFRA_API_KEY not set (or openai package missing)",
"No local STT available, using DeepInfra Whisper API",
),
}
# Explicit selections whose resolution is more than a probe + warning.
_EXPLICIT_RESOLVERS = {
"local": _resolve_explicit_local,
"local_command": _resolve_explicit_local_command,
"openai": _resolve_explicit_openai,
}
def _resolve_explicit_provider(provider: str) -> str:
"""Resolve an explicit ``stt.provider`` to a usable provider name or ``"none"``.
@@ -312,12 +309,9 @@ def _resolve_explicit_provider(provider: str) -> str:
Unknown names pass through untouched so the dispatcher can fail with the
provider-not-registered message.
"""
if provider == "local":
return _resolve_explicit_local()
if provider == "local_command":
return _resolve_explicit_local_command()
if provider == "openai":
return _resolve_explicit_openai()
resolver = _EXPLICIT_RESOLVERS.get(provider)
if resolver is not None:
return resolver()
spec = _CLOUD_PROVIDER_SPECS.get(provider)
if spec is None:
return provider
@@ -341,17 +335,15 @@ def _get_provider(stt_config: dict) -> str:
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 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. Only a raw
# config.yaml selection counts as explicit; otherwise autodetect (which
# prefers local first anyway).
# 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).
try:
from tools.tool_backend_helpers import read_selection
@@ -381,10 +373,7 @@ def _unload_local_model() -> None:
global _local_model, _local_model_name
with _local_model_lock:
if _local_model is not None:
logger.info(
"Unloading local whisper model '%s' after idle timeout",
_local_model_name or "unknown",
)
logger.info("Unloading local whisper model '%s' after idle timeout", _local_model_name or "unknown")
_local_model = None
_local_model_name = None
@@ -392,12 +381,12 @@ 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.
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:
@@ -412,9 +401,7 @@ def _start_idle_unload_watcher(timeout_seconds: int) -> None:
if _local_model is None:
break
try:
timeout = _get_idle_unload_seconds(
_load_stt_config().get("local") or {}
)
timeout = _get_idle_unload_seconds(_load_stt_config().get("local") or {})
except Exception: # noqa: BLE001 - keep the seed value
timeout = initial_timeout
if timeout <= 0:
@@ -424,9 +411,7 @@ def _start_idle_unload_watcher(timeout_seconds: int) -> None:
break
_idle_unload_stop.clear()
_idle_unload_thread = threading.Thread(
target=_watch, name="hermes-stt-idle-unload", daemon=True
)
_idle_unload_thread = threading.Thread(target=_watch, name="hermes-stt-idle-unload", daemon=True)
_idle_unload_thread.start()
@@ -503,9 +488,9 @@ 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 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(
@@ -538,9 +523,8 @@ def _transcribe_local(
def _read_block_error(file_path: str) -> Optional[Dict[str, Any]]:
"""Refuse to feed a credential / secret store (auth.json, .env, OAuth tokens, ...)
to an STT provider, which would ship its plaintext to a third-party API.
Mirrors the image-gen / video-gen read guards."""
"""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."""
from agent.file_safety import get_read_block_error
blocked = get_read_block_error(file_path)
return _error_result(blocked) if blocked else None
@@ -561,9 +545,8 @@ def _transcribe_prepared_audio(
if blocked:
return blocked
# 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.
# Validate before provider resolution so invalid files cannot 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
@@ -626,13 +609,6 @@ def _builtin_model_name(provider: str, stt_config: Dict[str, Any], model: Option
return cfg.get(key, default)
def _builtin_handler(provider: str):
"""Handler for a built-in provider, looked up in this module at call time so tests may patch ``_transcribe_*``."""
if provider not in BUILTIN_STT_PROVIDERS:
return None
return globals()[f"_transcribe_{provider}"]
def _dispatch_stt_provider(
file_path: str,
provider: str,
@@ -656,8 +632,9 @@ def _dispatch_stt_provider(
)
prompt = _enforce_prompt_length_limit(prompt, provider)
handler = _builtin_handler(provider)
if handler is not None:
if provider in BUILTIN_STT_PROVIDERS:
# Looked up in this module at call time so tests may patch ``_transcribe_*``.
handler = globals()[f"_transcribe_{provider}"]
model_name = _builtin_model_name(provider, stt_config, model)
if provider in ("local", "local_command"):
model_name = _normalize_local_model(model_name)
@@ -693,9 +670,8 @@ 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 instead of the all-provider hint.
# An explicit openai selection flattened to "none" carries a selection-specific
# reason (e.g. managed openai-audio gateway down); surface it with its remediation.
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:
@@ -770,5 +746,3 @@ def transcribe_audio_local_fallback(
if _has_local_command():
return _transcribe_local_command(file_path, _normalize_local_model(local_model))
return _error_result("No installed local STT backend is available.", provider="local")

View File

@@ -1,25 +1,20 @@
"""Shared runner for user-configured shell ("command") TTS/STT providers.
Both ``tools.tts_tool`` and ``tools.transcription_tools`` let users declare a
provider as a shell command template with ``{placeholders}``. This module owns
the shell-quote-aware template rendering and the idle-timeout process runner
they share, plus the TTS side's ``tts.providers.<name>`` config layer. Each
origin module re-imports these under its historical private names.
TTS config shape::
``tools.tts_tool`` and ``tools.transcription_tools`` both let users declare a
provider as a shell command template with ``{placeholders}`` (``{{``/``}}`` stay
literal; values are shell-quoted for their surrounding quote context). This
module owns the quote-aware rendering, the idle-timeout process runner and the
generic ``<section>.providers.<name>`` config readers; each origin module
re-imports them under its historical private names. TTS config shape::
tts:
provider: piper-en
providers:
piper-en:
type: command
command: "piper -m ~/model.onnx -f {output_path} < {input_path}"
output_format: wav
piper-en: {type: command, command: "piper -f {output_path} < {input_path}", output_format: wav}
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``.
TTS placeholders: ``{input_path}``/``{text_path}``, ``{output_path}``, ``{format}``,
``{voice}``, ``{model}``, ``{speed}``. Built-in provider names always win over a
same-named entry under ``providers``.
"""
from __future__ import annotations
@@ -33,7 +28,7 @@ import tempfile
import threading
import time
from pathlib import Path
from typing import Any, Dict, Optional
from typing import Any, Dict, FrozenSet, Optional
def shell_quote_context(command_template: str, position: int) -> Optional[str]:
@@ -68,43 +63,26 @@ def quote_command_placeholder(value: str, quote_context: Optional[str]) -> str:
if quote_context == "'":
return value.replace("'", r"'\''")
if quote_context == '"':
return (
value
.replace("\\", "\\\\")
.replace('"', r'\"')
.replace("$", r"\$")
.replace("`", r"\`")
)
return value.replace("\\", "\\\\").replace('"', r'\"').replace("$", r"\$").replace("`", r"\`")
if os.name == "nt":
return subprocess.list2cmdline([value])
return shlex.quote(value)
def render_command_template(
command_template: str,
placeholders: Dict[str, str],
) -> str:
def render_command_template(command_template: str, placeholders: Dict[str, str]) -> str:
"""Replace ``{name}`` placeholders (quote-aware) while preserving ``{{``/``}}``."""
names = "|".join(re.escape(name) for name in placeholders)
pattern = re.compile(
rf"(?<!\$)(?:\{{\{{(?P<double>{names})\}}\}}|\{{(?P<single>{names})\}})"
)
pattern = re.compile(rf"(?<!\$)(?:\{{\{{(?P<double>{names})\}}\}}|\{{(?P<single>{names})\}})")
replacements: list[tuple[str, str]] = []
def replace_match(match: re.Match[str]) -> str:
name = match.group("double") or match.group("single")
token = f"__HERMES_CMD_PLACEHOLDER_{len(replacements)}__"
replacements.append((
token,
quote_command_placeholder(
placeholders[name],
shell_quote_context(command_template, match.start()),
),
))
quoted = quote_command_placeholder(placeholders[name], shell_quote_context(command_template, match.start()))
replacements.append((token, quoted))
return token
rendered = pattern.sub(replace_match, command_template)
rendered = rendered.replace("{{", "{").replace("}}", "}")
rendered = pattern.sub(replace_match, command_template).replace("{{", "{").replace("}}", "}")
for token, value in replacements:
rendered = rendered.replace(token, value)
return rendered
@@ -130,20 +108,15 @@ def terminate_command_process_tree(proc: subprocess.Popen) -> None:
"""Best-effort termination of a shell process and all of its children."""
if proc.poll() is not None:
return
if os.name == "nt":
try:
subprocess.run(
["taskkill", "/F", "/T", "/PID", str(proc.pid)],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
timeout=5,
stdin=subprocess.DEVNULL,
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=5, stdin=subprocess.DEVNULL,
)
except Exception:
proc.kill()
return
try:
import psutil # type: ignore
except ImportError:
@@ -153,40 +126,33 @@ def terminate_command_process_tree(proc: subprocess.Popen) -> None:
except subprocess.TimeoutExpired:
proc.kill()
return
_signal_process_tree(psutil, proc, "terminate")
try:
proc.wait(timeout=2)
return
except subprocess.TimeoutExpired:
pass
_signal_process_tree(psutil, proc, "kill")
_signal_process_tree(psutil, proc, "kill")
def command_env_passthrough(config: Dict[str, Any]) -> list:
"""Return the provider's ``env_passthrough`` allowlist.
The child env is scrubbed of Hermes secrets by default; this list names
variables copied back from the parent env so a trusted template (e.g. a
curl one-liner using its own API key) keeps working.
"""
"""``env_passthrough`` allowlist: parent env vars copied back into the secret-scrubbed child env."""
raw = config.get("env_passthrough")
if not isinstance(raw, (list, tuple)):
return []
return [str(item).strip() for item in raw if str(item).strip()]
def run_command_provider(
command: str,
timeout: float,
env_passthrough: Optional[list] = None,
) -> subprocess.CompletedProcess:
def command_failure_detail(exc: subprocess.CalledProcessError) -> str:
"""``stderr: ...; stdout: ...`` for a failed command provider, or ``no command output``."""
parts = [f"{stream}: {text.strip()}" for stream, text in (("stderr", exc.stderr), ("stdout", exc.stdout)) if text]
return "; ".join(parts) or "no command output"
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 slow-but-alive provider survives, a silently stalled one is killed.
Child env is scrubbed of Hermes secrets while propagating delegated-child
lineage markers.
``timeout`` is an IDLE timeout, reset whenever the command emits output — a
slow-but-alive provider survives, a silently stalled one is killed. Child env
is scrubbed of Hermes secrets while propagating delegated-child lineage markers.
"""
from agent.delegation_context import delegated_child_subprocess_env
from tools.environments.local import hermes_subprocess_env
@@ -197,15 +163,9 @@ def run_command_provider(
if value is not None:
scrubbed[key] = value
popen_kwargs: Dict[str, Any] = {
"shell": True,
"stdout": subprocess.PIPE,
"stderr": subprocess.PIPE,
"text": True,
# Lossy UTF-8 decode: locale-mismatched bytes must not raise in the
# reader threads on non-UTF-8 Windows.
"encoding": "utf-8",
"errors": "replace",
"env": delegated_child_subprocess_env(scrubbed),
"shell": True, "stdout": subprocess.PIPE, "stderr": subprocess.PIPE, "text": True,
# Lossy UTF-8 decode: locale-mismatched bytes must not raise in the reader threads.
"encoding": "utf-8", "errors": "replace", "env": delegated_child_subprocess_env(scrubbed),
}
if os.name == "nt":
popen_kwargs["creationflags"] = getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0)
@@ -222,10 +182,7 @@ def run_command_provider(
read1 = getattr(getattr(stream, "buffer", None), "read1", None)
try:
while True:
if read1 is None:
chunk = stream.read(65536)
else:
chunk = read1(65536).decode(encoding, errors="replace")
chunk = stream.read(65536) if read1 is None else read1(65536).decode(encoding, errors="replace")
if not chunk:
break
output_queue.put((name, chunk))
@@ -273,63 +230,42 @@ def run_command_provider(
break
if chunk:
chunks[name].append(chunk)
stdout = "".join(chunks["stdout"])
stderr = "".join(chunks["stderr"])
stdout, stderr = "".join(chunks["stdout"]), "".join(chunks["stderr"])
try:
raise subprocess.TimeoutExpired(command, timeout)
except subprocess.TimeoutExpired as exc:
raise subprocess.TimeoutExpired(
command, timeout, output=stdout, stderr=stderr,
) from exc
stdout = "".join(chunks["stdout"])
stderr = "".join(chunks["stderr"])
raise subprocess.TimeoutExpired(command, timeout, output=stdout, stderr=stderr) from exc
stdout, stderr = "".join(chunks["stdout"]), "".join(chunks["stderr"])
if proc.returncode:
raise subprocess.CalledProcessError(
proc.returncode, command, output=stdout, stderr=stderr,
)
raise subprocess.CalledProcessError(proc.returncode, command, output=stdout, stderr=stderr)
return subprocess.CompletedProcess(command, proc.returncode, stdout, stderr)
# ===========================================================================
# TTS ``tts.providers.<name>`` config layer
# Generic ``<section>.providers.<name>`` config layer (TTS and STT share it)
# ===========================================================================
# Any ``tts.provider`` value NOT in this set refers to ``tts.providers.<name>``.
BUILTIN_TTS_PROVIDERS = frozenset({
"edge", "elevenlabs", "openai", "minimax", "xai", "mistral", "gemini",
"neutts", "kittentts", "piper", "deepinfra",
})
DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS = 120
DEFAULT_COMMAND_TTS_OUTPUT_FORMAT = "mp3"
COMMAND_TTS_OUTPUT_FORMATS = frozenset(
{"mp3", "wav", "ogg", "flac", "m4a", "aac", "amr", "opus"}
)
DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH = 5000
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."""
if not isinstance(tts_config, dict):
def _get_provider_section(config: Dict[str, Any], name: str) -> Dict[str, Any]:
"""Return ``config[name]`` if it's a dict, else an empty dict."""
if not isinstance(config, dict):
return {}
section = tts_config.get(name)
section = config.get(name)
return section if isinstance(section, dict) else {}
def _get_named_provider_config(tts_config: Dict[str, Any], name: str) -> Dict[str, Any]:
"""Config dict for a user-declared provider, or {}.
def _named_provider_config(config: Dict[str, Any], name: str, builtins: FrozenSet[str]) -> Dict[str, Any]:
"""``<section>.providers.<name>`` (canonical), else ``<section>.<name>`` for non-built-in names only.
``tts.providers.<name>`` is canonical; ``tts.<name>`` 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).
The back-compat form is refused for built-ins so a user's ``openai:`` block
still means the OpenAI provider, not a custom command.
"""
section = _get_provider_section(tts_config, "providers").get(name)
section = _get_provider_section(config, "providers").get(name)
if isinstance(section, dict):
return section
if name.lower() not in BUILTIN_TTS_PROVIDERS:
return _get_provider_section(tts_config, name)
if name.lower() not in builtins:
return _get_provider_section(config, name)
return {}
@@ -344,59 +280,74 @@ def _is_command_provider_config(config: Dict[str, Any]) -> bool:
return isinstance(command, str) and bool(command.strip())
def _resolve_command_provider_config(
provider: str,
tts_config: Dict[str, Any],
) -> Optional[Dict[str, Any]]:
"""The provider config when *provider* is a user-declared command provider.
None for built-in names (native handlers win), unknown names, or
non-command types.
"""
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."""
if not provider:
return None
key = provider.lower().strip()
if key in BUILTIN_TTS_PROVIDERS:
if key in reserved:
return None
config = _get_named_provider_config(tts_config, key)
return config if _is_command_provider_config(config) else None
named = _named_provider_config(config, key, reserved)
return named if _is_command_provider_config(named) else None
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*."""
raw = config.get("timeout", config.get("timeout_seconds", default))
try:
value = float(raw)
except (TypeError, ValueError):
return float(default)
return value if value > 0 else float(default)
def _command_output_format(config: Dict[str, Any], formats: FrozenSet[str], default: str) -> str:
"""Validated ``format``/``output_format`` from *config*, else *default*."""
raw = config.get("format") or config.get("output_format") or default
fmt = str(raw).lower().strip().lstrip(".")
return fmt if fmt in formats else default
# ---- TTS ``tts.providers.<name>`` layer -----------------------------------
# Any ``tts.provider`` value NOT in this set refers to ``tts.providers.<name>``.
BUILTIN_TTS_PROVIDERS = frozenset({
"edge", "elevenlabs", "openai", "minimax", "xai", "mistral", "gemini",
"neutts", "kittentts", "piper", "deepinfra",
})
DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS = 120
DEFAULT_COMMAND_TTS_OUTPUT_FORMAT = "mp3"
COMMAND_TTS_OUTPUT_FORMATS = frozenset({"mp3", "wav", "ogg", "flac", "m4a", "aac", "amr", "opus"})
DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH = 5000
def _get_named_provider_config(tts_config: Dict[str, Any], name: str) -> Dict[str, Any]:
return _named_provider_config(tts_config, name, BUILTIN_TTS_PROVIDERS)
def _resolve_command_provider_config(provider: str, tts_config: Dict[str, Any]) -> Optional[Dict[str, Any]]:
return _resolve_command_config(provider, tts_config, BUILTIN_TTS_PROVIDERS)
def _iter_command_providers(tts_config: Dict[str, Any]):
"""Yield (name, config) pairs for every declared command-type provider."""
for name, cfg in _get_provider_section(tts_config, "providers").items():
if (
isinstance(name, str)
and name.lower() not in BUILTIN_TTS_PROVIDERS
and _is_command_provider_config(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:
"""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)
except (TypeError, ValueError):
return float(DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS)
if value <= 0:
return float(DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS)
return value
return _command_timeout(config, DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS)
def _get_command_tts_output_format(
config: Dict[str, Any],
output_path: Optional[str] = None,
) -> str:
def _get_command_tts_output_format(config: Dict[str, Any], output_path: Optional[str] = None) -> str:
"""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:
return suffix
raw = config.get("format") or config.get("output_format") or DEFAULT_COMMAND_TTS_OUTPUT_FORMAT
fmt = str(raw).lower().strip().lstrip(".")
return fmt if fmt in COMMAND_TTS_OUTPUT_FORMATS else DEFAULT_COMMAND_TTS_OUTPUT_FORMAT
return _command_output_format(config, COMMAND_TTS_OUTPUT_FORMATS, DEFAULT_COMMAND_TTS_OUTPUT_FORMAT)
def _is_command_tts_voice_compatible(config: Dict[str, Any]) -> bool:
@@ -412,24 +363,15 @@ 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:
"""Generate speech by running a user-configured shell command.
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.
Returns the absolute path of the audio file the command wrote. Raises
``ValueError`` for invalid provider config and ``RuntimeError`` for
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:
raise ValueError(
f"tts.providers.{provider_name}.command is not configured"
)
raise ValueError(f"tts.providers.{provider_name}.command is not configured")
output = Path(output_path).expanduser()
output.parent.mkdir(parents=True, exist_ok=True)
@@ -443,46 +385,24 @@ def _generate_command_tts(
with tempfile.TemporaryDirectory() as tmpdir:
text_path = Path(tmpdir) / "input.txt"
text_path.write_text(text, encoding="utf-8")
placeholders = {
"input_path": str(text_path),
"text_path": str(text_path),
"output_path": str(output),
"format": output_format,
"voice": str(config.get("voice", "")),
"model": str(config.get("model", "")),
"speed": str(speed),
"input_path": str(text_path), "text_path": str(text_path), "output_path": str(output),
"format": output_format, "voice": str(config.get("voice", "")),
"model": str(config.get("model", "")), "speed": str(speed),
}
command = render_command_template(command_template, placeholders)
try:
# Resolved through the origin so tests patching
# ``tools.tts_tool._run_command_tts`` still intercept.
# Resolved through the origin so tests patching ``tools.tts_tool._run_command_tts`` still intercept.
from tools.tts_tool import _run_command_tts
_run_command_tts(
command,
timeout,
env_passthrough=command_env_passthrough(config),
)
_run_command_tts(command, timeout, env_passthrough=command_env_passthrough(config))
except subprocess.TimeoutExpired as exc:
raise RuntimeError(
f"TTS provider '{provider_name}' timed out after {timeout:g}s"
) from exc
raise RuntimeError(f"TTS provider '{provider_name}' timed out after {timeout:g}s") from exc
except subprocess.CalledProcessError as exc:
detail_parts = []
if exc.stderr:
detail_parts.append(f"stderr: {exc.stderr.strip()}")
if exc.stdout:
detail_parts.append(f"stdout: {exc.stdout.strip()}")
detail = "; ".join(detail_parts) or "no command output"
raise RuntimeError(
f"TTS provider '{provider_name}' exited with code "
f"{exc.returncode}: {detail}"
f"TTS provider '{provider_name}' exited with code {exc.returncode}: {command_failure_detail(exc)}"
) from exc
if not output.exists() or output.stat().st_size <= 0:
raise RuntimeError(
f"TTS provider '{provider_name}' produced no output at {output}"
)
raise RuntimeError(f"TTS provider '{provider_name}' produced no output at {output}")
return str(output)

View File

@@ -1,16 +1,11 @@
"""Provider-agnostic streaming TTS: sentence text → int16 PCM chunk iterator.
"""Provider-agnostic streaming TTS: sentence text → int16 mono PCM chunk iterator.
``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.
One contract (int16 mono PCM at ``sample_rate``): **true streamers**
(`StreamingTTSProvider.stream`) wrap chunked APIs (ElevenLabs pcm_24000, OpenAI
pcm, …); providers with no chunked API (edge, the default) still get per-
*sentence* playback via the sync ``text_to_speech_tool`` path in the dispatcher.
Adding a streamer is ``@register("name")`` on a subclass; the dispatcher, config
gate (``tts.<name>.streaming``) and resolver come free.
``stream_tts_to_speaker`` (``tools.tts_tool``) owns the sentence buffer, sounddevice
output and stop/queue protocol; this module owns the *provider* half so playback
starts on sentence one. True streamers (``StreamingTTSProvider.stream``) wrap chunked
APIs; providers with no chunked API (edge, the default) get per-sentence playback via
the sync ``text_to_speech_tool`` path. Adding a streamer is ``@register("name")`` on
a subclass; the dispatcher, config gate (``tts.<name>.streaming``) and resolver come free.
"""
from __future__ import annotations
@@ -32,10 +27,9 @@ _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); seam over ``tts_tool._resolve_provider_key``.
Monkeypatchable seam over ``tools.tts_tool._resolve_provider_key``. ALL
streaming-provider key lookups go through here — never bare ``get_env_value``.
ALL streaming-provider key lookups go through here — never bare ``get_env_value``.
"""
try:
from tools.tts_tool import _resolve_provider_key
@@ -49,17 +43,11 @@ 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, 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
# Interruption latch: when the user barges in on a spoken reply, the surface marks
# it; 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.]"
)
SPEECH_INTERRUPTED_NOTE = "[Note: the user interrupted your previous spoken reply before it finished.]"
_INTERRUPT_TTL_S = 120.0
_interrupted_at: Optional[float] = None
@@ -83,10 +71,9 @@ _THINK_BLOCK_RE = re.compile(r"<think[\s>].*?</think>", flags=re.DOTALL)
class SentenceChunker:
"""Incremental sentence cutter for LLM token deltas.
Shared by the speaker pipeline and the speak-stream WebSocket so every
surface cuts speech identically. Strips ``<think>`` blocks (even split
across deltas) and merges fragments shorter than *min_len* into the
following sentence, so "Ha!" rides along instead of stalling as a tiny clip.
Shared by the speaker pipeline and the speak-stream WebSocket so every surface
cuts speech identically. Strips ``<think>`` blocks (even split across deltas) and
merges fragments shorter than *min_len* into the following sentence.
"""
def __init__(self, min_len: int = 20):
@@ -117,10 +104,6 @@ class SentenceChunker:
return [tail] if tail else []
# ---------------------------------------------------------------------------
# ABC + registry
# ---------------------------------------------------------------------------
class StreamingTTSProvider(ABC):
"""Yields raw int16, little-endian, mono PCM chunks at ``sample_rate``."""
@@ -149,7 +132,6 @@ def register(name: str) -> Callable[[type[StreamingTTSProvider]], type[Streaming
def _wrap(cls: type[StreamingTTSProvider]) -> type[StreamingTTSProvider]:
_REGISTRY[name] = cls
return cls
return _wrap
@@ -177,27 +159,20 @@ def resolve_streaming_provider(
) -> Optional[StreamingTTSProvider]:
"""Return a ready streamer for the *configured* provider, else ``None``.
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.
``tts.streaming.provider`` when set: a name pins that exact streamer (``None``
if unusable); ``auto`` returns the first usable in ``_PROVIDER_PRIORITY``.
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()
pinned = str((tts_config.get("streaming") or {}).get("provider") or "").lower().strip()
if pinned == "auto":
for name in _PROVIDER_PRIORITY:
inst = _try_instantiate(name, tts_config)
if inst is not None:
return inst
return None
if pinned:
return _try_instantiate(pinned, tts_config)
name = (preferred or _get_provider(tts_config)).lower().strip()
return _try_instantiate(name, tts_config)
return _try_instantiate(pinned or (preferred or _get_provider(tts_config)).lower().strip(), tts_config)
def _capped(chunks: Iterator[bytes], label: str) -> Iterator[bytes]:
@@ -206,16 +181,11 @@ def _capped(chunks: Iterator[bytes], label: str) -> Iterator[bytes]:
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)
logger.warning("%s exceeded %d bytes for one sentence; truncating", label, _STREAM_SENTENCE_BYTE_CAP)
return
yield chunk
# ---------------------------------------------------------------------------
# Providers
# ---------------------------------------------------------------------------
@register("elevenlabs")
class ElevenLabsStreamer(StreamingTTSProvider):
"""ElevenLabs chunked HTTP → pcm_24000 (the original reference path)."""
@@ -229,25 +199,18 @@ class ElevenLabsStreamer(StreamingTTSProvider):
def stream(self, text: str) -> Iterator[bytes]:
from tools.tts_tool import _import_elevenlabs
from tools.tts_tool_providers import (
DEFAULT_ELEVENLABS_STREAMING_MODEL_ID,
DEFAULT_ELEVENLABS_VOICE_ID,
_elevenlabs_environment_kwargs,
DEFAULT_ELEVENLABS_STREAMING_MODEL_ID, DEFAULT_ELEVENLABS_VOICE_ID, _elevenlabs_environment_kwargs,
)
client = _import_elevenlabs()(
api_key=_resolve_key("ELEVENLABS_API_KEY", "elevenlabs"),
**_elevenlabs_environment_kwargs(self.section),
api_key=_resolve_key("ELEVENLABS_API_KEY", "elevenlabs"), **_elevenlabs_environment_kwargs(self.section),
)
voice_id = self.section.get("voice_id", DEFAULT_ELEVENLABS_VOICE_ID)
model_id = self.section.get(
"streaming_model_id",
self.section.get("model_id", DEFAULT_ELEVENLABS_STREAMING_MODEL_ID),
"streaming_model_id", self.section.get("model_id", DEFAULT_ELEVENLABS_STREAMING_MODEL_ID),
)
yield from client.text_to_speech.convert(
text=text,
voice_id=voice_id,
model_id=model_id,
output_format="pcm_24000",
text=text, voice_id=voice_id, model_id=model_id, output_format="pcm_24000",
)
@@ -275,30 +238,18 @@ class OpenAIStreamer(StreamingTTSProvider):
client = OpenAI(
api_key=(self.section.get("api_key") or resolve_openai_audio_api_key()),
base_url=(
self.section.get("base_url")
or get_env_value("OPENAI_BASE_URL")
or None
),
base_url=(self.section.get("base_url") or get_env_value("OPENAI_BASE_URL") or None),
)
model = self.section.get("model", "gpt-4o-mini-tts")
voice = self.section.get("voice", "alloy")
with client.audio.speech.with_streaming_response.create(
model=model,
voice=voice,
input=text,
response_format="pcm",
model=self.section.get("model", "gpt-4o-mini-tts"), voice=self.section.get("voice", "alloy"),
input=text, response_format="pcm",
) as response:
yield from _capped(response.iter_bytes(), "OpenAI streaming TTS")
@register("gemini")
class GeminiStreamer(StreamingTTSProvider):
"""Gemini ``streamGenerateContent?alt=sse`` → base64 PCM chunks (24 kHz).
``?alt=sse`` flips the response from one JSON blob to an SSE feed of
base64 PCM chunks. Uses requests with a bounded streamed body.
"""
"""Gemini ``streamGenerateContent?alt=sse`` → SSE feed of base64 PCM chunks (24 kHz), bounded streamed body."""
sample_rate = 24000
@@ -312,42 +263,25 @@ 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
voice = str(self.section.get("voice", DEFAULT_GEMINI_TTS_VOICE)).strip() or DEFAULT_GEMINI_TTS_VOICE
base_url = str(
self.section.get("base_url")
or get_env_value("GEMINI_BASE_URL")
or DEFAULT_GEMINI_TTS_BASE_URL
self.section.get("base_url") or get_env_value("GEMINI_BASE_URL") or DEFAULT_GEMINI_TTS_BASE_URL
).strip().rstrip("/")
payload = {
"contents": [{"parts": [{"text": text}]}],
"generationConfig": {
"responseModalities": ["AUDIO"],
"speechConfig": {
"voiceConfig": {
"prebuiltVoiceConfig": {"voiceName": voice},
},
},
"speechConfig": {"voiceConfig": {"prebuiltVoiceConfig": {"voiceName": voice}}},
},
}
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: "):
@@ -374,9 +308,9 @@ class GeminiStreamer(StreamingTTSProvider):
class XAIStreamer(StreamingTTSProvider):
"""xAI WebSocket TTS (``wss://api.x.ai/v1/tts``) → binary PCM frames (24 kHz mono int16).
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.
Credentials route through ``resolve_xai_http_credentials`` (OAuth or XAI_API_KEY),
same as the sync path. ``_collect_async`` bridges the async WS loop to the sync
iterator contract — the seam unit tests patch.
"""
sample_rate = 24000
@@ -394,8 +328,6 @@ class XAIStreamer(StreamingTTSProvider):
def stream(self, text: str) -> Iterator[bytes]:
yield from _capped(iter(self._collect_async(text)), "xAI streaming TTS")
# -- async→sync bridge (test seam) ------------------------------------
def _collect_async(self, text: str) -> List[bytes]:
import asyncio
@@ -420,18 +352,10 @@ class XAIStreamer(StreamingTTSProvider):
if not api_key:
raise RuntimeError("No xAI credentials for streaming TTS")
voice = str(self.section.get("voice_id", DEFAULT_XAI_VOICE_ID)).strip() or DEFAULT_XAI_VOICE_ID
ws_url = str(
self.section.get("streaming_url") or "wss://api.x.ai/v1/tts"
).strip()
ws_url = str(self.section.get("streaming_url") or "wss://api.x.ai/v1/tts").strip()
async with websockets.connect(
ws_url, extra_headers={"Authorization": f"Bearer {api_key}"}
) as ws:
await ws.send(_json.dumps({
"text": text,
"voice_id": voice,
"response_format": "pcm",
}))
async with websockets.connect(ws_url, extra_headers={"Authorization": f"Bearer {api_key}"}) as ws:
await ws.send(_json.dumps({"text": text, "voice_id": voice, "response_format": "pcm"}))
try:
while True:
message = await ws.recv()
@@ -448,8 +372,9 @@ class XAIStreamer(StreamingTTSProvider):
if etype == "done":
return
if etype == "error":
logger.warning("xAI WS error envelope: %s",
envelope.get("error") or envelope.get("message") or envelope)
logger.warning(
"xAI WS error envelope: %s", envelope.get("error") or envelope.get("message") or envelope,
)
return
except Exception as exc:
if exc.__class__.__name__ == "ConnectionClosed":

View File

@@ -1,11 +1,8 @@
"""Utilities for preparing assistant text for speech synthesis.
"""Deterministic cleanup turning assistant Markdown into a spoken script.
The TTS provider should receive a spoken script, not raw chat Markdown. This
module centralises the lightweight, deterministic cleanup used by explicit TTS
calls and gateway auto-TTS replies.
Non-ASCII characters are written as escapes on purpose so the file stays free of
invisible/look-alike glyphs.
Shared by explicit TTS calls, gateway auto-TTS, voice-mode streaming and the web
dashboard. Non-ASCII characters are written as escapes on purpose so the file
stays free of invisible/look-alike glyphs.
"""
from __future__ import annotations
@@ -13,9 +10,9 @@ from __future__ import annotations
import html
import re
# Sentinel appended to former heading lines so smooth_whitespace_for_tts can
# fold a heading into the sentence that follows it ("Weather, it will be sunny")
# rather than leaving a bare "Weather." label that reads abruptly aloud.
# Sentinel appended to former heading lines so smooth_whitespace_for_tts folds the
# heading into the sentence after it ("Weather, it will be sunny") instead of a bare
# "Weather." label.
_HEAD = "\x00"
_MD_CODE_BLOCK_RE = re.compile(r"```[\s\S]*?```")
@@ -34,21 +31,22 @@ _MD_HR_RE = re.compile(r"^\s*[-*_]{3,}\s*$", flags=re.MULTILINE)
_MD_TABLE_PIPE_RE = re.compile(r"\s*\|\s*")
_URL_RE = re.compile(r"https?://\S+")
# Broad emoji / pictograph cleanup. Voice providers vary a lot here; most read
# emojis as awkward labels, so keep the speech script calm and literal.
# Unit suffix (regex, after a digit) -> spoken word; km/h variants before the bare "m".
_UNIT_WORDS = (
(r"km\s*/\s*h", "kilometres per hour"), (r"km/h", "kilometres per hour"),
(r"mm", "millimetres"), (r"cm", "centimetres"), (r"m", "metres"),
)
# Currency prefix (regex) -> spoken word; order matters (NZ$/A$/US$ before bare $).
_CURRENCY_WORDS = (
(r"NZ\$", "New Zealand dollars", re.IGNORECASE), (r"A\$", "Australian dollars", re.IGNORECASE),
(r"US\$", "US dollars", re.IGNORECASE), ("€", "euros", 0), ("£", "pounds", 0), (r"\$", "dollars", 0),
)
# Broad emoji / pictograph cleanup: most voice providers read emojis as awkward labels.
_EMOJI_RE = re.compile(
"["
"\U0001F1E6-\U0001F1FF"
"\U0001F300-\U0001F5FF"
"\U0001F600-\U0001F64F"
"\U0001F680-\U0001F6FF"
"\U0001F700-\U0001F77F"
"\U0001F780-\U0001F7FF"
"\U0001F800-\U0001F8FF"
"\U0001F900-\U0001F9FF"
"\U0001FA00-\U0001FAFF"
"☀-➿"
"]+",
"[\U0001F1E6-\U0001F1FF\U0001F300-\U0001F5FF\U0001F600-\U0001F64F\U0001F680-\U0001F6FF"
"\U0001F700-\U0001F77F\U0001F780-\U0001F7FF\U0001F800-\U0001F8FF\U0001F900-\U0001F9FF"
"\U0001FA00-\U0001FAFF☀-➿]+",
flags=re.UNICODE,
)
_VARIATION_SELECTOR_RE = re.compile("[︎️]")
@@ -70,34 +68,26 @@ def strip_markdown_for_tts(text: str) -> str:
text = _MD_ITALIC_RE.sub(r"\1", text)
text = _MD_UNDERSCORE_ITALIC_RE.sub(r"\1", text)
text = _MD_STRIKE_RE.sub(r"\1", text)
# Mark headings (do not just delete the marker): the whitespace pass folds a
# heading into the sentence after it so speech says "Weather, it will be
# sunny" instead of a clipped "Weather." then a separate sentence.
# Mark headings (do not just delete the marker): see _HEAD.
text = _MD_HEADING_LINE_RE.sub(lambda m: m.group(1).rstrip() + _HEAD, text)
text = _MD_BLOCKQUOTE_RE.sub("", text)
text = _MD_LIST_ITEM_RE.sub("", text)
text = _MD_HR_RE.sub("", text)
# Pipe tables are terrible read aloud. Turn any leftover pipes into pauses
# instead of letting a provider speak "vertical bar".
# Leftover table pipes become pauses instead of a spoken "vertical bar".
text = _MD_TABLE_PIPE_RE.sub("; ", text)
return text
def _normalize_temperature_ranges(text: str) -> str:
# 11-17 degrees C -> "11 to 17 degrees Celsius" (en/em dash or hyphen).
text = re.sub(
r"(?<!\w)([-+\u2212]?\d+(?:\.\d+)?)\s*[\u2013\u2014-]\s*([-+\u2212]?\d+(?:\.\d+)?)\s*°\s*C\b",
lambda m: f"{m.group(1).replace(chr(0x2212), '-')} to {m.group(2).replace(chr(0x2212), '-')} degrees Celsius",
text,
flags=re.IGNORECASE,
)
text = re.sub(
r"(?<!\w)([-+\u2212]?\d+(?:\.\d+)?)\s*[\u2013\u2014-]\s*([-+\u2212]?\d+(?:\.\d+)?)\s*°\s*F\b",
lambda m: f"{m.group(1).replace(chr(0x2212), '-')} to {m.group(2).replace(chr(0x2212), '-')} degrees Fahrenheit",
text,
flags=re.IGNORECASE,
)
"""``11-17°C`` -> ``11 to 17 degrees Celsius`` (en/em dash or hyphen; unicode minus normalized)."""
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}",
text,
flags=re.IGNORECASE,
)
return text
@@ -112,44 +102,35 @@ def normalize_symbols_for_tts(text: str) -> str:
text = text.replace("…", "...") # ellipsis
text = _normalize_temperature_ranges(text)
# Temperatures with a number. Do this before generic degree handling.
text = re.sub(r"(?<!\w)([-+]?\d+(?:\.\d+)?)\s*°\s*C\b", r"\1 degrees Celsius", text, flags=re.IGNORECASE)
text = re.sub(r"(?<!\w)([-+]?\d+(?:\.\d+)?)\s*°\s*F\b", r"\1 degrees Fahrenheit", text, flags=re.IGNORECASE)
# Bare units with no leading number ("measured in degrees C").
text = re.sub(r"°\s*C\b", "degrees Celsius", text, flags=re.IGNORECASE)
text = re.sub(r"°\s*F\b", "degrees Fahrenheit", text, flags=re.IGNORECASE)
# Any remaining degree symbol (angles, stray cases).
# 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)
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)
text = text.replace("°", " degrees")
# Common weather/travel units.
text = re.sub(r"(?<=\d)\s*km\s*/\s*h\b", " kilometres per hour", text, flags=re.IGNORECASE)
text = re.sub(r"(?<=\d)\s*km/h\b", " kilometres per hour", text, flags=re.IGNORECASE)
text = re.sub(r"(?<=\d)\s*mm\b", " millimetres", text, flags=re.IGNORECASE)
text = re.sub(r"(?<=\d)\s*cm\b", " centimetres", text, flags=re.IGNORECASE)
text = re.sub(r"(?<=\d)\s*m\b", " metres", text, flags=re.IGNORECASE)
for pattern, word in _UNIT_WORDS:
text = re.sub(r"(?<=\d)\s*" + pattern + r"\b", " " + word, text, flags=re.IGNORECASE)
# Numeric rates only ("5/month" -> "5 per month"). Requiring digit-then-letter
# keeps "and/or", "N/A", "TCP/IP" and dates like "2026/06" intact.
text = re.sub(r"(?<=\d)\s*/\s*(?=[A-Za-z])", " per ", text)
# Money and percentages. The integer part must END in a digit so a trailing
# comma ("A$50, ...") is not swallowed into the spoken amount.
text = re.sub(r"NZ\$\s*([\d,]*\d(?:\.\d+)?)", r"\1 New Zealand dollars", text, flags=re.IGNORECASE)
text = re.sub(r"A\$\s*([\d,]*\d(?:\.\d+)?)", r"\1 Australian dollars", text, flags=re.IGNORECASE)
text = re.sub(r"US\$\s*([\d,]*\d(?:\.\d+)?)", r"\1 US dollars", text, flags=re.IGNORECASE)
text = re.sub(r"€\s*([\d,]*\d(?:\.\d+)?)", r"\1 euros", text)
text = re.sub(r"£\s*([\d,]*\d(?:\.\d+)?)", r"\1 pounds", text)
text = re.sub(r"\$\s*([\d,]*\d(?:\.\d+)?)", r"\1 dollars", text)
# Money and percentages. The integer part must END in a digit so a trailing
# comma ("A$50, ...") is not swallowed into the spoken amount. Prefixed
# currencies run first so "$" doesn't eat "NZ$".
for symbol, word, flags in _CURRENCY_WORDS:
text = re.sub(symbol + r"\s*([\d,]*\d(?:\.\d+)?)", r"\1 " + word, text, flags=flags)
text = re.sub(r"(?<=\d)\s*%", " percent", text)
# Operators and separators that commonly leak from formatted answers.
text = text.replace("&", " and ")
text = re.sub("[•◦▪▫]", " ", text) # bullet glyphs
text = text.replace("→", " to ") # ->
text = text.replace("⇒", " to ") # =>
text = text.replace("≈", " about ") # almost equal
text = text.replace("~", " about ")
for symbol, word in (("→", " to "), ("⇒", " to "), ("≈", " about "), ("~", " about ")):
text = text.replace(symbol, word)
text = _VARIATION_SELECTOR_RE.sub("", text)
text = _EMOJI_RE.sub("", text)
@@ -159,10 +140,8 @@ def normalize_symbols_for_tts(text: str) -> str:
def smooth_whitespace_for_tts(text: str) -> str:
"""Collapse visual formatting into calm spoken paragraphs.
A former heading line (marked with the _HEAD sentinel) folds into the next
content line as a spoken lead-in: "Weather" + "It will be sunny" becomes
"Weather, It will be sunny." A heading with no content after it becomes its
own short sentence.
A _HEAD-marked heading folds into the next content line as a lead-in ("Weather,
It will be sunny."); a heading with no content after it becomes its own sentence.
"""
if not text:
return ""
@@ -182,8 +161,7 @@ def smooth_whitespace_for_tts(text: str) -> str:
is_heading = raw_line.rstrip().endswith(_HEAD)
line = raw_line.replace(_HEAD, "").strip()
if not line:
# Hold a pending heading across blank lines so it still folds into
# the next real content line; otherwise just collapse the blank.
# Hold a pending heading across blank lines so it still folds into the next content line.
if pending_heading is None and lines and lines[-1] != "":
lines.append("")
continue
@@ -209,44 +187,30 @@ def smooth_whitespace_for_tts(text: str) -> str:
return text.strip()
# Reasoning blocks: models with ``/reasoning show`` enabled emit
# ``<think>...</think>`` blocks in the final assistant message. Users want to
# SEE reasoning, not hear it read aloud (#34213).
# ``/reasoning show`` emits ``<think>...</think>`` in the final message: users want to
# SEE reasoning, not hear it. An unterminated block (streaming cut-off) is also silenced.
_THINK_BLOCK_RE = re.compile(r"<think[\s>].*?</think>", flags=re.DOTALL | re.IGNORECASE)
# An unterminated block (streaming cut-off) should still not be spoken.
_THINK_BLOCK_OPEN_RE = re.compile(r"<think[\s>].*\Z", flags=re.DOTALL | re.IGNORECASE)
# Turn-end file-mutation verifier footer appended by run_agent.py
# (``_format_file_mutation_failure_footer``). It's a UI affordance — reading
# "warning file mutation verifier, 2 files were NOT modified..." aloud is
# noise (#40772). The footer is a ``⚠️ File-mutation verifier:`` header line
# followed by indented ``•`` bullet lines; strip the whole block.
_VERIFIER_FOOTER_RE = re.compile(
r"^\s*⚠️?\s*File-mutation verifier:.*(?:\n[ \t]+•.*)*",
flags=re.MULTILINE,
)
# run_agent.py's turn-end file-mutation verifier footer (a ``⚠️ File-mutation verifier:``
# header line plus indented ``•`` bullets) is a UI affordance, not speech.
_VERIFIER_FOOTER_RE = re.compile(r"^\s*⚠️?\s*File-mutation verifier:.*(?:\n[ \t]+•.*)*", flags=re.MULTILINE)
def strip_nonspoken_blocks(text: str) -> str:
"""Remove blocks that must never reach a speech provider.
Currently: ``<think>`` reasoning blocks and the end-of-turn
file-mutation verifier footer.
"""
"""Remove ``<think>`` reasoning blocks and the file-mutation verifier footer."""
if not text:
return ""
text = _THINK_BLOCK_RE.sub(" ", text)
text = _THINK_BLOCK_OPEN_RE.sub(" ", text)
text = _VERIFIER_FOOTER_RE.sub(" ", text)
for pattern in (_THINK_BLOCK_RE, _THINK_BLOCK_OPEN_RE, _VERIFIER_FOOTER_RE):
text = pattern.sub(" ", text)
return text
def flatten_newlines_for_payload(text: str) -> str:
"""Collapse newlines into sentence breaks for single-line TTS payloads.
Some OpenAI-compatible backends (e.g. Kokoro) truncate synthesis at the
first newline (#9004). The smoothing pass already terminates each line
with punctuation, so newlines can safely become plain spaces.
Some OpenAI-compatible backends (e.g. Kokoro) truncate at the first newline; the
smoothing pass already terminates each line with punctuation, so this is safe.
"""
if not text:
return ""
@@ -259,28 +223,20 @@ def flatten_newlines_for_payload(text: str) -> str:
def prepare_spoken_text(text: str, max_chars: int | None = 4000) -> str:
"""Return a TTS-friendly script from assistant text.
"""Return a TTS-friendly script from assistant text (deterministic cleanup, not a rewrite).
Deterministic cleanup, not a semantic rewrite: it removes ``<think>``
reasoning blocks and the file-mutation verifier footer, removes Markdown,
expands common symbols such as a degree-Celsius sign to "degrees Celsius",
turns visual line formatting into speakable sentence pauses, and flattens
the result to a single line so newline-sensitive providers (Kokoro) speak
the whole script.
Pipeline: non-spoken blocks > Markdown > symbols/units > line formatting into
sentence pauses > single line (for newline-sensitive providers), then ``max_chars``.
"""
spoken = strip_nonspoken_blocks(text)
spoken = strip_markdown_for_tts(spoken)
spoken = normalize_symbols_for_tts(spoken)
spoken = smooth_whitespace_for_tts(spoken)
spoken = flatten_newlines_for_payload(spoken)
spoken = text
for step in (strip_nonspoken_blocks, strip_markdown_for_tts, normalize_symbols_for_tts,
smooth_whitespace_for_tts, flatten_newlines_for_payload):
spoken = step(spoken)
if max_chars is not None and max_chars > 0 and len(spoken) > max_chars:
spoken = spoken[:max_chars].rstrip()
return spoken
# ===========================================================================
# Speech text cleanup (shared by voice-mode streaming and gateway auto-TTS)
# ===========================================================================
# Legacy regex fallback, only used if the shared normalizer raises.
_LEGACY_TTS_STRIP_STEPS = (
(re.compile(r'<think[\s>].*?</think>', flags=re.DOTALL), ' '),
@@ -300,14 +256,7 @@ _LEGACY_TTS_STRIP_STEPS = (
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 <think> 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.
"""
"""``prepare_spoken_text`` without a length cap; falls back to the legacy regex pipeline if it raises."""
try:
return prepare_spoken_text(text, max_chars=None)
except Exception: