448 lines
17 KiB
Python
448 lines
17 KiB
Python
"""User-declared STT command providers (stt.providers.<name>: type: command) and plugin-registered TranscriptionProvider dispatch, plus the pre_transcription hook that threads prompt/language/model overrides into every backend.
|
|
|
|
Split out of ``tools/transcription_tools.py``; every moved name is re-imported there,
|
|
so ``tools.transcription_tools.<name>`` keeps resolving (and monkeypatching) as before.
|
|
Origin helpers are imported lazily inside functions so patches on the origin intercept.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import subprocess
|
|
import tempfile
|
|
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,
|
|
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,
|
|
)
|
|
|
|
# Log-record parity with the origin module.
|
|
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.
|
|
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"})
|
|
|
|
|
|
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 _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 _get_command_stt_timeout(config: Dict[str, Any]) -> float:
|
|
"""Return timeout in seconds, falling back when invalid."""
|
|
raw = config.get("timeout", config.get("timeout_seconds", DEFAULT_COMMAND_STT_TIMEOUT_SECONDS))
|
|
try:
|
|
value = float(raw)
|
|
except (TypeError, ValueError):
|
|
return float(DEFAULT_COMMAND_STT_TIMEOUT_SECONDS)
|
|
return value if value > 0 else float(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
|
|
|
|
|
|
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.
|
|
"""
|
|
if output_path.exists():
|
|
try:
|
|
content = output_path.read_text(encoding="utf-8").strip()
|
|
except UnicodeDecodeError:
|
|
content = output_path.read_bytes().decode("utf-8", errors="replace").strip()
|
|
if content:
|
|
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"
|
|
)
|
|
|
|
|
|
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 (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,
|
|
``{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")
|
|
|
|
audio = Path(file_path).expanduser()
|
|
if not audio.exists():
|
|
return fail(f"Audio file not found: {file_path}")
|
|
|
|
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
|
|
)
|
|
model = model_override or config.get("model") or ""
|
|
|
|
try:
|
|
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),
|
|
}
|
|
command = _render_command_stt_template(command_template, placeholders)
|
|
logger.info(
|
|
"Transcribing %s via command STT provider '%s'...",
|
|
audio.name, provider_name,
|
|
)
|
|
try:
|
|
result = _run_command_stt(
|
|
command,
|
|
timeout,
|
|
env_passthrough=_command_stt_env_passthrough(config),
|
|
)
|
|
except subprocess.TimeoutExpired:
|
|
return fail(f"STT command provider '{provider_name}' timed out after {timeout:g}s")
|
|
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"
|
|
return fail(
|
|
f"STT command provider '{provider_name}' exited with code "
|
|
f"{exc.returncode}: {detail}"
|
|
)
|
|
|
|
try:
|
|
transcript_text = _read_command_stt_output(
|
|
output_path, result.stdout or "", output_format,
|
|
)
|
|
except RuntimeError as exc:
|
|
return fail(str(exc))
|
|
|
|
except OSError as exc:
|
|
return fail(f"STT command provider '{provider_name}' failed: {exc}")
|
|
|
|
logger.info(
|
|
"Transcribed %s via command STT provider '%s' (%d chars)",
|
|
audio.name, provider_name, len(transcript_text),
|
|
)
|
|
return _ok_result(transcript_text, provider_name)
|
|
|
|
|
|
def _unregistered_stt_provider_error(provider: str) -> Dict[str, Any]:
|
|
key = str(provider or "").strip()
|
|
return _error_result(
|
|
f"stt.provider='{key}' is set but no built-in, command, or plugin "
|
|
"provider registered that name. Run `hermes plugins list` to see "
|
|
"installed STT plugins, or configure a command provider under "
|
|
f"`stt.providers.{key}.command`.",
|
|
provider=key,
|
|
error_type="provider_not_registered",
|
|
)
|
|
|
|
|
|
def _dispatch_to_plugin_provider(
|
|
file_path: str,
|
|
provider: str,
|
|
stt_config: Optional[Dict[str, Any]] = None,
|
|
*,
|
|
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 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.
|
|
"""
|
|
if not provider:
|
|
return None
|
|
key = provider.lower().strip()
|
|
if key in BUILTIN_STT_PROVIDERS or key == "none":
|
|
return None
|
|
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
|
|
from hermes_cli.plugins import _ensure_plugins_discovered
|
|
|
|
_ensure_plugins_discovered()
|
|
plugin_provider = get_provider(key)
|
|
if plugin_provider is None:
|
|
# Long-lived sessions may have discovered plugins before a backend
|
|
# was patched in or config changed — retry once with a forced refresh.
|
|
_ensure_plugins_discovered(force=True)
|
|
plugin_provider = get_provider(key)
|
|
except Exception as exc: # noqa: BLE001 — discovery failure is non-fatal
|
|
logger.debug("STT plugin dispatch skipped (discovery failed): %s", exc)
|
|
return None
|
|
if plugin_provider is None:
|
|
return None
|
|
|
|
# ``is_available()`` MUST NOT raise per the ABC contract; defend anyway so
|
|
# a buggy plugin can't break dispatch for everyone.
|
|
try:
|
|
available = plugin_provider.is_available()
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning(
|
|
"STT plugin provider '%s' is_available() raised: %s — "
|
|
"treating as unavailable", key, exc, exc_info=True,
|
|
)
|
|
available = False
|
|
if not available:
|
|
logger.info(
|
|
"STT plugin provider '%s' reports not available; returning "
|
|
"unavailability envelope.", key,
|
|
)
|
|
return _error_result(
|
|
f"STT plugin '{key}' is not available — check that its "
|
|
"required credentials / dependencies are configured.",
|
|
provider=key,
|
|
)
|
|
|
|
logger.info("Transcribing with plugin STT provider '%s'...", key)
|
|
# The prompt travels via the ABC's ``**extra`` kwargs and is only sent when
|
|
# set, so pre-prompt providers see byte-identical calls on the no-prompt path.
|
|
extra_kwargs: Dict[str, Any] = {}
|
|
if prompt is not None:
|
|
extra_kwargs["prompt"] = prompt
|
|
try:
|
|
result = plugin_provider.transcribe(
|
|
file_path,
|
|
model=model,
|
|
language=language,
|
|
**extra_kwargs,
|
|
)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning(
|
|
"STT plugin provider '%s' raised: %s", key, exc, exc_info=True,
|
|
)
|
|
return _error_result(f"STT plugin '{key}' raised: {exc}", provider=key)
|
|
|
|
if not isinstance(result, dict):
|
|
return _error_result(f"STT plugin '{key}' returned a non-dict result", provider=key)
|
|
result.setdefault("provider", key)
|
|
return result
|
|
|
|
|
|
# Fields a pre_transcription hook may mutate. ``file_path`` is read-only —
|
|
# 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_PROMPT_TOKEN_CAP = 224
|
|
|
|
|
|
_PROMPT_CHARS_PER_TOKEN = 4
|
|
|
|
|
|
_WHISPER_PROMPT_CAPPED_PROVIDERS = frozenset(
|
|
{"local", "openai", "groq", "deepinfra"}
|
|
)
|
|
|
|
|
|
def _enforce_prompt_length_limit(
|
|
prompt: Optional[str], provider: str
|
|
) -> Optional[str]:
|
|
"""Truncate *prompt* to the whisper-family token cap, keeping the TAIL (fail-open).
|
|
|
|
Whisper conditions on the final context window, so the most recently
|
|
appended hints survive. Other providers own their own validation.
|
|
"""
|
|
if not prompt or provider not in _WHISPER_PROMPT_CAPPED_PROVIDERS:
|
|
return prompt
|
|
max_chars = _WHISPER_PROMPT_TOKEN_CAP * _PROMPT_CHARS_PER_TOKEN
|
|
if len(prompt) <= max_chars:
|
|
return prompt
|
|
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,
|
|
)
|
|
return prompt[-max_chars:]
|
|
|
|
|
|
def _apply_pre_transcription_hook(
|
|
*,
|
|
file_path: str,
|
|
provider: str,
|
|
model: Optional[str],
|
|
language: Optional[str],
|
|
prompt: Optional[str],
|
|
source: Optional[str],
|
|
) -> tuple[Optional[str], Optional[str], Optional[str]]:
|
|
"""Fire the ``pre_transcription`` plugin hook and merge its results.
|
|
|
|
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
|
|
normalization a caller-supplied model would.
|
|
|
|
Returns ``(model, language_override, prompt)``; ``language_override`` is
|
|
None unless a hook explicitly set ``language``, so backends keep their own
|
|
config/env language resolution.
|
|
"""
|
|
try:
|
|
from hermes_cli.plugins import has_hook, invoke_hook
|
|
|
|
if not has_hook("pre_transcription"):
|
|
return model, None, prompt
|
|
|
|
hook_results = invoke_hook(
|
|
"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:
|
|
if not isinstance(hook_result, dict):
|
|
continue
|
|
for key, value in hook_result.items():
|
|
if key == "file_path":
|
|
logger.warning(
|
|
"pre_transcription hook attempted to change "
|
|
"file_path (read-only) — ignoring the attempt."
|
|
)
|
|
continue
|
|
if key not in _PRE_TRANSCRIPTION_MUTABLE_FIELDS:
|
|
logger.debug(
|
|
"pre_transcription hook returned unsupported field "
|
|
"%r — ignoring.", key,
|
|
)
|
|
continue
|
|
if not isinstance(value, str):
|
|
logger.debug(
|
|
"pre_transcription hook returned non-string value "
|
|
"%r for field %r — ignoring.", value, key,
|
|
)
|
|
continue
|
|
overrides[key] = value
|
|
|
|
if "model" in overrides:
|
|
model = overrides["model"]
|
|
if "prompt" in overrides:
|
|
# Hooks win over the static ``stt.prompt`` config; "" clears it.
|
|
prompt = overrides["prompt"] or None
|
|
return model, overrides.get("language") or None, prompt
|
|
except Exception as _hook_err: # noqa: BLE001 — hook plumbing is fail-open
|
|
logger.debug("pre_transcription hook error: %s", _hook_err)
|
|
return model, None, prompt
|