diff --git a/tests/tools/test_tts_command_providers.py b/tests/tools/test_tts_command_providers.py index 631d7d183a..abe540c268 100644 --- a/tests/tools/test_tts_command_providers.py +++ b/tests/tools/test_tts_command_providers.py @@ -32,7 +32,6 @@ from tools.tts_tool import ( _get_command_tts_output_format, _get_command_tts_timeout, _get_named_provider_config, - _has_any_command_tts_provider, _is_command_provider_config, _is_command_tts_voice_compatible, _iter_command_providers, @@ -168,7 +167,7 @@ class TestIsCommandProviderConfig: # --------------------------------------------------------------------------- -# _iter_command_providers / _has_any_command_tts_provider +# _iter_command_providers # --------------------------------------------------------------------------- class TestIterCommandProviders: @@ -185,11 +184,6 @@ class TestIterCommandProviders: assert names == ["piper-cli", "voxcpm"] - def test_has_any_command_provider_when_none(self): - assert _has_any_command_tts_provider({"providers": {}}) is False - assert _has_any_command_tts_provider({}) is False - - # --------------------------------------------------------------------------- # config getters # --------------------------------------------------------------------------- diff --git a/tests/tools/test_tts_mistral.py b/tests/tools/test_tts_mistral.py index 0f8d4432c4..6917546253 100644 --- a/tests/tools/test_tts_mistral.py +++ b/tests/tools/test_tts_mistral.py @@ -163,6 +163,5 @@ class TestCheckTtsRequirementsMistral: patch("tools.tts_tool._import_openai_client", side_effect=ImportError), \ patch("tools.tts_tool._check_neutts_available", return_value=False), \ patch("tools.tts_tool._check_kittentts_available", return_value=False), \ - patch("tools.tts_tool._check_piper_available", return_value=False), \ - patch("tools.tts_tool._has_any_command_tts_provider", return_value=False): + patch("tools.tts_tool._check_piper_available", return_value=False): assert check_tts_requirements() is False diff --git a/tests/tools/test_tts_piper.py b/tests/tools/test_tts_piper.py index 33bb4feb47..52ed2cacff 100644 --- a/tests/tools/test_tts_piper.py +++ b/tests/tools/test_tts_piper.py @@ -232,7 +232,6 @@ class TestCheckTtsRequirementsPiper: monkeypatch.setattr(tts_tool, "_import_mistral_client", lambda: (_ for _ in ()).throw(ImportError())) monkeypatch.setattr(tts_tool, "_check_neutts_available", lambda: False) monkeypatch.setattr(tts_tool, "_check_kittentts_available", lambda: False) - monkeypatch.setattr(tts_tool, "_has_any_command_tts_provider", lambda: False) monkeypatch.setattr(tts_tool, "_has_openai_audio_backend", lambda: False) for env in ("MINIMAX_API_KEY", "XAI_API_KEY", "GEMINI_API_KEY", "GOOGLE_API_KEY", "MISTRAL_API_KEY", "ELEVENLABS_API_KEY"): diff --git a/tools/tts_tool.py b/tools/tts_tool.py index 5485725b6d..d5ab28f54f 100644 --- a/tools/tts_tool.py +++ b/tools/tts_tool.py @@ -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.`` 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.`` -keeps resolving and tests patching ``tools.tts_tool.`` 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.`` +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.`` 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:`` 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="🔊") diff --git a/tools/tts_tool_delivery.py b/tools/tts_tool_delivery.py index 34e3a83486..6836b693ee 100644 --- a/tools/tts_tool_delivery.py +++ b/tools/tts_tool_delivery.py @@ -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..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.`` 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..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 `` 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: diff --git a/tools/tts_tool_lifecycle.py b/tools/tts_tool_lifecycle.py index cb2ca73ac7..7c30bb059c 100644 --- a/tools/tts_tool_lifecycle.py +++ b/tools/tts_tool_lifecycle.py @@ -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 ``_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]: diff --git a/tools/tts_tool_local.py b/tools/tts_tool_local.py index db1de1b4df..96a08d2445 100644 --- a/tools/tts_tool_local.py +++ b/tools/tts_tool_local.py @@ -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: """``/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) diff --git a/tools/tts_tool_openai.py b/tools/tts_tool_openai.py index 8952bffe79..a872798053 100644 --- a/tools/tts_tool_openai.py +++ b/tools/tts_tool_openai.py @@ -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.`` 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)))) diff --git a/tools/tts_tool_plugins.py b/tools/tts_tool_plugins.py index 5e8a690f42..19960ab67e 100644 --- a/tools/tts_tool_plugins.py +++ b/tools/tts_tool_plugins.py @@ -1,9 +1,7 @@ -"""Plugin-registered TTS providers for ``tools.tts_tool``. - -Routes ``tts.provider: `` 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: `` 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) diff --git a/tools/tts_tool_providers.py b/tools/tts_tool_providers.py index 43c925d3d5..2e1dc2a8ec 100644 --- a/tools/tts_tool_providers.py +++ b/tools/tts_tool_providers.py @@ -1,16 +1,16 @@ """Cloud TTS backends for ``tools.tts_tool``: Edge, ElevenLabs, xAI, MiniMax, Mistral, Gemini. -Each ``_generate_(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_(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.`` 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")\]|)", - flags=re.IGNORECASE, -) + rf"(\[(?:{'|'.join(_XAI_INLINE_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 `...` — 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.`` overrides global ``tts.``; 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=`` 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=`` 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: diff --git a/tools/tts_tool_speaker.py b/tools/tts_tool_speaker.py index a1cebd7f2f..f095b46aec 100644 --- a/tools/tts_tool_speaker.py +++ b/tools/tts_tool_speaker.py @@ -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=" 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()