Files
hermes-agent/tools/tts_command_provider.py
Teknium d913175067 refactor(tts): move command-provider config layer into tts_command_provider
Completes the half-done move snapshotted by the parent (wip commit 283f63b):
the tts.providers.<name> config layer (_get_named_provider_config,
_resolve_command_provider_config, _generate_command_tts, ...) now lives in
tools/tts_command_provider.py; tts_tool re-imports the historical names.
2026-09-02 16:44:05 -07:00

489 lines
17 KiB
Python

"""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.<name>`` 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"(?<!\$)(?:\{{\{{(?P<double>{names})\}}\}}|\{{(?P<single>{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.<name>`` config layer
# ===========================================================================
# Any ``tts.provider`` value NOT in this set refers to ``tts.providers.<name>``.
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.<name>`` is canonical; ``tts.<name>`` 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)