refactor(stt): shared lazy-install helper, compact trim/validation/local-kwargs bodies, repack re-export block

This commit is contained in:
Teknium
2026-09-02 16:04:57 -07:00
parent 5844839639
commit 063517e745
5 changed files with 66 additions and 153 deletions

View File

@@ -23,6 +23,7 @@ from tools.transcription_common import (
SUPPORTED_FORMATS,
_config_number,
_error_result,
_lazy_ensure_quietly,
_process_error_detail,
)
@@ -68,9 +69,7 @@ _STT_M4A_ENCODE_ARGS = (
)
def _run_ffmpeg_stt_encode(
ffmpeg: str, input_path: str, output_path: str, *, audio_filter: Optional[str] = None
) -> None:
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).
@@ -119,11 +118,7 @@ def _validate_audio_file_size(audio_path: Path, *, enforce_size_limit: bool = Tr
return None
def _validate_audio_source_file(
file_path: str,
*,
enforce_size_limit: bool = True,
) -> Optional[Dict[str, Any]]:
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)
@@ -136,15 +131,9 @@ def _validate_audio_source_file(
return _validate_audio_file_size(audio_path, enforce_size_limit=enforce_size_limit)
def _validate_audio_file(
file_path: str,
*,
enforce_size_limit: bool = True,
) -> Optional[Dict[str, Any]]:
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
)
source_error = _validate_audio_source_file(file_path, enforce_size_limit=enforce_size_limit)
if source_error:
return source_error
@@ -156,9 +145,7 @@ def _validate_audio_file(
return None
def _prepare_audio_for_transcription(
file_path: str,
) -> tuple[Optional[str], Optional[str], Optional[Dict[str, Any]]]:
def _prepare_audio_for_transcription(file_path: str) -> tuple[Optional[str], Optional[str], Optional[Dict[str, Any]]]:
"""Convert a decoder-safe .silk source to a temporary supported WAV file."""
from tools.transcription_tools import _HAS_PILK, _safe_find_spec
audio_path = Path(file_path)
@@ -167,11 +154,7 @@ def _prepare_audio_for_transcription(
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.
try:
from tools.lazy_deps import ensure as _lazy_ensure
_lazy_ensure("stt.silk", prompt=False)
except Exception:
pass
_lazy_ensure_quietly("stt.silk")
if not _safe_find_spec("pilk"):
return None, None, _error_result(
"Unsupported format: .silk. Install the optional 'pilk' dependency to enable WeChat voice transcription."
@@ -294,9 +277,7 @@ def _cloud_trim_settings(stt_config: Dict[str, Any]) -> tuple[bool, int, int]:
return enabled, threshold_db, max(keep_ms, 0)
def _trim_silence_for_cloud_stt(
file_path: str, stt_config: Dict[str, Any]
) -> Optional[str]:
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.
``None`` always means "upload the original" (disabled, tools missing, clip
@@ -315,10 +296,11 @@ def _trim_silence_for_cloud_stt(
if not original_duration or original_duration <= 0:
logger.debug("Cloud STT silence trim skipped: could not probe %s", file_path)
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",
Path(file_path).name, original_duration, _CLOUD_TRIM_MIN_INPUT_SECONDS,
name, original_duration, _CLOUD_TRIM_MIN_INPUT_SECONDS,
)
return None
@@ -341,21 +323,18 @@ def _trim_silence_for_cloud_stt(
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)",
Path(file_path).name, trimmed_duration or 0.0,
"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)",
Path(file_path).name, _CLOUD_TRIM_MIN_SAVING * 100,
original_duration, trimmed_duration,
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%%)",
Path(file_path).name, original_duration, trimmed_duration,
round((1 - trimmed_duration / original_duration) * 100),
name, original_duration, trimmed_duration, round((1 - trimmed_duration / original_duration) * 100),
)
keep_result = True
return trimmed_path

View File

@@ -26,6 +26,7 @@ from tools.transcription_common import (
XAI_STT_BASE_URL,
_error_result,
_get_stt_section,
_lazy_ensure_quietly,
_log_prompt_unsupported,
_ok_result,
)
@@ -247,11 +248,7 @@ def _transcribe_mistral(
return _error_result("MISTRAL_API_KEY not set")
try:
try:
from tools.lazy_deps import ensure as _lazy_ensure
_lazy_ensure("stt.mistral", prompt=False)
except Exception:
pass
_lazy_ensure_quietly("stt.mistral")
from mistralai.client import Mistral
with Mistral(api_key=api_key) as client:

View File

@@ -67,6 +67,19 @@ def _get_stt_section(stt_config: Dict[str, Any], name: str) -> Dict[str, Any]:
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.
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)
except Exception:
pass
def _process_error_detail(exc: "subprocess.CalledProcessError") -> str:
"""stderr > stdout > str(exc) for a failed helper binary."""
return exc.stderr.strip() or exc.stdout.strip() or str(exc)

View File

@@ -61,8 +61,7 @@ def _normalize_local_model(model_name: Optional[str]) -> str:
"STT model '%s' is a cloud-only name and cannot be used with the local "
"provider. Falling back to '%s'. Set stt.local.model to a valid "
"faster-whisper size (tiny, base, small, medium, large-v3).",
model_name,
DEFAULT_LOCAL_MODEL,
model_name, DEFAULT_LOCAL_MODEL,
)
return DEFAULT_LOCAL_MODEL
return model_name
@@ -78,10 +77,7 @@ def _try_lazy_install_stt() -> bool:
ensure("stt.faster_whisper", prompt=False)
if _ilu.find_spec("faster_whisper"):
return True
logger.warning(
"faster-whisper was installed but importlib still cannot find it "
"(may require Python restart)"
)
logger.warning("faster-whisper was installed but importlib still cannot find it (may require Python restart)")
except Exception as exc:
logger.warning(
"Lazy install of faster-whisper failed: %s. "
@@ -199,32 +195,24 @@ def build_local_transcribe_kwargs(stt_config: Optional[Dict[str, Any]] = None) -
stt_config = stt_config if isinstance(stt_config, dict) else _load_stt_config()
local_cfg = stt_config.get("local") or {}
# ``vad: null`` in YAML means "default on".
vad_enabled = local_cfg.get("vad", True)
kwargs: Dict[str, Any] = {
"beam_size": 5,
"condition_on_previous_text": False,
"vad_filter": vad_enabled is None or bool(vad_enabled),
}
vad_enabled = local_cfg.get("vad", True)
if vad_enabled is None:
vad_enabled = True
if bool(vad_enabled):
kwargs["vad_filter"] = True
if kwargs["vad_filter"]:
kwargs["vad_parameters"] = {
"min_silence_duration_ms": _config_number(
local_cfg, "vad_min_silence_ms", _VAD_MIN_SILENCE_MS_DEFAULT, int
)
"min_silence_duration_ms": _config_number(local_cfg, "vad_min_silence_ms", _VAD_MIN_SILENCE_MS_DEFAULT, int)
}
else:
kwargs["vad_filter"] = False
# 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.
no_speech_threshold, log_prob_threshold = _confidence_thresholds(local_cfg)
kwargs["no_speech_threshold"] = no_speech_threshold
kwargs["log_prob_threshold"] = log_prob_threshold
kwargs["no_speech_threshold"], kwargs["log_prob_threshold"] = _confidence_thresholds(local_cfg)
forced_lang = _resolve_stt_language("local", stt_config)
if forced_lang:
@@ -252,14 +240,10 @@ def _is_hallucinated_segment(segment: Any, no_speech_threshold: float, logprob_t
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.
"""
no_speech_prob = getattr(segment, "no_speech_prob", None)
avg_logprob = getattr(segment, "avg_logprob", None)
if no_speech_prob is None or avg_logprob is None:
return False
try:
no_speech_prob = float(no_speech_prob)
avg_logprob = float(avg_logprob)
except (TypeError, ValueError):
no_speech_prob = float(getattr(segment, "no_speech_prob"))
avg_logprob = float(getattr(segment, "avg_logprob"))
except (AttributeError, TypeError, ValueError):
return False
return no_speech_prob > no_speech_threshold and avg_logprob < logprob_threshold
@@ -295,9 +279,7 @@ def _transcribe_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"
)
return _error_result(f"{LOCAL_STT_COMMAND_ENV} not configured and no local whisper binary was found")
# Language: hook override > stt.local.language > stt.language > env > "en".
language = language or _resolve_stt_language("local") or DEFAULT_LOCAL_STT_LANGUAGE
@@ -318,10 +300,7 @@ def _transcribe_local_command(
# Scrub Hermes secrets from the child env (same policy as _run_command_stt).
from tools.environments.local import hermes_subprocess_env
_run_quiet(
shlex.split(command), timeout=300,
env=hermes_subprocess_env(inherit_credentials=False),
)
_run_quiet(shlex.split(command), timeout=300, env=hermes_subprocess_env(inherit_credentials=False))
txt_files = sorted(Path(output_dir).glob("*.txt"))
if not txt_files:
@@ -330,9 +309,7 @@ def _transcribe_local_command(
transcript_text = txt_files[0].read_text(encoding="utf-8").strip()
logger.info(
"Transcribed %s via local STT command (%s, %d chars)",
Path(file_path).name,
normalized_model,
len(transcript_text),
Path(file_path).name, normalized_model, len(transcript_text),
)
return _ok_result(transcript_text, "local_command")

View File

@@ -23,94 +23,41 @@ from hermes_cli._subprocess_compat import windows_hide_flags # noqa: F401 (imp
from utils import is_truthy_value
from tools.managed_tool_gateway import resolve_managed_tool_gateway # noqa: F401 (patched by tests)
from tools.tool_backend_helpers import ( # noqa: F401 (patched by tests; read lazily by transcription_cloud)
managed_nous_tools_enabled,
nous_tool_gateway_unavailable_message,
resolve_openai_audio_api_key,
managed_nous_tools_enabled, nous_tool_gateway_unavailable_message, resolve_openai_audio_api_key,
)
from tools.transcription_common import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
BUILTIN_STT_PROVIDERS,
CLOUD_STT_PROVIDERS,
DEFAULT_ELEVENLABS_STT_MODEL,
DEFAULT_GROQ_STT_MODEL,
DEFAULT_LOCAL_MODEL,
DEFAULT_MISTRAL_STT_MODEL,
DEFAULT_PROVIDER,
DEFAULT_STT_MODEL,
ELEVENLABS_STT_BASE_URL,
GROQ_MODELS,
LOCAL_STT_COMMAND_ENV,
LOCAL_STT_LANGUAGE_ENV,
MAX_FILE_SIZE,
OPENAI_MODELS,
SUPPORTED_FORMATS,
XAI_STT_BASE_URL,
_error_result,
_get_stt_section,
_ok_result,
BUILTIN_STT_PROVIDERS, CLOUD_STT_PROVIDERS, DEFAULT_ELEVENLABS_STT_MODEL,
DEFAULT_GROQ_STT_MODEL, DEFAULT_LOCAL_MODEL, DEFAULT_MISTRAL_STT_MODEL, DEFAULT_PROVIDER,
DEFAULT_STT_MODEL, ELEVENLABS_STT_BASE_URL, GROQ_MODELS, LOCAL_STT_COMMAND_ENV,
LOCAL_STT_LANGUAGE_ENV, MAX_FILE_SIZE, OPENAI_MODELS, SUPPORTED_FORMATS, XAI_STT_BASE_URL,
_error_result, _get_stt_section, _ok_result,
)
from tools.transcription_audio import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
_CLOUD_TRIM_KEEP_MS_DEFAULT,
_CLOUD_TRIM_MIN_INPUT_SECONDS,
_CLOUD_TRIM_THRESHOLD_DB_DEFAULT,
_cloud_trim_settings,
_convert_caf_to_wav,
_find_ffmpeg_binary,
_find_ffprobe_binary,
_find_whisper_binary,
_prepare_audio_for_transcription,
_prepare_local_audio,
_probe_audio_duration,
_run_ffmpeg_stt_encode,
_trim_silence_for_cloud_stt,
_validate_audio_file,
_validate_audio_file_size,
_validate_audio_source_file,
_CLOUD_TRIM_KEEP_MS_DEFAULT, _CLOUD_TRIM_MIN_INPUT_SECONDS, _CLOUD_TRIM_THRESHOLD_DB_DEFAULT,
_cloud_trim_settings, _convert_caf_to_wav, _find_ffmpeg_binary, _find_ffprobe_binary,
_find_whisper_binary, _prepare_audio_for_transcription, _prepare_local_audio,
_probe_audio_duration, _run_ffmpeg_stt_encode, _trim_silence_for_cloud_stt,
_validate_audio_file, _validate_audio_file_size, _validate_audio_source_file,
)
from tools.transcription_local import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
_LOGPROB_THRESHOLD_DEFAULT,
_NO_SPEECH_PROB_THRESHOLD_DEFAULT,
_get_idle_unload_seconds,
_get_local_command_template,
_has_local_command,
_is_hallucinated_segment,
_join_confident_segments,
_load_local_whisper_model,
_looks_like_cuda_lib_error,
_normalize_local_model,
_transcribe_local_command,
_try_lazy_install_stt,
_LOGPROB_THRESHOLD_DEFAULT, _NO_SPEECH_PROB_THRESHOLD_DEFAULT, _get_idle_unload_seconds,
_get_local_command_template, _has_local_command, _is_hallucinated_segment,
_join_confident_segments, _load_local_whisper_model, _looks_like_cuda_lib_error,
_normalize_local_model, _transcribe_local_command, _try_lazy_install_stt,
build_local_transcribe_kwargs,
)
from tools.transcription_cloud import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
_extract_transcript_text,
_has_xai_stt_credentials,
_is_local_or_private_url,
_resolve_openai_audio_client_config,
_transcribe_deepinfra,
_transcribe_elevenlabs,
_transcribe_groq,
_transcribe_mistral,
_transcribe_openai,
_transcribe_xai,
_extract_transcript_text, _has_xai_stt_credentials, _is_local_or_private_url,
_resolve_openai_audio_client_config, _transcribe_deepinfra, _transcribe_elevenlabs,
_transcribe_groq, _transcribe_mistral, _transcribe_openai, _transcribe_xai,
)
from tools.transcription_command import ( # noqa: F401 (re-exported; tests patch tools.transcription_tools.<name>)
COMMAND_STT_OUTPUT_FORMATS,
DEFAULT_COMMAND_STT_LANGUAGE,
DEFAULT_COMMAND_STT_OUTPUT_FORMAT,
DEFAULT_COMMAND_STT_TIMEOUT_SECONDS,
_PROMPT_CHARS_PER_TOKEN,
_WHISPER_PROMPT_TOKEN_CAP,
_apply_pre_transcription_hook,
_dispatch_to_plugin_provider,
_enforce_prompt_length_limit,
_get_command_stt_output_format,
_get_command_stt_timeout,
_get_named_stt_provider_config,
_render_command_stt_template,
_resolve_command_stt_provider_config,
_run_command_stt,
_transcribe_command_stt,
_unregistered_stt_provider_error,
COMMAND_STT_OUTPUT_FORMATS, DEFAULT_COMMAND_STT_LANGUAGE, DEFAULT_COMMAND_STT_OUTPUT_FORMAT,
DEFAULT_COMMAND_STT_TIMEOUT_SECONDS, _PROMPT_CHARS_PER_TOKEN, _WHISPER_PROMPT_TOKEN_CAP,
_apply_pre_transcription_hook, _dispatch_to_plugin_provider, _enforce_prompt_length_limit,
_get_command_stt_output_format, _get_command_stt_timeout, _get_named_stt_provider_config,
_render_command_stt_template, _resolve_command_stt_provider_config, _run_command_stt,
_transcribe_command_stt, _unregistered_stt_provider_error,
)
logger = logging.getLogger(__name__)