From 6a6ee5bb67b7d4655404add14c52b100daf265b9 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 00:44:58 -0700 Subject: [PATCH 01/16] =?UTF-8?q?refactor(tools):=20process=5Fregistry=20s?= =?UTF-8?q?lice=20-25%=20LOC=20=E2=80=94=20shared=20exit/status/stdin/read?= =?UTF-8?q?er=20helpers,=20early-return=20watch=20limiter,=20suppress(),?= =?UTF-8?q?=20compact=20docs;=20schema=20byte-identical?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/process_registry.py | 2049 ++++++++--------------- tools/process_registry_notifications.py | 247 +-- 2 files changed, 778 insertions(+), 1518 deletions(-) diff --git a/tools/process_registry.py b/tools/process_registry.py index 950329e38e..efbdd77a05 100644 --- a/tools/process_registry.py +++ b/tools/process_registry.py @@ -1,15 +1,11 @@ -""" -Process Registry -- in-memory registry for background processes spawned via -terminal(background=true): rolling 200KB output buffer, poll/log/wait/kill, -crash recovery via a JSON checkpoint, and session-scoped tracking for gateway -reset protection. - -Background processes execute THROUGH the environment interface -- nothing runs -on the host unless TERMINAL_ENV=local; for Docker/Singularity/Modal/Daytona/SSH -the command runs inside the sandbox. +"""Process Registry -- in-memory registry for background processes spawned via +terminal(background=true): rolling 200KB output buffer, poll/log/wait/kill, JSON +checkpoint for crash recovery, session-scoped tracking for gateway reset protection. +Nothing runs on the host unless TERMINAL_ENV=local; other backends run in their sandbox. """ import codecs +from contextlib import suppress import json import logging import os @@ -23,9 +19,8 @@ import uuid from pathlib import Path _IS_WINDOWS = platform.system() == "Windows" -# systemd transient scopes exist only on Linux. Gate every scope-path branch -# on this constant (not merely "not Windows") so macOS and other POSIX -# platforms provably never touch systemd code (#70716 cross-platform audit). +# systemd transient scopes exist only on Linux; gate every scope-path branch on this +# (not merely "not Windows") so macOS and other POSIX platforms never touch systemd. _IS_LINUX = platform.system() == "Linux" from tools.environments.local import _find_shell, _resolve_safe_cwd, _sanitize_subprocess_env from hermes_cli._subprocess_compat import windows_hide_flags @@ -38,45 +33,35 @@ from agent.redact import redact_sensitive_text logger = logging.getLogger(__name__) - -# Checkpoint file for crash recovery (gateway only) +# Crash-recovery checkpoint (gateway only) CHECKPOINT_PATH = get_hermes_home() / "processes.json" -# Limits -MAX_OUTPUT_CHARS = 200_000 # 200KB rolling output buffer -FINISHED_TTL_SECONDS = 1800 # Keep finished processes for 30 minutes -MAX_PROCESSES = 64 # Max concurrent tracked processes (LRU pruning) +MAX_OUTPUT_CHARS = 200_000 # rolling output buffer +FINISHED_TTL_SECONDS = 1800 # keep finished processes 30 minutes +MAX_PROCESSES = 64 # max tracked processes (LRU pruning) -# Watch pattern rate limiting — PER SESSION. At most ONE watch-match notification -# every WATCH_MIN_INTERVAL_SECONDS; a match inside the cooldown is dropped and counts -# as one strike per window. After WATCH_STRIKE_LIMIT consecutive strike windows the -# session's watch_patterns are permanently disabled and it falls back to -# notify_on_complete semantics (one notification when the process exits). +# Watch-pattern rate limiting, PER SESSION: one watch-match notification per +# WATCH_MIN_INTERVAL_SECONDS; a match inside the cooldown is dropped and counts as one +# strike per window; WATCH_STRIKE_LIMIT consecutive strike windows permanently disable +# watching and fall back to notify_on_complete semantics. WATCH_MIN_INTERVAL_SECONDS = 15 WATCH_STRIKE_LIMIT = 3 - -# Lifetime cap — independent of the strike counter. A pattern recurring at a cadence -# just above the cooldown never trips the strike limit yet forces a full-context agent -# turn every time; watch_patterns is documented as "ONLY for rare one-shot signals", -# so after this many delivered matches we disable it and fall back to notify_on_complete. +# Lifetime cap, independent of strikes: a pattern recurring just above the cooldown never +# strikes yet forces a full-context agent turn each time; watch_patterns is "ONLY for +# rare one-shot signals", so after this many deliveries fall back to notify_on_complete. WATCH_LIFETIME_MAX_HITS = 8 - -# Global circuit breaker across all sessions — secondary safety net so concurrent -# siblings can't collectively flood the user even when each is under its own cap. +# Global circuit breaker across all sessions so concurrent siblings can't collectively +# flood the user even when each is under its own cap. WATCH_GLOBAL_MAX_PER_WINDOW = 15 WATCH_GLOBAL_WINDOW_SECONDS = 10 WATCH_GLOBAL_COOLDOWN_SECONDS = 30 -# --------------------------------------------------------------------------- -# systemd cgroup isolation for gateway-spawned local executors -# --------------------------------------------------------------------------- -# Under a systemd gateway with MemoryMax, local background commands inherit the -# gateway's cgroup, so a memory-heavy executor can get the ENTIRE gateway killed by -# systemd-oomd. Wrapping the spawn in ``systemd-run --user --scope`` gives the worker -# its own transient cgroup. Usability is probed once (the binary can exist while the -# user D-Bus session is absent — system services, containers) and cached. - +# --- systemd cgroup isolation for gateway-spawned local executors ------------------ +# Under a systemd gateway with MemoryMax, local background commands inherit the gateway's +# cgroup, so a memory-heavy executor can get the ENTIRE gateway killed by systemd-oomd; +# ``systemd-run --user --scope`` gives the worker its own transient cgroup. Usability is +# probed once (binary present but user D-Bus absent in system services/containers). _SYSTEMD_SCOPE_AVAILABLE: Optional[bool] = None _SYSTEMD_SCOPE_PROBE_LOCK = threading.Lock() _SYSTEMD_SCOPE_PROBED_AT = 0.0 @@ -88,11 +73,9 @@ _WORKER_MEMORY_MAX_CAP_BYTES = 4 * 1024 * 1024 * 1024 def _worker_memory_max_bytes() -> int: """Finite per-worker cgroup limit that can never widen host risk. - ``TERMINAL_LOCAL_MEMORY_MAX_MB`` is honored only when it *tightens* the safe bound (min of the gateway's cgroup-v2 ``memory.max`` and half of physical RAM, - capped at 4 GiB), so an oversized override cannot exceed the enclosing slice. - """ + capped at 4 GiB), so an oversized override cannot exceed the enclosing slice.""" override_bound: Optional[int] = None override = os.getenv("TERMINAL_LOCAL_MEMORY_MAX_MB", "").strip() if override: @@ -106,90 +89,59 @@ def _worker_memory_max_bytes() -> int: logger.warning( "Ignoring invalid TERMINAL_LOCAL_MEMORY_MAX_MB=%r; " "expected an integer representing at least %d MiB", - override, - _MIN_WORKER_MEMORY_MAX_BYTES // (1024 * 1024), - ) - + override, _MIN_WORKER_MEMORY_MAX_BYTES // (1024 * 1024)) candidates: List[int] = [] - try: - for line in Path("/proc/self/cgroup").read_text(encoding="utf-8").splitlines(): - if line.startswith("0::"): - relative = line.partition("::")[2].lstrip("/") - raw_limit = ( - Path("/sys/fs/cgroup") / relative / "memory.max" - ).read_text(encoding="utf-8").strip() - if raw_limit.isdigit() and int(raw_limit) >= _MIN_WORKER_MEMORY_MAX_BYTES: - candidates.append(int(raw_limit)) - break - except (OSError, ValueError): - pass - - try: + with suppress(OSError, ValueError): + lines = Path("/proc/self/cgroup").read_text(encoding="utf-8").splitlines() + v2 = next((ln for ln in lines if ln.startswith("0::")), None) + if v2 is not None: + relative = v2.partition("::")[2].lstrip("/") + raw_limit = (Path("/sys/fs/cgroup") / relative / "memory.max").read_text(encoding="utf-8").strip() + if raw_limit.isdigit() and int(raw_limit) >= _MIN_WORKER_MEMORY_MAX_BYTES: + candidates.append(int(raw_limit)) + with suppress(OSError, ValueError, TypeError): physical_bytes = int(os.sysconf("SC_PHYS_PAGES")) * int(os.sysconf("SC_PAGE_SIZE")) - candidates.append(min( - _WORKER_MEMORY_MAX_CAP_BYTES, - max(_MIN_WORKER_MEMORY_MAX_BYTES, physical_bytes // 2), - )) - except (OSError, ValueError, TypeError): - pass - + candidates.append(min(_WORKER_MEMORY_MAX_CAP_BYTES, max(_MIN_WORKER_MEMORY_MAX_BYTES, physical_bytes // 2))) safe_bound = min(candidates) if candidates else _DEFAULT_WORKER_MEMORY_MAX_BYTES return min(override_bound, safe_bound) if override_bound else safe_bound def _systemd_scope_argv(binary: str, unit_name: str, *argv: str) -> List[str]: - """``systemd-run --user --scope`` command line shared by the probe and real spawns. - - ``--collect`` makes the transient scope self-clean after exit; ``--unit`` gives it - a recognisable name for ``systemctl --user status`` / journalctl. - """ + """``systemd-run --user --scope`` argv shared by the probe and real spawns. + ``--collect`` self-cleans the scope after exit; ``--unit`` names it for systemctl.""" return [ - binary, "--user", "--scope", "--quiet", - "--unit", unit_name, - "--collect", + binary, "--user", "--scope", "--quiet", "--unit", unit_name, "--collect", "--property", "MemoryAccounting=yes", "--property", f"MemoryMax={_worker_memory_max_bytes()}", "--property", "OOMPolicy=kill", - "--", - *argv, + "--", *argv, ] def _systemd_scope_cached() -> Optional[bool]: - """Cached probe verdict, or None when a (re)probe is due. - - A True verdict is permanent; a False one expires after - ``_SYSTEMD_SCOPE_FAILURE_TTL_SECONDS`` so a transient D-Bus outage isn't sticky. - """ - cached = _SYSTEMD_SCOPE_AVAILABLE - if cached is True: + """Cached probe verdict, or None when a (re)probe is due. True is permanent; False + expires after ``_SYSTEMD_SCOPE_FAILURE_TTL_SECONDS`` so a D-Bus blip isn't sticky.""" + if _SYSTEMD_SCOPE_AVAILABLE is True: return True - if cached is False and time.monotonic() - _SYSTEMD_SCOPE_PROBED_AT < _SYSTEMD_SCOPE_FAILURE_TTL_SECONDS: - return False - return None + stale = time.monotonic() - _SYSTEMD_SCOPE_PROBED_AT >= _SYSTEMD_SCOPE_FAILURE_TTL_SECONDS + return None if _SYSTEMD_SCOPE_AVAILABLE is None or stale else False def _systemd_run_user_scope_available() -> bool: - """Return True if ``systemd-run --user --scope`` can create a cgroup. - - ``shutil.which`` alone is insufficient: system-service deployments and containers - may lack the user D-Bus session bus even though the binary is on PATH, so every - spawn would fail with ``Failed to connect to user bus``. We run a cheap no-op - probe (``systemd-run --user --scope --unit=… -- /bin/true``) and cache the outcome. - """ + """True if ``systemd-run --user --scope`` can create a cgroup. + ``shutil.which`` alone is insufficient: system services and containers may lack + the user D-Bus bus even with the binary on PATH (every spawn would fail with + ``Failed to connect to user bus``), so a cheap ``/bin/true`` probe is run and cached.""" global _SYSTEMD_SCOPE_AVAILABLE, _SYSTEMD_SCOPE_PROBED_AT verdict = _systemd_scope_cached() if verdict is not None: return verdict - - # Double-checked locking keeps concurrent first-use spawns from observing a - # temporary False while the definitive probe is still in flight — such a race - # would launch the losing workload back inside the gateway cgroup. + # Double-checked locking: a concurrent first-use spawn must not observe a temporary + # False mid-probe, or it would launch back inside the gateway cgroup. with _SYSTEMD_SCOPE_PROBE_LOCK: verdict = _systemd_scope_cached() if verdict is not None: return verdict - available = False if _IS_LINUX: try: @@ -197,23 +149,19 @@ def _systemd_run_user_scope_available() -> bool: binary = shutil.which("systemd-run") if binary: - # A unique unit avoids collisions; the timeout bounds D-Bus. + # Unique unit avoids collisions; the timeout bounds D-Bus. probe_unit = f"hermes-probe-scope-{os.getpid()}-{uuid.uuid4().hex[:8]}" result = subprocess.run( - _systemd_scope_argv(binary, probe_unit, "/bin/true"), - capture_output=True, - timeout=3, + _systemd_scope_argv(binary, probe_unit, "/bin/true"), capture_output=True, timeout=3, ) available = result.returncode == 0 if not available: logger.debug( "systemd-run --user --scope probe failed (rc=%s): %s", - result.returncode, - (result.stderr or b"").decode("utf-8", "replace").strip(), + result.returncode, (result.stderr or b"").decode("utf-8", "replace").strip(), ) except Exception as exc: logger.debug("systemd-run --user --scope probe error: %s", exc) - _SYSTEMD_SCOPE_AVAILABLE = available _SYSTEMD_SCOPE_PROBED_AT = time.monotonic() return available @@ -221,79 +169,51 @@ def _systemd_run_user_scope_available() -> bool: def _is_supervised_gateway_process() -> bool: """Whether this process is the live, supervised Hermes gateway itself. - - Supervisor markers and ``_HERMES_GATEWAY`` are inherited by every descendant - (and importing ``gateway.run`` sets the latter), so also require ownership of - the live gateway PID file — transient scopes are for the gateway, not terminal - children or unrelated CLIs in the same supervised tree. - """ + Supervisor markers and ``_HERMES_GATEWAY`` are inherited by every descendant (and + importing ``gateway.run`` sets the latter), so also require ownership of the live + gateway PID file — scopes are for the gateway, not terminal children or CLIs.""" if os.environ.get("_HERMES_GATEWAY") != "1": return False - try: from gateway.restart import is_gateway_supervisor_process from gateway.status import get_running_pid - return ( - is_gateway_supervisor_process() - and get_running_pid(cleanup_stale=False) == os.getpid() - ) + return is_gateway_supervisor_process() and get_running_pid(cleanup_stale=False) == os.getpid() except Exception as exc: logger.debug("Could not verify supervised gateway process identity: %s", exc) return False -def _build_systemd_scope_argv( - shell_argv: List[str], - unit_suffix: str, -) -> List[str]: +def _build_systemd_scope_argv(shell_argv: List[str], unit_suffix: str) -> List[str]: """Wrap *shell_argv* in a ``systemd-run --user --scope`` invocation with its own memory accounting, so an OOM in the worker cannot kill the gateway cgroup.""" import shutil binary = shutil.which("systemd-run") if binary is None: - # Caller should have checked _systemd_run_user_scope_available(); - # guard anyway so we never pass None into Popen. + # Caller should have probed availability; never pass None into Popen anyway. return shell_argv return _systemd_scope_argv(binary, f"hermes-worker-{unit_suffix}", *shell_argv) def _stop_systemd_unit(unit_name: str) -> bool: """Stop a transient systemd user scope by unit name. - - Reaps the *entire* cgroup — catching double-forked descendants that survive a - plain PID signal because they were reparented to init inside the scope. - ``systemctl --user stop`` SIGTERMs every process in the cgroup and escalates to - SIGKILL after ``TimeoutStopSec``. - - Returns True if the unit was stopped (or was already gone), False if - ``systemctl`` is unavailable or the stop command failed. - """ + Reaps the *entire* cgroup — catching double-forked descendants reparented to init + inside the scope that survive a plain PID signal (SIGTERM all, SIGKILL after + ``TimeoutStopSec``). True if stopped or already gone; False if ``systemctl`` is + unavailable or the stop failed.""" import shutil binary = shutil.which("systemctl") if binary is None: return False try: - result = subprocess.run( - [binary, "--user", "stop", unit_name], - capture_output=True, - timeout=15, - ) + result = subprocess.run([binary, "--user", "stop", unit_name], capture_output=True, timeout=15) if result.returncode != 0: stderr = (result.stderr or b"").decode(errors="replace").strip() - stderr_lower = stderr.lower() - if any( - marker in stderr_lower - for marker in ("not loaded", "not found", "does not exist") - ): + if any(marker in stderr.lower() for marker in ("not loaded", "not found", "does not exist")): return True - logger.debug( - "systemctl --user stop %s exited %d: %s", - unit_name, result.returncode, - stderr, - ) + logger.debug("systemctl --user stop %s exited %d: %s", unit_name, result.returncode, stderr) return False return True except Exception as exc: @@ -312,32 +232,43 @@ def format_uptime_short(seconds: int) -> str: return f"{hours}h {mins}m" +def _not_found(session_id: str) -> dict: + return {"status": "not_found", "error": f"No process with ID {session_id}"} + + +def _output_tail(session: "ProcessSession", n: int) -> str: + """Last *n* chars of the session output with ANSI sequences stripped.""" + from tools.ansi_strip import strip_ansi + + return strip_ansi(session.output_buffer[-n:]) + + @dataclass class ProcessSession: """A tracked background process with output buffering.""" - id: str # Unique session ID ("proc_xxxxxxxxxxxx") - command: str # Original command string - task_id: str = "" # Task/sandbox isolation key - owner_task_id: str = "" # RAW spawning task id (e.g. "sa-..."); task_id is the - # CONTAINER key (may be collapsed by _resolve_container_task_id) - # so ownership checks must use this field - session_key: str = "" # Gateway session key (for reset protection) - pid: Optional[int] = None # OS process ID + id: str # "proc_xxxxxxxxxxxx" + command: str + task_id: str = "" # Task/sandbox isolation key (CONTAINER key, + # may be collapsed by _resolve_container_task_id) + owner_task_id: str = "" # RAW spawning task id ("sa-..."); ownership + # checks must use this, not task_id + session_key: str = "" # Gateway session key (reset protection) + pid: Optional[int] = None process: Optional[subprocess.Popen] = None # Popen handle (local only) - env_ref: Any = None # Reference to the environment object - cwd: Optional[str] = None # Working directory - started_at: float = 0.0 # time.time() of spawn (wall clock) + env_ref: Any = None # Environment object (sandbox spawns) + cwd: Optional[str] = None + started_at: float = 0.0 # time.time() of spawn host_start_time: Optional[int] = None # kernel start ticks (/proc//stat f22) — PID-reuse guard - exited: bool = False # Whether the process has finished - exit_code: Optional[int] = None # Exit code (None if still running) + exited: bool = False + exit_code: Optional[int] = None # None while running completion_reason: str = "exited" # exited|killed|lost|failed_start|already_exited termination_source: str = "" # process.kill|kill_all|backend_lost|failed_start - output_buffer: str = "" # Rolling output (last MAX_OUTPUT_CHARS) + output_buffer: str = "" # Rolling tail (last max_output_chars) max_output_chars: int = MAX_OUTPUT_CHARS - detached: bool = False # True if recovered from crash (no pipe) + detached: bool = False # Recovered from checkpoint (no pipe) pid_scope: str = "host" # "host" for local/PTY PIDs, "sandbox" for env-local PIDs - systemd_unit: str = "" # transient scope unit name when spawned under systemd-run (#70716) - # Watcher/notification metadata (persisted for crash recovery) + systemd_unit: str = "" # transient scope unit name when spawned under systemd-run + # Watcher/notification routing (persisted for crash recovery) watcher_platform: str = "" watcher_chat_id: str = "" watcher_user_id: str = "" @@ -345,25 +276,22 @@ class ProcessSession: watcher_thread_id: str = "" watcher_message_id: str = "" # Triggering message id — reply anchor for topic routing watcher_interval: int = 0 # 0 = no watcher configured - # Session-db id of the spawning conversation; lets the gateway drop completions - # whose session was closed at a user boundary (/new) instead of injecting them - # into the chat's NEW session. + # Session-db id of the spawning conversation; lets the gateway drop completions whose + # session was closed at a user boundary (/new) instead of injecting into the NEW one. parent_session_id: str = "" - notify_on_complete: bool = False # Queue agent notification on exit + notify_on_complete: bool = False # Queue agent notification on exit watch_patterns: List[str] = field(default_factory=list) _watch_hits: int = field(default=0, repr=False) # total matches delivered _watch_suppressed: int = field(default=0, repr=False) # matches dropped by rate limit _watch_disabled: bool = field(default=False, repr=False) # permanently killed after strike limit - # Per-session rate-limit state (see WATCH_* constants). A strike is a WINDOW with - # drops, not a dropped match. - _watch_last_emit_at: float = field(default=0.0, repr=False) + # Rate-limit window state (see WATCH_*). A strike is a WINDOW with drops, not a drop. _watch_cooldown_until: float = field(default=0.0, repr=False) _watch_strike_candidate: bool = field(default=False, repr=False) _watch_consecutive_strikes: int = field(default=0, repr=False) _completion_event: threading.Event = field(default_factory=threading.Event, repr=False) _lock: threading.Lock = field(default_factory=threading.Lock) _reader_thread: Optional[threading.Thread] = field(default=None, repr=False) - _pty: Any = field(default=None, repr=False) # ptyprocess handle (when use_pty=True) + _pty: Any = field(default=None, repr=False) # ptyprocess handle (use_pty=True) def append_output(self, text: str) -> None: """Append to the rolling output buffer under the session lock, keeping the tail.""" @@ -372,16 +300,26 @@ class ProcessSession: if len(self.output_buffer) > self.max_output_chars: self.output_buffer = self.output_buffer[-self.max_output_chars:] + def mark_exited(self, exit_code, reason: str = "exited", source: str = "") -> None: + """Record an exit. A kill that raced the observer already recorded its own + exit_code/reason; never overwrite it.""" + self.exited = True + if self.completion_reason != "killed": + self.exit_code = exit_code + self.completion_reason = reason + if source: + self.termination_source = source + +# Watcher routing fields, in event-dict key order (``watcher_`` on the session). +_WATCHER_ROUTE_KEYS = ("platform", "chat_id", "user_id", "user_name", "thread_id", "message_id") # Session fields persisted verbatim in the crash-recovery checkpoint (plus # ``session_id``; ``command`` is redacted and ``owner_task_id`` defaulted on write). _CHECKPOINT_FIELDS = ( "command", "pid", "pid_scope", "host_start_time", "systemd_unit", "cwd", "started_at", "task_id", "owner_task_id", "session_key", - "watcher_platform", "watcher_chat_id", "watcher_user_id", "watcher_user_name", - "watcher_thread_id", "watcher_message_id", "watcher_interval", - "parent_session_id", "notify_on_complete", "watch_patterns", -) + *(f"watcher_{k}" for k in _WATCHER_ROUTE_KEYS), "watcher_interval", + "parent_session_id", "notify_on_complete", "watch_patterns") _CHECKPOINT_DEFAULTS = { f.name: ([] if f.name == "watch_patterns" else f.default) for f in ProcessSession.__dataclass_fields__.values() @@ -391,30 +329,21 @@ _CHECKPOINT_DEFAULTS = { class ProcessRegistry: """In-memory registry of running and finished background processes. - Thread-safe: accessed from executor threads (terminal_tool, process handlers), - the gateway asyncio loop (watchers, reset checks) and the cleanup thread. - """ + the gateway asyncio loop (watchers, reset checks) and the cleanup thread.""" _SHELL_NOISE_SUBSTRINGS = ( - "bash: cannot set terminal process group", - "bash: no job control in this shell", - "no job control in this shell", - "cannot set terminal process group", - "tcsetattr: Inappropriate ioctl for device", - ) + "no job control in this shell", "cannot set terminal process group", + "tcsetattr: Inappropriate ioctl for device") def __init__(self): self._running: Dict[str, ProcessSession] = {} self._finished: Dict[str, ProcessSession] = {} self._lock = threading.Lock() - # Side-channel for check_interval watchers (gateway reads after agent run) self.pending_watchers: List[Dict[str, Any]] = [] - - # Unified queue for all background events (completion, watch_match, - # async_delegation...; distinguished by "type"). CLI process_loop and the - # gateway drain it after each agent turn to auto-trigger new turns. + # Unified queue for all background events (distinguished by "type"); the CLI + # process_loop and the gateway drain it after each agent turn to trigger new turns. import queue as _queue_mod self.completion_queue: _queue_mod.Queue = _queue_mod.Queue() # Rehydrate durable delegation completions once, at registry startup. @@ -423,23 +352,19 @@ class ProcessRegistry: restore_undelivered_completions(self.completion_queue) except Exception as exc: logger.warning("Could not restore async delegation completions: %s", exc) - - # Sessions whose completion the agent already consumed via wait()/read_log() - # — it has the output in hand, so drain loops AND gateway/tui watchers skip. + # Completions the agent already consumed via wait()/read_log() (output in + # hand): drain loops AND gateway/tui watchers skip them. self._completion_consumed: set = set() - # Sessions merely *observed* exited via poll(). poll() is read-only and must - # NOT mark consumed (a status check would suppress the watcher's autonomous - # delivery turn), but on the CLI the poll result is inline in the same turn, - # so drain_notifications() skips these to avoid a duplicate [SYSTEM: ...]; + # Sessions merely *observed* exited via poll(). poll() is read-only and must NOT + # mark consumed (a status check would suppress the watcher's autonomous delivery + # turn), but the CLI has the poll result inline in the same turn, so + # drain_notifications() skips these to avoid a duplicate [SYSTEM: ...]; # gateway/tui watchers deliberately ignore this set. self._poll_observed: set = set() - # Global watch-match circuit breaker across all sessions. self._global_watch_lock = threading.Lock() - self._global_watch_window_start: float = 0.0 - self._global_watch_window_hits: int = 0 - self._global_watch_tripped_until: float = 0.0 - self._global_watch_suppressed_during_trip: int = 0 + self._global_watch_window_start = self._global_watch_tripped_until = 0.0 + self._global_watch_window_hits = self._global_watch_suppressed_during_trip = 0 # Driver-installed sinks (desktop gateway): on_output(session, chunk) streams # live output from reader threads; on_close(session_or_none, process_id) drops # a read-only terminal tab without killing the process. @@ -455,137 +380,97 @@ class ProcessRegistry: return "\n".join(lines) def _emit_output(self, session: ProcessSession, chunk: str) -> None: - """Forward a freshly-read chunk to the live-output sink, if one is set. - Called from reader threads; never raise into the read loop.""" + """Forward a chunk to the live-output sink; called from reader threads, never raises.""" sink = self.on_output if sink is None or not chunk: return - try: + with suppress(Exception): sink(session, chunk) - except Exception: - pass def _check_watch_patterns(self, session: ProcessSession, new_text: str) -> None: """Scan a freshly-read chunk for watch patterns and queue notifications. - - Rate limiting per session (see the WATCH_* constants): one match per cooldown - window, a match inside the window is one strike, WATCH_STRIKE_LIMIT - consecutive strikes or WATCH_LIFETIME_MAX_HITS total deliveries disable - watching and promote the session to notify_on_complete. - """ + Per-session rate limiting (see WATCH_* constants): one match per cooldown + window, a match inside the window is one strike, WATCH_STRIKE_LIMIT consecutive + strikes or WATCH_LIFETIME_MAX_HITS total deliveries disable watching and + promote the session to notify_on_complete.""" if not session.watch_patterns or session._watch_disabled: return # Late chunks after the reader declared exit are post-exit noise; dropping them # avoids stale notifications minutes after the process ended. if session.exited: return - - # Scan new text line-by-line for pattern matches - matched_lines = [] - matched_pattern = None - for line in new_text.splitlines(): - for pat in session.watch_patterns: - if pat in line: - matched_lines.append(line.rstrip()) - if matched_pattern is None: - matched_pattern = pat - break # one match per line is enough - - if not matched_lines: + hits = [ # (first matching pattern, line) — one match per line + (next(p for p in session.watch_patterns if p in line), line.rstrip()) + for line in new_text.splitlines() if any(p in line for p in session.watch_patterns)] + if not hits: return - + matched_pattern = hits[0][0] + matched_lines = [line for _, line in hits] now = time.time() - should_disable = False - lifetime_exhausted = False with session._lock: - # Case 1: inside the cooldown — drop, count one strike per window, and - # disable + promote once the strike limit is hit. if session._watch_cooldown_until and now < session._watch_cooldown_until: + # Inside the cooldown: drop, count one strike per window, disable + + # promote once the strike limit is hit. session._watch_suppressed += len(matched_lines) - if not session._watch_strike_candidate: - # First drop in this window — count one strike. - session._watch_strike_candidate = True - session._watch_consecutive_strikes += 1 - if session._watch_consecutive_strikes >= WATCH_STRIKE_LIMIT: - session._watch_disabled = True - # Promote to notify_on_complete so the agent still gets - # exactly one notification when the process actually ends. - session.notify_on_complete = True - should_disable = True - return_early = True - else: - # Case 2: cooldown expired. A prior window with no drops resets the - # consecutive-strike counter (healthy cadence again). - if session._watch_cooldown_until and not session._watch_strike_candidate: - session._watch_consecutive_strikes = 0 - session._watch_strike_candidate = False - - # Emit the notification and start a new cooldown window. - session._watch_last_emit_at = now - session._watch_cooldown_until = now + WATCH_MIN_INTERVAL_SECONDS - session._watch_hits += 1 - suppressed = session._watch_suppressed - session._watch_suppressed = 0 - return_early = False - # Lifetime cap: this match is still delivered, but no further ones. - lifetime_exhausted = session._watch_hits >= WATCH_LIFETIME_MAX_HITS - if lifetime_exhausted: - session._watch_disabled = True - session.notify_on_complete = True - - if return_early: - if should_disable: - # Exactly one summary so the agent/user sees why things went quiet. - self.completion_queue.put({ - **self._watch_event_base(session), - "type": "watch_disabled", - "suppressed": session._watch_suppressed, - "message": ( - f"Watch patterns disabled for process {session.id} — " - f"{WATCH_STRIKE_LIMIT} consecutive rate-limit windows triggered " - f"(min spacing {WATCH_MIN_INTERVAL_SECONDS}s). " - f"Falling back to notify_on_complete semantics; you'll get " - f"exactly one notification when the process exits." - ), - }) - return - - # Trim matched output to a reasonable size + if session._watch_strike_candidate: + return + session._watch_strike_candidate = True + session._watch_consecutive_strikes += 1 + if session._watch_consecutive_strikes < WATCH_STRIKE_LIMIT: + return + session._watch_disabled = True + # Promote so the agent still gets exactly one notification on exit, + # plus exactly one summary so it sees why things went quiet. + session.notify_on_complete = True + self._emit_watch_disabled( + session, session._watch_suppressed, + f"{WATCH_STRIKE_LIMIT} consecutive rate-limit windows triggered " + f"(min spacing {WATCH_MIN_INTERVAL_SECONDS}s). ") + return + # Cooldown expired. A prior window with no drops resets the + # consecutive-strike counter (healthy cadence again). + if session._watch_cooldown_until and not session._watch_strike_candidate: + session._watch_consecutive_strikes = 0 + session._watch_strike_candidate = False + # Emit and start a new cooldown window. + session._watch_cooldown_until = now + WATCH_MIN_INTERVAL_SECONDS + session._watch_hits += 1 + suppressed = session._watch_suppressed + session._watch_suppressed = 0 + # Lifetime cap: this match is still delivered, but no further ones. + lifetime_exhausted = session._watch_hits >= WATCH_LIFETIME_MAX_HITS + if lifetime_exhausted: + session._watch_disabled = True + session.notify_on_complete = True output = "\n".join(matched_lines[:20]) if len(output) > 2000: output = output[:2000] + "\n...(truncated)" - - if not self._global_watch_admit(now): - # Even when the breaker drops the final match, still explain the silence. - if lifetime_exhausted: - self._emit_lifetime_watch_disabled(session) - return - - notification = { - **self._watch_event_base(session), - "type": "watch_match", - "pattern": matched_pattern, - "output": output, - "suppressed": suppressed, - } - _redact_process_result(notification) - self.completion_queue.put(notification) - + if self._global_watch_admit(now): + notification = { + **self._watch_event_base(session), + "type": "watch_match", + "pattern": matched_pattern, + "output": output, + "suppressed": suppressed, + } + _redact_process_result(notification) + self.completion_queue.put(notification) + # Even when the breaker drops the final match, still explain the silence. if lifetime_exhausted: - self._emit_lifetime_watch_disabled(session) + self._emit_watch_disabled( + session, 0, f"reached the lifetime cap of {WATCH_LIFETIME_MAX_HITS} delivered matches. ", + ) - def _emit_lifetime_watch_disabled(self, session: ProcessSession) -> None: - """Queue the watch_disabled summary for the lifetime-cap path.""" + def _emit_watch_disabled(self, session: ProcessSession, suppressed: int, why: str) -> None: + """Queue the one-shot watch_disabled summary (strike-limit or lifetime-cap path).""" self.completion_queue.put({ **self._watch_event_base(session), "type": "watch_disabled", - "suppressed": 0, + "suppressed": suppressed, "message": ( - f"Watch patterns disabled for process {session.id} — " - f"reached the lifetime cap of {WATCH_LIFETIME_MAX_HITS} delivered " - f"matches. Falling back to notify_on_complete semantics; you'll get " - f"exactly one notification when the process exits." - ), + f"Watch patterns disabled for process {session.id} — {why}" + f"Falling back to notify_on_complete semantics; you'll get " + f"exactly one notification when the process exits."), }) @staticmethod @@ -597,89 +482,58 @@ class ProcessRegistry: "task_id": session.task_id, "owner_task_id": session.owner_task_id or session.task_id, "command": session.command, - "platform": session.watcher_platform, - "chat_id": session.watcher_chat_id, - "user_id": session.watcher_user_id, - "user_name": session.watcher_user_name, - "thread_id": session.watcher_thread_id, - "message_id": session.watcher_message_id, + **{key: getattr(session, f"watcher_{key}") for key in _WATCHER_ROUTE_KEYS}, } @staticmethod def _global_watch_event(type_: str, message: str, **extra) -> dict: """Unaddressed (all-sessions) watch breaker event.""" return { - "session_id": "", - "session_key": "", - "command": "", - "type": type_, - **extra, + "session_id": "", "session_key": "", "command": "", "type": type_, **extra, "message": message, - "platform": "", - "chat_id": "", - "user_id": "", - "user_name": "", - "thread_id": "", + "platform": "", "chat_id": "", "user_id": "", "user_name": "", "thread_id": "", } def _global_watch_admit(self, now: float) -> bool: """True if this watch_match may pass the global breaker. - In cooldown: drop and count. Otherwise slide the rolling window; exceeding the cap trips the breaker for WATCH_GLOBAL_COOLDOWN_SECONDS with ONE - "tripped" summary, and the cooldown's end emits ONE "released" summary. - """ - release_msg = None + "tripped" summary, and the cooldown's end emits ONE "released" summary.""" + events = [] # summary events, queued outside the lock with self._global_watch_lock: # Handle cooldown expiry first so we can emit the release summary. if self._global_watch_tripped_until and now >= self._global_watch_tripped_until: suppressed = self._global_watch_suppressed_during_trip self._global_watch_tripped_until = 0.0 self._global_watch_suppressed_during_trip = 0 - self._global_watch_window_start = now - self._global_watch_window_hits = 0 + self._global_watch_window_start, self._global_watch_window_hits = now, 0 if suppressed > 0: - # Queued outside the lock (below). - release_msg = self._global_watch_event( + events.append(self._global_watch_event( "watch_overflow_released", f"Watch-pattern notifications resumed. " f"{suppressed} match event(s) were suppressed during the flood.", - suppressed=suppressed, - ) - - # Still in cooldown — drop and count. + suppressed=suppressed)) if self._global_watch_tripped_until and now < self._global_watch_tripped_until: + # Still in cooldown — drop and count. self._global_watch_suppressed_during_trip += 1 admit = False - trip_now = None else: - # Slide the window. if now - self._global_watch_window_start >= WATCH_GLOBAL_WINDOW_SECONDS: - self._global_watch_window_start = now - self._global_watch_window_hits = 0 - - if self._global_watch_window_hits >= WATCH_GLOBAL_MAX_PER_WINDOW: - # Trip the breaker. + self._global_watch_window_start, self._global_watch_window_hits = now, 0 + admit = self._global_watch_window_hits < WATCH_GLOBAL_MAX_PER_WINDOW + if admit: + self._global_watch_window_hits += 1 + else: self._global_watch_tripped_until = now + WATCH_GLOBAL_COOLDOWN_SECONDS self._global_watch_suppressed_during_trip += 1 - trip_now = now - admit = False - else: - self._global_watch_window_hits += 1 - trip_now = None - admit = True - - # Queue summary events outside the lock. - if release_msg is not None: - self.completion_queue.put(release_msg) - if trip_now is not None: - self.completion_queue.put(self._global_watch_event( - "watch_overflow_tripped", - f"Watch-pattern overflow: >{WATCH_GLOBAL_MAX_PER_WINDOW} " - f"notifications in {WATCH_GLOBAL_WINDOW_SECONDS}s across all processes. " - f"Suppressing further watch_match events for " - f"{WATCH_GLOBAL_COOLDOWN_SECONDS}s.", - )) + events.append(self._global_watch_event( + "watch_overflow_tripped", + f"Watch-pattern overflow: >{WATCH_GLOBAL_MAX_PER_WINDOW} " + f"notifications in {WATCH_GLOBAL_WINDOW_SECONDS}s across all processes. " + f"Suppressing further watch_match events for " + f"{WATCH_GLOBAL_COOLDOWN_SECONDS}s.")) + for msg in events: + self.completion_queue.put(msg) return admit @staticmethod @@ -687,141 +541,108 @@ class ProcessRegistry: """Best-effort liveness check for host-visible PIDs.""" if not pid: return False - # ``os.kill(pid, 0)`` is NOT a no-op on Windows (bpo-14484) — use - # the cross-platform existence check. + # ``os.kill(pid, 0)`` is NOT a no-op on Windows (bpo-14484) — use the + # cross-platform existence check. from gateway.status import _pid_exists return _pid_exists(pid) @staticmethod def _safe_host_start_time(pid: Optional[int]) -> Optional[int]: """Kernel start ticks for a host PID, or None when unavailable.""" - if not pid: - return None try: from gateway.status import get_process_start_time - return get_process_start_time(pid) + return get_process_start_time(pid) if pid else None except Exception: return None @classmethod def _host_pid_is_ours(cls, pid: Optional[int], expected_start: Optional[int]) -> bool: """True only if ``pid`` is alive AND still the process we spawned. - The kernel recycles PIDs, so a stored number can later name an unrelated process (seen in the wild: a browser's session leader tree-killed). The kernel start time captured at spawn must match the live one; with no baseline - (legacy checkpoints, no ``/proc``) degrade to a bare liveness check. - """ - if not cls._is_host_pid_alive(pid): - return False - if expected_start is None: - return True - return cls._safe_host_start_time(pid) == expected_start + (legacy checkpoints, no ``/proc``) degrade to a bare liveness check.""" + return cls._is_host_pid_alive(pid) and ( + expected_start is None or cls._safe_host_start_time(pid) == expected_start) def _refresh_detached_session(self, session: Optional[ProcessSession]) -> Optional[ProcessSession]: """Update recovered host-PID sessions when the underlying process has exited.""" if session is None or session.exited or not session.detached or session.pid_scope != "host": return session - # A recycled PID (alive but not ours) counts as "our process exited" so a # later kill() can never tree-kill the stranger. if self._host_pid_is_ours(session.pid, session.host_start_time): return session - with session._lock: if session.exited: return session - session.exited = True - # Recovered sessions no longer have a waitable handle, so the real - # exit code is unavailable once the original process object is gone. - session.exit_code = None - + # No waitable handle survives recovery, so the real exit code is unknown. + session.exited, session.exit_code = True, None self._move_to_finished(session) return session @staticmethod def _proc_alive(proc) -> bool: - """True if a psutil.Process is running and not a zombie. - - A zombie is already dead (just unreaped), so there's nothing to SIGKILL. - """ + """True if a psutil.Process is running and not a zombie (already dead, just unreaped).""" try: import psutil - if not proc.is_running(): - return False - return proc.status() != psutil.STATUS_ZOMBIE + return proc.is_running() and proc.status() != psutil.STATUS_ZOMBIE except Exception: return False @staticmethod def _config_value(section: str, key: str, fallback): """``config.yaml`` value for ``section.key``, else the DEFAULT_CONFIG value. - Raises if config is unreadable; callers wrap with their own hard fallback so - registry code paths never crash on a broken config file. - """ + registry code paths never crash on a broken config file.""" from hermes_cli.config import DEFAULT_CONFIG, cfg_get, read_raw_config val = cfg_get(read_raw_config(), section, key) return DEFAULT_CONFIG[section][key] if val is None else val @staticmethod - def _daemon_term_grace_seconds() -> float: - """Grace window (s) between SIGTERM and escalated SIGKILL, floored at 0 - (0 disables escalation). ``terminal.daemon_term_grace_seconds``; 2.0 if - config is unreadable.""" + def _config_seconds(key: str, fallback: float) -> float: + """``terminal.`` as a non-negative float (0 disables); *fallback* if unreadable.""" try: - return max(float(ProcessRegistry._config_value("terminal", "daemon_term_grace_seconds", 2.0)), 0.0) + return max(float(ProcessRegistry._config_value("terminal", key, fallback)), 0.0) except Exception: - return 2.0 + return fallback + + @staticmethod + def _daemon_term_grace_seconds() -> float: + """Grace (s) between SIGTERM and escalated SIGKILL; 0 disables escalation.""" + return ProcessRegistry._config_seconds("daemon_term_grace_seconds", 2.0) @classmethod def _terminate_host_pid(cls, pid: int, expected_start: Optional[int] = None) -> None: """Terminate a host-visible PID and its descendants. - - ``expected_start`` (kernel start time at spawn) is re-validated first; a - mismatch or dead PID means the number was recycled onto a stranger and we - refuse to touch it — a leaked orphan beats tree-killing someone's browser. - - POSIX: psutil walks the tree and SIGTERMs children before the parent so - subprocess trees (Chromium renderers under an agent-browser daemon) aren't - reparented to init and survive. After ``terminal.daemon_term_grace_seconds`` - any survivor is SIGKILLed (0 disables escalation). - - Windows: ``taskkill /PID /T /F`` (same primitive as - ``gateway.status.terminate_pid``; ``/F`` is already a hard kill). The psutil - path is unusable there: PPID links go stale so ``children(recursive=True)`` - misses orphans, and ``terminate()`` is ``TerminateProcess()`` on one handle — - nothing cascades like a SIGTERM to a process group. The bare ``os.kill`` - fallback covers OSError/PermissionError and a missing ``taskkill.exe``. - """ + ``expected_start`` (kernel start time at spawn) is re-validated first: a mismatch + or dead PID means the number was recycled onto a stranger and we refuse to touch + it — a leaked orphan beats tree-killing someone's browser. POSIX: psutil SIGTERMs + children before the parent (so trees aren't reparented to init and survive), then + SIGKILLs survivors after ``terminal.daemon_term_grace_seconds``. Windows: + ``taskkill /T /F`` (psutil's stale PPID links miss orphans there); ``os.kill`` + is the fallback.""" if expected_start is not None and not cls._host_pid_is_ours(pid, expected_start): logger.warning( "Refusing to terminate host pid %d: start-time mismatch — " - "PID was recycled onto an unrelated process.", pid, - ) + "PID was recycled onto an unrelated process.", pid) return - def _sigterm_quietly(): - try: - os.kill(pid, signal.SIGTERM) - except (OSError, ProcessLookupError, PermissionError): - pass + def _sigterm_quietly(): + with suppress(OSError, ProcessLookupError, PermissionError): + os.kill(pid, signal.SIGTERM) if _IS_WINDOWS: try: subprocess.run( - ["taskkill", "/PID", str(pid), "/T", "/F"], - capture_output=True, - text=True, encoding='utf-8', errors='replace', - timeout=10, - creationflags=windows_hide_flags(), - stdin=subprocess.DEVNULL, - ) + ["taskkill", "/PID", str(pid), "/T", "/F"], capture_output=True, text=True, + encoding='utf-8', errors='replace', timeout=10, creationflags=windows_hide_flags(), + stdin=subprocess.DEVNULL) except (FileNotFoundError, subprocess.TimeoutExpired, OSError): _sigterm_quietly() return - import psutil + gone = (psutil.NoSuchProcess, psutil.AccessDenied, OSError) try: parent = psutil.Process(pid) except psutil.NoSuchProcess: @@ -829,47 +650,41 @@ class ProcessRegistry: except (OSError, PermissionError): _sigterm_quietly() return - # Snapshot the whole tree (children before parent) and SIGTERM each. try: targets = parent.children(recursive=True) - except (psutil.NoSuchProcess, psutil.AccessDenied, OSError): + except gone: targets = [] targets.append(parent) - for proc in targets: - try: + with suppress(gone): proc.terminate() - except (psutil.NoSuchProcess, psutil.AccessDenied, OSError): - pass - - # Escalate to SIGKILL for anything that ignored SIGTERM within the grace - # window. We deliberately do NOT trust ``psutil.wait_procs``' gone/alive - # partition: it reaps via ``Process.wait()`` and mis-partitions across - # zombie transitions in a parent/child tree, leaving survivors un-killed. - # A direct liveness re-probe of every target is deterministic. + # Escalate to SIGKILL for anything that ignored SIGTERM within the grace window. + # ``psutil.wait_procs``' gone/alive partition is deliberately NOT trusted: it + # reaps via ``Process.wait()`` and mis-partitions across zombie transitions in a + # parent/child tree, leaving survivors un-killed. Re-probing every target is + # deterministic. grace = cls._daemon_term_grace_seconds() if grace <= 0: return deadline = time.monotonic() + grace - while time.monotonic() < deadline: - if not any(cls._proc_alive(_p) for _p in targets): - break + while time.monotonic() < deadline and any(cls._proc_alive(_p) for _p in targets): time.sleep(0.05) for proc in targets: - try: - if not cls._proc_alive(proc): - continue - proc.kill() # SIGKILL on POSIX - logger.info( - "Escalated to SIGKILL for pid %d (ignored SIGTERM within " - "%.1fs grace)", proc.pid, grace, - ) - except (psutil.NoSuchProcess, psutil.AccessDenied, OSError): - pass + with suppress(gone): + if cls._proc_alive(proc): + proc.kill() # SIGKILL on POSIX + logger.info("Escalated to SIGKILL for pid %d (ignored SIGTERM within %.1fs grace)", proc.pid, grace) # ----- Spawn ----- + @staticmethod + def _new_session(command, task_id, owner_task_id, session_key, cwd, **extra) -> ProcessSession: + return ProcessSession( + id=f"proc_{uuid.uuid4().hex[:12]}", command=command, task_id=task_id, + owner_task_id=owner_task_id or task_id, session_key=session_key, cwd=cwd, + started_at=time.time(), **extra) + @staticmethod def _env_temp_dir(env: Any) -> str: """Return the writable sandbox temp dir for env-backed background tasks.""" @@ -883,31 +698,35 @@ class ProcessRegistry: logger.debug("Could not resolve environment temp dir: %s", exc) return "/tmp" - def _scope_argv(self, session: ProcessSession, argv: List[str], unit_suffix: str, label: str): - """Wrap *argv* in a transient systemd scope when we are the supervised gateway. - - Returns ``(argv, scoped)``. A scoped worker gets its own cgroup so an OOM kills - only the worker, not the gateway (and its messaging control plane). - """ + def _scope_argv(self, session: ProcessSession, safe_command: str, unit_suffix: str, label: str) -> List[str]: + """Login-shell argv for *safe_command* (parity with LocalEnvironment: rc files + sourced, user tools on PATH), wrapped in a transient systemd scope when we are + the supervised gateway (own cgroup: an OOM kills only the worker, not the + gateway and its messaging control plane).""" + argv = [_find_shell(), "-lic", f"set +m; {safe_command}"] in_supervised_gateway = _IS_LINUX and _is_supervised_gateway_process() if in_supervised_gateway and _systemd_run_user_scope_available(): session.systemd_unit = f"hermes-worker-{unit_suffix}.scope" - return _build_systemd_scope_argv(argv, unit_suffix=unit_suffix), True + return _build_systemd_scope_argv(argv, unit_suffix=unit_suffix) if in_supervised_gateway: - # Under a supervisor but no private cgroup: an OOM in the worker can - # still take the whole gateway down. + # Under a supervisor but no private cgroup: a worker OOM can still take + # the whole gateway down. logger.debug( "%s background executor not isolated in a systemd scope " - "(systemd-run --user unavailable); worker shares the gateway cgroup.", - label, - ) - return argv, False + "(systemd-run --user unavailable); worker shares the gateway cgroup.", label) + return argv - def _track_started(self, session: ProcessSession, reader_target, reader_name: str) -> None: + @staticmethod + def _spawn_env(env_vars: dict) -> dict: + """Sanitized child env; PYTHONUNBUFFERED so tqdm/datasets-style buffering + doesn't hide progress from process(action="poll").""" + env = _sanitize_subprocess_env(os.environ, env_vars) + env["PYTHONUNBUFFERED"] = "1" + return env + + def _track_started(self, session: ProcessSession, reader_target, reader_name: str, extra_args=()) -> None: """Start the output reader thread, register the session and checkpoint it.""" - reader = threading.Thread( - target=reader_target, args=(session,), daemon=True, name=reader_name, - ) + reader = threading.Thread(target=reader_target, args=(session, *extra_args), daemon=True, name=reader_name) session._reader_thread = reader reader.start() with self._lock: @@ -917,30 +736,19 @@ class ProcessRegistry: def _spawn_local_pty(self, session: ProcessSession, safe_command: str, env_vars: dict) -> ProcessSession: """PTY spawn for interactive CLI tools (Codex, Claude Code, REPLs). - Raises ImportError when no PTY backend is installed and re-raises any spawn - failure; ``spawn_local`` falls back to pipe mode in both cases. - """ + failure; ``spawn_local`` falls back to pipe mode in both cases.""" if _IS_WINDOWS: from winpty import PtyProcess as _PtyProcessCls else: from ptyprocess import PtyProcess as _PtyProcessCls - user_shell = _find_shell() - pty_env = _sanitize_subprocess_env(os.environ, env_vars) - pty_env["PYTHONUNBUFFERED"] = "1" + pty_env = self._spawn_env(env_vars) # A PTY is a real TTY, so pager-happy tools (git log/diff, man) WILL page and # hang waiting for `q` — default them to cat, honoring any pager the user set. pty_env.setdefault("GIT_PAGER", "cat") pty_env.setdefault("PAGER", "cat") - pty_argv, _ = self._scope_argv( - session, [user_shell, "-lic", f"set +m; {safe_command}"], session.id, "PTY", - ) - pty_proc = _PtyProcessCls.spawn( - pty_argv, - cwd=session.cwd, - env=pty_env, - dimensions=(30, 120), - ) + pty_argv = self._scope_argv(session, safe_command, session.id, "PTY") + pty_proc = _PtyProcessCls.spawn(pty_argv, cwd=session.cwd, env=pty_env, dimensions=(30, 120)) session.pid = pty_proc.pid session.host_start_time = self._safe_host_start_time(session.pid) session._pty = pty_proc @@ -948,38 +756,18 @@ class ProcessRegistry: return session def spawn_local( - self, - command: str, - cwd: str = None, - task_id: str = "", - session_key: str = "", - env_vars: dict = None, - use_pty: bool = False, - owner_task_id: str = "", - ) -> ProcessSession: - """Spawn a background process locally (TERMINAL_ENV=local only; other - backends use spawn_via_env()). - - ``use_pty`` requests a pseudo-terminal via ptyprocess/pywinpty for interactive - CLI tools; it falls back to a plain pipe when that is unavailable or fails. - """ + self, command: str, cwd: str = None, task_id: str = "", session_key: str = "", + env_vars: dict = None, use_pty: bool = False, owner_task_id: str = "") -> ProcessSession: + """Spawn a background process locally (TERMINAL_ENV=local; other backends use + spawn_via_env()). ``use_pty`` requests a pseudo-terminal via ptyprocess/pywinpty + for interactive CLIs, falling back to a plain pipe when unavailable or failing.""" # Bash parses ``A && B &`` as ``(A && B) &`` — a subshell that holds our stdout # pipe open forever when B is a long-running server. The rewriter turns it into # ``A && { B & }``. Lazy import: terminal_tool imports this module. from tools.terminal_tool import _rewrite_compound_background as _rewrite_bg safe_command = _rewrite_bg(command) - - session = ProcessSession( - id=f"proc_{uuid.uuid4().hex[:12]}", - command=command, - task_id=task_id, - owner_task_id=owner_task_id or task_id, - session_key=session_key, - cwd=_resolve_safe_cwd(cwd or os.getcwd()), - started_at=time.time(), - ) - + session = self._new_session(command, task_id, owner_task_id, session_key, _resolve_safe_cwd(cwd or os.getcwd())) pty_scope_attempted = False if use_pty: try: @@ -996,189 +784,99 @@ class ProcessRegistry: "to avoid duplicate command execution" ) from e session.systemd_unit = "" - - # Pipe path (non-PTY or PTY fallback). The user's login shell keeps parity with - # LocalEnvironment (rc files sourced, user tools on PATH). PYTHONUNBUFFERED so - # tqdm/datasets-style buffering doesn't hide progress from process(action="poll"). - user_shell = _find_shell() - bg_env = _sanitize_subprocess_env(os.environ, env_vars) - bg_env["PYTHONUNBUFFERED"] = "1" + # Pipe path (non-PTY or PTY fallback). _popen_kwargs = {"creationflags": windows_hide_flags()} if _IS_WINDOWS else {} - unit_suffix = f"{session.id}-pipe-fallback" if pty_scope_attempted else session.id - spawn_argv, _ = self._scope_argv( - session, [user_shell, "-lic", f"set +m; {safe_command}"], unit_suffix, "Local", - ) - + spawn_argv = self._scope_argv(session, safe_command, unit_suffix, "Local") # start_new_session is REQUIRED with systemd-run --scope too: the scope does not # give the worker a new session, so from an interactive TUI the worker would # share the foreground process group and background spawns would stop the whole # session (observed as dead TUIs in state T). Cgroup isolation is unaffected — # the scope attaches to the invoked process, not the spawning session. proc = subprocess.Popen( - spawn_argv, - text=True, - cwd=session.cwd, - env=bg_env, - encoding="utf-8", - errors="replace", - stdout=subprocess.PIPE, - stderr=subprocess.STDOUT, - stdin=subprocess.DEVNULL, - start_new_session=True, - **_popen_kwargs, - ) - + spawn_argv, text=True, cwd=session.cwd, env=self._spawn_env(env_vars), encoding="utf-8", + errors="replace", stdout=subprocess.PIPE, stderr=subprocess.STDOUT, stdin=subprocess.DEVNULL, + start_new_session=True, **_popen_kwargs) session.process = proc session.pid = proc.pid session.host_start_time = self._safe_host_start_time(session.pid) - try: self._track_started(session, self._reader_loop, f"proc-reader-{session.id}") except Exception: - # Post-Popen setup failed — kill the orphaned subprocess (and any setsid - # descendants) before re-raising so nothing leaks untracked. - try: - if session.systemd_unit: - # Scope teardown is the authoritative cleanup for the worker cgroup - # (never killpg here); the wrapper PID is terminated as fallback. - _stop_systemd_unit(session.systemd_unit) - self._terminate_host_pid(proc.pid, session.host_start_time) - elif not _IS_WINDOWS: - try: - kill_signal = getattr(signal, "SIGKILL", signal.SIGTERM) - os.killpg(os.getpgid(proc.pid), kill_signal) # windows-footgun: ok - guarded by _IS_WINDOWS above - except (ProcessLookupError, PermissionError, OSError): - proc.kill() - else: - proc.kill() - except Exception: - pass - try: - proc.wait(timeout=5) - except Exception: - pass + self._reap_untracked(session, proc) raise - return session + def _reap_untracked(self, session: ProcessSession, proc: subprocess.Popen) -> None: + """Post-Popen setup failed: kill the orphaned subprocess (and any setsid + descendants) so nothing leaks untracked.""" + with suppress(Exception): + if session.systemd_unit: + # Scope teardown is the authoritative cleanup for the worker cgroup + # (never killpg here); the wrapper PID is terminated as fallback. + _stop_systemd_unit(session.systemd_unit) + self._terminate_host_pid(proc.pid, session.host_start_time) + elif not _IS_WINDOWS: + try: + kill_signal = getattr(signal, "SIGKILL", signal.SIGTERM) + os.killpg(os.getpgid(proc.pid), kill_signal) # windows-footgun: ok - guarded by _IS_WINDOWS above + except (ProcessLookupError, PermissionError, OSError): + proc.kill() + else: + proc.kill() + with suppress(Exception): + proc.wait(timeout=5) + def spawn_via_env( - self, - env: Any, - command: str, - cwd: str = None, - task_id: str = "", - session_key: str = "", - timeout: int = 10, - owner_task_id: str = "", - ) -> ProcessSession: + self, env: Any, command: str, cwd: str = None, task_id: str = "", session_key: str = "", + timeout: int = 10, owner_task_id: str = "") -> ProcessSession: """Spawn a background process inside a non-local backend's sandbox. - - The command is wrapped to capture its in-sandbox PID and redirect output to - a log file, which later execute() calls poll. Less capable than local spawn - (no live pipe, no stdin) but runs in the correct sandbox context. - """ - session = ProcessSession( - id=f"proc_{uuid.uuid4().hex[:12]}", - command=command, - task_id=task_id, - owner_task_id=owner_task_id or task_id, - session_key=session_key, - cwd=cwd, - started_at=time.time(), - env_ref=env, - pid_scope="sandbox", - ) - - # Run the command in the sandbox with output capture + The command is wrapped to capture its in-sandbox PID and redirect output to a + log file that later execute() calls poll. No live pipe or stdin, but it runs in + the correct sandbox context.""" + session = self._new_session(command, task_id, owner_task_id, session_key, cwd, env_ref=env, pid_scope="sandbox") temp_dir = self._env_temp_dir(env) - log_path = f"{temp_dir}/hermes_bg_{session.id}.log" - pid_path = f"{temp_dir}/hermes_bg_{session.id}.pid" - exit_path = f"{temp_dir}/hermes_bg_{session.id}.exit" - quoted_command = shlex.quote(command) - quoted_temp_dir = shlex.quote(temp_dir) - quoted_log_path = shlex.quote(log_path) - quoted_pid_path = shlex.quote(pid_path) - quoted_exit_path = shlex.quote(exit_path) + log_path, pid_path, exit_path = (f"{temp_dir}/hermes_bg_{session.id}.{ext}" for ext in ("log", "pid", "exit")) + q = shlex.quote bg_command = ( - f"mkdir -p {quoted_temp_dir} && " - f"( nohup bash -lc {quoted_command} > {quoted_log_path} 2>&1; " - f"rc=$?; printf '%s\\n' \"$rc\" > {quoted_exit_path} ) & " - f"echo $! > {quoted_pid_path} && cat {quoted_pid_path}" - ) - + f"mkdir -p {q(temp_dir)} && " + f"( nohup bash -lc {q(command)} > {q(log_path)} 2>&1; " + f"rc=$?; printf '%s\\n' \"$rc\" > {q(exit_path)} ) & " + f"echo $! > {q(pid_path)} && cat {q(pid_path)}") try: - result = env.execute( - bg_command, - timeout=timeout, - rewrite_compound_background=False, - ) + result = env.execute(bg_command, timeout=timeout, rewrite_compound_background=False) output = result.get("output", "").strip() - # Try to extract the PID from the output - for line in output.splitlines(): - line = line.strip() - if line.isdigit(): - session.pid = int(line) - break + session.pid = next((int(ln) for ln in map(str.strip, output.splitlines()) if ln.isdigit()), None) # No PID from the wrapper (syntax error, broken redirect): a failed launch, # not a fake running session. if session.pid is None: - session.exited = True - session.exit_code = int(result.get("returncode", -1)) - if session.exit_code == 0: - session.exit_code = -1 - session.completion_reason = "failed_start" - session.termination_source = "failed_start" - session.output_buffer = result.get("output", "").strip() + session.mark_exited(int(result.get("returncode", -1)) or -1, "failed_start", "failed_start") + session.output_buffer = output except Exception as e: - session.exited = True - session.exit_code = -1 - session.completion_reason = "failed_start" - session.termination_source = "failed_start" + session.mark_exited(-1, "failed_start", "failed_start") session.output_buffer = f"Failed to start: {e}" - - if not session.exited: - # Start a poller thread that periodically reads the log file - reader = threading.Thread( - target=self._env_poller_loop, - args=(session, env, log_path, pid_path, exit_path), - daemon=True, - name=f"proc-poller-{session.id}", - ) - session._reader_thread = reader - reader.start() - - with self._lock: - self._prune_if_needed() - if not session.exited: - self._running[session.id] = session - - if not session.exited: - self._write_checkpoint() - + if session.exited: + with self._lock: + self._prune_if_needed() + else: + self._track_started( + session, self._env_poller_loop, f"proc-poller-{session.id}", (env, log_path, pid_path, exit_path)) return session # ----- Reader / Poller Threads ----- def _reader_loop(self, session: ProcessSession): """Background thread: read stdout from a local Popen process. - - Uses ``buffer.read1(4096)`` not ``TextIOWrapper.read(4096)``: on pipes the - latter blocks until EOF, landing "live" output in one burst at exit. - - Orphaned-pipe guard: a backgrounded grandchild (``node server.js &``) - inherits our pipe's write end, so EOF never arrives while it lives — a - blocking read would park this thread, ``session.exited`` would never flip - and ``notify_on_complete`` never fire (``_reconcile_local_exit`` only runs - lazily from poll()/wait()). On POSIX we ``select()`` with a short interval - and stop draining shortly after the direct child exits, mirroring - ``tools/environments/base.py::_wait_for_process``. Windows pipes lack - select(); the blocking path stays and the lazy reconcile is the safety net. - """ + ``buffer.read1(4096)`` not ``TextIOWrapper.read(4096)``: on pipes the latter + blocks until EOF, landing "live" output in one burst at exit. Orphaned-pipe + guard: a backgrounded grandchild (``node server.js &``) inherits our pipe's write + end so EOF never arrives while it lives, which would park this thread and never + fire ``notify_on_complete``; on POSIX we ``select()`` and stop draining shortly + after the direct child exits (mirrors ``environments/base.py::_wait_for_process``). + Windows pipes lack select(), so the lazy ``_reconcile_local_exit`` is the net.""" first_chunk = True - # A multibyte UTF-8 char split across read1() chunks would become U+FFFD - # mojibake with stateless decoding; the incremental decoder holds the partial - # sequence until the continuation bytes arrive. + # A split multibyte UTF-8 char would become U+FFFD with stateless decoding; the + # incremental decoder holds the partial sequence until the rest arrives. decoder = codecs.getincrementaldecoder("utf-8")(errors="replace") def _append_chunk(chunk: str): @@ -1187,101 +885,80 @@ class ProcessRegistry: chunk = self._clean_shell_noise(chunk) first_chunk = False self._ingest_output(session, chunk) - try: proc = session.process if proc is None or proc.stdout is None: return stdout = proc.stdout - raw_read = getattr(getattr(stdout, "buffer", None), "read1", None) + def _read_once(): + """One 4 KiB read: decoded text ('' for a partial multibyte tail), None at EOF.""" + if raw_read is None: # mocked/alternate streams without a raw buffer: less "live" + return stdout.read(4096) or None + raw = raw_read(4096) + return decoder.decode(raw) if raw else None # select() needs a real OS fd; mocked streams (tests, adapters) may lack - # fileno() and use the blocking loop instead. - fd = None - if raw_read is not None and not _IS_WINDOWS: - fileno = getattr(stdout, "fileno", None) - try: - candidate = fileno() if callable(fileno) else None - except Exception: - candidate = None - if isinstance(candidate, int) and candidate >= 0: - fd = candidate - + # fileno() and use the blocking read instead. + try: + fd = stdout.fileno() if raw_read is not None and not _IS_WINDOWS else None + except Exception: + fd = None + if not (isinstance(fd, int) and fd >= 0): + fd = None if fd is not None: import select as _select - - idle_after_exit = 0 - while True: + idle_after_exit = 0 + while True: + if fd is not None: try: ready, _, _ = _select.select([fd], [], [], 0.2) except (ValueError, OSError): break # fd already closed - if ready: - raw = raw_read(4096) - if not raw: - break # true EOF — all writers closed - chunk = decoder.decode(raw) - if chunk: - _append_chunk(chunk) - idle_after_exit = 0 - elif proc.poll() is not None: - # Direct child gone and pipe idle ~200ms: a few more cycles - # for a buffered tail, then stop rather than wait forever on - # an orphaned grandchild's pipe. - idle_after_exit += 1 + if not ready: + # Direct child gone and pipe idle ~200ms: a few more cycles for a + # buffered tail, then stop rather than wait forever on an orphaned + # grandchild's pipe. + if proc.poll() is not None: + idle_after_exit += 1 if idle_after_exit >= 3: break - else: - while True: - if raw_read is not None: - raw = raw_read(4096) - if not raw: - break - chunk = decoder.decode(raw) - if not chunk: - continue # partial multibyte sequence — wait for more bytes - else: - # Mocked/alternate streams without a raw buffer: less "live". - chunk = stdout.read(4096) - if not chunk: - break - + continue + chunk = _read_once() + if chunk is None: + break # true EOF — all writers closed + if chunk: _append_chunk(chunk) + idle_after_exit = 0 except Exception as e: logger.debug("Process stdout reader ended: %s", e) finally: - # Flush the decoder: a truncated multibyte sequence at EOF becomes one - # U+FFFD instead of vanishing. - try: - tail = decoder.decode(b"", final=True) - if tail: - _append_chunk(tail) - except Exception: - pass - # Always reap the child to prevent zombie processes. - try: - session.process.wait(timeout=5) - except Exception as e: - logger.debug("Process wait timed out or failed: %s", e) - self._finish_exited(session, session.process.returncode) + self._finish_reader( + session, decoder, _append_chunk, "Process", + lambda: session.process.wait(timeout=5), lambda: session.process.returncode) - def _env_poller_loop( - self, session: ProcessSession, env: Any, log_path: str, pid_path: str, exit_path: str - ): + def _finish_reader(self, session, decoder, append, label, wait, exit_code) -> None: + """Reader-thread teardown: flush the decoder (a truncated multibyte tail becomes + one U+FFFD instead of vanishing), reap the child (no zombies), record the exit.""" + with suppress(Exception): + tail = decoder.decode(b"", final=True) + if tail: + append(tail) + try: + wait() + except Exception as e: + logger.debug("%s wait timed out or failed: %s", label, e) + self._finish_exited(session, exit_code()) + + def _env_poller_loop(self, session: ProcessSession, env: Any, log_path: str, pid_path: str, exit_path: str): """Background thread: poll a sandbox log file for non-local backends.""" - quoted_log_path = shlex.quote(log_path) - quoted_pid_path = shlex.quote(pid_path) - quoted_exit_path = shlex.quote(exit_path) - prev_output_len = 0 # track delta for watch pattern scanning + q = shlex.quote + prev_output_len = 0 # delta tracking for watch-pattern scanning while not session.exited: - time.sleep(2) # Poll every 2 seconds + time.sleep(2) try: - # Read new output from the log file - result = env.execute(f"cat {quoted_log_path} 2>/dev/null", timeout=10) - new_output = result.get("output", "") + new_output = env.execute(f"cat {q(log_path)} 2>/dev/null", timeout=10).get("output", "") if new_output: - # Delta since the previous read feeds watch-pattern scanning. delta = new_output[prev_output_len:] if len(new_output) > prev_output_len else "" prev_output_len = len(new_output) with session._lock: @@ -1289,36 +966,23 @@ class ProcessRegistry: if delta: self._check_watch_patterns(session, delta) self._emit_output(session, delta) - - # Check if process is still running check = env.execute( - f"kill -0 \"$(cat {quoted_pid_path} 2>/dev/null)\" 2>/dev/null; echo $?", - timeout=5, - ) + f"kill -0 \"$(cat {q(pid_path)} 2>/dev/null)\" 2>/dev/null; echo $?", timeout=5) check_output = check.get("output", "").strip() if check_output and check_output.splitlines()[-1].strip() != "0": - # Process has exited -- get exit code captured by the wrapper shell. - exit_result = env.execute( - f"cat {quoted_exit_path} 2>/dev/null", - timeout=5, - ) - exit_str = exit_result.get("output", "").strip() + # Exited -- read the exit code captured by the wrapper shell. + exit_str = env.execute(f"cat {q(exit_path)} 2>/dev/null", timeout=5).get("output", "").strip() try: - session.exit_code = int(exit_str.splitlines()[-1].strip()) + exit_code = int(exit_str.splitlines()[-1].strip()) except (ValueError, IndexError): - session.exit_code = -1 - session.exited = True - if session.completion_reason != "killed": - session.completion_reason = "exited" - self._move_to_finished(session) + exit_code = -1 + session.exit_code = exit_code # unlike mark_exited, a raced kill still takes this code + self._finish_exited(session, exit_code) return - except Exception: # Environment might be gone (sandbox reaped, etc.) - session.exited = True - session.exit_code = -1 - session.completion_reason = "lost" - session.termination_source = "backend_lost" + session.exited, session.exit_code = True, -1 + session.completion_reason, session.termination_source = "lost", "backend_lost" self._move_to_finished(session) return @@ -1327,9 +991,6 @@ class ProcessRegistry: pty = session._pty # Same split-multibyte handling as _reader_loop. decoder = codecs.getincrementaldecoder("utf-8")(errors="replace") - - _append_text = lambda text: self._ingest_output(session, text) # noqa: E731 - try: while pty.isalive(): try: @@ -1338,25 +999,14 @@ class ProcessRegistry: # ptyprocess returns bytes; pywinpty returns str text = chunk if isinstance(chunk, str) else decoder.decode(chunk) if text: - _append_text(text) + self._ingest_output(session, text) except Exception: # EOFError included break except Exception as e: logger.debug("PTY stdout reader ended: %s", e) - - # Flush any partial multibyte sequence held by the decoder. - try: - tail = decoder.decode(b"", final=True) - if tail: - _append_text(tail) - except Exception: - pass - - try: - pty.wait() - except Exception as e: - logger.debug("PTY wait timed out or failed: %s", e) - self._finish_exited(session, pty.exitstatus if hasattr(pty, 'exitstatus') else -1) + self._finish_reader( + session, decoder, lambda t: self._ingest_output(session, t), "PTY", + pty.wait, lambda: pty.exitstatus if hasattr(pty, 'exitstatus') else -1) def _ingest_output(self, session: ProcessSession, text: str) -> None: """Buffer a freshly-read chunk, then scan watch patterns and stream it live.""" @@ -1365,32 +1015,20 @@ class ProcessRegistry: self._emit_output(session, text) def _finish_exited(self, session: ProcessSession, exit_code) -> None: - """Mark a reader-observed exit and move the session to finished. - - A kill that raced the reader already recorded its own exit_code/reason; - don't overwrite it. - """ - session.exited = True - if session.completion_reason != "killed": - session.exit_code = exit_code - session.completion_reason = "exited" + """Mark a reader-observed exit (a raced kill keeps its own code/reason) and finish.""" + session.mark_exited(exit_code) self._move_to_finished(session) def _move_to_finished(self, session: ProcessSession): """Move a session from running to finished. - Idempotent: kill_process() and the reader thread can both call this; only - the FIRST move enqueues the completion notification, so no duplicates. - """ + the FIRST move enqueues the completion notification, so no duplicates.""" with self._lock: was_running = self._running.pop(session.id, None) is not None self._finished[session.id] = session session._completion_event.set() self._write_checkpoint() - if was_running and session.notify_on_complete: - from tools.ansi_strip import strip_ansi - output_tail = strip_ansi(session.output_buffer[-2000:]) if session.output_buffer else "" notification = { "type": "completion", "session_id": session.id, @@ -1398,10 +1036,8 @@ class ProcessRegistry: "task_id": session.task_id, "owner_task_id": session.owner_task_id or session.task_id, "command": session.command, - "exit_code": session.exit_code, - "completion_reason": session.completion_reason, - "termination_source": session.termination_source, - "output": output_tail, + **self._exit_fields(session), + "output": _output_tail(session, 2000), # Stable producer identity across checkpoint recovery (unlike a # consumer-observed completion timestamp). "started_at": session.started_at, @@ -1409,6 +1045,14 @@ class ProcessRegistry: _redact_process_result(notification) self.completion_queue.put(notification) + @staticmethod + def _exit_fields(session: ProcessSession) -> dict: + return { + "exit_code": session.exit_code, + "completion_reason": session.completion_reason, + "termination_source": session.termination_source, + } + # ----- Query Methods ----- def is_completion_consumed(self, session_id: str) -> bool: @@ -1416,57 +1060,38 @@ class ProcessRegistry: return session_id in self._completion_consumed def is_session_waiting(self, session_id: str) -> bool: - """Whether a goal loop (``hermes_cli.goals`` wait barrier) should stay parked - on this session: still running AND, if it has ``watch_patterns``, none has - matched yet (a long-lived watcher unblocks on its trigger, not on exit). - Unknown/exited/already-fired sessions return False so a stale barrier can - never wedge the loop.""" - if not session_id: - return False + """Whether a goal loop (``hermes_cli.goals`` wait barrier) should stay parked on + this session: still running AND, with ``watch_patterns``, none matched yet (a + long-lived watcher unblocks on its trigger, not on exit). Unknown/exited/ + already-fired sessions return False so a stale barrier can never wedge the loop.""" with self._lock: - session = self._running.get(session_id) or self._finished.get(session_id) + session = (self._running.get(session_id) or self._finished.get(session_id)) if session_id else None if session is None: return False - try: + with suppress(Exception): self._refresh_detached_session(session) - except Exception: - pass - if session.exited: - return False - return not (session.watch_patterns and not session._watch_disabled and session._watch_hits > 0) + return not session.exited and not ( + session.watch_patterns and not session._watch_disabled and session._watch_hits > 0) def wait_for_pending_completions( - self, - task_id: Optional[str] = None, - *, - timeout: float | None = None, - poll_interval: float = 1.0, + self, task_id: Optional[str] = None, *, timeout: float | None = None, poll_interval: float = 1.0, ) -> dict: """Bounded linger for ``notify_on_complete`` background processes at one-shot exit. - - A one-shot CLI run (``hermes -q/-Q/-z``) exits when its turn ends; any - background process it spawned still holds a stdout pipe owned by the dying - parent and dies of SIGPIPE seconds later (Bot Mode handoff replies sent via - message_agent/bot_relay were the visible casualty). Only ``notify_on_complete`` - processes carry a completion contract — servers/daemons/watchers aren't the - parent's to wait for. - - ``task_id=None`` waits on every tracked process (a one-shot process hosts one - agent). ``timeout=None`` reads ``terminal.oneshot_completion_wait_seconds``; - ``<= 0`` disables. Each ``poll_interval`` pass re-reconciles child state so an - orphaned-pipe exit can't wedge the linger. Returns - ``{"waited": [...], "completed": [...], "timed_out": [...]}`` of session ids. - """ + A one-shot CLI run (``hermes -q/-Q/-z``) exits when its turn ends; a background + process it spawned still holds a stdout pipe owned by the dying parent and dies of + SIGPIPE seconds later (Bot Mode handoff replies were the visible casualty). Only + ``notify_on_complete`` processes carry a completion contract — servers/daemons/ + watchers aren't the parent's to wait for. ``task_id=None`` waits on every tracked + process; ``timeout=None`` reads ``terminal.oneshot_completion_wait_seconds`` (``<= 0`` + disables). Each pass re-reconciles child state so an orphaned-pipe exit can't wedge + the linger. Returns ``{"waited", "completed", "timed_out"}`` id lists.""" if timeout is None: timeout = self._oneshot_completion_wait_seconds() result: dict = {"waited": [], "completed": [], "timed_out": []} with self._lock: pending = [ - s - for s in self._running.values() - if s.notify_on_complete - and not s.exited - and (task_id is None or s.task_id == task_id) + s for s in self._running.values() + if s.notify_on_complete and not s.exited and (task_id is None or s.task_id == task_id) ] if not pending or timeout <= 0: return result @@ -1474,17 +1099,13 @@ class ProcessRegistry: logger.info( "One-shot exit lingering (bounded %ss) for %d notify_on_complete " "background process(es): %s", - timeout, - len(pending), - ", ".join(s.id for s in pending), - ) + timeout, len(pending), ", ".join(s.id for s in pending)) deadline = time.monotonic() + max(float(timeout), 0.0) interval = max(float(poll_interval), 0.05) try: from tools.interrupt import is_interrupted as _is_interrupted except Exception: - def _is_interrupted() -> bool: - return False + _is_interrupted = lambda: False # noqa: E731 interrupted = False for session in pending: try: @@ -1496,11 +1117,9 @@ class ProcessRegistry: if remaining <= 0: break # Reconcile first so orphaned-pipe and detached exits fire the event. - try: + with suppress(Exception): self._reconcile_local_exit(session) self._refresh_detached_session(session) - except Exception: - pass if session.exited: break session._completion_event.wait(min(remaining, interval)) @@ -1508,73 +1127,62 @@ class ProcessRegistry: # Stop waiting, but never let the interrupt skip the caller's durable # teardown (session flush, end_session) that follows. interrupted = True - if session.exited: - result["completed"].append(session.id) - else: - result["timed_out"].append(session.id) + result["completed" if session.exited else "timed_out"].append(session.id) if result["timed_out"]: logger.warning( "One-shot exit linger timed out after %ss with %d background " "process(es) still running: %s — they may be killed when this " "process exits.", - timeout, - len(result["timed_out"]), - ", ".join(result["timed_out"]), - ) + timeout, len(result["timed_out"]), ", ".join(result["timed_out"])) return result @staticmethod def _oneshot_completion_wait_seconds() -> float: - """Bounded linger (s) for one-shot exits with pending notify_on_complete - processes: ``terminal.oneshot_completion_wait_seconds`` (0 disables), 600 - if config is unreadable.""" - try: - return max(float(ProcessRegistry._config_value("terminal", "oneshot_completion_wait_seconds", 600.0)), 0.0) - except Exception: - return 600.0 + """Linger (s) for one-shot exits with pending notify_on_complete processes; 0 disables.""" + return ProcessRegistry._config_seconds("oneshot_completion_wait_seconds", 600.0) - def _drain_should_skip( - self, session_id: str, *, skip_poll_observed: bool = True - ) -> bool: - """Skip a completion the CLI agent already has this turn — consumed via - wait/log or observed inline via poll(). Gateway/tui watchers check only - ``is_completion_consumed`` so a read-only poll never suppresses their - autonomous delivery turn.""" - return session_id in self._completion_consumed or ( - skip_poll_observed and session_id in self._poll_observed - ) + def _drain_should_skip(self, session_id: str, *, skip_poll_observed: bool = True) -> bool: + """Skip a completion the CLI agent already has this turn — consumed via wait/log + or observed inline via poll(). Gateway/tui watchers check only + ``is_completion_consumed`` so a read-only poll never suppresses their turn.""" + return session_id in self._completion_consumed or (skip_poll_observed and session_id in self._poll_observed) @staticmethod def _surface_child_process_notifications() -> bool: - """Whether subagent-owned process notifications surface in the parent - (``delegation.surface_child_process_notifications``; suppress on any config - error — never crash the drain loop).""" + """``delegation.surface_child_process_notifications``; False on any config + error — never crash the drain loop.""" try: return bool(ProcessRegistry._config_value("delegation", "surface_child_process_notifications", False)) except Exception: return False + @staticmethod + def _owns_event(evt: dict, session_key: str, owns_event, is_async_delegation: bool) -> bool: + """Routing verdict for one drained event (see drain_notifications); False = requeue.""" + evt_session_key = str(evt.get("session_key") or "") + requires_positive_proof = is_async_delegation or bool(evt_session_key or evt.get("origin_ui_session_id")) + if owns_event is not None and requires_positive_proof: + try: + return bool(owns_event(evt)) + except Exception: + return False # fail closed — never leak on a broken check + if session_key and requires_positive_proof: + return evt_session_key == session_key + # Restored payloads from a previous process: an unfiltered drain cannot prove + # ownership, so leave them for the owner. + return not (is_async_delegation and evt.get("restored")) + def drain_notifications( - self, - session_key: str = "", - owns_event=None, - *, - skip_poll_observed: bool = True, + self, session_key: str = "", owns_event=None, *, skip_poll_observed: bool = True, ) -> "list[tuple[dict, str]]": """Pop all pending events and return ``(raw_event, formatted_text)`` pairs. - - Skips completions per ``_drain_should_skip``; gateway/TUI callers pass - ``skip_poll_observed=False``. - - Routing: async-delegation events always need ownership proof; ordinary - events need it once they carry ``session_key`` or ``origin_ui_session_id``. - ``owns_event(evt)`` (strongest; the TUI passes a compression-chain-aware - check so a post-compression session still claims its pre-compression - dispatches) consumes ONLY on True; ``session_key`` uses plain equality. - Non-owned routed events are re-queued for their owner. With no filter every - event is consumed (legacy single-session), except restored delegation - payloads, which stay fail-closed. - """ + Skips completions per ``_drain_should_skip`` (gateway/TUI pass + ``skip_poll_observed=False``). Routing (``_owns_event``): async-delegation events + always need ownership proof, ordinary events once they carry ``session_key`` or + ``origin_ui_session_id``; ``owns_event(evt)`` (strongest; the TUI passes a + compression-chain-aware check) consumes ONLY on True, ``session_key`` uses plain + equality; non-owned events are re-queued for their owner. No filter consumes + everything (legacy single-session) except restored delegation payloads (fail-closed).""" results: "list[tuple[dict, str]]" = [] requeue: "list[dict]" = [] # delegation.surface_child_process_notifications, read at most once per drain @@ -1586,45 +1194,22 @@ class ProcessRegistry: except Exception: break is_async_delegation = evt.get("type") == "async_delegation" - evt_session_key = str(evt.get("session_key") or "") - evt_origin_sid = str(evt.get("origin_ui_session_id") or "") - requires_positive_proof = is_async_delegation or bool( - evt_session_key or evt_origin_sid - ) - if owns_event is not None and requires_positive_proof: - try: - owned = bool(owns_event(evt)) - except Exception: - owned = False # fail closed — never leak on a broken check - if not owned: - requeue.append(evt) - continue - elif session_key and requires_positive_proof: - if evt_session_key != session_key: - requeue.append(evt) - continue - elif is_async_delegation and evt.get("restored"): - # Restored payloads from a previous process: an unfiltered drain - # cannot prove ownership, so leave them for the owner. + if not self._owns_event(evt, session_key, owns_event, is_async_delegation): requeue.append(evt) continue # Routing happened first so a foreign session cannot drop the owner's # event via its own consumed/observed state. _evt_sid = evt.get("session_id", "") if evt.get("type") == "completion" and self._drain_should_skip( - _evt_sid, skip_poll_observed=skip_poll_observed - ): + _evt_sid, skip_poll_observed=skip_poll_observed): continue - # Subagent-owned process notifications are suppressed by default — the # child's delegation result is the deliverable. Judge ownership on # owner_task_id (RAW spawning id; task_id is the container key, collapsed # by _resolve_container_task_id). Dropped, NOT requeued: children never # drain, so a requeue would pin the event forever. 'async_delegation' # is the result itself and is NEVER suppressed. - _evt_task_id = str( - evt.get("owner_task_id") or evt.get("task_id") or "" - ) + _evt_task_id = str(evt.get("owner_task_id") or evt.get("task_id") or "") if not is_async_delegation and _evt_task_id.startswith("sa-"): if surface_child is None: surface_child = self._surface_child_process_notifications() @@ -1633,67 +1218,48 @@ class ProcessRegistry: "Suppressed subagent-owned process notification " "(delegation.surface_child_process_notifications=false): " "type=%s session_id=%s task_id=%s", - evt.get("type", "completion"), - _evt_sid, - _evt_task_id, - ) + evt.get("type", "completion"), _evt_sid, _evt_task_id) continue - - text = format_process_notification(evt) - if text: + if text := format_process_notification(evt): results.append((evt, text)) for evt in requeue: self.completion_queue.put(evt) return results - # Minimum characters of the random suffix required for prefix resolution. - # Short prefixes ("p", "pr", "proc_1") are too collision-prone to act on. + # Minimum suffix chars for prefix resolution; "p"/"proc_1" are too collision-prone. _MIN_PREFIX_CHARS = 4 def get(self, session_id: str) -> Optional[ProcessSession]: - """Get a session by full ID or unique prefix (``proc_4dae`` / bare ``4dae``, - like git short hashes). Ambiguous or too-short prefixes resolve to None, - never to an arbitrary pick.""" + """Session by full ID or unique prefix (``proc_4dae`` / bare ``4dae``, like git + short hashes); ambiguous or too-short prefixes resolve to None, never a guess.""" with self._lock: session = self._running.get(session_id) or self._finished.get(session_id) - if session is None: - session = self._resolve_prefix(session_id) - return self._refresh_detached_session(session) + return self._refresh_detached_session(session if session is not None else self._resolve_prefix(session_id)) def _resolve_prefix(self, session_id: str) -> Optional[ProcessSession]: - """Resolve a unique session-ID prefix (prefix-only, unique hit; a bare hex - tail is normalized to ``proc_``). :meth:`get` tries exact first.""" - if not session_id or not isinstance(session_id, str): - return None - query = session_id.strip() + """Resolve a unique session-ID prefix (a bare hex tail is normalized to + ``proc_``); :meth:`get` tries exact first.""" + query = session_id.strip() if isinstance(session_id, str) else "" if not query: return None - # Allow the bare suffix form: "4dae56" -> "proc_4dae56". if not query.startswith("proc_"): query = f"proc_{query}" - suffix = query[len("proc_"):] - if len(suffix) < self._MIN_PREFIX_CHARS: + if len(query) - len("proc_") < self._MIN_PREFIX_CHARS: return None with self._lock: matches = [ - s - for store in (self._running, self._finished) - for sid, s in store.items() - if sid.startswith(query) + s for store in (self._running, self._finished) + for sid, s in store.items() if sid.startswith(query) ] - if len(matches) == 1: - return matches[0] - return None + return matches[0] if len(matches) == 1 else None def _reconcile_local_exit(self, session: "ProcessSession") -> None: """Reconcile ``session.exited`` against the real child state. - - The reader flips ``exited`` only at EOF; when the direct child has exited but - a descendant (e.g. a daemon from ``hermes update``) holds the pipe open, poll() + The reader flips ``exited`` only at EOF; when the direct child has exited but a + descendant (e.g. a daemon from ``hermes update``) holds the pipe open, poll() would report "running" forever. If ``Popen.poll()`` has an exit code, drain - readable bytes non-blocking and flip ``exited``; the stuck daemon reader thread - is reaped with the process. No-op for env/PTY, exited and detached sessions. - """ + readable bytes non-blocking and flip ``exited``. No-op for env/PTY, exited and + detached sessions.""" if session is None or session.exited: return proc = getattr(session, "process", None) @@ -1705,9 +1271,7 @@ class ProcessRegistry: return if rc is None: return # Direct child still running — reader block is legitimate. - # Best-effort non-blocking drain of whatever the reader hasn't consumed. - drained = "" stdout = getattr(proc, "stdout", None) if stdout is not None and not _IS_WINDOWS: try: @@ -1716,67 +1280,46 @@ class ProcessRegistry: flags = fcntl.fcntl(fd, fcntl.F_GETFL) fcntl.fcntl(fd, fcntl.F_SETFL, flags | os.O_NONBLOCK) try: - chunk = stdout.read() - if chunk: - drained = chunk if isinstance(chunk, str) else chunk.decode("utf-8", errors="replace") - except (BlockingIOError, OSError, ValueError): - pass + with suppress(BlockingIOError, OSError, ValueError): + chunk = stdout.read() + if chunk: + session.append_output(chunk if isinstance(chunk, str) else chunk.decode("utf-8", errors="replace")) finally: - try: + with suppress(Exception): fcntl.fcntl(fd, fcntl.F_SETFL, flags) - except Exception: - pass except Exception as e: logger.debug("Non-blocking drain failed for %s: %s", session.id, e) - with session._lock: - if drained: - session.output_buffer += drained - if len(session.output_buffer) > session.max_output_chars: - session.output_buffer = session.output_buffer[-session.max_output_chars:] - session.exited = True - if session.completion_reason != "killed": - session.exit_code = rc - session.completion_reason = "exited" + session.mark_exited(rc) logger.info( "Reconciled session %s: direct child exited with code %s but reader " "was still blocked (orphaned pipe). Flipped to exited.", - session.id, rc, - ) + session.id, rc) self._move_to_finished(session) + @staticmethod + def _status_head(session: ProcessSession) -> dict: + return {"session_id": session.id, "command": session.command, "status": "exited" if session.exited else "running"} + def poll(self, session_id: str) -> dict: """Check status and get new output for a background process.""" - from tools.ansi_strip import strip_ansi - session = self.get(session_id) if session is None: - return {"status": "not_found", "error": f"No process with ID {session_id}"} - + return _not_found(session_id) self._reconcile_local_exit(session) # orphaned-pipe reader guard - with session._lock: - output_preview = strip_ansi(session.output_buffer[-1000:]) if session.output_buffer else "" - + output_preview = _output_tail(session, 1000) result = { - "session_id": session.id, - "command": session.command, - "status": "exited" if session.exited else "running", - "pid": session.pid, - "uptime_seconds": int(time.time() - session.started_at), - "output_preview": output_preview, - } + **self._status_head(session), "pid": session.pid, + "uptime_seconds": int(time.time() - session.started_at), "output_preview": output_preview} if session.exited: - result["exit_code"] = session.exit_code - result["completion_reason"] = session.completion_reason - result["termination_source"] = session.termination_source + result.update(self._exit_fields(session)) # Read-only: record in _poll_observed (CLI inline dedup) but NOT in # _completion_consumed, or a status check would suppress the watcher's # autonomous delivery turn. See __init__. self._poll_observed.add(session_id) if session.detached: - result["detached"] = True - result["note"] = "Process recovered after restart -- output history unavailable" + result.update(detached=True, note="Process recovered after restart -- output history unavailable") return result def read_log(self, session_id: str, offset: int | None = None, limit: int = 200) -> dict: @@ -1785,14 +1328,11 @@ class ProcessRegistry: session = self.get(session_id) if session is None: - return {"status": "not_found", "error": f"No process with ID {session_id}"} - + return _not_found(session_id) with session._lock: full_output = strip_ansi(session.output_buffer) - lines = full_output.splitlines() total_lines = len(lines) - # offset=None -> last N lines; an explicit offset=0 means the HEAD (don't # conflate the two via falsiness). if offset is None and limit > 0: @@ -1802,155 +1342,92 @@ class ProcessRegistry: offset = offset or 0 selected = lines[offset:offset + limit] stop = slice(offset, offset + limit).indices(total_lines)[1] - observed_completion_output = ( - total_lines == 0 or (bool(selected) and stop == total_lines) - ) - + observed_completion_output = total_lines == 0 or (bool(selected) and stop == total_lines) result = { - "session_id": session.id, - "command": session.command, - "status": "exited" if session.exited else "running", - "output": "\n".join(selected), - "total_lines": total_lines, - "showing": f"{len(selected)} lines", - } + **self._status_head(session), "output": "\n".join(selected), + "total_lines": total_lines, "showing": f"{len(selected)} lines"} if session.exited and observed_completion_output: self._completion_consumed.add(session_id) return result def wait(self, session_id: str, timeout: int = None) -> dict: """Block until the process exits, the timeout elapses, or the user interrupts. - ``timeout`` defaults to (and is clamped by) TERMINAL_TIMEOUT. Returns a dict - with status exited|timeout|interrupted|not_found|error and an output snapshot. - """ - from tools.ansi_strip import strip_ansi + with status exited|timeout|interrupted|not_found|error and an output snapshot.""" from tools.interrupt import is_interrupted as _is_interrupted try: - default_timeout = int(os.getenv("TERMINAL_TIMEOUT", "180")) + max_timeout = int(os.getenv("TERMINAL_TIMEOUT", "180")) except (ValueError, TypeError): - default_timeout = 180 - max_timeout = default_timeout - requested_timeout = timeout - timeout_note = None - + max_timeout = 180 # The schema says minimum=1 but not every caller enforces it; timeout=0 is # falsy and would silently fall through to the default wait. - if requested_timeout is not None and requested_timeout <= 0: - return { - "status": "error", - "error": f"timeout must be positive (got {requested_timeout})", - } - - if requested_timeout and requested_timeout > max_timeout: + if timeout is not None and timeout <= 0: + return {"status": "error", "error": f"timeout must be positive (got {timeout})"} + timeout_note = None + effective_timeout = timeout or max_timeout + if timeout and timeout > max_timeout: effective_timeout = max_timeout - timeout_note = ( - f"Requested wait of {requested_timeout}s was clamped " - f"to configured limit of {max_timeout}s" - ) - else: - effective_timeout = requested_timeout or max_timeout - + timeout_note = f"Requested wait of {timeout}s was clamped to configured limit of {max_timeout}s" session = self.get(session_id) if session is None: - return {"status": "not_found", "error": f"No process with ID {session_id}"} - + return _not_found(session_id) deadline = time.monotonic() + effective_timeout - while time.monotonic() < deadline: session = self._refresh_detached_session(session) if session is None: - return {"status": "not_found", "error": f"No process with ID {session_id}"} + return _not_found(session_id) self._reconcile_local_exit(session) # orphaned-pipe reader guard + result = None if session.exited: self._completion_consumed.add(session_id) result = self._exit_snapshot(session, "exited") elif _is_interrupted(): result = { - "status": "interrupted", - "command": session.command, - "output": strip_ansi(session.output_buffer[-1000:]), - "note": "User sent a new message -- wait interrupted", - } - else: - result = None + "status": "interrupted", "command": session.command, "output": _output_tail(session, 1000), + "note": "User sent a new message -- wait interrupted"} if result is not None: if timeout_note: result["timeout_note"] = timeout_note return result - remaining = deadline - time.monotonic() if remaining <= 0: break session._completion_event.wait(timeout=min(1.0, remaining)) - result = { - "status": "timeout", - "command": session.command, - "output": strip_ansi(session.output_buffer[-1000:]), - # Not a failure — models re-issued identical waits after misreading - # this result as an error. - "process_running": True, - } - uptime = time.time() - session.started_at if session.started_at else None + "status": "timeout", "command": session.command, "output": _output_tail(session, 1000), + # Not a failure — models re-issued identical waits after misreading this as an error. + "process_running": True} base_note = ( - f"Wait window of {effective_timeout}s elapsed — the process is " - "still running. This is not an error." - ) - if uptime is not None: - base_note += f" Uptime: {int(uptime)}s." - if session.notify_on_complete: - base_note += ( - " notify_on_complete is set: you will be notified on exit — " - "do more work instead of waiting again." - ) - else: - base_note += ( - " Poll again later or use terminal(background=true, " - "notify_on_complete=true) next time for automatic notification." - ) - if timeout_note: - result["timeout_note"] = f"{timeout_note}. {base_note}" - else: - result["timeout_note"] = base_note + f"Wait window of {effective_timeout}s elapsed — the process is still running. This is not an error.") + if session.started_at: + base_note += f" Uptime: {int(time.time() - session.started_at)}s." + base_note += ( + " notify_on_complete is set: you will be notified on exit — do more work instead of waiting again." + if session.notify_on_complete else + " Poll again later or use terminal(background=true, " + "notify_on_complete=true) next time for automatic notification.") + result["timeout_note"] = f"{timeout_note}. {base_note}" if timeout_note else base_note return result @staticmethod def _exit_snapshot(session: ProcessSession, status: str) -> dict: """Result dict for an exited session: exit metadata + last 2000 chars of output.""" - from tools.ansi_strip import strip_ansi - return { - "status": status, - "command": session.command, - "exit_code": session.exit_code, - "completion_reason": session.completion_reason, - "termination_source": session.termination_source, - "output": strip_ansi(session.output_buffer[-2000:]), - } + "status": status, "command": session.command, + **ProcessRegistry._exit_fields(session), "output": _output_tail(session, 2000)} def kill_process( - self, - session_id: str, - *, - source: str = "process.kill", - consume_output: bool = True, + self, session_id: str, *, source: str = "process.kill", consume_output: bool = True, ) -> dict: """Kill a background process and return its output snapshot. - ``consume_output`` is true for explicit tool/RPC kills (the caller sees the output). Bulk cleanup passes false so it doesn't suppress an autonomous - completion notification — except abandoned-turn reaping - (``kill_started_since``), which passes true so a killed abandoned process - can't enqueue a follow-up reviving work the timeout stopped. - """ - from tools.ansi_strip import strip_ansi - + completion notification — except abandoned-turn reaping (``kill_started_since``), + which passes true so a killed abandoned process can't revive stopped work.""" session = self.get(session_id) if session is None: - return {"status": "not_found", "error": f"No process with ID {session_id}"} - + return _not_found(session_id) if session.exited: # A double-forked descendant may still be alive in the systemd scope even # though the main process exited — stop the scope to reap survivors. @@ -1963,12 +1440,10 @@ class ProcessRegistry: if consume_output: self._completion_consumed.add(session_id) return result - try: early = self._signal_kill(session, session_id, consume_output) if early is not None: return early - # Additive to the PID kill: stopping the scope reaps double-forked # descendants reparented inside the cgroup. if session.systemd_unit: @@ -1976,7 +1451,7 @@ class ProcessRegistry: # Capture output, mark consumed, THEN expose ``exited`` to watcher tasks — # closes the delayed-notification race without losing the transcript. with session._lock: - output = strip_ansi(session.output_buffer[-2000:]) + output = _output_tail(session, 2000) if consume_output: self._completion_consumed.add(session_id) session.exited = True @@ -1986,12 +1461,8 @@ class ProcessRegistry: self._move_to_finished(session) self._write_checkpoint() return { - "status": "killed", - "session_id": session.id, - "completion_reason": session.completion_reason, - "termination_source": session.termination_source, - "output": output, - } + "status": "killed", "session_id": session.id, "completion_reason": session.completion_reason, + "termination_source": session.termination_source, "output": output} except Exception as e: return {"status": "error", "error": str(e)} @@ -1999,8 +1470,6 @@ class ProcessRegistry: """Deliver the kill via PTY, local Popen tree, sandbox exec or recovered host PID. Returns a final result dict when the kill cannot proceed (recycled/dead recovered PID, or no runtime handle), else None.""" - from tools.ansi_strip import strip_ansi - if session._pty: try: session._pty.terminate(force=True) @@ -2023,7 +1492,7 @@ class ProcessRegistry: with session._lock: session.exited = True session.exit_code = None - output = strip_ansi(session.output_buffer[-2000:]) + output = _output_tail(session, 2000) if consume_output: self._completion_consumed.add(session_id) self._move_to_finished(session) @@ -2032,140 +1501,97 @@ class ProcessRegistry: else: return { "status": "error", - "error": ( - "Recovered process cannot be killed after restart because " - "its original runtime handle is no longer available" - ), + "error": "Recovered process cannot be killed after restart because " + "its original runtime handle is no longer available", } return None - def _live_session(self, session_id: str): - """``(session, None)`` for a running session, else ``(None, error_result)``.""" + def _stdin_op(self, session_id: str, pty_op, pipe_op, ok: dict) -> dict: + """Run a stdin operation on a running session — ``pty_op(pty)`` under PTY mode, + else ``pipe_op(stdin)`` on the Popen pipe — and return *ok* on success.""" session = self.get(session_id) if session is None: - return None, {"status": "not_found", "error": f"No process with ID {session_id}"} + return _not_found(session_id) if session.exited: - return None, {"status": "already_exited", "error": "Process has already finished"} - return session, None + return {"status": "already_exited", "error": "Process has already finished"} + try: + if session._pty: + pty_op(session._pty) + elif not session.process or not session.process.stdin: + return {"status": "error", "error": "Process stdin not available (non-local backend or stdin closed)"} + else: + pipe_op(session.process.stdin) + return ok + except Exception as e: + return {"status": "error", "error": str(e)} def write_stdin(self, session_id: str, data: str) -> dict: """Send raw data to a running process's stdin (no newline appended).""" - session, err = self._live_session(session_id) - if err: - return err - # PTY mode -- write through pty handle. - if session._pty: - try: - # pywinpty expects str on Windows; ptyprocess expects bytes on POSIX. - if _IS_WINDOWS: - pty_data = data.decode("utf-8") if isinstance(data, bytes) else str(data) - else: - # surrogateescape: a PTY is a byte stream — round-trip the - # original bytes instead of crashing on surrogate content. - pty_data = data.encode("utf-8", "surrogateescape") if isinstance(data, str) else data - session._pty.write(pty_data) - return {"status": "ok", "bytes_written": len(data)} - except Exception as e: - return {"status": "error", "error": str(e)} + def via_pty(pty): + # pywinpty expects str on Windows; ptyprocess expects bytes on POSIX. + if _IS_WINDOWS: + pty.write(data.decode("utf-8") if isinstance(data, bytes) else str(data)) + else: + # surrogateescape: a PTY is a byte stream — round-trip the original + # bytes instead of crashing on surrogate content. + pty.write(data.encode("utf-8", "surrogateescape") if isinstance(data, str) else data) - # Popen mode -- write through stdin pipe - if not session.process or not session.process.stdin: - return {"status": "error", "error": "Process stdin not available (non-local backend or stdin closed)"} - try: - session.process.stdin.write(data) - session.process.stdin.flush() - return {"status": "ok", "bytes_written": len(data)} - except Exception as e: - return {"status": "error", "error": str(e)} + def via_pipe(stdin): + stdin.write(data) + stdin.flush() + return self._stdin_op(session_id, via_pty, via_pipe, {"status": "ok", "bytes_written": len(data)}) def submit_stdin(self, session_id: str, data: str = "") -> dict: """Send data + newline to stdin (like pressing Enter). - On a Windows PTY, Enter is a carriage return: ConPTY treats ``\\r`` as - end-of-line and a bare ``\\n`` through pywinpty is NOT a line terminator — - the child's blocking line read (``readline()``, Go ``bufio.Scanner`` in - ``gh auth login``) never returns and the process hangs looking healthy. - ``\\r\\n`` gives it both; POSIX PTYs and pipes keep ``\\n``. - """ + end-of-line and a bare ``\\n`` through pywinpty is NOT a line terminator — the + child's blocking line read (``readline()``, Go ``bufio.Scanner``) never returns + and the process hangs looking healthy. ``\\r\\n`` gives it both; POSIX keeps ``\\n``.""" session = self.get(session_id) - is_windows_pty = bool(_IS_WINDOWS and session is not None and session._pty) - return self.write_stdin(session_id, data + ("\r\n" if is_windows_pty else "\n")) + return self.write_stdin(session_id, data + ("\r\n" if _IS_WINDOWS and session and session._pty else "\n")) def request_close_terminal(self, session_id: str) -> dict: - """Ask the desktop GUI to close this process's read-only terminal tab. - - Does NOT kill the process — output keeps buffering and the tab can be - reopened from the status stack. Errors when no UI close sink is wired.""" - sink = self.on_close - if sink is None: - return { - "status": "error", - "error": "close_terminal is only available in the Hermes desktop app.", - } + """Ask the desktop GUI to close this process's read-only terminal tab. Does NOT + kill the process — output keeps buffering and the tab can be reopened from the + status stack. Errors when no UI close sink is wired.""" + if self.on_close is None: + return {"status": "error", "error": "close_terminal is only available in the Hermes desktop app."} # The session may already be finished (or pruned) — the tab can still # linger and be closed, so a missing session is not an error here. - session = self.get(session_id) try: - sink(session, session_id) + self.on_close(self.get(session_id), session_id) except Exception as e: return {"status": "error", "error": str(e)} return { - "status": "ok", - "closed": session_id, - "note": ( - "Closed the read-only terminal tab. The process was not killed; " - "its output remains available and the user can reopen the tab " - "from the status stack." - ), - } + "status": "ok", "closed": session_id, + "note": "Closed the read-only terminal tab. The process was not killed; " + "its output remains available and the user can reopen the tab " + "from the status stack."} def close_stdin(self, session_id: str) -> dict: """Close a running process's stdin / send EOF without killing the process.""" - session, err = self._live_session(session_id) - if err: - return err - - if session._pty: - try: - session._pty.sendeof() - return {"status": "ok", "message": "EOF sent"} - except Exception as e: - return {"status": "error", "error": str(e)} - - if not session.process or not session.process.stdin: - return {"status": "error", "error": "Process stdin not available (non-local backend or stdin closed)"} - try: - session.process.stdin.close() - return {"status": "ok", "message": "stdin closed"} - except Exception as e: - return {"status": "error", "error": str(e)} + session = self.get(session_id) + msg = "EOF sent" if session is not None and session._pty else "stdin closed" + return self._stdin_op( + session_id, lambda pty: pty.sendeof(), lambda stdin: stdin.close(), {"status": "ok", "message": msg}) def count_running(self) -> int: - """O(1) count of running processes for status-bar polling; CPython dict - ``len()`` is atomic so no lock is needed.""" - try: - return len(self._running) - except Exception: - return 0 + """O(1) running count for status-bar polling; dict ``len()`` is atomic, no lock.""" + return len(self._running) def list_sessions(self, task_id: str = None, session_key: str = None) -> list: - """List running and recently-finished processes for ``task_id`` and/or - ``session_key``. Cross-task entries that share the gateway session (a - forgotten preview server blocking session reset) are flagged - ``"session_scoped": true``.""" + """Running and recently-finished processes for ``task_id`` and/or ``session_key``; + cross-task entries sharing the gateway session (a forgotten preview server + blocking session reset) are flagged ``"session_scoped": true``.""" with self._lock: all_sessions = list(self._running.values()) + list(self._finished.values()) - all_sessions = [self._refresh_detached_session(s) for s in all_sessions] - if task_id or session_key: all_sessions = [ s for s in all_sessions - if (task_id and s.task_id == task_id) - or (session_key and s.session_key == session_key) + if (task_id and s.task_id == task_id) or (session_key and s.session_key == session_key) ] - result = [] for s in all_sessions: entry = { @@ -2182,8 +1608,7 @@ class ProcessRegistry: entry["session_scoped"] = True # Trigger metadata for goal-loop judges (a watcher may never exit). if s.watch_patterns and not s._watch_disabled: - entry["watch_patterns"] = list(s.watch_patterns) - entry["watch_hit"] = s._watch_hits > 0 + entry.update(watch_patterns=list(s.watch_patterns), watch_hit=s._watch_hits > 0) if s.notify_on_complete: entry["notify_on_complete"] = True if s.exited: @@ -2206,21 +1631,18 @@ class ProcessRegistry: return any(not s.exited and predicate(s) for s in self._running.values()) def has_active_processes(self, task_id: str) -> bool: - """Check if there are active (running) processes for a task_id.""" + """Whether any process for ``task_id`` is still running.""" return self._any_running(lambda s: s.task_id == task_id) - def has_active_for_session( - self, session_key: str, max_active_age: Optional[float] = None, - ) -> bool: + def has_active_for_session(self, session_key: str, max_active_age: Optional[float] = None) -> bool: """Active processes for a gateway session key. Processes older than - ``max_active_age`` seconds are ignored as stale so a forgotten - ``http.server`` can't freeze session idle/daily reset forever; ``None`` - keeps legacy behaviour (any running process blocks).""" + ``max_active_age`` seconds are ignored as stale so a forgotten ``http.server`` + can't freeze session idle/daily reset forever; ``None`` keeps legacy behaviour + (any running process blocks).""" now = time.time() return self._any_running( lambda s: s.session_key == session_key - and (max_active_age is None or (now - s.started_at) < max_active_age) - ) + and (max_active_age is None or (now - s.started_at) < max_active_age)) def has_any_active(self) -> bool: """Whether ANY background process is running — scale-to-zero must not @@ -2232,72 +1654,38 @@ class ProcessRegistry: only processes absent from the starting snapshot belong to the abandoned turn; older ones intentionally span turns and must survive.""" with self._lock: - return frozenset( - s.id - for s in self._running.values() - if s.task_id == task_id and not s.exited - ) + return frozenset(s.id for s in self._running.values() if s.task_id == task_id and not s.exited) - def kill_started_since( - self, - task_id: str, - baseline_ids, - *, - source: str, - ) -> int: + def kill_started_since(self, task_id: str, baseline_ids, *, source: str) -> int: """Kill ``task_id`` processes created after ``baseline_ids``. Output is consumed so an abandoned turn can't enqueue a follow-up reviving work the timeout deliberately stopped.""" - return self.kill_all( - task_id, - exclude_ids=frozenset(baseline_ids or ()), - source=source, - consume_output=True, - ) + return self.kill_all(task_id, exclude_ids=frozenset(baseline_ids or ()), source=source, consume_output=True) def kill_all( - self, - task_id: Optional[str] = None, - *, - exclude_ids: frozenset = frozenset(), - source: str = "kill_all", - consume_output: bool = False, - ) -> int: + self, task_id: Optional[str] = None, *, exclude_ids: frozenset = frozenset(), + source: str = "kill_all", consume_output: bool = False) -> int: """Kill all running processes, optionally filtered by task_id. Returns count killed.""" with self._lock: targets = [ s for s in self._running.values() - if (task_id is None or s.task_id == task_id) - and s.id not in exclude_ids - and not s.exited + if (task_id is None or s.task_id == task_id) and s.id not in exclude_ids and not s.exited ] - - killed = 0 - for session in targets: - result = self.kill_process( - session.id, - source=source, - consume_output=consume_output, - ) - if result.get("status") in {"killed", "already_exited"}: - killed += 1 - return killed + return sum( + self.kill_process(s.id, source=source, consume_output=consume_output).get("status") + in {"killed", "already_exited"} + for s in targets) # ----- Cleanup / Pruning ----- def _prune_if_needed(self): - """Remove oldest finished sessions if over MAX_PROCESSES. Must hold _lock.""" - # First prune expired finished sessions + """Drop expired finished sessions, then the oldest survivor while over + MAX_PROCESSES. Must hold _lock.""" now = time.time() - expired = [ - sid for sid, s in self._finished.items() - if (now - s.started_at) > FINISHED_TTL_SECONDS - ] - if len(self._running) + len(self._finished) - len(expired) >= MAX_PROCESSES: - # Still over the limit: also drop the oldest surviving finished session. - survivors = [sid for sid in self._finished if sid not in expired] - if survivors: - expired.append(min(survivors, key=lambda sid: self._finished[sid].started_at)) + expired = [sid for sid, s in self._finished.items() if (now - s.started_at) > FINISHED_TTL_SECONDS] + over_cap = len(self._running) + len(self._finished) - len(expired) >= MAX_PROCESSES + if over_cap and (survivors := [sid for sid in self._finished if sid not in expired]): + expired.append(min(survivors, key=lambda sid: self._finished[sid].started_at)) for sid in expired: del self._finished[sid] # Belt-and-suspenders against module-lifetime growth: forget consumed / @@ -2308,159 +1696,110 @@ class ProcessRegistry: # ----- Checkpoint (crash recovery) ----- - def _write_checkpoint( - self, - extra_entries: Optional[List[Dict[str, Any]]] = None, - ): - """Write running process metadata to checkpoint file atomically.""" + def _write_checkpoint(self, extra_entries: Optional[List[Dict[str, Any]]] = None): + """Write running process metadata to the checkpoint file atomically.""" try: with self._lock: entries = [] for s in self._running.values(): - if not s.exited: - # Backfill the start time so recovery can detect PID recycling - # even for sessions spawned before this field existed. - if s.host_start_time is None and s.pid_scope == "host" and s.pid: - s.host_start_time = self._safe_host_start_time(s.pid) - entry = {"session_id": s.id, **{f: getattr(s, f) for f in _CHECKPOINT_FIELDS}} - # Redact inline credentials before persisting: the file lives - # at ~/.hermes/processes.json. Recovery uses command only for - # display (adoption re-validates the PID, never re-runs it), - # so masking is lossless. - entry["command"] = redact_sensitive_text(s.command, code_file=True) - entry["owner_task_id"] = s.owner_task_id or s.task_id - entries.append(entry) + if s.exited: + continue + # Backfill the start time so recovery can detect PID recycling + # even for sessions spawned before this field existed. + if s.host_start_time is None and s.pid_scope == "host" and s.pid: + s.host_start_time = self._safe_host_start_time(s.pid) + entry = {"session_id": s.id, **{f: getattr(s, f) for f in _CHECKPOINT_FIELDS}} + # Redact inline credentials before persisting (~/.hermes/processes.json). + # Recovery uses command only for display (adoption re-validates the + # PID, never re-runs it), so masking is lossless. + entry["command"] = redact_sensitive_text(s.command, code_file=True) + entry["owner_task_id"] = s.owner_task_id or s.task_id + entries.append(entry) if extra_entries: tracked_ids = {item.get("session_id") for item in entries} - entries.extend( - item - for item in extra_entries - if item.get("session_id") not in tracked_ids - ) - - # Atomic write to avoid corruption on crash + entries.extend(item for item in extra_entries if item.get("session_id") not in tracked_ids) from utils import atomic_json_write atomic_json_write(CHECKPOINT_PATH, entries) except Exception as e: logger.debug("Failed to write checkpoint file: %s", e, exc_info=True) def recover_from_checkpoint(self) -> int: - """ - On gateway startup, probe PIDs from checkpoint file. - - Returns the number of processes recovered as detached. - """ + """On gateway startup, probe PIDs from the checkpoint file; returns how many + were recovered as detached sessions.""" if not CHECKPOINT_PATH.exists(): return 0 - try: entries = json.loads(CHECKPOINT_PATH.read_text(encoding="utf-8")) except Exception: return 0 - recovered = 0 unresolved_scope_entries: List[Dict[str, Any]] = [] for entry in entries: - pid = entry.get("pid") + pid, pid_scope = entry.get("pid"), entry.get("pid_scope", "host") if not pid: continue - - pid_scope = entry.get("pid_scope", "host") - if pid_scope != "host": - # In-sandbox PIDs mean nothing once the environment handle is gone. + if pid_scope != "host": # in-sandbox PIDs mean nothing once the env handle is gone logger.info( "Skipping recovery for non-host process: %s (pid=%s, scope=%s)", - entry.get("command", "unknown")[:60], - pid, - pid_scope, - ) + entry.get("command", "unknown")[:60], pid, pid_scope) continue - # Alive AND the same process: across a restart the kernel may have # recycled the PID onto a stranger, and adopting it would let a later # kill tree-kill e.g. a browser. - recorded_start = entry.get("host_start_time") - if not self._host_pid_is_ours(pid, recorded_start): + if not self._host_pid_is_ours(pid, entry.get("host_start_time")): if self._is_host_pid_alive(pid): logger.info( "Not recovering session %s: pid %d is alive but its " "start time no longer matches — PID was recycled onto " "an unrelated process; refusing to adopt it.", - entry.get("session_id", "?"), pid, - ) + entry.get("session_id", "?"), pid) systemd_unit = entry.get("systemd_unit", "") if systemd_unit and not _stop_systemd_unit(systemd_unit): logger.warning( "Could not reap persisted scope %s for dead wrapper pid %s; " "retaining checkpoint entry for the next startup", - systemd_unit, - pid, - ) + systemd_unit, pid) unresolved_scope_entries.append(entry) continue - fields = {f: entry.get(f, _CHECKPOINT_DEFAULTS[f]) for f in _CHECKPOINT_FIELDS} fields.update( command=entry.get("command", "unknown"), owner_task_id=entry.get("owner_task_id", "") or entry.get("task_id", ""), - pid=pid, - host_start_time=recorded_start, - pid_scope=pid_scope, - started_at=entry.get("started_at", time.time()), - ) - session = ProcessSession( - id=entry["session_id"], - detached=True, # Can't read output, but can report status + kill - **fields, - ) + started_at=entry.get("started_at", time.time())) + # detached: can't read output, but can report status + kill + session = ProcessSession(id=entry["session_id"], detached=True, **fields) with self._lock: self._running[session.id] = session recovered += 1 logger.info("Recovered detached process: %s (pid=%d)", session.command[:60], pid) - # Re-enqueue watcher so gateway can resume notifications if session.watcher_interval > 0: self.pending_watchers.append({ "session_id": session.id, "check_interval": session.watcher_interval, "session_key": session.session_key, - "platform": session.watcher_platform, - "chat_id": session.watcher_chat_id, - "user_id": session.watcher_user_id, - "user_name": session.watcher_user_name, - "thread_id": session.watcher_thread_id, - "message_id": session.watcher_message_id, + **{key: getattr(session, f"watcher_{key}") for key in _WATCHER_ROUTE_KEYS}, "notify_on_complete": session.notify_on_complete, "parent_session_id": session.parent_session_id, }) - self._write_checkpoint(extra_entries=unresolved_scope_entries) - return recovered -# Module-level singleton process_registry = ProcessRegistry() -# Notification rendering lives in tools.process_registry_notifications; the names are -# re-exported here so `from tools.process_registry import format_process_notification` -# and `patch("tools.process_registry._x")` keep resolving. +# Notification rendering lives in tools.process_registry_notifications; re-exported so +# `from tools.process_registry import format_process_notification` and +# `patch("tools.process_registry._x")` keep resolving. from tools.process_registry_notifications import ( # noqa: F401,E402 - _delegation_attribution_line, - _delegation_config, - _delegation_model_not_found, - _delegation_model_not_found_notice, - _format_age, - _format_async_delegation, - _model_not_found_patterns, - format_process_notification, + _delegation_attribution_line, _delegation_config, _delegation_model_not_found, + _delegation_model_not_found_notice, _format_age, _format_async_delegation, + _model_not_found_patterns, format_process_notification, ) -# --------------------------------------------------------------------------- -# Registry -- the "process" tool schema + handler -# --------------------------------------------------------------------------- +# --- the "process_manage" tool schema + handler ----------------------------------- from tools.registry import registry, tool_error PROCESS_SCHEMA = { @@ -2514,37 +1853,30 @@ def _redact_process_result(result: dict) -> dict: """Redact secrets from background-process output before it reaches the model, session.db and CLI, mirroring the foreground ``terminal`` redaction so the two surfaces can't diverge. Respects ``security.redact_secrets``; ``redact_terminal_output`` - picks ``code_file`` from the recorded command. The command itself is redacted too. - """ + picks ``code_file`` from the recorded command. The command itself is redacted too.""" if not isinstance(result, dict): return result from agent.redact import redact_sensitive_text, redact_terminal_output command = result.get("command") or "" - for field in ("output", "output_preview"): - value = result.get(field) - if isinstance(value, str) and value: - result[field] = redact_terminal_output(value, command) - if isinstance(result.get("command"), str) and result["command"]: - result["command"] = redact_sensitive_text(result["command"], code_file=True) + for key in ("output", "output_preview"): + if isinstance(value := result.get(key), str) and value: + result[key] = redact_terminal_output(value, command) + if isinstance(command, str) and command: + result["command"] = redact_sensitive_text(command, code_file=True) return result def _list_processes(task_id) -> dict: - # Surface session-scoped background processes (e.g. a forgotten preview - # server) in addition to this task's own — they share the gateway - # session_key and can block session reset. - try: + # Also surface session-scoped background processes (e.g. a forgotten preview + # server): they share the gateway session_key and can block session reset. + session_key = "" + with suppress(Exception): from tools.approval import get_current_session_key session_key = get_current_session_key(default="") or "" - except Exception: - session_key = "" - return { - "processes": [ - _redact_process_result(p) - for p in process_registry.list_sessions(task_id=task_id, session_key=session_key or None) - ] - } + return {"processes": [ + _redact_process_result(p) + for p in process_registry.list_sessions(task_id=task_id, session_key=session_key or None)]} # action -> (handler(session_id, args) -> dict, redact output?). Output-bearing @@ -2564,7 +1896,6 @@ def _handle_process(args, **kw): action = args.get("action", "") # Coerce to string — some models send session_id as an integer session_id = str(args.get("session_id", "")) if args.get("session_id") is not None else "" - if action == "list": return json.dumps(_list_processes(kw.get("task_id")), ensure_ascii=False) if action in _SESSION_ACTIONS: diff --git a/tools/process_registry_notifications.py b/tools/process_registry_notifications.py index ccd8d2fd60..b470be43a3 100644 --- a/tools/process_registry_notifications.py +++ b/tools/process_registry_notifications.py @@ -7,6 +7,7 @@ conversation by the CLI drain loop, the gateway, and the TUI. """ import time +from contextlib import suppress def _format_age(seconds: float) -> str: @@ -19,19 +20,15 @@ def _format_age(seconds: float) -> str: return f"{s}s" m, s = divmod(s, 60) if m < 60: - return f"{m}m" if s == 0 else f"{m}m{s}s" + return f"{m}m" + (f"{s}s" if s else "") h, m = divmod(m, 60) - return f"{h}h" if m == 0 else f"{h}h{m}m" + return f"{h}h" + (f"{m}m" if m else "") def _model_not_found_patterns() -> "list[str]": - """Model-not-found phrases shared with the failover classifier. - - Imported from ``agent.error_classifier`` so the batch renderer applies the - same classification the failover path uses (no hand-copied list to drift). - Fails open to a minimal built-in set so an import problem never hides the - per-task blocks. - """ + """Model-not-found phrases from ``agent.error_classifier`` (same classification + the failover path uses, no hand-copied list to drift); a minimal built-in set + if the import fails so per-task blocks are never hidden.""" try: from agent.error_classifier import _MODEL_NOT_FOUND_PATTERNS @@ -42,11 +39,8 @@ def _model_not_found_patterns() -> "list[str]": def _delegation_config() -> dict: """Active delegation config (model/provider/fallbacks); ``{}`` on any error. - - Mirrors ``tools.delegate_tool._load_config`` lazily so the renderer sees the - same model/provider the dispatcher used without importing the heavy - delegation module at import time. - """ + Lazy ``tools.delegate_tool._load_config`` so the renderer sees the dispatcher's + model/provider without importing the heavy delegation module at import time.""" try: from tools.delegate_tool import _load_config as _cfg @@ -57,24 +51,15 @@ def _delegation_config() -> dict: def _delegation_model_not_found(results, config) -> bool: """True when a result reflects a config-level model_not_found rejection. - Requires both a model-not-found phrase AND the currently-configured model name in the same error/summary text, so a stale task failing on a - different (removed) model is not mis-attributed to the config. - """ - model = (config or {}).get("model") + different (removed) model is not mis-attributed to the config.""" + model = str((config or {}).get("model") or "").lower() if not model: return False - model = str(model).lower() - for r in results or []: - text = " ".join( - str(part) for part in (r.get("error"), r.get("summary")) if part - ).lower() - if not text or model not in text: - continue - if any(p in text for p in _model_not_found_patterns()): - return True - return False + patterns = _model_not_found_patterns() + texts = (" ".join(str(x) for x in (r.get("error"), r.get("summary")) if x).lower() for r in results or []) + return any(model in text and any(p in text for p in patterns) for text in texts) def _delegation_model_not_found_notice(results) -> "list[str] | None": @@ -89,37 +74,34 @@ def _delegation_model_not_found_notice(results) -> "list[str] | None": f'"{model}" was rejected by provider "{provider}" ' "(HTTP 400: not a valid model ID).", "Every task in this batch failed for this reason before doing any work.", - "Check Settings → Advanced → Subagent Model (or: " - "hermes config get delegation.model).", + "Check Settings → Advanced → Subagent Model (or: hermes config get delegation.model).", ] - try: + with suppress(Exception): from hermes_cli.fallback_config import get_fallback_chain if not get_fallback_chain(config): - lines.append( - "No fallback chain is configured, so no failover was attempted." - ) - except Exception: - pass + lines.append("No fallback chain is configured, so no failover was attempted.") return lines _TRUNCATED_SUMMARY_NOTE = ( "[TRUNCATED — subagent hit its iteration cap; the summary below " "may be incomplete. Verify before relying on it, or re-dispatch " - "the unfinished part.]" -) + "the unfinished part.]") def _is_truncated(entry: dict) -> bool: return bool(entry.get("truncated") or entry.get("exit_reason") == "max_iterations") -def _dispatched_line(dispatched_at, completed_at) -> "str | None": - if not isinstance(dispatched_at, (int, float)): - return None - ts = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(dispatched_at)) - return f"Dispatched: {ts} ({_format_age(completed_at - dispatched_at)} ago)" +def _header_lines(evt: dict, title: str, intro: str, completed_at: float) -> "list[str]": + """Shared preamble: title, intro, blank, dispatch time and task-source lines.""" + lines = [title, intro, ""] + dispatched_at = evt.get("dispatched_at") + if isinstance(dispatched_at, (int, float)): + ts = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(dispatched_at)) + lines.append(f"Dispatched: {ts} ({_format_age(completed_at - dispatched_at)} ago)") + return lines def _task_source_lines(evt: dict) -> "list[str]": @@ -131,6 +113,10 @@ def _task_source_lines(evt: dict) -> "list[str]": return lines +def _role_model(evt: dict) -> str: + return f"Role: {evt.get('role') or 'leaf'} Model: {evt.get('model') or '?'}" + + def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> str: """Consolidated block for a delegate_task fan-out that finished as one unit.""" results = evt.get("results") or [] @@ -138,32 +124,24 @@ def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> s n = len(results) if results else len(goals) total_dur = evt.get("total_duration_seconds", evt.get("duration_seconds", "?")) error = evt.get("error") - lines = [ + lines = _header_lines( + evt, f"[ASYNC DELEGATION BATCH COMPLETE — {deleg_id}]", f"A background fan-out of {n} subagent(s) you dispatched earlier " "has finished. All ran in parallel and waited on each other; their " "consolidated results are below. You may have moved on since " "dispatching — act on these or re-dispatch if things have changed.", - "", - ] - dispatched = _dispatched_line(evt.get("dispatched_at"), completed_at) - if dispatched: - lines.append(dispatched) + completed_at) lines.extend(_task_source_lines(evt)) - lines.append( - f"Role: {evt.get('role') or 'leaf'} Model: {evt.get('model') or '?'}" - f" Total duration: {total_dur}s" - ) + lines.append(f"{_role_model(evt)} Total duration: {total_dur}s") if error and not results: - lines.append("--- ERROR ---") - lines.append(f"The batch did not complete successfully: {error}") + lines += ["--- ERROR ---", f"The batch did not complete successfully: {error}"] return "\n".join(lines) # Config-level rejection notice BEFORE the per-task wall — a rejected # delegation model fails every task identically and must not stay buried. _notice = _delegation_model_not_found_notice(results) if _notice: - lines.append("") - lines.extend(_notice) + lines += ["", *_notice] for r in sorted(results, key=lambda x: x.get("task_index", 0)): idx = r.get("task_index", 0) r_status = r.get("status", "?") @@ -172,19 +150,14 @@ def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> s r_goal = goals[idx] if idx < len(goals) else r.get("goal", "") r_truncated = _is_truncated(r) icon = "⚠" if r_truncated else ("✓" if r_status in ("completed", "success") else "✗") - lines.append("") - header = f"--- {icon} TASK {idx + 1}/{n}" - if r_goal: - header += f": {r_goal}" - header += f" (status={r_status}" + header = f"--- {icon} TASK {idx + 1}/{n}" + (f": {r_goal}" if r_goal else "") + f" (status={r_status}" if r.get("api_calls"): header += f", api_calls={r['api_calls']}" if r.get("duration_seconds") is not None: header += f", {r['duration_seconds']}s" if r_truncated: header += ", TRUNCATED: hit max_iterations — work may be incomplete" - header += ") ---" - lines.append(header) + lines += ["", header + ") ---"] if r_status in ("completed", "success") and r_summary: if r_truncated: lines.append(_TRUNCATED_SUMMARY_NOTE) @@ -192,102 +165,77 @@ def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> s elif r_summary: if r_error: lines.append(f"({r_status}: {r_error})") - lines.append("Partial output:") - lines.append(r_summary) + lines += ["Partial output:", r_summary] else: - lines.append( - f"(no summary — status={r_status}" - + (f": {r_error}" if r_error else "") - + ")" - ) - r_live = r.get("live_transcript") - if r_live: - lines.append( - f"Full live transcript (complete tool/assistant trace): {r_live}" - ) + lines.append(f"(no summary — status={r_status}" + (f": {r_error}" if r_error else "") + ")") + if r.get("live_transcript"): + lines.append(f"Full live transcript (complete tool/assistant trace): {r['live_transcript']}") return "\n".join(lines) def _format_async_delegation(evt: dict) -> str: """Format an async-delegation completion into a self-contained re-injection. - - Carries the FULL original task source (goal, context, toolsets, role, - model) plus dispatch time, status, and the complete result summary: when - this re-enters the conversation the agent may be deep in unrelated context - and must be able to use the result OR re-dispatch without remembering why - the subagent existed. - """ + Carries the FULL original task source (goal, context, toolsets, role, model) plus + dispatch time, status, and the complete result summary: when this re-enters the + conversation the agent may be deep in unrelated context and must be able to use + the result OR re-dispatch without remembering why the subagent existed.""" deleg_id = evt.get("delegation_id", "unknown") completed_at = evt.get("completed_at") or time.time() if evt.get("is_batch") or isinstance(evt.get("results"), list): return _format_batch_delegation(evt, deleg_id, completed_at) - status = evt.get("status") or "completed" summary = evt.get("summary") error = evt.get("error") truncated = _is_truncated(evt) - lines = [ + lines = _header_lines( + evt, f"[ASYNC DELEGATION COMPLETE — {deleg_id}]", "A background subagent you dispatched earlier has finished. You may " "have moved on since dispatching it; the full task source is below so " "you can act on the result or re-dispatch if things have changed.", - "", - ] - dispatched = _dispatched_line(evt.get("dispatched_at"), completed_at) - if dispatched: - lines.append(dispatched) + completed_at) lines.append(f"Original goal: {evt.get('goal', '') or ''}") lines.extend(_task_source_lines(evt)) - lines.append(f"Role: {evt.get('role') or 'leaf'} Model: {evt.get('model') or '?'}") + lines.append(_role_model(evt)) _notice = _delegation_model_not_found_notice([evt]) if _notice: - lines.append("") - lines.extend(_notice) + lines += ["", *_notice] _trunc = " [TRUNCATED: hit max_iterations — work may be incomplete]" if truncated else "" - lines.append( + lines += [ f"Status: {status} API calls: {evt.get('api_calls', 0)} " - f"Duration: {evt.get('duration_seconds', '?')}s{_trunc}" - ) - lines.append("--- RESULT ---") + f"Duration: {evt.get('duration_seconds', '?')}s{_trunc}", + "--- RESULT ---", + ] if status in ("completed", "success") and summary: if truncated: lines.append(_TRUNCATED_SUMMARY_NOTE) lines.append(summary) else: if status == "interrupted": - lines.append( - "The subagent was interrupted before completing" - + (f": {error}" if error else ".") - ) + lines.append("The subagent was interrupted before completing" + (f": {error}" if error else ".")) else: # error / timeout / failed lines.append( - f"The subagent did not complete successfully (status={status})." - + (f"\n{error}" if error else "") + f"The subagent did not complete successfully (status={status})." + (f"\n{error}" if error else "") ) if summary: - lines.append("Partial output:") - lines.append(summary) + lines += ["Partial output:", summary] return "\n".join(lines) def _delegation_attribution_line(evt: dict) -> "str | None": """One-line provenance for a subagent-owned process event, else None. - - Subagents run terminal sessions under ``task_id == subagent_id``; a - background process they started outlives the child and is routed to the - PARENT conversation, which otherwise sees an anonymous raw output wall. + A background process a subagent started outlives the child and is routed to + the PARENT conversation, which otherwise sees an anonymous raw output wall. Judged on ``owner_task_id`` (the raw spawning id) — ``task_id`` is the - container key and may be collapsed to the session key. - """ + container key and may be collapsed to the session key.""" task_id = str(evt.get("owner_task_id") or evt.get("task_id") or "") if not task_id.startswith("sa-"): return None - try: + info = None + with suppress(Exception): from tools.delegate_tool import get_subagent_attribution info = get_subagent_attribution(task_id) - except Exception: - info = None if not info: # Registry entry aged out — still attribute generically, not anonymously. return f"Started by subagent {task_id} (delegate_task)." @@ -295,27 +243,20 @@ def _delegation_attribution_line(evt: dict) -> "str | None": if len(goal) > 120: goal = goal[:117] + "..." deleg = info.get("delegation_id") - parts = [f"Started by subagent {task_id}"] - if deleg: - parts.append(f"of delegation {deleg}") - line = " ".join(parts) + "." - if goal: - line += f' Task: "{goal}"' - return line + line = f"Started by subagent {task_id}" + (f" of delegation {deleg}" if deleg else "") + "." + return line + (f' Task: "{goal}"' if goal else "") + + +_REASON_STATUS = {"lost": "marked lost because the process backend disappeared", "failed_start": "failed to start"} def _completion_status(evt: dict) -> str: - _exit = evt.get("exit_code", "?") - _reason = evt.get("completion_reason") or "exited" - if _reason == "killed": + reason = evt.get("completion_reason") or "exited" + if reason == "killed": return f"terminated by {evt.get('termination_source') or 'Hermes'}" - if _reason == "lost": - return "marked lost because the process backend disappeared" - if _reason == "failed_start": - return "failed to start" - if _exit == 0: - return "completed normally" - return "exited" + if reason in _REASON_STATUS: + return _REASON_STATUS[reason] + return "completed normally" if evt.get("exit_code", "?") == 0 else "exited" def format_process_notification(evt: dict) -> "str | None": @@ -325,45 +266,33 @@ def format_process_notification(evt: dict) -> "str | None": _cmd = evt.get("command", "unknown") _attribution = _delegation_attribution_line(evt) - # watch_disabled and overflow events carry their own human-readable - # `message`; without this branch overflow events would fall through to the - # completion formatter as a phantom "process exited (exit code ?)". + # watch_disabled and overflow events carry their own human-readable `message`; + # otherwise overflow events would fall through to the completion formatter as a + # phantom "process exited (exit code ?)". if evt_type in ("watch_disabled", "watch_overflow_tripped", "watch_overflow_released"): return f"[IMPORTANT: {evt.get('message', '')}]" - + attribution = f"{_attribution}\n" if _attribution else "" if evt_type == "watch_match": _sup = evt.get("suppressed", 0) text = ( - f"[IMPORTANT: Background process {_sid} matched " - f"watch pattern \"{evt.get('pattern', '?')}\".\n" - ) - if _attribution: - text += f"{_attribution}\n" - text += f"Command: {_cmd}\nMatched output:\n{evt.get('output', '')}" + f"[IMPORTANT: Background process {_sid} matched watch pattern \"{evt.get('pattern', '?')}\".\n" + f"{attribution}Command: {_cmd}\nMatched output:\n{evt.get('output', '')}") if _sup: text += f"\n({_sup} earlier matches were suppressed by rate limit)" return text + "]" - if evt_type == "async_delegation": return _format_async_delegation(evt) _exit = evt.get("exit_code", "?") _out = evt.get("output", "") _signal = ", SIGTERM" if _exit in {-15, 143, "-15", "143"} else "" - text = ( - f"[IMPORTANT: Background process {_sid} {_completion_status(evt)} " - f"(exit code {_exit}{_signal}).\n" - ) - if _attribution: - text += f"{_attribution}\n" - # A subagent-owned process's full output belongs in the child's - # transcript, not as a raw wall in the parent — trim hard but keep - # enough tail to recognise failures. - if isinstance(_out, str) and len(_out) > 600: - _out = ( - "...(output trimmed — subagent-owned process; see the " - "delegation's live transcript for full output)\n" - + _out[-600:] - ) - text += f"Command: {_cmd}\nOutput:\n{_out}]" - return text + # A subagent-owned process's full output belongs in the child's transcript, not as + # a raw wall in the parent — trim hard but keep enough tail to recognise failures. + if _attribution and isinstance(_out, str) and len(_out) > 600: + _out = ( + "...(output trimmed — subagent-owned process; see the " + "delegation's live transcript for full output)\n" + + _out[-600:]) + return ( + f"[IMPORTANT: Background process {_sid} {_completion_status(evt)} (exit code {_exit}{_signal}).\n" + f"{attribution}Command: {_cmd}\nOutput:\n{_out}]") From d9d25c89ba9950188054e060df56cd72e70302a9 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:06:07 -0700 Subject: [PATCH 02/16] refactor(tools): unify skill_manage guard/ledger boilerplate; compact batch + ledger helpers --- tools/skill_ledger.py | 87 ++++++++--------------- tools/skill_manager_batch.py | 85 +++++++++++------------ tools/skill_manager_guards.py | 111 +++++++++++++++--------------- tools/skill_manager_tool.py | 126 ++++++++++++---------------------- 4 files changed, 163 insertions(+), 246 deletions(-) diff --git a/tools/skill_ledger.py b/tools/skill_ledger.py index 6509db5639..ffe617b7ee 100644 --- a/tools/skill_ledger.py +++ b/tools/skill_ledger.py @@ -3,10 +3,9 @@ Every skill mutation (any actor) appends one JSONL entry to ``~/.hermes/skills/.curator_ledger.jsonl`` with before/after file manifests whose contents are stored content-addressed (sha256-deduped) under -``~/.hermes/.curator_backups/blobs/``. JSONL, not the state DB: durable, -human-greppable, survives DB resets. The ledger is TELEMETRY, NOT A GATE: every -public write path swallows and logs. The one exception is ``rollback_entry``, -which FAILS CLOSED when its own pre-rollback safety capture fails. +``~/.hermes/.curator_backups/blobs/``. JSONL, not the state DB: durable, greppable, +survives DB resets. TELEMETRY, NOT A GATE: every public write path swallows and +logs — except ``rollback_entry``, which FAILS CLOSED when its safety capture fails. """ from __future__ import annotations @@ -28,14 +27,12 @@ from hermes_constants import get_hermes_home logger = logging.getLogger(__name__) -# Snapshot-id shape used by agent.curator_backup (duplicated so the ledger can -# read the newest skills.tar.gz without importing the backup stack). +# Snapshot-id shape of agent.curator_backup (duplicated to avoid importing the backup stack). _BACKUP_ID_RE = re.compile(r"^\d{4}-\d{2}-\d{2}T\d{2}-\d{2}-\d{2}Z(-\d{2})?$") # ".archive/-YYYYMMDDHHMMSS" collision suffix added by archive_skill. _ARCHIVE_TS_SUFFIX_RE = re.compile(r"^(.+)-\d{14}$") -# Actions whose rollback must restore a COMPLETE package: consolidation may -# have re-homed support files out of the tree first, so a disk-only capture -# would make rollback restore a hollow skill. +# Rollback of these must restore a COMPLETE package: consolidation may have re-homed +# support files first, so a disk-only capture would restore a hollow skill. _PACKAGE_RESTORE_ACTIONS = frozenset({"delete", "archive", "purge"}) _VALID_ACTORS = {"curator", "agent", "user"} _NON_PACKAGE_TOPS = {".curator_backups", ".hub", ".archive"} @@ -96,17 +93,13 @@ def _rel_posix(path: Path | str, root: Path) -> Optional[str]: """POSIX path of ``path`` relative to ``root`` (both normalized), or None when outside.""" try: return _norm(path).relative_to(_norm(root)).as_posix() - except ValueError: + except (ValueError, TypeError): return None def _is_within(root: Path, path: Path) -> bool: """True when *path* (normalized, no symlink resolution) sits under *root*.""" - try: - root_r, path_r = _norm(root), _norm(path) - return path_r == root_r or root_r in path_r.parents - except Exception: - return False + return _rel_posix(path, root) is not None def _store_blob(data: bytes) -> str: @@ -150,9 +143,7 @@ def snapshot_paths(root: Optional[Path], *, complete_package: bool = False) -> L else: return [] out = [{"path": str(f), "sha256": _store_blob(f.read_bytes())} for f in files] - if complete_package: - out = fill_snapshot_from_curator_backup(root, out) - return out + return fill_snapshot_from_curator_backup(root, out) if complete_package else out def _package_rel(root: Path) -> Optional[str]: @@ -208,9 +199,7 @@ def _latest_skills_tarball() -> Optional[Path]: def _read_package_files_from_latest_backup(prefixes: List[str]) -> Dict[str, bytes]: """``{posix-relpath: bytes}`` under *prefixes* in the newest snapshot; malicious member names (absolute, ``..`` traversal) are rejected.""" - if not prefixes: - return {} - archive = _latest_skills_tarball() + archive = _latest_skills_tarball() if prefixes else None if archive is None: return {} prefixed = tuple(p if p.endswith("/") else p + "/" for p in prefixes) @@ -219,12 +208,10 @@ def _read_package_files_from_latest_backup(prefixes: List[str]) -> Dict[str, byt try: with tarfile.open(archive, "r:gz") as tf: for member in tf.getmembers(): - if not member.isfile(): - continue name = member.name.replace("\\", "/").lstrip("./") - if not name or name.startswith("/") or ".." in Path(name).parts: - continue - if name not in exact and not name.startswith(prefixed): + if (not member.isfile() or not name or name.startswith("/") + or ".." in Path(name).parts + or (name not in exact and not name.startswith(prefixed))): continue extracted = tf.extractfile(member) if extracted is not None: @@ -293,14 +280,10 @@ def append_entry( return None try: entry = { - "id": uuid.uuid4().hex[:12], - "ts": datetime.now(timezone.utc).isoformat(), + "id": uuid.uuid4().hex[:12], "ts": datetime.now(timezone.utc).isoformat(), "actor": actor if actor in _VALID_ACTORS else derive_actor(), - "action": action, - "skill": skill, - "evidence": evidence or {}, - "before": before or [], - "after": after or []} + "action": action, "skill": skill, "evidence": evidence or {}, + "before": before or [], "after": after or []} path = ledger_path() path.parent.mkdir(parents=True, exist_ok=True) with open(path, "a", encoding="utf-8") as fh: @@ -327,9 +310,8 @@ def record_mutation( before = snapshot_paths(before_root, complete_package=_complete) elif _complete: before = fill_snapshot_from_curator_backup(before_root, before, skill=skill) - after = snapshot_paths(after_root) - return append_entry( - action, skill, before=before, after=after, actor=actor, evidence=evidence) + return append_entry(action, skill, before=before, after=snapshot_paths(after_root), + actor=actor, evidence=evidence) except Exception as e: logger.warning("skill_ledger: record_mutation failed (%s) — mutation unaffected", e) return None @@ -361,20 +343,16 @@ def list_entries(skill: Optional[str] = None, limit: Optional[int] = None) -> Li try: with open(path, "r", encoding="utf-8") as fh: for line in fh: - try: + with suppress(json.JSONDecodeError): row = json.loads(line) if line.strip() else None - except json.JSONDecodeError: - continue - if isinstance(row, dict): - rows.append(row) + if isinstance(row, dict): + rows.append(row) except OSError: return [] if skill: rows = [r for r in rows if r.get("skill") == skill] rows.reverse() - if limit is not None and limit >= 0: - rows = rows[:limit] - return rows + return rows[:limit] if limit is not None and limit >= 0 else rows def get_entry(entry_id: str) -> Optional[Dict[str, Any]]: @@ -403,14 +381,11 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]: entry = get_entry(entry_id) if entry is None: return False, f"no ledger entry with id '{entry_id}'" - - path_err = _validate_entry_paths(entry) - if path_err: + if path_err := _validate_entry_paths(entry): return False, f"refusing rollback: {path_err}" before = list(entry.get("before") or []) after = list(entry.get("after") or []) - # Historical hollow delete/archive/purge entries (SKILL.md only): fill from the # newest curator backup so the complete package is restored. Entry hashes win; # only missing paths are added, and the filled set is re-validated. @@ -418,10 +393,8 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]: before = fill_snapshot_from_curator_backup( next(iter(_skill_md_parents(before)), None), before, skill=str(entry.get("skill") or "") or None) - path_err = _validate_entry_paths({**entry, "before": before, "after": after}) - if path_err: + if path_err := _validate_entry_paths({**entry, "before": before, "after": after}): return False, f"refusing rollback: {path_err}" - # Pre-check every blob we need so we never fail mid-restore. for item in before: if read_blob(str(item.get("sha256", ""))) is None: @@ -431,11 +404,8 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]: # Safety entry: CURRENT state of every touched path, so the rollback itself is undoable. touched = {str(i["path"]) for i in before + after if i.get("path")} try: - safety_before: List[Dict[str, str]] = [] - for p in sorted(touched): - fp = Path(p) - if fp.is_file(): - safety_before.append({"path": p, "sha256": _store_blob(fp.read_bytes())}) + safety_before = [{"path": p, "sha256": _store_blob(Path(p).read_bytes())} + for p in sorted(touched) if Path(p).is_file()] safety_id = append_entry( "pre-rollback", entry.get("skill", "?"), before=safety_before, after=safety_before, evidence={"rollback_target": entry_id}) @@ -456,10 +426,9 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]: for item in after: p = str(item.get("path", "")) if p and p not in before_paths: - fp = Path(p) try: - if fp.is_file(): - fp.unlink() + if Path(p).is_file(): + Path(p).unlink() removed += 1 except OSError as e: logger.warning("skill_ledger: could not remove %s during rollback: %s", p, e) diff --git a/tools/skill_manager_batch.py b/tools/skill_manager_batch.py index 5d1e863239..c0a2f4eb8c 100644 --- a/tools/skill_manager_batch.py +++ b/tools/skill_manager_batch.py @@ -19,24 +19,23 @@ def _validate_batch_ops(operations, default_name, tool_error): """Shape checks with no side effects. Returns (names, None) or (None, error_json).""" from tools.skill_manager_guards import _background_review_preflight + def fail(i, msg): + return None, tool_error(f"operations[{i}]{msg}", success=False) + names = [] for i, op in enumerate(operations): if not isinstance(op, dict) or not op.get("action"): - return None, tool_error(f"operations[{i}] needs an 'action'.", success=False) + return fail(i, " needs an 'action'.") act = op["action"] if act not in _BATCH_OP_ACTIONS: - return None, tool_error( - f"operations[{i}]: unknown action '{act}'. Batchable: " - f"{', '.join(sorted(_BATCH_OP_ACTIONS))}; delete must be sole.", - success=False) + return fail(i, f": unknown action '{act}'. Batchable: " + f"{', '.join(sorted(_BATCH_OP_ACTIONS))}; delete must be sole.") nm = op.get("name") or default_name if not nm: - return None, tool_error(f"operations[{i}] needs a 'name' (the skill it targets).", success=False) + return fail(i, " needs a 'name' (the skill it targets).") names.append(nm) if act == "create" and nm in names[:-1]: - return None, tool_error( - f"operations[{i}]: create for '{nm}' must precede that skill's other ops.", - success=False) + return fail(i, f": create for '{nm}' must precede that skill's other ops.") preflight = _background_review_preflight(act, nm) if preflight is not None: return None, json.dumps(preflight, ensure_ascii=False) @@ -46,22 +45,18 @@ def _validate_batch_ops(operations, default_name, tool_error): # Additive patches are always legal. Paths are normalized against spelling variants. touched_files = set() for i, op in enumerate(operations): - act = op["action"] - nm = names[i] + act, nm = op["action"], names[i] # create and full-rewrite patch (content) always hit SKILL.md. full_rewrite = act == "patch" and bool(op.get("content")) fp = (op.get("file_path") or "").strip() target = ("SKILL.md" if (act == "create" or full_rewrite or not fp) else posixpath.normpath(fp.lstrip("/"))) key = (nm, target) - destructive = act in ("create", "write_file", "remove_file") or full_rewrite - if destructive and key in touched_files: - return None, tool_error( - f"operations[{i}]: {act} on '{target}' of skill '{nm}' — an earlier op in this " - f"batch already touched that file, and this op would silently discard its work. " - f"One destructive op (write_file/remove_file/full rewrite) per file per batch; put " - f"it first, or fold the change in. Patch chains are fine.", - success=False) + if (act in ("create", "write_file", "remove_file") or full_rewrite) and key in touched_files: + return fail(i, f": {act} on '{target}' of skill '{nm}' — an earlier op in this " + f"batch already touched that file, and this op would silently discard its work. " + f"One destructive op (write_file/remove_file/full rewrite) per file per batch; put " + f"it first, or fold the change in. Patch chains are fine.") touched_files.add(key) return names, None @@ -84,26 +79,27 @@ def _snapshot_skills(names, snap_root, find_skill): def _restore_snapshot(pre_dir, snap, post_dir) -> None: - if snap is not None: - if post_dir is not None and post_dir.is_dir(): - # Move the broken state aside and delete it only after the snapshot is - # back, so a failed copytree (disk full, locked file) can't mean total loss. - aside = post_dir.with_name(post_dir.name + ".rollback-broken") - shutil.rmtree(aside, ignore_errors=True) - post_dir.rename(aside) - try: - shutil.copytree(snap, pre_dir) - except Exception: - # Restore failed: put the half-applied state back rather than nothing. - shutil.rmtree(pre_dir, ignore_errors=True) - aside.rename(pre_dir) - raise - shutil.rmtree(aside, ignore_errors=True) - else: - shutil.copytree(snap, pre_dir) - elif post_dir is not None and post_dir.is_dir(): - # Batch created this skill: remove the partial result. - shutil.rmtree(post_dir) + post_exists = post_dir is not None and post_dir.is_dir() + if snap is None: + if post_exists: # Batch created this skill: remove the partial result. + shutil.rmtree(post_dir) + return + if not post_exists: + shutil.copytree(snap, pre_dir) + return + # Move the broken state aside and delete it only after the snapshot is + # back, so a failed copytree (disk full, locked file) can't mean total loss. + aside = post_dir.with_name(post_dir.name + ".rollback-broken") + shutil.rmtree(aside, ignore_errors=True) + post_dir.rename(aside) + try: + shutil.copytree(snap, pre_dir) + except Exception: + # Restore failed: put the half-applied state back rather than nothing. + shutil.rmtree(pre_dir, ignore_errors=True) + aside.rename(pre_dir) + raise + shutil.rmtree(aside, ignore_errors=True) def _rollback(snapshots, find_skill): @@ -114,9 +110,8 @@ def _rollback(snapshots, find_skill): post = find_skill(nm) _restore_snapshot(pre_dir, snap, Path(post["path"]) if post else None) except Exception as exc: # noqa: BLE001 - notes.append( - f"ROLLBACK FAILED for '{nm}' ({exc}); snapshot preserved at '{snap}'" - if snap is not None else f"ROLLBACK FAILED for '{nm}' ({exc})") + notes.append(f"ROLLBACK FAILED for '{nm}' ({exc})" + + (f"; snapshot preserved at '{snap}'" if snap is not None else "")) return ("; ".join(notes) if notes else "all touched skills rolled back"), bool(notes) @@ -136,10 +131,8 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non return tool_error(f"operations is capped at {_BATCH_MAX_OPS} ops per call.", success=False) if any(isinstance(op, dict) and op.get("action") == "delete" for op in operations): if len(operations) != 1: - return tool_error( - "delete must be the SOLE op in its call — it doesn't " - "compose with other ops' rollback.", - success=False) + return tool_error("delete must be the SOLE op in its call — it doesn't " + "compose with other ops' rollback.", success=False) op = operations[0] nm = op.get("name") or default_name if not nm: diff --git a/tools/skill_manager_guards.py b/tools/skill_manager_guards.py index ad616e8253..41e8a02b39 100644 --- a/tools/skill_manager_guards.py +++ b/tools/skill_manager_guards.py @@ -8,6 +8,7 @@ reached lazily through ``tools.skill_manager_tool`` so test patches keep working import contextvars as _ctxvars import logging import threading +from contextlib import suppress from pathlib import Path from typing import Any, Dict, Optional @@ -77,22 +78,28 @@ def _reset_background_review_read_marks() -> None: _background_review_read_paths.set(_BackgroundReviewReadMarks()) -def _containing_skills_root(skill_path: Path) -> Path: - """Skills root (local or external_dirs) containing ``skill_path``; local dir if none match.""" +def _resolved_roots(skill_path: Path): + """``(resolved skill_path, [(root, resolved_root), ...])`` over every skills root + whose resolve() succeeds; an unresolvable skill_path is used as-is.""" from agent.skill_utils import get_all_skills_dirs - from tools import skill_manager_tool as _smt try: resolved = skill_path.resolve() except OSError: resolved = skill_path + roots = [] for root in get_all_skills_dirs(): - try: - if resolved.is_relative_to(root.resolve()): - return root - except OSError: - continue - return _smt._skills_dir() + with suppress(OSError): + roots.append((root, root.resolve())) + return resolved, roots + + +def _containing_skills_root(skill_path: Path) -> Path: + """Skills root (local or external_dirs) containing ``skill_path``; local dir if none match.""" + from tools import skill_manager_tool as _smt + + resolved, roots = _resolved_roots(skill_path) + return next((root for root, r in roots if resolved.is_relative_to(r)), _smt._skills_dir()) def _is_path_redirect(path: Path) -> bool: @@ -108,22 +115,16 @@ def _validate_delete_target(skill_dir: Path) -> Optional[str]: """Last-line guard before ``shutil.rmtree(skill_dir)``: even a poisoned tree must never delete (1) a path outside every known skills root, (2) a skills root itself, or (3) a symlink/junction (rmtree would follow it).""" - from agent.skill_utils import get_all_skills_dirs - if _is_path_redirect(skill_dir): return ( f"Refusing to delete '{skill_dir}': the skill directory is a " f"symlink/junction. Remove the link target manually if intended.") try: - resolved = skill_dir.resolve() + skill_dir.resolve() except OSError as exc: return f"Refusing to delete '{skill_dir}': could not resolve path ({exc})." - - for root in get_all_skills_dirs(): - try: - root = root.resolve() - except OSError: - continue + resolved, roots = _resolved_roots(skill_dir) + for _root, root in roots: if resolved == root: return ( f"Refusing to delete '{skill_dir}': resolves to the skills root " @@ -135,6 +136,16 @@ def _validate_delete_target(skill_dir: Path) -> Optional[str]: f"known skills root.") +def _is_pinned(name: str, what: str) -> Optional[bool]: + """skill_usage pinned flag; None (logged at debug) when the record is unreadable.""" + try: + from tools import skill_usage + return bool(skill_usage.get_record(name).get("pinned")) + except Exception: + logger.debug("%s lookup failed for %s", what, name, exc_info=True) + return None + + def _pinned_guard(name: str) -> Optional[str]: """Refusal message if *name* is pinned or essential, else None. @@ -150,15 +161,11 @@ def _pinned_guard(name: str) -> Optional[str]: f"cannot be deleted. Patches and edits are still allowed.") except Exception: logger.debug("essential-guard lookup failed for %s", name, exc_info=True) - try: - from tools import skill_usage - if skill_usage.get_record(name).get("pinned"): - return ( - f"Skill '{name}' is pinned and cannot be deleted by skill_manage. Ask the user to " - f"run `hermes curator unpin {name}` if they want to delete it. Patches and edits " - f"are allowed on pinned skills; only deletion is blocked.") - except Exception: - logger.debug("pinned-guard lookup failed for %s", name, exc_info=True) + if _is_pinned(name, "pinned-guard"): + return ( + f"Skill '{name}' is pinned and cannot be deleted by skill_manage. Ask the user to " + f"run `hermes curator unpin {name}` if they want to delete it. Patches and edits " + f"are allowed on pinned skills; only deletion is blocked.") return None @@ -170,22 +177,17 @@ def _background_review_write_guard( agents it is also blocked on pinned/external/bundled/hub skills.""" if not _is_background_review(): return None - - try: - from tools import skill_usage - if skill_usage.get_record(name).get("pinned"): - return _refusal( - f"Refusing background curator {action} for pinned skill '{name}': pinned skills " - f"are off-limits to autonomous maintenance. Ask the user to run `hermes curator " - f"unpin {name}` if they want it changed.") - except Exception: - logger.debug("pinned skill guard lookup failed for %s", name, exc_info=True) - + refuse = f"Refusing background curator {action} for" + if _is_pinned(name, "pinned skill guard"): + return _refusal( + f"{refuse} pinned skill '{name}': pinned skills " + f"are off-limits to autonomous maintenance. Ask the user to run `hermes curator " + f"unpin {name}` if they want it changed.") try: from agent.skill_utils import is_external_skill_path if is_external_skill_path(skill_dir): return _refusal( - f"Refusing background curator {action} for skill '{name}': " + f"{refuse} skill '{name}': " "the skill lives in skills.external_dirs, which are " "externally owned and read-only to autonomous curation.") except Exception: @@ -198,8 +200,7 @@ def _background_review_write_guard( (skill_usage.is_hub_installed, "hub-installed"), (skill_usage.is_bundled, "bundled")): if predicate(name): - return _refusal( - f"Refusing background curator {action} for {label} skill '{name}'.") + return _refusal(f"{refuse} {label} skill '{name}'.") # Not curator-managed (no `created_by: "agent"`) => user-owned. A MISSING # record and an explicit `created_by: null` must resolve IDENTICALLY (keying # on presence made the policy depend on the guard's own side effect: the @@ -209,13 +210,13 @@ def _background_review_write_guard( _detail = (f"created_by={usage_rec.get('created_by')!r}" if isinstance(usage_rec, dict) else "no usage record") return _refusal( - f"Refusing background curator {action} for skill '{name}': the skill is not " + f"{refuse} skill '{name}': the skill is not " f"curator-managed ({_detail}). User-owned skills are off-limits to autonomous " f"curation. Run `hermes curator adopt {name}` to opt it in.") except Exception: logger.warning("owned skill guard lookup failed for %s", name, exc_info=True) return _refusal( - f"Refusing background curator {action} for skill '{name}': agent ownership could not " + f"{refuse} skill '{name}': agent ownership could not " f"be verified because the provenance record is unavailable or unreadable.") return None @@ -249,31 +250,31 @@ def _curator_consolidation_delete_guard( via ``absorbed_into=`` (existence validated in ``_delete_skill``). The deterministic inactivity prune never calls ``skill_manage``, so a bare delete here can only be the LLM pass pruning without evidence: refuse it.""" - if not _is_background_review(): - return None - if isinstance(absorbed_into, str) and absorbed_into.strip(): + if not _is_background_review() or (isinstance(absorbed_into, str) and absorbed_into.strip()): return None return _refusal( f"Refusing background curator delete of skill '{name}': the consolidation pass may only " f"archive a skill it has absorbed into an umbrella. Pass absorbed_into= (the " f"umbrella must already exist) to record a verified consolidation. Pruning a skill with no " f"forwarding target is not permitted here — the deterministic inactivity prune handles " - f"staleness archival " - "separately. Keeping '{name}' active.".format(name=name), + f"staleness archival separately. Keeping '{name}' active.", _fail_closed=True) +def _is_org_mirror(skill_path: Path) -> bool: + from agent.skill_utils import is_org_mirror_path + from tools import skill_manager_tool as _smt + return is_org_mirror_path(skill_path, _smt._skills_dir()) + + def _maybe_auto_propose_org_edit(name: str, skill_path: Path) -> Optional[str]: """Submit an org-skill edit upstream when `sync.org_auto_propose` is on. Returns a note for the tool result or None; never raises (the edit is already saved locally and can be proposed later).""" - from tools import skill_manager_tool as _smt - try: - from agent.skill_utils import is_org_mirror_path from tools import skills_sync_client as ssc - if not is_org_mirror_path(skill_path, _smt._skills_dir()): + if not _is_org_mirror(skill_path): return None if not ssc.sync_org_auto_propose(): return ( @@ -302,12 +303,8 @@ def _org_mirror_write_guard(name: str, skill_path: Path, action: str) -> Optiona comes back, and removing for everyone is an admin action.""" if action not in {"delete", "remove_file"}: return None - from tools import skill_manager_tool as _smt - try: - from agent.skill_utils import is_org_mirror_path - - if is_org_mirror_path(skill_path, _smt._skills_dir()): + if _is_org_mirror(skill_path): return _refusal( f"Cannot {action} '{name}' locally: it is shared by your organisation, so a local " f"delete would just come back on the next sync. Ask an org admin to remove it for " diff --git a/tools/skill_manager_tool.py b/tools/skill_manager_tool.py index b72acaa9d6..e1c45f3f49 100644 --- a/tools/skill_manager_tool.py +++ b/tools/skill_manager_tool.py @@ -34,7 +34,7 @@ from tools.skill_manager_guards import ( # noqa: F401 — re-exported for calle _background_review_write_guard, _containing_skills_root, _curator_consolidation_delete_guard, _is_path_redirect, _maybe_auto_propose_org_edit, _org_mirror_write_guard, _pinned_guard, _reset_background_review_read_marks, _validate_delete_target, _is_background_review, - mark_background_review_skill_read) + mark_background_review_skill_read, _refusal as _err) from tools.skill_manager_batch import ( # noqa: F401 _BATCH_MAX_OPS, _BATCH_OP_ACTIONS, _skill_manage_batch) from tools.skills_guard import scan_skill, should_allow_install, format_scan_report @@ -104,10 +104,6 @@ def _display_create_dir() -> str: return f"{display_hermes_home()}/skills/" -def _err(message: str) -> Dict[str, Any]: - return {"success": False, "error": message} - - # --- Validation helpers ------------------------------------------------------- def _validate_name(name: str) -> Optional[str]: @@ -157,10 +153,9 @@ def _validate_frontmatter(content: str, *, new_skill: bool = False) -> Optional[ return f"YAML frontmatter parse error: {e}" if not isinstance(parsed, dict): return "Frontmatter must be a YAML mapping (key: value pairs)." - if "name" not in parsed: - return "Frontmatter must include 'name' field." - if "description" not in parsed: - return "Frontmatter must include 'description' field." + for field in ("name", "description"): + if field not in parsed: + return f"Frontmatter must include '{field}' field." desc = str(parsed["description"]) if len(desc) > MAX_DESCRIPTION_LENGTH: return f"Description exceeds {MAX_DESCRIPTION_LENGTH} characters." @@ -199,9 +194,7 @@ def _resolve_skill_dir(name: str, category: str = None) -> Path: base = _skills_dir() try: from agent.skill_utils import get_skill_create_dir - create_dir = get_skill_create_dir() - if create_dir is not None: - base = create_dir + base = get_skill_create_dir() or base except Exception: logger.debug("skills.create_dir lookup failed", exc_info=True) return base / category / name if category else base / name @@ -267,14 +260,12 @@ def _find_skill_in_other_profiles(name: str) -> List[Tuple[str, Path]]: if profiles_root.is_dir(): candidates += [(e.name, e / "skills") for e in profiles_root.iterdir() if e.is_dir()] for profile_name, skills_dir in candidates: - try: + with suppress(OSError, RuntimeError): if skills_dir.resolve() == active_dir or not skills_dir.is_dir(): continue hit = next((d for d in _iter_skill_dirs(skills_dir) if d.name == name), None) if hit is not None: matches.append((profile_name, hit)) # one match per profile is enough - except (OSError, RuntimeError): - continue return matches @@ -325,28 +316,22 @@ def _resolve_supporting_file(skill_dir: Path, file_path: str): -> ``(target, None)`` | ``(None, error_dict)``.""" from tools.path_security import validate_within_dir - err = _validate_file_path(file_path) - if err: - return None, _err(err) - target = skill_dir / file_path - err = validate_within_dir(target, skill_dir) - if err: - return None, _err(err) - return target, None + target = skill_dir / (file_path or "") + err = _validate_file_path(file_path) or validate_within_dir(target, skill_dir) + return (None, _err(err)) if err else (target, None) -def _locate_for_write(name: str, action: str, not_found_suffix: str = ""): - """Find the skill and run the org-mirror + background-review write guards - -> ``(skill_dir, None)`` | ``(None, error_dict)``.""" +def _locate_for_write(name: str, action: str, not_found_suffix: str = "", *, + org_guard: bool = True): + """Find the skill and run the org-mirror (unless ``org_guard=False``) + + background-review write guards -> ``(skill_dir, None)`` | ``(None, error_dict)``.""" existing = _find_skill(name) if not existing: return None, _err(_skill_not_found_error(name, not_found_suffix)) skill_dir = existing["path"] - guard = (_org_mirror_write_guard(name, skill_dir, action) + guard = ((org_guard and _org_mirror_write_guard(name, skill_dir, action)) or _background_review_write_guard(name, skill_dir, action)) - if guard: - return None, guard - return skill_dir, None + return (None, guard) if guard else (skill_dir, None) def _guarded_write(name: str, skill_dir: Path, target: Path, action: str, label: str, @@ -356,8 +341,7 @@ def _guarded_write(name: str, skill_dir: Path, target: Path, action: str, label: Returns an error dict or None.""" original = None if target.exists(): - read_guard = _background_review_read_before_write_guard(name, target, action, label) - if read_guard: + if read_guard := _background_review_read_before_write_guard(name, target, action, label): return read_guard original = target.read_text(encoding="utf-8") target.parent.mkdir(parents=True, exist_ok=True) @@ -373,8 +357,7 @@ def _guarded_write(name: str, skill_dir: Path, target: Path, action: str, label: def _attach_org_note(result: Dict[str, Any], name: str, skill_dir: Path) -> None: - org_note = _maybe_auto_propose_org_edit(name, skill_dir) - if org_note: + if org_note := _maybe_auto_propose_org_edit(name, skill_dir): result["org_sharing"] = org_note result["message"] = f"{result['message']} {org_note}" @@ -403,6 +386,10 @@ def _attach_lint_findings(result: Dict[str, Any], skill_md: Path) -> None: "— fix them with skill_manage(action='patch') to match Hermes skill standards.") +def _clip(text: str, n: int, ellipsis: str) -> str: + return text[:n] + (ellipsis if len(text) > n else "") + + # --- Core actions ------------------------------------------------------------- def _create_skill(name: str, content: str, category: str = None) -> Dict[str, Any]: @@ -410,8 +397,7 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An or _validate_frontmatter(content, new_skill=True) or _validate_content_size(content)) if err: return _err(err) - existing = _find_skill(name) - if existing: + if existing := _find_skill(name): return _err(f"A skill named '{name}' already exists at {existing['path']}.") skill_dir = _resolve_skill_dir(name, category) @@ -427,8 +413,7 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An display = skill_dir.relative_to(root) if skill_dir.is_relative_to(root) else skill_dir result = { "success": True, "message": f"Skill '{name}' created.", "path": str(display), - "skill_md": str(skill_md), "_change": {"description": _description_preview(content)}, - } + "skill_md": str(skill_md), "_change": {"description": _description_preview(content)}} if category: result["category"] = category result["hint"] = ( @@ -442,15 +427,12 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An def _edit_skill(name: str, content: str) -> Dict[str, Any]: """Replace the SKILL.md of any existing skill (full rewrite).""" - err = _validate_frontmatter(content) or _validate_content_size(content) - if err: + if err := _validate_frontmatter(content) or _validate_content_size(content): return _err(err) skill_dir, guard = _locate_for_write(name, "edit") - if guard: - return guard # SKILL.md always exists here (_find_skill requires it), so a blocked scan restores it. - guard = _guarded_write(name, skill_dir, skill_dir / "SKILL.md", "edit", "SKILL.md", content) - if guard: + if guard := guard or _guarded_write( + name, skill_dir, skill_dir / "SKILL.md", "edit", "SKILL.md", content): return guard result = { "success": True, "message": f"Skill '{name}' updated (full rewrite).", @@ -480,7 +462,6 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N skill_dir, guard = _locate_for_write(name, "patch") if guard: return guard - if file_path: target, err = _resolve_supporting_file(skill_dir, file_path) if err: @@ -490,8 +471,7 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N if not target.exists(): return _err(f"File not found: {target.relative_to(skill_dir)}") target_label = file_path or "SKILL.md" - read_guard = _background_review_read_before_write_guard(name, target, "patch", target_label) - if read_guard: + if read_guard := _background_review_read_before_write_guard(name, target, "patch", target_label): return read_guard content = target.read_text(encoding="utf-8") @@ -504,24 +484,18 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N with suppress(Exception): from tools.fuzzy_match import format_no_match_hint match_error += format_no_match_hint(match_error, match_count, old_string, content) - return _err(match_error) | {"file_preview": content[:500] + ("..." if len(content) > 500 else "")} + return _err(match_error) | {"file_preview": _clip(content, 500, "...")} - err = _validate_content_size(new_content, label=target_label) - if err: + if err := _validate_content_size(new_content, label=target_label): return _err(err) - if not file_path: - err = _validate_frontmatter(new_content) - if err: - return _err(f"Patch would break SKILL.md structure: {err}") - - guard = _guarded_write(name, skill_dir, target, "patch", target_label, new_content) - if guard: + if not file_path and (err := _validate_frontmatter(new_content)): + return _err(f"Patch would break SKILL.md structure: {err}") + if guard := _guarded_write(name, skill_dir, target, "patch", target_label, new_content): return guard result = { "success": True, "message": f"Patched {target_label} in skill '{name}' ({match_count} replacement{'s' if match_count > 1 else ''}).", - "_change": {"old": old_string[:200] + ("…" if len(old_string) > 200 else ""), - "new": new_string[:200] + ("…" if len(new_string) > 200 else "")}} + "_change": {"old": _clip(old_string, 200, "…"), "new": _clip(new_string, 200, "…")}} _attach_org_note(result, name, skill_dir) return result @@ -531,13 +505,9 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A explicit prune; "" = absorbed into that umbrella, which must exist on disk (validated here so the model can't claim a nonexistent umbrella).""" skill_dir, guard = _locate_for_write(name, "delete") - if guard: + if guard := guard or _curator_consolidation_delete_guard(name, absorbed_into): return guard - fail_closed = _curator_consolidation_delete_guard(name, absorbed_into) - if fail_closed: - return fail_closed - pinned_err = _pinned_guard(name) - if pinned_err: + if pinned_err := _pinned_guard(name): return _err(pinned_err) absorbed_target = absorbed_into.strip() if isinstance(absorbed_into, str) else "" @@ -549,8 +519,7 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A f"Create or patch the umbrella skill first, then retry the delete.") skills_root = _containing_skills_root(skill_dir) - unsafe = _validate_delete_target(skill_dir) # defense-in-depth before rmtree - if unsafe: + if unsafe := _validate_delete_target(skill_dir): # defense-in-depth before rmtree return _err(unsafe) # Curator consolidations must be RECOVERABLE (`hermes curator restore`): archive @@ -580,8 +549,7 @@ def _rmdir_if_empty(parent: Path, stop: Path) -> None: def _write_file(name: str, file_path: str, file_content: str) -> Dict[str, Any]: """Add or overwrite a supporting file within any skill directory.""" - err = _validate_file_path(file_path) - if err: + if err := _validate_file_path(file_path): return _err(err) if not file_content and file_content != "": return _err("file_content is required.") @@ -589,18 +557,14 @@ def _write_file(name: str, file_path: str, file_content: str) -> Dict[str, Any]: if content_bytes > MAX_SKILL_FILE_BYTES: return _err(f"File content is {content_bytes:,} bytes (limit: {MAX_SKILL_FILE_BYTES:,} " f"bytes / 1 MiB). Consider splitting into smaller files.") - err = _validate_content_size(file_content, label=file_path) - if err: + if err := _validate_content_size(file_content, label=file_path): return _err(err) skill_dir, guard = _locate_for_write(name, "write_file", " Create it first with action='create'.") if guard: return guard target, err = _resolve_supporting_file(skill_dir, file_path) - if err: - return err - guard = _guarded_write(name, skill_dir, target, "write_file", file_path, file_content) - if guard: + if guard := err or _guarded_write(name, skill_dir, target, "write_file", file_path, file_content): return guard result = {"success": True, "message": f"File '{file_path}' written to skill '{name}'.", "path": str(target)} @@ -610,14 +574,9 @@ def _write_file(name: str, file_path: str, file_content: str) -> Dict[str, Any]: def _remove_file(name: str, file_path: str) -> Dict[str, Any]: """Remove a supporting file from any skill directory.""" - err = _validate_file_path(file_path) - if err: + if err := _validate_file_path(file_path): return _err(err) - existing = _find_skill(name) - if not existing: - return _err(_skill_not_found_error(name)) - skill_dir = existing["path"] - guard = _background_review_write_guard(name, skill_dir, "remove_file") + skill_dir, guard = _locate_for_write(name, "remove_file", org_guard=False) if guard: return guard target, err = _resolve_supporting_file(skill_dir, file_path) @@ -630,8 +589,7 @@ def _remove_file(name: str, file_path: str) -> Dict[str, Any]: for f in (skill_dir / subdir).rglob("*") if f.is_file()] return {"success": False, "error": f"File '{file_path}' not found in skill '{name}'.", "available_files": available if available else None} - read_guard = _background_review_read_before_write_guard(name, target, "remove_file", file_path) - if read_guard: + if read_guard := _background_review_read_before_write_guard(name, target, "remove_file", file_path): return read_guard target.unlink() From eeb86e4b6ff8ddaf4c2881e8059a6136218d2b5b Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:17:57 -0700 Subject: [PATCH 03/16] =?UTF-8?q?refactor(tools):=20skill=5Fmanage=20?= =?UTF-8?q?=E2=80=94=20fold=20gate/flat-op=20plumbing,=20shared=20root/pin?= =?UTF-8?q?/org=20helpers=20in=20guards,=20ledger=20tarball=20lookup=20inl?= =?UTF-8?q?ine;=20compact=20docstrings?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/skill_ledger.py | 77 +++++++--------- tools/skill_manager_batch.py | 2 +- tools/skill_manager_guards.py | 67 ++++++-------- tools/skill_manager_tool.py | 163 ++++++++++++++-------------------- 4 files changed, 124 insertions(+), 185 deletions(-) diff --git a/tools/skill_ledger.py b/tools/skill_ledger.py index ffe617b7ee..8b4512ee4f 100644 --- a/tools/skill_ledger.py +++ b/tools/skill_ledger.py @@ -126,11 +126,10 @@ def read_blob(sha256: str) -> Optional[bytes]: def snapshot_paths(root: Optional[Path], *, complete_package: bool = False) -> List[Dict[str, str]]: - """{path, sha256} for every file under *root*, each stored as a blob. - - Empty when root is None/missing. Raises on I/O failure — callers decide whether - that is fatal (rollback safety capture) or swallowed (telemetry). - ``complete_package`` unions in the newest curator tarball's files (disk hashes win).""" + """{path, sha256} for every file under *root*, each stored as a blob; [] when root is + None/missing. Raises on I/O failure — callers decide whether that is fatal (rollback safety + capture) or swallowed (telemetry). ``complete_package`` unions in the newest curator + tarball's files (disk hashes win).""" if root is None: return [] root = Path(root) @@ -174,34 +173,24 @@ def package_prefixes( candidates = [_package_rel(Path(root)) if root is not None else None] candidates += [_package_rel(p) for p in _skill_md_parents(before)] candidates += [skill, _strip_archive_timestamp(skill) if skill else None] - found: List[str] = [] - for prefix in candidates: - prefix = (prefix or "").strip("/") - if prefix and prefix not in found: - found.append(prefix) - return found - - -def _latest_skills_tarball() -> Optional[Path]: - """Newest ``skills.tar.gz`` under ``skills/.curator_backups/``.""" - backups = _skills_dir() / ".curator_backups" - try: - children = list(backups.iterdir()) if backups.is_dir() else [] - except OSError: - return None - candidates = [ - child / "skills.tar.gz" for child in children - if child.is_dir() and _BACKUP_ID_RE.match(child.name) and (child / "skills.tar.gz").is_file()] - # Parent dirs sort lexicographically == chronologically for the id shape. - return max(candidates, key=lambda p: p.parent.name) if candidates else None + return list(dict.fromkeys(p for p in ((c or "").strip("/") for c in candidates) if p)) def _read_package_files_from_latest_backup(prefixes: List[str]) -> Dict[str, bytes]: - """``{posix-relpath: bytes}`` under *prefixes* in the newest snapshot; malicious - member names (absolute, ``..`` traversal) are rejected.""" - archive = _latest_skills_tarball() if prefixes else None - if archive is None: + """``{posix-relpath: bytes}`` under *prefixes* in the newest ``skills/.curator_backups/*/ + skills.tar.gz``; malicious member names (absolute, ``..`` traversal) are rejected.""" + backups = _skills_dir() / ".curator_backups" + try: + children = list(backups.iterdir()) if prefixes and backups.is_dir() else [] + except OSError: return {} + candidates = [ + child / "skills.tar.gz" for child in children + if child.is_dir() and _BACKUP_ID_RE.match(child.name) and (child / "skills.tar.gz").is_file()] + if not candidates: + return {} + # Parent dirs sort lexicographically == chronologically for the id shape. + archive = max(candidates, key=lambda p: p.parent.name) prefixed = tuple(p if p.endswith("/") else p + "/" for p in prefixes) exact = set(prefixes) out: Dict[str, bytes] = {} @@ -225,14 +214,12 @@ def _read_package_files_from_latest_backup(prefixes: List[str]) -> Dict[str, byt def fill_snapshot_from_curator_backup( root: Optional[Path], existing: Optional[List[Dict[str, str]]] = None, *, skill: Optional[str] = None) -> List[Dict[str, str]]: - """Union missing skill-package files from the newest curator snapshot. - - Completeness fill, not a gate: failures return *existing* unchanged, and only - ABSENT paths are filled. Fill targets go where rollback must restore them: - under *root* when known (for purge that is ``.archive//``, NOT the live - tree), else the live skills dir; the tar's leading package-dir segment is - stripped when *root* already names the package. Every target must stay under - ``skills/`` and HERMES_HOME.""" + """Union missing skill-package files from the newest curator snapshot. Completeness fill, not + a gate: failures return *existing* unchanged, and only ABSENT paths are filled. Fill targets go + where rollback must restore them: under *root* when known (for purge that is + ``.archive//``, NOT the live tree), else the live skills dir; the tar's leading + package-dir segment is stripped when *root* already names the package. Every target must stay + under ``skills/`` and HERMES_HOME.""" out = list(existing or []) prefixes = package_prefixes(root, skill, out) if not prefixes: @@ -247,9 +234,7 @@ def fill_snapshot_from_curator_backup( skills = _skills_dir() dest_root = Path(root) if root is not None else None pkg_names = {dest_root.name, _strip_archive_timestamp(dest_root.name)} if dest_root else set() - have = { - rel for rel in (_rel_posix(str(item.get("path", "")), skills) for item in out) if rel is not None - } + have = {rel for rel in (_rel_posix(str(i.get("path", "")), skills) for i in out) if rel is not None} for rel, data in extra.items(): parts = rel.split("/") if dest_root is not None and parts and parts[0] in pkg_names: @@ -298,10 +283,9 @@ def record_mutation( action: str, skill: str, before_root: Optional[Path] = None, before: Optional[List[Dict[str, str]]] = None, after_root: Optional[Path] = None, actor: Optional[str] = None, evidence: Optional[Dict[str, Any]] = None) -> Optional[str]: - """Mutation hook: after-state from *after_root* (before = pre-captured list or - captured from *before_root*), then append. NEVER raises. delete/archive/purge - capture a COMPLETE package (filled from the newest curator backup) so - rollback never restores a shell.""" + """Mutation hook: after-state from *after_root* (before = pre-captured list or captured from + *before_root*), then append. NEVER raises. delete/archive/purge capture a COMPLETE package + (filled from the newest curator backup) so rollback never restores a shell.""" if not ledger_enabled(): return None try: @@ -375,9 +359,8 @@ def _validate_entry_paths(entry: Dict[str, Any]) -> Optional[str]: def rollback_entry(entry_id: str) -> Tuple[bool, str]: """Restore the before-state of mutation *entry_id*. Fail-closed (mirrors - agent/curator_backup.rollback): every before-blob must exist BEFORE any - change, and a pre-rollback safety entry of every touched path's CURRENT - state is appended first — if that fails, nothing is changed.""" + agent/curator_backup.rollback): every before-blob must exist BEFORE any change, and a + pre-rollback safety entry of every touched path's CURRENT state is appended first.""" entry = get_entry(entry_id) if entry is None: return False, f"no ledger entry with id '{entry_id}'" diff --git a/tools/skill_manager_batch.py b/tools/skill_manager_batch.py index c0a2f4eb8c..79fa30341c 100644 --- a/tools/skill_manager_batch.py +++ b/tools/skill_manager_batch.py @@ -169,7 +169,7 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non try: for i, op in enumerate(operations): raw = _smt._skill_manage_from( - {**op, "name": names[i]}, task_id=task_id, session_id=session_id) + {**op, "name": names[i], "operations": None}, task_id=task_id, session_id=session_id) try: parsed = json.loads(raw) except Exception: # noqa: BLE001 diff --git a/tools/skill_manager_guards.py b/tools/skill_manager_guards.py index 41e8a02b39..143c308238 100644 --- a/tools/skill_manager_guards.py +++ b/tools/skill_manager_guards.py @@ -1,9 +1,7 @@ -"""Write/delete guards for ``skill_manage``. - -Every guard returns ``None`` when the operation may proceed, otherwise a refusal -(error dict or message). Origin-owned state (``_find_skill``, ``_skills_dir``) is -reached lazily through ``tools.skill_manager_tool`` so test patches keep working. -""" +"""Write/delete guards for ``skill_manage``. Every guard returns ``None`` when the +operation may proceed, else a refusal (error dict or message). Origin-owned state +(``_find_skill``, ``_skills_dir``) is reached lazily via ``tools.skill_manager_tool`` +so test patches keep working.""" import contextvars as _ctxvars import logging @@ -56,10 +54,9 @@ _background_review_read_paths: "_ctxvars.ContextVar[Optional[_BackgroundReviewRe def mark_background_review_skill_read(path: Path) -> None: - """Record that the active background-review fork has read a skill file. - - The fork must not patch content it only inferred from the transcript: - skill_view/read_file call this, and the write guards require the mark.""" + """Record that the active background-review fork has read a skill file. The fork must not + patch content it only inferred from the transcript: skill_view/read_file call this, and + the write guards require the mark.""" if not _is_background_review(): return marks = _background_review_read_paths.get() @@ -79,8 +76,7 @@ def _reset_background_review_read_marks() -> None: def _resolved_roots(skill_path: Path): - """``(resolved skill_path, [(root, resolved_root), ...])`` over every skills root - whose resolve() succeeds; an unresolvable skill_path is used as-is.""" + """``(resolved skill_path, [(root, resolved_root), ...])`` over every resolvable skills root.""" from agent.skill_utils import get_all_skills_dirs try: @@ -103,8 +99,7 @@ def _containing_skills_root(skill_path: Path) -> Path: def _is_path_redirect(path: Path) -> bool: - """Symlink or (Windows 3.12+) junction — either lets a poisoned tree redirect - ``shutil.rmtree`` outside the skills root.""" + """Symlink or (Windows 3.12+) junction — either lets a poisoned tree redirect rmtree outside.""" try: return path.is_symlink() or (hasattr(path, "is_junction") and path.is_junction()) except OSError: @@ -112,9 +107,8 @@ def _is_path_redirect(path: Path) -> bool: def _validate_delete_target(skill_dir: Path) -> Optional[str]: - """Last-line guard before ``shutil.rmtree(skill_dir)``: even a poisoned tree - must never delete (1) a path outside every known skills root, (2) a skills - root itself, or (3) a symlink/junction (rmtree would follow it).""" + """Last-line guard before rmtree: even a poisoned tree must never delete (1) a path outside + every known skills root, (2) a skills root itself, (3) a symlink/junction (rmtree follows it).""" if _is_path_redirect(skill_dir): return ( f"Refusing to delete '{skill_dir}': the skill directory is a " @@ -147,11 +141,9 @@ def _is_pinned(name: str, what: str) -> Optional[bool]: def _pinned_guard(name: str) -> Optional[str]: - """Refusal message if *name* is pinned or essential, else None. - - Pin only guards **deletion**; patches/edits stay allowed. ESSENTIAL_SKILLS are - permanently pinned (the system prompt references them). Best-effort: an - unreadable sidecar lets the delete through.""" + """Refusal message if *name* is pinned or essential, else None. Pin only guards DELETION; + patches/edits stay allowed. ESSENTIAL_SKILLS are permanently pinned (the system prompt + references them). Best-effort: an unreadable sidecar lets the delete through.""" try: from agent.skill_utils import ESSENTIAL_SKILLS if name in ESSENTIAL_SKILLS: @@ -171,10 +163,8 @@ def _pinned_guard(name: str) -> Optional[str]: def _background_review_write_guard( name: str, skill_dir: Path, action: str) -> Optional[Dict[str, Any]]: - """Refuse autonomous curator writes to anything but curator-owned sediment. - - The background review fork has no user in the loop, so unlike foreground - agents it is also blocked on pinned/external/bundled/hub skills.""" + """Refuse autonomous curator writes to anything but curator-owned sediment. The review fork + has no user in the loop, so it is also blocked on pinned/external/bundled/hub skills.""" if not _is_background_review(): return None refuse = f"Refusing background curator {action} for" @@ -244,12 +234,10 @@ def _background_review_preflight(action: str, name: str) -> Optional[Dict[str, A def _curator_consolidation_delete_guard( name: str, absorbed_into: Optional[str]) -> Optional[Dict[str, Any]]: - """Fail closed on unverified deletes during the curator consolidation pass. - - The review fork's only legitimate delete is a verified consolidation declared - via ``absorbed_into=`` (existence validated in ``_delete_skill``). - The deterministic inactivity prune never calls ``skill_manage``, so a bare - delete here can only be the LLM pass pruning without evidence: refuse it.""" + """Fail closed on unverified deletes during the curator consolidation pass. The fork's only + legitimate delete is a consolidation declared via ``absorbed_into=`` (existence + validated in ``_delete_skill``); the deterministic inactivity prune never calls skill_manage, + so a bare delete here can only be the LLM pass pruning without evidence.""" if not _is_background_review() or (isinstance(absorbed_into, str) and absorbed_into.strip()): return None return _refusal( @@ -268,9 +256,8 @@ def _is_org_mirror(skill_path: Path) -> bool: def _maybe_auto_propose_org_edit(name: str, skill_path: Path) -> Optional[str]: - """Submit an org-skill edit upstream when `sync.org_auto_propose` is on. - Returns a note for the tool result or None; never raises (the edit is - already saved locally and can be proposed later).""" + """Submit an org-skill edit upstream when `sync.org_auto_propose` is on. Returns a note for + the tool result or None; never raises (the edit is saved locally and can be proposed later).""" try: from tools import skills_sync_client as ssc @@ -295,12 +282,10 @@ def _maybe_auto_propose_org_edit(name: str, skill_path: Path) -> Optional[str]: def _org_mirror_write_guard(name: str, skill_path: Path, action: str) -> Optional[Dict[str, Any]]: - """Org-shared skills are EDITABLE IN PLACE — this only blocks deletion. - - Edits land in the mirror, survive the next org pull (baseline sidecar in - skills_sync_client) and reach the org via `hermes sync propose`. Deletion - stays refused: the mirror is a view of org HEAD, so a local delete just - comes back, and removing for everyone is an admin action.""" + """Org-shared skills are EDITABLE IN PLACE — this only blocks deletion. Edits land in the + mirror, survive the next org pull (baseline sidecar in skills_sync_client) and reach the org + via `hermes sync propose`. Deletion stays refused: the mirror is a view of org HEAD, so a + local delete just comes back, and removing for everyone is an admin action.""" if action not in {"delete", "remove_file"}: return None try: diff --git a/tools/skill_manager_tool.py b/tools/skill_manager_tool.py index e1c45f3f49..d99d977435 100644 --- a/tools/skill_manager_tool.py +++ b/tools/skill_manager_tool.py @@ -1,11 +1,10 @@ #!/usr/bin/env python3 """Skill Manager Tool — agent-managed skill creation & editing. -Skills are the agent's procedural memory (narrow "how to do X"), as opposed to -MEMORY.md/USER.md (broad, declarative). New skills land in ~/.hermes/skills/ -(or ``skills.create_dir``); existing skills (bundled, hub, user) are modified in -place. Layout: ``/[category/]/SKILL.md`` + optional -``references/ templates/ scripts/ assets/``. +Skills are the agent's procedural memory (narrow "how to do X"; MEMORY.md/USER.md are +broad, declarative). New skills land in ~/.hermes/skills/ (or ``skills.create_dir``); +existing skills (bundled, hub, user) are modified in place. Layout: +``/[category/]/SKILL.md`` + optional ``references/ templates/ scripts/ assets/``. """ import contextvars as _ctxvars @@ -43,8 +42,7 @@ logger = logging.getLogger(__name__) def _guard_agent_created_enabled() -> bool: - """skills.guard_agent_created (default False): the agent can already run the same - code via terminal() ungated, so the scan is opt-in belt-and-suspenders.""" + """skills.guard_agent_created (default False): opt-in — terminal() runs the same code ungated.""" try: from hermes_cli.config import load_config return is_truthy_value( @@ -54,9 +52,8 @@ def _guard_agent_created_enabled() -> bool: def _security_scan_skill(skill_dir: Path) -> Optional[str]: - """Post-write scan; error string if blocked, else None. No-op unless - skills.guard_agent_created. An "ask" verdict (dangerous findings) is surfaced - as an error so the agent can retry with the flagged content removed.""" + """Post-write scan (opt-in); error string if blocked, else None. An "ask" verdict + (dangerous findings) is surfaced as an error so the agent can retry without them.""" if not _guard_agent_created_enabled(): return None try: @@ -78,9 +75,8 @@ _SKILLS_DIR_AT_IMPORT = SKILLS_DIR def _skills_dir() -> Path: - """Active profile's skills dir at call time: multi-profile runtimes import once - and bind a different profile per session. An explicitly patched module-level - ``SKILLS_DIR`` (tests) wins, otherwise resolve from the live HERMES_HOME.""" + """Active profile's skills dir at call time (multi-profile runtimes rebind per session). + An explicitly patched module-level ``SKILLS_DIR`` (tests) wins over the live HERMES_HOME.""" configured = Path(SKILLS_DIR) return configured if configured != _SKILLS_DIR_AT_IMPORT else get_hermes_home() / "skills" @@ -134,11 +130,9 @@ def _validate_category(category: Optional[str]) -> Optional[str]: def _validate_frontmatter(content: str, *, new_skill: bool = False) -> Optional[str]: - """Validate frontmatter (name + description) and a non-empty body. - - ``new_skill`` (create only) also enforces SKILL_PROMPT_DESC_LIMIT so new skills - never lose routing signal to index truncation; edit/patch skip it so existing - over-limit skills remain maintainable.""" + """Validate frontmatter (name + description) and a non-empty body. ``new_skill`` (create + only) also enforces SKILL_PROMPT_DESC_LIMIT so new skills never lose routing signal to + index truncation; edit/patch skip it so existing over-limit skills stay maintainable.""" if not content.strip(): return "Content cannot be empty." content = content.lstrip("\ufeff") # tolerate a Windows UTF-8 BOM @@ -197,7 +191,7 @@ def _resolve_skill_dir(name: str, category: str = None) -> Path: base = get_skill_create_dir() or base except Exception: logger.debug("skills.create_dir lookup failed", exc_info=True) - return base / category / name if category else base / name + return base / (category or "") / name def _iter_skill_dirs(root: Path): @@ -208,16 +202,13 @@ def _iter_skill_dirs(root: Path): def _find_skill(name: str) -> Optional[Dict[str, Any]]: - """Find a skill across the local skills dir then skills.external_dirs. + """Find a skill (local skills dir, then skills.external_dirs) -> ``{"path": Path}`` | None. - Accepts the bare dir name (``axolotl``) and the categorized relative path - (``mlops/axolotl``) — the two forms skill_view resolves. Bare lookups compare - the skill's own dir name so category-nested skills still match. - Returns ``{"path": Path}`` or None.""" + Accepts the bare dir name (``axolotl``; matches category-nested skills too) and the + categorized relative path (``mlops/axolotl``) — the two forms skill_view resolves. The + categorized form matches RELATIVE to the local root only (relative_to raises for external dirs).""" from agent.skill_utils import get_all_skills_dirs - # The categorized form matches RELATIVE to the local root only (relative_to - # raises for external dirs). local_root = None if "/" in name or "\\" in name: try: @@ -243,8 +234,8 @@ def _find_skill(name: str) -> Optional[Dict[str, Any]]: def _find_skill_in_other_profiles(name: str) -> List[Tuple[str, Path]]: - """``(profile, skill_dir)`` pairs for OTHER profiles holding ``name`` (so the - not-found error can explain a wrong-profile mistake). Fail-quiet.""" + """``(profile, skill_dir)`` pairs for OTHER profiles holding ``name`` (so the not-found + error can explain a wrong-profile mistake). Fail-quiet.""" matches: List[Tuple[str, Path]] = [] try: from hermes_constants import get_default_hermes_root @@ -255,10 +246,9 @@ def _find_skill_in_other_profiles(name: str) -> List[Tuple[str, Path]]: active_dir = _active.resolve() if _active.exists() else _active # Every profile's skills dir EXCEPT the active one (already searched). candidates: List[Tuple[str, Path]] = [("default", root / "skills")] - profiles_root = root / "profiles" with suppress(OSError): - if profiles_root.is_dir(): - candidates += [(e.name, e / "skills") for e in profiles_root.iterdir() if e.is_dir()] + if (root / "profiles").is_dir(): + candidates += [(e.name, e / "skills") for e in (root / "profiles").iterdir() if e.is_dir()] for profile_name, skills_dir in candidates: with suppress(OSError, RuntimeError): if skills_dir.resolve() == active_dir or not skills_dir.is_dir(): @@ -323,8 +313,8 @@ def _resolve_supporting_file(skill_dir: Path, file_path: str): def _locate_for_write(name: str, action: str, not_found_suffix: str = "", *, org_guard: bool = True): - """Find the skill and run the org-mirror (unless ``org_guard=False``) + - background-review write guards -> ``(skill_dir, None)`` | ``(None, error_dict)``.""" + """Find the skill; run the org-mirror (unless ``org_guard=False``) and background-review + write guards -> ``(skill_dir, None)`` | ``(None, error_dict)``.""" existing = _find_skill(name) if not existing: return None, _err(_skill_not_found_error(name, not_found_suffix)) @@ -336,9 +326,8 @@ def _locate_for_write(name: str, action: str, not_found_suffix: str = "", *, def _guarded_write(name: str, skill_dir: Path, target: Path, action: str, label: str, content: str) -> Optional[Dict[str, Any]]: - """Read-before-write guard (existing targets only), atomic write, then the - security scan; a blocked scan restores the original (or unlinks a new file). - Returns an error dict or None.""" + """Read-before-write guard (existing targets only), atomic write, then the security scan; + a blocked scan restores the original (or unlinks a new file). Error dict or None.""" original = None if target.exists(): if read_guard := _background_review_read_before_write_guard(name, target, action, label): @@ -413,13 +402,11 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An display = skill_dir.relative_to(root) if skill_dir.is_relative_to(root) else skill_dir result = { "success": True, "message": f"Skill '{name}' created.", "path": str(display), - "skill_md": str(skill_md), "_change": {"description": _description_preview(content)}} - if category: - result["category"] = category - result["hint"] = ( - "To add reference files, templates, or scripts, use " - "skill_manage(action='write_file', name='{}', file_path='references/example.md', file_content='...')".format(name) - ) + "skill_md": str(skill_md), "_change": {"description": _description_preview(content)}, + **({"category": category} if category else {}), + "hint": "To add reference files, templates, or scripts, use " + f"skill_manage(action='write_file', name='{name}', file_path='references/example.md', " + "file_content='...')"} _add_description_prompt_preview(result, content) _attach_lint_findings(result, skill_md) return result @@ -444,11 +431,10 @@ def _edit_skill(name: str, content: str) -> Dict[str, Any]: def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = None, replace_all: bool = False) -> Dict[str, Any]: - """Targeted find-and-replace within SKILL.md (default) or a supporting file; - requires a unique match unless replace_all.""" + """Targeted find-and-replace in SKILL.md (default) or a supporting file; unique match unless replace_all.""" if not old_string: - # A bare "required" error is a dead end: the model retries blindly and - # often escapes to action='write_file', clobbering the whole file. + # A bare "required" error is a dead end: the model retries blindly and often + # escapes to action='write_file', clobbering the whole file. return _err( "old_string is required for 'patch' and must be the EXACT text currently in the file. " "Read the target file first (read_file on the skill's SKILL.md, or the file named by " @@ -456,12 +442,12 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N "action='write_file' — that rewrites the entire file and destroys unrelated content.") if new_string is None: return _err("new_string is required for 'patch'. Use an empty string to delete matched text.") - # No old_string == new_string guard here: fuzzy_find_and_replace rejects - # that with a richer error (file_preview) this layer cannot produce. - + # No old_string == new_string guard here: fuzzy_find_and_replace rejects that with a + # richer error (file_preview) this layer cannot produce. skill_dir, guard = _locate_for_write(name, "patch") if guard: return guard + target_label = file_path or "SKILL.md" if file_path: target, err = _resolve_supporting_file(skill_dir, file_path) if err: @@ -470,13 +456,12 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N target = skill_dir / "SKILL.md" if not target.exists(): return _err(f"File not found: {target.relative_to(skill_dir)}") - target_label = file_path or "SKILL.md" if read_guard := _background_review_read_before_write_guard(name, target, "patch", target_label): return read_guard content = target.read_text(encoding="utf-8") - # Same fuzzy engine as the file patch tool (whitespace/indent/escape - # normalization, block anchors) so minor formatting mismatches don't fail. + # Same fuzzy engine as the file patch tool (whitespace/indent/escape normalization, + # block anchors) so minor formatting mismatches don't fail. from tools.fuzzy_match import fuzzy_find_and_replace new_content, match_count, _strategy, match_error = fuzzy_find_and_replace( content, old_string, new_string, replace_all) @@ -501,9 +486,8 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, Any]: - """Delete a skill. ``absorbed_into``: None = undeclared (legacy, accepted); "" = - explicit prune; "" = absorbed into that umbrella, which must exist on - disk (validated here so the model can't claim a nonexistent umbrella).""" + """Delete a skill. ``absorbed_into``: None = undeclared (legacy, accepted); "" = explicit prune; + "" = absorbed into that umbrella, which must exist (so the model can't claim one).""" skill_dir, guard = _locate_for_write(name, "delete") if guard := guard or _curator_consolidation_delete_guard(name, absorbed_into): return guard @@ -521,9 +505,8 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A skills_root = _containing_skills_root(skill_dir) if unsafe := _validate_delete_target(skill_dir): # defense-in-depth before rmtree return _err(unsafe) - - # Curator consolidations must be RECOVERABLE (`hermes curator restore`): archive - # instead of rmtree. Foreground deletes keep hard-delete semantics. + # Curator consolidations must be RECOVERABLE (`hermes curator restore`): archive instead + # of rmtree. Foreground deletes keep hard-delete semantics. absorbed_note = f" Content absorbed into '{absorbed_target}'." if absorbed_target else "" if _is_background_review(): try: @@ -533,9 +516,8 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A return _err(f"failed to archive '{name}': {e}") if not ok: return _err(archive_msg) - return {"success": True, - "message": f"Skill '{name}' archived ({archive_msg}).{absorbed_note}", - "_archived": True} + return {"success": True, "_archived": True, + "message": f"Skill '{name}' archived ({archive_msg}).{absorbed_note}"} shutil.rmtree(skill_dir) _rmdir_if_empty(skill_dir.parent, skills_root) # empty category dir, never the root @@ -582,7 +564,7 @@ def _remove_file(name: str, file_path: str) -> Dict[str, Any]: target, err = _resolve_supporting_file(skill_dir, file_path) if err: return err - if not target.exists(): + if not target.exists(): # list what IS there so the model can pick the right path available = [ str(f.relative_to(skill_dir)) for subdir in ALLOWED_SUBDIRS if (skill_dir / subdir).exists() @@ -604,13 +586,11 @@ def _remove_file(name: str, file_path: str) -> Dict[str, Any]: _skill_gate_bypass: "_ctxvars.ContextVar[bool]" = _ctxvars.ContextVar( "skill_gate_bypass", default=False) -_GATED_ACTIONS = {"create", "edit", "patch", "delete", "write_file", "remove_file"} - def _run_write_gate(build_staging): """Shared write gate: None to proceed, else a JSON tool result (blocked/staged). - ``build_staging(wa) -> (payload, gist)`` runs only when staging. Fails open - if write_approval cannot be imported.""" + ``build_staging(wa) -> (payload, gist)`` runs only when staging. Fails open if + write_approval cannot be imported.""" try: from tools import write_approval as wa except Exception: @@ -627,9 +607,8 @@ def _run_write_gate(build_staging): def _apply_skill_write_gate(action, name, **payload_kwargs): - """Flat-shape gate: stage the full kwargs so approval can replay them; - bypassed during approved-pending replay.""" - if action not in _GATED_ACTIONS or _skill_gate_bypass.get(): + """Flat-shape gate: stage the full kwargs so approval can replay them; bypassed during replay.""" + if action not in _ACTION_HANDLERS or _skill_gate_bypass.get(): return None def _staging(wa): @@ -642,11 +621,12 @@ def _apply_skill_write_gate(action, name, **payload_kwargs): return _run_write_gate(_staging) -_FLAT_OP_KEYS = ("content", "category", "file_path", "file_content", "old_string", "new_string") +_FLAT_OP_KEYS = ("content", "category", "file_path", "file_content", "old_string", "new_string", + "absorbed_into", "operations") def _skill_manage_from(payload: Dict[str, Any], **extra) -> str: - """Call ``skill_manage`` with the flat-shape fields taken from ``payload``.""" + """Call ``skill_manage`` with the flat-shape fields (and absorbed_into/operations) of ``payload``.""" return skill_manage( action=payload.get("action", ""), name=payload.get("name", ""), replace_all=payload.get("replace_all", False), @@ -657,9 +637,7 @@ def apply_skill_pending(payload: Dict[str, Any]) -> str: """Replay a staged skill write, bypassing the gate (the /skills approve handler).""" token = _skill_gate_bypass.set(True) try: - return _skill_manage_from( - payload, absorbed_into=payload.get("absorbed_into"), - operations=payload.get("operations")) + return _skill_manage_from(payload) finally: _skill_gate_bypass.reset(token) @@ -672,9 +650,8 @@ _SYNC_PUSH_DEBOUNCE_S = 5.0 def _maybe_debounced_sync_push(skill_name: str) -> None: - """Debounced best-effort sync push after a skill write; never blocks the caller. - Skills not opted into sync do nothing (no auth, no network); the push itself - (``skills_sync_client.maybe_push_skills``) enforces the access gate.""" + """Debounced best-effort sync push after a skill write; never blocks the caller. Skills not + opted into sync do nothing (no auth/network); ``maybe_push_skills`` enforces the access gate.""" global _sync_push_timer try: from tools.skill_usage import is_sync_enabled @@ -733,20 +710,15 @@ _REQUIRED_ARGS = { def _record_success(action, name, result, *, file_path, absorbed_into, task_id, session_id, ledger_before) -> None: - """Best-effort post-mutation side effects (never break the tool): audit ledger, - prompt-cache clear, curator telemetry, debounced sync push.""" + """Best-effort post-mutation side effects (never break the tool): ledger, prompt-cache + clear, curator telemetry, debounced sync push.""" with suppress(Exception): from tools import skill_ledger as _ledger _post = _find_skill(name) - _evidence = {} - if action == "delete": - # consolidation vs prune, and whether the recoverable archive handled it - _evidence["absorbed_into"] = absorbed_into - _evidence["archived"] = bool(result.get("_archived")) - if session_id: - _evidence["session_id"] = session_id - if file_path: - _evidence["file_path"] = file_path + # delete: consolidation vs prune, and whether the recoverable archive handled it + _evidence = ({"absorbed_into": absorbed_into, "archived": bool(result.get("_archived"))} + if action == "delete" else {}) + _evidence.update({k: v for k, v in (("session_id", session_id), ("file_path", file_path)) if v}) _ledger.record_mutation( action, name, before=ledger_before if ledger_before is not None else [], after_root=_post["path"] if _post else None, evidence=_evidence) @@ -777,8 +749,8 @@ def skill_manage( file_content: str = None, old_string: str = None, new_string: str = None, replace_all: bool = False, absorbed_into: str = None, task_id: str = None, session_id: str = None, operations=None) -> str: - """Dispatch to the action handler; returns a JSON string. ``operations`` (batch - shape, applied atomically by _skill_manage_batch) overrides the flat fields.""" + """Dispatch to the action handler -> JSON string. ``operations`` (atomic batch shape, + see _skill_manage_batch) overrides the flat fields.""" if operations is not None: return _skill_manage_batch( operations, default_name=name or None, task_id=task_id, session_id=session_id) @@ -794,9 +766,9 @@ def skill_manage( if (gate_result := _apply_skill_write_gate(action, name, **args)) is not None: return gate_result - # Audit ledger pre-capture: telemetry, not a gate — failures must NEVER block the - # mutation. delete destroys the whole package (and consolidation may have re-homed - # support files first), so complete it from the newest curator backup or a restore is hollow. + # Ledger pre-capture: telemetry, not a gate — failures must NEVER block the mutation. delete + # destroys the whole package (consolidation may have re-homed support files first), so + # complete it from the newest curator backup or a restore is hollow. _ledger_before = None with suppress(Exception): from tools import skill_ledger as _ledger @@ -924,5 +896,4 @@ from tools.registry import registry, tool_error registry.register( name="skill_manage", toolset="skills", schema=SKILL_MANAGE_SCHEMA, emoji="📝", handler=lambda args, **kw: _skill_manage_from( - args, absorbed_into=args.get("absorbed_into"), operations=args.get("operations"), - task_id=kw.get("task_id"), session_id=kw.get("session_id"))) + args, task_id=kw.get("task_id"), session_id=kw.get("session_id"))) From 760ac0b609ba885086bb2190021e378aed0075c9 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:18:04 -0700 Subject: [PATCH 04/16] =?UTF-8?q?refactor(tools):=20MCP=20handlers/errors/?= =?UTF-8?q?transport/config/health/registration/oauth-manager=20compaction?= =?UTF-8?q?=20=E2=80=94=20shared=20=5Fdispatch,=20folded=20utility=20facto?= =?UTF-8?q?ry,=20auth-type=20cache=20tuple?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/mcp_oauth_manager.py | 91 ++++++--------- tools/mcp_tool_config.py | 110 ++++++++---------- tools/mcp_tool_errors.py | 64 +++++------ tools/mcp_tool_handlers.py | 200 +++++++++++++-------------------- tools/mcp_tool_health.py | 26 ++--- tools/mcp_tool_registration.py | 64 +++++------ tools/mcp_tool_transport.py | 186 ++++++++++++++---------------- 7 files changed, 311 insertions(+), 430 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index c0141159a8..2ca0fe2cd1 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -1,12 +1,10 @@ -"""Central manager for per-server MCP OAuth state (one instance per process). - -Holds per-server providers and coordinates cross-process token reload (mtime-based disk watch, -so tokens refreshed by cron/another CLI are picked up without a restart), 401 deduplication -(N concurrent 401s with the same access_token trigger one recovery) and reconnect signalling -(``MCPServerTask`` drives the reconnect; the manager decides when). The ONLY place that -instantiates the SDK's ``OAuthClientProvider`` for runtime use. We rely on the SDK's lazy -refresh: one ``stat()`` per tool call is cheaper than an await + refresh round-trip. -""" +"""Central manager for per-server MCP OAuth state (one instance per process): per-server +providers, cross-process token reload (mtime-based disk watch so tokens refreshed by cron/another +CLI are picked up without a restart), 401 deduplication (N concurrent 401s with the same +access_token trigger one recovery) and reconnect signalling (``MCPServerTask`` drives the +reconnect; the manager decides when). The ONLY place that instantiates the SDK's +``OAuthClientProvider`` for runtime use; refresh stays lazy in the SDK — one ``stat()`` per tool +call is cheaper than an await + refresh round-trip.""" from __future__ import annotations @@ -67,20 +65,18 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): def _hermes_storage(self): """The context storage when it is a ``HermesTokenStorage``, else None.""" from tools.mcp_oauth import HermesTokenStorage - storage = self.context.storage - return storage if isinstance(storage, HermesTokenStorage) else None + return self.context.storage if isinstance(self.context.storage, HermesTokenStorage) else None def _log_nonfatal(self, what: str, exc: BaseException) -> None: logger.debug("MCP OAuth '%s': %s failed (non-fatal): %s", self._hermes_server_name, what, exc) async def _initialize(self) -> None: """Load stored state, seed ``token_expiry_time``, restore/prefetch metadata. The SDK's - ``_initialize`` never calls ``update_token_expiry``, so ``is_token_valid()`` is True for - any loaded token regardless of age and a restarted process ships stale Bearer tokens; - seeding the expiry (``HermesTokenStorage`` persists absolute ``expires_at``) makes the SDK - refresh first. Metadata is restored from disk, else discovered pre-flight when we hold - tokens but no metadata: otherwise ``_refresh_token`` guesses ``{server_url}/token`` - (wrong for split-origin providers), 404s, and we fall through to browser reauth.""" + ``_initialize`` never calls ``update_token_expiry``, so a restarted process would ship stale + Bearer tokens as "valid"; seeding the expiry (``HermesTokenStorage`` persists absolute + ``expires_at``) makes the SDK refresh first. Metadata is restored from disk, else discovered + pre-flight when we hold tokens but no metadata: otherwise ``_refresh_token`` guesses + ``{server_url}/token`` (wrong for split-origin providers), 404s, and we fall to browser reauth.""" await super()._initialize() tokens = self.context.current_tokens if tokens is not None and tokens.expires_in is not None: @@ -160,8 +156,7 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): async def _is_invalid_client_at_token_endpoint(self, response: Any) -> bool: """True when *response* is the token endpoint (same scheme/host/path, query ignored) rejecting our client_id with ``invalid_client`` — whole word, so RFC 7591's - ``invalid_client_metadata`` does not trip it. The body is read only after the endpoint - matches.""" + ``invalid_client_metadata`` does not trip it. The body is read only after the endpoint matches.""" from urllib.parse import urlsplit token_endpoint = getattr(getattr(self.context, "oauth_metadata", None), "token_endpoint", None) req = getattr(response, "request", None) @@ -171,11 +166,9 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): pa, pb = urlsplit(str(req.url)), urlsplit(str(token_endpoint)) except ValueError: # pragma: no cover — malformed URL return False - if not (pa.scheme == pb.scheme and pa.netloc.lower() == pb.netloc.lower() - and pa.path.rstrip("/") == pb.path.rstrip("/")): + if (pa.scheme, pa.netloc.lower(), pa.path.rstrip("/")) != (pb.scheme, pb.netloc.lower(), pb.path.rstrip("/")): return False - body = await response.aread() - return re.search(rb"\binvalid_client\b", body.lower()) is not None + return re.search(rb"\binvalid_client\b", (await response.aread()).lower()) is not None async def _maybe_flag_poisoned_client(self, response: Any) -> None: """An ``invalid_client`` rejection of our ``client_id`` at the token endpoint proves the @@ -185,15 +178,13 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): in the body; pre-registered clients are never poisoned; any failure is swallowed. The browser-side "Redirect URI Mismatch" case has no HTTP signal (``hermes mcp reauth``).""" try: - if self._hermes_preregistered or getattr(response, "status_code", None) not in (400, 401): - return - if not await self._is_invalid_client_at_token_endpoint(response): + if (self._hermes_preregistered or getattr(response, "status_code", None) not in (400, 401) + or not await self._is_invalid_client_at_token_endpoint(response)): return storage = self._hermes_storage() - # If the rejected client_id was our CIMD URL, re-presenting it would loop (the - # server already fetched and refused it). Drop the URL so the retry takes DCR, and - # mark it on disk so the next process doesn't walk back into the same refusal - # (`hermes mcp login` clears the marker). + # A rejected CIMD URL would loop if re-presented (the server already fetched and refused + # it): drop it so the retry takes DCR, and mark it on disk so the next process doesn't walk + # back into the same refusal (`hermes mcp login` clears the marker). cimd_url = getattr(self.context, "client_metadata_url", None) if cimd_url and getattr(self.context.client_info, "client_id", None) == cimd_url: logger.warning("MCP OAuth '%s': authorization server rejected our Client ID Metadata Document (%s) " @@ -211,18 +202,15 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): self._log_nonfatal("invalid_client detection", exc) async def async_auth_flow(self, request): # type: ignore[override] - # Pre-flow hook: reload from disk if it changed (non-fatal on error). - try: + try: # pre-flow hook: reload from disk if it changed (non-fatal on error) await get_manager().invalidate_if_disk_changed(self._hermes_server_name, hermes_home=self._hermes_home) except Exception as exc: # pragma: no cover — defensive self._log_nonfatal("pre-flow disk-watch", exc) - # Bridge the bidirectional generator by hand: a naive ``async for item in inner: yield # item`` DISCARDS the responses httpx sends back via ``asend``, and the SDK crashes on None. inner = super().async_auth_flow(request) - resource_lock_released = False + resource_lock_released = retry_after_concurrent_auth = False sent_access_token = None - retry_after_concurrent_auth = False try: outgoing = await inner.__anext__() while True: @@ -252,8 +240,8 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): self._persist_oauth_metadata_if_changed() # metadata discovered lazily in the 401 branch finally: if resource_lock_released: - # Balance the SDK's surrounding ``async with`` even when HTTPX cancels/closes - # the flow mid-request. Shield only this local bookkeeping. + # Balance the SDK's surrounding ``async with`` even when HTTPX cancels/closes the + # flow mid-request; shield only this local bookkeeping. import anyio with anyio.CancelScope(shield=True): await self.context.lock.acquire() @@ -274,8 +262,7 @@ class MCPOAuthManager: def __init__(self) -> None: self._entries: dict[tuple[str, str], _ProviderEntry] = {} self._entries_lock = threading.Lock() - # Strong refs to in-flight 401 tasks so the loop's weak bookkeeping cannot GC them - # mid-run and leave `await pending` hanging forever. + # Strong refs to in-flight 401 tasks so the loop's weak bookkeeping cannot GC them mid-run. self._inflight_tasks: set[asyncio.Task] = set() def get_or_build_provider(self, server_name: str, server_url: str, oauth_config: Optional[dict]) -> Optional[Any]: @@ -306,8 +293,7 @@ class MCPOAuthManager: if _HERMES_PROVIDER_CLS is None: logger.warning("MCP OAuth '%s': SDK auth module unavailable", server_name) return None - # Local imports avoid circular deps at module import time. - from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow + from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow # lazy: circular at import time from tools.mcp_oauth import _OAUTH_AVAILABLE, OAuthNonInteractiveError, _is_interactive from tools.mcp_oauth_provider import build_provider_kwargs, prepare_oauth_config if not _OAUTH_AVAILABLE: @@ -350,16 +336,14 @@ class MCPOAuthManager: if entry is None or entry.provider is None: return False async with entry.lock: - tokens_path = _get_token_dir(hermes_home) / f"{_safe_filename(server_name)}.json" try: - mtime_ns = tokens_path.stat().st_mtime_ns - except (FileNotFoundError, OSError): + mtime_ns = (_get_token_dir(hermes_home) / f"{_safe_filename(server_name)}.json").stat().st_mtime_ns + except OSError: return False if mtime_ns == entry.last_mtime_ns: return False old, entry.last_mtime_ns = entry.last_mtime_ns, mtime_ns - # `_initialized` is private SDK API but stable across the versions we pin - # (>=1.26.0); resetting it forces a reload. + # `_initialized` is private SDK API but stable across the pinned versions (>=1.26.0). if hasattr(entry.provider, "_initialized"): entry.provider._initialized = False # noqa: SLF001 logger.info("MCP OAuth '%s': tokens file changed (mtime %d -> %d), forcing reload", server_name, old, mtime_ns) @@ -367,22 +351,22 @@ class MCPOAuthManager: async def _recover_401(self, server_name: str, entry: _ProviderEntry, key: str, pending: asyncio.Future) -> None: """Single recovery attempt behind *pending*; always clears the dedup slot.""" + can_refresh = False try: # Disk changed (external refresh)? Else: if the SDK can refresh in place, let the # caller retry (the httpx.Auth flow refreshes on the next request). - can_refresh = True - if not await self.invalidate_if_disk_changed(server_name): + if await self.invalidate_if_disk_changed(server_name): + can_refresh = True + else: try: can_refresh = bool(entry.provider.context.can_refresh_token()) except Exception: # no context / not callable / probe failed can_refresh = False - if not pending.done(): - pending.set_result(can_refresh) except Exception as exc: # pragma: no cover — defensive logger.warning("MCP OAuth '%s': 401 handler failed: %s", server_name, exc) - if not pending.done(): - pending.set_result(False) finally: + if not pending.done(): + pending.set_result(can_refresh) entry.pending_401.pop(key, None) async def handle_401(self, server_name: str, failed_access_token: Optional[str] = None) -> bool: @@ -394,11 +378,10 @@ class MCPOAuthManager: if entry is None or entry.provider is None: return False key = failed_access_token or "" - loop = asyncio.get_running_loop() async with entry.lock: pending = entry.pending_401.get(key) if pending is None: - pending = entry.pending_401[key] = loop.create_future() + pending = entry.pending_401[key] = asyncio.get_running_loop().create_future() task = asyncio.create_task(self._recover_401(server_name, entry, key, pending)) self._inflight_tasks.add(task) task.add_done_callback(self._inflight_tasks.discard) diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py index 0e39bf1a50..34ae1a716b 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -19,27 +19,25 @@ _mcp_stderr_log_lock = threading.Lock() def _get_mcp_stderr_log() -> Any: - """Shared append-mode handle for MCP subprocess stderr, opened once per process. Must - expose a real fd (``fileno()``) because asyncio wires the child's stderr directly to it. - Falls back to ``/dev/null``, then real stderr.""" + """Shared append-mode handle for MCP subprocess stderr, opened once per process. Must expose a + real fd (asyncio wires the child's stderr to it); falls back to ``/dev/null``, then real stderr.""" global _mcp_stderr_log_fh with _mcp_stderr_log_lock: - if _mcp_stderr_log_fh is not None: - return _mcp_stderr_log_fh - try: - from hermes_constants import get_hermes_home - log_dir = get_hermes_home() / "logs" - log_dir.mkdir(parents=True, exist_ok=True) - # Line-buffered so output lands promptly; errors="replace" tolerates garbled binary. - fh = open(log_dir / "mcp-stderr.log", "a", encoding="utf-8", errors="replace", buffering=1) - fh.fileno() # confirm a real fd before committing - _mcp_stderr_log_fh = fh - except Exception as exc: # pragma: no cover — best-effort fallback - logger.debug("Failed to open MCP stderr log, using devnull: %s", exc) + if _mcp_stderr_log_fh is None: try: - _mcp_stderr_log_fh = open(os.devnull, "w", encoding="utf-8") - except Exception: - _mcp_stderr_log_fh = sys.stderr + from hermes_constants import get_hermes_home + log_dir = get_hermes_home() / "logs" + log_dir.mkdir(parents=True, exist_ok=True) + # Line-buffered so output lands promptly; errors="replace" tolerates garbled binary. + fh = open(log_dir / "mcp-stderr.log", "a", encoding="utf-8", errors="replace", buffering=1) + fh.fileno() # confirm a real fd before committing + _mcp_stderr_log_fh = fh + except Exception as exc: # pragma: no cover — best-effort fallback + logger.debug("Failed to open MCP stderr log, using devnull: %s", exc) + try: + _mcp_stderr_log_fh = open(os.devnull, "w", encoding="utf-8") + except Exception: + _mcp_stderr_log_fh = sys.stderr return _mcp_stderr_log_fh @@ -48,8 +46,7 @@ def _write_stderr_log_header(server_name: str) -> None: (per-line prefixes would need a pipe + reader thread).""" fh = _core._get_mcp_stderr_log() try: - ts = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - fh.write(f"\n===== [{ts}] starting MCP server '{server_name}' =====\n") + fh.write(f"\n===== [{datetime.now():%Y-%m-%d %H:%M:%S}] starting MCP server '{server_name}' =====\n") fh.flush() except Exception: pass @@ -58,16 +55,15 @@ def _write_stderr_log_header(server_name: str) -> None: # Env vars safe to pass to stdio subprocesses (no secrets). _SAFE_ENV_KEYS = frozenset({"PATH", "HOME", "USER", "LANG", "LC_ALL", "TERM", "SHELL", "TMPDIR"}) -# Windows process/location vars needed by launcher-style tools (e.g. Docker -# Desktop's MCP plugin discovery); none carry secrets. +# Windows process/location vars needed by launcher-style tools (e.g. Docker Desktop's MCP plugin +# discovery); none carry secrets. _SAFE_ENV_KEYS_CASE_INSENSITIVE = frozenset({ "ALLUSERSPROFILE", "APPDATA", "COMMONPROGRAMFILES", "COMMONPROGRAMFILES(X86)", "COMMONPROGRAMW6432", "COMPUTERNAME", "COMSPEC", "HOMEDRIVE", "HOMEPATH", "LOCALAPPDATA", "NUMBER_OF_PROCESSORS", "OS", "PATHEXT", "PROCESSOR_ARCHITECTURE", "PROGRAMDATA", "PROGRAMFILES", "PROGRAMFILES(X86)", "PROGRAMW6432", "PUBLIC", "SYSTEMDRIVE", "SYSTEMROOT", "TEMP", "TMP", "USERDOMAIN", "USERNAME", - "USERPROFILE", "WINDIR", -}) + "USERPROFILE", "WINDIR"}) # ${VAR_NAME} interpolation; any non-} chars allowed so MY-VAR / my.var work. _ENV_VAR_PATTERN = re.compile(r"\$\{([^}]+)\}") @@ -79,11 +75,9 @@ def _workspace_folder() -> str: try: from tools.file_tools import _authoritative_workspace_root root = _authoritative_workspace_root() - if root: - return root except Exception: - pass - return os.getcwd() + root = None + return root or os.getcwd() def _workspace_basename() -> str: @@ -93,12 +87,8 @@ def _workspace_basename() -> str: # Cursor's case-sensitive context vars -> resolver. _CONTEXT_VAR_RESOLVERS = { - "userHome": lambda: os.path.expanduser("~"), - "workspaceFolder": lambda: _core._workspace_folder(), - "workspaceFolderBasename": _workspace_basename, - "pathSeparator": lambda: os.sep, - "/": lambda: os.sep, -} + "userHome": lambda: os.path.expanduser("~"), "workspaceFolder": lambda: _core._workspace_folder(), + "workspaceFolderBasename": _workspace_basename, "pathSeparator": lambda: os.sep, "/": lambda: os.sep} def _build_safe_env(user_env: Optional[dict]) -> dict: @@ -140,16 +130,11 @@ def _node_fallback(command: str) -> str: failed; *command* unchanged when none is executable.""" home = os.path.expanduser("~") hermes_home = os.path.expanduser(os.getenv("HERMES_HOME", os.path.join(home, ".hermes"))) - candidates = [ - os.path.join(hermes_home, "node", "bin", command), - os.path.join(home, ".local", "bin", command), - # Canonical Node location (from-source Linux, Hermes Docker image, Intel Homebrew). Needed - # when a hand-authored env.PATH omits it: npx's shebang re-execs /usr/bin/env node. - os.path.join(os.sep, "usr", "local", "bin", command)] - for candidate in candidates: - if os.path.isfile(candidate) and os.access(candidate, os.X_OK): - return candidate - return command + # /usr/local/bin: canonical Node location (from-source Linux, Hermes Docker image, Intel Homebrew), + # needed when a hand-authored env.PATH omits it — npx's shebang re-execs /usr/bin/env node. + candidates = (os.path.join(hermes_home, "node", "bin", command), os.path.join(home, ".local", "bin", command), + os.path.join(os.sep, "usr", "local", "bin", command)) + return next((c for c in candidates if os.path.isfile(c) and os.access(c, os.X_OK)), command) def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: @@ -157,7 +142,6 @@ def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: ``npx``/``npm``/``node`` work under a filtered PATH.""" resolved_command = os.path.expanduser(str(command).strip()) resolved_env = dict(env or {}) - if os.sep not in resolved_command: path_arg = resolved_env.get("PATH") which_hit = shutil.which(resolved_command, path=path_arg) @@ -178,9 +162,9 @@ def _wrap_command_with_watchdog(command: str, args: list) -> tuple[str, list]: """Wrap a stdio command in the parent-death watchdog (POSIX only — it relies on process groups, same scope as the killpg-based orphan cleanup). Unchanged on non-POSIX or if the PID cannot be read — watchdog bookkeeping must never block a connection.""" - if os.name != "posix": - return command, args try: + if os.name != "posix": + return command, args my_pid = os.getpid() except Exception: return command, args @@ -205,15 +189,15 @@ def _interpolate_env_vars(value): return value -# (server_name, dotted key path) pairs already warned about; config loads -# happen on every discovery pass, so warn once per process. +# (server_name, dotted key path) pairs already warned about: config loads happen on every discovery +# pass, so warn once per process. _whitespace_warned: Set[Tuple[str, str]] = set() def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: """Warn once per (server, key path) about string values with leading/trailing whitespace (a - pasted newline causes opaque auth failures, invisible in config.yaml). Advisory only: values - are never mutated (could be intentional) nor logged (often secrets). Returns flagged paths.""" + pasted newline causes opaque auth failures). Advisory only: values are never mutated (could be + intentional) nor logged (often secrets). Returns flagged paths.""" flagged: List[str] = [] def _walk(value: Any, path: str) -> None: @@ -228,15 +212,12 @@ def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: _walk(config, "") for key_path in flagged: - if (server_name, key_path) in _whitespace_warned: - continue - _whitespace_warned.add((server_name, key_path)) - logger.warning( - "MCP server '%s': config value '%s' has hidden leading or " - "trailing whitespace — this often causes authentication or " - "connection failures. Check for stray spaces/newlines in " - "config.yaml (or the referenced env var).", - server_name, key_path) + if (server_name, key_path) not in _whitespace_warned: + _whitespace_warned.add((server_name, key_path)) + logger.warning( + "MCP server '%s': config value '%s' has hidden leading or trailing whitespace — this often " + "causes authentication or connection failures. Check for stray spaces/newlines in config.yaml " + "(or the referenced env var).", server_name, key_path) return flagged @@ -246,14 +227,13 @@ def _filter_suspicious_mcp_servers(servers: Dict[str, dict]) -> Dict[str, dict]: from hermes_cli.mcp_security import validate_mcp_server_entry except Exception: return servers - safe_servers = {} for name, cfg in servers.items(): issues = validate_mcp_server_entry(name, cfg) if isinstance(cfg, dict) else None if issues: logger.warning("Skipping suspicious MCP server '%s': %s", name, "; ".join(issues)) - continue - safe_servers[name] = cfg + else: + safe_servers[name] = cfg return safe_servers @@ -267,8 +247,8 @@ def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None: for name, cfg in _core._filter_suspicious_mcp_servers(portable).items(): if name in safe_servers: logger.warning("Portable MCP server '%s' conflicts with native config; skipping", name) - continue - safe_servers[name] = dict(cfg) + else: + safe_servers[name] = dict(cfg) except Exception: logger.debug("Failed to load portable MCP servers", exc_info=True) diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index 91ff039518..9ab2e50e89 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -19,24 +19,20 @@ logger = logging.getLogger("tools.mcp_tool") _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022 -def _jsonrpc_code(exc: BaseException): - """Structural ``MCPError.error.code`` (None when absent).""" - return getattr(getattr(exc, "error", None), "code", None) - - -def _jsonrpc_matches(exc: BaseException, code, codes: tuple, markers: tuple) -> bool: - """Structural *code* in *codes*, else any *marker* in ``str(exc).lower()``. Never ``isinstance`` - on SDK exception types: they arrive wrapped in ExceptionGroups and drift across generations.""" +def _jsonrpc_matches(exc: BaseException, codes: tuple, markers: tuple, code=None) -> bool: + """Structural ``MCPError.error.code`` (or *code*) in *codes*, else any *marker* in + ``str(exc).lower()``. Never ``isinstance`` on SDK exception types: they arrive wrapped in + ExceptionGroups and drift across generations.""" + code = getattr(getattr(exc, "error", None), "code", None) or code return code in codes or any(marker in str(exc).lower() for marker in markers) def _handshake_rejected_as_modern(exc: BaseException) -> bool: """True when a failed ``initialize`` signals a stateless-only (2026-07-28) server.""" return _jsonrpc_matches( - exc, _jsonrpc_code(exc) or getattr(exc, "code", None), - (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND), + exc, (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND), ("unsupported protocol version", str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION)), - ) or _is_method_not_found_error(exc) + code=getattr(exc, "code", None)) or _is_method_not_found_error(exc) def _is_method_not_found_error(exc: BaseException) -> bool: @@ -44,7 +40,7 @@ def _is_method_not_found_error(exc: BaseException) -> bool: substring fallback includes "Unknown method: " — without it the ping→list_tools keepalive fallback never latches and reconnect-loops.""" return _jsonrpc_matches( - exc, _jsonrpc_code(exc), (_core._JSONRPC_METHOD_NOT_FOUND,), + exc, (_core._JSONRPC_METHOD_NOT_FOUND,), (str(_core._JSONRPC_METHOD_NOT_FOUND), "method not found", "unknown method", "not found: ping")) @@ -138,8 +134,7 @@ def _resolve_client_cert(server_name: str, config: dict): cert_path = _expand(raw_cert, "client_cert") return (cert_path, _expand(raw_key, "client_key")) if raw_key is not None else cert_path # combined PEM if raw_key is not None: - raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR " - f"client_cert + client_key, not both") + raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR client_cert + client_key, not both") if len(raw_cert) not in (2, 3): raise ValueError(f"{prefix}client_cert list form must have 2 or 3 elements (got {len(raw_cert)})") pair = (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]")) @@ -248,11 +243,6 @@ def _format_connect_error(exc: BaseException) -> str: return _sanitize_error(message) -# Lazily-built caches so this module imports without the SDK OAuth module. -_AUTH_ERROR_TYPES: tuple = () -_HTTP_STATUS_ERROR_TYPES: Optional[tuple] = None - - def _optional_types(module: str, *names: str) -> list: """``[module.name, ...]`` or ``[]`` when the module/attribute is unavailable.""" try: @@ -262,35 +252,33 @@ def _optional_types(module: str, *names: str) -> list: return [] -def _http_status_error_types() -> tuple: - """``HTTPStatusError`` from both httpx flavours: a 401 may come from the SDK's own stack - (``httpx2`` on mcp >= 2.0) or Hermes' pinned ``httpx``; the classes are unrelated.""" - global _HTTP_STATUS_ERROR_TYPES - if _HTTP_STATUS_ERROR_TYPES is None: - sdk_mod = _core.sdk_httpx() - _HTTP_STATUS_ERROR_TYPES = tuple(dict.fromkeys( - ([sdk_mod.HTTPStatusError] if sdk_mod is not None else []) + _optional_types("httpx", "HTTPStatusError"))) - return _HTTP_STATUS_ERROR_TYPES +# Lazily-built ``(auth_types, http_status_types)`` so this module imports without the SDK OAuth module. +_AUTH_ERROR_TYPES: Optional[tuple] = None def _get_auth_error_types() -> tuple: - """Cached MCP OAuth failure types: SDK ``OAuthFlowError``/``OAuthTokenError`` (+ legacy - ``UnauthorizedError``), our ``OAuthNonInteractiveError``, and both ``HTTPStatusError`` flavours - (which still need the 401 check in :func:`_is_auth_error`).""" + """Cached ``(auth_types, http_status_types)``: SDK ``OAuthFlowError``/``OAuthTokenError`` (+ legacy + ``UnauthorizedError``), our ``OAuthNonInteractiveError``, and ``HTTPStatusError`` from both httpx + flavours — a 401 may come from the SDK's own stack (``httpx2`` on mcp >= 2.0) or Hermes' pinned + ``httpx``; the classes are unrelated and still need the 401 check in :func:`_is_auth_error`.""" global _AUTH_ERROR_TYPES - if not _AUTH_ERROR_TYPES: - _AUTH_ERROR_TYPES = (*_optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError"), - *_optional_types("mcp.client.auth", "UnauthorizedError"), # older SDKs - *_optional_types("tools.mcp_oauth", "OAuthNonInteractiveError"), - *_http_status_error_types()) + if not (_AUTH_ERROR_TYPES and _AUTH_ERROR_TYPES[0]): # retry while empty (SDK may import later) + sdk_mod = _core.sdk_httpx() + http_types = tuple(dict.fromkeys( + ([sdk_mod.HTTPStatusError] if sdk_mod is not None else []) + _optional_types("httpx", "HTTPStatusError"))) + auth_types = (*_optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError"), + *_optional_types("mcp.client.auth", "UnauthorizedError"), # older SDKs + *_optional_types("tools.mcp_oauth", "OAuthNonInteractiveError"), *http_types) + _AUTH_ERROR_TYPES = (auth_types, http_types) return _AUTH_ERROR_TYPES def _is_auth_error(exc: BaseException) -> bool: """True if ``exc`` indicates an MCP OAuth failure; ``HTTPStatusError`` counts only with status 401.""" - if not isinstance(exc, _get_auth_error_types()): + auth_types, http_types = _get_auth_error_types() + if not isinstance(exc, auth_types): return False - return getattr(exc.response, "status_code", None) == 401 if isinstance(exc, _http_status_error_types()) else True + return getattr(exc.response, "status_code", None) == 401 if isinstance(exc, http_types) else True # Lower-cased substrings meaning the transport session expired / was GC'd (OAuth token still valid). diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 1a5695800c..832d0c7066 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -20,6 +20,16 @@ from tools.mcp_tool_content import ( from tools.mcp_tool_errors import _is_session_expired_error logger = logging.getLogger("tools.mcp_tool") +_MISSING = object() + +_NEEDS_REAUTH_MSG = ("MCP server '{s}' requires re-authentication. Run `hermes mcp login {s}` (or delete the tokens " + "file under ~/.hermes/mcp-tokens/ and restart). Do NOT retry this tool — ask the user to re-authenticate.") +_STDIO_NO_RESPAWN_MSG = ("MCP server '{s}' stdio subprocess had exited (this is not a timeout — the call never reached the " + "server). A respawn was requested but no fresh session came back within {t:.0f}s. Wait a few " + "seconds before retrying; if it keeps failing the server is not starting and needs the user.") +_STDIO_DIED_AGAIN_MSG = ("MCP server '{s}' respawned its stdio subprocess and it exited again immediately. The server is not " + "starting cleanly — do NOT retry this tool; ask the user to check the server's command and its " + "stderr log.") # --------------------------------------------------------------- pre-call gates @@ -30,8 +40,7 @@ def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]: trust = _core._server_trust_levels.get(server_name, _core._TRUST_FULL) if trust != _core._TRUST_UNTRUSTED or _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True: return None - # Lazy import: tools.approval routes the prompt to whichever surface owns the session. - try: + try: # lazy: tools.approval routes the prompt to whichever surface owns the session from tools.approval import request_elicitation_consent answer = request_elicitation_consent( f"MCP tool '{tool_name}' on UNTRUSTED server '{server_name}' wants to run. This tool is write-capable " @@ -55,15 +64,12 @@ def _check_circuit_breaker(server_name: str) -> Optional[str]: """Open-breaker error, or None when calls may proceed. After the cooldown the breaker is half-open: the next call probes; success resets, failure re-bumps and re-arms the cooldown.""" failures = _core._server_error_counts.get(server_name, 0) - if failures < _core._CIRCUIT_BREAKER_THRESHOLD: - return None age = time.monotonic() - _core._server_breaker_opened_at.get(server_name, 0.0) - if age >= _core._CIRCUIT_BREAKER_COOLDOWN_SEC: + if failures < _core._CIRCUIT_BREAKER_THRESHOLD or age >= _core._CIRCUIT_BREAKER_COOLDOWN_SEC: return None - remaining = max(1, int(_core._CIRCUIT_BREAKER_COOLDOWN_SEC - age)) return tool_error(f"MCP server '{server_name}' is unreachable after {failures} consecutive failures. " - f"Auto-retry available in ~{remaining}s. Do NOT retry this tool yet — use alternative " - f"approaches or ask the user to check the MCP server.") + f"Auto-retry available in ~{max(1, int(_core._CIRCUIT_BREAKER_COOLDOWN_SEC - age))}s. Do NOT retry " + f"this tool yet — use alternative approaches or ask the user to check the MCP server.") def _acquire_call_server(server_name: str, tool_timeout: float): @@ -72,13 +78,10 @@ def _acquire_call_server(server_name: str, tool_timeout: float): server task to rebuild (probing a dead transport would re-arm the breaker forever).""" not_connected = tool_error(f"MCP server '{server_name}' is not connected") server = _core._get_connected_server_for_call(server_name) - if not server: - _core._bump_server_error(server_name) - return None, not_connected - if server.session or _core._wait_for_server_session_ready(server, timeout=min(5.0, float(tool_timeout or 5.0))): + if server and (server.session or _core._wait_for_server_session_ready(server, timeout=min(5.0, float(tool_timeout or 5.0)))): return server, None _core._bump_server_error(server_name) - if _core._signal_reconnect(server): + if server and _core._signal_reconnect(server): return None, tool_error(f"MCP server '{server_name}' transport is down; reconnect requested. Do NOT retry this " f"tool immediately — give it a few seconds to come back.") return None, not_connected @@ -96,10 +99,7 @@ def _result_is_error(result) -> bool: def _record_call_outcome(server_name: str, result) -> Any: """Breaker bookkeeping: an error payload from the tool itself still counts as a strike.""" - if _result_is_error(result): - _core._bump_server_error(server_name) - else: - _core._reset_server_error(server_name) + (_core._bump_server_error if _result_is_error(result) else _core._reset_server_error)(server_name) return result @@ -109,18 +109,17 @@ def _strike(server_name: str, message: str, **extra) -> str: return tool_error(message, **extra) +def _mcp_loop_running() -> bool: + return _core._mcp_loop is not None and _core._mcp_loop.is_running() + + def _lookup_reconnectable_server(server_name: str, require_loop: bool = False): """The registered server object when it can be signalled to reconnect, else None. With *require_loop*, also None unless the MCP loop is running (nothing to wait on).""" with _core._lock: srv = _core._servers.get(server_name) - if srv is None or not hasattr(srv, "_reconnect_event") or (require_loop and not _mcp_loop_running()): - return None - return srv - - -def _mcp_loop_running() -> bool: - return _core._mcp_loop is not None and _core._mcp_loop.is_running() + ok = srv is not None and hasattr(srv, "_reconnect_event") and (_mcp_loop_running() or not require_loop) + return srv if ok else None def _retry_once(server_name: str, retry_call, op_description: str, what: str): @@ -146,9 +145,8 @@ def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_cal if not _core._is_auth_error(exc): return None from tools.mcp_oauth_manager import get_manager - manager = get_manager() try: - recovered = _core._run_on_mcp_loop(lambda: manager.handle_401(server_name, None), timeout=10) + recovered = _core._run_on_mcp_loop(lambda: get_manager().handle_401(server_name, None), timeout=10) except Exception as rec_exc: logger.warning("MCP OAuth '%s': recovery attempt failed: %s", server_name, rec_exc) recovered = False @@ -162,20 +160,13 @@ def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_cal result = _retry_once(server_name, retry_call, op_description, "auth recovery") if result is not None: return result - return _strike( - server_name, - f"MCP server '{server_name}' requires re-authentication. Run `hermes mcp login " - f"{server_name}` (or delete the tokens file under ~/.hermes/mcp-tokens/ and restart). Do " - f"NOT retry this tool — ask the user to re-authenticate.", - needs_reauth=True, server=server_name) + return _strike(server_name, _NEEDS_REAUTH_MSG.format(s=server_name), needs_reauth=True, server=server_name) def _handle_session_expired_and_retry(server_name: str, exc: BaseException, retry_call, op_description: str): """Transport reconnect + one retry on session expiry; None to fall through. Skips ``handle_401``: the token is valid, only the server-side session is stale.""" - if not _is_session_expired_error(exc): - return None - srv = _lookup_reconnectable_server(server_name, require_loop=True) + srv = _lookup_reconnectable_server(server_name, require_loop=True) if _is_session_expired_error(exc) else None if srv is None: return None logger.info("MCP server '%s': %s failed with session-expired error (%s); signalling transport reconnect " @@ -205,27 +196,17 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry if _mcp_loop_running(): reconnected = _core._signal_reconnect_and_wait( server_name, srv, op_description=op_description, timeout=_core._STDIO_RESPAWN_WAIT_SEC) - else: - # No MCP loop to wait on (non-async adapters, tests): still request the respawn. + else: # No MCP loop to wait on (non-async adapters, tests): still request the respawn. _core._signal_reconnect(srv) if not reconnected: - return _strike( - server_name, - f"MCP server '{server_name}' stdio subprocess had exited (this is not a timeout — the " - f"call never reached the server). A respawn was requested but no fresh session came " - f"back within {_core._STDIO_RESPAWN_WAIT_SEC:.0f}s. Wait a few seconds before retrying; " - f"if it keeps failing the server is not starting and needs the user.") + return _strike(server_name, _STDIO_NO_RESPAWN_MSG.format(s=server_name, t=_core._STDIO_RESPAWN_WAIT_SEC)) try: return _record_call_outcome(server_name, retry_call()) except _StdioChildExited as retry_exc: # Died again right after respawn: broken server; run()'s budget takes it to the park. logger.warning("MCP server '%s': %s stdio subprocess exited again right after respawn (%s); not retrying " "further.", server_name, op_description, retry_exc) - return _strike( - server_name, - f"MCP server '{server_name}' respawned its stdio subprocess and it exited again " - f"immediately. The server is not starting cleanly — do NOT retry this tool; ask the " - f"user to check the server's command and its stderr log.") + return _strike(server_name, _STDIO_DIED_AGAIN_MSG.format(s=server_name)) except Exception as retry_exc: logger.warning("MCP %s/%s retry after stdio respawn failed: %s", server_name, op_description, retry_exc) return _strike(server_name, _sanitize_error( @@ -254,14 +235,18 @@ def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: st return tool_error(_sanitize_error(f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}")) -# ------------------------------------------------------------- the RPC itself - -def _mark_server_call_started(server: Any) -> None: - """Record a user-visible MCP operation when the server supports it.""" +def _dispatch(server_name: str, server: Any, op: str, call, tool_timeout: float, recoverers, + on_final_failure: Callable[[BaseException], None], record_outcome: bool = False) -> str: + """Mark the call started on *server* (doubles may lack ``mark_tool_call``), run coroutine + function *call* on the MCP loop and walk the recovery ladder (:func:`_invoke_with_recovery`).""" if callable(getattr(server, "mark_tool_call", None)): server.mark_tool_call() + return _invoke_with_recovery(server_name, lambda: _core._run_on_mcp_loop(call, timeout=tool_timeout), op, + recoverers, on_final_failure, record_outcome=record_outcome) +# ------------------------------------------------------------- the RPC itself + @asynccontextmanager async def _track_inflight_rpc(server: Any, server_name: str, op: str): """Register the running RPC so teardown can fail it fast. A deliberate teardown @@ -285,10 +270,9 @@ async def _track_inflight_rpc(server: Any, server_name: str, op: str): async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str, args: dict): """``session.call_tool`` that fails fast when the stdio child is/gets dead: pre-call (a dead - child must not hold the slot for the full timeout; ``server.session`` is stale) and mid-call - (race against ``_watch_stdio_children``). Both raise :class:`_StdioChildExited` for the - respawn path, which owns the reconnect signal. callable()/``is True`` because MagicMock - attributes are truthy.""" + child must not hold the slot for the full timeout) and mid-call (race against + ``_watch_stdio_children``). Both raise :class:`_StdioChildExited` for the respawn path, which + owns the reconnect signal. callable()/``is True`` because MagicMock attributes are truthy.""" _stdio_dead = getattr(server, "_stdio_children_dead", None) if callable(_stdio_dead) and _stdio_dead() is True: raise _StdioChildExited(f"MCP stdio subprocess for '{server_name}' had already exited when the call was dispatched") @@ -339,8 +323,7 @@ def _render_content_blocks(result, server_name: str) -> str: logger.debug("MCP %s: content block type %r rendered empty", server_name, block_type) else: logger.warning("MCP %s: dropping unsupported content block type %r", server_name, block_type) - # Hard-cap pathological payloads; ordinary large results pass to spillover. - return _truncate_mcp_text_result("\n".join(parts)) + return _truncate_mcp_text_result("\n".join(parts)) # hard-cap pathological payloads; spillover handles the rest def _capped_structured_content(result): @@ -374,8 +357,7 @@ def _render_call_tool_result(result, server_name: str) -> str: payload.setdefault("result", text_result) try: return json.dumps(payload, ensure_ascii=False) - except (TypeError, ValueError): - # Non-serializable metadata: drop the extras, keep the call. + except (TypeError, ValueError): # Non-serializable metadata: drop the extras, keep the call. return json.dumps({"result": text_result}, ensure_ascii=False) @@ -386,8 +368,7 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): op = f"tools/call {tool_name}" def _handler(args: dict, **kwargs) -> str: - # Security boundary: untrusted-server write tools need approval before ANY transport work - # (including the lazy first-use spawn). + # Security boundary: untrusted-server write tools need approval before ANY transport work (incl. lazy spawn). error = _trust_gate_check(server_name, tool_name) or _check_circuit_breaker(server_name) if error is not None: return error @@ -396,16 +377,13 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): return error async def _call(): - _mark_server_call_started(server) async with server._rpc_lock, _track_inflight_rpc(server, server_name, op): - # Snapshot contextvars for the elicitation callback (MCP recv loop doesn't inherit them). - server._pending_call_context = contextvars.copy_context() + server._pending_call_context = contextvars.copy_context() # for the elicitation callback try: result = await _call_tool_racing_stdio_death(server, server_name, tool_name, args) finally: server._pending_call_context = None - # Round-trip completed: transport is healthy even if the tool returned isError. - if getattr(server, "_mark_session_proven", None) is not None: + if getattr(server, "_mark_session_proven", None) is not None: # round-trip done: transport healthy server._mark_session_proven() return _render_call_tool_result(result, server_name) @@ -413,39 +391,39 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): _core._bump_server_error(server_name) logger.error("MCP tool %s/%s call failed: %s", server_name, tool_name, exc) - return _invoke_with_recovery( - server_name, lambda: _core._run_on_mcp_loop(_call, timeout=tool_timeout), op, + return _dispatch( + server_name, server, op, _call, tool_timeout, (_handle_stdio_child_exited_and_retry, _handle_auth_error_and_retry, _handle_session_expired_and_retry), _on_failure, record_outcome=True) return _handler -def _make_utility_handler(server_name: str, tool_timeout: float, op: str, log_label: str, - rpc, render, required: Optional[str] = None): - """Shared shape of the four utility handlers: ``rpc(session, args, server_name)`` awaited - under ``_rpc_lock``, ``render(result, server_name)`` -> JSON-able payload, ``required`` - validated before any transport work; owns the connected check and recovery ladder.""" +def _make_utility_handler(op: str, log_label: str, rpc, render, required: Optional[str] = None): + """``(server_name, tool_timeout) -> sync handler`` for one utility tool: ``rpc(session, args, + server_name)`` awaited under ``_rpc_lock``, ``render(result, server_name)`` -> JSON-able + payload, ``required`` validated before any transport work.""" + def _factory(server_name: str, tool_timeout: float): + def _handler(args: dict, **kwargs) -> str: + server = _core._get_connected_server_for_call(server_name) + if not server or not server.session: + return tool_error(f"MCP server '{server_name}' is not connected") + if required and not args.get(required): + return tool_error(f"Missing required parameter '{required}'") - def _handler(args: dict, **kwargs) -> str: - server = _core._get_connected_server_for_call(server_name) - if not server or not server.session: - return tool_error(f"MCP server '{server_name}' is not connected") - if required and not args.get(required): - return tool_error(f"Missing required parameter '{required}'") + async def _call(): + async with server._rpc_lock: + result = await rpc(server.session, args, server_name) + return json.dumps(render(result, server_name), ensure_ascii=False) - async def _call(): - _mark_server_call_started(server) - async with server._rpc_lock: - result = await rpc(server.session, args, server_name) - return json.dumps(render(result, server_name), ensure_ascii=False) + return _dispatch( + server_name, server, op, _call, tool_timeout, + (_handle_auth_error_and_retry, _handle_session_expired_and_retry), + lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc)) - return _invoke_with_recovery( - server_name, lambda: _core._run_on_mcp_loop(_call, timeout=tool_timeout), op, - (_handle_auth_error_and_retry, _handle_session_expired_and_retry), - lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc)) + return _handler - return _handler + return _factory def _pick(obj, *specs) -> dict: @@ -453,10 +431,8 @@ def _pick(obj, *specs) -> dict: so SDK models and stubs behave alike; ``truthy`` also skips falsy). Key order = spec order.""" entry = {} for out_key, attr, *truthy in specs: - if not hasattr(obj, attr): - continue - value = getattr(obj, attr) - if value or not (truthy and truthy[0]): + value = getattr(obj, attr, _MISSING) + if value is not _MISSING and (value or not (truthy and truthy[0])): entry[out_key] = value return entry @@ -479,8 +455,7 @@ def _render_read_resource(result, server_name: str) -> dict: for block in getattr(result, "contents", []): if getattr(block, "text", None) is not None: parts.append(strip_unicode_tags(block.text)) - elif getattr(block, "blob", None) is not None: - # Binary contents go to the document cache (same contract as EmbeddedResource blocks). + elif getattr(block, "blob", None) is not None: # binary -> document cache, like EmbeddedResource blocks rendered = _render_mcp_resource_block(SimpleNamespace(type="resource", resource=block), server_name) parts.append(rendered or f"[binary data, {len(block.blob)} bytes]") return {"result": "\n".join(parts)} @@ -507,41 +482,28 @@ def _render_get_prompt(result, server_name: str) -> dict: return {"messages": messages, **_pick(result, ("description", "description", True))} -def _utility_factory(op: str, log_label: str, rpc, render, required: Optional[str] = None): - """``(server_name, tool_timeout) -> sync handler`` for one utility tool.""" - def _factory(server_name: str, tool_timeout: float): - return _make_utility_handler(server_name, tool_timeout, op, log_label, rpc, render, required) - - return _factory - - -_make_list_resources_handler = _utility_factory( +_make_list_resources_handler = _make_utility_handler( "resources/list", "list_resources", - lambda session, args, sn: _core._paginate_full_list(session.list_resources, "resources", sn), - _render_resource_list) -_make_read_resource_handler = _utility_factory( + lambda session, args, sn: _core._paginate_full_list(session.list_resources, "resources", sn), _render_resource_list) +_make_read_resource_handler = _make_utility_handler( "resources/read", "read_resource", lambda session, args, sn: session.read_resource(args["uri"]), _render_read_resource, required="uri") -_make_list_prompts_handler = _utility_factory( +_make_list_prompts_handler = _make_utility_handler( "prompts/list", "list_prompts", - lambda session, args, sn: _core._paginate_full_list(session.list_prompts, "prompts", sn), - _render_prompt_list) -_make_get_prompt_handler = _utility_factory( + lambda session, args, sn: _core._paginate_full_list(session.list_prompts, "prompts", sn), _render_prompt_list) +_make_get_prompt_handler = _make_utility_handler( "prompts/get", "get_prompt", lambda session, args, sn: session.get_prompt(args["name"], arguments=args.get("arguments", {})), _render_get_prompt, required="name") def _make_check_fn(server_name: str): - """Check function that verifies the MCP connection is alive.""" - + """Check function that verifies the MCP connection is alive. Lazy (schema-cache registered) + servers count as available: the first real call spawns/connects them.""" def _check() -> bool: with _core._lock: server = _core._servers.get(server_name) - if server is not None and (server.session is not None or server._is_recycled_stdio()): - return True - # Lazy (schema-cache registered) servers count as available: the first real - # call spawns/connects them. - return server_name in _core._lazy_server_configs + return ((server is not None and (server.session is not None or server._is_recycled_stdio())) + or server_name in _core._lazy_server_configs) return _check diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index d213d7df0e..c90db405b2 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -1,6 +1,6 @@ """Session health for MCPServerTask: dynamic tool refresh on list_changed notifications, server log forwarding, keepalive probes, suspect-mark / lazy-verify, in-flight call fail-fast, stdio -child liveness and stdio idle/lifetime recycling. Split from tools/mcp_tool.py.""" +child liveness and stdio idle/lifetime recycling.""" import asyncio import json @@ -88,8 +88,7 @@ class MCPServerHealthMixin: if len(data) > 2000: # cap payloads so a chatty server can't flood agent.log data = data[:2000] + "... [truncated]" logger_name = getattr(params, "logger", None) - origin = f"{self.name}/{logger_name}" if logger_name else self.name - logger.log(level, "MCP server log [%s]: %s", origin, data) + logger.log(level, "MCP server log [%s]: %s", f"{self.name}/{logger_name}" if logger_name else self.name, data) except Exception: logger.debug("Failed to handle MCP log notification from '%s'", self.name, exc_info=True) return _on_log @@ -125,12 +124,10 @@ class MCPServerHealthMixin: """Deregister *tool_names* this server's toolset still owns. Never removes a colliding name currently owned by another server.""" from tools.registry import registry - toolset_name = f"mcp-{self.name}" for tool_name in tool_names: - if registry.get_toolset_for_tool(tool_name) != toolset_name: - continue - registry.deregister(tool_name, scope=_core._server_registry_scope(self.name)) - _forget_mcp_tool_server(tool_name) + if registry.get_toolset_for_tool(tool_name) == f"mcp-{self.name}": + registry.deregister(tool_name, scope=_core._server_registry_scope(self.name)) + _forget_mcp_tool_server(tool_name) async def _refresh_tools(self): """Re-fetch tools on ``tools/list_changed`` and update the registry. The lock serializes @@ -165,6 +162,9 @@ class MCPServerHealthMixin: """Exercise the session; raise on a genuine connection failure. ``ping`` first (cheap, OPTIONAL); on -32601 latch ``_ping_unsupported`` (reset per transport connection) and fall back to ``list_tools`` when the server advertises tools, else the -32601 propagates.""" + async def list_tools(): + await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT) + if not self._ping_unsupported: try: await asyncio.wait_for(self.session.send_ping(), timeout=_KEEPALIVE_RPC_TIMEOUT) @@ -180,7 +180,7 @@ class MCPServerHealthMixin: # A server that silently drops ping looks like a dead transport: confirm with # list_tools before declaring it dead, else propagate the original failure. try: - await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT) + await list_tools() except Exception: raise exc from None self._ping_unsupported = True # latch so later keepalives skip the 30s wait @@ -189,7 +189,7 @@ class MCPServerHealthMixin: return else: raise # closed transport, expired session, etc. — real failure - await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT) + await list_tools() def _mark_session_proven(self) -> None: """Record that the session demonstrated real health (keepalive or tool-call success). @@ -204,8 +204,7 @@ class MCPServerHealthMixin: logger.warning("MCP server '%s': revived — session healthy again after " "parking (state: parked → connected)", self.name) # A proven fresh transport clears the one-time permanent-failure grace and any race bookkeeping. - self._permanent_grace_used = False - self._teardown_race = False + self._permanent_grace_used = self._teardown_race = False def mark_suspect(self, reason: str) -> None: """Latch a suspicion (no I/O). The NEXT call verifies via :meth:`ensure_healthy` and @@ -253,8 +252,7 @@ class MCPServerHealthMixin: victims = [t for t in self._inflight_tasks if not t.done()] if not victims: return - self._reconnecting = True - self._teardown_race = True + self._reconnecting = self._teardown_race = True self.mark_suspect(f"{reason} tore down {len(victims)} in-flight call(s)") for task in victims: task.cancel() diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index b9c994df5a..7595873656 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -1,8 +1,7 @@ """Registering a connected (or schema-cached) MCP server's tools into the tool registry: -include/exclude filtering, trust-tier metadata capture, utility-tool selection, -name-collision resolution and the schema-cache write-through. Both entry points -(``_register_server_tools`` live, ``_register_from_cache_sync`` lazy) build ``_Candidate`` -records and feed the single ``_register_candidates`` loop.""" +include/exclude filtering, trust-tier metadata capture, utility-tool selection, name-collision +resolution and the schema-cache write-through. Both entry points (``_register_server_tools`` +live, ``_register_from_cache_sync`` lazy) build ``_Candidate`` records for ``_register_candidates``.""" import logging from dataclasses import dataclass @@ -34,8 +33,7 @@ def _normalize_server_trust(value: Any) -> str: text = str(value).strip().lower() if text in (_core._TRUST_FULL, _core._TRUST_UNTRUSTED): return text - logger.warning("MCP trust: unrecognized trust value %r — treating as 'untrusted' (valid values: full, untrusted)", - value) + logger.warning("MCP trust: unrecognized trust value %r — treating as 'untrusted' (valid values: full, untrusted)", value) return _core._TRUST_UNTRUSTED @@ -43,9 +41,8 @@ def _annotation_read_only_hint(mcp_tool: Any) -> bool: """True only when annotations (SDK object or schema-cache dict) carry ``readOnlyHint is True``; unknown metadata means write-capable.""" annotations = getattr(mcp_tool, "annotations", None) - if isinstance(annotations, dict): - return annotations.get("readOnlyHint") is True - return getattr(annotations, "readOnlyHint", None) is True + hint = annotations.get("readOnlyHint") if isinstance(annotations, dict) else getattr(annotations, "readOnlyHint", None) + return hint is True def _record_tool_trust_metadata(server_name: str, config: dict, tools: List[Any]) -> None: @@ -55,10 +52,7 @@ def _record_tool_trust_metadata(server_name: str, config: dict, tools: List[Any] with _core._lock: _core._server_trust_levels[server_name] = _normalize_server_trust((config or {}).get("trust")) hints = _core._tool_read_only_hints.setdefault(server_name, {}) - for tool in tools: - name = getattr(tool, "name", None) - if name: - hints[name] = _annotation_read_only_hint(tool) + hints.update({t.name: _annotation_read_only_hint(t) for t in tools if getattr(t, "name", None)}) def _track_mcp_tool_server(tool_name: str, server_name: str) -> None: @@ -80,17 +74,14 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d filters anything since ClientSession defines all four methods.""" tools_filter = config.get("tools") or {} enabled = {f: _parse_boolish(tools_filter.get(f), default=True) for f in ("resources", "prompts")} - init_result = getattr(server, "initialize_result", None) - advertised = getattr(init_result, "capabilities", None) if init_result is not None else None + advertised = getattr(getattr(server, "initialize_result", None), "capabilities", None) def _skip_reason(handler_key: str) -> Optional[str]: family = _UTILITY_CAPABILITY_ATTRS[handler_key] if not enabled[family]: return f"{family} disabled" if advertised is not None: - if getattr(advertised, family, None) is None: - return f"server does not advertise '{family}' capability" - return None + return None if getattr(advertised, family, None) is not None else f"server does not advertise '{family}' capability" # Legacy gate (no initialize_result): the ClientSession method shares the handler key. return None if hasattr(server.session, handler_key) else f"session lacks {handler_key}" @@ -176,8 +167,7 @@ def _tool_candidates(name: str, tools: Iterable[Any], should_register: Callable[ continue _core._scan_mcp_description(name, t.name, t.description or "") schema = _core._convert_mcp_schema(name, t) - handler = _core._make_tool_handler(name, t.name, tool_timeout) - out.append(_Candidate(schema["name"], f"tool {t.name!r}", schema, handler)) + out.append(_Candidate(schema["name"], f"tool {t.name!r}", schema, _core._make_tool_handler(name, t.name, tool_timeout))) return out @@ -187,8 +177,8 @@ def _utility_candidates(name: str, entries: Iterable[Any], tool_timeout) -> List for raw in entries: schema, key = (raw.get("schema"), raw.get("handler_key")) if isinstance(raw, dict) else (None, None) if isinstance(schema, dict) and key in _UTILITY_HANDLER_FACTORIES and schema.get("name"): - handler = _UTILITY_HANDLER_FACTORIES[key](name, tool_timeout) - out.append(_Candidate(schema["name"], f"{_UTILITY_ORIGIN_PREFIX}{key!r}", schema, handler)) + out.append(_Candidate(schema["name"], f"{_UTILITY_ORIGIN_PREFIX}{key!r}", schema, + _UTILITY_HANDLER_FACTORIES[key](name, tool_timeout))) return out @@ -197,16 +187,15 @@ def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_C a native tool's name is shadowed (native wins); any other multi-origin collision skips every colliding entry (fail closed). Returns survivors in order.""" unique: List[_Candidate] = [] - seen: set[tuple[str, str]] = set() origins_by_name: Dict[str, set[str]] = {} for c in candidates: - if (c.registry_name, c.origin) in seen: + origins = origins_by_name.setdefault(c.registry_name, set()) + if c.origin in origins: logger.debug("MCP server '%s': duplicate registration candidate %s for '%s'; keeping one", name, c.origin, c.registry_name) continue - seen.add((c.registry_name, c.origin)) + origins.add(c.origin) unique.append(c) - origins_by_name.setdefault(c.registry_name, set()).add(c.origin) ambiguous: Dict[str, List[str]] = {} shadowed: set[tuple[str, str]] = set() for registry_name, origins in origins_by_name.items(): @@ -220,8 +209,8 @@ def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_C "MCP server '%s': generated utility %s normalizes onto server-native %s — keeping the native tool " "and dropping the utility (the utility only applies when the server has no such tool of its own)", name, ", ".join(utility_origins), native_origins[0]) - continue - ambiguous[registry_name] = sorted(origins) + else: + ambiguous[registry_name] = sorted(origins) for registry_name, origins in sorted(ambiguous.items()): logger.error("MCP server '%s': name normalization collision for '%s' from %s; skipping every colliding " "entry instead of choosing an arbitrary handler", name, registry_name, ", ".join(origins)) @@ -234,8 +223,7 @@ def _log_foreign_owner(name: str, c: _Candidate, existing_toolset: str, lazy: bo if not c.is_utility: logger.warning("MCP server '%s' (lazy): cached tool '%s' collides with toolset '%s' — skipping", name, c.registry_name, existing_toolset) - return - if existing_toolset.startswith("mcp-"): + elif existing_toolset.startswith("mcp-"): logger.error("MCP server '%s': %s normalizes to '%s', already owned by MCP toolset '%s' — skipping to " "preserve the existing owner", name, c.origin, c.registry_name, existing_toolset) else: @@ -259,13 +247,12 @@ def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: C registry.register( name=c.registry_name, toolset=toolset_name, schema=c.schema, handler=c.handler, check_fn=check_fn, is_async=False, description=c.schema.get("description") or "", scope=scope()) - if registry.get_toolset_for_tool(c.registry_name) != toolset_name: - if not lazy: - logger.error("MCP server '%s': registration of %s as '%s' was rejected by the registry; " - "skipping provenance/count updates", name, c.origin, c.registry_name) - continue - _core._track_mcp_tool_server(c.registry_name, name) - registered.append(c.registry_name) + if registry.get_toolset_for_tool(c.registry_name) == toolset_name: + _core._track_mcp_tool_server(c.registry_name, name) + registered.append(c.registry_name) + elif not lazy: + logger.error("MCP server '%s': registration of %s as '%s' was rejected by the registry; " + "skipping provenance/count updates", name, c.origin, c.registry_name) if registered: registry.register_toolset_alias(name, toolset_name) return registered @@ -279,8 +266,7 @@ def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should tools_payload = [{ "name": t.name, "description": t.description or "", "inputSchema": t.inputSchema if isinstance(getattr(t, "inputSchema", None), dict) else {}, - # Persisted so the lazy path trust-gates identically next startup. - "annotations": {"readOnlyHint": _annotation_read_only_hint(t)}, + "annotations": {"readOnlyHint": _annotation_read_only_hint(t)}, # lazy path trust-gates identically } for t in server._tools if should_register(t.name)] utility_payload = [{"schema": e["schema"], "handler_key": e["handler_key"]} for e in _select_utility_schemas(name, server, config)] diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index 99eef3fd38..29721bb977 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -1,6 +1,6 @@ """Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID ledger), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth), -protocol negotiation and initial tool discovery. Split from tools/mcp_tool.py.""" +protocol negotiation and initial tool discovery.""" import logging import asyncio @@ -29,15 +29,17 @@ def _is_2xx(resp) -> bool: return 200 <= resp.status_code < 300 +def _present(**kwargs) -> dict: + """*kwargs* minus the ``None`` values (optional httpx client arguments).""" + return {k: v for k, v in kwargs.items() if v is not None} + + def _pgroup_alive(pgid: Optional[int]) -> bool: """Signal 0 to the group succeeds iff any member is alive (POSIX only).""" - _killpg = getattr(os, "killpg", None) - if pgid is None or _killpg is None: - return False try: - _killpg(pgid, 0) + os.killpg(pgid, 0) return True - except (ProcessLookupError, PermissionError, OSError): + except (AttributeError, TypeError, OSError): # non-POSIX / pgid None / gone return False @@ -63,16 +65,17 @@ class MCPServerTransportMixin: __slots__ = () def _advertises_tools(self) -> bool: - """Whether the server advertises ``tools`` (prompt-/resource-only servers omit it and - ``tools/list`` raises -32601). True when no capability info was captured (legacy fallback).""" + """False only when captured capabilities omit ``tools`` (prompt-/resource-only servers, + where ``tools/list`` raises -32601); True without capability info (legacy fallback).""" caps = getattr(self.initialize_result, "capabilities", None) return caps is None or getattr(caps, "tools", None) is not None def _session_kwargs(self) -> dict: """ClientSession kwargs: sampling, elicitation, notification + logging callbacks.""" - kwargs = self._sampling.session_kwargs() if self._sampling else {} - if self._elicitation: - kwargs.update(self._elicitation.session_kwargs()) + kwargs = {} + for handler in (self._sampling, self._elicitation): + if handler: + kwargs.update(handler.session_kwargs()) if _core._MCP_NOTIFICATION_TYPES and _core._MCP_MESSAGE_HANDLER_SUPPORTED: kwargs["message_handler"] = self._make_message_handler() if _core._MCP_LOGGING_CALLBACK_SUPPORTED: @@ -80,12 +83,11 @@ class MCPServerTransportMixin: return kwargs async def _negotiate_session(self, session, connect_timeout: float): - """Negotiate the protocol era (``initialize`` vs ``server/discover``); both results expose - ``.capabilities``. ``protocol: auto`` (default) tries the legacy handshake FIRST and falls back - to discover only on a modern-only signal (-32022 / initialize -32601) — deliberately the - reverse of the SDK's discover-first mode: zero extra round-trips for the handshake-era - servers that dominate today. ``stateless`` probes discover first (one legacy retry on any - error); ``legacy`` is handshake only. A handshake TIMEOUT never falls back — it propagates.""" + """Negotiate the protocol era (``initialize`` vs ``server/discover``; both expose + ``.capabilities``). ``auto`` tries the legacy handshake FIRST, falling back to discover only + on a modern-only signal (-32022 / initialize -32601) — the reverse of the SDK's discover-first + mode, so handshake-era servers pay zero extra round-trips. ``stateless`` probes discover first + (one legacy retry on any error); ``legacy`` is handshake only. A TIMEOUT never falls back.""" def call(method: str): return asyncio.wait_for(getattr(session, method)(), timeout=connect_timeout) @@ -111,14 +113,13 @@ class MCPServerTransportMixin: # mcp 1.x has no server/discover client — nothing to fall back to. return await attempt( "initialize", "discover", lambda exc: _handshake_rejected_as_modern(exc) and hasattr(session, "discover"), - "MCP server '%s': legacy handshake rejected (%s) — " - "retrying via server/discover (2026-07-28 stateless server)") + "MCP server '%s': legacy handshake rejected (%s) — retrying via server/discover (2026-07-28 stateless server)") async def _serve_session(self, session, connect_timeout: float, label: str = "", mark_lifecycle: bool = False) -> str: """Handshake, discover, publish readiness, then serve until a lifecycle event. Clears stale breaker state but leaves the session UNPROVEN: flapping transports handshake fine and drop - moments later, so only keepalive or tool-call success clears the reconnect budget.""" + moments later, so only keepalive/tool-call success clears the reconnect budget.""" self.initialize_result = await self._negotiate_session(session, connect_timeout) self.session = session if mark_lifecycle: @@ -135,8 +136,7 @@ class MCPServerTransportMixin: async def _serve_transport(self, transport_cm, label: str, connect_timeout: float) -> str: """Open *transport_cm*, wrap its streams in a ClientSession and serve it. Streams are indexed, - not unpacked: mcp 1.x yields ``(read, write, get_session_id)``, 2.x ``(read, write)``. - A transport TaskGroup drop maps to ``"reconnect"`` instead of backoff/park.""" + not unpacked (mcp 1.x yields a 3-tuple, 2.x a pair); a TaskGroup drop maps to ``"reconnect"``.""" try: async with transport_cm as _streams: async with _core.ClientSession(_streams[0], _streams[1], **self._session_kwargs()) as session: @@ -148,12 +148,12 @@ class MCPServerTransportMixin: def _track_spawned_children(self, new_pids: Set[int]) -> None: """Ledger the freshly spawned stdio children (pids, pgids, machine spawn ledger). pgids are - captured while alive (getpgid fails once it exits; the sweep needs it for reparented descendants).""" + captured while alive (getpgid fails after exit; the sweep needs them for reparented descendants).""" new_pgids: Dict[int, int] = {} for pid in new_pids: try: new_pgids[pid] = os.getpgid(pid) - except (AttributeError, ProcessLookupError, OSError): # Windows / already exited + except (AttributeError, OSError): # Windows / already exited pass with _core._lock: _stdio_pids.update(dict.fromkeys(new_pids, self.name)) @@ -173,8 +173,7 @@ class MCPServerTransportMixin: with _core._lock: for pid in new_pids: _stdio_pids.pop(pid, None) - # ``os.kill(pid, 0)`` is NOT a no-op on Windows; the child may be gone while - # descendants remain in its pgroup. + # Windows-safe pid probe; the child may be gone while descendants remain in its pgroup. if _pid_exists(pid) or _pgroup_alive(_stdio_pgids.get(pid)): _orphan_stdio_pids.add(pid) _orphan_stdio_pid_servers[pid] = self.name @@ -183,8 +182,7 @@ class MCPServerTransportMixin: async def _run_stdio(self, config: dict): """Run the server using stdio transport.""" - if config.get("identity_header") is not None: - # No headers on stdio — warn so a copy-pasted HTTP block doesn't mislead. + if config.get("identity_header") is not None: # copy-pasted HTTP block: warn, don't mislead logger.warning("MCP server '%s': identity_header is only supported on " "HTTP/SSE transports — ignored for stdio servers", self.name) if not _core._ensure_mcp_sdk(): @@ -195,39 +193,34 @@ class MCPServerTransportMixin: raise ValueError(f"MCP server '{self.name}' has no 'command' in config") command, safe_env = _core._resolve_stdio_command(command, _core._build_safe_env(config.get("env"))) await _osv_malware_preflight(self.name, command, config.get("args", [])) - # Parent-death watchdog so kill -9 / crash can't leave the child tree running (POSIX-only). - # AFTER the OSV preflight so the check inspects the real package. + # Parent-death watchdog (POSIX) so kill -9 can't leave the child tree running; AFTER the + # OSV preflight so the check inspects the real package. command, args = _wrap_command_with_watchdog(command, config.get("args", [])) server_params = _core.StdioServerParameters( command=command, args=args, env=safe_env or None, cwd=config.get("cwd"), # Windows pipes can split non-UTF-8 bytes at chunk boundaries; substitute, don't raise. encoding_error_handler="replace") - session_kwargs = self._session_kwargs() - # Reap orphans of prior attempts first, else each retry piles up zombie pairs. Unscoped on - # purpose (also reaps servers that never reconnect). Off-loop: the reaper blocks up to 2s. + # Reap orphans of prior attempts first (else retries pile up zombie pairs); unscoped on purpose; + # off-loop because the reaper blocks up to 2s. await asyncio.to_thread(_core._kill_orphaned_mcp_children) pids_before = _core._snapshot_child_pids() # so the new child can be identified after spawn new_pids: set = set() # Subprocess stderr goes to ~/.hermes/logs/mcp-stderr.log so banners can't corrupt the TUI. _core._write_stderr_log_header(self.name) try: - errlog = _core._get_mcp_stderr_log() - async with _core.stdio_client(server_params, errlog=errlog) as (read_stream, write_stream): - # New PIDs for force-kill cleanup, minus non-MCP children (slash_worker, LSP) that - # race into the window: they share the TUI's pgid, so leaking them into _stdio_pgids - # would make the shutdown killpg() kill the TUI itself. + async with _core.stdio_client(server_params, errlog=_core._get_mcp_stderr_log()) as (read_stream, write_stream): + # New PIDs for force-kill cleanup, minus non-MCP children (slash_worker, LSP) racing + # into the window: they share the TUI's pgid — leaking them would killpg() the TUI. new_pids = _filter_mcp_children(_core._snapshot_child_pids() - pids_before) if new_pids: self._track_spawned_children(new_pids) self._stdio_child_pids = set(new_pids) # so in-flight calls fail fast when the child dies - async with _core.ClientSession(read_stream, write_stream, **session_kwargs) as session: - # Bound the handshake here: ``connect_timeout`` only bounds the caller's ``.result()``. - # A server that never answers ``initialize`` would otherwise hang forever, skip the - # ``finally`` and leak child + pipes on every retry until EMFILE. + async with _core.ClientSession(read_stream, write_stream, **self._session_kwargs()) as session: + # Bound the handshake here (``connect_timeout`` only bounds the caller's ``.result()``): + # a server that never answers ``initialize`` would leak child + pipes per retry until EMFILE. connect_timeout = float(config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)) return await self._serve_session(session, connect_timeout, mark_lifecycle=True) - finally: - # Runs on clean exit, exceptions AND cancellation. + finally: # clean exit, exceptions AND cancellation if new_pids: self._release_spawned_children(new_pids) @@ -235,29 +228,33 @@ class MCPServerTransportMixin: async def _preflight_content_type(self, url: str, *, headers: Optional[dict] = None, ssl_verify: bool = True, client_cert=None, timeout: float = 5.0) -> None: - """Probe *url* before the SDK connects: a plain web page makes the SDK sit out the full + """Probe *url* before the SDK connects: a plain web page would make the SDK sit out the full ``connect_timeout`` before an opaque ``CancelledError``; this raises NonMcpEndpointError within - ``timeout`` instead. Allow-list based: only a 2xx with a definite non-MCP content type is - rejected, and only after a JSON-RPC ``initialize`` POST also fails to look like MCP (some - servers serve a UI on GET but speak MCP via POST). Missing content type, non-2xx or transport - errors pass silently — the real handshake stays the source of truth. Own httpx client, OUTSIDE - the SDK's anyio task group, so the error isn't wrapped in an ExceptionGroup.""" + ``timeout``. Allow-list based: only a 2xx with a definite non-MCP content type is rejected, and + only after a JSON-RPC ``initialize`` POST also fails to look like MCP (some servers serve a UI + on GET but speak MCP via POST). Anything else passes — the handshake stays the source of truth. + Own httpx client, OUTSIDE the SDK's anyio task group, so the error isn't group-wrapped.""" try: import httpx as _httpx except ImportError: return # No httpx → skip probe; SDK import would have failed first. + def _non_mcp_2xx(resp) -> bool: + # Only judge 2xx (4xx/5xx may be an auth challenge the handshake handles); no content + # type advertised → don't second-guess the SDK. + ct = _content_type_base(resp) + return _is_2xx(resp) and bool(ct) and ct not in self._MCP_CONTENT_TYPES + probe_headers = dict(headers) if headers else {} try: async with _httpx.AsyncClient(verify=ssl_verify, follow_redirects=True, timeout=_httpx.Timeout(timeout), - **({"cert": client_cert} if client_cert is not None else {})) as client: + **_present(cert=client_cert)) as client: # HEAD is cheapest; fall back to GET on 405/501. resp = await client.head(url, headers=probe_headers) if resp.status_code in (405, 501): resp = await client.get(url, headers=probe_headers) # Non-MCP content type on HEAD/GET: try a JSON-RPC POST so POST-only servers pass. - ct = _content_type_base(resp) - if ct and ct not in self._MCP_CONTENT_TYPES and _is_2xx(resp): + if _non_mcp_2xx(resp): post_resp = await client.post( url, content=_PROBE_INITIALIZE_BODY, headers={**probe_headers, "Content-Type": "application/json", @@ -266,12 +263,9 @@ class MCPServerTransportMixin: resp = post_resp except _httpx.HTTPError: return # DNS/connect/timeout/transport error — let the SDK try. - - # Only judge 2xx (4xx/5xx may be an auth challenge the handshake handles); no content type - # advertised → don't second-guess the SDK. - ct_base = _content_type_base(resp) - if not _is_2xx(resp) or not ct_base or ct_base in self._MCP_CONTENT_TYPES: + if not _non_mcp_2xx(resp): return + ct_base = _content_type_base(resp) raise NonMcpEndpointError( f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP " f"response (expected one of: {', '.join(self._MCP_CONTENT_TYPES)}). The URL most likely " @@ -281,10 +275,10 @@ class MCPServerTransportMixin: def _reconnect_or_reraise_group(self, eg: BaseExceptionGroup) -> str: """Map an SDK transport TaskGroup failure to a clean ``"reconnect"``: HTTP/SSE stream pumps run in an anyio TaskGroup, so a transient drop escapes as a ``BaseExceptionGroup`` that would - otherwise back off and park the server for 300s over a sub-second glitch. Re-raise when it is - not a transient drop: shutdown in progress (``_shutdown_event`` is set before cancel), the - group carries KeyboardInterrupt/SystemExit or a real CancelledError, or no live session was - reached this attempt (``_ready`` unset — connect failures must back off, not hot-loop).""" + otherwise park the server for 300s over a sub-second glitch. Re-raise when it is not one: + shutdown in progress (``_shutdown_event`` is set before cancel), KeyboardInterrupt/SystemExit + or a real CancelledError in the group, or no live session this attempt (``_ready`` unset — + connect failures must back off, not hot-loop).""" if (self._shutdown_event.is_set() or eg.split((KeyboardInterrupt, SystemExit))[0] is not None or eg.split(asyncio.CancelledError)[0] is not None @@ -296,8 +290,7 @@ class MCPServerTransportMixin: def _build_oauth_auth(self, url: str, config: dict): """OAuth 2.1 PKCE via the central MCPOAuthManager (one provider reused across reconnects and - CLI paths). Setup failures (e.g. non-interactive without cached tokens) re-raise so only this - server is reported failed.""" + CLI paths). Setup failures re-raise (after a warning) so only this server is reported failed.""" if self._auth_type != "oauth": return None try: @@ -310,28 +303,25 @@ class MCPServerTransportMixin: def _sse_transport(self, url: str, headers: dict, connect_timeout: float, ssl_verify, client_cert, oauth_auth, strict_cfg_headers: bool): """``sse_client`` context manager for ``transport: sse`` entries.""" - if strict_cfg_headers: - # Fail closed: SSE cannot enforce the redirect boundary. + if strict_cfg_headers: # fail closed: SSE cannot enforce the redirect boundary raise ValueError(f"MCP server '{self.name}': strict_redirect_headers is " "not supported on the SSE transport.") if _core.sse_client is None: raise ImportError(f"MCP server '{self.name}' requires SSE transport but " "mcp.client.sse.sse_client is not available. " "Upgrade the mcp package to get SSE support.") - # sse_read_timeout bounds the gap between events: SSE servers idle for minutes, so 300s - # (matching the Streamable HTTP read timeout), not tool_timeout. ``auth`` must be forwarded - # or OAuth SSE servers 401 silently. + # sse_read_timeout bounds the gap between events: SSE servers idle for minutes, so 300s (the + # Streamable HTTP read timeout), not tool_timeout. ``auth`` must be forwarded or OAuth SSE 401s silently. sse_kwargs: dict = {"url": url, "headers": headers or None, "timeout": float(connect_timeout), - "sse_read_timeout": 300.0, **({"auth": oauth_auth} if oauth_auth is not None else {})} + "sse_read_timeout": 300.0, **_present(auth=oauth_auth)} if client_cert is not None or ssl_verify is not True: - # sse_client has no verify/cert kwargs: an httpx_client_factory forwards the SDK's - # (headers, auth, timeout) and layers TLS on top. The client MUST come from the SDK's - # own httpx module (httpx2 on mcp >= 2.0) — see sdk_httpx(). + # sse_client has no verify/cert kwargs: an httpx_client_factory forwards the SDK's (headers, + # auth, timeout) and layers TLS on top. Client MUST come from the SDK's httpx (httpx2 on mcp >= 2.0). _httpx_mod = _core.sdk_httpx() sse_kwargs["httpx_client_factory"] = lambda headers=None, timeout=None, auth=None: _httpx_mod.AsyncClient( follow_redirects=True, verify=ssl_verify, timeout=timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0), - **{k: v for k, v in (("headers", headers), ("auth", auth), ("cert", client_cert)) if v is not None}) + **_present(headers=headers, auth=auth, cert=client_cert)) return _core.sse_client(**sse_kwargs) def _streamable_http_transport(self, url: str, headers: dict, connect_timeout: float, @@ -340,26 +330,24 @@ class MCPServerTransportMixin: """Streamable HTTP context manager: mcp >= 1.24.0 gets a caller-owned httpx client; on the deprecated API (mcp < 1.24.0) the SDK owns the client.""" if not _core._MCP_NEW_HTTP: - if strict_cfg_headers: - # Fail closed: without an owned client redirects can't be hooked. + if strict_cfg_headers: # fail closed: without an owned client redirects can't be hooked raise ImportError(f"MCP server '{self.name}' requires mcp >= 1.24.0 to " "enforce the portable redirect-header boundary " "(strict_redirect_headers). Upgrade the mcp package.") return _core.streamablehttp_client(url, headers=headers, timeout=float(connect_timeout), verify=ssl_verify, - **({"auth": oauth_auth} if oauth_auth is not None else {})) - # Explicit AsyncClient matching the SDK's create_mcp_http_client defaults; MUST come from - # the SDK's httpx module (httpx2 on mcp >= 2.0) since the SDK sends its own Requests through it. + **_present(auth=oauth_auth)) + # Explicit AsyncClient matching the SDK's create_mcp_http_client defaults; MUST come from the + # SDK's httpx (httpx2 on mcp >= 2.0) since the SDK sends its own Requests through it. httpx = _core.sdk_httpx() _strip_auth_on_cross_origin_redirect = _make_redirect_header_stripper( httpx.URL(url), strict=strict_cfg_headers, configured_header_names=configured_header_names) client_kwargs: dict = {"follow_redirects": True, "timeout": httpx.Timeout(float(connect_timeout), read=300.0), "verify": ssl_verify, **({"headers": headers} if headers else {}), "event_hooks": {"response": [_strip_auth_on_cross_origin_redirect]}, - **{k: v for k, v in (("auth", oauth_auth), ("cert", client_cert)) if v is not None}} + **_present(auth=oauth_auth, cert=client_cert)} @asynccontextmanager - async def _owned_client_streams(): - # Caller owns the client lifecycle — the SDK skips cleanup when http_client is provided. + async def _owned_client_streams(): # the SDK skips cleanup when http_client is provided async with httpx.AsyncClient(**client_kwargs) as http_client: async with _core.streamable_http_client(url, http_client=http_client) as streams: yield streams @@ -376,13 +364,11 @@ class MCPServerTransportMixin: url = config["url"] headers = dict(config.get("headers") or {}) # Agent Plugins v1 strict_redirect_headers: configured headers MUST NOT follow a cross-origin - # redirect. Capture their names BEFORE client-generated headers are merged in. + # redirect — capture their names BEFORE client-generated headers are merged in. configured_header_names = {key.lower() for key in headers} - # Optional per-user identity header; explicit headers of the same name win. - headers = _apply_identity_header(self.name, config, headers) - # Some servers require MCP-Protocol-Version on the first request; seed it (user override - # wins) from the HANDSHAKE version, not the latest: a 2026-07-28 header would route the - # handshake-era ``initialize()`` body onto the per-request-envelope ladder, which rejects it. + headers = _apply_identity_header(self.name, config, headers) # explicit same-name headers win + # Seed MCP-Protocol-Version (user override wins) from the HANDSHAKE version, not the latest: a + # 2026-07-28 header routes the handshake-era ``initialize()`` onto the envelope ladder, which rejects it. if not any(key.lower() == "mcp-protocol-version" for key in headers): headers["mcp-protocol-version"] = _core.LATEST_HANDSHAKE_VERSION connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT) @@ -400,8 +386,7 @@ class MCPServerTransportMixin: async def _discover_tools(self): """Discover tools from the connected session. Capability-gated: prompt-/resource-only servers raise ``MCPError(-32601)`` on ``tools/list``, which would abort the connection.""" - # Fresh transport: re-probe ``ping`` in case the server gained support across the reconnect. - self._ping_unsupported = False + self._ping_unsupported = False # fresh transport: re-probe ``ping`` across the reconnect if self.session is None: return if not self._advertises_tools(): @@ -416,19 +401,18 @@ class MCPServerTransportMixin: self._register_discovered_tools_if_needed() def _register_discovered_tools_if_needed(self) -> None: - """Publish freshly discovered tools for a registry-owned server if none are registered - (initial registration normally happens in ``_discover_and_register_server``). On reconnect, - outage handling may clear ``_ready`` and deregister stale tools; ownership via ``_servers`` - authorizes publishing before readiness is restored so a revival never comes back with zero - tools — likewise a server retained after a recoverable initial failure.""" + """Publish freshly discovered tools when none are registered (initial registration normally + happens in ``_discover_and_register_server``). Outage handling may clear ``_ready`` and + deregister stale tools; ownership via ``_servers`` authorizes publishing before readiness is + restored so a revival (or a server retained after a recoverable initial failure) never comes + back with zero tools.""" if self._registered_tool_names: return - if not self._ready.is_set(): - with _core._lock: - if _core._servers.get(self.name) is not self: - return - self._registered_tool_names = _core._register_server_tools(self.name, self, self._config) - # A retained initial-failure server that just published tools has recovered. with _core._lock: + owned = _core._servers.get(self.name) is self + if not owned and not self._ready.is_set(): + return + self._registered_tool_names = _core._register_server_tools(self.name, self, self._config) + with _core._lock: # a retained initial-failure server that just published tools has recovered if _core._servers.get(self.name) is self: _core._server_connect_errors.pop(self.name, None) From c66fbdc2ca1a17fdb5b66b2f8721776bf04dcd50 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:20:23 -0700 Subject: [PATCH 05/16] =?UTF-8?q?refactor(tools):=20skill=5Fmanager=5Fguar?= =?UTF-8?q?ds=20=E2=80=94=20compact=20delete-target=20refusals?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/skill_manager_guards.py | 14 +++++--------- 1 file changed, 5 insertions(+), 9 deletions(-) diff --git a/tools/skill_manager_guards.py b/tools/skill_manager_guards.py index 143c308238..d7de879d3a 100644 --- a/tools/skill_manager_guards.py +++ b/tools/skill_manager_guards.py @@ -110,9 +110,8 @@ def _validate_delete_target(skill_dir: Path) -> Optional[str]: """Last-line guard before rmtree: even a poisoned tree must never delete (1) a path outside every known skills root, (2) a skills root itself, (3) a symlink/junction (rmtree follows it).""" if _is_path_redirect(skill_dir): - return ( - f"Refusing to delete '{skill_dir}': the skill directory is a " - f"symlink/junction. Remove the link target manually if intended.") + return (f"Refusing to delete '{skill_dir}': the skill directory is a " + f"symlink/junction. Remove the link target manually if intended.") try: skill_dir.resolve() except OSError as exc: @@ -120,14 +119,11 @@ def _validate_delete_target(skill_dir: Path) -> Optional[str]: resolved, roots = _resolved_roots(skill_dir) for _root, root in roots: if resolved == root: - return ( - f"Refusing to delete '{skill_dir}': resolves to the skills root " - f"itself, which would remove every installed skill.") + return (f"Refusing to delete '{skill_dir}': resolves to the skills root " + f"itself, which would remove every installed skill.") if resolved.is_relative_to(root): return None - return ( - f"Refusing to delete '{skill_dir}': path does not resolve inside any " - f"known skills root.") + return f"Refusing to delete '{skill_dir}': path does not resolve inside any known skills root." def _is_pinned(name: str, what: str) -> Optional[bool]: From 9458f4c27a4243ac93a177056183b986905f34b2 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:20:43 -0700 Subject: [PATCH 06/16] refactor(tools): MCP group L docstring rewrap --- tools/mcp_oauth_manager.py | 31 ++++++++++++++----------------- tools/mcp_tool_errors.py | 16 +++++++--------- tools/mcp_tool_handlers.py | 5 ++--- tools/mcp_tool_health.py | 5 ++--- tools/mcp_tool_registration.py | 10 ++++------ tools/mcp_tool_transport.py | 25 +++++++++++-------------- 6 files changed, 40 insertions(+), 52 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index 2ca0fe2cd1..f5cb5f52ea 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -1,10 +1,9 @@ -"""Central manager for per-server MCP OAuth state (one instance per process): per-server -providers, cross-process token reload (mtime-based disk watch so tokens refreshed by cron/another -CLI are picked up without a restart), 401 deduplication (N concurrent 401s with the same -access_token trigger one recovery) and reconnect signalling (``MCPServerTask`` drives the -reconnect; the manager decides when). The ONLY place that instantiates the SDK's -``OAuthClientProvider`` for runtime use; refresh stays lazy in the SDK — one ``stat()`` per tool -call is cheaper than an await + refresh round-trip.""" +"""Central manager for per-server MCP OAuth state (one instance per process): per-server providers, cross-process +token reload (mtime-based disk watch so tokens refreshed by cron/another CLI are picked up without a restart), 401 +deduplication (N concurrent 401s with the same access_token trigger one recovery) and reconnect signalling +(``MCPServerTask`` drives the reconnect; the manager decides when). The ONLY place that instantiates the SDK's +``OAuthClientProvider`` for runtime use; refresh stays lazy in the SDK — one ``stat()`` per tool call is cheaper +than an await + refresh round-trip.""" from __future__ import annotations @@ -171,12 +170,11 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): return re.search(rb"\binvalid_client\b", (await response.aread()).lower()) is not None async def _maybe_flag_poisoned_client(self, response: Any) -> None: - """An ``invalid_client`` rejection of our ``client_id`` at the token endpoint proves the - cached registration is dead server-side: delete ``client.json`` (+ stale metadata) so the - SDK re-runs DCR next flow. Conservative: acts ONLY on 400/401 at the discovered - ``token_endpoint`` (the only request carrying our ``client_id``) with ``invalid_client`` - in the body; pre-registered clients are never poisoned; any failure is swallowed. The - browser-side "Redirect URI Mismatch" case has no HTTP signal (``hermes mcp reauth``).""" + """An ``invalid_client`` rejection of our ``client_id`` at the token endpoint proves the cached registration + is dead server-side: delete ``client.json`` (+ stale metadata) so the SDK re-runs DCR next flow. + Conservative: acts ONLY on 400/401 at the discovered ``token_endpoint`` (the only request carrying our + ``client_id``) with ``invalid_client`` in the body; pre-registered clients are never poisoned; any failure + is swallowed. The browser-side "Redirect URI Mismatch" case has no HTTP signal (``hermes mcp reauth``).""" try: if (self._hermes_preregistered or getattr(response, "status_code", None) not in (400, 401) or not await self._is_invalid_client_at_token_endpoint(response)): @@ -370,10 +368,9 @@ class MCPOAuthManager: entry.pending_401.pop(key, None) async def handle_401(self, server_name: str, failed_access_token: Optional[str] = None) -> bool: - """Handle a 401 from a tool call. True: a (possibly new) token is available — reconnect - and retry. False: no recovery path — surface ``needs_reauth`` so the model stops - hallucinating manual refreshes. Concurrent 401s with the same ``failed_access_token`` - fire one recovery attempt; the rest await its future.""" + """Handle a 401 from a tool call. True: a (possibly new) token is available — reconnect and retry. False: no + recovery path — surface ``needs_reauth`` so the model stops hallucinating manual refreshes. Concurrent 401s + with the same ``failed_access_token`` fire one recovery attempt; the rest await its future.""" entry = self._entries.get(self._key(server_name)) if entry is None or entry.provider is None: return False diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index 9ab2e50e89..e56bd4cc88 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -20,9 +20,8 @@ _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022 def _jsonrpc_matches(exc: BaseException, codes: tuple, markers: tuple, code=None) -> bool: - """Structural ``MCPError.error.code`` (or *code*) in *codes*, else any *marker* in - ``str(exc).lower()``. Never ``isinstance`` on SDK exception types: they arrive wrapped in - ExceptionGroups and drift across generations.""" + """Structural ``MCPError.error.code`` (or *code*) in *codes*, else any *marker* in ``str(exc).lower()``. Never + ``isinstance`` on SDK exception types: they arrive wrapped in ExceptionGroups and drift across generations.""" code = getattr(getattr(exc, "error", None), "code", None) or code return code in codes or any(marker in str(exc).lower() for marker in markers) @@ -293,12 +292,11 @@ _EXC_TRAVERSAL_MAX_NODES = 10_000 def _is_session_expired_error(exc: BaseException) -> bool: - """True if ``exc`` looks like a transport session expiry (Streamable-HTTP servers GC session - state on idle TTL / restart / pod rotation while the OAuth token stays valid) — the fix is a - transport reconnect, not an OAuth refresh. Iterative walk over ``exceptions`` / ``__cause__`` / - ``__context__`` with a visited set AND a node budget; every reachable node is inspected so an - InterruptedError anywhere overrides transport markers, and the chain walk matters because SDK - wrappers raise a generic RuntimeError *from* a message-less ClosedResourceError.""" + """True if ``exc`` looks like a transport session expiry (Streamable-HTTP servers GC session state on idle TTL / + restart / pod rotation while the OAuth token stays valid) — the fix is a transport reconnect, not an OAuth + refresh. Iterative walk over ``exceptions`` / ``__cause__`` / ``__context__`` with a visited set AND a node + budget; every reachable node is inspected so an InterruptedError anywhere overrides transport markers, and the + chain walk matters because SDK wrappers raise a generic RuntimeError *from* a message-less ClosedResourceError.""" # AnyIO stream exceptions are often message-less, so type checks complement marker matching. transport_error_types = tuple(_optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream")) stack: "list[BaseException | None]" = [exc] diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 832d0c7066..69b5c40341 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -1,6 +1,5 @@ -"""Registry-facing sync handlers for MCP tools and utility tools (resources/prompts), plus -the per-call recovery ladder: trust gating, circuit breaker, auth (401) refresh, -session-expired reconnect and dead-stdio respawn retry.""" +"""Registry-facing sync handlers for MCP tools and utility tools (resources/prompts), plus the per-call recovery +ladder: trust gating, circuit breaker, auth (401) refresh, session-expired reconnect and dead-stdio respawn retry.""" import logging import asyncio diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index c90db405b2..93d8576af3 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -130,9 +130,8 @@ class MCPServerHealthMixin: _forget_mcp_tool_server(tool_name) async def _refresh_tools(self): - """Re-fetch tools on ``tools/list_changed`` and update the registry. The lock serializes - rapid-fire notifications; after the list_tools ``await`` all mutations are synchronous — - atomic on the event loop.""" + """Re-fetch tools on ``tools/list_changed`` and update the registry. The lock serializes rapid-fire + notifications; after the list_tools ``await`` all mutations are synchronous — atomic on the event loop.""" if not self._advertises_tools(): return # tools/list would raise MCPError(-32601) async with self._refresh_lock: diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index 7595873656..a09f393107 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -46,9 +46,8 @@ def _annotation_read_only_hint(mcp_tool: Any) -> bool: def _record_tool_trust_metadata(server_name: str, config: dict, tools: List[Any]) -> None: - """Capture per-server trust and per-tool readOnlyHint at discovery — the security - boundary: the call-time gate classifies from data we control, never re-read - server-supplied state.""" + """Capture per-server trust and per-tool readOnlyHint at discovery — the security boundary: the call-time gate + classifies from data we control, never re-read server-supplied state.""" with _core._lock: _core._server_trust_levels[server_name] = _normalize_server_trust((config or {}).get("trust")) hints = _core._tool_read_only_hints.setdefault(server_name, {}) @@ -109,9 +108,8 @@ def _existing_tool_names() -> List[str]: def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]: - """Include/exclude predicate for a server's tool names: ``tools.include`` is a whitelist - (``[]`` = register nothing), ``tools.exclude`` a blacklist; entries are exact names or - fnmatch globs; include wins over exclude.""" + """Include/exclude predicate for a server's tool names: ``tools.include`` is a whitelist (``[]`` = register + nothing), ``tools.exclude`` a blacklist; entries are exact names or fnmatch globs; include wins over exclude.""" tools_filter = config.get("tools") or {} include_raw = tools_filter.get("include") include_set = _normalize_name_filter(include_raw, f"mcp_servers.{name}.tools.include") diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index 29721bb977..5eb98b1481 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -1,6 +1,5 @@ -"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID -ledger), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth), -protocol negotiation and initial tool discovery.""" +"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID ledger), Streamable HTTP +/ SSE connect (preflight, identity header, client certs, OAuth), protocol negotiation and initial tool discovery.""" import logging import asyncio @@ -273,12 +272,11 @@ class MCPServerTransportMixin: "HTTP / SSE endpoint (e.g. https://host/mcp, not https://host/).") def _reconnect_or_reraise_group(self, eg: BaseExceptionGroup) -> str: - """Map an SDK transport TaskGroup failure to a clean ``"reconnect"``: HTTP/SSE stream pumps - run in an anyio TaskGroup, so a transient drop escapes as a ``BaseExceptionGroup`` that would - otherwise park the server for 300s over a sub-second glitch. Re-raise when it is not one: - shutdown in progress (``_shutdown_event`` is set before cancel), KeyboardInterrupt/SystemExit - or a real CancelledError in the group, or no live session this attempt (``_ready`` unset — - connect failures must back off, not hot-loop).""" + """Map an SDK transport TaskGroup failure to a clean ``"reconnect"``: HTTP/SSE stream pumps run in an anyio + TaskGroup, so a transient drop escapes as a ``BaseExceptionGroup`` that would otherwise park the server for + 300s over a sub-second glitch. Re-raise when it is not one: shutdown in progress (``_shutdown_event`` is + set before cancel), KeyboardInterrupt/SystemExit or a real CancelledError in the group, or no live session + this attempt (``_ready`` unset — connect failures must back off, not hot-loop).""" if (self._shutdown_event.is_set() or eg.split((KeyboardInterrupt, SystemExit))[0] is not None or eg.split(asyncio.CancelledError)[0] is not None @@ -401,11 +399,10 @@ class MCPServerTransportMixin: self._register_discovered_tools_if_needed() def _register_discovered_tools_if_needed(self) -> None: - """Publish freshly discovered tools when none are registered (initial registration normally - happens in ``_discover_and_register_server``). Outage handling may clear ``_ready`` and - deregister stale tools; ownership via ``_servers`` authorizes publishing before readiness is - restored so a revival (or a server retained after a recoverable initial failure) never comes - back with zero tools.""" + """Publish freshly discovered tools when none are registered (initial registration normally happens in + ``_discover_and_register_server``). Outage handling may clear ``_ready`` and deregister stale tools; + ownership via ``_servers`` authorizes publishing before readiness is restored so a revival (or a server + retained after a recoverable initial failure) never comes back with zero tools.""" if self._registered_tool_names: return with _core._lock: From 4ba77a6c136ef921c036e5c10cdccc4e90387cd2 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:23:09 -0700 Subject: [PATCH 07/16] =?UTF-8?q?refactor(tools):=20skill=5Fmanage=20?= =?UTF-8?q?=E2=80=94=20shared=20missing-arg=20predicates,=20compact=20cate?= =?UTF-8?q?gory=20validation=20and=20comments?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/skill_manager_batch.py | 18 +++++------- tools/skill_manager_tool.py | 53 ++++++++++++++++-------------------- 2 files changed, 30 insertions(+), 41 deletions(-) diff --git a/tools/skill_manager_batch.py b/tools/skill_manager_batch.py index 79fa30341c..b23bc50c01 100644 --- a/tools/skill_manager_batch.py +++ b/tools/skill_manager_batch.py @@ -133,13 +133,11 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non if len(operations) != 1: return tool_error("delete must be the SOLE op in its call — it doesn't " "compose with other ops' rollback.", success=False) - op = operations[0] - nm = op.get("name") or default_name + nm = operations[0].get("name") or default_name if not nm: return tool_error("operations[0] (delete) needs a 'name'.", success=False) - return _smt.skill_manage( - action="delete", name=nm, absorbed_into=op.get("absorbed_into"), - task_id=task_id, session_id=session_id) + return _smt.skill_manage(action="delete", name=nm, task_id=task_id, session_id=session_id, + absorbed_into=operations[0].get("absorbed_into")) names, err = _validate_batch_ops(operations, default_name, tool_error) if err is not None: @@ -176,13 +174,11 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non parsed = {"success": False, "error": "unparseable op result"} if not parsed.get("success"): note, rollback_failed = _rollback(snapshots, _smt._find_skill) - fail = { + fail = { # key order is wire-visible "success": False, - "error": ( - f"operations[{i}] ({op['action']} on '{names[i]}') failed: " - f"{parsed.get('error', 'unknown error')} — batch aborted, {note}."), - "failed_index": i, - "completed_before_failure": i} + "error": (f"operations[{i}] ({op['action']} on '{names[i]}') failed: " + f"{parsed.get('error', 'unknown error')} — batch aborted, {note}."), + "failed_index": i, "completed_before_failure": i} # Carry the failing op's teaching payload (patch's file_preview / # fuzzy-match hints) through — without it the model recovers blind. for k, v in parsed.items(): diff --git a/tools/skill_manager_tool.py b/tools/skill_manager_tool.py index d99d977435..2cb78edf31 100644 --- a/tools/skill_manager_tool.py +++ b/tools/skill_manager_tool.py @@ -113,13 +113,11 @@ def _validate_name(name: str) -> Optional[str]: def _validate_category(category: Optional[str]) -> Optional[str]: - if category is None: + if category is None or (isinstance(category, str) and not category.strip()): return None if not isinstance(category, str): return "Category must be a string." category = category.strip() - if not category: - return None invalid = (f"Invalid category '{category}'. {_NAME_RULE} " "Categories must be a single directory name.") if "/" in category or "\\" in category: @@ -516,8 +514,9 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A return _err(f"failed to archive '{name}': {e}") if not ok: return _err(archive_msg) - return {"success": True, "_archived": True, - "message": f"Skill '{name}' archived ({archive_msg}).{absorbed_note}"} + return {"success": True, + "message": f"Skill '{name}' archived ({archive_msg}).{absorbed_note}", + "_archived": True} shutil.rmtree(skill_dir) _rmdir_if_empty(skill_dir.parent, skills_root) # empty category dir, never the root @@ -581,8 +580,7 @@ def _remove_file(name: str, file_path: str) -> Dict[str, Any]: # --- Main entry point --------------------------------------------------------- -# Set while replaying an already-approved staged skill write so skill_manage() -# does not re-gate (and re-stage) it. +# Set while replaying an approved staged skill write so skill_manage() does not re-gate it. _skill_gate_bypass: "_ctxvars.ContextVar[bool]" = _ctxvars.ContextVar( "skill_gate_bypass", default=False) @@ -642,8 +640,7 @@ def apply_skill_pending(payload: Dict[str, Any]) -> str: _skill_gate_bypass.reset(token) -# Debounce state for the sync push hook: a burst of skill_manage writes -# collapses into one push after a quiet window, on a daemon timer. +# Sync push debounce: a burst of skill_manage writes collapses into one push on a daemon timer. _sync_push_timer = None _sync_push_lock = threading.Lock() _SYNC_PUSH_DEBOUNCE_S = 5.0 @@ -674,21 +671,18 @@ def _maybe_debounced_sync_push(skill_name: str) -> None: def _act_patch(a): - # Two shapes: old_string/new_string = targeted replacement; - # content (alone) = full SKILL.md rewrite (absorbs the old 'edit'). + """Two shapes: old_string/new_string = targeted replacement (validated in _patch_skill so the + tool and the helper give the same guidance); content alone = full rewrite (the old 'edit').""" if a["content"] and (a["old_string"] or a["new_string"] is not None): return tool_error("Pass EITHER content (full SKILL.md rewrite) OR " "old_string/new_string (targeted replacement), not both.", success=False) if a["content"]: return _edit_skill(a["name"], a["content"]) - # Targeted-replacement validation lives in _patch_skill so the public - # tool and the helper return the same actionable guidance. return _patch_skill(a["name"], a["old_string"], a["new_string"], a["file_path"], a["replace_all"]) -# action -> handler(args dict). Handlers return a result dict, or a JSON string -# (tool_error) for argument-shape errors. "edit" is a legacy alias for a full -# rewrite (old transcripts/callers; not in the schema). +# action -> handler(args dict) returning a result dict, or a tool_error JSON string for +# argument-shape errors. "edit" is a legacy alias for a full rewrite (not in the schema). _ACTION_HANDLERS = { "create": lambda a: _create_skill(a["name"], a["content"], a["category"]), "edit": lambda a: _edit_skill(a["name"], a["content"]), @@ -697,15 +691,16 @@ _ACTION_HANDLERS = { "write_file": lambda a: _write_file(a["name"], a["file_path"], a["file_content"]), "remove_file": lambda a: _remove_file(a["name"], a["file_path"])} # action -> (arg, is_missing, error) argument-shape checks run before the handler. +_MISSING, _IS_NONE = (lambda v: not v), (lambda v: v is None) _REQUIRED_ARGS = { - "create": [("content", lambda v: not v, + "create": [("content", _MISSING, "content is required for 'create'. Provide the full SKILL.md text (frontmatter + body).")], - "edit": [("content", lambda v: not v, + "edit": [("content", _MISSING, "content is required for a full rewrite. Provide the full updated SKILL.md text.")], - "write_file": [("file_path", lambda v: not v, - "file_path is required for 'write_file'. Example: 'references/api-guide.md'"), - ("file_content", lambda v: v is None, "file_content is required for 'write_file'.")], - "remove_file": [("file_path", lambda v: not v, "file_path is required for 'remove_file'.")]} + "write_file": [ + ("file_path", _MISSING, "file_path is required for 'write_file'. Example: 'references/api-guide.md'"), + ("file_content", _IS_NONE, "file_content is required for 'write_file'.")], + "remove_file": [("file_path", _MISSING, "file_path is required for 'remove_file'.")]} def _record_success(action, name, result, *, file_path, absorbed_into, task_id, @@ -738,8 +733,7 @@ def _record_success(action, name, result, *, file_path, absorbed_into, task_id, bump_patch(name, action=action, task_id=task_id, session_id=session_id) elif action == "delete" and not result.get("_archived"): forget(name) - # Runs only AFTER the write gate passed (staged writes returned early), so - # un-reviewed content is never pushed. + # Only AFTER the write gate passed (staged writes returned early): never push un-reviewed content. with suppress(Exception): _maybe_debounced_sync_push(name) @@ -757,12 +751,11 @@ def skill_manage( if (preflight := _background_review_preflight(action, name)) is not None: return json.dumps(preflight, ensure_ascii=False) - # Approval gate: skills are too large to review inline, so they always stage - # regardless of origin; bypassed when replaying an approved staged write. - args = dict( - content=content, category=category, file_path=file_path, - file_content=file_content, old_string=old_string, new_string=new_string, - replace_all=replace_all, absorbed_into=absorbed_into) + # Approval gate: skills are too large to review inline, so they always stage regardless + # of origin; bypassed when replaying an approved staged write. + args = dict(content=content, category=category, file_path=file_path, file_content=file_content, + old_string=old_string, new_string=new_string, replace_all=replace_all, + absorbed_into=absorbed_into) if (gate_result := _apply_skill_write_gate(action, name, **args)) is not None: return gate_result From a5bee5eaab99fa783dae90f7d84fad0ee4bd3f9a Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:25:48 -0700 Subject: [PATCH 08/16] refactor(tools): MCP registration cache stand-ins -> SimpleNamespace factory; string-fragment joins --- tools/mcp_tool_registration.py | 32 +++++++++++--------------------- tools/mcp_tool_transport.py | 3 +-- 2 files changed, 12 insertions(+), 23 deletions(-) diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index a09f393107..1984422109 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -5,6 +5,7 @@ live, ``_register_from_cache_sync`` lazy) build ``_Candidate`` records for ``_re import logging from dataclasses import dataclass +from types import SimpleNamespace from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional from tools.mcp_tool_common import _parse_boolish, _core, _resolve_tool_timeout from tools.mcp_tool_handlers import ( @@ -119,24 +120,13 @@ def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]: return lambda tool_name: not (exclude_set and matches_name_filter(tool_name, exclude_set)) -class _CachedMCPTool: - """Stand-in for MCP Tool objects loaded from the schema cache. Missing or non-dict - ``annotations`` (older cache files) fail closed to write-capable.""" - - __slots__ = ("name", "description", "inputSchema", "annotations") - - def __init__(self, name: str, description: str, inputSchema: dict, annotations: Optional[dict] = None): - self.name = name - self.description = description - self.inputSchema = inputSchema or {} - self.annotations = annotations if isinstance(annotations, dict) else None - - @classmethod - def from_cache_dicts(cls, raws: Iterable[Any]) -> List["_CachedMCPTool"]: - """Cached rows -> stand-ins; rows that are not dicts or lack a name are dropped.""" - return [cls(raw["name"], raw.get("description") or "", - raw["inputSchema"] if isinstance(raw.get("inputSchema"), dict) else {}, raw.get("annotations")) - for raw in raws if isinstance(raw, dict) and raw.get("name")] +def _cached_tools(raws: Iterable[Any]) -> List[SimpleNamespace]: + """Schema-cache rows -> stand-ins for MCP Tool objects; rows that are not dicts or lack a name + are dropped. Missing or non-dict ``annotations`` (older cache files) fail closed to write-capable.""" + return [SimpleNamespace(name=raw["name"], description=raw.get("description") or "", + inputSchema=raw["inputSchema"] if isinstance(raw.get("inputSchema"), dict) else {}, + annotations=raw["annotations"] if isinstance(raw.get("annotations"), dict) else None) + for raw in raws if isinstance(raw, dict) and raw.get("name")] @dataclass @@ -156,8 +146,8 @@ class _Candidate: def _tool_candidates(name: str, tools: Iterable[Any], should_register: Callable[[str], bool], tool_timeout) -> List[_Candidate]: - """Native tools (live SDK objects or ``_CachedMCPTool``) -> candidates. The injection scan - runs on BOTH paths: the cache file is user-writable JSON.""" + """Native tools (live SDK objects or cache stand-ins) -> candidates. The injection scan runs on + BOTH paths: the cache file is user-writable JSON.""" out: List[_Candidate] = [] for t in tools: if not should_register(t.name): @@ -297,7 +287,7 @@ def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str] call-time gate is identical for live and cached registrations.""" from tools.mcp_schema_cache import config_fingerprint, tools_from_cache_entry, utility_tools_from_cache_entry tool_timeout = _resolve_tool_timeout(config) - cached_tools = _CachedMCPTool.from_cache_dicts(tools_from_cache_entry(entry)) + cached_tools = _cached_tools(tools_from_cache_entry(entry)) _record_tool_trust_metadata(name, config, cached_tools) candidates = _tool_candidates(name, cached_tools, _make_tool_filter(name, config), tool_timeout) candidates += _utility_candidates(name, utility_tools_from_cache_entry(entry), tool_timeout) diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index 5eb98b1481..1d16c998f7 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -265,8 +265,7 @@ class MCPServerTransportMixin: if not _non_mcp_2xx(resp): return ct_base = _content_type_base(resp) - raise NonMcpEndpointError( - f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP " + raise NonMcpEndpointError(f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP " f"response (expected one of: {', '.join(self._MCP_CONTENT_TYPES)}). The URL most likely " "points at a web page rather than an MCP endpoint — check it resolves to a Streamable " "HTTP / SSE endpoint (e.g. https://host/mcp, not https://host/).") From 370c6a245671e7f824e237653623b58e4a509e25 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:26:03 -0700 Subject: [PATCH 09/16] =?UTF-8?q?refactor(tools):=20skill=5Fmanage=20?= =?UTF-8?q?=E2=80=94=20dispatch=20default=20for=20unknown=20action,=20simp?= =?UTF-8?q?ler=20ledger=20read,=20squeeze=20body=20blanks?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/skill_ledger.py | 23 +++++++------------- tools/skill_manager_batch.py | 15 ++----------- tools/skill_manager_guards.py | 9 ++------ tools/skill_manager_tool.py | 41 ++++++++--------------------------- 4 files changed, 21 insertions(+), 67 deletions(-) diff --git a/tools/skill_ledger.py b/tools/skill_ledger.py index 8b4512ee4f..637f26f2e2 100644 --- a/tools/skill_ledger.py +++ b/tools/skill_ledger.py @@ -320,19 +320,16 @@ def capture_before( def list_entries(skill: Optional[str] = None, limit: Optional[int] = None) -> List[Dict[str, Any]]: """Read the ledger, newest first. Malformed lines are skipped.""" - path = ledger_path() - if not path.exists(): + try: + lines = ledger_path().read_text(encoding="utf-8").splitlines() + except OSError: # missing or unreadable ledger == empty return [] rows: List[Dict[str, Any]] = [] - try: - with open(path, "r", encoding="utf-8") as fh: - for line in fh: - with suppress(json.JSONDecodeError): - row = json.loads(line) if line.strip() else None - if isinstance(row, dict): - rows.append(row) - except OSError: - return [] + for line in lines: + with suppress(json.JSONDecodeError): + row = json.loads(line) if line.strip() else None + if isinstance(row, dict): + rows.append(row) if skill: rows = [r for r in rows if r.get("skill") == skill] rows.reverse() @@ -366,7 +363,6 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]: return False, f"no ledger entry with id '{entry_id}'" if path_err := _validate_entry_paths(entry): return False, f"refusing rollback: {path_err}" - before = list(entry.get("before") or []) after = list(entry.get("after") or []) # Historical hollow delete/archive/purge entries (SKILL.md only): fill from the @@ -383,7 +379,6 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]: if read_blob(str(item.get("sha256", ""))) is None: return False, (f"missing blob {item.get('sha256')} for {item.get('path')}; " "rollback aborted, nothing was changed") - # Safety entry: CURRENT state of every touched path, so the rollback itself is undoable. touched = {str(i["path"]) for i in before + after if i.get("path")} try: @@ -398,7 +393,6 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]: if safety_id is None: return False, ("pre-rollback safety capture failed (ledger disabled or " "unwritable); rollback aborted and current skills were not changed") - # Restore: write every before-file, remove files the mutation created. before_paths = {str(i["path"]) for i in before} for item in before: @@ -415,7 +409,6 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]: removed += 1 except OSError as e: logger.warning("skill_ledger: could not remove %s during rollback: %s", p, e) - append_entry( "rollback", entry.get("skill", "?"), before=safety_before, after=before, evidence={"rollback_target": entry_id, "restored": restored, "removed": removed}) diff --git a/tools/skill_manager_batch.py b/tools/skill_manager_batch.py index b23bc50c01..c5fe07118b 100644 --- a/tools/skill_manager_batch.py +++ b/tools/skill_manager_batch.py @@ -18,10 +18,8 @@ _BATCH_MAX_OPS = 20 def _validate_batch_ops(operations, default_name, tool_error): """Shape checks with no side effects. Returns (names, None) or (None, error_json).""" from tools.skill_manager_guards import _background_review_preflight - def fail(i, msg): return None, tool_error(f"operations[{i}]{msg}", success=False) - names = [] for i, op in enumerate(operations): if not isinstance(op, dict) or not op.get("action"): @@ -39,7 +37,6 @@ def _validate_batch_ops(operations, default_name, tool_error): preflight = _background_review_preflight(act, nm) if preflight is not None: return None, json.dumps(preflight, ensure_ascii=False) - # Clobber guard: a DESTRUCTIVE op (create/write_file/remove_file/full rewrite) on # a file an earlier op touched would SILENTLY discard its work — reject it. # Additive patches are always legal. Paths are normalized against spelling variants. @@ -67,9 +64,8 @@ def _snapshot_skills(names, snap_root, find_skill): for nm in dict.fromkeys(names): # ordered unique pre = find_skill(nm) pre_dir = Path(pre["path"]) if pre else None - snap = None - if pre_dir is not None and pre_dir.is_dir(): - snap = snap_root / nm + snap = snap_root / nm if pre_dir is not None and pre_dir.is_dir() else None + if snap is not None: try: shutil.copytree(pre_dir, snap) except Exception as exc: # noqa: BLE001 — no snapshot, no atomicity @@ -124,7 +120,6 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non top-level ``name`` fallback (staged replay).""" from tools import skill_manager_tool as _smt from tools.registry import tool_error - if not isinstance(operations, list) or not operations: return tool_error("operations must be a non-empty array.", success=False) if len(operations) > _BATCH_MAX_OPS: @@ -138,28 +133,23 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non return tool_error("operations[0] (delete) needs a 'name'.", success=False) return _smt.skill_manage(action="delete", name=nm, task_id=task_id, session_id=session_id, absorbed_into=operations[0].get("absorbed_into")) - names, err = _validate_batch_ops(operations, default_name, tool_error) if err is not None: return err - if not _smt._skill_gate_bypass.get(): # Approval gate for the WHOLE batch as one pending write. def _staging(wa): acts = ", ".join(op["action"] for op in operations) gist = f"batch({len(operations)} ops: {acts}) on {', '.join(sorted(set(names)))}" return {"action": "batch", "operations": operations}, gist - staged = _smt._run_write_gate(_staging) if staged is not None: return staged - snap_root = Path(tempfile.mkdtemp(prefix="skill_batch_")) snapshots, snap_err = _snapshot_skills(names, snap_root, _smt._find_skill) if snap_err is not None: shutil.rmtree(snap_root, ignore_errors=True) return tool_error(snap_err, success=False) - # Single-op path with the gate bypassed (the batch already cleared/staged it). results = [] rollback_failed = False @@ -194,7 +184,6 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non logger.warning("skill_manage batch rollback failed, snapshots kept at %s", snap_root) else: shutil.rmtree(snap_root, ignore_errors=True) - return json.dumps( {"success": True, "operations_applied": len(results), "results": results}, ensure_ascii=False) diff --git a/tools/skill_manager_guards.py b/tools/skill_manager_guards.py index d7de879d3a..3b4d913a1b 100644 --- a/tools/skill_manager_guards.py +++ b/tools/skill_manager_guards.py @@ -78,7 +78,6 @@ def _reset_background_review_read_marks() -> None: def _resolved_roots(skill_path: Path): """``(resolved skill_path, [(root, resolved_root), ...])`` over every resolvable skills root.""" from agent.skill_utils import get_all_skills_dirs - try: resolved = skill_path.resolve() except OSError: @@ -93,7 +92,6 @@ def _resolved_roots(skill_path: Path): def _containing_skills_root(skill_path: Path) -> Path: """Skills root (local or external_dirs) containing ``skill_path``; local dir if none match.""" from tools import skill_manager_tool as _smt - resolved, roots = _resolved_roots(skill_path) return next((root for root, r in roots if resolved.is_relative_to(r)), _smt._skills_dir()) @@ -173,12 +171,10 @@ def _background_review_write_guard( from agent.skill_utils import is_external_skill_path if is_external_skill_path(skill_dir): return _refusal( - f"{refuse} skill '{name}': " - "the skill lives in skills.external_dirs, which are " - "externally owned and read-only to autonomous curation.") + f"{refuse} skill '{name}': the skill lives in skills.external_dirs, which are " + f"externally owned and read-only to autonomous curation.") except Exception: logger.debug("external skill guard lookup failed for %s", name, exc_info=True) - try: from tools import skill_usage for predicate, label in ( @@ -256,7 +252,6 @@ def _maybe_auto_propose_org_edit(name: str, skill_path: Path) -> Optional[str]: the tool result or None; never raises (the edit is saved locally and can be proposed later).""" try: from tools import skills_sync_client as ssc - if not _is_org_mirror(skill_path): return None if not ssc.sync_org_auto_propose(): diff --git a/tools/skill_manager_tool.py b/tools/skill_manager_tool.py index 2cb78edf31..a0e42bfb29 100644 --- a/tools/skill_manager_tool.py +++ b/tools/skill_manager_tool.py @@ -206,7 +206,6 @@ def _find_skill(name: str) -> Optional[Dict[str, Any]]: categorized relative path (``mlops/axolotl``) — the two forms skill_view resolves. The categorized form matches RELATIVE to the local root only (relative_to raises for external dirs).""" from agent.skill_utils import get_all_skills_dirs - local_root = None if "/" in name or "\\" in name: try: @@ -216,7 +215,6 @@ def _find_skill(name: str) -> Optional[Dict[str, Any]]: "skills dir resolve failed; categorized lookups fall back to the unresolved path", exc_info=True) local_root = _skills_dir() - for skills_dir in get_all_skills_dirs(): if not skills_dir.exists(): continue @@ -281,7 +279,6 @@ def _skill_not_found_error(name: str, suffix: str = "") -> str: def _validate_file_path(file_path: str) -> Optional[str]: """Validate a write_file/remove_file path: under an allowed subdir, no escape.""" from tools.path_security import has_traversal_component - if not file_path: return "file_path is required." parts = Path(file_path).parts @@ -303,7 +300,6 @@ def _resolve_supporting_file(skill_dir: Path, file_path: str): """Validate ``file_path`` and resolve it inside ``skill_dir`` -> ``(target, None)`` | ``(None, error_dict)``.""" from tools.path_security import validate_within_dir - target = skill_dir / (file_path or "") err = _validate_file_path(file_path) or validate_within_dir(target, skill_dir) return (None, _err(err)) if err else (target, None) @@ -363,7 +359,7 @@ def _attach_lint_findings(result: Dict[str, Any], skill_md: Path) -> None: from tools.skill_linter import lint_skill # local import: optional path findings = lint_skill(skill_md) except Exception: - return + findings = None if not findings: return result["lint_warnings"] = [ @@ -386,7 +382,6 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An return _err(err) if existing := _find_skill(name): return _err(f"A skill named '{name}' already exists at {existing['path']}.") - skill_dir = _resolve_skill_dir(name, category) skill_dir.mkdir(parents=True, exist_ok=True) skill_md = skill_dir / "SKILL.md" @@ -394,7 +389,6 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An if scan_error := _security_scan_skill(skill_dir): shutil.rmtree(skill_dir, ignore_errors=True) return _err(scan_error) - root = _skills_dir() # Relative when under the profile dir; absolute when created under skills.create_dir. display = skill_dir.relative_to(root) if skill_dir.is_relative_to(root) else skill_dir @@ -456,7 +450,6 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N return _err(f"File not found: {target.relative_to(skill_dir)}") if read_guard := _background_review_read_before_write_guard(name, target, "patch", target_label): return read_guard - content = target.read_text(encoding="utf-8") # Same fuzzy engine as the file patch tool (whitespace/indent/escape normalization, # block anchors) so minor formatting mismatches don't fail. @@ -468,7 +461,6 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N from tools.fuzzy_match import format_no_match_hint match_error += format_no_match_hint(match_error, match_count, old_string, content) return _err(match_error) | {"file_preview": _clip(content, 500, "...")} - if err := _validate_content_size(new_content, label=target_label): return _err(err) if not file_path and (err := _validate_frontmatter(new_content)): @@ -491,7 +483,6 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A return guard if pinned_err := _pinned_guard(name): return _err(pinned_err) - absorbed_target = absorbed_into.strip() if isinstance(absorbed_into, str) else "" if absorbed_target: if absorbed_target == name: @@ -499,7 +490,6 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A if not _find_skill(absorbed_target): return _err(f"absorbed_into='{absorbed_target}' does not exist. " f"Create or patch the umbrella skill first, then retry the delete.") - skills_root = _containing_skills_root(skill_dir) if unsafe := _validate_delete_target(skill_dir): # defense-in-depth before rmtree return _err(unsafe) @@ -517,7 +507,6 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A return {"success": True, "message": f"Skill '{name}' archived ({archive_msg}).{absorbed_note}", "_archived": True} - shutil.rmtree(skill_dir) _rmdir_if_empty(skill_dir.parent, skills_root) # empty category dir, never the root return {"success": True, "message": f"Skill '{name}' deleted.{absorbed_note}"} @@ -540,7 +529,6 @@ def _write_file(name: str, file_path: str, file_content: str) -> Dict[str, Any]: f"bytes / 1 MiB). Consider splitting into smaller files.") if err := _validate_content_size(file_content, label=file_path): return _err(err) - skill_dir, guard = _locate_for_write(name, "write_file", " Create it first with action='create'.") if guard: return guard @@ -572,7 +560,6 @@ def _remove_file(name: str, file_path: str) -> Dict[str, Any]: "available_files": available if available else None} if read_guard := _background_review_read_before_write_guard(name, target, "remove_file", file_path): return read_guard - target.unlink() _rmdir_if_empty(target.parent, skill_dir) return {"success": True, "message": f"File '{file_path}' removed from skill '{name}'."} @@ -608,14 +595,12 @@ def _apply_skill_write_gate(action, name, **payload_kwargs): """Flat-shape gate: stage the full kwargs so approval can replay them; bypassed during replay.""" if action not in _ACTION_HANDLERS or _skill_gate_bypass.get(): return None - def _staging(wa): payload = {"action": action, "name": name, **{k: v for k, v in payload_kwargs.items() if v is not None}} gist_kw = {k: payload_kwargs.get(k) or "" for k in ("content", "file_path", "old_string", "new_string")} return payload, wa.skill_gist(action, name, **gist_kw) - return _run_write_gate(_staging) @@ -656,12 +641,10 @@ def _maybe_debounced_sync_push(skill_name: str) -> None: return except Exception: return - def _fire(): with suppress(Exception): from tools.skills_sync_client import maybe_push_skills maybe_push_skills(message=f"sync: {skill_name}") - with _sync_push_lock: if _sync_push_timer is not None: _sync_push_timer.cancel() # only sets an Event; never raises @@ -750,7 +733,6 @@ def skill_manage( operations, default_name=name or None, task_id=task_id, session_id=session_id) if (preflight := _background_review_preflight(action, name)) is not None: return json.dumps(preflight, ensure_ascii=False) - # Approval gate: skills are too large to review inline, so they always stage regardless # of origin; bypassed when replaying an approved staged write. args = dict(content=content, category=category, file_path=file_path, file_content=file_content, @@ -758,7 +740,6 @@ def skill_manage( absorbed_into=absorbed_into) if (gate_result := _apply_skill_write_gate(action, name, **args)) is not None: return gate_result - # Ledger pre-capture: telemetry, not a gate — failures must NEVER block the mutation. delete # destroys the whole package (consolidation may have re-homed support files first), so # complete it from the newest curator backup or a restore is hollow. @@ -768,18 +749,14 @@ def skill_manage( _pre = _find_skill(name) _ledger_before = _ledger.capture_before( _pre["path"] if _pre else None, complete_package=(action == "delete"), skill=name) - - handler = _ACTION_HANDLERS.get(action) - if handler is None: - result = _err(f"Unknown action '{action}'. Use: create, edit, patch, delete, write_file, remove_file") - else: - for arg, missing, message in _REQUIRED_ARGS.get(action, ()): - if missing(args[arg]): - return tool_error(message, success=False) - result = handler({"name": name, **args}) - if isinstance(result, str): - return result # tool_error JSON for argument-shape problems (patch) - + for arg, missing, message in _REQUIRED_ARGS.get(action, ()): + if missing(args[arg]): + return tool_error(message, success=False) + handler = _ACTION_HANDLERS.get(action, lambda a: _err( + f"Unknown action '{action}'. Use: create, edit, patch, delete, write_file, remove_file")) + result = handler({"name": name, **args}) + if isinstance(result, str): + return result # tool_error JSON for argument-shape problems (patch) if result.get("success"): _record_success( action, name, result, file_path=file_path, absorbed_into=absorbed_into, From 58a993a54d9e3cc710af4ea5fc6967742f20c112 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:30:13 -0700 Subject: [PATCH 10/16] refactor(tools): MCP group L docstring/comment compaction, >118-col fixes --- tools/mcp_oauth_manager.py | 28 ++++++++----------- tools/mcp_tool_config.py | 21 +++++--------- tools/mcp_tool_errors.py | 17 +++++------- tools/mcp_tool_handlers.py | 51 ++++++++++++++++------------------ tools/mcp_tool_health.py | 21 ++++++-------- tools/mcp_tool_registration.py | 25 ++++++++--------- tools/mcp_tool_transport.py | 16 +++++------ 7 files changed, 76 insertions(+), 103 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index f5cb5f52ea..7e3dca6af3 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -50,15 +50,13 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): def __init__(self, *args: Any, server_name: str = "", preregistered: bool = False, **kwargs: Any): super().__init__(*args, **kwargs) - # mcp 2.0 uses a task-owned anyio.Lock held across the yielded resource request (a - # session-long GET blocks every POST; HTTPX may close the generator from another task). - # A binary semaphore keeps mutual exclusion without task ownership. + # mcp 2.0 uses a task-owned anyio.Lock held across the yielded resource request (a session-long GET blocks + # every POST; HTTPX may close the generator from another task). A binary semaphore drops task ownership. import anyio self.context.lock = anyio.Semaphore(1, max_value=1) self._hermes_server_name = server_name self._hermes_home = "" - # A config-supplied client_id rejected as invalid_client means the *config* is wrong — - # re-registration can't help, so only dynamically-registered clients auto-heal. + # A config-supplied client_id rejected as invalid_client means the *config* is wrong — only DCR clients auto-heal. self._hermes_preregistered = preregistered def _hermes_storage(self): @@ -94,10 +92,9 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): self._log_nonfatal("pre-flight metadata discovery", exc) async def _prefetch_oauth_metadata(self) -> None: - """Fetch PRM + ASM from the well-known endpoints before the first request, using the - SDK's own URL builders/response handlers so we track whatever the pinned SDK expects.""" - # The SDK's httpx flavour, not Hermes' — mcp 2.0 builds on httpx2 and - # `create_oauth_metadata_request` returns *its* Request objects. + """Fetch PRM + ASM from the well-known endpoints before the first request, via the SDK's own URL + builders/response handlers so we track whatever the pinned SDK expects.""" + # The SDK's httpx flavour, not Hermes': `create_oauth_metadata_request` returns *its* (httpx2) Request objects. from tools.mcp_tool import sdk_httpx httpx = sdk_httpx() if httpx is None: # pragma: no cover — SDK import would have failed @@ -264,8 +261,8 @@ class MCPOAuthManager: self._inflight_tasks: set[asyncio.Task] = set() def get_or_build_provider(self, server_name: str, server_url: str, oauth_config: Optional[dict]) -> Optional[Any]: - """Cached OAuth provider for ``server_name``, built on first use (rebuilt when - ``server_url`` changes). None if the MCP SDK's OAuth support is unavailable.""" + """Cached OAuth provider for ``server_name``, built on first use (rebuilt when ``server_url`` changes); + None if the MCP SDK's OAuth support is unavailable.""" key = self._key(server_name) with self._entries_lock: entry = self._entries.get(key) @@ -306,8 +303,7 @@ class MCPOAuthManager: **build_provider_kwargs(cfg, storage, ssh_proxy_hint=False)) def remove(self, server_name: str, *, hermes_home: str | Path | None = None) -> _ProviderEntry | None: - """Evict the provider from cache AND delete tokens from disk (``hermes mcp remove``, - and ``hermes mcp login`` during forced re-auth).""" + """Evict the provider from cache AND delete tokens from disk (``hermes mcp remove`` / forced re-auth).""" entry = self.evict(server_name, hermes_home=hermes_home) from tools.mcp_oauth import remove_oauth_tokens remove_oauth_tokens(server_name, hermes_home=hermes_home) @@ -327,8 +323,7 @@ class MCPOAuthManager: return self._entries.pop(self._key(server_name, hermes_home), None) async def invalidate_if_disk_changed(self, server_name: str, *, hermes_home: str | Path | None = None) -> bool: - """Force the SDK provider to reload when the tokens file mtime changed; True if - invalidated. A cron job writes fresh tokens and the next tool call picks them up.""" + """Force the SDK provider to reload when the tokens file mtime changed (e.g. a cron refresh); True if so.""" from tools.mcp_oauth import _get_token_dir, _safe_filename entry = self._entries.get(self._key(server_name, hermes_home)) if entry is None or entry.provider is None: @@ -351,8 +346,7 @@ class MCPOAuthManager: """Single recovery attempt behind *pending*; always clears the dedup slot.""" can_refresh = False try: - # Disk changed (external refresh)? Else: if the SDK can refresh in place, let the - # caller retry (the httpx.Auth flow refreshes on the next request). + # Disk changed (external refresh)? Else: if the SDK can refresh in place, let the caller retry. if await self.invalidate_if_disk_changed(server_name): can_refresh = True else: diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py index 34ae1a716b..16491814d3 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -55,8 +55,7 @@ def _write_stderr_log_header(server_name: str) -> None: # Env vars safe to pass to stdio subprocesses (no secrets). _SAFE_ENV_KEYS = frozenset({"PATH", "HOME", "USER", "LANG", "LC_ALL", "TERM", "SHELL", "TMPDIR"}) -# Windows process/location vars needed by launcher-style tools (e.g. Docker Desktop's MCP plugin -# discovery); none carry secrets. +# Windows process/location vars needed by launcher-style tools (e.g. Docker Desktop's MCP plugin discovery). _SAFE_ENV_KEYS_CASE_INSENSITIVE = frozenset({ "ALLUSERSPROFILE", "APPDATA", "COMMONPROGRAMFILES", "COMMONPROGRAMFILES(X86)", "COMMONPROGRAMW6432", "COMPUTERNAME", "COMSPEC", "HOMEDRIVE", "HOMEPATH", @@ -109,8 +108,7 @@ def _build_safe_env(user_env: Optional[dict]) -> dict: def _which_with_config_pathext(command: str, path_arg, env: dict): - """``shutil.which`` retried under the config env's PATHEXT (Windows only): - ``which(path=...)`` uses the PARENT's PATHEXT, not the config env's.""" + """``shutil.which`` retried under the config env's PATHEXT (Windows only; ``which`` uses the PARENT's).""" cfg_pathext = next((v for k, v in env.items() if k.upper() == "PATHEXT" and isinstance(v, str) and v.strip()), None) if not cfg_pathext or cfg_pathext == os.environ.get("PATHEXT"): return None @@ -126,8 +124,7 @@ def _which_with_config_pathext(command: str, path_arg, env: dict): def _node_fallback(command: str) -> str: - """Well-known Node install locations for bare ``npx``/``npm``/``node`` when PATH lookup - failed; *command* unchanged when none is executable.""" + """Well-known Node install locations for bare ``npx``/``npm``/``node``; *command* unchanged when none exists.""" home = os.path.expanduser("~") hermes_home = os.path.expanduser(os.getenv("HERMES_HOME", os.path.join(home, ".hermes"))) # /usr/local/bin: canonical Node location (from-source Linux, Hermes Docker image, Intel Homebrew), @@ -138,8 +135,7 @@ def _node_fallback(command: str) -> str: def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: - """Resolve a stdio command against the exact subprocess env, mainly so bare - ``npx``/``npm``/``node`` work under a filtered PATH.""" + """Resolve a stdio command against the exact subprocess env (bare ``npx``/``npm``/``node`` under a filtered PATH).""" resolved_command = os.path.expanduser(str(command).strip()) resolved_env = dict(env or {}) if os.sep not in resolved_command: @@ -189,8 +185,7 @@ def _interpolate_env_vars(value): return value -# (server_name, dotted key path) pairs already warned about: config loads happen on every discovery -# pass, so warn once per process. +# (server_name, dotted key path) pairs already warned about: config loads repeat per discovery pass. _whitespace_warned: Set[Tuple[str, str]] = set() @@ -238,8 +233,7 @@ def _filter_suspicious_mcp_servers(servers: Dict[str, dict]) -> Dict[str, dict]: def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None: - """Merge plugin-provided (portable) MCP servers into *safe_servers*; native config wins - on a name clash. Never raises.""" + """Merge plugin-provided (portable) MCP servers into *safe_servers*; native config wins on a clash. Never raises.""" try: from hermes_cli.plugins import discover_plugins, get_plugin_manager discover_plugins() @@ -254,8 +248,7 @@ def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None: def _load_mcp_config() -> Dict[str, dict]: - """Read ``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error or in safe - mode); ``${VAR}`` placeholders are interpolated after ``.env`` is loaded.""" + """``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error / safe mode), ``${VAR}`` interpolated.""" try: from hermes_cli.config import load_config from utils import env_var_enabled as _env_enabled diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index e56bd4cc88..74573d046b 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -14,8 +14,7 @@ from tools.mcp_tool_common import _sanitize_error, _core logger = logging.getLogger("tools.mcp_tool") -# Stateless (2026-07-28) servers reject a legacy ``initialize`` with -# UnsupportedProtocolVersion (-32022) or plain method-not-found. +# Stateless (2026-07-28) servers reject a legacy ``initialize`` with this or plain method-not-found. _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022 @@ -44,8 +43,7 @@ def _is_method_not_found_error(exc: BaseException) -> bool: class InvalidMcpUrlError(ValueError): - """A remote MCP server's ``url`` is not parseable http(s):// — validated once at startup so we - fail fast instead of burning the reconnect-backoff loop.""" + """A remote MCP server's ``url`` is not parseable http(s):// — validated once at startup to fail fast.""" class NonMcpEndpointError(ConnectionError): @@ -88,8 +86,8 @@ def _classify_mcp_failure(exc: BaseException) -> str: def _validate_remote_mcp_url(server_name: str, url: Any) -> str: - """The stripped URL if it is a valid http(s) URL; else InvalidMcpUrlError naming the server - (non-string, other scheme — stdio servers use ``command`` — or empty host).""" + """The stripped URL if valid http(s); else InvalidMcpUrlError naming the server (non-string, other scheme — + stdio servers use ``command`` — or empty host).""" def _bad(detail: str) -> InvalidMcpUrlError: return InvalidMcpUrlError(f"Invalid MCP URL for '{server_name}': {detail}") @@ -286,8 +284,8 @@ _SESSION_EXPIRED_MARKERS: tuple = ( "unknown session", "session terminated", "closedresourceerror", "closed resource", "transport is closed", "connection closed", "broken pipe", "end of file") -# Node budget for ``_is_session_expired_error`` (the visited set breaks cycles; this bounds acyclic -# blow-ups). Well above ``sys.getrecursionlimit()`` so deep task-group nesting is fully scanned. +# Node budget for ``_is_session_expired_error`` (the visited set breaks cycles; this bounds acyclic blow-ups). +# Well above ``sys.getrecursionlimit()`` so deep task-group nesting is fully scanned. _EXC_TRAVERSAL_MAX_NODES = 10_000 @@ -311,8 +309,7 @@ def _is_session_expired_error(exc: BaseException) -> bool: budget -= 1 if isinstance(current, InterruptedError): return False - # Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids - # false positives. + # Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids false positives. msg = str(current).lower() found = found or isinstance(current, transport_error_types) or any(m in msg for m in _SESSION_EXPIRED_MARKERS) stack.extend(getattr(current, "exceptions", ())) diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 69b5c40341..51ee49393a 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -21,14 +21,16 @@ from tools.mcp_tool_errors import _is_session_expired_error logger = logging.getLogger("tools.mcp_tool") _MISSING = object() -_NEEDS_REAUTH_MSG = ("MCP server '{s}' requires re-authentication. Run `hermes mcp login {s}` (or delete the tokens " - "file under ~/.hermes/mcp-tokens/ and restart). Do NOT retry this tool — ask the user to re-authenticate.") -_STDIO_NO_RESPAWN_MSG = ("MCP server '{s}' stdio subprocess had exited (this is not a timeout — the call never reached the " - "server). A respawn was requested but no fresh session came back within {t:.0f}s. Wait a few " - "seconds before retrying; if it keeps failing the server is not starting and needs the user.") -_STDIO_DIED_AGAIN_MSG = ("MCP server '{s}' respawned its stdio subprocess and it exited again immediately. The server is not " - "starting cleanly — do NOT retry this tool; ask the user to check the server's command and its " - "stderr log.") +_NEEDS_REAUTH_MSG = ( + "MCP server '{s}' requires re-authentication. Run `hermes mcp login {s}` (or delete the tokens file under " + "~/.hermes/mcp-tokens/ and restart). Do NOT retry this tool — ask the user to re-authenticate.") +_STDIO_NO_RESPAWN_MSG = ( + "MCP server '{s}' stdio subprocess had exited (this is not a timeout — the call never reached the server). A " + "respawn was requested but no fresh session came back within {t:.0f}s. Wait a few seconds before retrying; if it " + "keeps failing the server is not starting and needs the user.") +_STDIO_DIED_AGAIN_MSG = ( + "MCP server '{s}' respawned its stdio subprocess and it exited again immediately. The server is not starting " + "cleanly — do NOT retry this tool; ask the user to check the server's command and its stderr log.") # --------------------------------------------------------------- pre-call gates @@ -77,7 +79,8 @@ def _acquire_call_server(server_name: str, tool_timeout: float): server task to rebuild (probing a dead transport would re-arm the breaker forever).""" not_connected = tool_error(f"MCP server '{server_name}' is not connected") server = _core._get_connected_server_for_call(server_name) - if server and (server.session or _core._wait_for_server_session_ready(server, timeout=min(5.0, float(tool_timeout or 5.0)))): + wait = min(5.0, float(tool_timeout or 5.0)) + if server and (server.session or _core._wait_for_server_session_ready(server, timeout=wait)): return server, None _core._bump_server_error(server_name) if server and _core._signal_reconnect(server): @@ -151,8 +154,8 @@ def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_cal recovered = False if recovered: srv = _lookup_reconnectable_server(server_name) - # Recovery + reconnect is independent evidence of viability: close the breaker here, not - # only on retry success (else a failing retry pins it open forever). + # Recovery + reconnect is independent evidence of viability: close the breaker here, not only on + # retry success (else a failing retry pins it open forever). if srv is not None and _core._signal_reconnect_and_wait( server_name, srv, op_description=f"{op_description} after OAuth recovery", timeout=15): _core._reset_server_error(server_name) @@ -298,15 +301,13 @@ async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str # ---------------------------------------------------------- result rendering def _error_result_text(result) -> str: - """Concatenated text of an ``isError`` result's blocks (EmbeddedResource error payloads - carry text under ``.resource.text``).""" + """Concatenated text of an ``isError`` result's blocks (EmbeddedResource payloads: ``.resource.text``).""" texts = (getattr(b, "text", None) or getattr(getattr(b, "resource", None), "text", None) for b in (result.content or [])) return "".join(str(t) for t in texts if t) def _render_content_blocks(result, server_name: str) -> str: - """Text passes through; image/audio blocks are cached (MEDIA: tags); resource blocks are - materialized rather than silently dropped.""" + """Text passes through; image/audio blocks are cached (MEDIA: tags); resource blocks are materialized.""" parts: List[str] = [] for block in (result.content or []): if getattr(block, "text", None): @@ -316,9 +317,8 @@ def _render_content_blocks(result, server_name: str) -> str: if rendered: parts.append(rendered) continue - # Benign empty renders log at debug; warn only for unknown shapes. block_type = getattr(block, "type", None) or type(block).__name__ - if block_type in {"text", "resource", "audio", "image"}: + if block_type in {"text", "resource", "audio", "image"}: # benign empty render logger.debug("MCP %s: content block type %r rendered empty", server_name, block_type) else: logger.warning("MCP %s: dropping unsupported content block type %r", server_name, block_type) @@ -326,8 +326,7 @@ def _render_content_blocks(result, server_name: str) -> str: def _capped_structured_content(result): - """``structuredContent`` (or None); over the hard cap it degrades to the head+tail - truncated JSON string (multi-MB JSON flood guard).""" + """``structuredContent`` (or None); over the hard cap it degrades to the truncated JSON string (flood guard).""" structured = mcp_field(result, "structured_content", "structuredContent") try: as_json = json.dumps(structured, ensure_ascii=False, default=str) if structured is not None else "" @@ -337,8 +336,8 @@ def _capped_structured_content(result): def _render_call_tool_result(result, server_name: str) -> str: - """Pure: ``CallToolResult`` -> handler JSON. ``content`` is primary; ``structuredContent`` - supplements it (or becomes ``result`` without text); ``_meta`` minus reserved keys.""" + """Pure: ``CallToolResult`` -> handler JSON. ``content`` is primary; ``structuredContent`` supplements it (or + becomes ``result`` without text); ``_meta`` minus reserved keys.""" if mcp_field(result, "is_error", "isError", False): return tool_error(_sanitize_error(_truncate_mcp_text_result(_error_result_text(result) or "MCP tool returned an error"))) text_result = _render_content_blocks(result, server_name) @@ -346,8 +345,7 @@ def _render_call_tool_result(result, server_name: str) -> str: meta = _strip_reserved_meta_keys(mcp_field(result, "meta", "meta")) if structured is None and meta is None: return json.dumps({"result": text_result}, ensure_ascii=False) - # Key order is part of the output: "result" leads when there is text, otherwise "_meta" - # precedes the (empty) "result". + # Key order is part of the output: "result" leads when there is text, otherwise "_meta" precedes it. payload: Dict[str, Any] = {"result": text_result} if text_result else {} if structured is not None: payload["structuredContent" if text_result else "result"] = structured @@ -426,8 +424,8 @@ def _make_utility_handler(op: str, log_label: str, rpc, render, required: Option def _pick(obj, *specs) -> dict: - """``{out_key: value}`` for each ``(out_key, attr[, truthy])`` present on *obj* (``hasattr`` - so SDK models and stubs behave alike; ``truthy`` also skips falsy). Key order = spec order.""" + """``{out_key: value}`` for each ``(out_key, attr[, truthy])`` present on *obj* (presence check so SDK models + and stubs behave alike; ``truthy`` also skips falsy). Key order = spec order.""" entry = {} for out_key, attr, *truthy in specs: value = getattr(obj, attr, _MISSING) @@ -497,8 +495,7 @@ _make_get_prompt_handler = _make_utility_handler( def _make_check_fn(server_name: str): - """Check function that verifies the MCP connection is alive. Lazy (schema-cache registered) - servers count as available: the first real call spawns/connects them.""" + """Connection-alive check; lazy (schema-cache registered) servers count as available.""" def _check() -> bool: with _core._lock: server = _core._servers.get(server_name) diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index 93d8576af3..14b7f1962f 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -38,8 +38,7 @@ class MCPServerHealthMixin: self._recycled_reason = None def _stdio_recycle_deadlines(self): - """``[(deadline, reason), ...]`` for the configured lifetime/idle limits; empty for HTTP - servers or while an RPC holds the lock.""" + """``[(deadline, reason), ...]`` for the lifetime/idle limits; empty for HTTP or while an RPC holds the lock.""" if self._is_http() or self._rpc_lock.locked(): return [] limits = ((self._lifecycle_started_at, self._max_lifetime_seconds, "max_lifetime_seconds"), @@ -74,8 +73,7 @@ class MCPServerHealthMixin: return task def _make_logging_callback(self): - """``logging_callback`` forwarding server ``notifications/message`` into Hermes logging - tagged with the server name (the SDK default drops them).""" + """``logging_callback`` forwarding server ``notifications/message`` into Hermes logging (SDK default drops them).""" async def _on_log(params): try: level = _core._MCP_LOG_LEVEL_MAP.get(str(getattr(params, "level", "info")).lower(), logging.INFO) @@ -88,14 +86,14 @@ class MCPServerHealthMixin: if len(data) > 2000: # cap payloads so a chatty server can't flood agent.log data = data[:2000] + "... [truncated]" logger_name = getattr(params, "logger", None) - logger.log(level, "MCP server log [%s]: %s", f"{self.name}/{logger_name}" if logger_name else self.name, data) + origin = f"{self.name}/{logger_name}" if logger_name else self.name + logger.log(level, "MCP server log [%s]: %s", origin, data) except Exception: logger.debug("Failed to handle MCP log notification from '%s'", self.name, exc_info=True) return _on_log def _make_message_handler(self): - """``message_handler`` for ``ClientSession``: only ``ToolListChangedNotification`` triggers - a refresh; prompt/resource changes are logged.""" + """``message_handler``: only ``ToolListChangedNotification`` triggers a refresh; prompt/resource changes log.""" async def _handler(message): try: if isinstance(message, Exception): @@ -121,8 +119,7 @@ class MCPServerHealthMixin: return _handler def _deregister_owned(self, tool_names: Iterable[str]) -> None: - """Deregister *tool_names* this server's toolset still owns. Never removes a colliding - name currently owned by another server.""" + """Deregister *tool_names* this server's toolset still owns (never a colliding name owned by another server).""" from tools.registry import registry for tool_name in tool_names: if registry.get_toolset_for_tool(tool_name) == f"mcp-{self.name}": @@ -206,8 +203,7 @@ class MCPServerHealthMixin: self._permanent_grace_used = self._teardown_race = False def mark_suspect(self, reason: str) -> None: - """Latch a suspicion (no I/O). The NEXT call verifies via :meth:`ensure_healthy` and - recycles the transport if the probe fails.""" + """Latch a suspicion (no I/O); the NEXT call verifies via :meth:`ensure_healthy` and recycles on failure.""" if self._suspect_reason is None and reason: logger.warning("MCP server '%s': connection marked suspect (%s); next call will health-check it", self.name, reason) @@ -269,7 +265,6 @@ class MCPServerHealthMixin: return False async def _watch_stdio_children(self) -> None: - """Poll child liveness while a stdio RPC is in flight; resolves when a tracked child dies - so the caller cancels the RPC instead of waiting out the timeout.""" + """Poll child liveness during a stdio RPC; resolves when a tracked child dies so the caller cancels the RPC.""" while not self._stdio_children_dead(): await asyncio.sleep(0.25) diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index 1984422109..18e4c82092 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -27,20 +27,19 @@ _UTILITY_HANDLER_FACTORIES = { def _normalize_server_trust(value: Any) -> str: - """Config ``trust`` -> tier. None -> ``full`` (backward-compatible default); an - unrecognized string -> ``untrusted`` so a misspelled tier fails closed.""" + """Config ``trust`` -> tier. None -> ``full`` (compat default); unrecognized -> ``untrusted`` (fail closed).""" if value is None: return _core._TRUST_FULL text = str(value).strip().lower() if text in (_core._TRUST_FULL, _core._TRUST_UNTRUSTED): return text - logger.warning("MCP trust: unrecognized trust value %r — treating as 'untrusted' (valid values: full, untrusted)", value) + logger.warning("MCP trust: unrecognized trust value %r — treating as 'untrusted' (valid values: full, untrusted)", + value) return _core._TRUST_UNTRUSTED def _annotation_read_only_hint(mcp_tool: Any) -> bool: - """True only when annotations (SDK object or schema-cache dict) carry ``readOnlyHint is - True``; unknown metadata means write-capable.""" + """True only when annotations (SDK object or cache dict) carry ``readOnlyHint is True``; unknown = write-capable.""" annotations = getattr(mcp_tool, "annotations", None) hint = annotations.get("readOnlyHint") if isinstance(annotations, dict) else getattr(annotations, "readOnlyHint", None) return hint is True @@ -81,7 +80,9 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d if not enabled[family]: return f"{family} disabled" if advertised is not None: - return None if getattr(advertised, family, None) is not None else f"server does not advertise '{family}' capability" + if getattr(advertised, family, None) is None: + return f"server does not advertise '{family}' capability" + return None # Legacy gate (no initialize_result): the ClientSession method shares the handler key. return None if hasattr(server.session, handler_key) else f"session lacks {handler_key}" @@ -96,8 +97,7 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d def _existing_tool_names() -> List[str]: - """Tool names for all currently connected servers plus lazy (cache-registered) servers, - whose tools live only in the registry.""" + """Tool names for all connected servers plus lazy (cache-registered) servers, whose tools live only in the registry.""" names: List[str] = [] for server in _core._servers.values(): names.extend(server._registered_tool_names if hasattr(server, "_registered_tool_names") @@ -131,8 +131,7 @@ def _cached_tools(raws: Iterable[Any]) -> List[SimpleNamespace]: @dataclass class _Candidate: - """One registration attempt: a native tool or a generated utility. ``origin`` is the - provenance text used in collision diagnostics.""" + """One registration attempt (native tool or generated utility); ``origin`` is the provenance text in diagnostics.""" registry_name: str origin: str @@ -155,7 +154,8 @@ def _tool_candidates(name: str, tools: Iterable[Any], should_register: Callable[ continue _core._scan_mcp_description(name, t.name, t.description or "") schema = _core._convert_mcp_schema(name, t) - out.append(_Candidate(schema["name"], f"tool {t.name!r}", schema, _core._make_tool_handler(name, t.name, tool_timeout))) + handler = _core._make_tool_handler(name, t.name, tool_timeout) + out.append(_Candidate(schema["name"], f"tool {t.name!r}", schema, handler)) return out @@ -247,8 +247,7 @@ def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: C def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should_register) -> None: - """Write-through: persist the manifest so the next startup can register this server - lazily without spawning it. Never raises.""" + """Write-through: persist the manifest so the next startup registers this server lazily (no spawn). Never raises.""" try: from tools.mcp_schema_cache import config_fingerprint, write_cache_entry tools_payload = [{ diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index 1d16c998f7..b523c0fb1f 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -13,8 +13,7 @@ from tools.mcp_tool_common import _core logger = logging.getLogger("tools.mcp_tool") -# JSON-RPC ``initialize`` body used by the content-type preflight POST. -_PROBE_INITIALIZE_BODY = ( +_PROBE_INITIALIZE_BODY = ( # JSON-RPC ``initialize`` body for the content-type preflight POST '{"jsonrpc":"2.0","id":"_probe","method":"initialize","params":{"protocolVersion":"2025-03-26",' '"capabilities":{},"clientInfo":{"name":"hermes-probe","version":"0.1"}}}') @@ -43,8 +42,8 @@ def _pgroup_alive(pgid: Optional[int]) -> bool: async def _osv_malware_preflight(server_name: str, command: str, args: list) -> None: - """OSV malware preflight, off-loop with a wall-clock bound (fail-open on timeout). Must run on - the REAL command/args — the watchdog wrap rewrites argv to the supervisor (check becomes a no-op).""" + """OSV malware preflight, off-loop with a wall-clock bound (fail-open on timeout). Must run on the REAL + command/args — the watchdog wrap rewrites argv to the supervisor (check becomes a no-op).""" from tools.osv_check import check_package_for_malware try: malware_error = await asyncio.wait_for( @@ -207,7 +206,8 @@ class MCPServerTransportMixin: # Subprocess stderr goes to ~/.hermes/logs/mcp-stderr.log so banners can't corrupt the TUI. _core._write_stderr_log_header(self.name) try: - async with _core.stdio_client(server_params, errlog=_core._get_mcp_stderr_log()) as (read_stream, write_stream): + errlog = _core._get_mcp_stderr_log() + async with _core.stdio_client(server_params, errlog=errlog) as (read_stream, write_stream): # New PIDs for force-kill cleanup, minus non-MCP children (slash_worker, LSP) racing # into the window: they share the TUI's pgid — leaking them would killpg() the TUI. new_pids = _filter_mcp_children(_core._snapshot_child_pids() - pids_before) @@ -239,8 +239,7 @@ class MCPServerTransportMixin: return # No httpx → skip probe; SDK import would have failed first. def _non_mcp_2xx(resp) -> bool: - # Only judge 2xx (4xx/5xx may be an auth challenge the handshake handles); no content - # type advertised → don't second-guess the SDK. + # Only judge 2xx (4xx/5xx may be an auth challenge); no content type advertised → trust the SDK. ct = _content_type_base(resp) return _is_2xx(resp) and bool(ct) and ct not in self._MCP_CONTENT_TYPES @@ -248,8 +247,7 @@ class MCPServerTransportMixin: try: async with _httpx.AsyncClient(verify=ssl_verify, follow_redirects=True, timeout=_httpx.Timeout(timeout), **_present(cert=client_cert)) as client: - # HEAD is cheapest; fall back to GET on 405/501. - resp = await client.head(url, headers=probe_headers) + resp = await client.head(url, headers=probe_headers) # cheapest; GET on 405/501 if resp.status_code in (405, 501): resp = await client.get(url, headers=probe_headers) # Non-MCP content type on HEAD/GET: try a JSON-RPC POST so POST-only servers pass. From e6048ff6aabf9a9336dbe656b52b828104236d13 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:30:26 -0700 Subject: [PATCH 11/16] =?UTF-8?q?refactor(tools):=20skill=5Fmanage/ledger?= =?UTF-8?q?=20=E2=80=94=20shared=20identifier=20check,=20chainable=20resul?= =?UTF-8?q?t=20decorators,=20ledger=20read/filter=20fold?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/skill_ledger.py | 51 +++++++++------------------ tools/skill_manager_batch.py | 7 ++-- tools/skill_manager_guards.py | 8 ++--- tools/skill_manager_tool.py | 65 +++++++++++++++-------------------- 4 files changed, 49 insertions(+), 82 deletions(-) diff --git a/tools/skill_ledger.py b/tools/skill_ledger.py index 637f26f2e2..baf281ace6 100644 --- a/tools/skill_ledger.py +++ b/tools/skill_ledger.py @@ -63,18 +63,18 @@ def derive_actor() -> str: return "agent" +def _skills_dir() -> Path: + return get_hermes_home() / "skills" + + def ledger_path() -> Path: - return get_hermes_home() / "skills" / ".curator_ledger.jsonl" + return _skills_dir() / ".curator_ledger.jsonl" def blobs_dir() -> Path: return get_hermes_home() / ".curator_backups" / "blobs" -def _skills_dir() -> Path: - return get_hermes_home() / "skills" - - def ledger_enabled() -> bool: """Config gate ``skills.ledger`` (default True); lazy import keeps this importable without the CLI.""" try: @@ -85,14 +85,10 @@ def ledger_enabled() -> bool: return True -def _norm(path: Path | str) -> Path: - return Path(os.path.normpath(str(path))) - - def _rel_posix(path: Path | str, root: Path) -> Optional[str]: """POSIX path of ``path`` relative to ``root`` (both normalized), or None when outside.""" try: - return _norm(path).relative_to(_norm(root)).as_posix() + return Path(os.path.normpath(str(path))).relative_to(os.path.normpath(str(root))).as_posix() except (ValueError, TypeError): return None @@ -118,11 +114,9 @@ def read_blob(sha256: str) -> Optional[bytes]: """Return blob content or None when missing/invalid.""" if not sha256 or not all(c in "0123456789abcdef" for c in sha256): return None - try: - p = blobs_dir() / sha256 - return p.read_bytes() if p.exists() else None - except OSError: - return None + with suppress(OSError): + return (blobs_dir() / sha256).read_bytes() if (blobs_dir() / sha256).exists() else None + return None def snapshot_paths(root: Optional[Path], *, complete_package: bool = False) -> List[Dict[str, str]]: @@ -132,15 +126,9 @@ def snapshot_paths(root: Optional[Path], *, complete_package: bool = False) -> L tarball's files (disk hashes win).""" if root is None: return [] - root = Path(root) - if root.is_file(): - files = [root] - elif root.is_dir(): - files = sorted(p for p in root.rglob("*") if p.is_file()) - elif complete_package: - files = [] # gone from disk; the backup fill below may still recover it - else: - return [] + root = Path(root) # gone from disk -> []; the complete_package fill may still recover it + files = ([root] if root.is_file() + else sorted(p for p in root.rglob("*") if p.is_file()) if root.is_dir() else []) out = [{"path": str(f), "sha256": _store_blob(f.read_bytes())} for f in files] return fill_snapshot_from_curator_backup(root, out) if complete_package else out @@ -160,8 +148,7 @@ def _strip_archive_timestamp(name: str) -> str: def _skill_md_parents(items: Optional[List[Dict[str, str]]]) -> List[Path]: - paths = [Path(str(item.get("path", ""))) for item in items or []] - return [p.parent for p in paths if p.name == "SKILL.md"] + return [p.parent for p in (Path(str(i.get("path", ""))) for i in items or []) if p.name == "SKILL.md"] def package_prefixes( @@ -310,9 +297,7 @@ def capture_before( return None try: captured = snapshot_paths(root) - if complete_package: - captured = fill_snapshot_from_curator_backup(root, captured, skill=skill) - return captured + return fill_snapshot_from_curator_backup(root, captured, skill=skill) if complete_package else captured except Exception as e: logger.warning("skill_ledger: before-capture failed (%s) — mutation unaffected", e) return None @@ -328,18 +313,14 @@ def list_entries(skill: Optional[str] = None, limit: Optional[int] = None) -> Li for line in lines: with suppress(json.JSONDecodeError): row = json.loads(line) if line.strip() else None - if isinstance(row, dict): + if isinstance(row, dict) and (not skill or row.get("skill") == skill): rows.append(row) - if skill: - rows = [r for r in rows if r.get("skill") == skill] rows.reverse() return rows[:limit] if limit is not None and limit >= 0 else rows def get_entry(entry_id: str) -> Optional[Dict[str, Any]]: - if not entry_id: - return None - return next((row for row in list_entries() if row.get("id") == entry_id), None) + return next((r for r in list_entries() if r.get("id") == entry_id), None) if entry_id else None def _validate_entry_paths(entry: Dict[str, Any]) -> Optional[str]: diff --git a/tools/skill_manager_batch.py b/tools/skill_manager_batch.py index c5fe07118b..6eae86588d 100644 --- a/tools/skill_manager_batch.py +++ b/tools/skill_manager_batch.py @@ -34,8 +34,7 @@ def _validate_batch_ops(operations, default_name, tool_error): names.append(nm) if act == "create" and nm in names[:-1]: return fail(i, f": create for '{nm}' must precede that skill's other ops.") - preflight = _background_review_preflight(act, nm) - if preflight is not None: + if (preflight := _background_review_preflight(act, nm)) is not None: return None, json.dumps(preflight, ensure_ascii=False) # Clobber guard: a DESTRUCTIVE op (create/write_file/remove_file/full rewrite) on # a file an earlier op touched would SILENTLY discard its work — reject it. @@ -156,8 +155,8 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non token = _smt._skill_gate_bypass.set(True) try: for i, op in enumerate(operations): - raw = _smt._skill_manage_from( - {**op, "name": names[i], "operations": None}, task_id=task_id, session_id=session_id) + raw = _smt._skill_manage_from({**op, "name": names[i], "operations": None}, + task_id=task_id, session_id=session_id) try: parsed = json.loads(raw) except Exception: # noqa: BLE001 diff --git a/tools/skill_manager_guards.py b/tools/skill_manager_guards.py index 3b4d913a1b..12e321669e 100644 --- a/tools/skill_manager_guards.py +++ b/tools/skill_manager_guards.py @@ -27,10 +27,9 @@ def _is_background_review() -> bool: def _resolved_str(path: Path) -> str: - try: + with suppress(Exception): return str(path.resolve()) - except Exception: - return str(path) + return str(path) class _BackgroundReviewReadMarks: @@ -59,8 +58,7 @@ def mark_background_review_skill_read(path: Path) -> None: the write guards require the mark.""" if not _is_background_review(): return - marks = _background_review_read_paths.get() - if marks is None: + if (marks := _background_review_read_paths.get()) is None: _background_review_read_paths.set(marks := _BackgroundReviewReadMarks()) marks.add(_resolved_str(path)) diff --git a/tools/skill_manager_tool.py b/tools/skill_manager_tool.py index a0e42bfb29..c8e2f2e850 100644 --- a/tools/skill_manager_tool.py +++ b/tools/skill_manager_tool.py @@ -45,8 +45,7 @@ def _guard_agent_created_enabled() -> bool: """skills.guard_agent_created (default False): opt-in — terminal() runs the same code ungated.""" try: from hermes_cli.config import load_config - return is_truthy_value( - cfg_get(load_config(), "skills", "guard_agent_created"), default=False) + return is_truthy_value(cfg_get(load_config(), "skills", "guard_agent_created"), default=False) except Exception: return False @@ -102,14 +101,17 @@ def _display_create_dir() -> str: # --- Validation helpers ------------------------------------------------------- +def _check_identifier(value: str, label: str, invalid: str) -> Optional[str]: + if len(value) > MAX_NAME_LENGTH: + return f"{label} exceeds {MAX_NAME_LENGTH} characters." + return None if VALID_NAME_RE.match(value) else invalid + + def _validate_name(name: str) -> Optional[str]: if not name: return "Skill name is required." - if len(name) > MAX_NAME_LENGTH: - return f"Skill name exceeds {MAX_NAME_LENGTH} characters." - if not VALID_NAME_RE.match(name): - return f"Invalid skill name '{name}'. {_NAME_RULE} Must start with a letter or digit." - return None + return _check_identifier( + name, "Skill name", f"Invalid skill name '{name}'. {_NAME_RULE} Must start with a letter or digit.") def _validate_category(category: Optional[str]) -> Optional[str]: @@ -122,9 +124,7 @@ def _validate_category(category: Optional[str]) -> Optional[str]: "Categories must be a single directory name.") if "/" in category or "\\" in category: return invalid - if len(category) > MAX_NAME_LENGTH: - return f"Category exceeds {MAX_NAME_LENGTH} characters." - return None if VALID_NAME_RE.match(category) else invalid + return _check_identifier(category, "Category", invalid) def _validate_frontmatter(content: str, *, new_skill: bool = False) -> Optional[str]: @@ -339,18 +339,20 @@ def _guarded_write(name: str, skill_dir: Path, target: Path, action: str, label: return _err(scan_error) -def _attach_org_note(result: Dict[str, Any], name: str, skill_dir: Path) -> None: +def _attach_org_note(result: Dict[str, Any], name: str, skill_dir: Path) -> Dict[str, Any]: if org_note := _maybe_auto_propose_org_edit(name, skill_dir): result["org_sharing"] = org_note result["message"] = f"{result['message']} {org_note}" + return result -def _add_description_prompt_preview(result: Dict[str, Any], content: str) -> None: +def _add_description_prompt_preview(result: Dict[str, Any], content: str) -> Dict[str, Any]: fm, _ = _parse_frontmatter(content) if is_skill_description_truncated_for_prompt(fm): result["system_prompt_preview"] = ( f"System prompt will show: \"{extract_skill_description(fm)}\" — keep the trigger " f"self-contained in the first {SKILL_PROMPT_DESC_LIMIT - 3} chars.") + return result def _attach_lint_findings(result: Dict[str, Any], skill_md: Path) -> None: @@ -376,9 +378,8 @@ def _clip(text: str, n: int, ellipsis: str) -> str: # --- Core actions ------------------------------------------------------------- def _create_skill(name: str, content: str, category: str = None) -> Dict[str, Any]: - err = (_validate_name(name) or _validate_category(category) - or _validate_frontmatter(content, new_skill=True) or _validate_content_size(content)) - if err: + if err := (_validate_name(name) or _validate_category(category) + or _validate_frontmatter(content, new_skill=True) or _validate_content_size(content)): return _err(err) if existing := _find_skill(name): return _err(f"A skill named '{name}' already exists at {existing['path']}.") @@ -389,8 +390,7 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An if scan_error := _security_scan_skill(skill_dir): shutil.rmtree(skill_dir, ignore_errors=True) return _err(scan_error) - root = _skills_dir() - # Relative when under the profile dir; absolute when created under skills.create_dir. + root = _skills_dir() # display relative under the profile dir; absolute under skills.create_dir display = skill_dir.relative_to(root) if skill_dir.is_relative_to(root) else skill_dir result = { "success": True, "message": f"Skill '{name}' created.", "path": str(display), @@ -399,8 +399,7 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An "hint": "To add reference files, templates, or scripts, use " f"skill_manage(action='write_file', name='{name}', file_path='references/example.md', " "file_content='...')"} - _add_description_prompt_preview(result, content) - _attach_lint_findings(result, skill_md) + _attach_lint_findings(_add_description_prompt_preview(result, content), skill_md) return result @@ -410,15 +409,12 @@ def _edit_skill(name: str, content: str) -> Dict[str, Any]: return _err(err) skill_dir, guard = _locate_for_write(name, "edit") # SKILL.md always exists here (_find_skill requires it), so a blocked scan restores it. - if guard := guard or _guarded_write( - name, skill_dir, skill_dir / "SKILL.md", "edit", "SKILL.md", content): + if guard := guard or _guarded_write(name, skill_dir, skill_dir / "SKILL.md", "edit", "SKILL.md", content): return guard result = { "success": True, "message": f"Skill '{name}' updated (full rewrite).", "path": str(skill_dir), "_change": {"description": _description_preview(content)}} - _attach_org_note(result, name, skill_dir) - _add_description_prompt_preview(result, content) - return result + return _add_description_prompt_preview(_attach_org_note(result, name, skill_dir), content) def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = None, @@ -471,8 +467,7 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N "success": True, "message": f"Patched {target_label} in skill '{name}' ({match_count} replacement{'s' if match_count > 1 else ''}).", "_change": {"old": _clip(old_string, 200, "…"), "new": _clip(new_string, 200, "…")}} - _attach_org_note(result, name, skill_dir) - return result + return _attach_org_note(result, name, skill_dir) def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, Any]: @@ -523,8 +518,7 @@ def _write_file(name: str, file_path: str, file_content: str) -> Dict[str, Any]: return _err(err) if not file_content and file_content != "": return _err("file_content is required.") - content_bytes = len(file_content.encode("utf-8")) - if content_bytes > MAX_SKILL_FILE_BYTES: + if (content_bytes := len(file_content.encode("utf-8"))) > MAX_SKILL_FILE_BYTES: return _err(f"File content is {content_bytes:,} bytes (limit: {MAX_SKILL_FILE_BYTES:,} " f"bytes / 1 MiB). Consider splitting into smaller files.") if err := _validate_content_size(file_content, label=file_path): @@ -535,10 +529,8 @@ def _write_file(name: str, file_path: str, file_content: str) -> Dict[str, Any]: target, err = _resolve_supporting_file(skill_dir, file_path) if guard := err or _guarded_write(name, skill_dir, target, "write_file", file_path, file_content): return guard - result = {"success": True, "message": f"File '{file_path}' written to skill '{name}'.", - "path": str(target)} - _attach_org_note(result, name, skill_dir) - return result + return _attach_org_note({"success": True, "message": f"File '{file_path}' written to skill '{name}'.", + "path": str(target)}, name, skill_dir) def _remove_file(name: str, file_path: str) -> Dict[str, Any]: @@ -552,12 +544,9 @@ def _remove_file(name: str, file_path: str) -> Dict[str, Any]: if err: return err if not target.exists(): # list what IS there so the model can pick the right path - available = [ - str(f.relative_to(skill_dir)) for subdir in ALLOWED_SUBDIRS - if (skill_dir / subdir).exists() - for f in (skill_dir / subdir).rglob("*") if f.is_file()] - return {"success": False, "error": f"File '{file_path}' not found in skill '{name}'.", - "available_files": available if available else None} + available = [str(f.relative_to(skill_dir)) for subdir in ALLOWED_SUBDIRS + if (skill_dir / subdir).exists() for f in (skill_dir / subdir).rglob("*") if f.is_file()] + return _err(f"File '{file_path}' not found in skill '{name}'.", available_files=available or None) if read_guard := _background_review_read_before_write_guard(name, target, "remove_file", file_path): return read_guard target.unlink() From bccfd1de261ea86046a5aaa548b5209aad298edc Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:31:18 -0700 Subject: [PATCH 12/16] refactor(tools): MCP handlers inline single-use render helpers, drop banners; registration foreign-owner log inlined; body blank squeeze --- tools/mcp_oauth_manager.py | 2 -- tools/mcp_tool_config.py | 2 -- tools/mcp_tool_errors.py | 5 ---- tools/mcp_tool_handlers.py | 49 +++++++++------------------------- tools/mcp_tool_health.py | 2 -- tools/mcp_tool_registration.py | 28 ++++++++----------- tools/mcp_tool_transport.py | 3 --- 7 files changed, 23 insertions(+), 68 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index 7e3dca6af3..c5cdda47ca 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -111,7 +111,6 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): except httpx.HTTPError as exc: logger.debug("MCP OAuth '%s': %s discovery to %s failed: %s", self._hermes_server_name, label, url, exc) return None - async with httpx.AsyncClient(timeout=10.0) as client: # PRM discovery to learn the authorization_server URL. for url in build_protected_resource_metadata_discovery_urls(None, server_url): @@ -240,7 +239,6 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): import anyio with anyio.CancelScope(shield=True): await self.context.lock.acquire() - if retry_after_concurrent_auth: yield request self._persist_oauth_metadata_if_changed() diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py index 16491814d3..bd70d416d0 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -147,7 +147,6 @@ def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: resolved_command = which_hit elif resolved_command in {"npx", "npm", "node"}: resolved_command = _node_fallback(resolved_command) - command_dir = os.path.dirname(resolved_command) if command_dir: resolved_env = _prepend_path(resolved_env, command_dir) @@ -204,7 +203,6 @@ def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: elif isinstance(value, list): for i, v in enumerate(value): _walk(v, f"{path}[{i}]") - _walk(config, "") for key_path in flagged: if (server_name, key_path) not in _whitespace_warned: diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index 74573d046b..5209782883 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -90,7 +90,6 @@ def _validate_remote_mcp_url(server_name: str, url: Any) -> str: stdio servers use ``command`` — or empty host).""" def _bad(detail: str) -> InvalidMcpUrlError: return InvalidMcpUrlError(f"Invalid MCP URL for '{server_name}': {detail}") - if not isinstance(url, str): raise _bad(f"expected a string, got {type(url).__name__}") stripped = url.strip() @@ -126,7 +125,6 @@ def _resolve_client_cert(server_name: str, config: dict): if not os.path.isfile(expanded): raise FileNotFoundError(f"{prefix}{label} not found at {expanded!r}") return expanded - if not isinstance(raw_cert, (list, tuple)): cert_path = _expand(raw_cert, "client_cert") return (cert_path, _expand(raw_key, "client_key")) if raw_key is not None else cert_path # combined PEM @@ -153,7 +151,6 @@ def _resolve_identity_header(server_name: str, config: dict): def _ignore(detail: str, *args): logger.warning("MCP server '%s': identity_header " + detail + " — ignoring", server_name, *args) return None - if not isinstance(raw, dict): return _ignore("must be a mapping with 'name' and 'value'/'value_from' keys (got %s)", type(raw).__name__) name = raw.get("name") @@ -202,7 +199,6 @@ def _make_redirect_header_stripper(original_url, *, strict: bool = False, for _name in configured_header_names if strict else (): while _name in headers: del headers[_name] - return _strip_on_cross_origin_redirect @@ -228,7 +224,6 @@ def _format_connect_error(exc: BaseException) -> str: text = "" if getattr(current, "exceptions", None) else str(current).strip() messages = ([text] if text else []) + [m for child in _exc_children(current) for m in _flatten_messages(child)] return messages or [current.__class__.__name__] - missing = _find_missing(exc) if not missing: return _sanitize_error("; ".join(list(dict.fromkeys(_flatten_messages(exc)))[:3])) diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 51ee49393a..22b02df029 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -33,8 +33,6 @@ _STDIO_DIED_AGAIN_MSG = ( "cleanly — do NOT retry this tool; ask the user to check the server's command and its stderr log.") -# --------------------------------------------------------------- pre-call gates - def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]: """Approval gate for write-capable tools on ``trust: untrusted`` servers. None to proceed, else a ``tool_error``. Fail-closed: approval-system errors block.""" @@ -89,8 +87,6 @@ def _acquire_call_server(server_name: str, tool_timeout: float): return None, not_connected -# ------------------------------------------------------------ breaker bookkeeping - def _result_is_error(result) -> bool: """True only for a JSON payload carrying an ``error`` key (non-JSON = success).""" try: @@ -138,8 +134,6 @@ def _retry_once(server_name: str, retry_call, op_description: str, what: str): return result -# --------------------------------------------------------------- recovery ladder - def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_call, op_description: str): """OAuth recovery + one retry; None when *exc* is not an auth error. ``handle_401`` decides viability; if viable, signal a reconnect (fresh credentials), wait ready, retry once. Any @@ -247,8 +241,6 @@ def _dispatch(server_name: str, server: Any, op: str, call, tool_timeout: float, recoverers, on_final_failure, record_outcome=record_outcome) -# ------------------------------------------------------------- the RPC itself - @asynccontextmanager async def _track_inflight_rpc(server: Any, server_name: str, op: str): """Register the running RPC so teardown can fail it fast. A deliberate teardown @@ -298,14 +290,6 @@ async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str await asyncio.gather(rpc_task, watch_task, return_exceptions=True) -# ---------------------------------------------------------- result rendering - -def _error_result_text(result) -> str: - """Concatenated text of an ``isError`` result's blocks (EmbeddedResource payloads: ``.resource.text``).""" - texts = (getattr(b, "text", None) or getattr(getattr(b, "resource", None), "text", None) for b in (result.content or [])) - return "".join(str(t) for t in texts if t) - - def _render_content_blocks(result, server_name: str) -> str: """Text passes through; image/audio blocks are cached (MEDIA: tags); resource blocks are materialized.""" parts: List[str] = [] @@ -325,23 +309,22 @@ def _render_content_blocks(result, server_name: str) -> str: return _truncate_mcp_text_result("\n".join(parts)) # hard-cap pathological payloads; spillover handles the rest -def _capped_structured_content(result): - """``structuredContent`` (or None); over the hard cap it degrades to the truncated JSON string (flood guard).""" +def _render_call_tool_result(result, server_name: str) -> str: + """Pure: ``CallToolResult`` -> handler JSON. ``content`` is primary; ``structuredContent`` supplements it (or + becomes ``result`` without text) and over the hard cap degrades to the truncated JSON string (flood guard); + ``_meta`` minus reserved keys. Error text also reads EmbeddedResource payloads (``.resource.text``).""" + if mcp_field(result, "is_error", "isError", False): + texts = (getattr(b, "text", None) or getattr(getattr(b, "resource", None), "text", None) for b in (result.content or [])) + error_text = "".join(str(t) for t in texts if t) or "MCP tool returned an error" + return tool_error(_sanitize_error(_truncate_mcp_text_result(error_text))) + text_result = _render_content_blocks(result, server_name) structured = mcp_field(result, "structured_content", "structuredContent") try: as_json = json.dumps(structured, ensure_ascii=False, default=str) if structured is not None else "" + if len(as_json) > _MCP_HARD_RESULT_CAP_CHARS: + structured = _truncate_mcp_text_result(as_json) except (TypeError, ValueError): - return structured - return _truncate_mcp_text_result(as_json) if len(as_json) > _MCP_HARD_RESULT_CAP_CHARS else structured - - -def _render_call_tool_result(result, server_name: str) -> str: - """Pure: ``CallToolResult`` -> handler JSON. ``content`` is primary; ``structuredContent`` supplements it (or - becomes ``result`` without text); ``_meta`` minus reserved keys.""" - if mcp_field(result, "is_error", "isError", False): - return tool_error(_sanitize_error(_truncate_mcp_text_result(_error_result_text(result) or "MCP tool returned an error"))) - text_result = _render_content_blocks(result, server_name) - structured = _capped_structured_content(result) + pass meta = _strip_reserved_meta_keys(mcp_field(result, "meta", "meta")) if structured is None and meta is None: return json.dumps({"result": text_result}, ensure_ascii=False) @@ -358,8 +341,6 @@ def _render_call_tool_result(result, server_name: str) -> str: return json.dumps({"result": text_result}, ensure_ascii=False) -# ------------------------------------------------------------------- handlers - def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): """Sync registry handler (``handler(args_dict, **kwargs) -> str``) calling an MCP tool via the background loop.""" op = f"tools/call {tool_name}" @@ -387,12 +368,10 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): def _on_failure(exc): _core._bump_server_error(server_name) logger.error("MCP tool %s/%s call failed: %s", server_name, tool_name, exc) - return _dispatch( server_name, server, op, _call, tool_timeout, (_handle_stdio_child_exited_and_retry, _handle_auth_error_and_retry, _handle_session_expired_and_retry), _on_failure, record_outcome=True) - return _handler @@ -412,14 +391,11 @@ def _make_utility_handler(op: str, log_label: str, rpc, render, required: Option async with server._rpc_lock: result = await rpc(server.session, args, server_name) return json.dumps(render(result, server_name), ensure_ascii=False) - return _dispatch( server_name, server, op, _call, tool_timeout, (_handle_auth_error_and_retry, _handle_session_expired_and_retry), lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc)) - return _handler - return _factory @@ -501,5 +477,4 @@ def _make_check_fn(server_name: str): server = _core._servers.get(server_name) return ((server is not None and (server.session is not None or server._is_recycled_stdio())) or server_name in _core._lazy_server_configs) - return _check diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index 14b7f1962f..f5ed602d1e 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -66,7 +66,6 @@ class MCPServerHealthMixin: await self._refresh_tools() except Exception: logger.exception("MCP server '%s': dynamic tool refresh failed", self.name) - task = asyncio.create_task(_run()) self._pending_refresh_tasks.add(task) task.add_done_callback(self._pending_refresh_tasks.discard) @@ -160,7 +159,6 @@ class MCPServerHealthMixin: back to ``list_tools`` when the server advertises tools, else the -32601 propagates.""" async def list_tools(): await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT) - if not self._ping_unsupported: try: await asyncio.wait_for(self.session.send_ping(), timeout=_KEEPALIVE_RPC_TIMEOUT) diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index 18e4c82092..17282af067 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -85,7 +85,6 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d return None # Legacy gate (no initialize_result): the ClientSession method shares the handler key. return None if hasattr(server.session, handler_key) else f"session lacks {handler_key}" - selected: List[dict] = [] for entry in _build_utility_schemas(server_name): reason = _skip_reason(entry["handler_key"]) @@ -205,20 +204,6 @@ def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_C return [c for c in unique if c.registry_name not in ambiguous and (c.registry_name, c.origin) not in shadowed] -def _log_foreign_owner(name: str, c: _Candidate, existing_toolset: str, lazy: bool) -> None: - """Diagnostics for a name already owned by another toolset (skipped to preserve the owner).""" - if lazy: - if not c.is_utility: - logger.warning("MCP server '%s' (lazy): cached tool '%s' collides with toolset '%s' — skipping", - name, c.registry_name, existing_toolset) - elif existing_toolset.startswith("mcp-"): - logger.error("MCP server '%s': %s normalizes to '%s', already owned by MCP toolset '%s' — skipping to " - "preserve the existing owner", name, c.origin, c.registry_name, existing_toolset) - else: - logger.warning("MCP server '%s': %s (→ '%s') collides with built-in tool in toolset '%s' — skipping to " - "preserve built-in", name, c.origin, c.registry_name, existing_toolset) - - def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: Callable, scope: Callable[[], Optional[str]], lazy: bool) -> List[str]: """Register candidates under toolset ``mcp-{name}``; returns the names that landed. The @@ -229,8 +214,17 @@ def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: C registered: List[str] = [] for c in candidates: existing_toolset = registry.get_toolset_for_tool(c.registry_name) - if existing_toolset and existing_toolset != toolset_name: - _log_foreign_owner(name, c, existing_toolset, lazy) + if existing_toolset and existing_toolset != toolset_name: # foreign owner: skip, preserve it + if lazy: + if not c.is_utility: + logger.warning("MCP server '%s' (lazy): cached tool '%s' collides with toolset '%s' — skipping", + name, c.registry_name, existing_toolset) + elif existing_toolset.startswith("mcp-"): + logger.error("MCP server '%s': %s normalizes to '%s', already owned by MCP toolset '%s' — skipping to " + "preserve the existing owner", name, c.origin, c.registry_name, existing_toolset) + else: + logger.warning("MCP server '%s': %s (→ '%s') collides with built-in tool in toolset '%s' — skipping to " + "preserve built-in", name, c.origin, c.registry_name, existing_toolset) continue registry.register( name=c.registry_name, toolset=toolset_name, schema=c.schema, handler=c.handler, check_fn=check_fn, diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index b523c0fb1f..081bedb1e9 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -97,7 +97,6 @@ class MCPServerTransportMixin: raise logger.info(log_fmt, self.name, exc, *log_extra) return await call(fallback) - mode = str((self._config or {}).get("protocol", "auto")).lower().strip() if mode in ("stateless", "modern", "2026-07-28"): return await attempt("discover", "initialize", lambda exc: True, @@ -242,7 +241,6 @@ class MCPServerTransportMixin: # Only judge 2xx (4xx/5xx may be an auth challenge); no content type advertised → trust the SDK. ct = _content_type_base(resp) return _is_2xx(resp) and bool(ct) and ct not in self._MCP_CONTENT_TYPES - probe_headers = dict(headers) if headers else {} try: async with _httpx.AsyncClient(verify=ssl_verify, follow_redirects=True, timeout=_httpx.Timeout(timeout), @@ -346,7 +344,6 @@ class MCPServerTransportMixin: async with httpx.AsyncClient(**client_kwargs) as http_client: async with _core.streamable_http_client(url, http_client=http_client) as streams: yield streams - return _owned_client_streams() async def _run_http(self, config: dict): From 691e2f7d9f986087d2c7ad16fc59d6a6fc12b774 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:36:04 -0700 Subject: [PATCH 13/16] =?UTF-8?q?refactor(tools):=20send=5Fmessage=20wave-?= =?UTF-8?q?2=20cut=20=E2=80=94=20table-driven=20telegram=20media,=20shared?= =?UTF-8?q?=20platform-module=20guard,=20flattened=20slack/target=20resolu?= =?UTF-8?q?tion,=20compact=20docs?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/react_to_message_tool.py | 60 ++---- tools/send_message_senders.py | 381 +++++++++++++-------------------- tools/send_message_targets.py | 152 +++++-------- tools/send_message_tool.py | 260 ++++++++-------------- 4 files changed, 321 insertions(+), 532 deletions(-) diff --git a/tools/react_to_message_tool.py b/tools/react_to_message_tool.py index 9b28c805cd..f8056fccae 100644 --- a/tools/react_to_message_tool.py +++ b/tools/react_to_message_tool.py @@ -4,6 +4,7 @@ costs nothing elsewhere (adapters expose reactions via ``send_message(action="react")``); defaults to the triggering message and emits ``message.reaction`` for live painting.""" +import contextlib import json from gateway.session_context import get_session_env @@ -20,56 +21,41 @@ def _open_session_db(): return None -def _react(emoji: str, message_row_id, messages_back, *, db, session_key: str) -> str: - row_id = message_row_id - target_role = "user" - if row_id is None: - # Default: the latest user message; `messages_back` steps to earlier user turns - # (ids aren't visible to the model; "two messages ago" is how a person thinks). - back = max(0, int(messages_back or 0)) - row_id = db.latest_message_row_id(session_key, role="user", offset=back) - if row_id is None: - return tool_error( - f"No user message found {back} back." if back else "No user message to react to yet.") - else: - target_role = db.get_message_role(session_key, int(row_id)) or "user" - - try: - reactions = db.set_message_reaction(session_key, int(row_id), emoji or None, author="agent") - except Exception as exc: - return tool_error(f"Failed to set the reaction: {exc}") - if reactions is None: - return tool_error(f"Message {row_id} is not part of this conversation.") - - # Paint it live; a missing bridge (non-desktop) is not an error — the reaction is - # persisted. `role` lets the renderer match a live message without a durable row id. - try: - desktop_ui.emit( - "message.reaction", {"row_id": int(row_id), "reactions": reactions, "role": target_role}) - except Exception: - pass - - return json.dumps({"success": True, "row_id": int(row_id), "reactions": reactions}, ensure_ascii=False) - - def react_to_message_tool(emoji: str, message_row_id=None, messages_back=None) -> str: """Attach (or with an empty ``emoji`` retract) the agent's reaction.""" emoji = (emoji or "").strip() session_key = get_session_env("HERMES_SESSION_KEY", "") or get_session_env("HERMES_SESSION_ID", "") if not session_key: return tool_error("No active session — reactions need a persisted conversation.") - db = _open_session_db() if db is None: return tool_error("Session storage is unavailable.") try: - return _react(emoji, message_row_id, messages_back, db=db, session_key=session_key) - finally: + row_id, target_role = message_row_id, "user" + if row_id is None: + # Default: the latest user message; `messages_back` steps to earlier user turns + # (ids aren't visible to the model; "two messages ago" is how a person thinks). + back = max(0, int(messages_back or 0)) + row_id = db.latest_message_row_id(session_key, role="user", offset=back) + if row_id is None: + return tool_error(f"No user message found {back} back." if back else "No user message to react to yet.") + else: + target_role = db.get_message_role(session_key, int(row_id)) or "user" try: + reactions = db.set_message_reaction(session_key, int(row_id), emoji or None, author="agent") + except Exception as exc: + return tool_error(f"Failed to set the reaction: {exc}") + if reactions is None: + return tool_error(f"Message {row_id} is not part of this conversation.") + # Paint it live; a missing bridge (non-desktop) is not an error — the reaction is + # persisted. `role` lets the renderer match a live message without a durable row id. + with contextlib.suppress(Exception): + desktop_ui.emit("message.reaction", {"row_id": int(row_id), "reactions": reactions, "role": target_role}) + return json.dumps({"success": True, "row_id": int(row_id), "reactions": reactions}, ensure_ascii=False) + finally: + with contextlib.suppress(Exception): from hermes_state import release_or_close release_or_close(db) - except Exception: - pass def check_react_requirements() -> bool: diff --git a/tools/send_message_senders.py b/tools/send_message_senders.py index e3721cb14b..d287f03359 100644 --- a/tools/send_message_senders.py +++ b/tools/send_message_senders.py @@ -1,6 +1,7 @@ """Standalone per-platform senders and error helpers for send_message.""" import asyncio +import contextlib import logging import os import re @@ -14,33 +15,24 @@ _IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".gif"} _VIDEO_EXTS = {".mp4", ".mov", ".avi", ".mkv", ".webm", ".3gp"} _AUDIO_EXTS = {".ogg", ".opus", ".mp3", ".m2a", ".wav", ".m4a", ".flac"} _VOICE_EXTS = {".ogg", ".opus"} -# Telegram's sendAudio only accepts MP3 / M4A; other audio goes via sendVoice (Opus/OGG) or document. -_TELEGRAM_SEND_AUDIO_EXTS = {".mp3", ".m4a"} - -# Extensions carrying a native caption on the media bubble. Voice/audio notes are excluded: -# a caption on a voice note reads as a separate label, so the text stays its own message. +_TELEGRAM_SEND_AUDIO_EXTS = {".mp3", ".m4a"} # sendAudio accepts only these; other audio -> sendVoice / document +# Captionable on the media bubble; voice/audio notes excluded (a caption there reads as a separate label). _CAPTIONABLE_EXTS = _IMAGE_EXTS | _VIDEO_EXTS | {".pdf", ".doc", ".docx", ".txt", ".md", ".csv", ".xlsx", ".zip"} - # Native caption limits (chars): Telegram caps photo/video at 1024; one conservative shared ceiling elsewhere. _TELEGRAM_CAPTION_LIMIT = 1024 _DEFAULT_CAPTION_LIMIT = 4096 def _media_caption_split(text, media_files, *, max_caption_len): - """Single chokepoint deciding whether text rides on the media bubble as its caption. - - ``(caption, "")`` only for exactly one captionable file (not a voice/audio note) - whose text fits ``max_caption_len``; otherwise ``(None, text)`` — multi-file - caption→file association is ambiguous. Length is codepoints, which never - under-counts Telegram's UTF-16 units for BMP text (over-counting fails safe); - the Telegram sender re-checks the *formatted* caption since escaping inflates it. - """ + """Single chokepoint deciding whether text rides on the media bubble as its caption: + ``(caption, "")`` only for exactly one captionable file (not a voice/audio note) whose + text fits ``max_caption_len``, else ``(None, text)`` — multi-file caption→file association + is ambiguous. Length is codepoints (never under-counts Telegram's UTF-16 units for BMP + text); the Telegram sender re-checks the *formatted* caption since escaping inflates it.""" stripped = (text or "").strip() media = media_files or [] - if not stripped or len(media) != 1 or len(stripped) > max_caption_len: - return None, text - media_path, is_voice = media[0] - if is_voice or os.path.splitext(media_path)[1].lower() not in _CAPTIONABLE_EXTS: + if (not stripped or len(media) != 1 or len(stripped) > max_caption_len or media[0][1] + or os.path.splitext(media[0][0])[1].lower() not in _CAPTIONABLE_EXTS): return None, text return stripped, "" @@ -68,20 +60,15 @@ def _success(platform: str, chat_id, warnings=None, **fields) -> dict: **({"warnings": warnings} if warnings else {})} -def _display_chat_id(platform_name: str, chat_id: str) -> str: - """Return a result-safe chat identifier for tool transcripts/log consumers.""" - return "group:***" if platform_name == "signal" and str(chat_id).startswith("group:") else chat_id - - _NO_DELIVERABLE = "No deliverable text or media remained after processing MEDIA tags" -_TELEGRAM_TRANSIENT_MARKERS = ("bad gateway", "502", "too many requests", "429", - "service unavailable", "503", "gateway timeout", "504") +_TELEGRAM_TRANSIENT_MARKERS = ("bad gateway", "502", "too many requests", "429", "service unavailable", "503", + "gateway timeout", "504") def _telegram_retry_delay(exc: Exception, attempt: int) -> float | None: - """Seconds to wait before retrying, or None when final. Honours ``retry_after``; - timeouts are never retried (the send may have gone through); 5xx/429 back off.""" + """Retry delay in seconds, or None when final: honours ``retry_after``; timeouts are + never retried (the send may have gone through); 5xx/429 back off exponentially.""" retry_after = getattr(exc, "retry_after", None) if retry_after is not None: try: @@ -114,22 +101,19 @@ def _is_telegram_thread_not_found(error: Exception) -> bool: def _telegram_bot(token): - """Bot honouring TELEGRAM_PROXY (``telegram.proxy_url``) — without it the standalone - path times out where api.telegram.org is blocked. Falls back to a direct connection.""" + """Bot honouring TELEGRAM_PROXY (standalone sends time out where api.telegram.org is + blocked); falls back to a direct connection.""" from telegram import Bot try: from gateway.platforms.base import resolve_proxy_url proxy = resolve_proxy_url("TELEGRAM_PROXY", target_hosts=["api.telegram.org"]) - except Exception: - proxy = None - if proxy: - try: - from telegram.request import HTTPXRequest - logger.info("send_message: standalone Telegram send routed through proxy %s", proxy) - return Bot(token=token, request=HTTPXRequest(proxy=proxy), - get_updates_request=HTTPXRequest(proxy=proxy)) - except Exception as proxy_err: - logger.warning("send_message: failed to attach Telegram proxy (%s), falling back to direct connection", proxy_err) + if not proxy: + return Bot(token=token) + from telegram.request import HTTPXRequest + logger.info("send_message: standalone Telegram send routed through proxy %s", proxy) + return Bot(token=token, request=HTTPXRequest(proxy=proxy), get_updates_request=HTTPXRequest(proxy=proxy)) + except Exception as proxy_err: + logger.warning("send_message: failed to attach Telegram proxy (%s), falling back to direct connection", proxy_err) return Bot(token=token) @@ -141,8 +125,7 @@ def _telegram_thread_kwargs(thread_id): try: from plugins.platforms.telegram.adapter import TelegramAdapter effective = TelegramAdapter._message_thread_id_for_send(str(thread_id)) - except Exception: - # Explicit mapping if the adapter import fails (python-telegram-bot missing). + except Exception: # adapter import failed (python-telegram-bot missing): explicit mapping effective = None if str(thread_id) == "1" else int(thread_id) return {} if effective is None else {"message_thread_id": effective} @@ -157,8 +140,8 @@ def _strip_mdv2_safe(text): def _adapter_media_method(ext, voice, force_document=False): - """Adapter media method name + kind for one file: document when forced, else image / - video / voice by extension (``voice`` already folds in the caller's audio rule).""" + """``(adapter method, kind)``: document when forced, else image / video / voice by + extension (``voice`` already folds in the caller's audio rule).""" if force_document: return "send_document", "document" if ext in _IMAGE_EXTS: @@ -169,33 +152,26 @@ def _adapter_media_method(ext, voice, force_document=False): async def _telegram_send_media(bot, chat_id, f, ext, is_voice, force_document, **kwargs): - """Bot API media method by extension: photo (unless forced document), video, voice - note, sendAudio (MP3/M4A only), else document.""" - if ext in _IMAGE_EXTS and not force_document: - return await bot.send_photo(chat_id=chat_id, photo=f, **kwargs) - if ext in _VIDEO_EXTS: - return await bot.send_video(chat_id=chat_id, video=f, **kwargs) - if ext in _VOICE_EXTS and is_voice: - return await bot.send_voice(chat_id=chat_id, voice=f, **kwargs) - if ext in _TELEGRAM_SEND_AUDIO_EXTS: - return await bot.send_audio(chat_id=chat_id, audio=f, **kwargs) - return await bot.send_document(chat_id=chat_id, document=f, **kwargs) + """Bot API media method by extension: photo (unless forced document), video, voice note, + sendAudio (MP3/M4A only), else document.""" + kind = next((k for exts, k in ((() if force_document else _IMAGE_EXTS, "photo"), (_VIDEO_EXTS, "video"), + (_VOICE_EXTS if is_voice else (), "voice"), (_TELEGRAM_SEND_AUDIO_EXTS, "audio")) + if ext in exts), "document") + return await getattr(bot, f"send_{kind}")(chat_id=chat_id, **{kind: f}, **kwargs) async def _telegram_send_text_chunk(bot, chat_id, chunk, parse_mode, has_html, text_kwargs): - """One formatted text chunk with adapter-matching fallbacks: thread-not-found -> retry - without ``message_thread_id`` (dropped from ``text_kwargs`` for later chunks too); - parse failure -> plain text.""" + """One text chunk with adapter-matching fallbacks: thread-not-found -> retry without + ``message_thread_id`` (dropped from ``text_kwargs`` for later chunks too); parse failure + -> plain text.""" async def send(text, mode): return await _send_telegram_message_with_retry(bot, chat_id=chat_id, text=text, parse_mode=mode, **text_kwargs) - try: return await send(chunk, parse_mode) except Exception as md_error: if _is_telegram_thread_not_found(md_error) and text_kwargs.get("message_thread_id") is not None: logger.warning("Thread %s not found in _send_telegram, retrying without message_thread_id", - text_kwargs.get("message_thread_id")) - text_kwargs.pop("message_thread_id", None) + text_kwargs.pop("message_thread_id")) return await send(chunk, parse_mode) err_text = str(md_error).lower() if "parse" in err_text or "markdown" in err_text or "html" in err_text: @@ -205,28 +181,22 @@ async def _telegram_send_text_chunk(bot, chat_id, chunk, parse_mode, has_html, t raise -async def _telegram_send_one_media( - bot, chat_id, media_path, is_voice, *, caption, parse_mode, has_html, thread_kwargs, force_document -): +async def _telegram_send_one_media(bot, chat_id, media_path, is_voice, *, caption, parse_mode, has_html, + thread_kwargs, force_document): """Upload one file with adapter-matching fallbacks (thread-not-found -> no - ``message_thread_id``; caption parse failure -> plain caption). Retries re-seek - the file because the first attempt consumed it.""" + ``message_thread_id``; caption parse failure -> plain caption); retries re-seek the file.""" ext = os.path.splitext(media_path)[1].lower() voice_note = ext in _VOICE_EXTS and is_voice - media_kwargs = dict(thread_kwargs) - # ``caption`` is only set for a single captionable file, so this never - # double-captions a multi-file send or a voice note. - if caption is not None and not voice_note: - media_kwargs.update(caption=caption, parse_mode=parse_mode) + # ``caption`` is only set for a single captionable file, so this never double-captions + # a multi-file send or a voice note. + media_kwargs = {**thread_kwargs, **({"caption": caption, "parse_mode": parse_mode} + if caption is not None and not voice_note else {})} if voice_note or ext in _TELEGRAM_SEND_AUDIO_EXTS: - try: + with contextlib.suppress(Exception): from plugins.platforms.telegram.adapter import _probe_voice_duration_seconds duration = await asyncio.to_thread(_probe_voice_duration_seconds, media_path) if duration is not None: media_kwargs["duration"] = duration - except Exception: - pass - with open(media_path, "rb") as f: try: return await _telegram_send_media(bot, chat_id, f, ext, is_voice, force_document, **media_kwargs) @@ -255,54 +225,42 @@ def _telegram_format(message): return message, ParseMode.HTML, True try: from plugins.platforms.telegram.adapter import TelegramAdapter - formatted = TelegramAdapter.__new__(TelegramAdapter).format_message(message) + return TelegramAdapter.__new__(TelegramAdapter).format_message(message), ParseMode.MARKDOWN_V2, False except Exception: - formatted = message # formatting unavailable: send as-is - return formatted, ParseMode.MARKDOWN_V2, False + return message, ParseMode.MARKDOWN_V2, False # formatting unavailable: send as-is async def _send_telegram(token, chat_id, message, media_files=None, thread_id=None, disable_link_previews=False, force_document=False): - """One-shot Telegram Bot API send; parse failures fall back to plain text so the - message still delivers.""" + """One-shot Telegram Bot API send; parse failures fall back to plain text.""" try: formatted, send_parse_mode, _has_html = _telegram_format(message) bot = _telegram_bot(token) from plugins.platforms.telegram.telegram_ids import normalize_telegram_chat_id + from gateway.platforms.base import BasePlatformAdapter, utf16_len # Telegram accepts a numeric chat_id OR an @username string; never force-int. int_chat_id = normalize_telegram_chat_id(chat_id) media_files = media_files or [] thread_kwargs = _telegram_thread_kwargs(thread_id) # disable_web_page_preview is only valid for send_message, not media sends. text_kwargs = {**thread_kwargs, **({"disable_web_page_preview": True} if disable_link_previews else {})} - last_msg, warnings = None, [] - - # MEDIA caption: a single captionable file + short text rides on the bubble as its - # *formatted* caption. Formatting can inflate a raw <1024 string past Telegram's - # cap, so re-check in UTF-16 units and fall back to a separate body. - _tg_caption = None - from gateway.platforms.base import BasePlatformAdapter, utf16_len + last_msg, warnings, _tg_caption = None, [], None + # MEDIA caption rides on the bubble as its *formatted* caption; formatting can inflate a + # raw <1024 string past Telegram's cap, so re-check in UTF-16 units. _cap, _ = _media_caption_split(message, media_files, max_caption_len=_TELEGRAM_CAPTION_LIMIT) if _cap is not None and utf16_len(formatted) <= _TELEGRAM_CAPTION_LIMIT: _tg_caption, formatted = formatted, "" # suppress the separate text send below - - if formatted.strip(): - # Chunk *after* formatting, in UTF-16 units: MarkdownV2/HTML escaping inflates - # text, so a raw-<4096 message can exceed the limit once formatted. - for chunk in BasePlatformAdapter.truncate_message(formatted, 4096, len_fn=utf16_len): - last_msg = await _telegram_send_text_chunk( - bot, int_chat_id, chunk, send_parse_mode, _has_html, text_kwargs) - + # Chunk *after* formatting, in UTF-16 units: escaping can push a raw-<4096 message over. + for chunk in BasePlatformAdapter.truncate_message(formatted, 4096, len_fn=utf16_len) if formatted.strip() else (): + last_msg = await _telegram_send_text_chunk(bot, int_chat_id, chunk, send_parse_mode, _has_html, text_kwargs) for media_path, is_voice in media_files: if not os.path.exists(media_path): warnings.append(f"Media file not found, skipping: {media_path}") logger.warning(warnings[-1]) - # Caption mode suppressed the text send; if the file it was meant to - # caption is gone, deliver the words on their own. + # Caption mode suppressed the text send; the file is gone, so deliver the words alone. if _tg_caption is not None and last_msg is None: try: last_msg = await _send_telegram_message_with_retry( - bot, chat_id=int_chat_id, text=_tg_caption, - parse_mode=send_parse_mode, **text_kwargs) + bot, chat_id=int_chat_id, text=_tg_caption, parse_mode=send_parse_mode, **text_kwargs) _tg_caption = None # delivered — don't re-caption a later file except Exception as _cap_err: logger.warning("Telegram caption-fallback send failed for missing media: %s", @@ -310,13 +268,11 @@ async def _send_telegram(token, chat_id, message, media_files=None, thread_id=No continue try: last_msg = await _telegram_send_one_media( - bot, int_chat_id, media_path, is_voice, - caption=_tg_caption, parse_mode=send_parse_mode, has_html=_has_html, - thread_kwargs=thread_kwargs, force_document=force_document) + bot, int_chat_id, media_path, is_voice, caption=_tg_caption, parse_mode=send_parse_mode, + has_html=_has_html, thread_kwargs=thread_kwargs, force_document=force_document) except Exception as e: warnings.append(_sanitize_error_text(f"Failed to send media {media_path}: {e}")) logger.error(warnings[-1]) - if last_msg is None: return {"error": _NO_DELIVERABLE, **({"warnings": warnings} if warnings else {})} return _success("telegram", chat_id, warnings, message_id=str(last_msg.message_id)) @@ -326,20 +282,15 @@ async def _send_telegram(token, chat_id, message, media_files=None, thread_id=No return _error(f"Telegram send failed: {e}") -def _live_runner(): - """Return the in-process gateway runner, or None (standalone/cron).""" +def _live_adapter(platform, *, lookup_failed_warning=None): + """``(runner, adapter)`` for the in-process gateway; ``(None, None)`` standalone (cron); + ``(runner, None)`` when the lookup fails — logged when a warning is given, never silently + swallowed (a silent fall-through could recreate a reconnect storm).""" try: from gateway.run import _gateway_runner_ref - return _gateway_runner_ref() + runner = _gateway_runner_ref() except Exception: - return None - - -def _live_adapter(platform, *, lookup_failed_warning=None): - """Return ``(runner, adapter)`` for the running gateway, or ``(runner, None)``. - A runner whose adapter lookup raises is logged when a warning is given, never - silently swallowed (a silent fall-through could recreate a reconnect storm).""" - runner = _live_runner() + runner = None if runner is None: return None, None try: @@ -366,16 +317,13 @@ def _plugin_standalone_sender(platform_name, *, label=None, discover=True): async def _registry_standalone_send(platform_name, pconfig, chat_id, message, thread_id=None): """One-shot text send through a plugin's ``standalone_sender_fn``.""" sender, err = _plugin_standalone_sender(platform_name) - if err: - return err - return await sender(pconfig, chat_id, message, thread_id=thread_id) + return err or await sender(pconfig, chat_id, message, thread_id=thread_id) async def _resolve_slack_user_target(token, chat_id): - """Resolve ``user:U...`` / ``user_name:`` to a D... DM conversation - (chat.postMessage needs a conversation ID). ``user_name:`` maps to a user id via - users.list first (stable handle match only); other ids pass through unchanged. - Returns ``(chat_id, None)`` or ``(None, error_dict)``.""" + """Resolve ``user:U...`` / ``user_name:`` to a D... DM conversation (chat.postMessage + needs a conversation ID); ``user_name:`` goes through users.list first (stable handle match + only); other ids pass through. ``(chat_id, None)`` or ``(None, error_dict)``.""" if not (chat_id.startswith("user:") or chat_id.startswith("user_name:")): return chat_id, None try: @@ -386,65 +334,52 @@ async def _resolve_slack_user_target(token, chat_id): from gateway.platforms.base import resolve_proxy_url, proxy_kwargs_for_aiohttp _sess_kw, _req_kw = proxy_kwargs_for_aiohttp(resolve_proxy_url()) headers = {"Authorization": f"Bearer {token}", "Content-Type": "application/json"} - - async def post_api(session, method, payload): - async with session.post( - f"https://slack.com/api/{method}", headers=headers, json=payload, **_req_kw) as resp: - return await resp.json() - - async def resolve_user_name(session, name): - query = name.strip().lstrip("@").lower() - matches, cursor = [], None - for _page in range(20): - payload = {"limit": 200, **({"cursor": cursor} if cursor else {})} - data = await post_api(session, "users.list", payload) - if not data.get("ok"): - return None, f"Slack users.list error: {data.get('error', 'unknown')}" - # Stable handle only: display/real names are mutable and non-unique. - matches += [m for m in data.get("members", []) - if not (m.get("deleted") or m.get("is_bot")) - and str(m.get("name", "")).strip().lower() == query] - cursor = (data.get("response_metadata") or {}).get("next_cursor") - if not cursor: - break - if not matches: - return None, f"Could not resolve Slack user '@{name}'." - if len(matches) > 1: - return None, f"Slack user '@{name}' matched multiple Slack users. Use a Slack user ID instead." - return matches[0].get("id"), None - async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=30), **_sess_kw) as session: - if chat_id.startswith("user_name:"): - user_id, error = await resolve_user_name(session, chat_id[len("user_name:"):]) - if error: - return None, _error(error) - chat_id = f"user:{user_id}" + async def post_api(method, payload): + async with session.post(f"https://slack.com/api/{method}", headers=headers, json=payload, + **_req_kw) as resp: + return await resp.json() - user_id = chat_id[len("user:"):] - opened = await post_api(session, "conversations.open", {"users": user_id}) + if chat_id.startswith("user_name:"): + name = chat_id[len("user_name:"):] + query = name.strip().lstrip("@").lower() + matches, cursor = [], None + for _page in range(20): + data = await post_api("users.list", {"limit": 200, **({"cursor": cursor} if cursor else {})}) + if not data.get("ok"): + return None, _error(f"Slack users.list error: {data.get('error', 'unknown')}") + # Stable handle only: display/real names are mutable and non-unique. + matches += [m for m in data.get("members", []) if not (m.get("deleted") or m.get("is_bot")) + and str(m.get("name", "")).strip().lower() == query] + cursor = (data.get("response_metadata") or {}).get("next_cursor") + if not cursor: + break + if not matches: + return None, _error(f"Could not resolve Slack user '@{name}'.") + if len(matches) > 1: + return None, _error(f"Slack user '@{name}' matched multiple Slack users. Use a Slack user ID instead.") + chat_id = f"user:{matches[0].get('id')}" + opened = await post_api("conversations.open", {"users": chat_id[len("user:"):]}) if not opened.get("ok"): return None, _error(f"Slack conversations.open error: {opened.get('error', 'unknown')}. " "Check bot permissions (im:write).") dm_id = (opened.get("channel") or {}).get("id") - if not dm_id: - return None, _error("Slack conversations.open did not return a DM channel ID") - return dm_id, None + return (dm_id, None) if dm_id else (None, _error("Slack conversations.open did not return a DM channel ID")) except Exception as e: return None, _error(f"Slack DM resolution failed: {e}") async def _signal_send_batch(post, scheduler, rl, idx, n_batches, att_batch, batch_message): - """One Signal batch under the scheduler with rate-limit retries. None on success, - False when retries were exhausted (batch lost), error dict for a non-rate-limit RPC error.""" + """One Signal batch under the scheduler with rate-limit retries: None on success, False when + retries were exhausted (batch lost), error dict for a non-rate-limit RPC error.""" n, max_attempts = len(att_batch), rl.SIGNAL_RATE_LIMIT_MAX_ATTEMPTS for attempt in range(1, max_attempts + 1): try: await scheduler.acquire(n) _rpc_t0 = time.monotonic() data = await post(att_batch, batch_message) - _rpc_duration = time.monotonic() - _rpc_t0 if "error" not in data: - await scheduler.report_rpc_duration(_rpc_duration, n) + await scheduler.report_rpc_duration(time.monotonic() - _rpc_t0, n) return None err = data["error"] if not rl._is_signal_rate_limit_error(err): @@ -453,12 +388,10 @@ async def _signal_send_batch(post, scheduler, rl, idx, n_batches, att_batch, bat scheduler.feedback(server_retry_after, n) retry_after_label = f"{server_retry_after:.0f}s" if server_retry_after else "unknown" if attempt >= max_attempts: - logger.error("Signal: rate-limit retries exhausted on batch %d/%d " - "(%d attachments lost, server retry_after=%s)", - idx + 1, n_batches, n, retry_after_label) + logger.error("Signal: rate-limit retries exhausted on batch %d/%d (%d attachments lost, " + "server retry_after=%s)", idx + 1, n_batches, n, retry_after_label) return False - logger.warning("Signal: rate-limited on batch %d/%d " - "(attempt %d/%d, server retry_after=%s); " + logger.warning("Signal: rate-limited on batch %d/%d (attempt %d/%d, server retry_after=%s); " "scheduler will pace the retry", idx + 1, n_batches, attempt, max_attempts, retry_after_label) except Exception as e: @@ -471,34 +404,29 @@ async def _signal_send_batch(post, scheduler, rl, idx, n_batches, att_batch, bat async def _send_signal(extra, chat_id, message, media_files=None): - """signal-cli JSON-RPC send. Attachments go in SIGNAL_MAX_ATTACHMENTS_PER_MSG batches - metered by the process-wide SignalAttachmentScheduler — the same bucket the gateway - adapter uses, so tool sends and inbound replies share rate-limit state.""" + """signal-cli JSON-RPC send; attachments go in SIGNAL_MAX_ATTACHMENTS_PER_MSG batches metered + by the process-wide SignalAttachmentScheduler (shared with the gateway adapter's rate-limit state).""" try: import httpx except ImportError: return {"error": "httpx not installed"} - from gateway.platforms import signal_rate_limit as rl from gateway.platforms.signal_format import markdown_to_signal try: - http_url = extra.get("http_url", "http://127.0.0.1:8080").rstrip("/") - account = extra.get("account", "") + http_url, account = extra.get("http_url", "http://127.0.0.1:8080").rstrip("/"), extra.get("account", "") if not account: return {"error": "Signal account not configured"} - valid_media = media_files or [] - attachment_paths = [path for path, _is_voice in valid_media if os.path.exists(path)] + attachment_paths = [] for media_path, _is_voice in valid_media: - if not os.path.exists(media_path): + if os.path.exists(media_path): + attachment_paths.append(media_path) + else: logger.warning("Signal media file not found, skipping: %s", media_path) - # No attachments still means one (text-only) batch; with attachments - # the text rides on batch #0 so it isn't repeated per batch. + # No attachments still means one (text-only) batch; text rides on batch #0 only. per_batch = rl.SIGNAL_MAX_ATTACHMENTS_PER_MSG - att_batches = [attachment_paths[i:i + per_batch] - for i in range(0, len(attachment_paths), per_batch)] or [[]] - n_batches = len(att_batches) - plain_text, text_styles = markdown_to_signal(message) + att_batches = [attachment_paths[i:i + per_batch] for i in range(0, len(attachment_paths), per_batch)] or [[]] + n_batches, (plain_text, text_styles) = len(att_batches), markdown_to_signal(message) recipient = {"groupId": chat_id[6:]} if chat_id.startswith("group:") else {"recipient": [chat_id]} async def _rpc_send(text, *, id_prefix, timeout, attachments=None, styled=False): @@ -518,20 +446,17 @@ async def _send_signal(extra, chat_id, message, media_files=None): timeout=rl._signal_send_timeout(len(batch_attachments))) resp.raise_for_status() return resp.json() - scheduler = rl.get_scheduler() logger.info("send_message Signal: scheduler state=%s, %d attachment(s) in %d batch(es)", scheduler.state(), len(attachment_paths), n_batches) failed_batches: list[int] = [] for idx, att_batch in enumerate(att_batches): n = len(att_batch) - estimated = scheduler.estimate_wait(n) if n > 0 else 0.0 - if n > 0 and estimated >= rl.SIGNAL_BATCH_PACING_NOTICE_THRESHOLD: + if n > 0 and (estimated := scheduler.estimate_wait(n)) >= rl.SIGNAL_BATCH_PACING_NOTICE_THRESHOLD: # Best-effort one-shot RPC for a user-facing pacing notice. - notice = (f"(More images coming — pausing ~{rl._format_wait(estimated)} " - f"for Signal rate limit, batch {idx + 1}/{n_batches}.)") try: - await _rpc_send(notice, id_prefix="notice", timeout=30.0) + await _rpc_send(f"(More images coming — pausing ~{rl._format_wait(estimated)} " + f"for Signal rate limit, batch {idx + 1}/{n_batches}.)", id_prefix="notice", timeout=30.0) except Exception as _e: logger.warning("Signal: inline notice failed: %s", _e) outcome = await _signal_send_batch(_post, scheduler, rl, idx, n_batches, att_batch, @@ -540,7 +465,6 @@ async def _send_signal(extra, chat_id, message, media_files=None): failed_batches.append(idx + 1) elif outcome is not None: return outcome - warnings = [] if len(attachment_paths) < len(valid_media): warnings.append("Some media files were skipped (not found on disk)") @@ -549,16 +473,16 @@ async def _send_signal(extra, chat_id, message, media_files=None): f"(#{', #'.join(str(b) for b in failed_batches)})") if failed_batches and len(failed_batches) == n_batches: return _error(f"Signal: every batch ({n_batches}) hit rate limit; no attachments delivered") - return _success("signal", _display_chat_id("signal", chat_id), warnings) + # Result-safe chat identifier for tool transcripts/log consumers. + return _success("signal", "group:***" if str(chat_id).startswith("group:") else chat_id, warnings) except Exception as e: return _error(f"Signal send failed: {e}") async def _send_matrix_via_adapter(pconfig, chat_id, message, media_files=None, thread_id=None): - """Matrix adapter send (native media preserved). Prefer the live gateway adapter's - persistent olm/megolm session: ephemeral per-send connects re-init E2EE and claim - one-time keys, which under bursts exhausts recipient OTKs and silently drops - messages — so the ephemeral connect/disconnect path is only for standalone/cron.""" + """Matrix adapter send (native media preserved). Prefer the live gateway adapter's persistent + olm/megolm session: ephemeral per-send connects re-init E2EE and claim one-time keys, which + under bursts exhausts recipient OTKs and silently drops messages — ephemeral is cron-only.""" media_files = media_files or [] metadata = {"thread_id": thread_id} if thread_id else None from gateway.config import Platform @@ -566,16 +490,12 @@ async def _send_matrix_via_adapter(pconfig, chat_id, message, media_files=None, "Matrix: live gateway adapter lookup failed; falling back to an " "ephemeral connect (may re-init E2EE per send)")) if live_adapter is not None: - # Owned by the gateway — must NOT be disconnected; return before the - # ephemeral adapter (and its ``finally`` disconnect) exists. + # Owned by the gateway — must NOT be disconnected (return before the ephemeral ``finally``). return await _matrix_send_core(live_adapter, chat_id, message, media_files, metadata) - - # --- Fallback: ephemeral adapter (standalone / cron context) --- try: from plugins.platforms.matrix.adapter import MatrixAdapter except ImportError: return {"error": "Matrix dependencies not installed. Run: pip install 'mautrix[encryption]'"} - adapter = MatrixAdapter(pconfig) try: if not await adapter.connect(): @@ -584,10 +504,8 @@ async def _send_matrix_via_adapter(pconfig, chat_id, message, media_files=None, except Exception as e: return _error(f"Matrix send failed: {e}") finally: - try: + with contextlib.suppress(Exception): await adapter.disconnect() - except Exception: - pass async def _matrix_send_core(adapter, chat_id, message, media_files, metadata): @@ -597,58 +515,58 @@ async def _matrix_send_core(adapter, chat_id, message, media_files, metadata): last_result = await adapter.send(chat_id, message, metadata=metadata) if not last_result.success: return _error(f"Matrix send failed: {last_result.error}") - for media_path, is_voice in media_files: if not os.path.exists(media_path): return _error(f"Media file not found: {media_path}") - ext = os.path.splitext(media_path)[1].lower() method, _ = _adapter_media_method(ext, (ext in _VOICE_EXTS and is_voice) or ext in _AUDIO_EXTS) last_result = await getattr(adapter, method)(chat_id, media_path, metadata=metadata) if not last_result.success: return _error(f"Matrix media send failed: {last_result.error}") + return {"error": _NO_DELIVERABLE} if last_result is None else _success("matrix", chat_id, message_id=last_result.message_id) - return {"error": _NO_DELIVERABLE} if last_result is None else _success( - "matrix", chat_id, message_id=last_result.message_id) + +def _gateway_platform_module(name, *, unavailable, unmet): + """``(gateway.platforms., None)`` once its ``check__requirements`` passes, else ``(None, error)``.""" + import importlib + try: + module = importlib.import_module(f"gateway.platforms.{name}") + except ImportError: + return None, {"error": unavailable} + return (module, None) if getattr(module, f"check_{name}_requirements")() else (None, {"error": unmet}) async def _send_weixin(pconfig, chat_id, message, media_files=None): """Send via Weixin iLink using the native adapter helper.""" + wx, err = _gateway_platform_module("weixin", unavailable="Weixin adapter not available.", + unmet="Weixin requirements not met. Need aiohttp + cryptography.") + if err: + return err try: - from gateway.platforms.weixin import check_weixin_requirements, send_weixin_direct - if not check_weixin_requirements(): - return {"error": "Weixin requirements not met. Need aiohttp + cryptography."} - except ImportError: - return {"error": "Weixin adapter not available."} - - try: - return await send_weixin_direct(extra=pconfig.extra, token=pconfig.token, chat_id=chat_id, - message=message, media_files=media_files) + return await wx.send_weixin_direct(extra=pconfig.extra, token=pconfig.token, chat_id=chat_id, + message=message, media_files=media_files) except Exception as e: return _error(f"Weixin send failed: {e}") async def _send_bluebubbles(extra, chat_id, message): """Send via BlueBubbles iMessage server using the adapter's REST API.""" - try: - from gateway.platforms.bluebubbles import BlueBubblesAdapter, check_bluebubbles_requirements - if not check_bluebubbles_requirements(): - return {"error": "BlueBubbles requirements not met (need aiohttp + httpx)."} - except ImportError: - return {"error": "BlueBubbles adapter not available."} - + bb, err = _gateway_platform_module("bluebubbles", unavailable="BlueBubbles adapter not available.", + unmet="BlueBubbles requirements not met (need aiohttp + httpx).") + if err: + return err try: from gateway.config import PlatformConfig - adapter = BlueBubblesAdapter(PlatformConfig(extra=extra)) + adapter = bb.BlueBubblesAdapter(PlatformConfig(extra=extra)) if not await adapter.connect(): return _error("BlueBubbles: failed to connect to server") try: result = await adapter.send(chat_id, message) - if not result.success: - return _error(f"BlueBubbles send failed: {result.error}") - return _success("bluebubbles", chat_id, message_id=result.message_id) finally: await adapter.disconnect() + if not result.success: + return _error(f"BlueBubbles send failed: {result.error}") + return _success("bluebubbles", chat_id, message_id=result.message_id) except Exception as e: return _error(f"BlueBubbles send failed: {e}") @@ -667,7 +585,6 @@ async def _send_qqbot(pconfig, chat_id, message): secret = pconfig.token or extra.get("client_secret") or _getenv("QQ_CLIENT_SECRET", "") if not appid or not secret: return _error("QQBot: QQ_APP_ID / QQ_CLIENT_SECRET not configured.") - try: async with httpx.AsyncClient(timeout=15) as client: token_resp = await client.post("https://bots.qq.com/app/getAppAccessToken", @@ -681,10 +598,9 @@ async def _send_qqbot(pconfig, chat_id, message): # Separate endpoints for guild channels, C2C (private) and groups; first 2xx wins. headers = {"Authorization": f"QQBot {access_token}", "Content-Type": "application/json"} payload = {"content": message[:4000], "msg_type": 0} - endpoints = ( - ("channel", f"https://api.sgroup.qq.com/channels/{chat_id}/messages"), - ("c2c", f"https://api.sgroup.qq.com/v2/users/{chat_id}/messages"), - ("group", f"https://api.sgroup.qq.com/v2/groups/{chat_id}/messages")) + endpoints = (("channel", f"https://api.sgroup.qq.com/channels/{chat_id}/messages"), + ("c2c", f"https://api.sgroup.qq.com/v2/users/{chat_id}/messages"), + ("group", f"https://api.sgroup.qq.com/v2/groups/{chat_id}/messages")) statuses = [] for kind, url in endpoints: resp = await client.post(url, json=payload, headers=headers) @@ -697,17 +613,14 @@ async def _send_qqbot(pconfig, chat_id, message): async def _send_yuanbao(chat_id, message, media_files=None): - """Send via the running Yuanbao adapter's persistent WebSocket (no throwaway client - possible). chat_id: ``group:``, ``direct:`` or ````.""" + """Send via the running Yuanbao adapter's persistent WebSocket (no throwaway client possible).""" try: from gateway.platforms.yuanbao import get_active_adapter, send_yuanbao_direct except ImportError: return _error("Yuanbao adapter module not available.") - adapter = get_active_adapter() if adapter is None: return _error("Yuanbao adapter is not running. Start the gateway with yuanbao platform enabled first.") - try: return await send_yuanbao_direct(adapter, chat_id, message, media_files=media_files) except Exception as e: diff --git a/tools/send_message_targets.py b/tools/send_message_targets.py index 4a3d1188d3..7e630b471e 100644 --- a/tools/send_message_targets.py +++ b/tools/send_message_targets.py @@ -6,10 +6,11 @@ import re logger = logging.getLogger("tools.send_message_tool") _TELEGRAM_TOPIC_TARGET_RE = re.compile(r"^\s*(-?\d+)(?::(\d+))?\s*$") +_NUMERIC_TOPIC_RE = _TELEGRAM_TOPIC_TARGET_RE # Discord snowflakes: numeric, same "[:]" shape _FEISHU_TARGET_RE = re.compile(r"^\s*((?:oc|ou|on|chat|open)_[-A-Za-z0-9]+)(?::([-A-Za-z0-9_]+))?\s*$") -# Slack conversation IDs: C (public), G (private/group), D (DM); uppercase alnum, 9+ chars. -# User IDs (U...) become ``user:U...`` and are opened as D... conversations first (posting -# straight to a U/W id fails); ``@handle`` -> ``user_name:...`` resolves via users.list. +# Slack conversation IDs: C (public), G (private/group), D (DM); uppercase alnum, 9+ chars. User IDs +# (U...) become ``user:U...`` and are opened as D... conversations first (posting straight to a U/W +# id fails); ``@handle`` -> ``user_name:...`` resolves via users.list. _SLACK_TARGET_RE = re.compile(r"^\s*([CGD][A-Z0-9]{8,})\s*$") _SLACK_USER_ID_RE = re.compile(r"^\s*(U[A-Z0-9]{8,})\s*$") _SLACK_USER_NAME_RE = re.compile(r"^\s*@([A-Za-z0-9._-]{1,80})\s*$") @@ -18,24 +19,17 @@ _SLACK_MENTION_RE = re.compile(r"^\s*<@(U[A-Z0-9]{8,})(?:\|[^>]+)?>\s*$") _SLACK_THREAD_TARGET_RE = re.compile(r"^\s*([CGD][A-Z0-9]{8,}):([^\s:]+)\s*$") _WEIXIN_TARGET_RE = re.compile(r"^\s*((?:wxid|gh|v\d+|wm|wb)_[A-Za-z0-9_-]+|[A-Za-z0-9._-]+@chatroom|filehelper)\s*$") _YUANBAO_TARGET_RE = re.compile(r"^\s*((?:group|direct):[^:]+)\s*$") -# Discord snowflake IDs are numeric, same regex pattern as Telegram topic targets. -_NUMERIC_TOPIC_RE = _TELEGRAM_TOPIC_TARGET_RE -# Platforms addressing recipients by E.164 phone number ("+1555..."): the '+' fails the -# isdigit() rule and channel-name resolution cannot resolve a raw number; keep the '+'. +# E.164 phone recipients ("+1555..."): the '+' fails the isdigit() rule and the channel directory +# cannot resolve a raw number, so keep the '+' and treat it as explicit. _PHONE_PLATFORMS = frozenset({"photon", "signal", "sms", "whatsapp"}) _E164_TARGET_RE = re.compile(r"^\s*\+(\d{7,15})\s*$") -# Photon DM chat GUID (mirrors _DM_CHAT_GUID_RE in the photon adapter). -_PHOTON_DM_GUID_RE = re.compile(r"^any;-;\+\d{6,}$") -# WhatsApp JIDs (@g.us groups, @s.whatsapp.net users, @lid, broadcast/newsletter): native -# targets the bridge accepts verbatim — never home-channel. -_WHATSAPP_JID_RE = re.compile( - r"^\s*[\w-]+@(?:g\.us|s\.whatsapp\.net|lid|broadcast|newsletter)\s*$", re.IGNORECASE) -# Buzz channels/DMs are native UUIDs: explicit targets, never the home channel. -_BUZZ_UUID_RE = re.compile( - r"^\s*[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}\s*$", re.IGNORECASE) -# A valid address is an explicit email target, not a channel name to resolve. +_PHOTON_DM_GUID_RE = re.compile(r"^any;-;\+\d{6,}$") # mirrors _DM_CHAT_GUID_RE in the photon adapter +# WhatsApp JIDs (@g.us, @s.whatsapp.net, @lid, broadcast/newsletter) and Buzz UUIDs are native targets +# the adapter accepts verbatim — explicit, never home-channel. A valid email address likewise. +_WHATSAPP_JID_RE = re.compile(r"^\s*[\w-]+@(?:g\.us|s\.whatsapp\.net|lid|broadcast|newsletter)\s*$", re.IGNORECASE) +_BUZZ_UUID_RE = re.compile(r"^\s*[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}\s*$", re.IGNORECASE) _EMAIL_TARGET_RE = re.compile(r"^\s*[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Za-z]{2,}\s*$") -# Exceptions to "_HOME_CHANNEL" (email reads EMAIL_HOME_ADDRESS) for error hints. +# Exceptions to "_HOME_CHANNEL" for error hints (email reads EMAIL_HOME_ADDRESS). _HOME_CHANNEL_ENV_OVERRIDES = {"email": "EMAIL_HOME_ADDRESS"} _UNRESOLVED = object() # sentinel: stop parsing, target is NOT explicit (skip generic rules) @@ -45,17 +39,19 @@ _UNRESOLVED = object() # sentinel: stop parsing, target is NOT explicit (skip g # through to the generic rules in _parse_target_ref, or _UNRESOLVED. def _parse_regex_groups(regex, *, thread_group=True): """Explicit when ``regex`` fully matches: chat_id = group 1, thread = group 2 (or None).""" - def parse(ref): - match = regex.fullmatch(ref) - return (match.group(1), match.group(2) if thread_group else None) if match else None - return parse + return lambda ref: ((m.group(1), m.group(2) if thread_group else None) + if (m := regex.fullmatch(ref)) else None) def _parse_regex_stripped(regex): """Explicit when ``regex`` fully matches; returns the stripped ref verbatim.""" - def parse(ref): - return (ref.strip(), None) if regex.fullmatch(ref) else None - return parse + return lambda ref: (ref.strip(), None) if regex.fullmatch(ref) else None + + +def _parse_nonempty(ref): + # ntfy topics and WeCom ids (the adapter picks the send command) are explicit when non-empty. + stripped = ref.strip() + return (stripped, None) if stripped else None def _parse_telegram(ref): @@ -68,20 +64,14 @@ def _parse_telegram(ref): # (regex, chat_id template, thread comes from group 2) — thread form before bare id. -_SLACK_FORMS = ( - (_SLACK_THREAD_TARGET_RE, "{}", True), - (_SLACK_TARGET_RE, "{}", False), - (_SLACK_USER_ID_RE, "user:{}", False), - (_SLACK_MENTION_RE, "user:{}", False), - (_SLACK_USER_NAME_RE, "user_name:{}", False)) +_SLACK_FORMS = ((_SLACK_THREAD_TARGET_RE, "{}", True), (_SLACK_TARGET_RE, "{}", False), + (_SLACK_USER_ID_RE, "user:{}", False), (_SLACK_MENTION_RE, "user:{}", False), + (_SLACK_USER_NAME_RE, "user_name:{}", False)) def _parse_slack(ref): - for regex, template, has_thread in _SLACK_FORMS: - match = regex.fullmatch(ref) - if match: - return template.format(match.group(1)), (match.group(2) if has_thread else None) - return None + return next(((template.format(m.group(1)), m.group(2) if has_thread else None) + for regex, template, has_thread in _SLACK_FORMS if (m := regex.fullmatch(ref))), None) def _parse_matrix(ref): @@ -89,9 +79,7 @@ def _parse_matrix(ref): # "@user" go via the generic rule so the numeric check keeps precedence. trimmed = ref.strip() split_idx = trimmed.rfind(":$") - if split_idx > 0: - return trimmed[:split_idx], trimmed[split_idx + 1 :] - return None + return (trimmed[:split_idx], trimmed[split_idx + 1:]) if split_idx > 0 else None def _parse_yuanbao(ref): @@ -99,24 +87,16 @@ def _parse_yuanbao(ref): match = _YUANBAO_TARGET_RE.fullmatch(ref) if match: return match.group(1), None - if ref.strip().isdigit(): - return f"group:{ref.strip()}", None - return _UNRESOLVED - - -def _parse_nonempty(ref): - # ntfy topics and WeCom ids (the adapter picks the send command) are explicit when non-empty. - stripped = ref.strip() - return (stripped, None) if stripped else None + return (f"group:{ref.strip()}", None) if ref.strip().isdigit() else _UNRESOLVED def _parse_signal(ref): # "group:" is a native group target; an empty id is not explicit. stripped = ref.strip() - if stripped.startswith("group:"): - group_id = stripped[len("group:"):].strip() - return (f"group:{group_id}", None) if group_id else _UNRESOLVED - return None + if not stripped.startswith("group:"): + return None + group_id = stripped[len("group:"):].strip() + return (f"group:{group_id}", None) if group_id else _UNRESOLVED _PLATFORM_PARSERS = { @@ -161,36 +141,28 @@ def resolve_send_target( platform_name: str, target_ref: str, *, pass_unresolved_references: bool = False ) -> tuple[str | None, str | None, str | None]: """Resolve one send target the same way for every caller (model tool, CLI, cron). - - Channel-directory IDs are trusted; plugin parsers are the authority on native syntax. - By default an unresolvable target is an error the model can read and pick a listed - target instead. ``pass_unresolved_references=True`` (no model in the loop: cron, - react/unreact on native message ids) hands an unresolvable target on a built-in - platform, or a plugin platform declaring no parser, to the adapter exactly as written; - a plugin platform WITH a parser stays strict for every caller. The optional validator - has the final say over parser-normalized, directory-resolved and passed-through IDs. - """ + Channel-directory IDs are trusted; plugin parsers are the authority on native syntax. By + default an unresolvable target is an error the model can act on. ``pass_unresolved_references`` + (no model in the loop: cron, react/unreact on native ids) hands an unresolvable target on a + built-in platform, or a plugin platform without a parser, to the adapter as written; a plugin + WITH a parser stays strict. The optional validator has the final say over every returned id.""" from gateway.config import Platform from gateway.platform_registry import platform_registry entry = platform_registry.get(platform_name) - def _validate(candidate: str) -> str | None: + def _validated(chat_id, thread_id): + """``(chat_id, thread_id, None)`` when the plugin validator (if any) accepts, else an error.""" if entry is None or entry.validate_target_ref_fn is None: - return None + return chat_id, thread_id, None try: - verdict = entry.validate_target_ref_fn(candidate) + verdict = entry.validate_target_ref_fn(chat_id) except Exception: logger.debug("Plugin target validator failed for %s", platform_name, exc_info=True) - return f"Target validator failed for platform '{platform_name}'" + return None, None, f"Target validator failed for platform '{platform_name}'" if verdict is True: - return None + return chat_id, thread_id, None detail = f": {verdict}" if isinstance(verdict, str) and verdict else "" - return f"Invalid target '{target_ref}' on {platform_name}{detail}" - - def _validated(chat_id, thread_id): - error = _validate(chat_id) - return (None, None, error) if error else (chat_id, thread_id, None) - + return None, None, f"Invalid target '{target_ref}' on {platform_name}{detail}" if entry is not None and entry.parse_target_ref_fn is not None: try: parsed = entry.parse_target_ref_fn(target_ref) @@ -202,11 +174,9 @@ def resolve_send_target( or not parsed[0] or (parsed[1] is not None and not isinstance(parsed[1], str))): return None, None, f"Target parser for platform '{platform_name}' returned an invalid result" return _validated(*parsed) - parsed_chat_id, parsed_thread_id, explicit = _parse_target_ref(platform_name, target_ref) if explicit and parsed_chat_id is not None: return _validated(parsed_chat_id, parsed_thread_id) - resolution_failed = False try: from gateway.channel_directory import resolve_channel_name @@ -217,27 +187,21 @@ def resolve_send_target( if resolved: parsed_chat_id, parsed_thread_id, _ = _parse_target_ref(platform_name, resolved) return _validated(parsed_chat_id or resolved, parsed_thread_id) - is_builtin = platform_name in {member.value for member in Platform} if entry is None and not is_builtin: return None, None, f"Unknown or unregistered plugin platform: {platform_name}" - - def _pass_through_unresolved(): - """Hand the raw target to the adapter unchanged (it validates).""" - error = _validate(target_ref) - if error: - return None, None, error - logger.debug("Handing unresolved target '%s' to the %s adapter unchanged " - "(the adapter validates it)", target_ref, platform_name) - return target_ref, None, None - - if entry is not None and entry.source == "plugin" and not is_builtin: - if pass_unresolved_references and entry.parse_target_ref_fn is None: - return _pass_through_unresolved() - return (None, None, f"Could not resolve '{target_ref}' on {platform_name}. " - "The plugin parser did not recognize it and no channel-directory entry matched.") - if pass_unresolved_references: - return _pass_through_unresolved() - hint = ("Try using a numeric channel ID instead." if resolution_failed - else "Use send_message(action='list') to see available targets.") + is_plugin = entry is not None and entry.source == "plugin" and not is_builtin + if pass_unresolved_references and (not is_plugin or entry.parse_target_ref_fn is None): + # Hand the raw target to the adapter unchanged (it validates). + chat_id, thread_id, error = _validated(target_ref, None) + if not error: + logger.debug("Handing unresolved target '%s' to the %s adapter unchanged " + "(the adapter validates it)", target_ref, platform_name) + return chat_id, thread_id, error + if is_plugin: + hint = "The plugin parser did not recognize it and no channel-directory entry matched." + elif resolution_failed: + hint = "Try using a numeric channel ID instead." + else: + hint = "Use send_message(action='list') to see available targets." return None, None, f"Could not resolve '{target_ref}' on {platform_name}. {hint}" diff --git a/tools/send_message_tool.py b/tools/send_message_tool.py index d4e3c08ea8..53c87729f8 100644 --- a/tools/send_message_tool.py +++ b/tools/send_message_tool.py @@ -44,20 +44,18 @@ def send_message_tool(args, **kw): def _resolve_tool_target(target: str, *, pass_unresolved_references: bool = False): - """``(platform_name, chat_id, thread_id, error)`` for a ``platform[:ref]`` target; - ``chat_id`` is None when no ref was given (caller falls back to the home channel).""" + """``(platform_name, chat_id, thread_id, error)``; ``chat_id`` is None when no ref was given + (caller falls back to the home channel).""" platform_name, _, target_ref = target.partition(":") platform_name, target_ref = platform_name.strip().lower(), target_ref.strip() or None prepare_send_message_platforms() if not target_ref: return platform_name, None, None, None - chat_id, thread_id, resolution_error = resolve_send_target( - platform_name, target_ref, pass_unresolved_references=pass_unresolved_references) - return platform_name, chat_id, thread_id, resolution_error + return platform_name, *resolve_send_target(platform_name, target_ref, + pass_unresolved_references=pass_unresolved_references) def _handle_list(): - """Return formatted list of available messaging targets.""" try: from gateway.channel_directory import format_directory_for_display return json.dumps({"targets": format_directory_for_display()}) @@ -66,37 +64,28 @@ def _handle_list(): def _handle_react(args, remove=False): - """Attach (``remove=True``: retract) an emoji reaction via a live gateway adapter's - ``add_reaction`` / ``remove_reaction``. No standalone fallback: reacting needs the - adapter's live message-id state.""" - target = args.get("target", "") - emoji = (args.get("emoji") or "").strip() + """Attach (``remove=True``: retract) an emoji reaction via the live gateway adapter; no + standalone fallback because reacting needs the adapter's live message-id state.""" + target, emoji = args.get("target", ""), (args.get("emoji") or "").strip() message_id = (args.get("message_id") or "").strip() or None if not target or (not remove and not emoji): return tool_error("'target' is required when action='unreact'" if remove else "Both 'target' and 'emoji' are required when action='react'") - # Platform-native ids (e.g. photon GUIDs) match no parser/directory entry; the - # adapter validates them. - platform_name, chat_id, _thread_id, resolution_error = _resolve_tool_target( - target, pass_unresolved_references=True) + # Platform-native ids (e.g. photon GUIDs) match no parser/directory entry; the adapter validates. + platform_name, chat_id, _thread_id, resolution_error = _resolve_tool_target(target, pass_unresolved_references=True) if resolution_error: return tool_error(resolution_error) - platform, err = _platform_enum(platform_name) if err: return tool_error(err) if not chat_id: try: from gateway.config import load_gateway_config - home = load_gateway_config().get_home_channel(platform) + chat_id = load_gateway_config().get_home_channel(platform).chat_id except Exception: - home = None - if not home: return tool_error(f"No chat specified and no home channel set for {platform_name}. " f"Use '{platform_name}:chat_id'.") - chat_id = home.chat_id - _, adapter = _live_adapter(platform) if adapter is None: return tool_error(f"Reactions require a live {platform_name} adapter in the running " @@ -104,45 +93,34 @@ def _handle_react(args, remove=False): react_fn = getattr(adapter, "remove_reaction" if remove else "add_reaction", None) if not callable(react_fn): return tool_error(f"Platform '{platform_name}' does not support message reactions.") - - kwargs = {"chat_id": chat_id, "message_id": message_id, **({} if remove else {"emoji": emoji})} try: from model_tools import _run_async - result = _run_async(react_fn(**kwargs)) + result = _run_async(react_fn(chat_id=chat_id, message_id=message_id, **({} if remove else {"emoji": emoji}))) except Exception as e: return json.dumps(_error(f"Reaction failed: {e}")) return json.dumps(result if isinstance(result, dict) else {"success": bool(result)}) def _handle_send(args): - """Send a message to a platform target.""" - target = args.get("target", "") - message = args.get("message", "") + target, message = args.get("target", ""), args.get("message", "") if not target or not message: return tool_error("Both 'target' and 'message' are required when action='send'") - platform_name, chat_id, thread_id, resolution_error = _resolve_tool_target(target) if resolution_error: return tool_error(resolution_error) - from tools.interrupt import is_interrupted if is_interrupted(): return tool_error("Interrupted") - try: from gateway.config import load_gateway_config config = load_gateway_config() except Exception as e: return json.dumps(_error(f"Failed to load gateway config: {e}")) - platform, pconfig, entry, err = _resolve_platform_config(platform_name, config) if err: return tool_error(err) - from gateway.platforms.base import BasePlatformAdapter - - # Capture [[as_document]] before extract_media strips it: images then go through - # send_document so the original bytes survive (Telegram's sendPhoto recompresses). + # Capture [[as_document]] before extract_media strips it (images keep original bytes via send_document). force_document_attachments = "[[as_document]]" in message media_files, cleaned_message = BasePlatformAdapter.extract_media(message) media_files = BasePlatformAdapter.filter_media_delivery_paths(media_files) @@ -152,30 +130,24 @@ def _handle_send(args): chat_id, err = _home_chat_id(config, platform, platform_name) if err: return tool_error(err) - - duplicate_skip = _maybe_skip_cron_duplicate_send(platform_name, chat_id, thread_id) - if duplicate_skip: + if duplicate_skip := _maybe_skip_cron_duplicate_send(platform_name, chat_id, thread_id): return json.dumps(duplicate_skip) - if platform_name == "slack" and chat_id: chat_id, resolve_err = _slack_dm_chat_id(pconfig, chat_id) if resolve_err: return json.dumps(resolve_err) - try: from model_tools import _run_async - send_kwargs = {"thread_id": thread_id, "media_files": media_files, - "force_document": force_document_attachments} # Only custom plugin handlers receive the complete typed request. - if entry is not None and entry.send_message_handler is not None: - send_kwargs["args"] = args - result = _run_async(_send_to_platform(platform, pconfig, chat_id, cleaned_message, **send_kwargs)) + handler_args = {"args": args} if entry is not None and entry.send_message_handler is not None else {} + result = _run_async(_send_to_platform(platform, pconfig, chat_id, cleaned_message, thread_id=thread_id, + media_files=media_files, force_document=force_document_attachments, + **handler_args)) if isinstance(result, dict) and result.get("success"): if used_home_channel: result["note"] = f"Sent to {platform_name} home channel (chat_id: {chat_id})" if mirror_text and _mirror_sent_message(platform_name, chat_id, mirror_text, thread_id): result["mirrored"] = True - if isinstance(result, dict) and "error" in result: result["error"] = _sanitize_error_text(result["error"]) return json.dumps(result) @@ -193,9 +165,8 @@ def _platform_enum(platform_name): def _resolve_platform_config(platform_name, config): - """``(platform, pconfig, registry_entry, error)`` for a send. Plugin platforms must be - registered; disabled/missing platforms error, except Weixin, which may be configured - purely via .env (synthesized pconfig so cron delivery works without gateway.yaml).""" + """``(platform, pconfig, registry_entry, error)``. Plugin platforms must be registered; + disabled/missing platforms error, except Weixin, which may be configured purely via .env.""" from gateway.config import Platform from gateway.platform_registry import platform_registry entry = platform_registry.get(platform_name) @@ -204,7 +175,6 @@ def _resolve_platform_config(platform_name, config): platform, err = _platform_enum(platform_name) if err: return None, None, None, err - pconfig = config.platforms.get(platform) if not pconfig or not pconfig.enabled: pconfig = _weixin_env_pconfig() if platform_name == "weixin" else None @@ -215,8 +185,7 @@ def _resolve_platform_config(platform_name, config): def _home_chat_id(config, platform, platform_name): - """Return ``(home chat_id, None)`` or ``(None, actionable error)``. - Weixin additionally honours the WEIXIN_HOME_CHANNEL env var.""" + """``(home chat_id, None)`` or ``(None, actionable error)``; Weixin also honours WEIXIN_HOME_CHANNEL.""" home = config.get_home_channel(platform) if home: return home.chat_id, None @@ -230,9 +199,8 @@ def _home_chat_id(config, platform, platform_name): def _slack_dm_chat_id(pconfig, chat_id): - """Open Slack user targets (``user:U...`` / ``user_name:@handle`` from the parser, or a - bare U... id from session metadata / home-channel config) as DM conversations — - chat.postMessage needs a conversation ID. ``(chat_id, None)`` or ``(None, error_dict)``.""" + """Open Slack user targets (``user:``/``user_name:`` from the parser, or a bare U... id from + session metadata / home-channel config) as DM conversations. ``(chat_id, None)`` or ``(None, error_dict)``.""" dm_target = f"user:{chat_id}" if chat_id.startswith("U") and _SLACK_USER_ID_RE.fullmatch(chat_id) else chat_id if not dm_target.startswith(("user:", "user_name:")): return chat_id, None @@ -261,8 +229,7 @@ def _weixin_env_pconfig(): return None from gateway.config import PlatformConfig return PlatformConfig(enabled=True, token=wx_token, extra={ - "account_id": wx_account, - "base_url": get_secret("WEIXIN_BASE_URL", "").strip(), + "account_id": wx_account, "base_url": get_secret("WEIXIN_BASE_URL", "").strip(), "cdn_base_url": get_secret("WEIXIN_CDN_BASE_URL", "").strip()}) @@ -286,14 +253,11 @@ def _maybe_skip_cron_duplicate_send(platform_name: str, chat_id: str, thread_id: from gateway.session_context import get_session_env auto_platform = get_session_env("HERMES_CRON_AUTO_DELIVER_PLATFORM", "").strip().lower() auto_chat_id = get_session_env("HERMES_CRON_AUTO_DELIVER_CHAT_ID", "").strip() - if not auto_platform or not auto_chat_id: - return None - auto_thread_id = get_session_env("HERMES_CRON_AUTO_DELIVER_THREAD_ID", "").strip() or None - if not (auto_platform == platform_name and auto_chat_id == str(chat_id) and auto_thread_id == thread_id): + if not (auto_platform and auto_chat_id and auto_platform == platform_name and auto_chat_id == str(chat_id) + and (get_session_env("HERMES_CRON_AUTO_DELIVER_THREAD_ID", "").strip() or None) == thread_id): return None target_label = f"{platform_name}:{chat_id}" + (f":{thread_id}" if thread_id is not None else "") - return { - "success": True, "skipped": True, "reason": "cron_auto_delivery_duplicate_target", "target": target_label, + return {"success": True, "skipped": True, "reason": "cron_auto_delivery_duplicate_target", "target": target_label, "note": (f"Skipped send_message to {target_label}. This cron job will already auto-deliver " "its final response to that same target. Put the intended user-facing content in " "your final response instead, or use a different target if you want an additional message.")} @@ -305,28 +269,25 @@ def _bounded_send_error(detail, max_chars=900): return text if len(text) <= max_chars else f"{text[: max_chars - 3]}..." -async def _send_live_adapter_media( - adapter, chat_id, message, media_files, *, thread_id=None, metadata=None, force_document=False): - """Deliver text and every media descriptor through adapter media APIs. Adapters that - only inherit the BasePlatformAdapter stub for a kind are unsupported, not no-op'd.""" - caption, separate_text = _media_caption_split( - message, media_files, max_caption_len=_DEFAULT_CAPTION_LIMIT) +async def _send_live_adapter_media(adapter, chat_id, message, media_files, *, thread_id=None, metadata=None, + force_document=False): + """Deliver text and every media descriptor through adapter media APIs; adapters that only + inherit the BasePlatformAdapter stub for a kind are unsupported, not no-op'd.""" + caption, separate_text = _media_caption_split(message, media_files, max_caption_len=_DEFAULT_CAPTION_LIMIT) last_result = None if separate_text and separate_text.strip(): last_result = await adapter.send(chat_id=chat_id, content=separate_text, metadata=metadata) if not last_result.success: return {"error": f"Adapter send failed: {_bounded_send_error(last_result.error)}"} - from gateway.platforms.base import BasePlatformAdapter total = len(media_files) for index, descriptor in enumerate(media_files): media_path = descriptor[0] if isinstance(descriptor, (list, tuple)) and descriptor else None if not isinstance(media_path, str) or not media_path: return {"error": f"Adapter media send failed: invalid media descriptor {index + 1}/{total}"} - is_voice = bool(descriptor[1]) if len(descriptor) > 1 else False + is_voice = len(descriptor) > 1 and bool(descriptor[1]) if not os.path.exists(media_path): return {"error": f"Adapter media send failed: media file {index + 1}/{total} was not found"} - ext = os.path.splitext(media_path)[1].lower() method_name, media_kind = _adapter_media_method(ext, is_voice or ext in _AUDIO_EXTS, force_document) adapter_method = getattr(type(adapter), method_name, None) @@ -345,22 +306,16 @@ async def _send_live_adapter_media( continue detail = _bounded_send_error(last_result.error or "media send failed") return {"error": f"Adapter media send failed after {index}/{total} files: {detail}"} - if last_result is None: return {"error": _NO_DELIVERABLE} return {"success": True, "message_id": last_result.message_id, "media_delivered": True} async def _dispatch_on_gateway_loop(runner, make_coro, log_message): - """Await ``make_coro()`` on the gateway's loop. adapter.send() uses queues/tasks bound - to that loop; awaiting it from another loop (the tool worker thread) deadlocks, so - cross-loop calls are scheduled threadsafe onto it.""" + """Await ``make_coro()`` on the gateway's loop: adapter.send() uses queues/tasks bound to it, + so awaiting from another loop (the tool worker thread) deadlocks.""" gateway_loop = getattr(runner, "_gateway_loop", None) - try: - current_loop = asyncio.get_running_loop() - except RuntimeError: - current_loop = None - if gateway_loop is None or current_loop is gateway_loop: + if gateway_loop is None or asyncio.get_running_loop() is gateway_loop: return await make_coro() # same loop / no gateway loop (CLI, tests) if not gateway_loop.is_running(): return {"error": "Gateway loop is not running; cannot dispatch adapter send"} @@ -368,31 +323,29 @@ async def _dispatch_on_gateway_loop(runner, make_coro, log_message): fut = safe_schedule_threadsafe(make_coro(), gateway_loop, logger=logger, log_message=log_message) if fut is None: return {"error": "Gateway loop unavailable for send dispatch"} - # shield: a cancelled caller (agent interrupt) must not cancel the enqueued send or a - # retry would duplicate it. No timeout: the adapter and outer _run_async bound the wait. + # shield: a cancelled caller must not cancel the enqueued send (a retry would duplicate it). + # No timeout: the adapter and outer _run_async bound the wait. return await asyncio.shield(asyncio.wrap_future(fut)) -async def _send_via_adapter( - platform, pconfig, chat_id, chunk, *, thread_id=None, media_files=None, force_document=False): - """Live in-process gateway adapter first, else the plugin's ``standalone_sender_fn`` - (gateway not in this process, e.g. cron), else an error naming both options. Media - goes through the adapter's native media APIs under the same cross-loop rules.""" +async def _send_via_adapter(platform, pconfig, chat_id, chunk, *, thread_id=None, media_files=None, + force_document=False): + """Live in-process gateway adapter first, else the plugin's ``standalone_sender_fn`` (cron), + else an error naming both; media uses the adapter's native media APIs under the same rules.""" platform_name = platform.value if hasattr(platform, "value") else str(platform) runner, adapter = _live_adapter(platform) if adapter is not None: try: metadata = {**({"thread_id": thread_id} if thread_id else {}), **({"publish_topic": chat_id} if platform_name == "ntfy" and chat_id else {})} or None - if media_files: - return await _dispatch_on_gateway_loop( - runner, lambda: _send_live_adapter_media( - adapter, chat_id, chunk, media_files, - thread_id=thread_id, metadata=metadata, force_document=force_document), - "send_message: failed to schedule media send on gateway loop") + if media_files: # always a dict result, returned as-is below + make_coro = lambda: _send_live_adapter_media( # noqa: E731 + adapter, chat_id, chunk, media_files, thread_id=thread_id, metadata=metadata, + force_document=force_document) + else: + make_coro = lambda: adapter.send(chat_id=chat_id, content=chunk, metadata=metadata) # noqa: E731 result = await _dispatch_on_gateway_loop( - runner, lambda: adapter.send(chat_id=chat_id, content=chunk, metadata=metadata), - "send_message: failed to schedule on gateway loop") + runner, make_coro, f"send_message: failed to schedule{' media send' if media_files else ''} on gateway loop") except asyncio.CancelledError: raise except Exception as e: @@ -402,48 +355,42 @@ async def _send_via_adapter( if result.success: return {"success": True, "message_id": result.message_id} return {"error": f"Adapter send failed: {_bounded_send_error(result.error)}"} - try: from gateway.platform_registry import platform_registry - entry = platform_registry.get(platform_name) + sender = platform_registry.get(platform_name).standalone_sender_fn except Exception: - entry = None - if entry is None or entry.standalone_sender_fn is None: - return {"error": ( - f"No live adapter for platform '{platform_name}'. Is the gateway running with this platform " - f"connected? For out-of-process delivery (e.g. cron in a separate process), the platform " - f"plugin must register a standalone_sender_fn on its PlatformEntry.")} + sender = None + if sender is None: + return {"error": (f"No live adapter for platform '{platform_name}'. Is the gateway running with this platform " + f"connected? For out-of-process delivery (e.g. cron in a separate process), the platform " + f"plugin must register a standalone_sender_fn on its PlatformEntry.")} try: - result = await entry.standalone_sender_fn(pconfig, chat_id, chunk, thread_id=thread_id, - media_files=media_files, force_document=force_document) + result = await sender(pconfig, chat_id, chunk, thread_id=thread_id, media_files=media_files, + force_document=force_document) except asyncio.CancelledError: raise except Exception as e: logger.debug("Plugin standalone send for %s raised", platform_name, exc_info=True) return {"error": f"Plugin standalone send failed: {_bounded_send_error(e)}"} if isinstance(result, dict) and (result.get("success") or result.get("error")): - if result.get("error"): - return {**result, "error": _bounded_send_error(result["error"])} - return result + return {**result, "error": _bounded_send_error(result["error"])} if result.get("error") else result return {"error": (f"Plugin standalone send for '{platform_name}' returned an invalid result: " f"expected a dict with 'success' or 'error' keys, got {type(result).__name__}")} async def _send_chunks(chunks, send_one): """``send_one(chunk, is_last)`` in order; stop at the first error dict, else last result.""" - last_result = None + result = None for i, chunk in enumerate(chunks): result = await send_one(chunk, i == len(chunks) - 1) if isinstance(result, dict) and result.get("error"): - return result - last_result = result - return last_result + break + return result def _platform_max_length(platform): - """Chunking limit: the adapter constant for Signal (its raw JSON-RPC path bypasses - SignalAdapter's own chunking), the registry's ``max_message_length`` for plugins, - else None (no chunking).""" + """Chunking limit: Signal's adapter constant (its raw JSON-RPC path bypasses the adapter's + chunking), the registry's ``max_message_length`` for plugins, else None (no chunking).""" from gateway.config import Platform if platform == Platform.SIGNAL: try: @@ -459,24 +406,18 @@ def _platform_max_length(platform): return None -# Plugin platforms whose media (Discord: all) sends bypass the live adapter on purpose for -# the registry ``standalone_sender_fn``: Discord's handles forum channels/threads/multipart -# uploads; Slack uploads via files_upload_v2; WhatsApp posts to the Baileys bridge -# /send-media so media arrive as native bubbles. -# platform -> (error label, run discover_plugins first, caption-capable, -# media_files sentinel for non-final chunks, forward force_document) -_PLUGIN_STANDALONE_MEDIA = { - "discord": ("Discord", False, True, [], False), - "feishu": ("Feishu", True, False, None, False), - "slack": ("Slack", True, True, [], False), - "whatsapp": ("WhatsApp", True, True, None, True)} +# Plugin platforms whose media (Discord: all) sends deliberately bypass the live adapter for the +# registry ``standalone_sender_fn`` (Discord: forums/threads/multipart; Slack: files_upload_v2; +# WhatsApp: Baileys /send-media). platform -> (error label, run discover_plugins first, +# caption-capable, media_files sentinel for non-final chunks, forward force_document) +_PLUGIN_STANDALONE_MEDIA = {"discord": ("Discord", False, True, [], False), "feishu": ("Feishu", True, False, None, False), + "slack": ("Slack", True, True, [], False), "whatsapp": ("WhatsApp", True, True, None, True)} -async def _send_plugin_standalone( - platform_name, pconfig, chat_id, message, chunks, media_files, *, thread_id, max_len, force_document -): - """Chunked send through a plugin's standalone_sender_fn; a single captionable file + - short text rides as the media caption.""" +async def _send_plugin_standalone(platform_name, pconfig, chat_id, message, chunks, media_files, *, thread_id, + max_len, force_document): + """Chunked send through a plugin's standalone_sender_fn; one captionable file + short text + rides as the media caption.""" label, discover, captionable, empty_media, pass_force = _PLUGIN_STANDALONE_MEDIA[platform_name] sender, err = _plugin_standalone_sender(platform_name, label=label, discover=discover) if err: @@ -492,30 +433,25 @@ async def _send_plugin_standalone( pconfig, chat_id, chunk, thread_id=thread_id, media_files=media_files if is_last else empty_media, **extra)) -# Native-media chunked routes for built-in platforms; media rides on the final chunk, -# non-final chunks get the empty-media sentinel. platform -> (media required, sentinel, -# sender(platform, pconfig, chat_id, chunk, media, thread_id, force_document)). -# Matrix: ALL sends use the native adapter so text is encrypted in E2EE rooms too. -# Signal: attachments ride the JSON-RPC ``attachments`` param. -# Yuanbao / WeCom: media needs the running gateway adapter. -# Slack (text; media intercepted above): prefer the live adapter — multi-workspace aware, -# honors gates like ignored_channels — else the plugin's standalone sender. -# Names resolve at call time so tests can monkeypatch e.g. ``_send_signal``. def _via_adapter_route(p, pc, cid, chunk, media, tid, fd): return _send_via_adapter(p, pc, cid, chunk, thread_id=tid, media_files=media, force_document=fd) +# Native-media chunked routes for built-in platforms; media rides on the final chunk, non-final +# chunks get the sentinel. platform -> (media required, sentinel, sender(platform, pconfig, +# chat_id, chunk, media, thread_id, force_document)). Matrix: ALL sends use the native adapter +# (E2EE text). Signal: attachments ride the JSON-RPC param. Yuanbao / WeCom: media needs the +# running gateway. Slack text: live adapter (multi-workspace, ignored_channels gates) else the +# plugin's standalone sender. Names resolve at call time so tests can monkeypatch ``_send_signal``. _CHUNKED_ROUTES = { "matrix": (False, [], lambda p, pc, cid, chunk, media, tid, fd: _send_matrix_via_adapter( pc, cid, chunk, media_files=media, thread_id=tid)), "signal": (True, [], lambda p, pc, cid, chunk, media, tid, fd: _send_signal( pc.extra, cid, chunk, media_files=media)), - "yuanbao": (True, None, lambda p, pc, cid, chunk, media, tid, fd: _send_yuanbao( - cid, chunk, media_files=media)), + "yuanbao": (True, None, lambda p, pc, cid, chunk, media, tid, fd: _send_yuanbao(cid, chunk, media_files=media)), "slack": (False, [], _via_adapter_route), "wecom": (True, None, _via_adapter_route)} - # Text-only senders for built-in platforms (generic path; media is dropped with a # warning). Signature: (pconfig, chat_id, chunk, thread_id) -> result. _TEXT_SENDERS = { @@ -530,66 +466,56 @@ _MEDIA_PLATFORMS_NOTE = "telegram, discord, matrix, weixin, signal, yuanbao, fei async def _send_to_platform(platform, pconfig, chat_id, message, thread_id=None, media_files=None, force_document=False, args=None): - """Route a message to the platform sender, chunking long text with the adapters' smart - splitter. Branch order matters: Weixin first (its native helper must not be blocked by - unrelated optional imports such as lark-oapi), Telegram (chunks itself), plugin - standalone media routes, native chunked routes, then the generic text path.""" + """Route to the platform sender, chunking long text with the adapters' splitter. Order matters: + Weixin first (its native helper must not be blocked by unrelated optional imports such as + lark-oapi), Telegram (chunks itself), plugin standalone media, native chunked, generic text.""" from gateway.config import Platform platform_name = platform.value if hasattr(platform, "value") else str(platform) media_files = media_files or [] if platform == Platform.WEIXIN: return await _send_weixin(pconfig, chat_id, message, media_files=media_files) - # Telegram chunks internally on the *formatted* text (escaping inflates length). if platform == Platform.TELEGRAM: - disable_link_previews = bool(getattr(pconfig, "extra", {}) and pconfig.extra.get("disable_link_previews")) - return await _send_telegram(pconfig.token, chat_id, message, media_files=media_files, thread_id=thread_id, - disable_link_previews=disable_link_previews, force_document=force_document) - + return await _send_telegram( + pconfig.token, chat_id, message, media_files=media_files, thread_id=thread_id, force_document=force_document, + disable_link_previews=bool(getattr(pconfig, "extra", {}) and pconfig.extra.get("disable_link_previews"))) from gateway.platforms.base import BasePlatformAdapter max_len = _platform_max_length(platform) chunks = BasePlatformAdapter.truncate_message(message, max_len) if max_len else [message] if platform_name == "discord" or (media_files and platform_name in _PLUGIN_STANDALONE_MEDIA): return await _send_plugin_standalone(platform_name, pconfig, chat_id, message, chunks, media_files, thread_id=thread_id, max_len=max_len, force_document=force_document) - route = _CHUNKED_ROUTES.get(platform_name) if route is not None and (media_files or not route[0]): _, empty_media, sender = route return await _send_chunks(chunks, lambda chunk, is_last: sender( platform, pconfig, chat_id, chunk, media_files if is_last else empty_media, thread_id, force_document)) - # Generic path: text only. Buzz has verified native media delivery through - # _send_via_adapter (media-only sends included), so it is exempt from the warning. + # Generic path: text only. Buzz delivers media natively via _send_via_adapter, so no warning. warning = None if media_files and platform_name != "buzz": if not message.strip(): - return {"error": ( - f"send_message MEDIA delivery is currently only supported for {_MEDIA_PLATFORMS_NOTE}; " - f"target {platform_name} had only media attachments")} - warning = ( - f"MEDIA attachments were omitted for {platform_name}; " - f"native send_message media delivery is currently only supported for {_MEDIA_PLATFORMS_NOTE}") - + return {"error": (f"send_message MEDIA delivery is currently only supported for {_MEDIA_PLATFORMS_NOTE}; " + f"target {platform_name} had only media attachments")} + warning = (f"MEDIA attachments were omitted for {platform_name}; " + f"native send_message media delivery is currently only supported for {_MEDIA_PLATFORMS_NOTE}") text_sender = _TEXT_SENDERS.get(platform_name) if text_sender is not None: send_one = lambda chunk, is_last: text_sender(pconfig, chat_id, chunk, thread_id) # noqa: E731 else: from gateway.platform_registry import platform_registry entry = platform_registry.get(platform_name) - handler = entry.send_message_handler if entry is not None else None - if handler is not None: + if entry is not None and entry.send_message_handler is not None: # Custom handler receives the full typed request once (not per chunk). try: import inspect - result = handler(args or {}, chat_id, platform_name, pconfig) + result = entry.send_message_handler(args or {}, chat_id, platform_name, pconfig) return await result if inspect.isawaitable(result) else result except Exception as e: return {"error": f"Plugin send_message handler failed: {e}"} # Plugin platform: live gateway adapter if available, else standalone_sender_fn. send_one = lambda chunk, is_last: _via_adapter_route( # noqa: E731 platform, pconfig, chat_id, chunk, media_files if is_last else [], thread_id, force_document) - last_result = await _send_chunks(chunks, send_one) if (warning and isinstance(last_result, dict) and last_result.get("success") and not last_result.get("media_delivered")): From 35b3888fc538a96cb1900a3d66d7ba7b44aa7c4c Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:37:03 -0700 Subject: [PATCH 14/16] refactor(tools): MCP _dispatch absorbs _invoke_with_recovery; oauth 401 recovery flattened; small predicate folds --- tools/mcp_oauth_manager.py | 11 +++++------ tools/mcp_tool_errors.py | 4 ++-- tools/mcp_tool_handlers.py | 39 +++++++++++++++++--------------------- tools/mcp_tool_health.py | 3 +-- 4 files changed, 25 insertions(+), 32 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index c5cdda47ca..9f749de74e 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -342,22 +342,21 @@ class MCPOAuthManager: async def _recover_401(self, server_name: str, entry: _ProviderEntry, key: str, pending: asyncio.Future) -> None: """Single recovery attempt behind *pending*; always clears the dedup slot.""" - can_refresh = False try: # Disk changed (external refresh)? Else: if the SDK can refresh in place, let the caller retry. - if await self.invalidate_if_disk_changed(server_name): - can_refresh = True - else: + can_refresh = await self.invalidate_if_disk_changed(server_name) + if not can_refresh: try: can_refresh = bool(entry.provider.context.can_refresh_token()) except Exception: # no context / not callable / probe failed can_refresh = False except Exception as exc: # pragma: no cover — defensive logger.warning("MCP OAuth '%s': 401 handler failed: %s", server_name, exc) + can_refresh = False finally: - if not pending.done(): - pending.set_result(can_refresh) entry.pending_401.pop(key, None) + if not pending.done(): + pending.set_result(can_refresh) async def handle_401(self, server_name: str, failed_access_token: Optional[str] = None) -> bool: """Handle a 401 from a tool call. True: a (possibly new) token is available — reconnect and retry. False: no diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index 5209782883..eb0aa5f3c2 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -307,6 +307,6 @@ def _is_session_expired_error(exc: BaseException) -> bool: # Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids false positives. msg = str(current).lower() found = found or isinstance(current, transport_error_types) or any(m in msg for m in _SESSION_EXPIRED_MARKERS) - stack.extend(getattr(current, "exceptions", ())) - stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None))) + stack.extend((*getattr(current, "exceptions", ()), getattr(current, "__cause__", None), + getattr(current, "__context__", None))) return found diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 22b02df029..f1c519bca3 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -36,8 +36,8 @@ _STDIO_DIED_AGAIN_MSG = ( def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]: """Approval gate for write-capable tools on ``trust: untrusted`` servers. None to proceed, else a ``tool_error``. Fail-closed: approval-system errors block.""" - trust = _core._server_trust_levels.get(server_name, _core._TRUST_FULL) - if trust != _core._TRUST_UNTRUSTED or _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True: + if (_core._server_trust_levels.get(server_name, _core._TRUST_FULL) != _core._TRUST_UNTRUSTED + or _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True): return None try: # lazy: tools.approval routes the prompt to whichever surface owns the session from tools.approval import request_elicitation_consent @@ -210,13 +210,18 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry f"{type(retry_exc).__name__}: {_exc_str(retry_exc)}")) -def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: str, - recoverers, on_final_failure: Callable[[BaseException], None], - record_outcome: bool = False) -> str: - """Run ``call_once``, walking ``recoverers`` (``(server_name, exc, retry_call, op) -> - Optional[str]``, None = not its kind; order matters) on failure. Unrecovered exceptions go - through ``on_final_failure`` and become the generic call-failed error. ``record_outcome`` - applies breaker bookkeeping to the FIRST attempt only; retries own theirs.""" +def _dispatch(server_name: str, server: Any, op: str, call, tool_timeout: float, recoverers, + on_final_failure: Callable[[BaseException], None], record_outcome: bool = False) -> str: + """Mark the call started on *server* (doubles may lack ``mark_tool_call``), run coroutine function *call* + on the MCP loop and, on failure, walk ``recoverers`` (``(server_name, exc, retry_call, op) -> Optional[str]``, + None = not its kind; order matters). Unrecovered exceptions go through ``on_final_failure`` and become the + generic call-failed error. ``record_outcome`` applies breaker bookkeeping to the FIRST attempt only.""" + if callable(getattr(server, "mark_tool_call", None)): + server.mark_tool_call() + + def call_once(): + return _core._run_on_mcp_loop(call, timeout=tool_timeout) + try: result = call_once() return _record_call_outcome(server_name, result) if record_outcome else result @@ -231,16 +236,6 @@ def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: st return tool_error(_sanitize_error(f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}")) -def _dispatch(server_name: str, server: Any, op: str, call, tool_timeout: float, recoverers, - on_final_failure: Callable[[BaseException], None], record_outcome: bool = False) -> str: - """Mark the call started on *server* (doubles may lack ``mark_tool_call``), run coroutine - function *call* on the MCP loop and walk the recovery ladder (:func:`_invoke_with_recovery`).""" - if callable(getattr(server, "mark_tool_call", None)): - server.mark_tool_call() - return _invoke_with_recovery(server_name, lambda: _core._run_on_mcp_loop(call, timeout=tool_timeout), op, - recoverers, on_final_failure, record_outcome=record_outcome) - - @asynccontextmanager async def _track_inflight_rpc(server: Any, server_name: str, op: str): """Register the running RPC so teardown can fail it fast. A deliberate teardown @@ -254,8 +249,8 @@ async def _track_inflight_rpc(server: Any, server_name: str, op: str): yield except asyncio.CancelledError: if getattr(server, "_reconnecting", False): - raise RuntimeError(f"MCP {op} on '{server_name}' was aborted by a reconnect " - f"teardown; retry the request on the rebuilt session") from None + raise RuntimeError(f"MCP {op} on '{server_name}' was aborted by a reconnect teardown; retry the " + f"request on the rebuilt session") from None raise finally: if tracked: @@ -272,7 +267,7 @@ async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str raise _StdioChildExited(f"MCP stdio subprocess for '{server_name}' had already exited when the call was dispatched") _call_coro = server.session.call_tool(tool_name, arguments=args) _watch_children = getattr(server, "_watch_stdio_children", None) - if not (_watch_children is not None and inspect.iscoroutinefunction(_watch_children) and asyncio.iscoroutine(_call_coro)): + if not (inspect.iscoroutinefunction(_watch_children) and asyncio.iscoroutine(_call_coro)): # Stubbed sessions return a non-awaitable, or there is no child-watcher to race: plain await. return await _call_coro if asyncio.iscoroutine(_call_coro) else _call_coro rpc_task = asyncio.ensure_future(_call_coro) diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index f5ed602d1e..4a8af5753a 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -51,8 +51,7 @@ class MCPServerHealthMixin: return next((reason for deadline, reason in self._stdio_recycle_deadlines() if now >= deadline), None) def _next_stdio_recycle_deadline(self) -> Optional[float]: - deadlines = self._stdio_recycle_deadlines() - return min(d for d, _ in deadlines) if deadlines else None + return min((d for d, _ in self._stdio_recycle_deadlines()), default=None) def _mark_stdio_recycled(self, reason: str) -> None: """Mark a stdio session dormant before its transport finishes closing.""" From f88c3285021dbf526bb490f48b6b8ff5d6acd08c Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:41:16 -0700 Subject: [PATCH 15/16] =?UTF-8?q?refactor(tools):=20wave-2=20group=20M=20?= =?UTF-8?q?=E2=80=94=20read=5Fextract/patch=5Fparser/shell=5Fheredoc/self?= =?UTF-8?q?=5Frepo=5Fguard/schema=5Fsanitizer=20-12%=20LOC=20(68%=20code)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - patch_parser: shared _fail/_written/_unified_diff apply helpers, single PatchResult build, flattened apply loop, merged inert/no-op hunk skips, overlay _remove, dropped never-set PatchOperation.content - self_repo_guard: conditional git-mutation predicates as a dispatch table over one _has_flag, unified git global-option/alias parse, suppress() for resolve/alias/shlex, folded scope and heredoc scanners - shell_heredoc: unit scanner and operator parser as single if/elif ladders, unified quote masking - read_extract: groupby gap runs, single try for anydoc, islice/suppress xlsx loops, folded notebook/docx render loops, _zip_xml optional -> empty element - schema_sanitizer: folded rename/type-array/recovery-strip loops, _dict_nodes/_carry_union_meta - docstrings/comments compacted by hand (every WHY kept) Verified: tool-schema dump byte-identical (SCHEMA-OK); golden corpus over pure functions and every registry schema + edge cases byte-identical vs 113f046 base (GOLDEN-OK); 95 test files 2517 passed (2 pre-existing reds in test_web_tools_config also fail on base). --- tools/patch_parser.py | 292 +++++++++++++++-------------------- tools/read_extract.py | 317 +++++++++++++++----------------------- tools/schema_sanitizer.py | 238 ++++++++++++---------------- tools/self_repo_guard.py | 276 ++++++++++++++------------------- tools/shell_heredoc.py | 173 ++++++++------------- 5 files changed, 515 insertions(+), 781 deletions(-) diff --git a/tools/patch_parser.py b/tools/patch_parser.py index 942941c6a1..74ce9fe331 100644 --- a/tools/patch_parser.py +++ b/tools/patch_parser.py @@ -1,14 +1,9 @@ -#!/usr/bin/env python3 -"""V4A patch format parser and applier (format used by codex, cline, etc.). - - *** Begin Patch / *** End Patch wrap the operations: - *** Update File: p.py then hunks: ``@@ hint @@``, `` ctx``, ``-old``, ``+new`` - *** Add File: n.py then ``+`` lines; *** Delete File: o.py; *** Move File: a -> b - - operations, error = parse_v4a_patch(patch_content) - result = apply_v4a_operations(operations, file_ops) -""" +"""V4A patch parser/applier (codex, cline). ``*** Begin Patch``/``*** End Patch`` wrap ops: +``*** Update File: p`` + hunks (``@@ hint @@``, `` ctx``, ``-old``, ``+new``); ``*** Add File: n`` ++ ``+`` lines; ``*** Delete File: o``; ``*** Move File: a -> b``. Entry points: +``parse_v4a_patch(text) -> (ops, error)`` and ``apply_v4a_operations(ops, file_ops)``.""" +import contextlib import difflib import inspect import re @@ -42,7 +37,6 @@ class PatchOperation: file_path: str new_path: Optional[str] = None # MOVE only hunks: List[Hunk] = field(default_factory=list) - content: Optional[str] = None # ADD only # Markers must occupy the whole line at column 0 so content lines that merely @@ -58,12 +52,9 @@ _HINT_RE = re.compile(r'@@\s*(.+?)\s*@@') def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[str]]: - """Parse a V4A patch -> ``(operations, None)`` (``[]`` for an empty patch is not an - error) or ``([], "Parse error: ...")`` for malformed operations.""" - # Tolerate CRLF bodies: a stray ``\r`` would otherwise end up in every - # HunkLine.content and defeat the anchored Begin/End markers. + """-> ``(operations, None)`` (empty patch = ``[]``, no error) or ``([], "Parse error: …")``.""" + # Tolerate CRLF: a stray ``\r`` would land in every HunkLine.content and defeat the markers. lines = [ln[:-1] if ln.endswith('\r') else ln for ln in patch_content.split('\n')] - start_idx = -1 # parse from the top when no Begin marker is present end_idx = len(lines) for i, line in enumerate(lines): @@ -72,7 +63,6 @@ def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[ elif _END_MARKER.match(line): end_idx = i break - operations: List[PatchOperation] = [] current_op: Optional[PatchOperation] = None current_hunk: Optional[Hunk] = None @@ -87,8 +77,7 @@ def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[ operations.append(current_op) for line in lines[start_idx + 1:end_idx]: - op_match = next( - ((kind, m) for kind, rx in _OP_MARKERS if (m := rx.match(line))), None) + op_match = next(((kind, m) for kind, rx in _OP_MARKERS if (m := rx.match(line))), None) if op_match: kind, m = op_match _flush() @@ -96,8 +85,8 @@ def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[ operation=kind, file_path=m.group(1).strip(), new_path=m.group(2).strip() if kind is OperationType.MOVE else None) - # UPDATE hunks start lazily (at '@@' or the first hunk line); ADD - # collects all '+' lines into one hunk; DELETE/MOVE are complete. + # UPDATE hunks start lazily ('@@' or first hunk line); ADD collects all '+' lines + # into one hunk; DELETE/MOVE are complete. current_hunk = Hunk() if kind is OperationType.ADD else None if kind in (OperationType.DELETE, OperationType.MOVE): operations.append(current_op) @@ -115,18 +104,16 @@ def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[ elif line[0] != '\\': # "\ No newline at end of file" marker is skipped current_hunk.lines.append(HunkLine(' ', line)) # implicit context line _flush() - parse_errors: List[str] = [] for op in operations: if not op.file_path: parse_errors.append("Operation with empty file path") - if op.operation == OperationType.UPDATE and not op.hunks: + if op.operation is OperationType.UPDATE and not op.hunks: parse_errors.append(f"UPDATE {op.file_path!r}: no hunks found") - if op.operation == OperationType.MOVE and not op.new_path: - parse_errors.append(f"MOVE {op.file_path!r}: missing destination path (expected 'src -> dst')") - if parse_errors: - return [], "Parse error: " + "; ".join(parse_errors) - return operations, None + if op.operation is OperationType.MOVE and not op.new_path: + parse_errors.append( + f"MOVE {op.file_path!r}: missing destination path (expected 'src -> dst')") + return ([], "Parse error: " + "; ".join(parse_errors)) if parse_errors else (operations, None) def _count_occurrences(text: str, pattern: str) -> int: @@ -136,66 +123,57 @@ def _count_occurrences(text: str, pattern: str) -> int: def _split_hunk(hunk: Hunk) -> Tuple[List[str], List[str]]: """``(search_lines, replace_lines)``: context+removed vs context+added.""" - search = [l.content for l in hunk.lines if l.prefix in {' ', '-'}] - replace = [l.content for l in hunk.lines if l.prefix in {' ', '+'}] - return search, replace + return ([l.content for l in hunk.lines if l.prefix != '+'], + [l.content for l in hunk.lines if l.prefix != '-']) def _no_match_hint(error: Optional[str], search_pattern: str, content: str) -> str: """Best-effort 'Did you mean...' suffix; never lets a hint failure mask the real error.""" - try: + with contextlib.suppress(Exception): from tools.fuzzy_match import format_no_match_hint return format_no_match_hint(error, 0, search_pattern, content) - except Exception: - return "" + return "" def _hint_ambiguity(content: str, hint: str, tail: str = "") -> Tuple[int, str]: """(occurrences, error) for an addition-only hunk's context hint; error is '' when unique.""" - occurrences = _count_occurrences(content, hint) - if occurrences > 1: - return occurrences, (f"context hint '{hint}' is ambiguous " - f"({occurrences} occurrences){tail}") - return occurrences, "" + n = _count_occurrences(content, hint) + return n, f"context hint '{hint}' is ambiguous ({n} occurrences){tail}" if n > 1 else "" def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> List[str]: - """Dry-run every operation; return error strings (empty list = safe to apply). UPDATE - hunks are simulated in order so later hunks see post-earlier-hunk content, as apply will.""" + """Dry-run every operation -> error strings (empty = safe). UPDATE hunks are simulated in + order so later hunks see post-earlier-hunk content, exactly as apply will.""" from tools.fuzzy_match import fuzzy_find_and_replace, is_already_applied - errors: List[str] = [] real_change_count = 0 - - # Virtual overlay so inter-op state validates (e.g. a MOVE creating the destination - # a later UPDATE targets): path -> content from an earlier op; paths MOVE/DELETE removed. + # Overlay so inter-op state validates (a MOVE creating the path a later UPDATE targets). pending_content: dict = {} removed_paths: set = set() - def _read(path: str): - if path in removed_paths and path not in pending_content: - return None, "file not found" + def _read(path: str) -> Tuple[Optional[str], Optional[str]]: if path in pending_content: return pending_content[path], None + if path in removed_paths: + return None, "file not found" r = file_ops.read_file_raw(path) return (None, r.error) if r.error else (r.content, None) def _validate_update(op: PatchOperation) -> None: nonlocal real_change_count - content, read_err = _read(op.file_path) + simulated, read_err = _read(op.file_path) if read_err: errors.append(f"{op.file_path}: {read_err}") return - simulated = content for hunk_index, hunk in enumerate(op.hunks, start=1): search_lines, replace_lines = _split_hunk(hunk) - if not any(l.prefix in '-+' for l in hunk.lines): - # Inert anchor hunk (context only) — models emit these - # between real changes; ignore without failing the patch. + if search_lines == replace_lines: + # Context-only anchor hunks (models emit these between changes) are inert; identical + # -/+ lines are skipped by apply as a no-op — neither may fail validation. + real_change_count += any(l.prefix in '-+' for l in hunk.lines) continue real_change_count += 1 - if not search_lines: - # Addition-only hunk: the context hint must be unique. + if not search_lines: # addition-only: the context hint must be unique if hunk.context_hint: occurrences, ambiguous = _hint_ambiguity(simulated, hunk.context_hint) if occurrences == 0: @@ -204,19 +182,13 @@ def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> Lis elif ambiguous: errors.append(f"{op.file_path}: addition-only hunk {ambiguous}") continue - search_pattern = '\n'.join(search_lines) - replacement = '\n'.join(replace_lines) - if search_lines == replace_lines: - # Identical -/+ lines: apply skips it as a no-op, so - # validation must not reject it with the identical-strings error. - continue + search_pattern, replacement = '\n'.join(search_lines), '\n'.join(replace_lines) new_simulated, count, _strategy, match_error = fuzzy_find_and_replace( simulated, search_pattern, replacement, replace_all=False) if count: simulated = new_simulated elif not is_already_applied(simulated or "", search_pattern, replacement): - # Already-applied hunks (edit landed in a prior call) are no-ops so - # multi-hunk patches don't fail wholesale; apply performs the same skip. + # Already-applied hunks are no-ops (apply performs the same skip). label = f"'{hunk.context_hint}'" if hunk.context_hint else "(no hint)" errors.append( f"{op.file_path}: hunk {hunk_index} {label} not found" @@ -224,18 +196,20 @@ def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> Lis + _no_match_hint(match_error, search_pattern, simulated)) pending_content[op.file_path] = simulated + def _remove(path: str) -> None: + removed_paths.add(path) + pending_content.pop(path, None) + for op in operations: if op.operation == OperationType.UPDATE: _validate_update(op) continue real_change_count += 1 if op.operation == OperationType.DELETE: - _content, read_err = _read(op.file_path) - if read_err: + if _read(op.file_path)[1]: errors.append(f"{op.file_path}: file not found for deletion") else: - removed_paths.add(op.file_path) - pending_content.pop(op.file_path, None) + _remove(op.file_path) elif op.operation == OperationType.MOVE: if not op.new_path: errors.append(f"{op.file_path}: MOVE operation missing destination path") @@ -243,16 +217,12 @@ def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> Lis src_content, src_err = _read(op.file_path) if src_err: errors.append(f"{op.file_path}: source file not found for move") - _dst, dst_err = _read(op.new_path) - if not dst_err: + if not _read(op.new_path)[1]: errors.append(f"{op.new_path}: destination already exists — move would overwrite") - # Only a cleanly-validated move updates the overlay. - if not src_err and dst_err: + elif not src_err: # only a cleanly-validated move updates the overlay pending_content[op.new_path] = src_content if src_content is not None else "" - pending_content.pop(op.file_path, None) - removed_paths.add(op.file_path) - # ADD: parent directory creation handled by write_file; no pre-check needed. - + _remove(op.file_path) + # ADD: write_file creates parent directories; no pre-check needed. if not errors and real_change_count == 0: errors.append("Patch contains no changes (only context lines were provided)") return errors @@ -262,104 +232,101 @@ def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> Lis ApplyResult = Tuple[bool, str, Optional[str], Optional[dict]] +def _fail(error: str) -> ApplyResult: + return False, error, None, None + + +def _written(result: Any, diff: str) -> ApplyResult: + """Outcome of a write: its error, else success with LSP/lint propagated from the WriteResult.""" + if result.error: + return _fail(result.error) + return True, diff, getattr(result, "lsp_diagnostics", None), getattr(result, "lint", None) + + +def _unified_diff(path: str, old: str, new: Optional[str]) -> str: + """Unified diff ``a/path`` -> ``b/path`` (``new=None`` = deletion, ``/dev/null``).""" + return ''.join(difflib.unified_diff( + old.splitlines(keepends=True), [] if new is None else new.splitlines(keepends=True), + fromfile=f"a/{path}", tofile="/dev/null" if new is None else f"b/{path}")) + + def apply_v4a_operations(operations: List[PatchOperation], file_ops: Any) -> 'PatchResult': - """Validate all operations, then apply them (two-phase, atomic on validation failure). - A phase-2 failure (validate/apply race) is reported with a ``git diff`` note since state - may be inconsistent. ``file_ops`` needs read_file_raw/write_file/delete_file/move_file.""" + """Two-phase: validate everything, then apply (atomic on validation failure). A phase-2 + failure (validate/apply race) carries a ``git diff`` note since state may be inconsistent. + ``file_ops`` needs read_file_raw/write_file/delete_file/move_file.""" from tools.file_operations import PatchResult # avoid circular import - validation_errors = _validate_operations(operations, file_ops) - if validation_errors: + def _bullets(errs: List[str]) -> str: + return "\n".join(f" • {e}" for e in errs) + + if errors := _validate_operations(operations, file_ops): return PatchResult( success=False, - error="Patch validation failed (no files were modified):\n" - + "\n".join(f" • {e}" for e in validation_errors)) - + error="Patch validation failed (no files were modified):\n" + _bullets(errors)) files: Dict[str, List[str]] = {"created": [], "deleted": [], "modified": []} all_diffs: List[str] = [] - # V4A bypasses the WriteResult/PatchResult plumbing that write_file uses, - # so LSP diagnostics and lint must be propagated explicitly per file. + # V4A bypasses write_file's WriteResult plumbing: LSP diagnostics and lint propagate per file. lsp_blocks: List[str] = [] - errors: List[str] = [] lint_results: Dict[str, dict] = {} - for op in operations: + handler, verb, bucket = _APPLY_DISPATCH[op.operation] try: - handler, verb, bucket = _APPLY_DISPATCH[op.operation] ok, payload, lsp, lint = handler(op, file_ops) - if not ok: - errors.append(f"Failed to {verb} {op.file_path}: {payload}") - continue - label = op.file_path - if op.operation is OperationType.MOVE: - label = f"{op.file_path} -> {op.new_path}" - files[bucket].append(label) - all_diffs.append(payload) - if lsp: - lsp_blocks.append(lsp) - if lint: - lint_results[op.file_path] = lint except Exception as e: - errors.append(f"Error processing {op.file_path}: {str(e)}") - - # Each LSP block carries its own header, so plain - # concatenation keeps per-file attribution. - result_kwargs = dict( + ok, payload = None, str(e) + if not ok: + prefix = f"Failed to {verb}" if ok is False else "Error processing" + errors.append(f"{prefix} {op.file_path}: {payload}") + continue + is_move = op.operation is OperationType.MOVE + files[bucket].append(f"{op.file_path} -> {op.new_path}" if is_move else op.file_path) + all_diffs.append(payload) + if lsp: + lsp_blocks.append(lsp) + if lint: + lint_results[op.file_path] = lint + # Each LSP block carries its own header; joining keeps attribution. + return PatchResult( + success=not errors, + error=("Apply phase failed (state may be inconsistent — run `git diff` to assess):\n" + + _bullets(errors)) if errors else None, diff='\n'.join(all_diffs), files_modified=files["modified"], files_created=files["created"], files_deleted=files["deleted"], - lint=lint_results if lint_results else None, - lsp_diagnostics="\n\n".join(lsp_blocks) if lsp_blocks else None) - if errors: - return PatchResult( - success=False, - error="Apply phase failed (state may be inconsistent — run `git diff` to assess):\n" - + "\n".join(f" • {e}" for e in errors), - **result_kwargs) - return PatchResult(success=True, **result_kwargs) + lint=lint_results or None, lsp_diagnostics="\n\n".join(lsp_blocks) or None) def _write_file_accepts_pre_content(file_ops: Any) -> bool: - """True when ``file_ops.write_file`` accepts ``pre_content``. Decided from the signature, - not by catching TypeError around the call, so a TypeError raised *inside* a capable - write_file propagates instead of triggering a duplicate write.""" - try: + """Whether ``file_ops.write_file`` accepts ``pre_content`` — read from the signature, not by + catching TypeError around the call, so a TypeError raised *inside* it can't double-write.""" + with contextlib.suppress(TypeError, ValueError): params = inspect.signature(file_ops.write_file).parameters - except (TypeError, ValueError): - return False - return "pre_content" in params or any( - p.kind is inspect.Parameter.VAR_KEYWORD for p in params.values()) + return "pre_content" in params or any( + p.kind is inspect.Parameter.VAR_KEYWORD for p in params.values()) + return False def _apply_add(op: PatchOperation, file_ops: Any) -> ApplyResult: """Create a file from the hunks' '+' lines.""" content_lines = [line.content for hunk in op.hunks for line in hunk.lines if line.prefix == '+'] result = file_ops.write_file(op.file_path, '\n'.join(content_lines)) - if result.error: - return False, result.error, None, None diff = f"--- /dev/null\n+++ b/{op.file_path}\n" + '\n'.join(f"+{line}" for line in content_lines) - return True, diff, getattr(result, "lsp_diagnostics", None), getattr(result, "lint", None) + return _written(result, diff) def _apply_delete(op: PatchOperation, file_ops: Any) -> ApplyResult: """Delete a file, producing a real unified diff of the removed content.""" - # Validation already confirmed existence; the re-read guards against races. - read_result = file_ops.read_file_raw(op.file_path) + read_result = file_ops.read_file_raw(op.file_path) # re-read guards validate/apply races if read_result.error: - return False, f"Cannot delete {op.file_path}: file not found", None, None + return _fail(f"Cannot delete {op.file_path}: file not found") result = file_ops.delete_file(op.file_path) - if result.error: - return False, result.error, None, None - diff = ''.join(difflib.unified_diff( - read_result.content.splitlines(keepends=True), [], - fromfile=f"a/{op.file_path}", tofile="/dev/null")) - return True, diff or f"# Deleted: {op.file_path}", None, None + diff = _unified_diff(op.file_path, read_result.content, None) or f"# Deleted: {op.file_path}" + return _fail(result.error) if result.error else (True, diff, None, None) def _apply_move(op: PatchOperation, file_ops: Any) -> ApplyResult: result = file_ops.move_file(op.file_path, op.new_path) - if result.error: - return False, result.error, None, None - return True, f"# Moved: {op.file_path} -> {op.new_path}", None, None + return _fail(result.error) if result.error else ( + True, f"# Moved: {op.file_path} -> {op.new_path}", None, None) def _insert_addition_only(new_content: str, hunk: Hunk, insert_text: str) -> Tuple[Optional[str], Optional[str]]: @@ -371,71 +338,54 @@ def _insert_addition_only(new_content: str, hunk: Hunk, insert_text: str) -> Tup return None, f"Addition-only hunk: {ambiguous}" if occurrences == 1: eol = new_content.find('\n', new_content.find(hunk.context_hint)) - if eol != -1: - return new_content[:eol + 1] + insert_text + '\n' + new_content[eol + 1:], None - return new_content + '\n' + insert_text, None - # Hint not found — append at end as a safe fallback. + if eol == -1: + return new_content + '\n' + insert_text, None + return new_content[:eol + 1] + insert_text + '\n' + new_content[eol + 1:], None + # No hint / hint not found — append at end as a safe fallback. return new_content.rstrip('\n') + '\n' + insert_text + '\n', None def _apply_update(op: PatchOperation, file_ops: Any) -> ApplyResult: """Apply each hunk via fuzzy replace, then write once.""" from tools.fuzzy_match import fuzzy_find_and_replace, is_already_applied - - # Raw read: no line-number prefixes or per-line truncation. - read_result = file_ops.read_file_raw(op.file_path) + read_result = file_ops.read_file_raw(op.file_path) # raw: no line numbers / truncation if read_result.error: - return False, f"Cannot read file: {read_result.error}", None, None - current_content = read_result.content - new_content = current_content - + return _fail(f"Cannot read file: {read_result.error}") + current_content = new_content = read_result.content for hunk in op.hunks: search_lines, replace_lines = _split_hunk(hunk) if search_lines and search_lines == replace_lines: continue + search_pattern, replacement = '\n'.join(search_lines), '\n'.join(replace_lines) if not search_lines: - new_content, err = _insert_addition_only(new_content, hunk, '\n'.join(replace_lines)) + new_content, err = _insert_addition_only(new_content, hunk, replacement) if err: - return False, err, None, None + return _fail(err) continue - - search_pattern = '\n'.join(search_lines) - replacement = '\n'.join(replace_lines) new_content, count, _strategy, error = fuzzy_find_and_replace( new_content, search_pattern, replacement, replace_all=False) if not (error and count == 0): continue - # Retry inside a window around the context hint, if any. hint_pos = new_content.find(hunk.context_hint) if hunk.context_hint else -1 if hint_pos != -1: window_start = max(0, hint_pos - 500) window_end = min(len(new_content), hint_pos + 2000) window_new, count, _strategy, error = fuzzy_find_and_replace( - new_content[window_start:window_end], search_pattern, replacement, replace_all=False - ) + new_content[window_start:window_end], search_pattern, replacement, replace_all=False) if count > 0: new_content = new_content[:window_start] + window_new + new_content[window_end:] error = None if error: - # Mirror the validation-phase already-applied skip, or the two - # phases disagree and the whole patch fails here. + # Mirror validation's already-applied skip, else the two phases disagree and fail here. if is_already_applied(new_content, search_pattern, replacement): continue - return False, f"Could not apply hunk: {error}" + _no_match_hint(error, search_pattern, new_content), None, None - + hint = _no_match_hint(error, search_pattern, new_content) + return _fail(f"Could not apply hunk: {error}" + hint) # Pass pre_content to skip a redundant re-read inside write_file when supported. - if _write_file_accepts_pre_content(file_ops): - write_result = file_ops.write_file(op.file_path, new_content, pre_content=current_content) - else: - write_result = file_ops.write_file(op.file_path, new_content) - if write_result.error: - return False, write_result.error, None, None - - diff = ''.join(difflib.unified_diff( - current_content.splitlines(keepends=True), new_content.splitlines(keepends=True), - fromfile=f"a/{op.file_path}", tofile=f"b/{op.file_path}")) - return True, diff, getattr(write_result, "lsp_diagnostics", None), getattr(write_result, "lint", None) + extra = {"pre_content": current_content} if _write_file_accepts_pre_content(file_ops) else {} + write_result = file_ops.write_file(op.file_path, new_content, **extra) + return _written(write_result, _unified_diff(op.file_path, current_content, new_content)) # operation -> (handler, verb for error text, files_* bucket) diff --git a/tools/read_extract.py b/tools/read_extract.py index 23a4299af2..a96e5431b3 100644 --- a/tools/read_extract.py +++ b/tools/read_extract.py @@ -1,16 +1,14 @@ -"""Stdlib document-to-text extraction for ``read_file``. - -Jupyter, DOCX and XLSX need no dependencies. The optional ``firecrawl-anydoc`` package -(imports as ``anydoc``) widens coverage to legacy Office, OpenDocument, RTF, EPUB and PDF. -The stdlib extractors stay authoritative for their three formats so behavior is identical -with or without anydoc. Malformed documents raise :class:`ExtractionError`; callers then -fall back to normal text/binary handling. -""" +"""Document-to-text extraction for ``read_file``: stdlib Jupyter/DOCX/XLSX (always +authoritative for those three), plus legacy Office/OpenDocument/RTF/EPUB/PDF when the +optional ``firecrawl-anydoc`` package (imports as ``anydoc``) is installed. Malformed +documents raise :class:`ExtractionError`; callers fall back to text/binary handling.""" from __future__ import annotations import contextlib +import functools import importlib +import itertools import json import os import posixpath @@ -29,12 +27,10 @@ __all__ = ["EXTRACTABLE_EXTENSIONS", "ExtractionError", "extract_document_bytes" "extract_document_text", "is_extractable_document"] EXTRACTABLE_EXTENSIONS = frozenset({".ipynb", ".docx", ".xlsx"}) -# Formats handled only when the optional anydoc converter is installed. ANYDOC_EXTENSIONS = frozenset({ ".doc", ".docm", ".ppt", ".pps", ".pot", ".pptx", ".pptm", ".ppsx", ".ppsm", ".xls", ".xlsm", ".xlsb", ".odt", ".ods", ".odp", ".rtf", ".epub", ".pdf"}) -# anydoc loads whole files with no streaming and the read_file char budget only applies -# after conversion — cap the input size. +# anydoc loads whole files (no streaming); read_file's char budget applies only post-conversion. MAX_ANYDOC_BYTES = 50 * 1024 * 1024 MAX_DOCUMENT_BYTES = 50 * 1024 * 1024 _MAX_XLSX_ROWS_PER_SHEET = 5000 @@ -59,15 +55,15 @@ def _extension(path: str) -> str: _ANYDOC_UNSET = object() _anydoc_module: Any = _ANYDOC_UNSET _anydoc_lock = threading.Lock() -# After a failed load, wait this long before retrying: the attempt can shell out -# to pip, so retrying every call would hammer the network where install can't succeed. +# Cooldown after a failed load: the attempt can shell out to pip, so retrying every call would +# hammer the network where install can't succeed. ANYDOC_RETRY_SECONDS = 300.0 _anydoc_failed_at: Optional[float] = None def _anydoc() -> Optional[Any]: - """Lazily import the optional anydoc converter; None when unavailable. A failed load is - retried after ANYDOC_RETRY_SECONDS so one transient pip/network blip does not stick.""" + """Lazily import the optional anydoc converter (None when unavailable; failures retried after + ANYDOC_RETRY_SECONDS so one transient pip/network blip does not stick).""" global _anydoc_module, _anydoc_failed_at if _anydoc_module is not _ANYDOC_UNSET: return _anydoc_module @@ -79,8 +75,7 @@ def _anydoc() -> Optional[Any]: return None try: from tools.lazy_deps import ensure as _lazy_ensure - # prompt=False: read_file must never block on an install prompt. - _lazy_ensure("tool.doc_extract", prompt=False) + _lazy_ensure("tool.doc_extract", prompt=False) # read_file must never block on a prompt _anydoc_module = importlib.import_module("anydoc") except Exception: # install failure, ImportError or a broken native binding _anydoc_failed_at = time.monotonic() @@ -101,23 +96,19 @@ def _check_size(size: int, limit: int) -> None: @contextlib.contextmanager def _temp_copy(data: bytes, suffix: str) -> Iterator[str]: """Materialize backend bytes in a private host temp file; removed even when parsing fails.""" - temp_path = "" + with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as fh: + fh.write(data) try: - with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as fh: - fh.write(data) - temp_path = fh.name - yield temp_path + yield fh.name finally: - if temp_path: - with contextlib.suppress(OSError): - os.unlink(temp_path) + with contextlib.suppress(OSError): + os.unlink(fh.name) def extract_document_text(path: str) -> str: ext = _extension(path) - extractor = _STDLIB_EXTRACTORS.get(ext) - if extractor is not None: - return extractor(path) + if ext in _STDLIB_EXTRACTORS: + return _STDLIB_EXTRACTORS[ext](path) if ext in ANYDOC_EXTENSIONS: return _extract_anydoc(path) raise ExtractionError(f"Unsupported document type: {path!r}") @@ -131,14 +122,12 @@ def extract_document_bytes(data: bytes, path: str) -> str: return _extract_anydoc_bytes(data, path) if ext not in EXTRACTABLE_EXTENSIONS: raise ExtractionError(f"Unsupported document type: {path!r}") - # The stdlib extractors are path-oriented. - with _temp_copy(data, ext) as temp_path: - return extract_document_text(temp_path) + with _temp_copy(data, ext) as temp_path: # the stdlib extractors are path-oriented + return _STDLIB_EXTRACTORS[ext](temp_path) def _anydoc_missing_error(path: str) -> str: - """Teaching error for anydoc-gated formats (deliberately absent from the schema so only - sessions that hit one pay for it).""" + """Teaching text for anydoc-gated formats (not in the schema: only sessions hitting one pay).""" return ( f"Cannot convert {path!r}: this format needs the optional anydoc " "converter, which is not installed (install blocked or first " @@ -149,10 +138,10 @@ def _anydoc_missing_error(path: str) -> str: def _hosted_ocr_config() -> tuple: - """Resolve hosted-OCR settings: (enabled, api_key, api_url). Never raises; no network. - Maintainer decision: the ONLY route is a direct ``FIRECRAWL_API_KEY`` (anydoc defaults the - api_url); the Nous gateway is NOT used — its Parse proxy live-probed broken (revisit when - it grows Parse support). ``file_tools.hosted_ocr: false`` disables even with a key.""" + """(enabled, api_key, api_url); never raises, no network. Maintainer decision: the ONLY route + is a direct ``FIRECRAWL_API_KEY`` (anydoc defaults api_url); the Nous gateway's Parse proxy + live-probed broken, so it is NOT used. ``file_tools.hosted_ocr: false`` disables even with a + key.""" api_key = os.environ.get("FIRECRAWL_API_KEY") or None enabled = api_key is not None with contextlib.suppress(Exception): @@ -164,35 +153,30 @@ def _hosted_ocr_config() -> tuple: def hosted_ocr_available() -> bool: - """Public probe for read_file's schema line; a key failing at conversion time lands in - the NEEDS-OCR warning instead.""" + """Probe for read_file's schema line; a key failing at conversion time surfaces in NEEDS-OCR.""" return _hosted_ocr_config()[0] def _needs_ocr_warning(path: str, pages, hosted_error: str = "") -> str: - """Result text when anydoc raises NeedsOcrError and hosted OCR is off/failed. Hints at - CHECKING for an OCR skill (never names one) and never advertises the hosted_ocr knob.""" + """NeedsOcrError result when hosted OCR is off/failed; hints at CHECKING for an OCR skill + (never names one) and never advertises the hosted_ocr knob.""" page_list = ", ".join(str(p) for p in pages) if pages else "unknown" - msg = ( + hosted = f"Hosted OCR was attempted and failed ({hosted_error}). " if hosted_error else "" + return ( f"[NEEDS OCR: pages {page_list} of this PDF are scanned images " - "with no text layer — their content is MISSING below. ") - if hosted_error: - msg += f"Hosted OCR was attempted and failed ({hosted_error}). " - msg += ( + f"with no text layer — their content is MISSING below. {hosted}" "If the missing pages matter: render just those pages with " f"`pdftoppm -jpeg -r 150 -f -l '{path}' /tmp/page` " "and inspect via vision_analyze, or check whether an OCR skill is " - "available (skills_list).") - return msg + "]\n" + "available (skills_list).]\n") def _finalize_anydoc_text(text: Any, path: str, pdf_note: Callable[[], str]) -> str: - """Normalize converter output and, for PDFs, PREPEND the coverage note (read_file - paginates: a footer may never be fetched). Covers PARTIAL gaps without NeedsOcrError.""" + """Normalize converter output; PDFs get the coverage note PREPENDED (read_file paginates, so a + footer may never be fetched) — this covers PARTIAL scan gaps that raise no NeedsOcrError.""" if not isinstance(text, str) or not text.strip(): raise ExtractionError("Document contains no extractable text") - note = pdf_note() if Path(path).suffix.lower() == ".pdf" else "" - return (note or "") + text.rstrip("\n") + "\n" + return (pdf_note() if Path(path).suffix.lower() == ".pdf" else "") + text.rstrip("\n") + "\n" def _ocr_scanned_pdf(mod: Any, path: str, exc: BaseException) -> str: @@ -206,40 +190,33 @@ def _ocr_scanned_pdf(mod: Any, path: str, exc: BaseException) -> str: return mod.to_markdown(path, ocr="hosted", **extra).rstrip("\n") + "\n" except Exception as hosted_exc: # noqa: BLE001 hosted_error = f"{type(hosted_exc).__name__}: {hosted_exc}" - # No route / disabled / hosted failed: whole doc is scans — the warning IS the result. - return _needs_ocr_warning(path, pages, hosted_error) - - -def _require_anydoc(path: str) -> Any: - mod = _anydoc() - if mod is None: - raise ExtractionError(_anydoc_missing_error(path)) - return mod + return _needs_ocr_warning(path, pages, hosted_error) # whole doc is scans: the warning IS it def _extract_anydoc(path: str) -> str: - mod = _require_anydoc(path) - try: - size = os.path.getsize(path) - except OSError as exc: - raise ExtractionError(str(exc)) from exc - _check_size(size, MAX_ANYDOC_BYTES) + mod = _anydoc() + if mod is None: + raise ExtractionError(_anydoc_missing_error(path)) try: + _check_size(os.path.getsize(path), MAX_ANYDOC_BYTES) text = mod.to_markdown(path) + except ExtractionError: + raise except OSError as exc: raise ExtractionError(str(exc)) from exc except Exception as exc: needs_ocr = getattr(mod, "NeedsOcrError", None) if needs_ocr is not None and isinstance(exc, needs_ocr): return _ocr_scanned_pdf(mod, path, exc) - # anydoc raises one ConvertError subclass per failure mode (Unsupported, Malformed, - # Encrypted, ResourceLimit, MissingPart); all mean "no meaningful text". + # Any ConvertError subclass (Unsupported/Malformed/Encrypted/...) = "no meaningful text". raise ExtractionError(f"{type(exc).__name__}: {exc}") from exc return _finalize_anydoc_text(text, path, lambda: _pdf_coverage_note(path)) def _extract_anydoc_bytes(data: bytes, path: str) -> str: - mod = _require_anydoc(path) + mod = _anydoc() + if mod is None: + raise ExtractionError(_anydoc_missing_error(path)) _check_size(len(data), MAX_ANYDOC_BYTES) try: text = mod.to_markdown_bytes(data) @@ -248,17 +225,13 @@ def _extract_anydoc_bytes(data: bytes, path: str) -> str: return _finalize_anydoc_text(text, path, lambda: _pdf_coverage_note_from_bytes(data, path)) -# ── Scanned-PDF coverage detection: text-layer extractors return nothing for scanned -# pages, so a mostly-scanned PDF converts "successfully" into headers with empty bodies — -# silent data loss. Count per-page text via pdftotext (form-feed separated) and warn. +# ── Scanned-PDF coverage: text-layer extractors return nothing for scanned pages, so a mostly +# scanned PDF converts "successfully" into silent data loss. Count per-page text via pdftotext. PDF_EMPTY_PAGE_CHARS = 20 # fewer extracted chars than this = empty page # Warn when empty pages reach both MIN_EMPTY and MIN_RATIO, or ABSOLUTE_EMPTY alone. -PDF_COVERAGE_MIN_EMPTY = 2 -PDF_COVERAGE_MIN_RATIO = 0.2 -PDF_COVERAGE_ABSOLUTE_EMPTY = 10 +PDF_COVERAGE_MIN_EMPTY, PDF_COVERAGE_MIN_RATIO, PDF_COVERAGE_ABSOLUTE_EMPTY = 2, 0.2, 10 PDF_PAGE_SCAN_TIMEOUT = 20.0 -# Cap the per-gap breakdown so alternating text/scan pages can't balloon the warning. -PDF_GAP_MAP_MAX_ENTRIES = 20 +PDF_GAP_MAP_MAX_ENTRIES = 20 # cap so alternating text/scan pages can't balloon the warning _GAP_CONTEXT_CHARS = 60 @@ -279,14 +252,10 @@ def _pdf_page_texts(path: str) -> Optional[list[str]]: def _gap_map(counts: list[int], texts: list[str], empty: list[int]) -> str: - """Per-gap breakdown, each empty range labeled with the last text seen before - it (usually a section header), so the agent can pick WHICH gaps to OCR.""" - ranges: list[list[int]] = [] # sorted 1-based page numbers -> [start, end] runs - for p in empty: - if ranges and p == ranges[-1][1] + 1: - ranges[-1][1] = p - else: - ranges.append([p, p]) + """Per-gap breakdown labeled with the text before each gap, so the agent picks which to OCR.""" + # Sorted 1-based page numbers -> (start, end) runs; consecutive pages share ``page - index``. + runs = [list(g) for _k, g in itertools.groupby(enumerate(empty), lambda e: e[1] - e[0])] + ranges = [(run[0][1], run[-1][1]) for run in runs] lines: list[str] = [] for a, b in ranges[:PDF_GAP_MAP_MAX_ENTRIES]: label = "" @@ -305,17 +274,17 @@ def _gap_map(counts: list[int], texts: list[str], empty: list[int]) -> str: def _pdf_coverage_note(path: str, display_path: Optional[str] = None) -> str: - """Warning header when many PDF pages produced no text, else ''. ``path`` is scanned - (may be a host temp file); ``display_path`` is what the recovery command shows.""" + """Warning header when many pages yielded no text, else ''. ``display_path`` (default ``path``, + which may be a host temp file) is what the recovery command shows.""" texts = _pdf_page_texts(path) if not texts or len(texts) < 2: return "" counts = [len(page.strip()) for page in texts] empty = [i + 1 for i, n in enumerate(counts) if n < PDF_EMPTY_PAGE_CHARS] total = len(counts) - if len(empty) < PDF_COVERAGE_MIN_EMPTY or ( - len(empty) / total < PDF_COVERAGE_MIN_RATIO and len(empty) < PDF_COVERAGE_ABSOLUTE_EMPTY - ): + n_empty = len(empty) + enough = n_empty / total >= PDF_COVERAGE_MIN_RATIO or n_empty >= PDF_COVERAGE_ABSOLUTE_EMPTY + if n_empty < PDF_COVERAGE_MIN_EMPTY or not enough: return "" shown = display_path or path return ( @@ -335,13 +304,10 @@ def _pdf_coverage_note(path: str, display_path: Optional[str] = None) -> str: def _pdf_coverage_note_from_bytes(data: bytes, display_path: str) -> str: - """Coverage note for backend-transferred PDF bytes via a host temp copy (pdftotext is - path-oriented); the recovery command still names ``display_path``.""" - try: - with _temp_copy(data, ".pdf") as temp_path: - return _pdf_coverage_note(temp_path, display_path=display_path) - except OSError: - return "" + """Coverage note for backend PDF bytes via a host temp copy (pdftotext needs a path).""" + with contextlib.suppress(OSError), _temp_copy(data, ".pdf") as temp_path: + return _pdf_coverage_note(temp_path, display_path=display_path) + return "" def _joined(lines: list[str], empty_error: str) -> str: @@ -352,11 +318,10 @@ def _joined(lines: list[str], empty_error: str) -> str: def _source_text(source) -> str: - if isinstance(source, str): - return source + """Notebook source/text fields are a str or a list of str fragments.""" if isinstance(source, list): - return "".join(item for item in source if isinstance(item, str)) - return "" + source = "".join(item for item in source if isinstance(item, str)) + return source if isinstance(source, str) else "" def _human_size(n_bytes: int) -> str: @@ -366,30 +331,24 @@ def _human_size(n_bytes: int) -> str: def _base64_bytes(payload: str) -> int: """Approximate decoded size of a base64 payload (whitespace ignored).""" clean = re.sub(r"[^0-9+/=A-Za-z]", "", payload) - padding = min(2, len(clean) - len(clean.rstrip("="))) - return max(0, (len(clean) * 3) // 4 - padding) + return max(0, (len(clean) * 3) // 4 - min(2, len(clean) - len(clean.rstrip("=")))) def _clean_stream_text(text: str) -> str: """Strip ANSI escapes; keep only the final ``\\r`` frame of each line (tqdm redraws).""" from tools.ansi_strip import strip_ansi - lines = [] - for line in strip_ansi(text).replace("\r\n", "\n").split("\n"): - frames = [frame for frame in line.split("\r") if frame] - lines.append(frames[-1] if frames else "") - return "\n".join(lines) + return "\n".join(([f for f in line.split("\r") if f] or [""])[-1] + for line in strip_ansi(text).replace("\r\n", "\n").split("\n")) -# Per-output-block truncation so one runaway training log cannot flood the extraction. -_MAX_OUTPUT_CHARS = 20_000 +_MAX_OUTPUT_CHARS = 20_000 # per code cell, so one runaway training log cannot flood the extraction # nbformat v3 stores mime data flat on the output dict under these keys. _V3_MIME_KEYS = (("png", "image/png"), ("jpeg", "image/jpeg"), ("svg", "image/svg+xml"), ("html", "text/html")) def _notebook_output_text(output: Any) -> str: - """Render one notebook output as compact text: stream text, tracebacks and textual - results kept; token-heavy payloads (images, HTML, widgets) become sized placeholders. - Handles nbformat v4 and legacy v3 (``pyout``/``pyerr``) shapes.""" + """One notebook output as compact text: stream/traceback/textual results kept; token-heavy + payloads (images, HTML, widgets) become sized placeholders. Handles v4 and legacy v3 shapes.""" if not isinstance(output, dict): return "" otype = output.get("output_type") @@ -397,52 +356,41 @@ def _notebook_output_text(output: Any) -> str: body = _clean_stream_text(_source_text(output.get("text", ""))) return body if body.strip() else "" if otype in {"error", "pyerr"}: - traceback = output.get("traceback") - tb_text = "" - if isinstance(traceback, list): - tb_text = _clean_stream_text( - "\n".join(line for line in traceback if isinstance(line, str))) + tb = output.get("traceback") + tb_text = _clean_stream_text("\n".join(filter(lambda l: isinstance(l, str), tb)) + if isinstance(tb, list) else "") header = f"Error: {output.get('ename', '')}: {output.get('evalue', '')}".rstrip(": ") return f"{header}\n{tb_text}".rstrip() if otype not in {"execute_result", "display_data", "pyout"}: return "" - data = output.get("data") if not isinstance(data, dict): # legacy v3: mime payloads sit flat on the output dict data = {"text/plain": output["text"]} if isinstance(output.get("text"), (str, list)) else {} data.update((mime, output[k]) for k, mime in _V3_MIME_KEYS if k in output) if "application/vnd.jupyter.widget-view+json" in data: return "[interactive widget — omitted]" - # Prefer readable text: models consume text/plain far better than markup. - for mime in ("text/plain", "text/markdown"): - if mime in data: - body = _clean_stream_text(_source_text(data[mime])) - if body.strip(): - return body + for mime in ("text/plain", "text/markdown"): # models consume text far better than markup + body = _clean_stream_text(_source_text(data[mime])) if mime in data else "" + if body.strip(): + return body for mime, value in data.items(): if isinstance(mime, str) and mime.startswith("image/"): - size = _base64_bytes(_source_text(value)) - return f"[{mime} output — {_human_size(size)}, omitted]" + return f"[{mime} output — {_human_size(_base64_bytes(_source_text(value)))}, omitted]" if "text/html" in data: - html = _source_text(data["text/html"]) - return f"[text/html output — {len(html):,} chars, omitted]" - mimes = ", ".join(str(m) for m in data) or "unknown" - return f"[{mimes} output — omitted]" + return f"[text/html output — {len(_source_text(data['text/html'])):,} chars, omitted]" + return f"[{', '.join(str(m) for m in data) or 'unknown'} output — omitted]" def _notebook_outputs(cell: dict, jq_pointer: str = "", filename: str = "") -> str: outputs = cell.get("outputs") if not isinstance(outputs, list): return "" - blocks = [text for text in (_notebook_output_text(o) for o in outputs) if text] - if not blocks: - return "" - joined = "\n".join(blocks) - if len(joined) > _MAX_OUTPUT_CHARS: - omitted = len(joined) - _MAX_OUTPUT_CHARS - hint = f" — full output: jq -r '{jq_pointer}' {filename}" if jq_pointer and filename else "" - joined = joined[:_MAX_OUTPUT_CHARS] + f"\n… [{omitted:,} output chars truncated{hint}]" - return joined + joined = "\n".join(filter(None, map(_notebook_output_text, outputs))) + if len(joined) <= _MAX_OUTPUT_CHARS: + return joined + hint = f" — full output: jq -r '{jq_pointer}' {filename}" if jq_pointer and filename else "" + omitted = len(joined) - _MAX_OUTPUT_CHARS + return joined[:_MAX_OUTPUT_CHARS] + f"\n… [{omitted:,} output chars truncated{hint}]" _CELL_LABELS = {"markdown": "Markdown", "code": "Code", "raw": "Raw"} @@ -459,7 +407,7 @@ def _extract_notebook(path: str) -> str: raw_cells = nb.get("cells") if isinstance(raw_cells, list): cells = [(f".cells[{i}].outputs", cell) for i, cell in enumerate(raw_cells)] - else: + else: # nbformat v3: cells live under worksheets cells = [ (f".worksheets[{wi}].cells[{ci}].outputs", cell) for wi, ws in enumerate(nb.get("worksheets", [])) if isinstance(ws, dict) @@ -470,18 +418,16 @@ def _extract_notebook(path: str) -> str: counts = dict.fromkeys(_CELL_LABELS, 0) out: list[str] = [] for jq_pointer, cell in cells: - if not isinstance(cell, dict): - continue - typ = cell.get("cell_type") + typ = cell.get("cell_type") if isinstance(cell, dict) else None if typ not in _CELL_LABELS: continue counts[typ] += 1 suffix = f" {counts[typ]}" if typ != "raw" else "" - out.extend((f"# ── {_CELL_LABELS[typ]} cell{suffix} ──", _source_text(cell.get("source", "")).rstrip("\n"), "")) - if typ == "code": - rendered = _notebook_outputs(cell, jq_pointer, nb_name) - if rendered: - out.extend((f"# ── Output (cell {counts[typ]}) ──", rendered.rstrip("\n"), "")) + source = _source_text(cell.get("source", "")).rstrip("\n") + out += [f"# ── {_CELL_LABELS[typ]} cell{suffix} ──", source, ""] + rendered = _notebook_outputs(cell, jq_pointer, nb_name) if typ == "code" else "" + if rendered: + out += [f"# ── Output (cell {counts[typ]}) ──", rendered.rstrip("\n"), ""] return _joined(out, "Notebook contains no readable cells") @@ -491,24 +437,21 @@ def _open_zip(path: str, kind: str) -> Iterator[zipfile.ZipFile]: try: with zipfile.ZipFile(path) as zf: yield zf - except zipfile.BadZipFile as exc: - raise ExtractionError(f"Not a valid {kind}: {exc}") from exc - except OSError as exc: - raise ExtractionError(str(exc)) from exc + except (zipfile.BadZipFile, OSError) as exc: + bad_zip = isinstance(exc, zipfile.BadZipFile) + raise ExtractionError(f"Not a valid {kind}: {exc}" if bad_zip else str(exc)) from exc def _zip_xml(zf: zipfile.ZipFile, name: str, optional: bool = False) -> Any: - """Parse a package part; ``optional`` parts yield None when absent or malformed.""" + """Parse a package part; ``optional`` parts yield an empty element when absent or malformed.""" try: return ET.fromstring(zf.read(name)) - except KeyError as exc: + except (KeyError, ET.ParseError) as exc: if optional: - return None - raise ExtractionError(f"Missing {name}") from exc - except ET.ParseError as exc: - if optional: - return None - raise ExtractionError(f"Malformed XML in {name}: {exc}") from exc + return ET.Element("missing") + raise ExtractionError( + f"Missing {name}" if isinstance(exc, KeyError) else f"Malformed XML in {name}: {exc}" + ) from exc def _extract_docx(path: str) -> str: @@ -518,9 +461,9 @@ def _extract_docx(path: str) -> str: breaks = {f"{w}tab": "\t", f"{w}br": "\n", f"{w}cr": "\n"} lines: list[str] = [] for para in root.iter(f"{w}p"): - buf = [(node.text or "") if node.tag == f"{w}t" else breaks.get(node.tag, "") - for node in para.iter()] - lines.extend("".join(buf).split("\n")) + text = "".join( + (n.text or "") if n.tag == f"{w}t" else breaks.get(n.tag, "") for n in para.iter()) + lines.extend(text.split("\n")) return _joined(lines, "DOCX contains no extractable text") @@ -529,38 +472,27 @@ def _extract_xlsx(path: str) -> str: with _open_zip(path, "XLSX") as zf: names = set(zf.namelist()) sst = _zip_xml(zf, "xl/sharedStrings.xml", optional=True) - shared = [] if sst is None else [ - "".join(t.text or "" for t in item.iter(f"{s}t")) for item in sst.iter(f"{s}si")] + shared = ["".join(t.text or "" for t in item.iter(f"{s}t")) for item in sst.iter(f"{s}si")] rels_root = _zip_xml(zf, "xl/_rels/workbook.xml.rels", optional=True) - rels = {} if rels_root is None else { - rel.get("Id", ""): rel.get("Target", "") - for rel in rels_root.iter(f"{pr}Relationship") if rel.get("Id")} + rels = {rel.get("Id", ""): rel.get("Target", "") + for rel in rels_root.iter(f"{pr}Relationship") if rel.get("Id")} out: list[str] = [] for sheet in _zip_xml(zf, "xl/workbook.xml").iter(f"{s}sheet"): - if sheet.get("state", "visible") in {"hidden", "veryHidden"}: - continue target = rels.get(sheet.get(f"{r}id", ""), "").lstrip("/") part = posixpath.normpath(target if target.startswith("xl/") else f"xl/{target}") - if part not in names: + if sheet.get("state", "visible") in {"hidden", "veryHidden"} or part not in names: continue - try: + with contextlib.suppress(ET.ParseError): rows = _sheet_rows(zf.read(part), shared) - except ET.ParseError: - continue - out.append(f"# ── Sheet: {sheet.get('name', 'Sheet')} ──") - out.extend("\t".join(row) for row in rows) - if not rows: - out.append("(empty)") - out.append("") + out += [f"# ── Sheet: {sheet.get('name', 'Sheet')} ──", + *(["\t".join(row) for row in rows] or ["(empty)"]), ""] return _joined(out, "XLSX has no visible sheets with content") def _col_index(ref: str) -> int: - idx = 0 - for ch in ref: - if not ch.isalpha(): - break - idx = idx * 26 + ord(ch.upper()) - ord("A") + 1 + """0-based column of a cell ref: ``A1`` -> 0, ``AB7`` -> 27 (bijective base-26 letters).""" + idx = functools.reduce(lambda acc, ch: acc * 26 + ord(ch.upper()) - ord("A") + 1, + itertools.takewhile(str.isalpha, ref), 0) return max(idx - 1, 0) @@ -568,18 +500,15 @@ def _sheet_rows(xml_bytes: bytes, shared: list[str]) -> list[list[str]]: root = ET.fromstring(xml_bytes) s = f"{{{_NS_S}}}" rows: list[list[str]] = [] - for row in root.iter(f"{s}row"): - if len(rows) >= _MAX_XLSX_ROWS_PER_SHEET: - break + for row in itertools.islice(root.iter(f"{s}row"), _MAX_XLSX_ROWS_PER_SHEET): cells: dict[int, str] = {} max_col = -1 for cell in row.iter(f"{s}c"): col = _col_index(cell.get("r", "")) if cell.get("r") else max_col + 1 - if col >= _MAX_XLSX_COLS: - continue - cells[col] = _cell_value(cell, shared, s) - max_col = max(max_col, col) - rows.append([cells.get(i, "") for i in range(max_col + 1)] if max_col >= 0 else []) + if col < _MAX_XLSX_COLS: + cells[col] = _cell_value(cell, shared, s) + max_col = max(max_col, col) + rows.append([cells.get(i, "") for i in range(max_col + 1)]) while rows and not any(value.strip() for value in rows[-1]): rows.pop() return rows diff --git a/tools/schema_sanitizer.py b/tools/schema_sanitizer.py index 39a3d6c937..c673998f28 100644 --- a/tools/schema_sanitizer.py +++ b/tools/schema_sanitizer.py @@ -1,11 +1,7 @@ -"""Sanitize tool JSON schemas for broad LLM-backend compatibility. - -Strict backends reject shapes OpenAI/Anthropic accept: llama.cpp's grammar converter fails -on ``{"type": "object"}`` without ``properties``, bare-string schemas and ``type`` arrays; -Anthropic rejects nullable ``anyOf`` at the top of ``input_schema``; Fireworks rejects -``default`` beside ``$ref``; Codex rejects top-level combinators. This module walks the -final schema tree on a deep copy and fixes only those shapes. -""" +"""Sanitize tool JSON schemas for strict LLM backends. llama.cpp's grammar converter fails on +``{"type": "object"}`` without ``properties``, bare-string schemas and ``type`` arrays; Anthropic +rejects nullable ``anyOf`` at the top of ``input_schema``; Fireworks rejects ``default`` beside +``$ref``; Codex rejects top-level combinators. Walks a deep copy and fixes only those shapes.""" from __future__ import annotations @@ -16,15 +12,12 @@ from typing import Any, Callable logger = logging.getLogger(__name__) - -# Anthropic (and Bedrock/Vertex/Azure fronting it) reject property keys not matching this; -# one bad key anywhere in the tools array 400s the request (Cloudflare's MCP ships 61). +# Anthropic (and Bedrock/Vertex/Azure fronting it) reject property keys not matching this; one bad +# key anywhere in the tools array 400s the request (Cloudflare's MCP ships 61). _PROP_KEY_RE = re.compile(r"^[a-zA-Z0-9_.-]{1,64}$") _PROP_KEY_BAD_CHARS = re.compile(r"[^a-zA-Z0-9_.-]") - _UNION_KEYS = ("anyOf", "oneOf") -# Outer-node metadata carried onto a union's replacement node. -_UNION_META_KEYS = ("title", "description", "default", "examples") +_UNION_META_KEYS = ("title", "description", "default", "examples") # copied onto replacements def _empty_object() -> dict: @@ -47,56 +40,45 @@ def sanitize_property_key(key: str) -> str: def _rename_property_keys(props: dict, path: str) -> dict[str, str]: """{original_key: conforming_key} for one properties dict (identity entries omitted). - Deterministic (insertion order, numeric suffixes on collision) so the model-visible - schema and the dispatch-time reverse map from the registry's original schema agree.""" + Deterministic (insertion order, numeric suffixes on collision) so the model-visible schema + and the dispatch-time reverse map from the registry's original schema agree.""" renames: dict[str, str] = {} taken = {k for k in props if _PROP_KEY_RE.match(k)} - for key in props: - if _PROP_KEY_RE.match(key): - continue + for key in (k for k in props if not _PROP_KEY_RE.match(k)): base = sanitize_property_key(key) candidate, i = base, 2 while candidate in taken: - suffix = f"_{i}" - candidate = base[: 64 - len(suffix)] + suffix - i += 1 + candidate, i = base[: 64 - len(f"_{i}")] + f"_{i}", i + 1 taken.add(candidate) renames[key] = candidate - logger.debug( - "schema_sanitizer[%s]: renamed property key %r -> %r " - "(provider key-pattern compat)", path, key, candidate) + logger.debug("schema_sanitizer[%s]: renamed property key %r -> %r " + "(provider key-pattern compat)", path, key, candidate) return renames def unrename_tool_args(params_schema: Any, args: Any) -> Any: - """Map sanitized property keys in model-emitted args back to wire names. ``params_schema`` - is the ORIGINAL registry schema; recurses into objects/array items; unknown keys pass.""" - if not isinstance(params_schema, dict) or not isinstance(args, dict): - return args - props = params_schema.get("properties") - if not isinstance(props, dict): + """Map sanitized keys in model-emitted args back to wire names. ``params_schema`` is the + ORIGINAL registry schema; recurses into objects/array items; unknown keys pass through.""" + props = params_schema.get("properties") if isinstance(params_schema, dict) else None + if not isinstance(props, dict) or not isinstance(args, dict): return args reverse = {v: k for k, v in _rename_property_keys(props, "").items()} out = {} for key, value in args.items(): orig = reverse.get(key, key) - subschema = props.get(orig) - if isinstance(subschema, dict): - if isinstance(value, dict): - value = unrename_tool_args(subschema, value) - elif isinstance(value, list) and isinstance(subschema.get("items"), dict): - value = [ - unrename_tool_args(subschema["items"], item) if isinstance(item, dict) else item - for item in value] + sub = props.get(orig) if isinstance(props.get(orig), dict) else {} + if isinstance(value, dict) and sub: + value = unrename_tool_args(sub, value) + elif isinstance(value, list) and isinstance(sub.get("items"), dict): + value = [unrename_tool_args(sub["items"], item) if isinstance(item, dict) else item + for item in value] out[orig] = value return out def sanitize_tool_schemas(tools: list[dict]) -> list[dict]: """Deep-copied ``tools`` (OpenAI format) with sanitized parameter schemas; safe to mutate.""" - if not tools: - return tools - return [_sanitize_single_tool(tool) for tool in tools] + return [_sanitize_single_tool(tool) for tool in tools] if tools else tools def _sanitize_single_tool(tool: dict) -> dict: @@ -110,9 +92,7 @@ def _sanitize_single_tool(tool: dict) -> dict: return out name = fn.get("name", "") top = _sanitize_node(params, path=name) - # Guarantee the top level is an object with properties. - if not isinstance(top, dict): - top = {} + top = top if isinstance(top, dict) else {} # guarantee an object with properties on top top["type"] = "object" if not isinstance(top.get("properties"), dict): top["properties"] = {} @@ -124,16 +104,14 @@ def _sanitize_single_tool(tool: dict) -> dict: return out -# Sibling keywords strict JSON Schema validators reject alongside ``$ref``. -_REF_FORBIDDEN_SIBLINGS = frozenset({"default"}) +_REF_FORBIDDEN_SIBLINGS = frozenset({"default"}) # strict validators reject these beside ``$ref`` def _strip_ref_siblings(node: Any) -> Any: """Recursively drop forbidden siblings of ``$ref`` (Fireworks rejects ``default`` there).""" def strip(out: dict) -> dict: - if "$ref" in out: - for key in _REF_FORBIDDEN_SIBLINGS: - out.pop(key, None) + for key in _REF_FORBIDDEN_SIBLINGS if "$ref" in out else (): + out.pop(key, None) return out return _rewrite(node, strip) @@ -143,18 +121,15 @@ _TOP_LEVEL_FORBIDDEN_KEYS = ("allOf", "anyOf", "oneOf", "enum", "not") def _strip_top_level_combinators(params: dict, *, path: str = "") -> dict: """Drop combinators from the TOP level only (Codex rejects them there). They are usually - conditional-required hints, so validity is unchanged (handlers re-validate); nested - combinators are preserved.""" + conditional-required hints, so validity is unchanged (handlers re-validate); nested ones + stay.""" if not isinstance(params, dict): return params out = dict(params) - for key in _TOP_LEVEL_FORBIDDEN_KEYS: - if key in out: - logger.debug( - "schema_sanitizer[%s]: stripped top-level %r combinator " - "from tool parameters (strict-backend compat)", - path, key) - out.pop(key, None) + for key in [k for k in _TOP_LEVEL_FORBIDDEN_KEYS if k in out]: + logger.debug("schema_sanitizer[%s]: stripped top-level %r combinator " + "from tool parameters (strict-backend compat)", path, key) + del out[key] return out @@ -163,20 +138,18 @@ def _is_null_branch(item: Any) -> bool: def _carry_union_meta(outer: dict, replacement: dict, *, skip_default_on_ref: bool) -> None: - """Copy outer-union metadata onto *replacement* where absent.""" + """Copy outer-union metadata onto *replacement* where absent (``default`` is illegal beside + ``$ref`` on strict backends, hence ``skip_default_on_ref``).""" for meta_key in _UNION_META_KEYS: - if meta_key in outer and meta_key not in replacement: - # ``default`` is illegal alongside ``$ref`` on strict backends. - if skip_default_on_ref and meta_key == "default" and "$ref" in replacement: - continue + if meta_key in outer and meta_key not in replacement and not ( + skip_default_on_ref and meta_key == "default" and "$ref" in replacement): replacement[meta_key] = outer[meta_key] def strip_nullable_unions(schema: Any, *, keep_nullable_hint: bool = True) -> Any: - """Collapse ``anyOf``/``oneOf`` nullable unions to the single non-null branch. MCP/Pydantic - optional fields arrive as ``{"anyOf": [{"type": "string"}, {"type": "null"}]}``; Anthropic - rejects the null branch and optionality is already in the parent's ``required``. Only - collapses when a null branch was dropped AND exactly one non-null branch survives. + """Collapse ``anyOf``/``oneOf`` nullable unions (MCP/Pydantic optional fields) to the single + non-null branch: Anthropic rejects the null branch and optionality already lives in the parent's + ``required``. Only when a null branch was dropped AND exactly one non-null branch survives. ``keep_nullable_hint`` sets ``nullable: true`` for runtime ``"null"`` → ``None`` coercion.""" def collapse(stripped: dict) -> Any: for key in _UNION_KEYS: @@ -189,7 +162,7 @@ def strip_nullable_unions(schema: Any, *, keep_nullable_hint: bool = True) -> An if keep_nullable_hint: replacement.setdefault("nullable", True) _carry_union_meta(stripped, replacement, skip_default_on_ref=True) - return _rewrite(replacement, collapse) + return _rewrite(replacement, collapse) # the survivor may itself be a union return stripped return _rewrite(schema, collapse) @@ -199,24 +172,23 @@ _CONST_PRIMITIVE_TYPES: dict[type, str] = { def _const_branch_type(branch: Any) -> str | None: - """JSON-Schema primitive type of a pure ``const`` branch, else None: a primitive ``const`` - whose declared ``type`` (if any) matches; only ``title``/``description`` may accompany it.""" + """Primitive JSON-Schema type of a pure ``const`` branch (declared ``type``, if any, must match; + only ``title``/``description`` may accompany it), else None.""" if not isinstance(branch, dict) or "const" not in branch \ or set(branch) - {"const", "type", "title", "description"}: return None # ``type(value)`` lookup (not isinstance): bool is a subclass of int. json_type = _CONST_PRIMITIVE_TYPES.get(type(branch["const"])) - return json_type if json_type is not None and branch.get("type") in (None, json_type) else None + return json_type if branch.get("type") in (None, json_type) else None def collapse_const_unions(schema: Any) -> Any: - """Collapse ``anyOf``/``oneOf`` unions of same-typed consts to ``enum`` (ported from - block/goose ``tool_schema_normalize.rs``, Apache-2.0). Rust/TS MCP servers emit - ``{"anyOf": [{"const": "red"}, {"const": "green"}]}``, which strict backends mishandle. - Applies only when EVERY non-null branch is a pure ``const`` of one primitive type - (``bool`` never merges with ``integer``); one ``{"type": "null"}`` branch is tolerated - and recorded as ``nullable: true``. Branch order is kept; outer metadata carried over; - input never mutated.""" + """Collapse ``anyOf``/``oneOf`` unions of same-typed consts (Rust/TS MCP servers emit + ``{"anyOf": [{"const": "red"}, {"const": "green"}]}``) to ``enum``; ported from block/goose + ``tool_schema_normalize.rs`` (Apache-2.0). Only when EVERY non-null branch is a pure ``const`` + of one primitive type (``bool`` never merges with ``integer``); one ``{"type": "null"}`` branch + is tolerated as ``nullable: true``. Branch order kept; outer metadata carried; input never + mutated.""" def collapse(out: dict) -> Any: for key in _UNION_KEYS: variants = out.get(key) @@ -240,52 +212,48 @@ def collapse_const_unions(schema: Any) -> Any: _BARE_TYPE_NAMES = frozenset({"object", "string", "number", "integer", "boolean", "array", "null"}) -# Keys whose values are NOT schemas (recursing would treat "path" as a bare-string schema). +# Values that are NOT schemas (recursing would treat a required name like "path" as a bare schema). _NON_SCHEMA_LIST_KEYS = frozenset({"required", "enum", "examples", "dependentRequired"}) def _normalize_type_array(value: list, out: dict) -> None: - """Normalize a ``type: [...]`` array into *out* (llama.cpp and Gemini-via-OpenAI reject - arrays). Per AI-SDK: one non-null type → ``type: X`` (+ ``nullable`` if ``null`` present); - several → ``anyOf`` of single-type schemas so EVERY branch survives; none → ``null`` or - the object fallback. Ported from anomalyco/opencode#31877.""" + """Normalize a ``type: [...]`` array into *out* (llama.cpp and Gemini-via-OpenAI reject arrays). + Per AI-SDK: one non-null type → ``type: X`` (+ ``nullable`` if ``null`` present); several → + ``anyOf`` of single-type schemas so EVERY branch survives; none → ``null``/object fallback.""" has_null = "null" in value non_null = [t for t in value if isinstance(t, str) and t != "null"] - if len(non_null) == 1: - out["type"] = non_null[0] - elif len(non_null) >= 2: - out["anyOf"] = [{"type": t} for t in non_null] - else: + if not non_null: out["type"] = "null" if has_null else "object" return + if len(non_null) == 1: + out["type"] = non_null[0] + else: + out["anyOf"] = [{"type": t} for t in non_null] if has_null: out.setdefault("nullable", True) def _sanitize_node(node: Any, path: str) -> Any: - """Recursively sanitize a JSON-Schema fragment: bare-string schemas become ``{"type": - }`` (unknown strings → permissive object); object nodes gain ``properties: {}``; - ``type`` arrays are normalized; property keys are renamed to the provider-safe pattern - and ``required`` follows, with entries missing from ``properties`` pruned.""" + """Recursively sanitize a JSON-Schema fragment: bare-string schemas → ``{"type": }`` + (unknown strings → permissive object); object nodes gain ``properties: {}``; ``type`` arrays + are normalized; property keys are renamed to the provider-safe pattern and ``required`` + follows, with entries missing from ``properties`` pruned.""" if isinstance(node, str): if node in _BARE_TYPE_NAMES: - logger.debug( - "schema_sanitizer[%s]: replacing bare-string schema %r with {'type': %r}", - path, node, node) + logger.debug("schema_sanitizer[%s]: replacing bare-string schema %r with {'type': %r}", + path, node, node) return _empty_object() if node == "object" else {"type": node} - logger.debug( - "schema_sanitizer[%s]: replacing non-schema string %r " - "with empty object schema", path, node) + logger.debug("schema_sanitizer[%s]: replacing non-schema string %r " + "with empty object schema", path, node) return _empty_object() if isinstance(node, list): return [_sanitize_node(item, f"{path}[{i}]") for i, item in enumerate(node)] if not isinstance(node, dict): return node - # Renames computed up front so ``required`` remaps even when it precedes ``properties``. - prop_renames: dict[str, str] = {} - if isinstance(node.get("properties"), dict): - prop_renames = _rename_property_keys(node["properties"], f"{path}.properties") + props_in = node.get("properties") + prop_renames = (_rename_property_keys(props_in, f"{path}.properties") + if isinstance(props_in, dict) else {}) out: dict = {} for key, value in node.items(): if key == "type" and isinstance(value, list): @@ -293,64 +261,55 @@ def _sanitize_node(node: Any, path: str) -> Any: elif key in {"properties", "$defs", "definitions"} and isinstance(value, dict): renames = prop_renames if key == "properties" else {} out[key] = { - renames.get(sub_k, sub_k): _sanitize_node(sub_v, f"{path}.{key}.{renames.get(sub_k, sub_k)}") - for sub_k, sub_v in value.items()} + renames.get(k, k): _sanitize_node(v, f"{path}.{key}.{renames.get(k, k)}") + for k, v in value.items()} elif key in {"items", "additionalProperties"}: # Bool ``additionalProperties`` is valid; bool ``items`` is non-standard but preserved. out[key] = value if isinstance(value, bool) else _sanitize_node(value, f"{path}.{key}") - elif key in {"anyOf", "oneOf", "allOf"} and isinstance(value, list): - out[key] = [_sanitize_node(item, f"{path}.{key}[{i}]") for i, item in enumerate(value)] elif key in _NON_SCHEMA_LIST_KEYS: if key == "required" and prop_renames and isinstance(value, list): out[key] = [prop_renames.get(r, r) if isinstance(r, str) else r for r in value] else: out[key] = copy.deepcopy(value) if isinstance(value, (list, dict)) else value - else: + else: # anyOf/oneOf/allOf and any other nested schema recurse (lists index the path) out[key] = _sanitize_node(value, f"{path}.{key}") if isinstance(value, (dict, list)) else value if out.get("type") == "object": if not isinstance(out.get("properties"), dict): out["properties"] = {} if isinstance(out.get("required"), list): - props = out.get("properties") or {} - valid = [r for r in out["required"] if isinstance(r, str) and r in props] - if not valid: - out.pop("required", None) - elif len(valid) != len(out["required"]): + valid = [r for r in out["required"] if isinstance(r, str) and r in out["properties"]] + if valid: out["required"] = valid + else: + del out["required"] return out -# ---- Reactive strips — only invoked after a backend rejects a schema ----------------------- +# ---- Reactive strips — only invoked after a backend rejects a schema ---- _STRIP_ON_RECOVERY_KEYS = frozenset({"pattern", "format"}) +_SCHEMA_MARKERS = frozenset({"type", "anyOf", "oneOf", "allOf"}) # a node with one IS a schema def _dict_nodes(node: Any): - """Pre-order walk yielding every dict node; each is yielded before its values are visited, - so a consumer may mutate it in place.""" + """Pre-order walk over every dict node (yielded before its values, so it may be mutated).""" if isinstance(node, dict): yield node - for v in node.values(): - yield from _dict_nodes(v) - elif isinstance(node, list): - for item in node: - yield from _dict_nodes(item) + children = node.values() if isinstance(node, dict) else node if isinstance(node, list) else () + for child in children: + yield from _dict_nodes(child) def _reactive_strip( tools: list[dict], strip_node: Callable[[dict], int], log_msg: str) -> tuple[list[dict], int]: - """Apply *strip_node* (returns keywords removed) to every dict node of each tool's - parameters, in place. Handles OpenAI (``{"function": {"parameters": ..}}``) and Responses - (``{"name": .., "parameters": ..}``) formats. Returns ``(tools, stripped_count)``.""" - if not tools: - return tools, 0 + """Apply *strip_node* (-> keywords removed) to every dict node of each tool's parameters, in + place; OpenAI (``{"function": {"parameters"}}``) and Responses (``{"parameters"}``) formats.""" stripped = 0 - for tool in tools: + for tool in tools or (): if not isinstance(tool, dict): continue fn = tool.get("function") params = fn.get("parameters") if isinstance(fn, dict) else None - if not isinstance(params, dict): - params = tool.get("parameters") + params = params if isinstance(params, dict) else tool.get("parameters") if isinstance(params, dict): stripped += sum(strip_node(node) for node in _dict_nodes(params)) if stripped: @@ -359,16 +318,14 @@ def _reactive_strip( def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]: - """Strip ``pattern``/``format`` keywords from tool schemas, in place. Reactive: only after - llama.cpp's grammar converter rejected a schema (its regex engine is a small ECMAScript - subset), since cloud providers use these as prompting hints. Only strips beside - ``type``/combinators, so a property literally *named* ``pattern`` is untouched.""" + """Strip ``pattern``/``format`` in place — reactive, only after llama.cpp's grammar converter + rejected a schema (its regex engine is a small ECMAScript subset); cloud providers use these as + prompting hints. Only beside ``type``/combinators, so a property *named* ``pattern`` stays.""" def _strip(node: dict) -> int: - if not ("type" in node or "anyOf" in node or "oneOf" in node or "allOf" in node): - return 0 - hits = [k for k in node if k in _STRIP_ON_RECOVERY_KEYS] + is_schema = bool(node.keys() & _SCHEMA_MARKERS) + hits = [k for k in node if k in _STRIP_ON_RECOVERY_KEYS] if is_schema else [] for k in hits: - node.pop(k, None) + del node[k] return len(hits) return _reactive_strip( tools, _strip, @@ -377,13 +334,12 @@ def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]: def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]: - """Strip ``enum`` keywords whose string values contain ``/``, in place. xAI compiles - schemas to a grammar that rejects ``/`` in enum values (HTTP 400 before any token) — - typically MCP enums of HuggingFace model IDs. The constraint is a prompting hint only.""" + """Strip ``enum`` keywords whose string values contain ``/``, in place: xAI's grammar compiler + rejects them (HTTP 400 before any token) — typically MCP enums of HuggingFace model IDs.""" def _strip(node: dict) -> int: enum_val = node.get("enum") if isinstance(enum_val, list) and any(isinstance(v, str) and "/" in v for v in enum_val): - node.pop("enum", None) + del node["enum"] return 1 return 0 return _reactive_strip( diff --git a/tools/self_repo_guard.py b/tools/self_repo_guard.py index 32b72a1fba..a97bb2ffb0 100644 --- a/tools/self_repo_guard.py +++ b/tools/self_repo_guard.py @@ -2,6 +2,7 @@ from __future__ import annotations +import contextlib import os import re import shlex @@ -11,9 +12,7 @@ from pathlib import Path from typing import Callable from tools.approval import ( - _bash_exec_payload, - _deobfuscate_shell_word_for_detection, - _iter_shell_command_starts, + _bash_exec_payload, _deobfuscate_shell_word_for_detection, _iter_shell_command_starts, _read_shell_word) # bisect drives repeated checkouts of the running root — the exact skew hazard guarded here. @@ -24,8 +23,7 @@ _WORKTREE_TARGET_ACTIONS = frozenset({"move", "remove"}) _STASH_SAFE_ACTIONS = frozenset({"list", "show", "create", "store", "drop", "clear"}) _RESET_WORKTREE_MODES = frozenset({"--hard", "--merge", "--keep"}) # `reset`/`stash`/`clean`/`restore` reach this set only in their SAFE forms (_mutates_worktree -# classifies the dangerous forms first); listing them just skips a pointless -# `git config --get alias.` subprocess for `stash list`, `reset --soft`, `clean -n`. +# runs first); listing them skips a pointless `git config --get alias.` subprocess. _KNOWN_GIT_BUILTINS = frozenset({ "add", "am", "apply", "blame", "branch", "bundle", "cat-file", "clean", "clone", "commit", "config", "describe", "diff", "fetch", "format-patch", "grep", "help", "init", "log", @@ -64,36 +62,31 @@ class _Heredoc: @dataclass -class _ShellContext: +class _ShellContext: # one `(` / `$(` / backtick nesting level and its live quote state kind: str opener: int quote: str | None = None def get_running_source_root() -> Path | None: - """Return the source checkout backing this process, if there is one.""" - try: + """The source checkout backing this process, if there is one.""" + with contextlib.suppress(OSError, RuntimeError): root = Path(__file__).resolve().parent.parent - except (OSError, RuntimeError): - return None - return root if (root / ".git").exists() else None + return root if (root / ".git").exists() else None + return None def _resolve(path_str: str, base: Path) -> Path: - path = Path(os.path.expanduser(path_str)) - if not path.is_absolute(): - path = base / path - try: + path = base / Path(os.path.expanduser(path_str)) # ``/`` keeps an absolute right operand + with contextlib.suppress(OSError, RuntimeError, ValueError): return path.resolve() - except (OSError, RuntimeError, ValueError): - return path + return path def _is_within(path: Path, root: Path) -> bool: - try: + with contextlib.suppress(OSError, RuntimeError, ValueError): return path == root or path.is_relative_to(root) - except (OSError, RuntimeError, ValueError): - return False + return False def _executable_name(value: str) -> str: @@ -101,6 +94,7 @@ def _executable_name(value: str) -> str: def _shell_words_at(command: str, start: int) -> list[str]: + """Deobfuscated words of the simple command at ``start`` (stops at a newline; max 64).""" words: list[str] = [] cursor = start for _ in range(64): @@ -116,13 +110,10 @@ def _consume_options( words: list[str], start: int, options_with_arg: frozenset[str] = _NO_OPTIONS) -> int: """Index of the first positional at/after ``start`` (``--`` ends options).""" index = start - while index < len(words): - option = words[index] - if option == "--": + while index < len(words) and words[index].startswith("-") and words[index] != "-": + if words[index] == "--": return index + 1 - if not option.startswith("-") or option == "-": - break - index += 2 if "=" not in option and option in options_with_arg else 1 + index += 2 if "=" not in words[index] and words[index] in options_with_arg else 1 return index @@ -140,9 +131,8 @@ def _command_parts(words: list[str]) -> tuple[dict[str, str], str | None, list[s wrapper_options = _WRAPPER_OPTIONS_WITH_ARG.get(executable) if wrapper_options is None: return env, words[index], words[index + 1 :] - # `command -v/-V` only reports; nothing runs. if executable == "command" and words[index + 1 : index + 2] in (["-v"], ["-V"]): - return env, None, [] + break # `command -v/-V` only reports; nothing runs index = _consume_options(words, index + 1, wrapper_options) return env, None, [] @@ -157,15 +147,14 @@ def _scope_keys(command: str, starts: list[int]) -> dict[int, tuple[int, ...]]: context = contexts[-1] quote = context.quote char = command[cursor] - nested = len(contexts) > 1 - if quote == "'": - if char == "'": - context.quote = None + closes = quote is None and len(contexts) > 1 # an unquoted closer may pop a scope + if quote is not None and char == quote: + context.quote = None + elif quote == "'": + pass # single quotes: no escapes, no substitutions elif char == "\\" and cursor + 1 < start: cursor += 1 - elif quote == '"' and char == '"': - context.quote = None - elif quote is None and char in {"'", '"'}: + elif quote is None and char in "'\"": context.quote = char # Unquoted or inside double quotes: substitutions still open scopes. elif command.startswith("$(", cursor): @@ -173,28 +162,27 @@ def _scope_keys(command: str, starts: list[int]) -> dict[int, tuple[int, ...]]: cursor += 1 elif quote is None and char == "(": contexts.append(_ShellContext("(", cursor)) - elif quote is None and char == ")" and nested and contexts[-1].kind in {"(", "$("}: + elif (char == ")" and closes and context.kind in {"(", "$("}) or ( + char == "`" and closes and context.kind == "`"): contexts.pop() elif char == "`": - if quote is None and nested and contexts[-1].kind == "`": - contexts.pop() - else: - contexts.append(_ShellContext("`", cursor)) + contexts.append(_ShellContext("`", cursor)) cursor += 1 scopes[start] = tuple(item.opener for item in contexts[1:]) return scopes def _operator_before(command: str, start: int) -> str | None: + """The list/grouping operator (or newline) immediately preceding a command start.""" head = command[:start].rstrip() - if head[-2:] in {"&&", "||"}: - return head[-2:] - if head[-1:] in {";", "|", "&", "(", "{"}: - return head[-1:] + for tail in (head[-2:], head[-1:]): + if tail in {"&&", "||", ";", "|", "&", "(", "{"}: + return tail return "\n" if "\n" in command[len(head):start] else None def _cd_target(executable: str, args: list[str], cwd: Path) -> Path | None: + """Directory a ``cd``/``pushd`` would land in (existing dirs only), else None.""" if _executable_name(executable) not in {"cd", "pushd"}: return None index = _consume_options(args, 0) @@ -205,17 +193,15 @@ def _cd_target(executable: str, args: list[str], cwd: Path) -> Path | None: def _shell_script_arg(args: list[str]) -> str | None: - """Return the script string owned by a shell's ``-c``, if present. approval.py's - ``_bash_exec_payload`` parses bash's real option grammar (``-o pipefail -c '