refactor(tools): AST-neutral layout compaction of STT/TTS modules

This commit is contained in:
Teknium
2026-09-02 22:45:22 -07:00
parent 4d37eddc2f
commit d34703fa86
5 changed files with 28 additions and 97 deletions

View File

@@ -91,11 +91,7 @@ def _sdk_prompt_kwargs(language: Optional[str], prompt: Optional[str]) -> Dict[s
def _transcribe_groq(
file_path: str,
model_name: str,
*,
language: Optional[str] = None,
prompt: Optional[str] = None,
file_path: str, model_name: str, *, language: Optional[str] = None, prompt: Optional[str] = None
) -> Dict[str, Any]:
"""Transcribe using Groq Whisper API (free tier available).
@@ -129,13 +125,8 @@ def _transcribe_groq(
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,
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.
@@ -216,11 +207,7 @@ def _transcribe_openai(
def _transcribe_mistral(
file_path: str,
model_name: str,
*,
language: Optional[str] = None,
prompt: Optional[str] = None,
file_path: str, model_name: str, *, language: Optional[str] = None, prompt: Optional[str] = None
) -> Dict[str, Any]:
"""Transcribe with the ``mistralai`` SDK (``/v1/audio/transcriptions``); requires ``MISTRAL_API_KEY``."""
from tools.transcription_tools import _resolve_provider_key, _resolve_stt_language
@@ -288,13 +275,8 @@ def _rest_transcript(response, label: str, extract_detail, extract_text):
def _rest_provider(
file_path: str,
provider: str,
label: str,
post: Callable[[], Any],
extract_detail,
extract_text,
log: Callable[[str, Dict[str, Any]], None],
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 REST flow: ``post()`` -> envelope/transcript -> ``log(text, body)`` -> ok; exceptions -> ``_cloud_failure``."""
try:
@@ -308,11 +290,7 @@ def _rest_provider(
def _transcribe_xai(
file_path: str,
model_name: str,
*,
language: Optional[str] = None,
prompt: Optional[str] = None,
file_path: str, model_name: str, *, language: Optional[str] = None, prompt: Optional[str] = None
) -> Dict[str, Any]:
"""Transcribe via xAI ``POST /v1/stt`` (multipart). Supports ITN, diarization, word timestamps."""
from tools.transcription_tools import _load_stt_config, _resolve_stt_language, get_env_value
@@ -390,10 +368,8 @@ def _transcribe_xai(
)
return _rest_provider(
file_path, "xai", "xAI STT", _post,
lambda body: body.get("error", {}).get("message", ""),
lambda body: body.get("text", "").strip(),
_log,
file_path, "xai", "xAI STT", _post, lambda body: body.get("error", {}).get("message", ""),
lambda body: body.get("text", "").strip(), _log,
)
@@ -405,11 +381,7 @@ def _elevenlabs_error_detail(err_body: Dict[str, Any]) -> str:
def _transcribe_elevenlabs(
file_path: str,
model_name: str,
*,
language: Optional[str] = None,
prompt: Optional[str] = None,
file_path: str, model_name: str, *, language: Optional[str] = None, prompt: Optional[str] = None
) -> Dict[str, Any]:
"""Transcribe using ElevenLabs Scribe STT API."""
from tools.transcription_tools import _load_stt_config, _resolve_provider_key, _resolve_stt_language, get_env_value
@@ -450,11 +422,7 @@ def _transcribe_elevenlabs(
def _transcribe_deepinfra(
file_path: str,
model_name: str,
*,
language: Optional[str] = None,
prompt: Optional[str] = None,
file_path: str, model_name: str, *, language: Optional[str] = None, prompt: Optional[str] = None
) -> Dict[str, Any]:
"""Resolve DeepInfra credentials/model (shared ``hermes_cli.models`` helpers), then delegate to :func:`_transcribe_openai`."""
from tools.transcription_tools import _load_stt_config, _resolve_provider_key, _transcribe_openai

View File

@@ -76,12 +76,8 @@ def _read_command_stt_output(output_path: Path, stdout: str, fmt: str) -> str:
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,
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``.
@@ -153,13 +149,8 @@ def _unregistered_stt_provider_error(provider: str) -> Dict[str, Any]:
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,
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.
@@ -257,13 +248,8 @@ def _enforce_prompt_length_limit(prompt: Optional[str], provider: str) -> Option
def _apply_pre_transcription_hook(
*,
file_path: str,
provider: str,
model: Optional[str],
language: Optional[str],
prompt: Optional[str],
source: Optional[str],
*, 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.

View File

@@ -249,11 +249,7 @@ def _join_confident_segments(segments: Any, local_cfg: Dict[str, Any]) -> str:
def _transcribe_local_command(
file_path: str,
model_name: str,
*,
language: Optional[str] = None,
prompt: Optional[str] = None,
file_path: str, model_name: str, *, language: Optional[str] = None, prompt: Optional[str] = None
) -> Dict[str, Any]:
"""Run the configured local STT command template and read back a .txt transcript."""
from tools.transcription_tools import _prepare_local_audio, _resolve_stt_language

View File

@@ -144,10 +144,7 @@ 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 = (),
provider_key: str, stt_config: Optional[Dict[str, Any]] = None, *, extra_keys: tuple = ()
) -> Optional[str]:
"""Resolve the language hint for an STT provider; first non-empty wins.
@@ -437,8 +434,7 @@ def _get_or_load_local_model(model_name: str, local_cfg: Dict[str, Any]):
# stt.local.device / compute_type let users pin a configuration
# where ``auto`` mis-detects; the loader keeps the CUDA→CPU fallback.
_local_model = _load_local_whisper_model(
model_name,
device=local_cfg.get("device", "auto"),
model_name, device=local_cfg.get("device", "auto"),
compute_type=local_cfg.get("compute_type", "auto"),
)
_local_model_name = model_name
@@ -458,11 +454,7 @@ def _replace_cached_model_on_cpu(model_name: str):
def _transcribe_local(
file_path: str,
model_name: str,
*,
language: Optional[str] = None,
prompt: Optional[str] = None,
file_path: str, model_name: str, *, language: Optional[str] = None, prompt: Optional[str] = None
) -> Dict[str, Any]:
"""Transcribe using faster-whisper (local, free)."""
if not _HAS_FASTER_WHISPER and not _try_lazy_install_stt():
@@ -531,9 +523,7 @@ 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,
file_path: str, model: Optional[str] = None, source: Optional[str] = None
) -> Dict[str, Any]:
"""Transcribe a validated audio file with the configured STT provider.
@@ -610,10 +600,7 @@ def _builtin_model_name(provider: str, stt_config: Dict[str, Any], model: Option
def _dispatch_stt_provider(
file_path: str,
provider: str,
stt_config: Dict[str, Any],
model: Optional[str] = None,
file_path: str, provider: str, stt_config: Dict[str, Any], model: Optional[str] = None,
source: Optional[str] = None,
) -> Dict[str, Any]:
"""Route *file_path* to the handler for *provider* (built-in > command > plugin)."""
@@ -654,10 +641,8 @@ def _dispatch_stt_provider(
# ``stt.<provider>`` like built-ins; the ``model`` argument overrides it.
plugin_cfg = _get_stt_section(stt_config, provider)
plugin_result = _dispatch_to_plugin_provider(
file_path, provider, stt_config,
model=model or plugin_cfg.get("model"),
language=language or _resolve_stt_language(provider, stt_config),
prompt=prompt,
file_path, provider, stt_config, model=model or plugin_cfg.get("model"),
language=language or _resolve_stt_language(provider, stt_config), prompt=prompt,
)
if plugin_result is not None:
return plugin_result
@@ -688,9 +673,7 @@ 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,
file_path: str, model: Optional[str] = None, source: Optional[str] = None
) -> Dict[str, Any]:
"""Safely validate, preprocess supported inputs, and dispatch transcription.
@@ -726,8 +709,7 @@ def transcribe_audio(
def transcribe_audio_local_fallback(
file_path: str,
model: Optional[str] = None,
file_path: str, model: Optional[str] = None
) -> Dict[str, Any]:
"""Try an already-installed local STT backend without changing config.

View File

@@ -154,8 +154,7 @@ _PROVIDER_PRIORITY: List[str] = ["elevenlabs", "gemini", "openai", "xai"]
def resolve_streaming_provider(
tts_config: Dict,
preferred: Optional[str] = None,
tts_config: Dict, preferred: Optional[str] = None
) -> Optional[StreamingTTSProvider]:
"""Return a ready streamer for the *configured* provider, else ``None``.