refactor(tools): drop test-only _has_any_command_tts_provider, fold tts provider write paths and speaker text pump

This commit is contained in:
Teknium
2026-09-03 00:51:58 -07:00
parent 45fb404019
commit 05c6da2985
6 changed files with 43 additions and 103 deletions
+1 -7
View File
@@ -32,7 +32,6 @@ from tools.tts_tool import (
_get_command_tts_output_format,
_get_command_tts_timeout,
_get_named_provider_config,
_has_any_command_tts_provider,
_is_command_provider_config,
_is_command_tts_voice_compatible,
_iter_command_providers,
@@ -168,7 +167,7 @@ class TestIsCommandProviderConfig:
# ---------------------------------------------------------------------------
# _iter_command_providers / _has_any_command_tts_provider
# _iter_command_providers
# ---------------------------------------------------------------------------
class TestIterCommandProviders:
@@ -185,11 +184,6 @@ class TestIterCommandProviders:
assert names == ["piper-cli", "voxcpm"]
def test_has_any_command_provider_when_none(self):
assert _has_any_command_tts_provider({"providers": {}}) is False
assert _has_any_command_tts_provider({}) is False
# ---------------------------------------------------------------------------
# config getters
# ---------------------------------------------------------------------------
+1 -2
View File
@@ -163,6 +163,5 @@ class TestCheckTtsRequirementsMistral:
patch("tools.tts_tool._import_openai_client", side_effect=ImportError), \
patch("tools.tts_tool._check_neutts_available", return_value=False), \
patch("tools.tts_tool._check_kittentts_available", return_value=False), \
patch("tools.tts_tool._check_piper_available", return_value=False), \
patch("tools.tts_tool._has_any_command_tts_provider", return_value=False):
patch("tools.tts_tool._check_piper_available", return_value=False):
assert check_tts_requirements() is False
-1
View File
@@ -232,7 +232,6 @@ class TestCheckTtsRequirementsPiper:
monkeypatch.setattr(tts_tool, "_import_mistral_client", lambda: (_ for _ in ()).throw(ImportError()))
monkeypatch.setattr(tts_tool, "_check_neutts_available", lambda: False)
monkeypatch.setattr(tts_tool, "_check_kittentts_available", lambda: False)
monkeypatch.setattr(tts_tool, "_has_any_command_tts_provider", lambda: False)
monkeypatch.setattr(tts_tool, "_has_openai_audio_backend", lambda: False)
for env in ("MINIMAX_API_KEY", "XAI_API_KEY", "GEMINI_API_KEY",
"GOOGLE_API_KEY", "MISTRAL_API_KEY", "ELEVENLABS_API_KEY"):
+5 -19
View File
@@ -192,13 +192,6 @@ _NATIVE_OPUS_PROVIDERS = frozenset({"openai", "elevenlabs", "mistral", "gemini"}
_FFMPEG_OPUS_PROVIDERS = frozenset({"edge", "neutts", "minimax", "xai", "kittentts", "piper"})
def _has_any_command_tts_provider(tts_config: Optional[Dict[str, Any]] = None) -> bool:
"""Return True when any command-type TTS provider is configured."""
if tts_config is None:
tts_config = _load_tts_config()
return any(True for _ in _iter_command_providers(tts_config))
# --- Built-in provider dispatch ---
# provider -> (availability predicate or None, log label, generator name, "package missing" error).
# Predicates and generator names resolve module globals at call time so tests that monkeypatch
@@ -249,10 +242,8 @@ def _select_builtin_engine(provider: str) -> tuple:
names take the Edge default; without edge-tts NeuTTS is the fallback (engine != provider)."""
entry = _BUILTIN_DISPATCH.get(provider)
if entry is not None:
available, missing_error = entry[0], entry[3]
if available is not None and not available():
return provider, _error_json(missing_error)
return provider, None
available, _label, _generator, missing_error = entry
return provider, (_error_json(missing_error) if available is not None and not available() else None)
if _importable(_import_edge_tts):
return provider, None # Edge default; the reported provider stays as configured
if _check_neutts_available():
@@ -476,10 +467,8 @@ def text_to_speech_tool(
if not chunks:
return tool_error("Text is required", success=False)
if len(chunks) > 1:
logger.info(
"TTS text for provider %s split into %d chunks (input=%d chars, cap=%d)",
provider, len(chunks), len(text), max_len,
)
logger.info("TTS text for provider %s split into %d chunks (input=%d chars, cap=%d)",
provider, len(chunks), len(text), max_len)
platform, want_opus = _session_platform()
delivery_profile = _resolve_audio_delivery_profile(platform, tts_config)
@@ -625,10 +614,7 @@ registry.register(
schema=TTS_SCHEMA,
handler=lambda args, **kw: text_to_speech_tool(
text=args.get("text", ""),
output_path=args.get("output_path"),
speed=args.get("speed"),
instructions=args.get("instructions"),
provider=args.get("provider")),
**{k: args.get(k) for k in ("output_path", "speed", "instructions", "provider")}),
check_fn=check_tts_requirements,
emoji="🔊",
)
+19 -42
View File
@@ -160,10 +160,6 @@ def _write_bytes(output_path: str, audio_bytes: bytes) -> str:
return output_path
def _write_tts_response_to_file(response: Any, output_path: str, *, label: str, limit: Optional[int] = None) -> None:
_write_bytes(output_path, _read_tts_response_bytes(response, label=label, limit=limit))
def _post_json(url: str, payload: Dict[str, Any], headers: Dict[str, str], **extra: Any):
"""Streaming ``requests.post`` with the shared 60s timeout (body read via the bounded readers)."""
import requests
@@ -355,20 +351,16 @@ def _generate_xai_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -
if auto_speech_tags:
text = _apply_xai_auto_speech_tags(text)
if creds.get("provider") == "xai-oauth":
base_url = str(creds.get("base_url") or DEFAULT_XAI_BASE_URL).strip().rstrip("/")
base_url = creds.get("base_url")
else:
base_url = str(
xai_config.get("base_url")
or creds.get("base_url")
or _origin().get_env_value("XAI_BASE_URL")
or DEFAULT_XAI_BASE_URL
).strip().rstrip("/")
base_url = xai_config.get("base_url") or creds.get("base_url") or _origin().get_env_value("XAI_BASE_URL")
base_url = str(base_url or DEFAULT_XAI_BASE_URL).strip().rstrip("/")
# Documented minimal POST /v1/tts shape; optional fields only when they
# differ from the API defaults.
codec = "wav" if output_path.endswith(".wav") else "mp3"
payload: Dict[str, Any] = {"text": text, "voice_id": voice_id, "language": language}
if codec != "mp3" or sample_rate != DEFAULT_XAI_SAMPLE_RATE or (codec == "mp3" and bit_rate != DEFAULT_XAI_BIT_RATE):
if codec != "mp3" or sample_rate != DEFAULT_XAI_SAMPLE_RATE or bit_rate != DEFAULT_XAI_BIT_RATE:
output_format: Dict[str, Any] = {"codec": codec}
if sample_rate:
output_format["sample_rate"] = sample_rate
@@ -377,7 +369,7 @@ def _generate_xai_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -
payload["output_format"] = output_format
if speed is not None and speed != DEFAULT_XAI_SPEED_DEFAULT:
payload["speed"] = speed
if optimize_streaming_latency is not None and optimize_streaming_latency != DEFAULT_XAI_OPTIMIZE_STREAMING_LATENCY_DEFAULT:
if optimize_streaming_latency not in (None, DEFAULT_XAI_OPTIMIZE_STREAMING_LATENCY_DEFAULT):
payload["optimize_streaming_latency"] = optimize_streaming_latency
if text_normalization:
payload["text_normalization"] = True
@@ -388,8 +380,7 @@ def _generate_xai_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -
"User-Agent": hermes_xai_user_agent(),
})
response.raise_for_status()
_write_tts_response_to_file(response, output_path, label="xAI TTS")
return output_path
return _write_bytes(output_path, _read_tts_response_bytes(response, label="xAI TTS"))
# --- MiniMax TTS ---
@@ -435,9 +426,7 @@ def _resolve_minimax_tts_runtime(tts_config: Dict[str, Any]) -> _MiniMaxTTSRunti
other_region = "cn" if region == "global" else "global"
if (urlparse(endpoint).hostname or "").lower() in _MINIMAX_OFFICIAL_HOSTS[other_region]:
raise ValueError(
f"tts.minimax.base_url points to the {other_region!r} MiniMax endpoint "
f"but region is {region!r}"
)
f"tts.minimax.base_url points to the {other_region!r} MiniMax endpoint but region is {region!r}")
return _MiniMaxTTSRuntime(region=region, endpoint=endpoint, credential_source=credential_source, api_key=api_key)
@@ -445,8 +434,8 @@ def _raise_minimax_api_error(result: Dict[str, Any]) -> None:
base_resp = result.get("base_resp", {})
status_code = base_resp.get("status_code", -1)
if status_code != 0:
status_msg = base_resp.get("status_msg", "unknown error")
raise RuntimeError(f"MiniMax TTS API error (code {status_code}): {status_msg}")
raise RuntimeError(
f"MiniMax TTS API error (code {status_code}): {base_resp.get('status_msg', 'unknown error')}")
def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
@@ -470,20 +459,14 @@ def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any
is_t2a_v2 = "t2a_v2" in base_url
if is_t2a_v2:
payload = {
"model": model,
"text": text,
"model": model, "text": text,
"voice_setting": {
"voice_id": voice_id,
"speed": mm_config.get("speed", 1.0),
"vol": mm_config.get("vol", 1.0),
"pitch": mm_config.get("pitch", 0),
"emotion": mm_config.get("emotion", "neutral"),
"voice_id": voice_id, "speed": mm_config.get("speed", 1.0), "vol": mm_config.get("vol", 1.0),
"pitch": mm_config.get("pitch", 0), "emotion": mm_config.get("emotion", "neutral"),
},
"audio_setting": {
"sample_rate": mm_config.get("sample_rate", 32000),
"bitrate": mm_config.get("bitrate", 128000),
"format": "mp3",
"channel": 1,
"sample_rate": mm_config.get("sample_rate", 32000), "bitrate": mm_config.get("bitrate", 128000),
"format": "mp3", "channel": 1,
},
}
else:
@@ -504,8 +487,7 @@ def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any
content_type = response.headers.get("Content-Type", "")
if "audio/" in content_type:
_write_tts_response_to_file(response, output_path, label="MiniMax TTS")
return output_path
return _write_bytes(output_path, _read_tts_response_bytes(response, label="MiniMax TTS"))
# Non-audio reply: surface the API error if the body is JSON.
raw_body = b""
@@ -515,9 +497,7 @@ def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any
except (json.JSONDecodeError, UnicodeDecodeError, TypeError):
response.raise_for_status()
raise RuntimeError(
f"MiniMax TTS returned unexpected Content-Type '{content_type}' "
f"({len(raw_body)} bytes)"
)
f"MiniMax TTS returned unexpected Content-Type '{content_type}' ({len(raw_body)} bytes)")
raise RuntimeError("MiniMax TTS returned no audio data")
@@ -647,10 +627,8 @@ def _generate_gemini_tts(text: str, output_path: str, tts_config: Dict[str, Any]
"""Generate audio via Gemini ``generateContent`` (``responseModalities=["AUDIO"]``). The reply is
base64 24kHz mono 16-bit PCM, wrapped as WAV and ffmpeg-converted to the requested container."""
origin = _origin()
api_key = (
origin._resolve_provider_key("GEMINI_API_KEY", "gemini")
or origin._resolve_provider_key("GOOGLE_API_KEY", "gemini")
)
api_key = origin._resolve_provider_key("GEMINI_API_KEY", "gemini") or origin._resolve_provider_key(
"GOOGLE_API_KEY", "gemini")
if not api_key:
raise ValueError("GEMINI_API_KEY not set. Get one at https://aistudio.google.com/app/apikey")
@@ -701,8 +679,7 @@ def _generate_gemini_tts(text: str, output_path: str, tts_config: Dict[str, Any]
audio_part = next((p for p in parts if "inlineData" in p or "inline_data" in p), None)
if audio_part is None:
raise RuntimeError("Gemini TTS response contained no audio data")
inline = audio_part.get("inlineData") or audio_part.get("inline_data") or {}
audio_b64 = inline.get("data", "")
audio_b64 = (audio_part.get("inlineData") or audio_part.get("inline_data") or {}).get("data", "")
except (KeyError, IndexError, TypeError) as e:
raise RuntimeError(f"Gemini TTS response was malformed: {e}") from e
if not audio_b64:
+17 -32
View File
@@ -184,11 +184,9 @@ class _StreamerPlayback:
def close_output_stream(self) -> None:
"""Always release the device so a later stream can open it."""
if self.output_stream is not None:
try:
with contextlib.suppress(Exception):
self.output_stream.stop()
self.output_stream.close()
except Exception:
pass
# -- prefetch ----------------------------------------------------------
@@ -202,9 +200,9 @@ class _StreamerPlayback:
self._prefetch_sem.acquire()
chunk_queue: "queue.Queue[Optional[bytes]]" = queue.Queue(maxsize=self._CHUNK_QUEUE_MAX)
self._audio_queue.put(chunk_queue)
t = threading.Thread(target=self._consume_to_queue, args=(audio_iter, chunk_queue), daemon=True)
self._prefetch_threads.append(t)
t.start()
self._prefetch_threads.append(
threading.Thread(target=self._consume_to_queue, args=(audio_iter, chunk_queue), daemon=True))
self._prefetch_threads[-1].start()
def _consume_to_queue(self, audio_iter: Iterator[bytes], chunk_queue: "queue.Queue[Optional[bytes]]") -> None:
try:
@@ -269,10 +267,8 @@ class _StreamerPlayback:
logger.warning("PortAudio write failed, attempting stream reinit: %s", write_exc)
if not self._recover_stream():
return
try:
with contextlib.suppress(Exception):
self._write_pcm(aligned)
except Exception:
pass
def _playback_worker(self) -> None:
"""Single consumer: play audio segments from the queue in order."""
@@ -339,8 +335,6 @@ def stream_tts_to_speaker(
playback = _StreamerPlayback(streamer, stop_event)
chunker = SentenceChunker()
long_flush_len = 100
queue_timeout = 0.5
spoken_sentences: list[str] = [] # skip duplicate/near-duplicate sentences (LLM repetition)
def _speak_sentence(sentence: str) -> None:
@@ -364,28 +358,21 @@ def stream_tts_to_speaker(
while not stop_event.is_set():
try:
delta = text_queue.get(timeout=queue_timeout)
delta = text_queue.get(timeout=0.5)
except queue.Empty:
# Idle producer: flush a long buffer instead of sitting on it
if len(chunker.buf) > long_flush_len:
for sentence in chunker.flush():
_speak_sentence(sentence)
continue
if delta is None:
for sentence in chunker.flush():
_speak_sentence(sentence)
break
for sentence in chunker.feed(delta):
delta = "" # idle producer: flush a long buffer instead of sitting on it
sentences = chunker.flush() if len(chunker.buf) > 100 else ()
else:
sentences = chunker.flush() if delta is None else chunker.feed(delta)
for sentence in sentences:
_speak_sentence(sentence)
while True:
try:
text_queue.get_nowait()
except queue.Empty:
if delta is None:
break
with contextlib.suppress(queue.Empty):
while True:
text_queue.get_nowait()
except Exception as exc:
logger.warning("Streaming TTS pipeline error: %s", exc)
finally:
@@ -393,10 +380,8 @@ def stream_tts_to_speaker(
# when stop_event is set) BEFORE tts_done_event fires, so continuous voice mode
# never reopens the mic over its own voice.
if sync_pipeline is not None:
try:
with contextlib.suppress(Exception):
sync_pipeline.close()
except Exception:
pass
# The end sentinel lives in finally: so an exception in the text pump still lets
# the playback worker exit.
if playback is not None: