refactor(tools): unify STT/TTS command-provider config helpers, REST STT flow, dispatch tables
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user