refactor(tools): drop test-only _has_any_command_tts_provider, fold tts provider write paths and speaker text pump
This commit is contained in:
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user