refactor(tools): compact STT module docstrings and section banners, keep every invariant

This commit is contained in:
Teknium
2026-09-03 01:13:03 -07:00
parent 198fe72a35
commit ff44eca74f
8 changed files with 83 additions and 126 deletions
+13 -17
View File
@@ -58,15 +58,14 @@ _STT_M4A_ENCODE_ARGS = ("-vn", "-ac", "1", "-ar", "16000", "-c:a", "aac", "-b:a"
def _run_ffmpeg_stt_encode(ffmpeg: str, input_path: str, output_path: str, *, audio_filter: Optional[str] = None) -> None:
"""Run the shared STT m4a encode, optionally with an ``-af`` filter. Raises on failure; callers own the semantics."""
filter_args = ["-af", audio_filter] if audio_filter else []
_run_quiet([ffmpeg, "-y", "-i", input_path, *filter_args, *_STT_M4A_ENCODE_ARGS, output_path], timeout=120)
_run_quiet([ffmpeg, "-y", "-i", input_path, *filter_args, *_STT_M4A_ENCODE_ARGS, output_path],
timeout=120)
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; ``(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.
"""
Newer OpenAI models reject containers ``whisper-1`` accepted (notably Ogg/Opus voice notes) and
gateway downloads may carry a misleading extension."""
from tools.transcription_tools import _find_ffmpeg_binary, _run_ffmpeg_stt_encode
ffmpeg = _find_ffmpeg_binary()
if not ffmpeg:
@@ -207,18 +206,17 @@ _CLOUD_TRIM_MIN_INPUT_SECONDS = 12.0
def _probe_audio_duration(file_path: str) -> Optional[float]:
"""Return the audio duration in seconds via ffprobe, or None.
Canonical sync probe; ``gateway/run.py._probe_audio_duration`` and the
Telegram adapter carry local variants — keep the command shape in sync.
"""
"""Return the audio duration in seconds via ffprobe, or None. Canonical sync probe;
``gateway/run.py._probe_audio_duration`` and the Telegram adapter carry local variants — keep
the command shape in sync."""
from tools.transcription_tools import _find_ffprobe_binary
ffprobe = _find_ffprobe_binary()
if not ffprobe:
return None
try:
return float(_run_quiet([ffprobe, "-v", "error", "-show_entries", "format=duration", "-of",
"default=noprint_wrappers=1:nokey=1", file_path], timeout=30).stdout.strip())
probe = _run_quiet([ffprobe, "-v", "error", "-show_entries", "format=duration", "-of",
"default=noprint_wrappers=1:nokey=1", file_path], timeout=30)
return float(probe.stdout.strip())
except Exception: # noqa: BLE001 - probe is best-effort
return None
@@ -235,9 +233,7 @@ 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 (= upload the original).
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)
if not enabled:
@@ -277,8 +273,8 @@ def _trim_silence_for_cloud_stt(file_path: str, stt_config: Dict[str, Any]) -> O
logger.debug("Cloud STT silence trim discarded for %s: saves <%.0f%% (%.1fs -> %.1fs)",
name, _CLOUD_TRIM_MIN_SAVING * 100, original_duration, trimmed_duration)
return None
logger.info("Trimmed silence from %s before cloud STT upload (%.1fs -> %.1fs, -%d%%)",
name, original_duration, trimmed_duration, round((1 - trimmed_duration / original_duration) * 100))
logger.info("Trimmed silence from %s before cloud STT upload (%.1fs -> %.1fs, -%d%%)", name,
original_duration, trimmed_duration, round((1 - trimmed_duration / original_duration) * 100))
keep_result = True
return trimmed_path
except Exception as exc: # noqa: BLE001 - trim is best-effort
+16 -25
View File
@@ -38,10 +38,8 @@ def _has_xai_stt_credentials() -> bool:
def _with_openai_client(api_key: str, base_url: Optional[str], file_path: str, log_label: str, body):
"""Run ``body(client)`` on a fresh OpenAI SDK client (30s timeout, no retries); always closed.
Errors map to the shared envelope. APIConnectionError is checked before APITimeoutError
(its subclass) so timeouts report as connection errors, as they always have.
"""
Errors map to the shared envelope. APIConnectionError is checked before APITimeoutError (its
subclass) so timeouts report as connection errors, as they always have."""
try:
from openai import OpenAI
client = OpenAI(api_key=api_key, base_url=base_url, timeout=30, max_retries=0)
@@ -96,13 +94,13 @@ def _transcribe_groq(
def _run(client):
with open(file_path, "rb") as audio_file:
transcription = client.audio.transcriptions.create(
file=audio_file, model=model_name, response_format="text", **_sdk_prompt_kwargs(language, prompt))
transcription = client.audio.transcriptions.create(file=audio_file, model=model_name,
response_format="text",
**_sdk_prompt_kwargs(language, prompt))
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)
@@ -110,11 +108,9 @@ def _transcribe_openai(
file_path: str, model_name: str, *, api_key: Optional[str] = None,
base_url: Optional[str] = None, provider_label: str = "openai", language: Optional[str] = None,
prompt: Optional[str] = None) -> Dict[str, Any]:
"""Transcribe via the OpenAI ``audio.transcriptions.create`` SDK shape.
Shared by every OpenAI-compatible endpoint (DeepInfra etc.): explicit ``api_key``/
``base_url`` skip the OpenAI-only auth chain; ``provider_label`` names the response's provider.
"""
"""Transcribe via the OpenAI ``audio.transcriptions.create`` SDK shape, shared by every
OpenAI-compatible endpoint (DeepInfra etc.): explicit ``api_key``/``base_url`` skip the
OpenAI-only auth chain; ``provider_label`` names the response's provider."""
from tools.transcription_tools import _HAS_OPENAI, _resolve_openai_audio_client_config, _resolve_stt_language
if api_key is None:
try:
@@ -149,7 +145,6 @@ def _transcribe_openai(
create_kwargs["prompt"] = prompt
with open(path, "rb") as audio_file:
return client.audio.transcriptions.create(file=audio_file, **create_kwargs)
with tempfile.TemporaryDirectory(prefix="hermes-stt-") as work_dir:
try:
transcription = _create_transcription(file_path)
@@ -167,7 +162,6 @@ def _transcribe_openai(
logger.info("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)
@@ -197,7 +191,6 @@ def _transcribe_mistral(
# ---- REST multipart backends (xAI, ElevenLabs) ----------------------------
def _post_audio_multipart(url: str, headers: Dict[str, str], file_path: str, data: Dict[str, str]):
import requests
with open(file_path, "rb") as audio_file:
@@ -208,12 +201,10 @@ def _post_audio_multipart(url: str, headers: Dict[str, str], file_path: str, dat
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 multipart REST flow: ``post()`` -> ``log(text, body)`` -> ok envelope.
Non-200 -> ``"<label> API error (HTTP n): detail"`` (JSON detail via *extract_detail*, else
the first 300 body chars); empty text -> the ``no_speech`` envelope (silence is non-fatal);
exceptions -> ``_cloud_failure``.
"""
"""Shared multipart REST flow: ``post()`` -> ``log(text, body)`` -> ok envelope. Non-200 ->
``"<label> API error (HTTP n): detail"`` (JSON detail via *extract_detail*, else the first 300
body chars); empty text -> the ``no_speech`` envelope (silence is non-fatal); exceptions ->
``_cloud_failure``."""
try:
response = post()
if response.status_code != 200:
@@ -355,7 +346,6 @@ def _transcribe_deepinfra(
# ---- OpenAI audio credential resolution -----------------------------------
def _is_local_or_private_url(url: str) -> bool:
"""True for loopback/RFC-1918/LAN-internal hosts, where an empty ``stt.openai.api_key`` is acceptable
(local OpenAI-compatible servers ignore the auth header — no sham ``api_key: not-needed`` needed)."""
@@ -404,15 +394,16 @@ def _resolve_openai_audio_client_config() -> tuple[str, str]:
if selected == NOUS_MANAGED_PROVIDER:
managed = _managed()
if managed is None:
raise ValueError(selection_error(
"stt", NOUS_MANAGED_PROVIDER, "the Nous Tool Gateway is not available (not entitled or unreachable)"))
raise ValueError(selection_error("stt", NOUS_MANAGED_PROVIDER,
"the Nous Tool Gateway is not available (not entitled or unreachable)"))
return managed
direct = _direct_openai_credentials(openai_cfg.get("api_key", ""), openai_cfg.get("base_url", ""))
if direct is not None:
return direct
if selected is not None:
raise ValueError(selection_error(
"stt", selected, "neither stt.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set"))
"stt", selected,
"neither stt.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set"))
managed = _managed()
if managed is None:
message = "Neither stt.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set"
+13 -19
View File
@@ -54,7 +54,8 @@ _get_command_stt_output_format = partial(_command_output_format, formats=COMMAND
def _read_command_stt_output(output_path: Path, stdout: str, fmt: str) -> str:
"""Transcript: non-empty output file > non-empty stdout (curl one-liners) > RuntimeError. JSON is returned raw."""
content = output_path.read_bytes().decode("utf-8", errors="replace").strip() if output_path.exists() else ""
content = (output_path.read_bytes().decode("utf-8", errors="replace").strip()
if output_path.exists() else "")
if content or (stdout or "").strip():
return content or stdout.strip()
raise RuntimeError(f"Command STT provider wrote no output file at {output_path} and produced no stdout")
@@ -64,19 +65,16 @@ def _transcribe_command_stt(
file_path: str, provider_name: str, config: Dict[str, Any], stt_config: Dict[str, Any],
model_override: Optional[str] = None, language_override: Optional[str] = None,
prompt: Optional[str] = None) -> Dict[str, Any]:
"""Transcribe via a user-declared ``stt.providers.<name>: type: command``.
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).
"""
"""Transcribe via a user-declared ``stt.providers.<name>: type: command``. 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
if prompt:
_log_prompt_unsupported(f"Command STT provider '{provider_name}'")
def fail(error: str) -> Dict[str, Any]:
return _error_result(error, provider=provider_name)
command_template = str(config.get("command") or "").strip()
if not command_template:
return fail(f"stt.providers.{provider_name}.command is not configured")
@@ -128,12 +126,10 @@ def _dispatch_to_plugin_provider(
model: Optional[str] = None, language: Optional[str] = None, prompt: Optional[str] = None,
) -> Optional[Dict[str, Any]]:
"""Route to a plugin-registered transcription provider; None when no plugin claims the name.
Invariants re-verified here so a caller refactor can't break them: built-in names never
reach the registry; a same-name command provider wins over a plugin. A matched plugin with
Invariants re-verified here so a caller refactor can't break them: built-in names never reach
the registry; a same-name command provider wins over a plugin. A matched plugin with
``is_available() == False`` returns an error envelope — not None — because the user
explicitly opted in via ``stt.provider``. Provider exceptions become the error envelope.
"""
explicitly opted in via ``stt.provider``. Provider exceptions become the error envelope."""
key = (provider or "").lower().strip()
if not key or key in _NON_COMMAND_STT_NAMES:
return None
@@ -212,12 +208,10 @@ def _apply_pre_transcription_hook(
prompt: Optional[str], source: Optional[str],
) -> tuple[Optional[str], Optional[str], Optional[str]]:
"""Fire the ``pre_transcription`` plugin hook; returns ``(model, language_override, prompt)``.
Gated on ``has_hook`` (the no-hook path never builds kwargs) and fail-open: any
plumbing error leaves the dispatch untouched. Results apply field-by-field in
registration order (last hook wins). ``language_override`` is None unless a hook
explicitly set ``language``, so backends keep their own config/env resolution.
"""
Gated on ``has_hook`` (the no-hook path never builds kwargs) and fail-open: any plumbing error
leaves the dispatch untouched. Results apply field-by-field in registration order (last hook
wins). ``language_override`` is None unless a hook explicitly set ``language``, so backends keep
their own config/env resolution."""
try:
from hermes_cli.plugins import has_hook, invoke_hook
if not has_hook("pre_transcription"):
+2 -4
View File
@@ -61,10 +61,8 @@ def _ok_result(transcript: str, provider: str) -> Dict[str, Any]:
def _lazy_ensure_quietly(dep: str) -> None:
"""Best-effort ``tools.lazy_deps.ensure(dep, prompt=False)``; failures are swallowed.
prompt=False: a bare input() deadlocks under the interactive CLI where
prompt_toolkit owns stdin; installs are gated by ``security.allow_lazy_installs``.
"""
prompt=False: a bare input() deadlocks under the interactive CLI where prompt_toolkit owns
stdin; installs are gated by ``security.allow_lazy_installs``."""
try:
from tools.lazy_deps import ensure
ensure(dep, prompt=False)
+7 -13
View File
@@ -121,13 +121,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). Try the requested
config first; on a CUDA library load failure fall back to CPU + int8. Pass
``stt.local.device`` / ``compute_type`` to pin.
"""
"""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): try the requested config first; on a CUDA library load failure fall back to
CPU + int8. Pass ``stt.local.device`` / ``compute_type`` to pin."""
force_cpu = _should_force_faster_whisper_cpu()
if force_cpu:
# Importing ctranslate2 can itself abort on Apple Silicon/Rosetta when
@@ -194,11 +191,9 @@ 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): non-speech AND low decode
confidence, so quiet-but-real speech survives. Unknown segment shapes are never dropped.
"""
"""True when a segment is very likely a silence hallucination. 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:
return (float(segment.no_speech_prob) > no_speech_threshold
and float(segment.avg_logprob) < logprob_threshold)
@@ -227,7 +222,6 @@ def _transcribe_local_command(
from tools.transcription_tools import _prepare_local_audio, _resolve_stt_language
if prompt:
_log_prompt_unsupported("STT provider 'local_command'")
command_template = _get_local_command_template()
if not command_template:
return _error_result(f"{LOCAL_STT_COMMAND_ENV} not configured and no local whisper binary was found")
+18 -37
View File
@@ -104,7 +104,6 @@ _IDLE_UNLOAD_CHECK_INTERVAL = 30 # seconds between idle checks
# ---- Config helpers -----------------------------------------------------
def _load_stt_config() -> dict:
"""Load the ``stt`` section from user config, falling back to defaults."""
try:
@@ -122,11 +121,9 @@ def is_stt_enabled(stt_config: Optional[dict] = None) -> bool:
def _resolve_stt_language(
provider_key: str, stt_config: Optional[Dict[str, Any]] = None, *, extra_keys: tuple = ()
) -> Optional[str]:
"""Language hint for an STT provider, first non-empty wins (never "").
``stt.<provider>.language`` (plus *extra_keys* aliases, e.g. ``language_code``) >
``stt.language`` > ``HERMES_LOCAL_STT_LANGUAGE`` env > None (provider auto-detects).
"""
"""Language hint for an STT provider, first non-empty wins (never ""): ``stt.<provider>.language``
(plus *extra_keys* aliases, e.g. ``language_code``) > ``stt.language`` > ``HERMES_LOCAL_STT_LANGUAGE``
env > None (provider auto-detects)."""
if stt_config is None:
stt_config = _load_stt_config()
provider_cfg = _get_stt_section(stt_config, provider_key)
@@ -156,7 +153,6 @@ def _is_local_stt_provider(provider: str, stt_config: Dict[str, Any]) -> bool:
# ---- Provider resolution ------------------------------------------------
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:
@@ -293,7 +289,6 @@ def _get_provider(stt_config: dict) -> str:
# ---- Provider: local (faster-whisper) -----------------------------------
def _unload_local_model() -> None:
"""Release the cached local whisper model. Thread-safe via the model lock."""
global _local_model, _local_model_name
@@ -305,13 +300,10 @@ def _unload_local_model() -> None:
def _start_idle_unload_watcher(timeout_seconds: int) -> None:
"""Ensure the single idle-unload watcher thread is running.
The loop re-reads ``stt.local.unload_after_idle_seconds`` every cycle so config edits
apply within one interval; ``timeout_seconds`` seeds the first cycle so a just-written
config is honored even if a concurrent read races. The thread exits after unloading,
when the timeout becomes 0, or when the model is already gone.
"""
"""Ensure the single idle-unload watcher thread is running. 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. Exits after unloading, when the timeout becomes 0, or when the model is gone."""
global _idle_unload_thread
with _idle_unload_mgmt_lock:
if _idle_unload_thread is not None and _idle_unload_thread.is_alive():
@@ -328,7 +320,6 @@ def _start_idle_unload_watcher(timeout_seconds: int) -> None:
if time.monotonic() - _last_transcription_time >= timeout:
_unload_local_model()
break
_idle_unload_stop.clear()
_idle_unload_thread = threading.Thread(target=_watch, name="hermes-stt-idle-unload", daemon=True)
_idle_unload_thread.start()
@@ -341,10 +332,8 @@ def _touch_transcription_time() -> None:
def _get_or_load_local_model(model_name: str, local_cfg: Dict[str, Any]):
"""Cached faster-whisper model, (re)loaded under a double-checked lock when needed.
The returned strong reference stays valid even if the idle watcher nulls the global mid-transcription.
"""
"""Cached faster-whisper model, (re)loaded under a double-checked lock when needed. The returned
strong reference stays valid even if the idle watcher nulls the global mid-transcription."""
global _local_model, _local_model_name
model = _local_model
if model is None or _local_model_name != model_name:
@@ -385,7 +374,8 @@ def _transcribe_local(
return _error_result("Local whisper model failed to load")
# pre_transcription hook overrides win over config-resolved values.
transcribe_kwargs = build_local_transcribe_kwargs(stt_config)
transcribe_kwargs.update({k: v for k, v in (("language", language), ("initial_prompt", prompt)) if v})
transcribe_kwargs.update({k: v for k, v in (("language", language), ("initial_prompt", prompt))
if v})
try:
segments, info = model.transcribe(file_path, **transcribe_kwargs)
except Exception as exc:
@@ -411,7 +401,6 @@ def _transcribe_local(
# ---- Public API ---------------------------------------------------------
def _read_block_error(file_path: str) -> Optional[Dict[str, Any]]:
"""Refuse to ship a credential store (auth.json, .env, OAuth tokens) to an STT provider.
Mirrors the image-gen / video-gen read guards."""
@@ -422,11 +411,9 @@ def _read_block_error(file_path: str) -> Optional[Dict[str, Any]]:
def _transcribe_prepared_audio(
file_path: str, model: Optional[str] = None, source: Optional[str] = None) -> Dict[str, Any]:
"""Transcribe a validated audio file with the configured STT provider.
``model`` overrides the config default; ``source`` is a caller-surface label
(``"gateway"``, ``"voice_mode"``) forwarded to the ``pre_transcription`` hook only.
"""
"""Transcribe a validated audio file with the configured STT provider. ``model`` overrides the
config default; ``source`` is a caller-surface label (``"gateway"``, ``"voice_mode"``) forwarded
to the ``pre_transcription`` hook only."""
# Validate before provider resolution so invalid files can't trigger provider setup
# or lazy installation; the remote-upload size cap applies to non-local only.
error = _read_block_error(file_path) or _validate_audio_file(file_path, enforce_size_limit=False)
@@ -536,11 +523,8 @@ def _no_provider_error(provider: str, stt_config: Dict[str, Any]) -> Dict[str, A
def transcribe_audio(
file_path: str, model: Optional[str] = None, source: Optional[str] = None) -> Dict[str, Any]:
"""Validate, preprocess supported inputs, and dispatch transcription.
``source`` is a caller-surface label (``"gateway"``, ``"voice_mode"``) forwarded to
the ``pre_transcription`` hook for observability only.
"""
"""Validate, preprocess supported inputs, and dispatch transcription. ``source`` is a caller-surface
label (``"gateway"``, ``"voice_mode"``) forwarded to the ``pre_transcription`` hook only."""
# Secret-store refusal runs before ANY validation so the error names the real reason.
blocked = _read_block_error(file_path)
if blocked:
@@ -563,11 +547,8 @@ def transcribe_audio(
def transcribe_audio_local_fallback(file_path: str, model: Optional[str] = None) -> Dict[str, Any]:
"""Try an already-installed local STT backend without changing config.
Passive inbound-media recovery after the configured provider failed: never
lazy-installs or falls through to a cloud provider.
"""
"""Try an already-installed local STT backend without changing config: passive inbound-media
recovery after the configured provider failed — never lazy-installs or falls through to cloud."""
error = _validate_audio_file(file_path)
if error:
return error
+9 -8
View File
@@ -111,7 +111,8 @@ def terminate_command_process_tree(proc: subprocess.Popen) -> None:
except ImportError:
psutil = None
# Without psutil only the shell itself is signalled (children may survive).
signal = (lambda m: getattr(proc, m)()) if psutil is None else (lambda m: _signal_process_tree(psutil, proc, m))
signal = ((lambda m: getattr(proc, m)()) if psutil is None
else (lambda m: _signal_process_tree(psutil, proc, m)))
signal("terminate")
try:
proc.wait(timeout=2)
@@ -122,7 +123,7 @@ def terminate_command_process_tree(proc: subprocess.Popen) -> None:
def command_env_passthrough(config: Dict[str, Any]) -> list:
"""``env_passthrough`` allowlist: parent env vars copied back into the secret-scrubbed child env."""
raw = config.get("env_passthrough")
return [str(item).strip() for item in raw if str(item).strip()] if isinstance(raw, (list, tuple)) else []
return [str(x).strip() for x in raw if str(x).strip()] if isinstance(raw, (list, tuple)) else []
def command_failure_detail(exc: subprocess.CalledProcessError) -> str:
@@ -211,7 +212,6 @@ def run_command_provider(
# ---- Generic ``<section>.providers.<name>`` config layer (TTS and STT share it) ----
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."""
section = config.get(name) if isinstance(config, dict) else None
@@ -219,8 +219,8 @@ def _get_provider_section(config: Dict[str, Any], name: str) -> Dict[str, Any]:
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 —
refused for built-ins so a user's ``openai:`` block still means the OpenAI provider, not a command."""
"""``<section>.providers.<name>`` (canonical), else ``<section>.<name>`` for non-built-in names
only — refused for built-ins so a user's ``openai:`` block still means OpenAI, not a command."""
section = _get_provider_section(config, "providers").get(name)
if isinstance(section, dict):
return section
@@ -310,7 +310,7 @@ def _generate_command_tts(
text: str, output_path: str, provider_name: str, config: Dict[str, Any], tts_config: Dict[str, Any],
) -> str:
"""Generate speech by running a user-configured shell command; returns the audio path it wrote.
Raises ``ValueError`` for invalid provider config, ``RuntimeError`` for timeouts / non-zero exits / empty output."""
Raises ``ValueError`` for bad provider config, ``RuntimeError`` for timeouts / bad exits / no output."""
command_template = str(config.get("command") or "").strip()
if not command_template:
raise ValueError(f"tts.providers.{provider_name}.command is not configured")
@@ -324,8 +324,9 @@ def _generate_command_tts(
text_path.write_text(text, encoding="utf-8")
placeholders = {
"input_path": str(text_path), "text_path": str(text_path), "output_path": str(output),
"format": _get_command_tts_output_format(config, str(output)), "voice": str(config.get("voice", "")),
"model": str(config.get("model", "")), "speed": str(config.get("speed", tts_config.get("speed", ""))),
"format": _get_command_tts_output_format(config, str(output)),
"voice": str(config.get("voice", "")), "model": str(config.get("model", "")),
"speed": str(config.get("speed", tts_config.get("speed", ""))),
}
command = render_command_template(command_template, placeholders)
try:
+5 -3
View File
@@ -96,7 +96,7 @@ class SentenceChunker:
class StreamingTTSProvider(ABC):
"""Yields raw int16, little-endian, mono PCM chunks at ``sample_rate`` (built-in streamers: 24 kHz)."""
"""Yields raw int16, little-endian, mono PCM chunks at ``sample_rate`` (built-ins: 24 kHz)."""
sample_rate: int = 24000
channels: int = 1
@@ -153,7 +153,8 @@ def resolve_streaming_provider(
providers just to get streaming."""
pinned = str((tts_config.get("streaming") or {}).get("provider") or "").lower().strip()
if pinned == "auto":
return next((inst for name in _PROVIDER_PRIORITY if (inst := _try_instantiate(name, tts_config))), None)
return next((inst for name in _PROVIDER_PRIORITY
if (inst := _try_instantiate(name, tts_config))), None)
return _try_instantiate(pinned or (preferred or _get_provider(tts_config)).lower().strip(), tts_config)
@@ -274,7 +275,8 @@ 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. ``_collect_async`` bridges the async WS loop to the sync iterator contract (test seam)."""
sync path. ``_collect_async`` bridges the async WS loop to the sync iterator contract (test
seam)."""
@staticmethod
def available() -> bool: