refactor(tools): fold command/plugin STT and audio prep bodies — inline single-use locals, merged guards, squeezed blanks

This commit is contained in:
Teknium
2026-09-03 00:48:24 -07:00
parent 35c32f6498
commit 97b554f200
2 changed files with 30 additions and 81 deletions

View File

@@ -65,11 +65,8 @@ _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."""
command = [ffmpeg, "-y", "-i", input_path]
if audio_filter:
command += ["-af", audio_filter]
command += [*_STT_M4A_ENCODE_ARGS, output_path]
_run_quiet(command, timeout=120)
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)
def _transcode_audio_for_stt(file_path: str, work_dir: str) -> tuple[Optional[str], Optional[str]]:
@@ -121,13 +118,10 @@ def _validate_audio_source_file(file_path: str, *, enforce_size_limit: bool = Tr
def _validate_audio_file(file_path: str, *, enforce_size_limit: bool = True) -> Optional[Dict[str, Any]]:
"""Validate a supported, decoder-safe audio file."""
source_error = _validate_audio_source_file(file_path, enforce_size_limit=enforce_size_limit)
if source_error:
return source_error
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 None
if source_error or suffix.lower() in SUPPORTED_FORMATS:
return source_error
return _error_result(f"Unsupported format: {suffix}. Supported: {', '.join(sorted(SUPPORTED_FORMATS))}")
def _prepare_audio_for_transcription(file_path: str) -> tuple[Optional[str], Optional[str], Optional[Dict[str, Any]]]:
@@ -143,12 +137,10 @@ def _prepare_audio_for_transcription(file_path: str) -> tuple[Optional[str], Opt
return None, None, _error_result(
"Unsupported format: .silk. Install the optional 'pilk' dependency to enable WeChat voice transcription."
)
temp_dir = tempfile.mkdtemp(prefix="hermes-silk-")
converted_path = os.path.join(temp_dir, f"{audio_path.stem}.wav")
try:
import pilk
pilk.silk_to_wav(file_path, converted_path)
if not Path(converted_path).is_file() or Path(converted_path).stat().st_size == 0:
raise RuntimeError("pilk did not produce a readable WAV file")
@@ -165,11 +157,9 @@ def _prepare_local_audio(file_path: str, work_dir: str) -> tuple[Optional[str],
audio_path = Path(file_path)
if audio_path.suffix.lower() in LOCAL_NATIVE_AUDIO_FORMATS:
return file_path, None
ffmpeg = _find_ffmpeg_binary()
if not ffmpeg:
return None, "Local STT fallback requires ffmpeg for non-WAV inputs, but ffmpeg was not found"
converted_path = os.path.join(work_dir, f"{audio_path.stem}.wav")
try:
_run_quiet([ffmpeg, "-y", "-i", file_path, converted_path], timeout=300)
@@ -194,9 +184,7 @@ def _convert_caf_to_wav(file_path: str) -> Optional[str]:
("ffmpeg", [ffmpeg, "-y", "-i", file_path, wav_path] if ffmpeg else None),
("afconvert", [afconvert, file_path, wav_path, "-d", "LEI16", "-f", "WAVE"] if afconvert else None),
)
for label, command in candidates:
if not command:
continue
for label, command in ((label, cmd) for label, cmd in candidates if cmd):
try:
_run_quiet(command, timeout=300)
return wav_path
@@ -236,11 +224,9 @@ def _probe_audio_duration(file_path: str) -> Optional[float]:
ffprobe = _find_ffprobe_binary()
if not ffprobe:
return None
command = [
ffprobe, "-v", "error", "-show_entries", "format=duration", "-of", "default=noprint_wrappers=1:nokey=1", file_path,
]
try:
return float(_run_quiet(command, timeout=30).stdout.strip())
return float(_run_quiet([ffprobe, "-v", "error", "-show_entries", "format=duration", "-of",
"default=noprint_wrappers=1:nokey=1", file_path], timeout=30).stdout.strip())
except Exception: # noqa: BLE001 - probe is best-effort
return None
@@ -274,12 +260,9 @@ def _trim_silence_for_cloud_stt(file_path: str, stt_config: Dict[str, Any]) -> O
return None
name = Path(file_path).name
if original_duration < _CLOUD_TRIM_MIN_INPUT_SECONDS:
logger.debug(
"Cloud STT silence trim skipped for %s: %.1fs is below the %.0fs gate",
name, original_duration, _CLOUD_TRIM_MIN_INPUT_SECONDS,
)
logger.debug("Cloud STT silence trim skipped for %s: %.1fs is below the %.0fs gate",
name, original_duration, _CLOUD_TRIM_MIN_INPUT_SECONDS)
return None
keep_seconds = keep_ms / 1000.0
# start_periods=1 strips leading silence; stop_periods=-1 collapses every interior/trailing silence.
filter_expr = (
@@ -296,20 +279,15 @@ def _trim_silence_for_cloud_stt(file_path: str, stt_config: Dict[str, Any]) -> O
_run_ffmpeg_stt_encode(ffmpeg, file_path, trimmed_path, audio_filter=filter_expr)
trimmed_duration = _probe_audio_duration(trimmed_path)
if not trimmed_duration or trimmed_duration < min_result_seconds:
logger.debug(
"Cloud STT silence trim discarded for %s: trimmed result ~empty (%.2fs)", name, trimmed_duration or 0.0,
)
logger.debug("Cloud STT silence trim discarded for %s: trimmed result ~empty (%.2fs)",
name, trimmed_duration or 0.0)
return None
if trimmed_duration > original_duration * (1 - _CLOUD_TRIM_MIN_SAVING):
logger.debug(
"Cloud STT silence trim discarded for %s: saves <%.0f%% (%.1fs -> %.1fs)",
name, _CLOUD_TRIM_MIN_SAVING * 100, original_duration, trimmed_duration,
)
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

View File

@@ -63,15 +63,9 @@ def _get_command_stt_output_format(config: Dict[str, Any]) -> str:
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."""
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()
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")
@@ -96,28 +90,21 @@ def _transcribe_command_stt(
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 ""
language = (language_override or config.get("language")
or _resolve_stt_language(provider_name, stt_config) or DEFAULT_COMMAND_STT_LANGUAGE)
try:
with tempfile.TemporaryDirectory(prefix=f"hermes-cmd-stt-{provider_name}-") as tmpdir:
output_path = Path(tmpdir) / f"transcript.{output_format}"
placeholders = {
command = _render_command_stt_template(command_template, {
"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)
"language": str(language), "model": str(model_override or config.get("model") or ""),
})
logger.info("Transcribing %s via command STT provider '%s'...", audio.name, provider_name)
result = _run_command_stt(command, timeout, env_passthrough=_command_stt_env_passthrough(config))
transcript_text = _read_command_stt_output(output_path, result.stdout or "", output_format)
@@ -131,7 +118,6 @@ def _transcribe_command_stt(
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)
@@ -159,17 +145,14 @@ def _dispatch_to_plugin_provider(
``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 _NON_COMMAND_STT_NAMES:
key = (provider or "").lower().strip()
if not key or 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)):
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:
@@ -182,7 +165,6 @@ def _dispatch_to_plugin_provider(
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:
@@ -198,7 +180,6 @@ def _dispatch_to_plugin_provider(
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.
@@ -208,7 +189,6 @@ def _dispatch_to_plugin_provider(
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)
@@ -229,10 +209,8 @@ _WHISPER_PROMPT_CAPPED_PROVIDERS = frozenset({"local", "openai", "groq", "deepin
def _enforce_prompt_length_limit(prompt: Optional[str], provider: str) -> Optional[str]:
"""Truncate *prompt* to the whisper-family token cap, keeping the TAIL (whisper conditions
on the final context window, so the newest hints survive). Other providers self-validate."""
if not prompt or provider not in _WHISPER_PROMPT_CAPPED_PROVIDERS:
return prompt
max_chars = _WHISPER_PROMPT_TOKEN_CAP * _PROMPT_CHARS_PER_TOKEN
if len(prompt) <= max_chars:
if not prompt or provider not in _WHISPER_PROMPT_CAPPED_PROVIDERS or len(prompt) <= max_chars:
return prompt
logger.warning(
"Transcription prompt is ~%d tokens; whisper-family provider '%s' "
@@ -255,19 +233,15 @@ def _apply_pre_transcription_hook(
"""
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():
for key, value in (hook_result.items() if isinstance(hook_result, dict) else ()):
if key == "file_path":
logger.warning(
"pre_transcription hook attempted to change "
@@ -281,13 +255,10 @@ def _apply_pre_transcription_hook(
)
else:
overrides[key] = value
if "model" in overrides:
model = overrides["model"]
# Hooks win over the static ``stt.prompt`` config; "" clears it.
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
return overrides.get("model", 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