merge simp/r3-34-F pass2 into simp/r3-34

This commit is contained in:
Teknium
2026-09-03 01:31:07 -07:00
11 changed files with 459 additions and 916 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"):
+84 -154
View File
@@ -1,16 +1,16 @@
#!/usr/bin/env python3
"""Text-to-speech tool: config resolution, built-in provider dispatch, output policy, registration.
Built-in providers: Edge (free default), ElevenLabs, OpenAI, DeepInfra, MiniMax, Mistral,
Gemini, xAI, and the local NeuTTS / KittenTTS / Piper engines; plus any ``type: command``
provider under ``tts.providers.<name>`` and plugin-registered providers. Output is Opus
(.ogg) for voice-bubble platforms, MP3 elsewhere. Sibling ``tts_tool_*`` modules hold the
backends/delivery/lifecycle; their names are re-imported here so ``tools.tts_tool.<name>``
keeps resolving and tests patching ``tools.tts_tool.<seam>`` still take effect (siblings
resolve those seams through ``_origin()`` at call time).
Built-ins: Edge (free default), ElevenLabs, OpenAI, DeepInfra, MiniMax, Mistral, Gemini, xAI,
local NeuTTS / KittenTTS / Piper; plus ``type: command`` providers under ``tts.providers.<name>``
and plugin-registered ones. Output is Opus (.ogg) on voice-bubble platforms, MP3 elsewhere.
Sibling ``tts_tool_*`` modules hold backends/delivery/lifecycle; their names are re-imported
here so ``tools.tts_tool.<name>`` resolves and test patches on this module still apply
(siblings read those seams through ``_origin()`` at call time).
"""
import asyncio
import contextlib
import datetime
import importlib.util
import json
@@ -53,62 +53,52 @@ from tools.tts_command_provider import ( # noqa: F401 — historical names re-e
_is_command_tts_voice_compatible, _iter_command_providers, _resolve_command_provider_config,
command_env_passthrough as _command_provider_env_passthrough,
render_command_template as _render_command_tts_template,
run_command_provider as _run_command_tts, shell_quote_context as _shell_quote_context,
)
run_command_provider as _run_command_tts, shell_quote_context as _shell_quote_context)
from tools.tool_backend_helpers import ( # noqa: F401 — seams patched by tests, resolved via tts_tool_openai._origin()
NOUS_MANAGED_PROVIDER, managed_nous_tools_enabled, read_selection, resolve_openai_audio_api_key,
)
NOUS_MANAGED_PROVIDER, managed_nous_tools_enabled, read_selection, resolve_openai_audio_api_key)
from tools.tts_tool_delivery import ( # noqa: F401 — historical names re-exported
FALLBACK_MAX_TEXT_LENGTH, PROVIDER_MAX_TEXT_LENGTH, _resolve_max_text_length,
AudioDeliveryProfile, _build_audio_delivery_files, _concat_audio_files, _convert_to_opus,
_pack_audio_files_for_delivery, _repair_ogg_container, _resolve_audio_delivery_profile,
_sniff_audio_container, _split_oversized_sentence, _split_text_for_tts, _wrap_pcm_as_wav,
)
_pack_audio_files_for_delivery, _remove_quietly, _repair_ogg_container,
_resolve_audio_delivery_profile, _sniff_audio_container, _split_oversized_sentence,
_split_text_for_tts, _wrap_pcm_as_wav)
from tools.tts_tool_providers import ( # noqa: F401 — historical names re-exported
DEFAULT_ELEVENLABS_MODEL_ID, DEFAULT_ELEVENLABS_VOICE_ID, DEFAULT_GEMINI_TTS_MODEL,
DEFAULT_GEMINI_TTS_VOICE, DEFAULT_MINIMAX_BASE_URL, DEFAULT_MINIMAX_CN_BASE_URL,
TTS_RESPONSE_BODY_LIMIT_BYTES, _XAI_FIRST_SENTENCE_RE, _XAI_INLINE_SPEECH_TAGS,
_XAI_WRAPPING_SPEECH_TAGS, _apply_xai_auto_speech_tags, _elevenlabs_environment_kwargs,
_generate_edge_tts, _generate_elevenlabs, _generate_gemini_tts, _generate_minimax_tts,
_generate_mistral_tts, _generate_xai_tts, _resolve_minimax_tts_runtime,
)
_generate_mistral_tts, _generate_xai_tts, _resolve_minimax_tts_runtime)
from tools.tts_tool_local import ( # noqa: F401 — historical names re-exported
DEFAULT_PIPER_VOICE, _LOCAL_TTS_MODEL_CACHES, _TTS_MODEL_CACHE_MAX, _generate_kittentts,
_generate_neutts, _generate_piper_tts, _kittentts_model_cache, _piper_voice_cache,
_resolve_piper_voice_path, _tts_cache_get_or_load,
)
_resolve_piper_voice_path, _tts_cache_get_or_load)
from tools.tts_tool_speaker import stream_tts_to_speaker # noqa: F401 — historical name re-exported
from tools.tts_text_normalize import _strip_markdown_for_tts # noqa: F401 — historical name re-exported
from tools.tts_tool_plugins import ( # noqa: F401 — historical names re-exported
_dispatch_to_plugin_provider, _plugin_provider_is_available,
_plugin_provider_is_voice_compatible,
)
_plugin_provider_is_voice_compatible)
from tools.tts_tool_openai import ( # noqa: F401 — historical names re-exported
DEFAULT_OPENAI_BASE_URL, DEFAULT_OPENAI_MODEL, DEFAULT_OPENAI_VOICE, MANAGED_OPENAI_TTS_MODELS,
_generate_deepinfra_tts, _generate_openai_tts, _has_openai_audio_backend,
_resolve_openai_audio_client_config,
)
_resolve_openai_audio_client_config)
from tools.tts_tool_lifecycle import ( # noqa: F401 — historical names re-exported
_local_tts_warmers, _reset_tts_leases_for_tests, acquire_tts_lease, release_tts_lease,
release_tts_provider, tts_lease_holders, warm_tts_provider,
)
release_tts_provider, tts_lease_holders, warm_tts_provider)
# --- Lazy SDK importers -- providers import only when used (headless boxes lack PortAudio etc.) ---
def _sdk_importer(module: str, attr: Optional[str] = None, feature: Optional[str] = None) -> Callable[[], Any]:
"""Lazy SDK importer: returns ``module`` (or ``module.attr``), raising ImportError when absent.
``feature`` names a ``tools.lazy_deps`` feature to best-effort install first (users who
enabled a provider in config.yaml never ran the post-setup hook); any failure there falls
through so the raw import still raises cleanly. sounddevice also raises OSError without PortAudio."""
``feature`` names a ``tools.lazy_deps`` feature to best-effort install first (users who enabled
a provider in config.yaml never ran the post-setup hook); any failure there falls through so
the raw import still raises cleanly. sounddevice also raises OSError without PortAudio."""
def _import():
if feature:
try:
with contextlib.suppress(Exception):
from tools.lazy_deps import ensure
ensure(feature, prompt=False)
except Exception:
pass
mod = importlib.import_module(module)
return getattr(mod, attr) if attr else mod
_import.__name__ = f"_import_{module.split('.')[0]}"
@@ -139,16 +129,9 @@ def _package_installed(name: str) -> bool:
return False
def _check_neutts_available() -> bool:
return _package_installed("neutts")
def _check_kittentts_available() -> bool:
return _package_installed("kittentts")
def _check_piper_available() -> bool:
return _package_installed("piper")
def _check_neutts_available() -> bool: return _package_installed("neutts")
def _check_kittentts_available() -> bool: return _package_installed("kittentts")
def _check_piper_available() -> bool: return _package_installed("piper")
# --- Defaults / config ---
@@ -160,8 +143,7 @@ def _get_default_output_dir() -> str:
return str(get_hermes_dir("cache/audio", "audio_cache"))
DEFAULT_OUTPUT_DIR = _get_default_output_dir()
_DEFAULT_OUTPUT_DIR_AT_IMPORT = DEFAULT_OUTPUT_DIR
DEFAULT_OUTPUT_DIR = _DEFAULT_OUTPUT_DIR_AT_IMPORT = _get_default_output_dir()
def _default_output_dir() -> str:
@@ -179,15 +161,14 @@ def _load_tts_config() -> Dict[str, Any]:
return load_config().get("tts") or {}
except ImportError:
logger.debug("hermes_cli.config not available, using default TTS config")
return {}
except Exception as e:
logger.warning("Failed to load TTS config: %s", e, exc_info=True)
return {}
return {}
def _get_provider(tts_config: Dict[str, Any]) -> str:
"""The configured TTS provider, or the free default — inference credentials never imply consent
to paid speech. ``nous`` is serviced by the OpenAI path via the managed openai-audio gateway."""
"""Configured provider or the free default (inference credentials never imply consent to paid
speech); ``nous`` is serviced by the OpenAI path through the managed openai-audio gateway."""
provider = (tts_config.get("provider") or DEFAULT_PROVIDER).lower().strip()
return "openai" if provider == NOUS_MANAGED_PROVIDER else provider
@@ -199,17 +180,9 @@ _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
# ``tools.tts_tool._import_x`` / ``_check_x`` / ``_generate_x`` apply.
# Predicates/generator names resolve module globals at call time so test monkeypatches apply.
_BUILTIN_DISPATCH: Dict[str, tuple] = {
"elevenlabs": (lambda: _importable(_import_elevenlabs), "ElevenLabs", "_generate_elevenlabs",
"ElevenLabs provider selected but 'elevenlabs' package not installed. Run: pip install elevenlabs"),
@@ -233,8 +206,7 @@ _BUILTIN_DISPATCH: Dict[str, tuple] = {
"piper": (lambda: _importable(_import_piper), "Piper (local)", "_generate_piper_tts",
"Piper provider selected but 'piper-tts' package not installed. "
"Run 'hermes tools' and select Piper under TTS, or install manually: "
"pip install piper-tts"),
}
"pip install piper-tts")}
def _error_json(message: str) -> str:
@@ -243,23 +215,22 @@ def _error_json(message: str) -> str:
def _run_edge_tts(text: str, file_str: str, tts_config: Dict[str, Any]) -> None:
"""Run the async Edge generator from sync code (worker thread; direct run if that fails)."""
run = lambda: asyncio.run(_generate_edge_tts(text, file_str, tts_config)) # noqa: E731
try:
import concurrent.futures
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
pool.submit(lambda: asyncio.run(_generate_edge_tts(text, file_str, tts_config))).result(timeout=60)
from concurrent.futures import ThreadPoolExecutor
with ThreadPoolExecutor(max_workers=1) as pool:
pool.submit(run).result(timeout=60)
except RuntimeError:
asyncio.run(_generate_edge_tts(text, file_str, tts_config))
run()
def _select_builtin_engine(provider: str) -> tuple:
"""Check a built-in provider's SDK -> ``(engine, None)`` or ``(provider, error_json)``. Unknown
names take the Edge default; without edge-tts NeuTTS is the fallback (engine != provider)."""
"""SDK check -> ``(engine, None)`` or ``(provider, error_json)``. Unknown 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():
@@ -267,8 +238,7 @@ def _select_builtin_engine(provider: str) -> tuple:
return "neutts", None
return provider, _error_json(
"No TTS provider available. Install edge-tts (pip install edge-tts) "
"or set up NeuTTS for local synthesis."
)
"or set up NeuTTS for local synthesis.")
def _synthesize_builtin(engine: str, text: str, file_str: str, tts_config: Dict[str, Any], instructions: Optional[str]) -> None:
@@ -286,7 +256,7 @@ def _synthesize_builtin(engine: str, text: str, file_str: str, tts_config: Dict[
def _finalize_voice_delivery(
file_str: str, provider: str, command_provider_config: Optional[Dict[str, Any]], want_opus: bool,
) -> tuple:
"""Decide voice-bubble eligibility, Opus-converting when needed -> ``(path, voice_compatible)``.
"""Voice-bubble eligibility (Opus-converting when needed) -> ``(path, voice_compatible)``.
Command/plugin providers are documents unless they opt in via ``voice_compatible``; native-Opus
built-ins qualify when the platform wants Opus and they wrote .ogg; MP3/WAV built-ins are
@@ -298,11 +268,9 @@ def _finalize_voice_delivery(
elif want_opus and provider in _FFMPEG_OPUS_PROVIDERS and not file_str.endswith(".ogg"):
opus_path = _convert_to_opus(file_str)
return (opus_path, True) if opus_path else (file_str, False)
elif provider in _NATIVE_OPUS_PROVIDERS:
return file_str, want_opus and file_str.endswith(".ogg")
else:
return file_str, False
native = provider in _NATIVE_OPUS_PROVIDERS
return file_str, native and want_opus and file_str.endswith(".ogg")
if not opted_in:
return file_str, False
if not file_str.endswith(".ogg"):
@@ -311,14 +279,12 @@ def _finalize_voice_delivery(
# --- Main tool function ---
def _apply_call_overrides(tts_config: Dict[str, Any], speed: Optional[float], provider: Optional[str]):
"""Apply per-call ``speed`` (clamped, on a shallow copy) and resolve the provider name."""
"""Apply per-call ``speed`` (clamped, on a shallow copy so the cached config isn't mutated) and
resolve the provider name."""
if speed is not None:
tts_config = dict(tts_config) # shallow copy to avoid mutating the cache
tts_config["speed"] = max(0.25, min(4.0, float(speed)))
provider = provider.lower().strip() if provider else _get_provider(tts_config)
return tts_config, provider
tts_config = {**tts_config, "speed": max(0.25, min(4.0, float(speed)))}
return tts_config, provider.lower().strip() if provider else _get_provider(tts_config)
def _session_platform() -> tuple:
@@ -341,9 +307,7 @@ def _resolve_output_base(
if has_traversal_component(output_path):
return None, _error_json(
f"output_path contains '..' traversal component: {output_path}. "
"Use an absolute path or one relative to the current directory "
"without '..'."
)
"Use an absolute path or one relative to the current directory without '..'.")
file_path = Path(output_path).expanduser()
if command_provider_config is not None:
file_path = _configured_command_tts_output_path(file_path, command_provider_config)
@@ -351,21 +315,24 @@ def _resolve_output_base(
if is_write_denied(str(file_path)) or is_write_approval_required(str(file_path)):
return None, _error_json(
f"output_path targets a protected credential or system path: "
f"{file_path}. Choose a normal audio output location."
)
f"{file_path}. Choose a normal audio output location.")
else:
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S_%f")
out_dir = Path(_default_output_dir())
out_dir.mkdir(parents=True, exist_ok=True)
if command_provider_config is not None:
ext = _get_command_tts_output_format(command_provider_config)
else:
ext = "ogg" if want_opus and provider in _NATIVE_OPUS_PROVIDERS else "mp3"
file_path = out_dir / f"tts_{timestamp}.{ext}"
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S_%f")
file_path = Path(_default_output_dir()) / f"tts_{timestamp}.{ext}"
file_path.parent.mkdir(parents=True, exist_ok=True)
return file_path, None
def _media_tag(paths: List[str], voice_compatible: bool) -> str:
"""``MEDIA:<path>`` lines; the ``[[audio_as_voice]]`` marker asks the platform for a voice bubble."""
media_tag = "\n".join(f"MEDIA:{path}" for path in paths)
return f"[[audio_as_voice]]\n{media_tag}" if voice_compatible else media_tag
def _tool_failure(prefix: str, provider: str, exc: BaseException) -> str:
"""Log and wrap a synthesis failure as the standard error envelope (traceback except for config errors)."""
error_msg = f"{prefix} ({provider}): {exc}"
@@ -377,7 +344,7 @@ def _text_to_speech_single(
text: str, file_str: str, *, provider: str, tts_config: Dict[str, Any],
command_provider_config: Optional[Dict[str, Any]], want_opus: bool, instructions: Optional[str],
) -> str:
"""Synthesize one normalized, provider-safe chunk into *file_str*; returns the result envelope.
"""Synthesize one provider-safe chunk into *file_str*; returns the result envelope.
Command providers resolve BEFORE built-in dispatch, but built-in names short-circuit so
``tts.providers.openai.command`` can't shadow OpenAI. Plugins fire only for names that are
@@ -385,7 +352,8 @@ def _text_to_speech_single(
try:
if command_provider_config is not None:
logger.info("Generating speech with command TTS provider '%s'...", provider)
file_str = _generate_command_tts(text, file_str, provider, command_provider_config, tts_config)
file_str = _generate_command_tts(
text, file_str, provider, command_provider_config, tts_config)
elif provider not in BUILTIN_TTS_PROVIDERS and (
_plugin_path := _dispatch_to_plugin_provider(text, file_str, provider, tts_config)
) is not None:
@@ -395,24 +363,17 @@ def _text_to_speech_single(
if error:
return error
_synthesize_builtin(provider, text, file_str, tts_config, instructions)
if not os.path.exists(file_str) or os.path.getsize(file_str) == 0:
return _error_json(f"TTS generation produced no output (provider: {provider})")
# Sniff once for every provider: MP3/WAV bytes in a .ogg path render as 0-second bubbles.
file_str = _repair_ogg_container(file_str)
file_str, voice_compatible = _finalize_voice_delivery(file_str, provider, command_provider_config, want_opus)
file_str, voice_compatible = _finalize_voice_delivery(
file_str, provider, command_provider_config, want_opus)
logger.info("TTS audio saved: %s (%s bytes, provider: %s)", file_str, f"{os.path.getsize(file_str):,}", provider)
media_tag = f"MEDIA:{file_str}"
if voice_compatible:
media_tag = f"[[audio_as_voice]]\n{media_tag}"
return json.dumps({
"success": True,
"file_path": file_str,
"media_tag": media_tag,
"provider": provider,
"voice_compatible": voice_compatible,
"success": True, "file_path": file_str, "media_tag": _media_tag([file_str], voice_compatible),
"provider": provider, "voice_compatible": voice_compatible,
}, ensure_ascii=False)
except ValueError as e:
return _tool_failure("TTS configuration error", provider, e)
@@ -458,8 +419,7 @@ def _synthesize_chunks(chunks: List[str], base_path: Path, generated_artifacts:
def text_to_speech_tool(
text: str, output_path: Optional[str] = None, speed: Optional[float] = None,
instructions: Optional[str] = None, provider: Optional[str] = None,
) -> str:
instructions: Optional[str] = None, provider: Optional[str] = None) -> str:
"""Convert text to speech with long-form chunking; returns the JSON result envelope.
Text is normalized, split into provider-safe chunks (never silently truncated), synthesized
@@ -467,16 +427,13 @@ def text_to_speech_tool(
separate valid files and no over-limit artifact is ever returned."""
if not text or not text.strip():
return tool_error("Text is required", success=False)
# Shared cleaner: markdown, emoji, think blocks, verifier footer, units, newlines.
try:
try: # shared cleaner: markdown, emoji, think blocks, verifier footer, units, newlines
from tools.tts_text_normalize import prepare_spoken_text
text = prepare_spoken_text(text, max_chars=None)
except Exception:
text = text.strip()
if not text:
return tool_error("Text is empty after TTS cleanup", success=False)
tts_config, provider = _apply_call_overrides(_load_tts_config(), speed, provider)
command_provider_config = _resolve_command_provider_config(provider, tts_config)
max_len = _resolve_max_text_length(provider, tts_config)
@@ -484,51 +441,36 @@ 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)
base_path, error = _resolve_output_base(output_path, provider, command_provider_config, want_opus)
base_path, error = _resolve_output_base(
output_path, provider, command_provider_config, want_opus)
if error:
return error
generated_artifacts: set[str] = set()
final_paths: List[str] = []
try:
encoded_paths, chunk_results = _synthesize_chunks(
chunks, base_path, generated_artifacts, provider=provider, tts_config=tts_config,
command_provider_config=command_provider_config, want_opus=want_opus,
instructions=instructions,
)
instructions=instructions)
voice_compatible = bool(chunk_results) and all(bool(r.get("voice_compatible")) for r in chunk_results)
delivery_base = base_path.with_suffix(Path(encoded_paths[0]).suffix)
final_paths, combined_chunks = _build_audio_delivery_files(
encoded_paths, str(delivery_base), delivery_profile, voice_compatible=voice_compatible,
)
encoded_paths, str(delivery_base), delivery_profile, voice_compatible=voice_compatible)
for path in final_paths:
logger.info("TTS audio saved: %s (%s bytes, provider: %s)", path, f"{os.path.getsize(path):,}", provider)
media_tag = "\n".join(f"MEDIA:{path}" for path in final_paths)
if voice_compatible:
media_tag = f"[[audio_as_voice]]\n{media_tag}"
return json.dumps({
"success": True,
"file_path": final_paths[0],
"file_paths": final_paths,
"media_tag": media_tag,
"provider": chunk_results[0].get("provider", provider),
"voice_compatible": voice_compatible,
"chunk_count": len(chunks),
"delivery_file_count": len(final_paths),
"success": True, "file_path": final_paths[0], "file_paths": final_paths,
"media_tag": _media_tag(final_paths, voice_compatible),
"provider": chunk_results[0].get("provider", provider), "voice_compatible": voice_compatible,
"chunk_count": len(chunks), "delivery_file_count": len(final_paths),
"combined_chunks": bool(combined_chunks),
"delivery_profile": {
"platform": delivery_profile.platform,
"max_file_bytes": delivery_profile.max_file_bytes,
"target_file_bytes": delivery_profile.target_file_bytes,
},
"platform": delivery_profile.platform, "max_file_bytes": delivery_profile.max_file_bytes,
"target_file_bytes": delivery_profile.target_file_bytes},
}, ensure_ascii=False)
except _ChunkFailed as exc:
return tool_error(str(exc), success=False)
@@ -540,26 +482,21 @@ def text_to_speech_tool(
final_absolute = {os.path.abspath(path) for path in final_paths}
for artifact in generated_artifacts:
if os.path.abspath(artifact) not in final_absolute:
try:
os.unlink(artifact)
except OSError:
pass
_remove_quietly(artifact)
# --- check_fn ---
def _minimax_requirements() -> bool:
try:
_resolve_minimax_tts_runtime(_load_tts_config())
return True
except ValueError:
return False
return True
def _xai_requirements() -> bool:
try:
from tools.xai_http import resolve_xai_http_credentials
return bool(resolve_xai_http_credentials().get("api_key"))
except Exception:
return False
@@ -578,8 +515,7 @@ _BUILTIN_REQUIREMENTS: Dict[str, Callable[[], bool]] = {
"mistral": lambda: _importable(_import_mistral_client) and bool(_resolve_provider_key("MISTRAL_API_KEY", "mistral")),
"neutts": lambda: _check_neutts_available(),
"kittentts": lambda: _check_kittentts_available(),
"piper": lambda: _check_piper_available(),
}
"piper": lambda: _check_piper_available()}
def check_tts_requirements() -> bool:
@@ -589,9 +525,7 @@ def check_tts_requirements() -> bool:
if _resolve_command_provider_config(provider, tts_config) is not None:
return True
check = _BUILTIN_REQUIREMENTS.get(provider)
if check is not None:
return check()
return _plugin_provider_is_available(provider)
return check() if check is not None else _plugin_provider_is_available(provider)
# --- Registry ---
@@ -645,10 +579,6 @@ 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="🔊",
)
emoji="🔊")
+98 -156
View File
@@ -1,10 +1,9 @@
"""Long-form chunking, ffmpeg encoding, container repair and delivery packing.
"""Long-form chunking, ffmpeg encoding, container repair and delivery packing (``tools.tts_tool``).
Provider-agnostic post-processing for ``tools.tts_tool``: split text under a
per-request cap, wrap raw PCM as WAV, convert WAV/MP3 to the target container,
sniff/repair mislabelled ``.ogg`` files, and combine final-encoded chunks under a
destination platform's upload limit. Origin module re-imports every name under
its historical spelling.
Provider-agnostic post-processing: split text under a per-request cap, wrap raw PCM as WAV,
convert WAV/MP3 to the target container, sniff/repair mislabelled ``.ogg`` files, and combine
final-encoded chunks under a destination platform's upload limit. Also home of the sibling
helpers ``_origin`` / ``_section`` / ``_remove_quietly``.
"""
from __future__ import annotations
@@ -24,19 +23,27 @@ from typing import Any, Dict, List, Optional, Tuple
from hermes_cli._subprocess_compat import windows_hide_flags
from tools.tts_command_provider import (
BUILTIN_TTS_PROVIDERS,
DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH,
_get_named_provider_config,
_is_command_provider_config,
)
BUILTIN_TTS_PROVIDERS, DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH, _get_named_provider_config,
_is_command_provider_config)
logger = logging.getLogger("tools.tts_tool")
# Final fallback when provider isn't recognised at all.
FALLBACK_MAX_TEXT_LENGTH = 4000
# Per-provider input-character caps (from official provider docs); override
# via ``tts.<provider>.max_text_length``.
def _origin():
"""``tools.tts_tool``, resolved per call so seams monkeypatched there still apply."""
from tools import tts_tool
return tts_tool
def _section(tts_config: Any, key: str) -> Dict[str, Any]:
"""``tts.<key>`` as a dict (``null``/non-dict sections read as empty)."""
section = tts_config.get(key) if isinstance(tts_config, dict) else None
return section if isinstance(section, dict) else {}
FALLBACK_MAX_TEXT_LENGTH = 4000 # provider not recognised at all
# Per-provider input-character caps (official docs); override: ``tts.<provider>.max_text_length``.
PROVIDER_MAX_TEXT_LENGTH: Dict[str, int] = {
"edge": 5000, # edge-tts practical sync limit
"openai": 4096, # https://platform.openai.com/docs/guides/text-to-speech
@@ -52,22 +59,15 @@ PROVIDER_MAX_TEXT_LENGTH: Dict[str, int] = {
# ElevenLabs caps vary by model_id. https://elevenlabs.io/docs/overview/models
ELEVENLABS_MODEL_MAX_TEXT_LENGTH: Dict[str, int] = {
"eleven_v3": 5000,
"eleven_ttv_v3": 5000,
"eleven_multilingual_v2": 10000,
"eleven_multilingual_v1": 10000,
"eleven_english_sts_v2": 10000,
"eleven_english_sts_v1": 10000,
"eleven_flash_v2": 30000,
"eleven_flash_v2_5": 40000,
}
"eleven_v3": 5000, "eleven_ttv_v3": 5000,
"eleven_multilingual_v2": 10000, "eleven_multilingual_v1": 10000,
"eleven_english_sts_v2": 10000, "eleven_english_sts_v1": 10000,
"eleven_flash_v2": 30000, "eleven_flash_v2_5": 40000}
def _positive_int(value: Any) -> Optional[int]:
"""*value* when it is a positive non-bool int, else None."""
if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
return None
return value
return value if isinstance(value, int) and not isinstance(value, bool) and value > 0 else None
def _resolve_max_text_length(provider: Optional[str], tts_config: Optional[Dict[str, Any]] = None) -> int:
@@ -79,16 +79,12 @@ def _resolve_max_text_length(provider: Optional[str], tts_config: Optional[Dict[
return FALLBACK_MAX_TEXT_LENGTH
key = provider.lower().strip()
cfg = tts_config or {}
prov_cfg = cfg.get(key)
if not isinstance(prov_cfg, dict):
prov_cfg = {}
prov_cfg = _section(cfg, key)
override = _positive_int(prov_cfg.get("max_text_length"))
if override:
return override
if key == "elevenlabs":
from tools.tts_tool_providers import DEFAULT_ELEVENLABS_MODEL_ID # providers imports this module
model_id = prov_cfg.get("model_id") or DEFAULT_ELEVENLABS_MODEL_ID
mapped = ELEVENLABS_MODEL_MAX_TEXT_LENGTH.get(str(model_id).strip())
if mapped:
@@ -103,19 +99,15 @@ def _resolve_max_text_length(provider: Optional[str], tts_config: Optional[Dict[
# PCM output specs for Gemini TTS (fixed by the API): 24kHz mono 16-bit (L16).
GEMINI_TTS_SAMPLE_RATE = 24000
GEMINI_TTS_CHANNELS = 1
GEMINI_TTS_SAMPLE_WIDTH = 2
GEMINI_TTS_SAMPLE_RATE, GEMINI_TTS_CHANNELS, GEMINI_TTS_SAMPLE_WIDTH = 24000, 1, 2
# ffmpeg args producing the Ogg/Opus voice-bubble encoding Telegram & co expect.
_OPUS_VOICE_ARGS = [
"-acodec", "libopus", "-ac", "1", "-b:a", "48k", "-vbr", "on",
"-application", "voip", "-compression_level", "10",
]
"-application", "voip", "-compression_level", "10"]
# --- Text chunking and delivery profiles ---
@dataclass(frozen=True)
class AudioDeliveryProfile:
"""Destination-platform constraints for generated TTS audio."""
@@ -133,34 +125,29 @@ class AudioDeliveryProfile:
_PLATFORM_AUDIO_DEFAULTS: Dict[str, Dict[str, Any]] = {
"discord": {"max_file_bytes": 10 * 1024 * 1024, "safety_ratio": 0.85},
"telegram": {"max_file_bytes": 50 * 1024 * 1024, "safety_ratio": 0.85},
"default": {"max_file_bytes": 10 * 1024 * 1024, "safety_ratio": 0.85},
}
"default": {"max_file_bytes": 10 * 1024 * 1024, "safety_ratio": 0.85}}
def _resolve_audio_delivery_profile(
platform: Optional[str], tts_config: Optional[Dict[str, Any]] = None,
) -> AudioDeliveryProfile:
platform: Optional[str], tts_config: Optional[Dict[str, Any]] = None) -> AudioDeliveryProfile:
"""Resolve upload constraints, including optional ``tts.delivery_profiles`` overrides."""
key = (platform or "default").lower().strip() or "default"
defaults = dict(_PLATFORM_AUDIO_DEFAULTS.get(key) or _PLATFORM_AUDIO_DEFAULTS["default"])
profiles = (tts_config or {}).get("delivery_profiles")
overrides = profiles.get(key, {}) if isinstance(profiles, dict) else {}
if isinstance(overrides, dict):
defaults.update({k: v for k, v in overrides.items() if v is not None})
max_file_bytes = _positive_int(defaults.get("max_file_bytes")) or _PLATFORM_AUDIO_DEFAULTS["default"]["max_file_bytes"]
overrides = _section(_section(tts_config, "delivery_profiles"), key)
defaults.update({k: v for k, v in overrides.items() if v is not None})
max_file_bytes = (_positive_int(defaults.get("max_file_bytes"))
or _PLATFORM_AUDIO_DEFAULTS["default"]["max_file_bytes"])
safety_ratio = defaults.get("safety_ratio", 0.85)
if isinstance(safety_ratio, bool) or not isinstance(safety_ratio, (int, float)) or not 0 < safety_ratio <= 1:
if (isinstance(safety_ratio, bool) or not isinstance(safety_ratio, (int, float))
or not 0 < safety_ratio <= 1):
safety_ratio = 0.85
return AudioDeliveryProfile(platform=key, max_file_bytes=max_file_bytes, safety_ratio=float(safety_ratio))
def _pack_under_cap(pieces: List[str], max_chars: int, *, slice_oversized: bool = False) -> List[str]:
"""Greedily join *pieces* with single spaces, starting a new chunk past *max_chars*.
With ``slice_oversized`` an over-long piece flushes the running chunk and emits its hard
slices as their own chunks (the tail slice is not merged with following pieces).
"""
"""Greedily join *pieces* with single spaces, starting a new chunk past *max_chars*. With
``slice_oversized`` an over-long piece flushes the running chunk and emits its hard slices as
their own chunks (the tail slice is not merged with following pieces)."""
chunks: List[str] = []
current = ""
for piece in pieces:
@@ -195,12 +182,8 @@ def _split_text_for_tts(text: str, max_chars: int) -> List[str]:
return []
if len(normalized) <= max_chars:
return [normalized]
expanded: List[str] = []
for sentence in re.split(r"(?<=[.!?;:,])\s+", normalized):
sentence = sentence.strip()
if not sentence:
continue
for sentence in filter(None, (s.strip() for s in re.split(r"(?<=[.!?;:,])\s+", normalized))):
if len(sentence) <= max_chars:
expanded.append(sentence)
else:
@@ -212,46 +195,38 @@ def _pack_audio_files_for_delivery(audio_paths: List[str], profile: AudioDeliver
"""Group final-encoded chunks under the size target; never mixes suffixes (can't concat-copy)."""
groups: List[List[str]] = []
current: List[str] = []
current_size = 0
current_suffix = ""
current_size, current_suffix = 0, ""
for path in audio_paths:
size = Path(path).stat().st_size
suffix = Path(path).suffix.lower()
size, suffix = Path(path).stat().st_size, Path(path).suffix.lower()
if current and (current_size + size > profile.target_file_bytes or suffix != current_suffix):
groups.append(current)
current, current_size = [], 0
current.append(path)
current_size += size
current_suffix = suffix
if current:
groups.append(current)
return groups
current_size, current_suffix = current_size + size, suffix
return groups + [current] if current else groups
# --- ffmpeg encoding helpers ---
def _ffmpeg_run(
ffmpeg: str, args: List[str], *, timeout: int = 30, check: bool = False, capture: bool = True,
) -> subprocess.CompletedProcess:
"""Run ``ffmpeg <args>`` headless (no stdin, hidden window on Windows)."""
return subprocess.run(
[ffmpeg, *args],
capture_output=capture, check=check, timeout=timeout, stdin=subprocess.DEVNULL, creationflags=windows_hide_flags(),
)
return subprocess.run([ffmpeg, *args], capture_output=capture, check=check, timeout=timeout,
stdin=subprocess.DEVNULL, creationflags=windows_hide_flags())
def _remove_quietly(path) -> None:
try:
os.remove(path)
except OSError:
pass
def _remove_quietly(path: Optional[str]) -> None:
"""Best-effort unlink (missing/locked files are ignored); None is a no-op."""
if path:
try:
os.remove(path)
except OSError:
pass
def _wav_sidecar_path(output_path: str) -> str:
"""Path a WAV-native engine writes to before conversion to *output_path*'s format."""
if output_path.endswith(".wav"):
return output_path
return output_path.rsplit(".", 1)[0] + ".wav"
return output_path if output_path.endswith(".wav") else output_path.rsplit(".", 1)[0] + ".wav"
def _finalize_wav_output(wav_path: str, output_path: str) -> str:
@@ -260,44 +235,37 @@ def _finalize_wav_output(wav_path: str, output_path: str) -> str:
if wav_path == output_path:
return output_path
ffmpeg = shutil.which("ffmpeg")
if ffmpeg:
_ffmpeg_run(ffmpeg, ["-i", wav_path, "-y", "-loglevel", "error", output_path], check=True, capture=False)
_remove_quietly(wav_path)
else:
if not ffmpeg:
os.rename(wav_path, output_path)
return output_path
_ffmpeg_run(ffmpeg, ["-i", wav_path, "-y", "-loglevel", "error", output_path],
check=True, capture=False)
_remove_quietly(wav_path)
return output_path
def _wrap_pcm_as_wav(
pcm_bytes: bytes,
sample_rate: int = GEMINI_TTS_SAMPLE_RATE,
channels: int = GEMINI_TTS_CHANNELS,
sample_width: int = GEMINI_TTS_SAMPLE_WIDTH,
) -> bytes:
pcm_bytes: bytes, sample_rate: int = GEMINI_TTS_SAMPLE_RATE,
channels: int = GEMINI_TTS_CHANNELS, sample_width: int = GEMINI_TTS_SAMPLE_WIDTH) -> bytes:
"""Wrap raw signed-little-endian PCM (e.g. Gemini's L16) with a minimal WAV RIFF header."""
byte_rate = sample_rate * channels * sample_width
block_align = channels * sample_width
data_size = len(pcm_bytes)
fmt_chunk = struct.pack(
"<4sIHHIIHH", b"fmt ", 16, 1, channels, sample_rate, byte_rate, block_align, sample_width * 8,
)
data_chunk_header = struct.pack("<4sI", b"data", data_size)
riff_size = 4 + len(fmt_chunk) + len(data_chunk_header) + data_size
fmt_chunk = struct.pack("<4sIHHIIHH", b"fmt ", 16, 1, channels, sample_rate,
sample_rate * block_align, block_align, sample_width * 8)
data_chunk_header = struct.pack("<4sI", b"data", len(pcm_bytes))
riff_size = 4 + len(fmt_chunk) + len(data_chunk_header) + len(pcm_bytes)
riff_header = struct.pack("<4sI4s", b"RIFF", riff_size, b"WAVE")
return riff_header + fmt_chunk + data_chunk_header + pcm_bytes
def _write_wav_bytes_as(wav_bytes: bytes, output_path: str) -> str:
"""Write in-memory WAV to *output_path*, ffmpeg-converting to its container.
``.ogg`` is forced to Opus (ffmpeg's .ogg default is Vorbis, which voice bubbles reject).
A failed conversion raises RuntimeError; without ffmpeg the raw WAV is written under
the requested name (misleading extension, but the audio still plays)."""
"""Write in-memory WAV to *output_path*, ffmpeg-converting to its container. ``.ogg`` is forced
to Opus (ffmpeg's .ogg default is Vorbis, which voice bubbles reject). A failed conversion
raises RuntimeError; without ffmpeg the raw WAV is written under the requested name
(misleading extension, but the audio still plays)."""
if output_path.lower().endswith(".wav"):
with open(output_path, "wb") as f:
f.write(wav_bytes)
return output_path
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
tmp.write(wav_bytes)
wav_path = tmp.name
@@ -305,7 +273,8 @@ def _write_wav_bytes_as(wav_bytes: bytes, output_path: str) -> str:
ffmpeg = shutil.which("ffmpeg")
if ffmpeg:
opus = _OPUS_VOICE_ARGS if output_path.lower().endswith(".ogg") else []
result = _ffmpeg_run(ffmpeg, ["-i", wav_path, *opus, "-y", "-loglevel", "error", output_path])
result = _ffmpeg_run(
ffmpeg, ["-i", wav_path, *opus, "-y", "-loglevel", "error", output_path])
if result.returncode != 0:
stderr = result.stderr.decode("utf-8", errors="ignore")[:300]
raise RuntimeError(f"ffmpeg conversion failed: {stderr}")
@@ -326,14 +295,13 @@ def _ffmpeg_transcode_to_opus(input_path: str, ogg_path: str) -> Optional[str]:
"""Transcode *input_path* to real Ogg/Opus at *ogg_path* (in-place safe via temp file); None on failure."""
if shutil.which("ffmpeg") is None:
return None
in_place = os.path.abspath(input_path) == os.path.abspath(ogg_path)
work_path = ogg_path + ".tmp.ogg" if in_place else ogg_path
try:
result = _ffmpeg_run("ffmpeg", ["-i", input_path, *_OPUS_VOICE_ARGS, "-f", "ogg", work_path, "-y"])
if result.returncode != 0:
logger.warning("ffmpeg conversion failed with return code %d: %s",
result.returncode, result.stderr.decode('utf-8', errors='ignore')[:200])
result.returncode, result.stderr.decode('utf-8', errors='ignore')[:200])
return None
if os.path.exists(work_path) and os.path.getsize(work_path) > 0:
if in_place:
@@ -352,86 +320,68 @@ def _ffmpeg_transcode_to_opus(input_path: str, ogg_path: str) -> Optional[str]:
# --- Container sniffing / repair ---
# Several backends silently ignore the requested opus format (Edge only emits MP3, Piper
# writes WAV, xAI writes MP3, some OpenAI-compatible servers ignore response_format="opus"),
# which breaks native voice bubbles. Sniff the magic bytes once after synthesis and repair
# when they don't match the extension.
# Several backends ignore the requested opus format (Edge/xAI emit MP3, Piper WAV, some
# OpenAI-compatible servers ignore response_format="opus"), which breaks native voice bubbles:
# sniff the magic bytes once after synthesis and repair when they don't match the extension.
def _sniff_audio_container(path: str) -> str:
"""Return a container id ('ogg', 'wav', 'mp3', 'flac', ...) or 'unknown'."""
from tools.audio_container import sniff_container
try:
with open(path, "rb") as fh:
head = fh.read(12)
return sniff_container(fh.read(12)) or "unknown"
except OSError:
return "unknown"
return sniff_container(head) or "unknown"
def _repair_ogg_container(file_str: str) -> str:
"""Ensure a ``.ogg`` path really holds Ogg: transcode in place, else rename to the sniffed
real extension so platforms get an honest file instead of a 0-second voice bubble."""
if not file_str.endswith(".ogg"):
return file_str
container = _sniff_audio_container(file_str)
container = _sniff_audio_container(file_str) if file_str.endswith(".ogg") else "ogg"
if container in ("ogg", "unknown"):
return file_str
logger.info("TTS wrote %s bytes into a .ogg path (%s) — transcoding to real Ogg/Opus", container, file_str)
repaired = _ffmpeg_transcode_to_opus(file_str, file_str)
if repaired:
return repaired
honest = file_str[:-4] + "." + container
honest = f"{file_str[:-4]}.{container}"
try:
os.replace(file_str, honest)
logger.warning(
"Could not transcode %s to Ogg/Opus — renamed to %s so the "
"file is delivered with its real format", file_str, honest,
)
return honest
except OSError:
return file_str
logger.warning("Could not transcode %s to Ogg/Opus — renamed to %s so the "
"file is delivered with its real format", file_str, honest)
return honest
# --- Long-form audio combination and delivery packing ---
def _concat_audio_files(audio_paths: List[str], output_path: str, *, voice_compatible: bool = False) -> Optional[str]:
"""Combine independently encoded chunks with ffmpeg (never byte-joined).
OGG/Opus is always re-encoded (even without voice opt-in); matching MP3 chunks keep their
frames (``-c:a copy``). None when ffmpeg is missing/fails so callers keep the valid parts."""
"""Combine independently encoded chunks with ffmpeg (never byte-joined). OGG/Opus is always
re-encoded (even without voice opt-in); matching MP3 chunks keep their frames (``-c:a copy``).
None when ffmpeg is missing/fails so callers keep the valid parts."""
if not audio_paths:
raise ValueError("No audio chunks to combine")
if len(audio_paths) == 1:
source = audio_paths[0]
if os.path.abspath(source) != os.path.abspath(output_path):
shutil.copyfile(source, output_path)
if os.path.abspath(audio_paths[0]) != os.path.abspath(output_path):
shutil.copyfile(audio_paths[0], output_path)
return output_path
ffmpeg = shutil.which("ffmpeg")
if not ffmpeg:
return None
destination = Path(output_path)
destination.parent.mkdir(parents=True, exist_ok=True)
concat_path = destination.with_name(f".{destination.name}.{uuid.uuid4().hex}.concat.txt")
temp_output = destination.with_name(f".{destination.stem}.{uuid.uuid4().hex}.combining{destination.suffix}")
try:
with concat_path.open("w", encoding="utf-8") as concat_file:
for path in audio_paths:
concat_file.write(f"file {shlex.quote(os.path.abspath(path))}\n")
entries = "".join(f"file {shlex.quote(os.path.abspath(p))}\n" for p in audio_paths)
concat_path.write_text(entries, encoding="utf-8")
args = ["-y", "-loglevel", "error", "-f", "concat", "-safe", "0", "-i", str(concat_path), "-vn"]
suffix = destination.suffix.lower()
if voice_compatible or suffix in {".ogg", ".opus"}:
args += ["-c:a", "libopus", "-ac", "1", "-b:a", "64k", "-vbr", "off"]
elif suffix == ".mp3" and all(Path(path).suffix.lower() == ".mp3" for path in audio_paths):
args += ["-c:a", "copy"]
args.append(str(temp_output))
result = _ffmpeg_run(ffmpeg, args, timeout=120)
result = _ffmpeg_run(ffmpeg, [*args, str(temp_output)], timeout=120)
if result.returncode == 0 and temp_output.exists() and temp_output.stat().st_size > 0:
os.replace(temp_output, destination)
return str(destination)
@@ -447,7 +397,7 @@ def _concat_audio_files(audio_paths: List[str], output_path: str, *, voice_compa
def _build_audio_delivery_files(
audio_paths: List[str], output_path: str, profile: AudioDeliveryProfile, *, voice_compatible: bool = False,
) -> Tuple[List[str], bool]:
"""Pack final-encoded chunks and enforce the hard upload limit; returns ``(final_paths, combined_any)``.
"""Pack final-encoded chunks under the hard upload limit -> ``(final_paths, combined_any)``.
Groups are packed against the conservative target, then each combined artifact is checked
at its real size; an over-limit group is split in half and retried. A failed combine
@@ -459,13 +409,10 @@ def _build_audio_delivery_files(
if size > profile.max_file_bytes:
raise ValueError(
f"Final-encoded TTS chunk exceeds {profile.platform} delivery "
f"limit ({size} > {profile.max_file_bytes} bytes): {path}"
)
f"limit ({size} > {profile.max_file_bytes} bytes): {path}")
base = Path(output_path)
scratch_outputs: List[str] = []
combined_any = False
combine_index = 0
combined_any, combine_index = False, 0
def emit(group: List[str]) -> List[str]:
nonlocal combined_any, combine_index
@@ -483,16 +430,12 @@ def _build_audio_delivery_files(
_remove_quietly(combined)
midpoint = max(1, len(group) // 2)
return emit(group[:midpoint]) + emit(group[midpoint:])
packed: List[str] = []
for group in _pack_audio_files_for_delivery(audio_paths, profile):
packed.extend(emit(group))
groups = _pack_audio_files_for_delivery(audio_paths, profile)
packed = [path for group in groups for path in emit(group)]
final_paths: List[str] = []
for index, source in enumerate(packed, start=1):
if len(packed) == 1:
destination = base
else:
destination = base
if len(packed) > 1:
destination = base.with_name(f"{base.stem}.part{index:02d}{Path(source).suffix or base.suffix}")
if os.path.abspath(source) != os.path.abspath(destination):
destination.parent.mkdir(parents=True, exist_ok=True)
@@ -500,7 +443,6 @@ def _build_audio_delivery_files(
if destination.stat().st_size > profile.max_file_bytes:
raise ValueError(f"Final TTS deliverable exceeds {profile.platform} delivery limit: {destination}")
final_paths.append(str(destination))
try:
return final_paths, combined_any
finally:
+25 -55
View File
@@ -1,11 +1,10 @@
"""Local-engine lifecycle for ``tools.tts_tool``: warm-up / release leases.
Local engines load their model lazily on first synthesis (dead air on the first spoken
reply) and then stay resident forever. Every surface that flips speech output on holds a
*lease* here (warming the configured engine); when the last lease is released the local
model caches are dropped, so one surface's "off" can't unload a model another surface
still needs. Cloud providers have nothing resident; warming only ensures the SDK imports.
Seams tests monkeypatch on the origin are resolved through :func:`_origin` at call time.
Local engines load lazily on first synthesis (dead air on the first spoken reply) and then stay
resident. Every surface that flips speech output on holds a *lease* here (warming the configured
engine); when the last lease is released the local model caches are dropped, so one surface's
"off" can't unload a model another surface still needs. Cloud providers have nothing resident;
warming only ensures the SDK imports. Origin seams are resolved through :func:`_origin` per call.
"""
from __future__ import annotations
@@ -16,28 +15,16 @@ import time
from typing import Any, Callable, Dict, List, Optional
from tools.tts_command_provider import (
BUILTIN_TTS_PROVIDERS,
_get_command_tts_timeout,
_get_named_provider_config,
_is_command_provider_config,
command_env_passthrough as _command_provider_env_passthrough,
render_command_template as _render_command_tts_template,
)
BUILTIN_TTS_PROVIDERS, _get_command_tts_timeout, _get_named_provider_config,
_is_command_provider_config, command_env_passthrough as _command_provider_env_passthrough,
render_command_template as _render_command_tts_template)
from tools.tts_tool_delivery import _origin
from tools.tts_tool_local import (
_LOCAL_TTS_MODEL_CACHES, _load_kittentts_model_for_config, _load_piper_voice_for_config,
)
_LOCAL_TTS_MODEL_CACHES, _load_kittentts_model_for_config, _load_piper_voice_for_config)
from tools.tts_tool_plugins import _lookup_plugin_provider
logger = logging.getLogger("tools.tts_tool")
def _origin():
"""``tools.tts_tool``, resolved per call so monkeypatched seams there still apply."""
from tools import tts_tool
return tts_tool
_tts_lease_lock = threading.Lock()
_tts_leases: set = set()
@@ -46,8 +33,7 @@ def _local_tts_warmers() -> Dict[str, Callable[[Dict[str, Any]], Any]]:
"""Provider name → loader populating that engine's cache slot (same key synthesis uses)."""
return {
"piper": lambda cfg: _load_piper_voice_for_config(cfg)[0],
"kittentts": lambda cfg: _load_kittentts_model_for_config(cfg)[0],
}
"kittentts": lambda cfg: _load_kittentts_model_for_config(cfg)[0]}
# tools.lazy_deps feature key for providers whose SDK installs on first use.
@@ -56,7 +42,6 @@ _LAZY_SDK_FEATURES = {"edge": "tts.edge", "elevenlabs": "tts.elevenlabs", "mistr
def _signal_user_tts_provider(name: str, tts_config: Dict[str, Any], hook: str) -> Optional[str]:
"""Forward a lease ``hook`` (``"warm"``/``"release"``) to a user-declared provider; returns the action.
Command providers run their optional ``<hook>_command`` (same template/env/timeout rules as
``command``) on a background thread so a toggle never waits on a model server; plugins get
:meth:`TTSProvider.warm`/``release``. Best-effort: failures are logged at debug."""
@@ -71,16 +56,15 @@ def _signal_user_tts_provider(name: str, tts_config: Dict[str, Any], hook: str)
command = _render_command_tts_template(template, {
"voice": str(cfg.get("voice", "")),
"model": str(cfg.get("model", "")),
"speed": str(cfg.get("speed", tts_config.get("speed", ""))),
})
"speed": str(cfg.get("speed", tts_config.get("speed", "")))})
def _run() -> None:
try:
_origin()._run_command_tts(command, _get_command_tts_timeout(cfg),
env_passthrough=_command_provider_env_passthrough(cfg))
_origin()._run_command_tts(
command, _get_command_tts_timeout(cfg),
env_passthrough=_command_provider_env_passthrough(cfg))
except Exception as exc: # noqa: BLE001 — best-effort hook
logger.debug("[TTS] %s_command for %s failed: %s", hook, name, exc)
threading.Thread(target=_run, name=f"tts-{hook}-{name}", daemon=True).start()
return hook
plugin_provider = _lookup_plugin_provider(name)
@@ -95,7 +79,6 @@ def _signal_user_tts_provider(name: str, tts_config: Dict[str, Any], hook: str)
def warm_tts_provider(tts_config: Optional[Dict[str, Any]] = None, provider: Optional[str] = None) -> Dict[str, Any]:
"""Pre-load the configured TTS provider so the next synthesis starts hot (blocking; never raises).
Local engines fill the same LRU slot synthesis reads (including first-use download); lazily
installed cloud SDKs are made importable; user-declared providers get their warm hook;
everything else is ``action: "noop"``. The result carries ``warmed`` / ``action`` / ``error``."""
@@ -103,38 +86,30 @@ def warm_tts_provider(tts_config: Optional[Dict[str, Any]] = None, provider: Opt
tts_config = _origin()._load_tts_config()
name = (provider or _origin()._get_provider(tts_config) or "").lower().strip()
result: Dict[str, Any] = {"provider": name, "warmed": False, "action": "noop"}
warmer = _local_tts_warmers().get(name)
if warmer is not None:
cache = _LOCAL_TTS_MODEL_CACHES.get(name)
before = len(cache) if cache is not None else 0
started = time.monotonic()
cache = _LOCAL_TTS_MODEL_CACHES.get(name, {})
before, started = len(cache), time.monotonic()
try:
warmer(tts_config)
except Exception as exc: # engine missing, download failed, bad voice…
logger.warning("[TTS] warm-up for %s failed: %s", name, exc)
result.update(action="error", error=str(exc))
return result
after = len(cache) if cache is not None else 0
result.update(
warmed=True,
action="loaded" if after > before else "cached",
elapsed_ms=int((time.monotonic() - started) * 1000),
)
warmed=True, action="loaded" if len(cache) > before else "cached",
elapsed_ms=int((time.monotonic() - started) * 1000))
logger.info("[TTS] warm-up %s: %s in %dms", name, result["action"], result["elapsed_ms"])
return result
signalled = _signal_user_tts_provider(name, tts_config, "warm")
if signalled is not None:
ok = signalled != "error"
result.update(warmed=ok, action="warmed" if ok else "error")
return result
feature = _LAZY_SDK_FEATURES.get(name)
if feature is not None:
try:
from tools.lazy_deps import ensure, is_available
if is_available(feature):
result.update(warmed=True, action="cached")
else:
@@ -155,10 +130,9 @@ def release_tts_provider(provider: Optional[str] = None) -> Dict[str, Any]:
_signal_user_tts_provider(_origin()._get_provider(tts_config), tts_config, "release")
released = 0
for cache_name, cache in _LOCAL_TTS_MODEL_CACHES.items():
if name and cache_name != name:
continue
released += len(cache)
cache.clear()
if not name or cache_name == name:
released += len(cache)
cache.clear()
if released:
logger.info("[TTS] released %d resident local model(s)", released)
return {"released": released}
@@ -170,9 +144,7 @@ def acquire_tts_lease(lease: str, tts_config: Optional[Dict[str, Any]] = None) -
with _tts_lease_lock:
_tts_leases.add(lease)
holders = len(_tts_leases)
result = _origin().warm_tts_provider(tts_config)
result["leases"] = holders
return result
return {**_origin().warm_tts_provider(tts_config), "leases": holders}
def release_tts_lease(lease: str) -> Dict[str, Any]:
@@ -181,10 +153,8 @@ def release_tts_lease(lease: str) -> Dict[str, Any]:
with _tts_lease_lock:
_tts_leases.discard(lease)
holders = len(_tts_leases)
result: Dict[str, Any] = {"leases": holders, "released": 0}
if holders == 0:
result["released"] = release_tts_provider()["released"]
return result
released = release_tts_provider()["released"] if holders == 0 else 0
return {"leases": holders, "released": released}
def tts_lease_holders() -> List[str]:
+25 -57
View File
@@ -1,10 +1,9 @@
"""Local on-device TTS engines for ``tools.tts_tool``: NeuTTS, Piper, KittenTTS.
All three synthesize WAV natively; :func:`_finalize_wav_output` then converts/renames to
the caller's requested container. Piper and KittenTTS keep loaded models in small LRU
caches registered in ``_LOCAL_TTS_MODEL_CACHES`` so the warm/release lifecycle can
pre-load or drop them. ``_import_piper`` / ``_import_kittentts`` are resolved through the
origin module at call time so test monkeypatches there apply.
All three synthesize WAV natively; :func:`_finalize_wav_output` converts/renames to the requested
container. Piper and KittenTTS keep loaded models in small LRU caches registered in
``_LOCAL_TTS_MODEL_CACHES`` so warm/release can pre-load or drop them. ``_import_piper`` /
``_import_kittentts`` are resolved through the origin module at call time (test monkeypatches).
"""
from __future__ import annotations
@@ -15,7 +14,7 @@ import sys
from pathlib import Path
from typing import Any, Callable, Dict, Tuple
from tools.tts_tool_delivery import _finalize_wav_output, _wav_sidecar_path
from tools.tts_tool_delivery import _finalize_wav_output, _origin, _section, _wav_sidecar_path
logger = logging.getLogger("tools.tts_tool")
@@ -24,25 +23,18 @@ DEFAULT_KITTENTTS_VOICE = "Jasper"
DEFAULT_PIPER_VOICE = "en_US-lessac-medium" # balanced size/quality
_NEUTTS_SAMPLES = Path(__file__).parent / "neutts_samples"
def _origin():
from tools import tts_tool
return tts_tool
# --- Bounded model caches ---
# Each cached entry is a whole loaded model (tens of MB); an unbounded dict would pin one
# per distinct voice for the process lifetime. Most sessions use one or two voices and a
# cold reload is cheap.
# Each entry is a whole loaded model (tens of MB); unbounded, one would be pinned per distinct
# voice for the process lifetime. Most sessions use one or two voices; a cold reload is cheap.
_TTS_MODEL_CACHE_MAX = 3
# Provider name → the model cache it populates (consulted by warm/release in
# tts_tool_lifecycle; a new local engine adds a row here plus a loader in _local_tts_warmers()).
# Piper voices keyed on absolute .onnx path (+cuda flag); KittenTTS on model name.
# Provider name -> the cache it populates (warm/release in tts_tool_lifecycle; a new local engine
# adds a row here plus a loader in _local_tts_warmers()). Piper keyed on absolute .onnx path
# (+cuda flag); KittenTTS on model name.
_piper_voice_cache: Dict[str, Any] = {}
_kittentts_model_cache: Dict[str, Any] = {}
_LOCAL_TTS_MODEL_CACHES: Dict[str, Dict[str, Any]] = {"piper": _piper_voice_cache, "kittentts": _kittentts_model_cache}
_LOCAL_TTS_MODEL_CACHES: Dict[str, Dict[str, Any]] = {
"piper": _piper_voice_cache, "kittentts": _kittentts_model_cache}
def _tts_cache_get_or_load(cache: Dict[str, Any], key: str, load: Callable[[], Any]) -> Any:
@@ -58,10 +50,6 @@ def _tts_cache_get_or_load(cache: Dict[str, Any], key: str, load: Callable[[], A
return value
def _section(tts_config: Any, key: str) -> Dict[str, Any]:
return (tts_config.get(key) or {}) if isinstance(tts_config, dict) else {}
def _run_helper(cmd: list, timeout: int) -> subprocess.CompletedProcess:
return subprocess.run(
cmd, capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=timeout, stdin=subprocess.DEVNULL,
@@ -69,7 +57,6 @@ def _run_helper(cmd: list, timeout: int) -> subprocess.CompletedProcess:
# --- NeuTTS (subprocess via tools/neutts_synth.py so the ~500MB model exits after use) ---
def _generate_neutts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
neutts_config = tts_config.get("neutts") or {}
wav_path = _wav_sidecar_path(output_path)
@@ -80,18 +67,15 @@ def _generate_neutts(text: str, output_path: str, tts_config: Dict[str, Any]) ->
"--ref-audio", neutts_config.get("ref_audio", "") or str(_NEUTTS_SAMPLES / "jo.wav"),
"--ref-text", neutts_config.get("ref_text", "") or str(_NEUTTS_SAMPLES / "jo.txt"),
"--model", neutts_config.get("model", "neuphonic/neutts-air-q4-gguf"),
"--device", neutts_config.get("device", "cpu"),
]
"--device", neutts_config.get("device", "cpu")]
result = _run_helper(cmd, 120)
if result.returncode != 0:
# The synth script reports success lines as "OK:" on stderr too.
if result.returncode != 0: # the synth script reports success lines as "OK:" on stderr too
error_lines = [l for l in result.stderr.strip().splitlines() if not l.startswith("OK:")]
raise RuntimeError(f"NeuTTS synthesis failed: {chr(10).join(error_lines) or 'unknown error'}")
return _finalize_wav_output(wav_path, output_path)
# --- Piper (local neural VITS, 44 languages) ---
def _get_piper_voices_dir() -> Path:
"""``<HERMES_HOME>/cache/piper-voices/`` so voice downloads follow profile boundaries."""
from hermes_constants import get_hermes_dir
@@ -107,11 +91,9 @@ def _resolve_piper_voice_path(voice: str, download_dir: Path) -> str:
candidate = Path(voice).expanduser()
if candidate.suffix.lower() == ".onnx" and candidate.exists():
return str(candidate)
cached = download_dir / f"{voice}.onnx"
if cached.exists() and (download_dir / f"{voice}.onnx.json").exists():
return str(cached)
logger.info("[Piper] Downloading voice '%s' to %s (first use)", voice, download_dir)
try:
result = _run_helper(
@@ -126,8 +108,7 @@ def _resolve_piper_voice_path(voice: str, download_dir: Path) -> str:
raise RuntimeError(
f"Piper voice download completed but {cached} is missing — "
f"check voice name (see: https://github.com/OHF-Voice/piper1-gpl/"
f"blob/main/docs/VOICES.md)"
)
f"blob/main/docs/VOICES.md)")
return str(cached)
@@ -140,11 +121,7 @@ def _load_piper_voice_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, Dict[
download_dir = Path(piper_config.get("voices_dir") or _get_piper_voices_dir()).expanduser()
download_dir.mkdir(parents=True, exist_ok=True)
use_cuda = bool(piper_config.get("use_cuda", False))
model_path = _resolve_piper_voice_path(voice_name, download_dir)
# speaker_id is applied per call via syn_config, so one PiperVoice instance serves
# every speaker and stays out of the cache key.
cache_key = f"{model_path}::cuda={use_cuda}"
def _load_piper_voice():
logger.info("[Piper] Loading voice: %s", model_path)
@@ -152,6 +129,8 @@ def _load_piper_voice_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, Dict[
logger.info("[Piper] Voice loaded")
return v
# speaker_id is applied per call via syn_config, so one instance serves every speaker.
cache_key = f"{model_path}::cuda={use_cuda}"
return _tts_cache_get_or_load(_piper_voice_cache, cache_key, _load_piper_voice), piper_config
@@ -160,16 +139,12 @@ _PIPER_ADVANCED_KNOBS = ("length_scale", "noise_scale", "noise_w_scale", "volume
def _generate_piper_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
import wave
voice, piper_config = _load_piper_voice_for_config(tts_config)
# Bad speaker_id input drops to 0 (Piper's default); booleans are rejected outright
# since True/False would silently coerce to 1/0.
# Bad speaker_id drops to 0 (Piper's default); bools are rejected (they'd coerce to 1/0).
_raw_speaker = piper_config.get("speaker_id", 0)
speaker_id = 0 if isinstance(_raw_speaker, bool) or not isinstance(_raw_speaker, int) else _raw_speaker
# Only build a SynthesisConfig when an advanced knob is configured, so we don't
# depend on a newer piper-tts than the user's unless we must.
speaker_id = _raw_speaker if type(_raw_speaker) is int else 0
# Only build a SynthesisConfig when an advanced knob is configured, so we don't depend on a
# newer piper-tts than the user's unless we must.
syn_config = None
if any(k in piper_config for k in _PIPER_ADVANCED_KNOBS):
try:
@@ -180,11 +155,9 @@ def _generate_piper_tts(text: str, output_path: str, tts_config: Dict[str, Any])
noise_w_scale=float(piper_config.get("noise_w_scale", 0.8)),
volume=float(piper_config.get("volume", 1.0)),
normalize_audio=bool(piper_config.get("normalize_audio", True)),
speaker_id=speaker_id,
)
speaker_id=speaker_id)
except ImportError:
logger.warning("[Piper] SynthesisConfig not available in this piper-tts version — advanced knobs ignored")
wav_path = _wav_sidecar_path(output_path)
with wave.open(wav_path, "wb") as wav_file:
if syn_config is not None:
@@ -195,7 +168,6 @@ def _generate_piper_tts(text: str, output_path: str, tts_config: Dict[str, Any])
# --- KittenTTS (local ONNX, 25-80MB models, CPU only) ---
def _load_kittentts_model_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, Dict[str, Any]]:
"""Load (or fetch from cache) the KittenTTS model; returns ``(model, kittentts_config)``."""
KittenTTS = _origin()._import_kittentts()
@@ -213,13 +185,9 @@ def _load_kittentts_model_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, D
def _generate_kittentts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
model, kt_config = _load_kittentts_model_for_config(tts_config)
audio = model.generate(
text,
voice=kt_config.get("voice", DEFAULT_KITTENTTS_VOICE),
speed=kt_config.get("speed", 1.0),
clean_text=kt_config.get("clean_text", True),
) # numpy array at 24kHz
audio = model.generate( # numpy array at 24kHz
text, voice=kt_config.get("voice", DEFAULT_KITTENTTS_VOICE),
speed=kt_config.get("speed", 1.0), clean_text=kt_config.get("clean_text", True))
import soundfile as sf
wav_path = _wav_sidecar_path(output_path)
sf.write(wav_path, audio, 24000)
+28 -75
View File
@@ -15,20 +15,12 @@ from typing import Any, Dict, Optional
from urllib.parse import urljoin
from tools.tool_backend_helpers import (
NOUS_MANAGED_PROVIDER, nous_tool_gateway_unavailable_message, selection_error,
)
NOUS_MANAGED_PROVIDER, nous_tool_gateway_unavailable_message, selection_error)
from tools.tts_tool_delivery import _origin, _section
from tools.tts_tool_providers import _tts_response_format_from_path
logger = logging.getLogger("tools.tts_tool")
def _origin():
"""``tools.tts_tool``, resolved per call so monkeypatched seams there still apply."""
from tools import tts_tool
return tts_tool
DEFAULT_OPENAI_MODEL = "gpt-4o-mini-tts"
# The managed OpenAI audio gateway only proxies these; anything else is 400 "Unsupported".
MANAGED_OPENAI_TTS_MODELS = frozenset({"gpt-4o-mini-tts"})
@@ -45,40 +37,29 @@ def _managed_openai_audio_route() -> Optional[tuple]:
return gateway.nous_user_token, urljoin(f"{gateway.gateway_origin.rstrip('/')}/", "v1"), True
def _openai_section(tts_config: Any, key: str) -> Dict[str, Any]:
"""``tts.<key>`` as a dict (``tts.openai: null`` in YAML yields None — coalesce so .get() is safe)."""
section = tts_config.get(key) if isinstance(tts_config, dict) else None
return section if isinstance(section, dict) else {}
def _resolve_openai_audio_client_config() -> tuple[str, str, bool]:
"""``(api_key, base_url, is_managed)`` for the OpenAI audio client (``is_managed`` = the restricted
Nous proxy, so callers coerce the request). Strict on the stored ``tts`` selection: ``"nous"``
→ managed ONLY (error if unavailable); any other → direct credentials ONLY (``tts.openai.api_key``
then ``VOICE_TOOLS_OPENAI_KEY``/``OPENAI_API_KEY``); unset → config key → env key → managed."""
origin = _origin()
openai_cfg = _openai_section(origin._load_tts_config(), "openai")
direct_base = openai_cfg.get("base_url") or DEFAULT_OPENAI_BASE_URL
openai_cfg = _section(origin._load_tts_config(), "openai")
selected = origin.read_selection("tts")
if selected == NOUS_MANAGED_PROVIDER:
route = _managed_openai_audio_route()
if route is None:
raise ValueError(selection_error(
"tts", NOUS_MANAGED_PROVIDER,
"the Nous Tool Gateway is not available (not entitled or unreachable)",
))
"the Nous Tool Gateway is not available (not entitled or unreachable)"))
return route
direct_api_key = openai_cfg.get("api_key") or origin.resolve_openai_audio_api_key()
if direct_api_key:
return direct_api_key, direct_base, False
return direct_api_key, openai_cfg.get("base_url") or DEFAULT_OPENAI_BASE_URL, False
if selected is not None:
raise ValueError(selection_error(
"tts", selected,
"neither tts.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set",
))
route = _managed_openai_audio_route()
if route is None:
message = "Neither tts.openai.api_key in config nor VOICE_TOOLS_OPENAI_KEY/OPENAI_API_KEY is set"
@@ -98,16 +79,9 @@ def _has_openai_audio_backend() -> bool:
def _generate_openai_tts(
text: str,
output_path: str,
tts_config: Dict[str, Any],
*,
api_key: Optional[str] = None,
base_url: Optional[str] = None,
model: Optional[str] = None,
voice: Optional[str] = None,
speed: Optional[float] = None,
instructions: Optional[str] = None) -> str:
text: str, output_path: str, tts_config: Dict[str, Any], *, api_key: Optional[str] = None,
base_url: Optional[str] = None, model: Optional[str] = None, voice: Optional[str] = None,
speed: Optional[float] = None, instructions: Optional[str] = None) -> str:
"""Generate audio via the OpenAI ``audio.speech.create`` SDK shape.
Explicit kwargs let OpenAI-compatible backends (DeepInfra) supply credentials/model/voice
@@ -119,21 +93,17 @@ def _generate_openai_tts(
explicit_base_url = base_url is not None
if api_key is None:
api_key, fallback_base, is_managed = _origin()._resolve_openai_audio_client_config()
oai_config = _openai_section(tts_config, "openai")
oai_config = _section(tts_config, "openai")
if model is None:
model = oai_config.get("model", DEFAULT_OPENAI_MODEL)
if voice is None:
voice = oai_config.get("voice", DEFAULT_OPENAI_VOICE)
config_base_url = oai_config.get("base_url")
if base_url is None:
# Config override beats the auth-chain fallback; an explicit arg (DeepInfra) always wins.
if base_url is None: # config override beats the auth-chain fallback; explicit arg wins
base_url = config_base_url or fallback_base or DEFAULT_OPENAI_BASE_URL
if speed is None:
speed_default = tts_config.get("speed", 1.0) if isinstance(tts_config, dict) else 1.0
speed = float(oai_config.get("speed", speed_default))
language = oai_config.get("language")
# The managed gateway only proxies MANAGED_OPENAI_TTS_MODELS; coerce a direct-OpenAI
# model (e.g. "tts-1-hd") unless the user redirected base_url to their own endpoint.
if is_managed and not explicit_base_url and not config_base_url and model not in MANAGED_OPENAI_TTS_MODELS:
@@ -141,29 +111,21 @@ def _generate_openai_tts(
"TTS: managed OpenAI audio gateway does not support model %r; "
"falling back to %s. Set VOICE_TOOLS_OPENAI_KEY or OPENAI_API_KEY "
"to use %r directly.",
model, DEFAULT_OPENAI_MODEL, model,
)
model, DEFAULT_OPENAI_MODEL, model)
model = DEFAULT_OPENAI_MODEL
OpenAIClient = _origin()._import_openai_client()
client = OpenAIClient(api_key=api_key, base_url=base_url)
create_kwargs: Dict[str, Any] = {
"model": model, "voice": voice, "input": text,
"response_format": _tts_response_format_from_path(output_path),
"extra_headers": {"x-idempotency-key": str(uuid.uuid4())}}
if speed != 1.0:
create_kwargs["speed"] = max(0.25, min(4.0, speed))
if instructions:
create_kwargs["instructions"] = instructions
if oai_config.get("language"):
create_kwargs["extra_body"] = {"lang_code": oai_config["language"]}
client = _origin()._import_openai_client()(api_key=api_key, base_url=base_url)
try:
create_kwargs: Dict[str, Any] = {
"model": model,
"voice": voice,
"input": text,
"response_format": _tts_response_format_from_path(output_path),
"extra_headers": {"x-idempotency-key": str(uuid.uuid4())},
}
if speed != 1.0:
create_kwargs["speed"] = max(0.25, min(4.0, speed))
if instructions:
create_kwargs["instructions"] = instructions
if language:
create_kwargs["extra_body"] = {"lang_code": language}
response = client.audio.speech.create(**create_kwargs)
response.stream_to_file(output_path)
client.audio.speech.create(**create_kwargs).stream_to_file(output_path)
return output_path
finally:
close = getattr(client, "close", None)
@@ -177,10 +139,8 @@ def _generate_deepinfra_tts(text: str, output_path: str, tts_config: Dict[str, A
api_key = _origin()._resolve_provider_key("DEEPINFRA_API_KEY", "deepinfra")
if not api_key:
raise ValueError("DEEPINFRA_API_KEY not set. Run `hermes setup` to configure, or set the env var directly.")
di_config = _openai_section(tts_config, "deepinfra")
di_config = _section(tts_config, "deepinfra")
from hermes_cli.models import deepinfra_base_url, deepinfra_model_ids
model = di_config.get("model")
if not isinstance(model, str) or not model.strip():
candidates = deepinfra_model_ids("tts")
@@ -188,16 +148,9 @@ def _generate_deepinfra_tts(text: str, output_path: str, tts_config: Dict[str, A
raise ValueError(
"No DeepInfra TTS model available. Pin one in config.yaml "
"under tts.deepinfra.model, or check connectivity to "
"api.deepinfra.com so the live catalog can be fetched."
)
"api.deepinfra.com so the live catalog can be fetched.")
model = candidates[0]
return _origin()._generate_openai_tts(
text,
output_path,
tts_config,
api_key=api_key,
base_url=deepinfra_base_url(di_config),
model=model,
voice=di_config.get("voice", DEFAULT_DEEPINFRA_TTS_VOICE),
speed=float(di_config.get("speed", tts_config.get("speed", 1.0))),
)
text, output_path, tts_config, api_key=api_key, base_url=deepinfra_base_url(di_config),
model=model, voice=di_config.get("voice", DEFAULT_DEEPINFRA_TTS_VOICE),
speed=float(di_config.get("speed", tts_config.get("speed", 1.0))))
+14 -33
View File
@@ -1,9 +1,7 @@
"""Plugin-registered TTS providers for ``tools.tts_tool``.
Routes ``tts.provider: <name>`` values that are neither built-in nor a ``type: command``
entry to a :class:`agent.tts_provider.TTSProvider` registered by a plugin. Discovery goes
through ``hermes_cli.plugins._ensure_plugins_discovered`` (imported lazily so the tool
module stays importable without the plugin machinery).
"""Plugin-registered TTS providers for ``tools.tts_tool``: routes ``tts.provider: <name>`` values
that are neither built-in nor ``type: command`` to a plugin :class:`agent.tts_provider.TTSProvider`.
Discovery goes through ``hermes_cli.plugins._ensure_plugins_discovered`` (imported lazily so the
tool module stays importable without the plugin machinery).
"""
from __future__ import annotations
@@ -12,11 +10,8 @@ import logging
from typing import Any, Dict, Optional
from tools.tts_command_provider import (
BUILTIN_TTS_PROVIDERS,
DEFAULT_COMMAND_TTS_OUTPUT_FORMAT,
_get_named_provider_config,
_is_command_provider_config,
)
BUILTIN_TTS_PROVIDERS, DEFAULT_COMMAND_TTS_OUTPUT_FORMAT, _get_named_provider_config,
_is_command_provider_config)
logger = logging.getLogger("tools.tts_tool")
@@ -26,10 +21,8 @@ def _lookup_plugin_provider(key: str, *, discover: bool = True, retry: bool = Fa
``retry`` re-discovers with ``force=True`` on a miss (a long-lived session may predate the
plugin's install). Raises on registry/discovery failure — callers decide if fatal."""
from agent.tts_registry import get_provider
if discover:
from hermes_cli.plugins import _ensure_plugins_discovered
_ensure_plugins_discovered()
plugin_provider = get_provider(key)
if plugin_provider is None and retry:
@@ -44,10 +37,8 @@ def _dispatch_to_plugin_provider(text: str, output_path: str, provider: str, tts
Invariants re-checked here so a caller refactor can't break them: built-in names never reach
the registry; a same-named ``type: command`` provider wins; only an exact registered name
dispatches. Plugin exceptions propagate to ``text_to_speech_tool``'s error envelope."""
if not provider:
return None
key = provider.lower().strip()
if key in BUILTIN_TTS_PROVIDERS:
key = (provider or "").lower().strip()
if not key or key in BUILTIN_TTS_PROVIDERS:
return None
if _is_command_provider_config(_get_named_provider_config(tts_config, key)):
return None
@@ -58,33 +49,23 @@ def _dispatch_to_plugin_provider(text: str, output_path: str, provider: str, tts
return None
if plugin_provider is None:
return None
# voice/model/speed/format are optional per the TTSProvider.synthesize contract;
# providers fall back to their own defaults on None.
# voice/model/speed/format are optional per TTSProvider.synthesize; providers default on None.
cfg = tts_config if isinstance(tts_config, dict) else {}
voice = cfg.get("voice")
model = cfg.get("model")
speed = cfg.get("speed")
voice, model, speed = cfg.get("voice"), cfg.get("model"), cfg.get("speed")
fmt = cfg.get("output_format", DEFAULT_COMMAND_TTS_OUTPUT_FORMAT)
logger.info("Generating speech with plugin TTS provider '%s'...", key)
written = plugin_provider.synthesize(
text, output_path,
voice=voice if isinstance(voice, str) and voice else None,
text, output_path, voice=voice if isinstance(voice, str) and voice else None,
model=model if isinstance(model, str) and model else None,
speed=float(speed) if isinstance(speed, (int, float)) else None,
format=str(fmt).lower() if fmt else "mp3",
)
# Contract: returns the (possibly rewritten) output path; tolerate None.
format=str(fmt).lower() if fmt else "mp3")
return written if isinstance(written, str) and written else output_path
def _plugin_provider_is_voice_compatible(provider: str) -> bool:
"""True when the registered plugin provider opts into voice-bubble delivery (any failure -> False)."""
if not provider:
return False
key = provider.lower().strip()
if key in BUILTIN_TTS_PROVIDERS:
key = (provider or "").lower().strip()
if not key or key in BUILTIN_TTS_PROVIDERS:
return False
try:
plugin_provider = _lookup_plugin_provider(key, discover=False)
+110 -232
View File
@@ -1,16 +1,16 @@
"""Cloud TTS backends for ``tools.tts_tool``: Edge, ElevenLabs, xAI, MiniMax, Mistral, Gemini.
Each ``_generate_<provider>(text, output_path, tts_config) -> path`` writes one
final-encoded file. Shared here: bounded upstream response reading (16 MiB cap so
a hostile endpoint can't feed unbounded audio) and the auxiliary-model speech-tag
rewrites. OpenAI/DeepInfra live in ``tts_tool_openai`` (managed-gateway routing).
Seams tests monkeypatch on the origin (``get_env_value``, ``_resolve_provider_key``,
``_import_*``) are resolved through :func:`_origin` at call time.
Each ``_generate_<provider>(text, output_path, tts_config) -> path`` writes one final-encoded
file. Shared here: bounded upstream response reading (16 MiB cap so a hostile endpoint can't
feed unbounded audio) and the auxiliary-model speech-tag rewrites. OpenAI/DeepInfra live in
``tts_tool_openai``. Origin seams (``get_env_value``, ``_resolve_provider_key``, ``_import_*``)
are resolved through :func:`_origin` at call time.
"""
from __future__ import annotations
import base64
import contextlib
import json
import logging
import os
@@ -20,19 +20,11 @@ from pathlib import Path
from typing import Any, Dict, Optional
from urllib.parse import urlparse
from tools.tts_tool_delivery import _wrap_pcm_as_wav, _write_wav_bytes_as
from tools.tts_tool_delivery import _origin, _section, _wrap_pcm_as_wav, _write_wav_bytes_as
from tools.xai_http import hermes_xai_user_agent
logger = logging.getLogger("tools.tts_tool")
def _origin():
"""``tools.tts_tool``, resolved per call so monkeypatched seams there still apply."""
from tools import tts_tool
return tts_tool
DEFAULT_EDGE_VOICE = "en-US-AriaNeural"
DEFAULT_ELEVENLABS_VOICE_ID = "pNInz6obpgDQGcFmaJgB" # Adam
DEFAULT_ELEVENLABS_MODEL_ID = "eleven_multilingual_v2"
@@ -71,33 +63,16 @@ _FALSE_WORDS = {"0", "false", "no", "off", "disabled"}
def _config_bool(value: Any, default: bool = False) -> bool:
"""Coerce common YAML/env bool spellings without treating random strings as true."""
if isinstance(value, bool):
return value
if value is None:
return default
if isinstance(value, (int, float)):
if isinstance(value, (bool, int, float)):
return bool(value)
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in _TRUE_WORDS:
return True
if normalized in _FALSE_WORDS:
return False
return default
normalized = value.strip().lower() if isinstance(value, str) else None
return normalized in _TRUE_WORDS if normalized in _TRUE_WORDS | _FALSE_WORDS else default
def _tts_response_format_from_path(output_path: str) -> str:
"""Pick an OpenAI-style response format (opus/wav/flac/mp3) from the output extension."""
for ext, fmt in ((".ogg", "opus"), (".wav", "wav"), (".flac", "flac")):
if output_path.endswith(ext):
return fmt
return "mp3"
def _section(tts_config: Dict[str, Any], key: str) -> Dict[str, Any]:
"""``tts.<key>`` as a dict (``null``/non-dict sections read as empty)."""
section = tts_config.get(key) if isinstance(tts_config, dict) else None
return section if isinstance(section, dict) else {}
formats = ((".ogg", "opus"), (".wav", "wav"), (".flac", "flac"))
return next((fmt for ext, fmt in formats if output_path.endswith(ext)), "mp3")
def _require_key(env_var: str, provider_id: str, hint: str) -> str:
@@ -109,7 +84,6 @@ def _require_key(env_var: str, provider_id: str, hint: str) -> str:
# --- Bounded upstream response reading ---
def _response_has_explicit_stream(response: Any) -> bool:
"""True for real ``requests`` responses (or doubles defining ``iter_content`` themselves)."""
if not callable(getattr(response, "iter_content", None)):
@@ -121,10 +95,8 @@ def _response_has_explicit_stream(response: Any) -> bool:
def _close_response(response: Any) -> None:
close = getattr(response, "close", None)
if callable(close):
try:
with contextlib.suppress(Exception):
close()
except Exception:
pass
def _read_tts_response_bytes(response: Any, *, label: str, limit: Optional[int] = None) -> bytes:
@@ -137,10 +109,7 @@ def _read_tts_response_bytes(response: Any, *, label: str, limit: Optional[int]
iterator = response.iter_content(chunk_size=TTS_RESPONSE_BODY_CHUNK_BYTES)
else:
content = vars(response).get("content", getattr(type(response), "content", b""))
if isinstance(content, str):
content = content.encode("utf-8", errors="replace")
iterator = (content,) if isinstance(content, (bytes, bytearray)) else ()
iterator = (content,) if isinstance(content, (str, bytes, bytearray)) else ()
for chunk in iterator:
if not chunk:
continue
@@ -157,12 +126,11 @@ def _read_tts_response_bytes(response: Any, *, label: str, limit: Optional[int]
_close_response(response)
def _read_tts_response_json(response: Any, *, label: str, limit: Optional[int] = None) -> Dict[str, Any]:
raw = _read_tts_response_bytes(response, label=label, limit=limit)
def _parse_json_body(response: Any, raw: bytes) -> Dict[str, Any]:
"""JSON from the already-read *raw* body. Unit-test doubles often only provide ``.json()``;
real ``requests`` responses took the streaming path, so production never buffers eagerly."""
if raw:
return json.loads(raw.decode("utf-8"))
# Unit-test doubles often only provide `.json()`; real requests.Response
# objects took the streaming path above, so production never buffers eagerly.
if not _response_has_explicit_stream(response):
json_reader = getattr(response, "json", None)
if callable(json_reader):
@@ -171,38 +139,31 @@ def _read_tts_response_json(response: Any, *, label: str, limit: Optional[int] =
return {}
def _read_tts_response_json(response: Any, *, label: str, limit: Optional[int] = None) -> Dict[str, Any]:
return _parse_json_body(response, _read_tts_response_bytes(response, label=label, limit=limit))
def _write_bytes(output_path: str, audio_bytes: bytes) -> str:
with open(output_path, "wb") as f:
f.write(audio_bytes)
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
return requests.post(url, headers=headers, json=payload, timeout=60, stream=True, **extra)
# --- Auxiliary-model speech-tag rewrites ---
def _extract_auxiliary_message_content(response: Any) -> str:
def _auxiliary_reply_text(response: Any) -> str:
"""The first choice's message content with any ```fence``` unwrapped ("" when unreadable)."""
try:
message = getattr(response.choices[0], "message", None)
if isinstance(message, dict):
return str(message.get("content") or "")
return str(getattr(message, "content", "") or "")
content = message.get("content") if isinstance(message, dict) else getattr(message, "content", "")
except Exception:
return ""
def _strip_code_fence(content: str) -> str:
"""Unwrap a ```fenced``` LLM reply; returns the stripped inner text."""
clean = (content or "").strip()
clean = str(content or "").strip()
fence = re.fullmatch(r"```(?:[A-Za-z0-9_-]+)?\s*(.*?)\s*```", clean, flags=re.DOTALL)
return fence.group(1).strip() if fence else clean
@@ -221,44 +182,37 @@ def _rewrite_with_auxiliary_model(
"""Ask the auxiliary model (task ``tts_audio_tags``) to rewrite a script; *fallback* on any failure/empty reply."""
try:
from agent.auxiliary_client import call_llm
response = call_llm(
task=GEMINI_AUDIO_TAG_REWRITE_TASK,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
temperature=0.7,
)
return _strip_code_fence(_extract_auxiliary_message_content(response)) or fallback
task=GEMINI_AUDIO_TAG_REWRITE_TASK, temperature=0.7,
messages=[{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt}])
return _auxiliary_reply_text(response) or fallback
except Exception as exc:
logger.log(level, "%s audio tag rewrite failed; using %s: %s", label, fallback_label, exc)
return fallback
# --- Edge TTS (free default) ---
async def _generate_edge_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
_edge_tts = _origin()._import_edge_tts()
edge_tts = _origin()._import_edge_tts()
edge_config = tts_config.get("edge") or {}
speed = float(edge_config.get("speed", tts_config.get("speed", 1.0)))
kwargs = {"voice": edge_config.get("voice", DEFAULT_EDGE_VOICE)}
if speed != 1.0:
kwargs["rate"] = f"{round((speed - 1.0) * 100):+d}%"
await _edge_tts.Communicate(text, **kwargs).save(output_path)
await edge_tts.Communicate(text, **kwargs).save(output_path)
return output_path
# --- ElevenLabs ---
def _elevenlabs_environment_kwargs(el_config: Dict[str, Any]) -> Dict[str, Any]:
"""SDK client kwargs for ``tts.elevenlabs.base_url``/``wss_url``; empty (SDK default environment)
without a base_url. ``wss_url`` defaults to the base_url host with a ``ws(s)://`` scheme."""
"""SDK client kwargs for ``tts.elevenlabs.base_url``/``wss_url``; empty (SDK default) without a
base_url. ``wss_url`` defaults to the base_url host with a ``ws(s)://`` scheme."""
base_url = (el_config.get("base_url") or "").rstrip("/")
if not base_url:
return {}
wss_url = (el_config.get("wss_url") or "").rstrip("/") or re.sub(r"^http", "ws", base_url)
from elevenlabs.environment import ElevenLabsEnvironment
wss_url = (el_config.get("wss_url") or "").rstrip("/") or re.sub(r"^http", "ws", base_url)
return {"environment": ElevenLabsEnvironment(base=base_url, wss=wss_url)}
@@ -267,30 +221,24 @@ def _generate_elevenlabs(text: str, output_path: str, tts_config: Dict[str, Any]
el_config = tts_config.get("elevenlabs") or {}
client = _origin()._import_elevenlabs()(api_key=api_key, **_elevenlabs_environment_kwargs(el_config))
audio_generator = client.text_to_speech.convert(
text=text,
voice_id=el_config.get("voice_id", DEFAULT_ELEVENLABS_VOICE_ID),
text=text, voice_id=el_config.get("voice_id", DEFAULT_ELEVENLABS_VOICE_ID),
model_id=el_config.get("model_id", DEFAULT_ELEVENLABS_MODEL_ID),
output_format="opus_48000_64" if output_path.endswith(".ogg") else "mp3_44100_128",
)
output_format="opus_48000_64" if output_path.endswith(".ogg") else "mp3_44100_128")
with open(output_path, "wb") as f:
for chunk in audio_generator:
f.write(chunk)
f.writelines(audio_generator)
return output_path
# --- xAI TTS (dedicated /v1/tts endpoint, not the OpenAI audio shape) ---
_XAI_INLINE_SPEECH_TAGS = (
"pause", "long-pause", "hum-tune", "laugh", "chuckle", "giggle", "cry", "tsk",
"tongue-click", "lip-smack", "breath", "inhale", "exhale", "sigh",
)
"tongue-click", "lip-smack", "breath", "inhale", "exhale", "sigh")
_XAI_WRAPPING_SPEECH_TAGS = (
"soft", "whisper", "loud", "build-intensity", "decrease-intensity", "higher-pitch",
"lower-pitch", "slow", "fast", "sing-song", "singing", "laugh-speak", "emphasis",
)
"lower-pitch", "slow", "fast", "sing-song", "singing", "laugh-speak", "emphasis")
_XAI_SPEECH_TAG_RE = re.compile(
r"(\[(?:" + "|".join(_XAI_INLINE_SPEECH_TAGS) + r")\]|</?(?:" + "|".join(_XAI_WRAPPING_SPEECH_TAGS) + r")>)",
flags=re.IGNORECASE,
)
rf"(\[(?:{'|'.join(_XAI_INLINE_SPEECH_TAGS)})\]|</?(?:{'|'.join(_XAI_WRAPPING_SPEECH_TAGS)})>)",
flags=re.IGNORECASE)
_XAI_FIRST_SENTENCE_RE = re.compile(r"^(.{12,120}?[.!?…])\s+(?=\S)", flags=re.DOTALL)
@@ -301,17 +249,12 @@ def _apply_xai_auto_speech_tags(text: str) -> str:
clean = text.strip()
if not clean:
return text
local = re.sub(r"\n\s*\n+", " [pause] ", clean)
local = re.sub(r"\s*\n\s*", " ", local)
local = re.sub(r"\s*\n\s*", " ", re.sub(r"\n\s*\n+", " [pause] ", clean))
if not _XAI_SPEECH_TAG_RE.search(local):
local = _XAI_FIRST_SENTENCE_RE.sub(r"\1 [pause] ", local, count=1)
local = re.sub(r"\s{2,}", " ", local).strip()
# Explicit user/model tags are trusted as-is.
if _XAI_SPEECH_TAG_RE.search(clean):
if _XAI_SPEECH_TAG_RE.search(clean): # explicit user/model tags are trusted as-is
return local
system_prompt = (
"You rewrite transcripts for the xAI /v1/tts endpoint by inserting "
"expressive speech tags.\n\n"
@@ -322,8 +265,7 @@ def _apply_xai_auto_speech_tags(text: str) -> str:
"- Use wrapping `[tag]...[/tag]` for sustained effects (whisper, soft, slow, fast, loud, etc.).\n"
"- Do not use angle-bracket tags like `<tag>...</tag>` — xAI uses BBCode-style closing tags with `[/tag]`.\n"
"- Do not use SSML.\n"
+ _TAG_REWRITE_TAIL
)
+ _TAG_REWRITE_TAIL)
return _rewrite_with_auxiliary_model(
system_prompt, f"TRANSCRIPT TO TAG:\n{local}", local, label="xAI TTS", fallback_label="locally-tagged text", level=logging.DEBUG,
)
@@ -351,41 +293,33 @@ def _generate_xai_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -
api_key = str(creds.get("api_key") or "").strip()
if not api_key:
raise ValueError("No xAI credentials found. Configure xAI OAuth in `hermes model` or set XAI_API_KEY.")
xai_config = tts_config.get("xai") or {}
voice_id = str(xai_config.get("voice_id", DEFAULT_XAI_VOICE_ID)).strip() or DEFAULT_XAI_VOICE_ID
language = str(xai_config.get("language", DEFAULT_XAI_LANGUAGE)).strip() or DEFAULT_XAI_LANGUAGE
sample_rate = int(xai_config.get("sample_rate", DEFAULT_XAI_SAMPLE_RATE))
bit_rate = int(xai_config.get("bit_rate", DEFAULT_XAI_BIT_RATE))
auto_speech_tags = _config_bool(
xai_config.get("auto_speech_tags", xai_config.get("speech_tags")), DEFAULT_XAI_AUTO_SPEECH_TAGS,
)
# ``tts.xai.speed`` overrides global ``tts.speed``; out-of-range values are
# clamped into the API's band rather than 400ing the request.
speed = _clamped_number(
xai_config.get("speed", tts_config.get("speed")), float, DEFAULT_XAI_SPEED_MIN, DEFAULT_XAI_SPEED_MAX,
)
optimize_streaming_latency = _clamped_number(
xai_config.get("optimize_streaming_latency", tts_config.get("optimize_streaming_latency")), int, 0, 2,
)
text_normalization = _config_bool(xai_config.get("text_normalization"), DEFAULT_XAI_TEXT_NORMALIZATION_DEFAULT)
if auto_speech_tags:
sample_rate, bit_rate = (int(xai_config.get("sample_rate", DEFAULT_XAI_SAMPLE_RATE)),
int(xai_config.get("bit_rate", DEFAULT_XAI_BIT_RATE)))
auto_speech_tags = xai_config.get("auto_speech_tags", xai_config.get("speech_tags"))
if _config_bool(auto_speech_tags, DEFAULT_XAI_AUTO_SPEECH_TAGS):
text = _apply_xai_auto_speech_tags(text)
# ``tts.xai.<knob>`` overrides global ``tts.<knob>``; out-of-range values are clamped into the
# API's band rather than 400ing the request.
speed = _clamped_number(xai_config.get("speed", tts_config.get("speed")), float,
DEFAULT_XAI_SPEED_MIN, DEFAULT_XAI_SPEED_MAX)
optimize_streaming_latency = _clamped_number(
xai_config.get("optimize_streaming_latency", tts_config.get("optimize_streaming_latency")),
int, 0, 2)
text_normalization = _config_bool(
xai_config.get("text_normalization"), DEFAULT_XAI_TEXT_NORMALIZATION_DEFAULT)
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.
# Documented minimal POST /v1/tts shape; optional fields only when they differ from 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
@@ -394,23 +328,18 @@ 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
response = _post_json(f"{base_url}/tts", payload, {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
"User-Agent": hermes_xai_user_agent(),
})
"Authorization": f"Bearer {api_key}", "Content-Type": "application/json",
"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 ---
@dataclass(frozen=True)
class _MiniMaxTTSRuntime:
"""A region-bound MiniMax endpoint and credential (key excluded from ``repr``)."""
@@ -424,8 +353,7 @@ class _MiniMaxTTSRuntime:
_MINIMAX_ENDPOINTS = {"global": DEFAULT_MINIMAX_BASE_URL, "cn": DEFAULT_MINIMAX_CN_BASE_URL}
_MINIMAX_OFFICIAL_HOSTS = {
"global": frozenset({"api.minimax.io", "api.minimax.chat"}),
"cn": frozenset({"api.minimaxi.com"}),
}
"cn": frozenset({"api.minimaxi.com"})}
def _resolve_minimax_tts_runtime(tts_config: Dict[str, Any]) -> _MiniMaxTTSRuntime:
@@ -434,27 +362,21 @@ def _resolve_minimax_tts_runtime(tts_config: Dict[str, Any]) -> _MiniMaxTTSRunti
mm_config = _section(tts_config, "minimax")
resolve_key = _origin()._resolve_provider_key
credentials = {
"global": ("MINIMAX_API_KEY", str(resolve_key("MINIMAX_API_KEY", "minimax") or "").strip()),
"cn": ("MINIMAX_CN_API_KEY", str(resolve_key("MINIMAX_CN_API_KEY", "minimax") or "").strip()),
}
region: (env_var, str(resolve_key(env_var, "minimax") or "").strip())
for region, env_var in (("global", "MINIMAX_API_KEY"), ("cn", "MINIMAX_CN_API_KEY"))}
region = str(mm_config.get("region") or "").strip().lower()
if region and region not in _MINIMAX_ENDPOINTS:
raise ValueError("tts.minimax.region must be 'global' or 'cn'")
if not region:
region = "cn" if credentials["cn"][1] and not credentials["global"][1] else "global"
credential_source, api_key = credentials[region]
if not api_key:
raise ValueError(f"{credential_source} not set for MiniMax TTS region {region!r}")
endpoint = str(mm_config.get("base_url") or _MINIMAX_ENDPOINTS[region]).strip()
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)
@@ -462,8 +384,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:
@@ -474,42 +396,29 @@ def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any
model = mm_config.get("model", DEFAULT_MINIMAX_MODEL)
voice_id = mm_config.get("voice_id", DEFAULT_MINIMAX_VOICE_ID)
base_url = runtime.endpoint
# MiniMax scopes TTS requests by GroupId (``?GroupId=<id>`` on the t2a_v2
# URL): config or MINIMAX_GROUP_ID, attached only when absent from the URL.
group_id = (
str(mm_config.get("group_id") or "").strip()
or (_origin().get_env_value("MINIMAX_GROUP_ID") or "").strip()
)
# MiniMax scopes TTS requests by GroupId (``?GroupId=<id>`` on the t2a_v2 URL): config or
# MINIMAX_GROUP_ID, attached only when absent from the URL.
group_id = (str(mm_config.get("group_id") or "").strip()
or (_origin().get_env_value("MINIMAX_GROUP_ID") or "").strip())
if group_id and "GroupId=" not in base_url:
base_url = f"{base_url}{'&' if '?' in base_url else '?'}GroupId={group_id}"
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:
payload = {"model": model, "text": text, "voice_id": voice_id}
response = _post_json(base_url, payload, {
"Content-Type": "application/json", "Authorization": f"Bearer {runtime.api_key}"
})
"Content-Type": "application/json", "Authorization": f"Bearer {runtime.api_key}"})
if is_t2a_v2:
response.raise_for_status()
result = _read_tts_response_json(response, label="MiniMax TTS")
@@ -518,12 +427,9 @@ def _generate_minimax_tts(text: str, output_path: str, tts_config: Dict[str, Any
if not hex_audio:
raise RuntimeError("MiniMax TTS returned empty audio data")
return _write_bytes(output_path, bytes.fromhex(hex_audio))
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""
try:
@@ -532,29 +438,24 @@ 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")
# --- Mistral (Voxtral TTS) — base64 audio, native Opus for voice bubbles ---
def _generate_mistral_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
api_key = _require_key("MISTRAL_API_KEY", "mistral", "Get one at https://console.mistral.ai/")
mi_config = tts_config.get("mistral") or {}
client_kwargs: Dict[str, Any] = {"api_key": api_key}
if mi_config.get("base_url"):
client_kwargs["server_url"] = mi_config["base_url"] # the Mistral SDK calls it server_url
Mistral = _origin()._import_mistral_client()
Mistral = _origin()._import_mistral_client() # ImportError must escape the RuntimeError wrap
try:
with Mistral(**client_kwargs) as client:
response = client.audio.speech.complete(
model=mi_config.get("model", DEFAULT_MISTRAL_TTS_MODEL),
input=text,
model=mi_config.get("model", DEFAULT_MISTRAL_TTS_MODEL), input=text,
voice_id=mi_config.get("voice_id") or DEFAULT_MISTRAL_TTS_VOICE_ID,
response_format=_tts_response_format_from_path(output_path),
)
response_format=_tts_response_format_from_path(output_path))
audio_bytes = base64.b64decode(response.audio_data)
except ValueError:
raise
@@ -565,7 +466,6 @@ def _generate_mistral_tts(text: str, output_path: str, tts_config: Dict[str, Any
# --- Google Gemini TTS ---
def _read_gemini_persona_prompt(gemini_config: Dict[str, Any]) -> str:
"""Read ``tts.gemini.persona_prompt_file`` (relative -> under HERMES_HOME), failing soft."""
raw = gemini_config.get("persona_prompt_file")
@@ -595,11 +495,8 @@ def _gemini_audio_tags_enabled(gemini_config: Dict[str, Any], model: str) -> boo
normalized = (model or "").strip().lower().rsplit("/", 1)[-1]
if "gemini-3.1" in normalized and "tts" in normalized:
return True
logger.warning(
"Gemini TTS audio_tags enabled, but model %s is not known to support "
"Gemini audio tags; skipping hidden tag rewrite",
model,
)
logger.warning("Gemini TTS audio_tags enabled, but model %s is not known to support "
"Gemini audio tags; skipping hidden tag rewrite", model)
return False
@@ -620,13 +517,11 @@ def _rewrite_gemini_tts_audio_tags(text: str, persona_prompt: str = "") -> str:
+ _TAG_REWRITE_RULES +
"- Use square brackets for every audio tag.\n"
"- Do not use SSML or XML tags.\n"
+ _TAG_REWRITE_TAIL
)
context = persona_prompt.strip() or "(none)"
user_prompt = f"PERSONA AND DIRECTOR CONTEXT:\n{context}\n\nTRANSCRIPT TO TAG:\n{transcript}"
return _rewrite_with_auxiliary_model(
system_prompt, user_prompt, text, label="Gemini TTS", fallback_label="untagged text", level=logging.WARNING,
)
+ _TAG_REWRITE_TAIL)
user_prompt = (f"PERSONA AND DIRECTOR CONTEXT:\n{persona_prompt.strip() or '(none)'}\n\n"
f"TRANSCRIPT TO TAG:\n{transcript}")
return _rewrite_with_auxiliary_model(system_prompt, user_prompt, text, label="Gemini TTS",
fallback_label="untagged text", level=logging.WARNING)
def _compose_gemini_tts_prompt(text: str, gemini_config: Dict[str, Any], persona_prompt: Optional[str] = None) -> str:
@@ -637,12 +532,10 @@ def _compose_gemini_tts_prompt(text: str, gemini_config: Dict[str, Any], persona
persona_prompt = _read_gemini_persona_prompt(gemini_config)
if not persona_prompt:
return transcript
preamble = (
"Synthesize speech from the TRANSCRIPT only. Treat AUDIO PROFILE, "
"SCENE, DIRECTOR'S NOTES, and SAMPLE CONTEXT as performance direction; "
"do not speak those sections aloud."
)
"do not speak those sections aloud.")
for pattern in (r"\{\{\s*transcript\s*\}\}", r"\{\s*transcript\s*\}"):
compiled = re.compile(pattern, flags=re.IGNORECASE)
if compiled.search(persona_prompt):
@@ -654,48 +547,38 @@ def _gemini_error_detail(response: Any) -> str:
"""Best-effort ``error.message`` from a non-200 Gemini reply, else the first 300 body chars."""
raw_body = _read_tts_response_bytes(response, label="Gemini TTS")
try:
if raw_body:
err = json.loads(raw_body.decode("utf-8")).get("error", {})
elif not _response_has_explicit_stream(response) and callable(getattr(response, "json", None)):
err = response.json().get("error", {})
else:
err = {}
return err.get("message") or raw_body.decode("utf-8", errors="replace")[:300]
message = _parse_json_body(response, raw_body).get("error", {}).get("message")
except Exception:
return raw_body.decode("utf-8", errors="replace")[:300]
message = None
return message or raw_body.decode("utf-8", errors="replace")[:300]
def _generate_gemini_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
"""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")
gemini_config = _section(tts_config, "gemini")
model = str(gemini_config.get("model", DEFAULT_GEMINI_TTS_MODEL)).strip() or DEFAULT_GEMINI_TTS_MODEL
voice = str(gemini_config.get("voice", DEFAULT_GEMINI_TTS_VOICE)).strip() or DEFAULT_GEMINI_TTS_VOICE
base_url = str(
gemini_config.get("base_url") or origin.get_env_value("GEMINI_BASE_URL") or DEFAULT_GEMINI_TTS_BASE_URL
).strip().rstrip("/")
base_url = str(gemini_config.get("base_url") or origin.get_env_value("GEMINI_BASE_URL")
or DEFAULT_GEMINI_TTS_BASE_URL).strip().rstrip("/")
persona_prompt = _read_gemini_persona_prompt(gemini_config)
tts_script = text
if _gemini_audio_tags_enabled(gemini_config, model):
tts_script = _rewrite_gemini_tts_audio_tags(text, persona_prompt=persona_prompt)
prompt_text = _compose_gemini_tts_prompt(tts_script, gemini_config, persona_prompt=persona_prompt)
prompt_text = _compose_gemini_tts_prompt(
tts_script, gemini_config, persona_prompt=persona_prompt)
max_len = origin._resolve_max_text_length("gemini", tts_config)
if len(prompt_text) > max_len:
raise ValueError(
"Gemini TTS composed prompt exceeds the provider request limit "
f"({len(prompt_text)} > {max_len} chars). Reduce the persona/audio-tag "
"prompt or lower tts.gemini.max_text_length so long-form text is "
"split with enough prompt headroom."
)
"split with enough prompt headroom.")
payload: Dict[str, Any] = {
"contents": [{"parts": [{"text": prompt_text}]}],
"generationConfig": {
@@ -706,26 +589,21 @@ def _generate_gemini_tts(text: str, output_path: str, tts_config: Dict[str, Any]
headers = {"Content-Type": "application/json"}
if urlparse(base_url).hostname == "generativelanguage.googleapis.com":
try:
import hermes_cli as _hermes_cli
_hermes_version = str(_hermes_cli.__version__)
import hermes_cli
version = str(hermes_cli.__version__)
except Exception:
_hermes_version = "0.0.0"
# Gemini partner-integration guidance: identify the client.
headers["X-Goog-Api-Client"] = f"hermes-agent/{_hermes_version}"
version = "0.0.0"
headers["X-Goog-Api-Client"] = f"hermes-agent/{version}" # partner-integration guidance
response = _post_json(f"{base_url}/models/{model}:generateContent", payload, headers, params={"key": api_key})
if response.status_code != 200:
raise RuntimeError(f"Gemini TTS API error (HTTP {response.status_code}): {_gemini_error_detail(response)}")
try:
data = _read_tts_response_json(response, label="Gemini TTS")
parts = data["candidates"][0]["content"]["parts"]
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:
+73 -144
View File
@@ -1,16 +1,16 @@
"""Speaker-side streaming pipeline for ``tools.tts_tool.stream_tts_to_speaker``.
Turns a queue of LLM text deltas into audio the moment each sentence is complete.
Two paths share the sentence cutter (``tools.tts_streaming``): :class:`_StreamerPlayback`
for a registered chunked streamer (prefetch thread per sentence, one FIFO playback worker
through a sounddevice OutputStream or temp WAV + system player) and
:class:`_SyncSentencePipeline` for every other provider (per-sentence ``text_to_speech_tool``
on a single-thread executor, overlapped with playback). Seams tests monkeypatch on the origin
module are resolved through :func:`_origin` at call time.
Turns a queue of LLM text deltas into audio the moment each sentence is complete. Two paths
share the sentence cutter (``tools.tts_streaming``): :class:`_StreamerPlayback` for a registered
chunked streamer (prefetch thread per sentence, one FIFO playback worker through a sounddevice
OutputStream or temp WAV + system player) and :class:`_SyncSentencePipeline` for every other
provider (per-sentence ``text_to_speech_tool`` on a single-thread executor, overlapped with
playback). Origin seams are resolved through :func:`_origin` at call time.
"""
from __future__ import annotations
import contextlib
import logging
import os
import platform
@@ -20,23 +20,10 @@ import threading
from concurrent.futures import Future, ThreadPoolExecutor
from typing import Callable, Iterable, Iterator, List, Optional
from tools.tts_tool_delivery import _origin, _remove_quietly as _unlink_quietly
logger = logging.getLogger("tools.tts_tool")
def _origin():
from tools import tts_tool
return tts_tool
def _unlink_quietly(path: Optional[str]) -> None:
if path:
try:
os.unlink(path)
except OSError:
pass
def _align_int16_chunks(chunks: Iterable[bytes], stop_evt: threading.Event, *, pad_tail: bool = True) -> Iterator[bytes]:
"""Yield int16-aligned byte chunks; a dangling odd byte is padded at the end (or dropped)."""
leftover = b""
@@ -47,15 +34,14 @@ def _align_int16_chunks(chunks: Iterable[bytes], stop_evt: threading.Event, *, p
aligned_len = len(buf) - (len(buf) % 2)
if aligned_len >= 2:
yield buf[:aligned_len]
leftover = buf[aligned_len:] if aligned_len < len(buf) else b""
leftover = buf[aligned_len:]
if leftover and pad_tail:
yield b"\x00"
def _play_via_tempfile(audio_iter: Iterable[bytes], stop_evt: threading.Event, sample_rate: int = 24000) -> None:
"""Write PCM chunks to a temp WAV file and play it with the system player."""
tmp = None
tmp_path = None
tmp = tmp_path = None
try:
import wave
tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)
@@ -75,21 +61,14 @@ def _play_via_tempfile(audio_iter: Iterable[bytes], stop_evt: threading.Event, s
logger.warning("Temp-file TTS fallback failed: %s", exc)
finally:
if tmp is not None:
try:
with contextlib.suppress(Exception):
tmp.close() # idempotent; ensures close on early error
except Exception:
pass
_unlink_quietly(tmp_path)
def _drain_chunks(chunk_queue: "queue.Queue[Optional[bytes]]") -> List[bytes]:
"""Collect one sentence's PCM chunks up to the ``None`` sentinel."""
chunks: List[bytes] = []
while True:
chunk = chunk_queue.get()
if chunk is None:
return chunks
chunks.append(chunk)
return list(iter(chunk_queue.get, None))
class _SyncSentencePipeline:
@@ -109,9 +88,8 @@ class _SyncSentencePipeline:
def speak(self, cleaned: str) -> None:
"""Queue one sentence. Blocks only when the lookahead bound is full."""
if self._stop.is_set():
return
self._queue.put((cleaned, self._executor.submit(self._synthesize_to_tmp, cleaned)))
if not self._stop.is_set():
self._queue.put((cleaned, self._executor.submit(self._synthesize_to_tmp, cleaned)))
def close(self) -> None:
"""Flush queued sentences in order (skipped if stopped), then join."""
@@ -134,11 +112,7 @@ class _SyncSentencePipeline:
return None
def _drain(self) -> None:
while True:
item = self._queue.get()
if item is None:
return
_sentence, future = item
for _sentence, future in iter(self._queue.get, None):
tmp_path = None
try:
tmp_path = future.result()
@@ -155,17 +129,15 @@ class _StreamerPlayback:
"""Prefetch + FIFO playback for a chunked :class:`StreamingTTSProvider`.
``speak(text)`` starts ``streamer.stream()`` immediately on a prefetch thread (at most 3 in
flight) buffering into a bounded per-sentence queue; one playback worker drains those in
order, so sentence N+1 arrives while N plays. Output is a PortAudio stream when one opened,
else temp WAV files; a failing write is retried on a reinitialized stream up to
``_MAX_REINIT`` times before falling back to temp files."""
flight) buffering into a bounded per-sentence queue; one playback worker drains those in order,
so sentence N+1 arrives while N plays. Output is a PortAudio stream when one opened, else temp
WAV files; a failing write is retried on a reinitialized stream up to ``_MAX_REINIT`` times."""
_MAX_REINIT = 3
_CHUNK_QUEUE_MAX = 64
def __init__(self, streamer, stop_event: threading.Event):
self.streamer = streamer
self.stop_event = stop_event
self.streamer, self.stop_event = streamer, stop_event
self.output_stream = self._open_output_stream()
self._audio_queue: "queue.Queue[Optional[queue.Queue[Optional[bytes]]]]" = queue.Queue()
self._prefetch_threads: List[threading.Thread] = []
@@ -173,18 +145,17 @@ class _StreamerPlayback:
self._worker = threading.Thread(target=self._playback_worker, daemon=True)
self._worker.start()
# -- PortAudio stream management ---------------------------------------
def _create_output_stream(self):
sd = _origin()._import_sounddevice()
stream = sd.OutputStream(samplerate=self.streamer.sample_rate, channels=self.streamer.channels, dtype="int16")
stream = sd.OutputStream(
samplerate=self.streamer.sample_rate, channels=self.streamer.channels, dtype="int16")
stream.start()
return stream
def _open_output_stream(self):
# On macOS skip sounddevice entirely: PortAudio/CoreAudio init triggers a
# kTCCServiceMediaLibrary permission prompt even though output needs no
# media-library access. None routes every sentence through tempfile -> afplay.
# macOS skips sounddevice entirely: PortAudio/CoreAudio init triggers a
# kTCCServiceMediaLibrary prompt though output needs no media-library access.
# None routes every sentence through tempfile -> afplay.
if platform.system() == "Darwin":
return None
try:
@@ -195,27 +166,12 @@ class _StreamerPlayback:
logger.warning("sounddevice OutputStream failed: %s", exc)
return None
def _reinit_output_stream(self):
"""Close the broken PortAudio stream and try to create a fresh one."""
self.close_output_stream()
try:
self.output_stream = self._create_output_stream()
logger.info("TTS: PortAudio output stream reinitialized after error")
except Exception as exc:
logger.warning("TTS: PortAudio stream reinit failed: %s", exc)
self.output_stream = None
return self.output_stream
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 ----------------------------------------------------------
def speak(self, text: str) -> None:
"""Start ``streamer.stream(text)`` and prefetch its chunks immediately."""
@@ -227,9 +183,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:
@@ -244,17 +200,12 @@ class _StreamerPlayback:
chunk_queue.put(None) # sentinel: no more chunks
self._prefetch_sem.release()
# -- playback ----------------------------------------------------------
def _play_sentence_via_tempfile(self, chunk_queue) -> None:
_play_via_tempfile(iter(_drain_chunks(chunk_queue)), self.stop_event, self.streamer.sample_rate)
_play_via_tempfile(_drain_chunks(chunk_queue), self.stop_event, self.streamer.sample_rate)
def _for_each_sentence(self, play: Callable[[queue.Queue], None]) -> None:
"""Feed queued sentences to *play* in order until the end sentinel; stopped sentences are skipped."""
while True:
chunk_queue = self._audio_queue.get()
if chunk_queue is None:
return
for chunk_queue in iter(self._audio_queue.get, None):
if not self.stop_event.is_set():
play(chunk_queue)
@@ -262,17 +213,24 @@ class _StreamerPlayback:
self._current_stream.write(self._np.frombuffer(buf, dtype="<i2").reshape(-1, 1))
def _recover_stream(self) -> bool:
"""Reinit the PortAudio stream after a failed write; False once ``_MAX_REINIT`` is exhausted."""
if self._reinit_count < self._MAX_REINIT:
self._reinit_count += 1
self._current_stream = self._reinit_output_stream()
return self._current_stream is not None
logger.warning(
"TTS: PortAudio reinit exhausted after %d attempts, falling back to tempfile for remaining sentences",
self._MAX_REINIT,
)
self._current_stream = None
return False
"""Close the broken PortAudio stream and open a fresh one after a failed write; False once
``_MAX_REINIT`` is exhausted (remaining sentences go through temp files)."""
if self._reinit_count >= self._MAX_REINIT:
logger.warning(
"TTS: PortAudio reinit exhausted after %d attempts, falling back to tempfile for remaining sentences",
self._MAX_REINIT)
self._current_stream = None
return False
self._reinit_count += 1
self.close_output_stream()
try:
self.output_stream = self._create_output_stream()
logger.info("TTS: PortAudio output stream reinitialized after error")
except Exception as exc:
logger.warning("TTS: PortAudio stream reinit failed: %s", exc)
self.output_stream = None
self._current_stream = self.output_stream
return self._current_stream is not None
def _play_sentence_via_stream(self, chunk_queue) -> None:
"""Write one sentence's PCM to PortAudio; after an unrecoverable write failure the rest is dropped."""
@@ -286,28 +244,20 @@ 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."""
if self.output_stream is None:
self._for_each_sentence(self._play_sentence_via_tempfile)
return
import numpy as _np
try:
from tools.voice_mode import mark_audio_output_active
except Exception:
def mark_audio_output_active(_active):
return None
self._np = _np
self._reinit_count = 0
self._current_stream = self.output_stream
mark_audio_output_active = lambda _active: None # noqa: E731
self._np, self._reinit_count, self._current_stream = _np, 0, self.output_stream
mark_audio_output_active(True)
try:
self._for_each_sentence(self._play_sentence_via_stream)
@@ -325,8 +275,7 @@ class _StreamerPlayback:
def stream_tts_to_speaker(
text_queue: queue.Queue, stop_event: threading.Event, tts_done_event: threading.Event,
display_callback: Optional[Callable[[str], None]] = None, provider: Optional[str] = None,
):
display_callback: Optional[Callable[[str], None]] = None, provider: Optional[str] = None):
"""Consume text deltas from *text_queue*, cut into sentences, speak each the moment it's ready.
A registered streaming provider plays chunked PCM; every other provider is spoken
@@ -337,28 +286,21 @@ def stream_tts_to_speaker(
origin = _origin()
sync_pipeline: Optional[_SyncSentencePipeline] = None
playback: Optional[_StreamerPlayback] = None
try:
tts_config = origin._load_tts_config()
# Prefer a chunked streamer for low time-to-first-audio; otherwise per-sentence
# sync synthesis (universal — edge + every non-streamer).
# Prefer a chunked streamer for low time-to-first-audio; otherwise per-sentence sync
# synthesis (universal — edge + every non-streamer).
from tools.tts_streaming import SentenceChunker, resolve_streaming_provider
streamer = resolve_streaming_provider(tts_config, preferred=provider)
stream_max_len = 0
if streamer is None:
sync_pipeline = _SyncSentencePipeline(stop_event)
else:
try:
stream_max_len = origin._resolve_max_text_length(provider or origin._get_provider(tts_config), tts_config)
except Exception:
stream_max_len = 0
with contextlib.suppress(Exception):
stream_max_len = origin._resolve_max_text_length(
provider or origin._get_provider(tts_config), tts_config)
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:
@@ -379,44 +321,31 @@ def stream_tts_to_speaker(
if stream_max_len and len(cleaned) > stream_max_len:
cleaned = cleaned[:stream_max_len]
playback.speak(cleaned)
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:
# Flush the sync pipeline first: queued sentences finish playing (or are skipped
# when stop_event is set) BEFORE tts_done_event fires, so continuous voice mode
# never reopens the mic over its own voice.
# Flush the sync pipeline first: queued sentences finish playing (or are skipped when
# stop_event is set) BEFORE tts_done_event fires, so continuous voice mode never reopens
# the mic over its own voice. The end sentinel lives in finally: so an exception in the
# text pump still lets the playback worker exit.
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:
playback.finish()
tts_done_event.set()