From 05c6da29852e1aa5c1dec5c3067a1d5efb173953 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 00:51:58 -0700 Subject: [PATCH] refactor(tools): drop test-only _has_any_command_tts_provider, fold tts provider write paths and speaker text pump --- tests/tools/test_tts_command_providers.py | 8 +-- tests/tools/test_tts_mistral.py | 3 +- tests/tools/test_tts_piper.py | 1 - tools/tts_tool.py | 24 ++------- tools/tts_tool_providers.py | 61 +++++++---------------- tools/tts_tool_speaker.py | 49 +++++++----------- 6 files changed, 43 insertions(+), 103 deletions(-) diff --git a/tests/tools/test_tts_command_providers.py b/tests/tools/test_tts_command_providers.py index 631d7d183a..abe540c268 100644 --- a/tests/tools/test_tts_command_providers.py +++ b/tests/tools/test_tts_command_providers.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/tools/test_tts_mistral.py b/tests/tools/test_tts_mistral.py index 0f8d4432c4..6917546253 100644 --- a/tests/tools/test_tts_mistral.py +++ b/tests/tools/test_tts_mistral.py @@ -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 diff --git a/tests/tools/test_tts_piper.py b/tests/tools/test_tts_piper.py index 33bb4feb47..52ed2cacff 100644 --- a/tests/tools/test_tts_piper.py +++ b/tests/tools/test_tts_piper.py @@ -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"): diff --git a/tools/tts_tool.py b/tools/tts_tool.py index 51c0538426..b41d4c3b77 100644 --- a/tools/tts_tool.py +++ b/tools/tts_tool.py @@ -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="🔊", ) diff --git a/tools/tts_tool_providers.py b/tools/tts_tool_providers.py index 114a8a95e6..53e0587bd8 100644 --- a/tools/tts_tool_providers.py +++ b/tools/tts_tool_providers.py @@ -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: diff --git a/tools/tts_tool_speaker.py b/tools/tts_tool_speaker.py index dcfad75f1a..fe13476107 100644 --- a/tools/tts_tool_speaker.py +++ b/tools/tts_tool_speaker.py @@ -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: