"""Shared runner for user-configured shell ("command") TTS/STT providers. Both ``tools.tts_tool`` and ``tools.transcription_tools`` let users declare a provider as a shell command template with ``{placeholders}``. This module owns the shell-quote-aware template rendering and the idle-timeout process runner they share, plus the TTS side's ``tts.providers.`` config layer. Each origin module re-imports these under its historical private names. TTS config shape:: tts: provider: piper-en providers: piper-en: type: command command: "piper -m ~/model.onnx -f {output_path} < {input_path}" output_format: wav Placeholders: ``{input_path}``, ``{text_path}`` (alias), ``{output_path}``, ``{format}``, ``{voice}``, ``{model}``, ``{speed}``; ``{{``/``}}`` for literal braces. Values are shell-quoted for their surrounding quote context. Built-in provider names always win over a same-named entry under ``tts.providers``. """ from __future__ import annotations import os import queue import re import shlex import subprocess import tempfile import threading import time from pathlib import Path from typing import Any, Dict, Optional def shell_quote_context(command_template: str, position: int) -> Optional[str]: """Return the shell quote char (``'``/``"``) active right before *position*, or None.""" quote: Optional[str] = None escaped = False i = 0 while i < position: char = command_template[i] if quote == "'": if char == "'": quote = None elif quote == '"': if escaped: escaped = False elif char == "\\": escaped = True elif char == '"': quote = None elif char == "'": quote = "'" elif char == '"': quote = '"' elif char == "\\": i += 1 i += 1 return quote def quote_command_placeholder(value: str, quote_context: Optional[str]) -> str: """Quote a placeholder value for its position in a shell command template.""" if quote_context == "'": return value.replace("'", r"'\''") if quote_context == '"': return ( value .replace("\\", "\\\\") .replace('"', r'\"') .replace("$", r"\$") .replace("`", r"\`") ) if os.name == "nt": return subprocess.list2cmdline([value]) return shlex.quote(value) def render_command_template( command_template: str, placeholders: Dict[str, str], ) -> str: """Replace ``{name}`` placeholders (quote-aware) while preserving ``{{``/``}}``.""" names = "|".join(re.escape(name) for name in placeholders) pattern = re.compile( rf"(?{names})\}}\}}|\{{(?P{names})\}})" ) replacements: list[tuple[str, str]] = [] def replace_match(match: re.Match[str]) -> str: name = match.group("double") or match.group("single") token = f"__HERMES_CMD_PLACEHOLDER_{len(replacements)}__" replacements.append(( token, quote_command_placeholder( placeholders[name], shell_quote_context(command_template, match.start()), ), )) return token rendered = pattern.sub(replace_match, command_template) rendered = rendered.replace("{{", "{").replace("}}", "}") for token, value in replacements: rendered = rendered.replace(token, value) return rendered def _signal_process_tree(psutil: Any, proc: subprocess.Popen, method: str) -> None: """Apply ``terminate``/``kill`` to *proc* and all descendants (best effort).""" try: parent = psutil.Process(proc.pid) for child in parent.children(recursive=True): try: getattr(child, method)() except psutil.NoSuchProcess: pass getattr(parent, method)() except psutil.NoSuchProcess: return except Exception: getattr(proc, method)() def terminate_command_process_tree(proc: subprocess.Popen) -> None: """Best-effort termination of a shell process and all of its children.""" if proc.poll() is not None: return if os.name == "nt": try: subprocess.run( ["taskkill", "/F", "/T", "/PID", str(proc.pid)], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=5, stdin=subprocess.DEVNULL, ) except Exception: proc.kill() return try: import psutil # type: ignore except ImportError: proc.terminate() try: proc.wait(timeout=2) except subprocess.TimeoutExpired: proc.kill() return _signal_process_tree(psutil, proc, "terminate") try: proc.wait(timeout=2) return except subprocess.TimeoutExpired: pass _signal_process_tree(psutil, proc, "kill") def command_env_passthrough(config: Dict[str, Any]) -> list: """Return the provider's ``env_passthrough`` allowlist. The child env is scrubbed of Hermes secrets by default; this list names variables copied back from the parent env so a trusted template (e.g. a curl one-liner using its own API key) keeps working. """ raw = config.get("env_passthrough") if not isinstance(raw, (list, tuple)): return [] return [str(item).strip() for item in raw if str(item).strip()] def run_command_provider( command: str, timeout: float, env_passthrough: Optional[list] = None, ) -> subprocess.CompletedProcess: """Run a command-provider shell command with process-tree idle cleanup. ``timeout`` is an IDLE timeout, reset whenever the command emits output — a slow-but-alive provider survives, a silently stalled one is killed. Child env is scrubbed of Hermes secrets while propagating delegated-child lineage markers. """ from agent.delegation_context import delegated_child_subprocess_env from tools.environments.local import hermes_subprocess_env scrubbed = hermes_subprocess_env(inherit_credentials=False) for key in env_passthrough or []: value = os.environ.get(key) if value is not None: scrubbed[key] = value popen_kwargs: Dict[str, Any] = { "shell": True, "stdout": subprocess.PIPE, "stderr": subprocess.PIPE, "text": True, # Lossy UTF-8 decode: locale-mismatched bytes must not raise in the # reader threads on non-UTF-8 Windows. "encoding": "utf-8", "errors": "replace", "env": delegated_child_subprocess_env(scrubbed), } if os.name == "nt": popen_kwargs["creationflags"] = getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0) else: popen_kwargs["start_new_session"] = True proc = subprocess.Popen(command, **popen_kwargs, stdin=subprocess.DEVNULL) output_queue: "queue.Queue[tuple[str, Optional[str]]]" = queue.Queue() chunks: Dict[str, list[str]] = {"stdout": [], "stderr": []} open_streams = {"stdout", "stderr"} def read_stream(name: str, stream: Any) -> None: encoding = getattr(stream, "encoding", None) or "utf-8" read1 = getattr(getattr(stream, "buffer", None), "read1", None) try: while True: if read1 is None: chunk = stream.read(65536) else: chunk = read1(65536).decode(encoding, errors="replace") if not chunk: break output_queue.put((name, chunk)) finally: output_queue.put((name, None)) readers = [ threading.Thread(target=read_stream, args=(name, stream), daemon=True) for name, stream in (("stdout", proc.stdout), ("stderr", proc.stderr)) ] for reader in readers: reader.start() deadline = time.monotonic() + timeout timed_out = False while open_streams: remaining = deadline - time.monotonic() if remaining <= 0: timed_out = True break try: name, chunk = output_queue.get(timeout=min(0.05, remaining)) except queue.Empty: continue if chunk is None: open_streams.discard(name) continue chunks[name].append(chunk) deadline = time.monotonic() + timeout if not timed_out: try: proc.wait(timeout=max(0.0, deadline - time.monotonic())) except subprocess.TimeoutExpired: timed_out = True if timed_out: terminate_command_process_tree(proc) for reader in readers: reader.join(timeout=0.5) while True: try: name, chunk = output_queue.get_nowait() except queue.Empty: break if chunk: chunks[name].append(chunk) stdout = "".join(chunks["stdout"]) stderr = "".join(chunks["stderr"]) try: raise subprocess.TimeoutExpired(command, timeout) except subprocess.TimeoutExpired as exc: raise subprocess.TimeoutExpired( command, timeout, output=stdout, stderr=stderr, ) from exc stdout = "".join(chunks["stdout"]) stderr = "".join(chunks["stderr"]) if proc.returncode: raise subprocess.CalledProcessError( proc.returncode, command, output=stdout, stderr=stderr, ) return subprocess.CompletedProcess(command, proc.returncode, stdout, stderr) # =========================================================================== # TTS ``tts.providers.`` config layer # =========================================================================== # Any ``tts.provider`` value NOT in this set refers to ``tts.providers.``. BUILTIN_TTS_PROVIDERS = frozenset({ "edge", "elevenlabs", "openai", "minimax", "xai", "mistral", "gemini", "neutts", "kittentts", "piper", "deepinfra", }) DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS = 120 DEFAULT_COMMAND_TTS_OUTPUT_FORMAT = "mp3" COMMAND_TTS_OUTPUT_FORMATS = frozenset( {"mp3", "wav", "ogg", "flac", "m4a", "aac", "amr", "opus"} ) DEFAULT_COMMAND_TTS_MAX_TEXT_LENGTH = 5000 def _get_provider_section(tts_config: Dict[str, Any], name: str) -> Dict[str, Any]: """Return a provider config block if it's a dict, else an empty dict.""" if not isinstance(tts_config, dict): return {} section = tts_config.get(name) return section if isinstance(section, dict) else {} def _get_named_provider_config(tts_config: Dict[str, Any], name: str) -> Dict[str, Any]: """Config dict for a user-declared provider, or {}. ``tts.providers.`` is canonical; ``tts.`` is accepted as back-compat only for non-built-in names (so a user's ``tts.openai`` block still means the OpenAI provider, not a custom command). """ section = _get_provider_section(tts_config, "providers").get(name) if isinstance(section, dict): return section if name.lower() not in BUILTIN_TTS_PROVIDERS: return _get_provider_section(tts_config, name) return {} def _is_command_provider_config(config: Dict[str, Any]) -> bool: """True when *config* declares a command-type provider (has a non-empty ``command``).""" if not isinstance(config, dict): return False ptype = str(config.get("type") or "").strip().lower() if ptype and ptype != "command": return False command = config.get("command") return isinstance(command, str) and bool(command.strip()) def _resolve_command_provider_config( provider: str, tts_config: Dict[str, Any], ) -> Optional[Dict[str, Any]]: """The provider config when *provider* is a user-declared command provider. None for built-in names (native handlers win), unknown names, or non-command types. """ if not provider: return None key = provider.lower().strip() if key in BUILTIN_TTS_PROVIDERS: return None config = _get_named_provider_config(tts_config, key) return config if _is_command_provider_config(config) else None def _iter_command_providers(tts_config: Dict[str, Any]): """Yield (name, config) pairs for every declared command-type provider.""" for name, cfg in _get_provider_section(tts_config, "providers").items(): if ( isinstance(name, str) and name.lower() not in BUILTIN_TTS_PROVIDERS and _is_command_provider_config(cfg) ): yield name, cfg def _get_command_tts_timeout(config: Dict[str, Any]) -> float: """Timeout in seconds; invalid or non-positive values fall back to the default.""" raw = config.get("timeout", config.get("timeout_seconds", DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS)) try: value = float(raw) except (TypeError, ValueError): return float(DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS) if value <= 0: return float(DEFAULT_COMMAND_TTS_TIMEOUT_SECONDS) return value def _get_command_tts_output_format( config: Dict[str, Any], output_path: Optional[str] = None, ) -> str: """Validated output format: the output path's suffix wins, then ``format``/``output_format``.""" if output_path: suffix = Path(output_path).suffix.lower().strip().lstrip(".") if suffix in COMMAND_TTS_OUTPUT_FORMATS: return suffix raw = config.get("format") or config.get("output_format") or DEFAULT_COMMAND_TTS_OUTPUT_FORMAT fmt = str(raw).lower().strip().lstrip(".") return fmt if fmt in COMMAND_TTS_OUTPUT_FORMATS else DEFAULT_COMMAND_TTS_OUTPUT_FORMAT def _is_command_tts_voice_compatible(config: Dict[str, Any]) -> bool: """True only when the user explicitly opted in to voice delivery.""" value = config.get("voice_compatible", False) if isinstance(value, str): return value.strip().lower() in {"1", "true", "yes", "on"} return bool(value) def _configured_command_tts_output_path(path: Path, config: Dict[str, Any]) -> Path: """Return an output path whose extension matches the provider's output_format.""" return path.with_suffix(f".{_get_command_tts_output_format(config)}") def _generate_command_tts( text: str, output_path: str, provider_name: str, config: Dict[str, Any], tts_config: Dict[str, Any], ) -> str: """Generate speech by running a user-configured shell command. Returns the absolute path of the audio file the command wrote. Raises ``ValueError`` for invalid provider config and ``RuntimeError`` for timeouts / non-zero exits / empty output. """ command_template = str(config.get("command") or "").strip() if not command_template: raise ValueError( f"tts.providers.{provider_name}.command is not configured" ) output = Path(output_path).expanduser() output.parent.mkdir(parents=True, exist_ok=True) if output.exists(): output.unlink() timeout = _get_command_tts_timeout(config) output_format = _get_command_tts_output_format(config, str(output)) speed = config.get("speed", tts_config.get("speed", "")) with tempfile.TemporaryDirectory() as tmpdir: text_path = Path(tmpdir) / "input.txt" text_path.write_text(text, encoding="utf-8") placeholders = { "input_path": str(text_path), "text_path": str(text_path), "output_path": str(output), "format": output_format, "voice": str(config.get("voice", "")), "model": str(config.get("model", "")), "speed": str(speed), } command = render_command_template(command_template, placeholders) try: # Resolved through the origin so tests patching # ``tools.tts_tool._run_command_tts`` still intercept. from tools.tts_tool import _run_command_tts _run_command_tts( command, timeout, env_passthrough=command_env_passthrough(config), ) except subprocess.TimeoutExpired as exc: raise RuntimeError( f"TTS provider '{provider_name}' timed out after {timeout:g}s" ) from exc except subprocess.CalledProcessError as exc: detail_parts = [] if exc.stderr: detail_parts.append(f"stderr: {exc.stderr.strip()}") if exc.stdout: detail_parts.append(f"stdout: {exc.stdout.strip()}") detail = "; ".join(detail_parts) or "no command output" raise RuntimeError( f"TTS provider '{provider_name}' exited with code " f"{exc.returncode}: {detail}" ) from exc if not output.exists() or output.stat().st_size <= 0: raise RuntimeError( f"TTS provider '{provider_name}' produced no output at {output}" ) return str(output)