merge simp/r3-34-F pass2 into simp/r3-34
This commit is contained in:
@@ -32,7 +32,6 @@ from tools.tts_tool import (
|
||||
_get_command_tts_output_format,
|
||||
_get_command_tts_timeout,
|
||||
_get_named_provider_config,
|
||||
_has_any_command_tts_provider,
|
||||
_is_command_provider_config,
|
||||
_is_command_tts_voice_compatible,
|
||||
_iter_command_providers,
|
||||
@@ -168,7 +167,7 @@ class TestIsCommandProviderConfig:
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _iter_command_providers / _has_any_command_tts_provider
|
||||
# _iter_command_providers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestIterCommandProviders:
|
||||
@@ -185,11 +184,6 @@ class TestIterCommandProviders:
|
||||
assert names == ["piper-cli", "voxcpm"]
|
||||
|
||||
|
||||
def test_has_any_command_provider_when_none(self):
|
||||
assert _has_any_command_tts_provider({"providers": {}}) is False
|
||||
assert _has_any_command_tts_provider({}) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# config getters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -163,6 +163,5 @@ class TestCheckTtsRequirementsMistral:
|
||||
patch("tools.tts_tool._import_openai_client", side_effect=ImportError), \
|
||||
patch("tools.tts_tool._check_neutts_available", return_value=False), \
|
||||
patch("tools.tts_tool._check_kittentts_available", return_value=False), \
|
||||
patch("tools.tts_tool._check_piper_available", return_value=False), \
|
||||
patch("tools.tts_tool._has_any_command_tts_provider", return_value=False):
|
||||
patch("tools.tts_tool._check_piper_available", return_value=False):
|
||||
assert check_tts_requirements() is False
|
||||
|
||||
@@ -232,7 +232,6 @@ class TestCheckTtsRequirementsPiper:
|
||||
monkeypatch.setattr(tts_tool, "_import_mistral_client", lambda: (_ for _ in ()).throw(ImportError()))
|
||||
monkeypatch.setattr(tts_tool, "_check_neutts_available", lambda: False)
|
||||
monkeypatch.setattr(tts_tool, "_check_kittentts_available", lambda: False)
|
||||
monkeypatch.setattr(tts_tool, "_has_any_command_tts_provider", lambda: False)
|
||||
monkeypatch.setattr(tts_tool, "_has_openai_audio_backend", lambda: False)
|
||||
for env in ("MINIMAX_API_KEY", "XAI_API_KEY", "GEMINI_API_KEY",
|
||||
"GOOGLE_API_KEY", "MISTRAL_API_KEY", "ELEVENLABS_API_KEY"):
|
||||
|
||||
+84
-154
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user