refactor(tools): AST-neutral layout compaction of STT/TTS modules
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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``.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user