refactor(tools): drop unreachable legacy TTS strip fallback and duplicated streamer sample rates; STT/TTS config wrappers as partials

This commit is contained in:
Teknium
2026-09-03 00:57:29 -07:00
parent d85b86e97a
commit 412ddb6772
5 changed files with 20 additions and 69 deletions

View File

@@ -14,6 +14,7 @@ import os
import shutil
import subprocess
import tempfile
from functools import partial
from pathlib import Path
from typing import Any, Dict, Optional
@@ -36,16 +37,9 @@ def _find_binary(binary_name: str) -> Optional[str]:
return shutil.which(binary_name)
def _find_ffmpeg_binary() -> Optional[str]:
return _find_binary("ffmpeg")
def _find_ffprobe_binary() -> Optional[str]:
return _find_binary("ffprobe")
def _find_whisper_binary() -> Optional[str]:
return _find_binary("whisper")
_find_ffmpeg_binary = partial(_find_binary, "ffmpeg")
_find_ffprobe_binary = partial(_find_binary, "ffprobe")
_find_whisper_binary = partial(_find_binary, "whisper")
def _run_quiet(command: list, *, timeout: float, env: Optional[dict] = None) -> subprocess.CompletedProcess:

View File

@@ -12,6 +12,7 @@ from __future__ import annotations
import logging
import subprocess
import tempfile
from functools import partial
from pathlib import Path
from typing import Any, Dict, Optional
@@ -41,22 +42,14 @@ COMMAND_STT_OUTPUT_FORMATS = frozenset({"txt", "json", "srt", "vtt"})
_NON_COMMAND_STT_NAMES = frozenset(BUILTIN_STT_PROVIDERS | {"none"})
def _get_named_stt_provider_config(stt_config: Dict[str, Any], name: str) -> Dict[str, Any]:
"""``stt.providers.<name>`` (canonical), else ``stt.<name>`` for non-built-in names only."""
return _named_provider_config(stt_config, name, BUILTIN_STT_PROVIDERS)
def _resolve_command_stt_provider_config(provider: str, stt_config: Dict[str, Any]) -> Optional[Dict[str, Any]]:
"""The provider config if *provider* is a command type; None for built-ins, ``none``, unknown."""
return _resolve_command_config(provider, stt_config, _NON_COMMAND_STT_NAMES)
def _get_command_stt_timeout(config: Dict[str, Any]) -> float:
return _command_timeout(config, DEFAULT_COMMAND_STT_TIMEOUT_SECONDS)
def _get_command_stt_output_format(config: Dict[str, Any]) -> str:
return _command_output_format(config, COMMAND_STT_OUTPUT_FORMATS, DEFAULT_COMMAND_STT_OUTPUT_FORMAT)
# ``stt.providers.<name>`` (canonical), else ``stt.<name>`` for non-built-in names only.
_get_named_stt_provider_config = partial(_named_provider_config, builtins=BUILTIN_STT_PROVIDERS)
# The provider config if it is a command type; None for built-ins, ``none``, unknown.
_resolve_command_stt_provider_config = partial(_resolve_command_config,
reserved=_NON_COMMAND_STT_NAMES)
_get_command_stt_timeout = partial(_command_timeout, default=DEFAULT_COMMAND_STT_TIMEOUT_SECONDS)
_get_command_stt_output_format = partial(_command_output_format, formats=COMMAND_STT_OUTPUT_FORMATS,
default=DEFAULT_COMMAND_STT_OUTPUT_FORMAT)
def _read_command_stt_output(output_path: Path, stdout: str, fmt: str) -> str:

View File

@@ -27,6 +27,7 @@ import subprocess
import tempfile
import threading
import time
from functools import partial
from pathlib import Path
from typing import Any, Dict, FrozenSet, Optional
@@ -294,16 +295,9 @@ COMMAND_TTS_OUTPUT_FORMATS = frozenset({"mp3", "wav", "ogg", "flac", "m4a", "aac
DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH = 5000
def _get_named_provider_config(tts_config: Dict[str, Any], name: str) -> Dict[str, Any]:
return _named_provider_config(tts_config, name, BUILTIN_TTS_PROVIDERS)
def _resolve_command_provider_config(provider: str, tts_config: Dict[str, Any]) -> Optional[Dict[str, Any]]:
return _resolve_command_config(provider, tts_config, BUILTIN_TTS_PROVIDERS)
def _get_command_tts_timeout(config: Dict[str, Any]) -> float:
return _command_timeout(config, DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS)
_get_named_provider_config = partial(_named_provider_config, builtins=BUILTIN_TTS_PROVIDERS)
_resolve_command_provider_config = partial(_resolve_command_config, reserved=BUILTIN_TTS_PROVIDERS)
_get_command_tts_timeout = partial(_command_timeout, default=DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS)
def _iter_command_providers(tts_config: Dict[str, Any]):

View File

@@ -102,7 +102,7 @@ class SentenceChunker:
class StreamingTTSProvider(ABC):
"""Yields raw int16, little-endian, mono PCM chunks at ``sample_rate``."""
"""Yields raw int16, little-endian, mono PCM chunks at ``sample_rate`` (built-in streamers: 24 kHz)."""
sample_rate: int = 24000
channels: int = 1
@@ -180,8 +180,6 @@ def _capped(chunks: Iterator[bytes], label: str) -> Iterator[bytes]:
class ElevenLabsStreamer(StreamingTTSProvider):
"""ElevenLabs chunked HTTP → pcm_24000 (the original reference path)."""
sample_rate = 24000
@staticmethod
def available() -> bool:
return bool(_resolve_key("ELEVENLABS_API_KEY", "elevenlabs"))
@@ -213,8 +211,6 @@ def _openai_config_api_key() -> str:
class OpenAIStreamer(StreamingTTSProvider):
"""OpenAI speech with ``response_format=pcm`` (24 kHz mono int16)."""
sample_rate = 24000
@staticmethod
def available() -> bool:
return bool(_openai_config_api_key() or resolve_openai_audio_api_key())
@@ -235,8 +231,6 @@ class OpenAIStreamer(StreamingTTSProvider):
class GeminiStreamer(StreamingTTSProvider):
"""Gemini ``streamGenerateContent?alt=sse`` → SSE feed of base64 PCM chunks (24 kHz), bounded streamed body."""
sample_rate = 24000
@staticmethod
def available() -> bool:
return bool(_gemini_key())
@@ -293,8 +287,6 @@ class XAIStreamer(StreamingTTSProvider):
iterator contract — the seam unit tests patch.
"""
sample_rate = 24000
@staticmethod
def available() -> bool:
try:

View File

@@ -216,28 +216,6 @@ def prepare_spoken_text(text: str, max_chars: int | None = 4000) -> str:
return spoken
# Legacy regex fallback, only used if the shared normalizer raises.
_LEGACY_TTS_STRIP_STEPS = (
(re.compile(r'<think[\s>].*?</think>', flags=re.DOTALL), ' '),
(re.compile(r'```[\s\S]*?```'), ' '),
(re.compile(r'\[([^\]]+)\]\([^)]+\)'), r'\1'),
(re.compile(r'https?://\S+'), ''),
(re.compile(r'\*\*(.+?)\*\*'), r'\1'),
(re.compile(r'\*(.+?)\*'), r'\1'),
(re.compile(r'`(.+?)`'), r'\1'),
(re.compile(r'^#+\s*', flags=re.MULTILINE), ''),
(re.compile(r'^\s*[-*]\s+', flags=re.MULTILINE), ''),
(re.compile(r'---+'), ''),
# Emoji + variation selectors/ZWJ: providers speak them as awkward labels.
(re.compile('[\U0001F000-\U0001FAFF\u2600-\u27BF\uFE0F\u200D\U000E0020-\U000E007F]+'), ' '),
(re.compile(r'\n{3,}'), '\n\n'))
def _strip_markdown_for_tts(text: str) -> str:
"""``prepare_spoken_text`` without a length cap; falls back to the legacy regex pipeline if it raises."""
try:
return prepare_spoken_text(text, max_chars=None)
except Exception:
for pattern, repl in _LEGACY_TTS_STRIP_STEPS:
text = pattern.sub(repl, text)
return text.strip()
"""``prepare_spoken_text`` without a length cap (``tts_tool`` compatibility name)."""
return prepare_spoken_text(text, max_chars=None)