From 97b554f2005aa3a0b36138b79f19e7b3d13b7cef Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 00:48:24 -0700 Subject: [PATCH] =?UTF-8?q?refactor(tools):=20fold=20command/plugin=20STT?= =?UTF-8?q?=20and=20audio=20prep=20bodies=20=E2=80=94=20inline=20single-us?= =?UTF-8?q?e=20locals,=20merged=20guards,=20squeezed=20blanks?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/transcription_audio.py | 54 ++++++++++---------------------- tools/transcription_command.py | 57 +++++++++------------------------- 2 files changed, 30 insertions(+), 81 deletions(-) diff --git a/tools/transcription_audio.py b/tools/transcription_audio.py index 15413f370a..51935a2b80 100644 --- a/tools/transcription_audio.py +++ b/tools/transcription_audio.py @@ -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 diff --git a/tools/transcription_command.py b/tools/transcription_command.py index 52e00c3dea..f26767163d 100644 --- a/tools/transcription_command.py +++ b/tools/transcription_command.py @@ -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