diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index c0141159a8..9f749de74e 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -1,12 +1,9 @@ -"""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 @@ -53,34 +50,30 @@ 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): """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: @@ -99,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 @@ -119,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): @@ -160,8 +151,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,29 +161,24 @@ 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 - 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): - 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 +196,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,12 +234,11 @@ 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() - if retry_after_concurrent_auth: yield request self._persist_oauth_metadata_if_changed() @@ -274,13 +255,12 @@ 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]: - """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 +286,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: @@ -322,8 +301,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) @@ -343,23 +321,20 @@ 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: 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) @@ -368,37 +343,33 @@ 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.""" 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): + # Disk changed (external refresh)? Else: if the SDK can refresh in place, let the caller retry. + 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 - 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) + can_refresh = False finally: 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 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 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 2c5c39d32f..57407eadea 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -20,27 +20,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 @@ -49,8 +47,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 @@ -59,16 +56,14 @@ 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", "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"\$\{([^}]+)\}") @@ -80,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: @@ -94,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: @@ -120,8 +109,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 @@ -137,28 +125,20 @@ 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"))) - 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]: - """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: path_arg = resolved_env.get("PATH") which_hit = shutil.which(resolved_command, path=path_arg) @@ -168,7 +148,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) @@ -269,15 +248,14 @@ 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() 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: @@ -289,18 +267,14 @@ 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) 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 @@ -310,20 +284,18 @@ 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 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() @@ -331,15 +303,14 @@ 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) 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 91ff039518..eb0aa5f3c2 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -14,29 +14,23 @@ 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 -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,13 +38,12 @@ 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")) 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): @@ -93,11 +86,10 @@ 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}") - if not isinstance(url, str): raise _bad(f"expected a string, got {type(url).__name__}") stripped = url.strip() @@ -133,13 +125,11 @@ 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 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]")) @@ -161,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") @@ -210,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 @@ -236,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])) @@ -248,11 +235,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 +244,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). @@ -299,18 +279,17 @@ _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 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] @@ -325,10 +304,9 @@ 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", ())) - 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 af16b4ed65..a799aa8b8f 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 @@ -21,18 +20,27 @@ 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 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 - # 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 " @@ -56,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): @@ -73,20 +78,16 @@ 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))): + 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 _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 -# ------------------------------------------------------------ breaker bookkeeping - def _result_is_error(result) -> bool: """True only for a JSON payload carrying an ``error`` key (non-JSON = success).""" try: @@ -97,10 +98,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 @@ -110,18 +108,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): @@ -138,8 +135,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 @@ -147,36 +142,28 @@ 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 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) 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 " @@ -206,27 +193,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( @@ -234,13 +211,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 @@ -255,14 +237,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)}")) -# ------------------------------------------------------------- the RPC itself - -def _mark_server_call_started(server: Any) -> None: - """Record a user-visible MCP operation when the server supports it.""" - if callable(getattr(server, "mark_tool_call", None)): - server.mark_tool_call() - - @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 @@ -276,8 +250,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: @@ -286,16 +260,15 @@ 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") _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) @@ -340,9 +313,8 @@ def _render_content_blocks(result, server_name: str) -> Tuple[str, int]: parts.append(rendered) usable_parts += 1 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) @@ -380,8 +352,7 @@ def _render_call_tool_result(result, server_name: str) -> str: structured = None # drop notices do not count as usable content 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 @@ -390,20 +361,16 @@ 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) -# ------------------------------------------------------------------- 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}" 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 @@ -412,67 +379,57 @@ 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) 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 _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(): - _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 _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 + 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) + 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 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: - 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 @@ -495,8 +452,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)} @@ -523,41 +479,26 @@ 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.""" - + """Connection-alive check; lazy (schema-cache registered) servers count as available.""" 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..4a8af5753a 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 @@ -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"), @@ -52,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.""" @@ -67,15 +65,13 @@ 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) 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) @@ -95,8 +91,7 @@ class MCPServerHealthMixin: 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): @@ -122,20 +117,16 @@ 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 - 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 - 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: @@ -165,6 +156,8 @@ 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 +173,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 +182,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,12 +197,10 @@ 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 - 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) @@ -253,8 +244,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() @@ -272,7 +262,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 b9c994df5a..17282af067 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -1,11 +1,11 @@ """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 +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 ( @@ -27,8 +27,7 @@ _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() @@ -40,25 +39,19 @@ def _normalize_server_trust(value: Any) -> str: 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) - 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: - """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, {}) - 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,8 +73,7 @@ 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] @@ -93,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"]) @@ -105,8 +96,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") @@ -118,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") @@ -130,30 +119,18 @@ 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 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 @@ -167,8 +144,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): @@ -187,8 +164,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 +174,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,29 +196,14 @@ 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)) 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) - return - if 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 @@ -253,34 +214,40 @@ 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, 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 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 = [{ "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)] @@ -313,7 +280,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 b8b702fb69..8cda945d47 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"}}}') @@ -28,15 +27,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 @@ -46,16 +47,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: @@ -63,12 +65,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) @@ -80,7 +81,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, @@ -94,14 +94,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: @@ -118,8 +117,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: @@ -131,7 +129,7 @@ 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: @@ -171,8 +169,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 @@ -184,8 +181,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(): @@ -201,9 +197,8 @@ class MCPServerTransportMixin: 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() @@ -212,21 +207,18 @@ class MCPServerTransportMixin: 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. + # 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) @@ -234,29 +226,30 @@ 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); 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), - **({"cert": client_cert} if client_cert is not None else {})) as client: - # HEAD is cheapest; fall back to GET on 405/501. - resp = await client.head(url, headers=probe_headers) + **_present(cert=client_cert)) as client: + 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. - 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", @@ -265,25 +258,20 @@ 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 - raise NonMcpEndpointError( - f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP " + 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 " "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/).") 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).""" + """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 @@ -295,8 +283,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: @@ -309,28 +296,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, @@ -339,30 +323,27 @@ 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 - return _owned_client_streams() async def _run_http(self, config: dict): @@ -375,13 +356,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) @@ -399,8 +378,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(): @@ -415,19 +393,17 @@ 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) diff --git a/tools/memory_tool.py b/tools/memory_tool.py index 57f97fa81a..8e42f22e57 100644 --- a/tools/memory_tool.py +++ b/tools/memory_tool.py @@ -23,7 +23,7 @@ try: except ImportError: fcntl = None try: - import msvcrt + import msvcrt # noqa: F401 except ImportError: pass @@ -32,8 +32,7 @@ logger = logging.getLogger(__name__) # One tool-definition pass must use ONE config decision for availability and the # dynamic target schema: the check_fn result flows to the immediately following # dynamic_schema_overrides call; ContextVar isolates concurrent profile builds. -_memory_surface_flags: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar( - "memory_surface_flags", default=None) +_memory_surface_flags: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar("memory_surface_flags", default=None) def get_memory_dir() -> Path: @@ -46,26 +45,22 @@ from tools.memory_tool_store import ( # noqa: E402,F401 (re-exports) def load_on_disk_store() -> "MemoryStore": - """Fresh on-disk MemoryStore with configured limits/flags for contexts with no - live agent (gateway, Desktop, ``/memory``) so approvals enforce the SAME caps - as ``agent_init``. Falls back to defaults if config can't load; never raises.""" + """Fresh on-disk MemoryStore with configured limits/flags for contexts with no live + agent (gateway, Desktop, ``/memory``) so approvals enforce the SAME caps as + ``agent_init``. Falls back to defaults if config can't load; never raises.""" try: from hermes_cli.config import load_config config = load_config() or {} mem_cfg = get_builtin_memory_config(config) memory_enabled, user_profile_enabled = get_builtin_memory_store_flags(config) - kwargs = {"memory_char_limit": int(mem_cfg.get("memory_char_limit", 2200)), - "user_char_limit": int(mem_cfg.get("user_char_limit", 1375)), - "memory_enabled": memory_enabled, "user_profile_enabled": user_profile_enabled} + store = MemoryStore(int(mem_cfg.get("memory_char_limit", 2200)), int(mem_cfg.get("user_char_limit", 1375)), + memory_enabled=memory_enabled, user_profile_enabled=user_profile_enabled) except Exception: - kwargs: Dict[str, Any] = {} # config optional — fall back to defaults rather than break /memory - store = MemoryStore(**kwargs) + store = MemoryStore() # config optional — fall back to defaults rather than break /memory store.load_from_disk() return store -# -- Write-approval gate -- - def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Optional[str]: """JSON tool-result string when the write must NOT proceed (blocked or staged for approval), None to proceed. Fails open if the gate module can't load.""" @@ -93,35 +88,30 @@ _STORE_ACTIONS = { lambda label, content, old_text: (f"remove from {label}", old_text or ""))} -def _apply_write_gate(action: str, target: str, content: Optional[str], old_text: Optional[str]) -> Optional[str]: - """Gate a single mutating op (add/replace/remove).""" - summary, detail = _STORE_ACTIONS[action][1]("user profile" if target == "user" else "memory", content, old_text) - return _gate_or_stage(summary, detail, +def _batch_op_line(op: Dict[str, Any]) -> str: + op = op or {} + act, content, old = op.get("action", "?"), op.get("content") or op.get("new_text") or "", op.get("old_text", "") + if act == "remove": + return f"- remove: {old}" + return f"- replace: {old} -> {content}" if act == "replace" else f"- {act}: {content}" + + +def _apply_write_gate(action: str, target: str, content: Optional[str], old_text: Optional[str], + operations: Optional[List[Dict[str, Any]]] = None) -> Optional[str]: + """Gate one mutating op, or (``operations`` set) a whole batch as a single unit.""" + label = "user profile" if target == "user" else "memory" + if operations is not None: + return _gate_or_stage(f"apply {len(operations)} op(s) to {label}", + "\n".join(_batch_op_line(op) for op in operations), + {"action": "batch", "target": target, "operations": operations}) + return _gate_or_stage(*_STORE_ACTIONS[action][1](label, content, old_text), {"action": action, "target": target, "content": content, "old_text": old_text}) -def _apply_batch_write_gate(target: str, operations: List[Dict[str, Any]]) -> Optional[str]: - """Gate a whole batch as a single unit.""" - summary = f"apply {len(operations)} op(s) to {'user profile' if target == 'user' else 'memory'}" - detail_lines = [] - for op in operations: - op = op or {} - act = op.get("action", "?") - content = op.get("content") or op.get("new_text") or "" - detail_lines.append(f"- remove: {op.get('old_text', '')}" if act == "remove" - else f"- replace: {op.get('old_text', '')} -> {content}" if act == "replace" - else f"- {act}: {content}") - return _gate_or_stage(summary, "\n".join(detail_lines), - {"action": "batch", "target": target, "operations": operations}) - - -# -- Tool entry point -- - def _validate_single_op(store, action, target, content, old_text) -> Optional[str]: - """Validate BEFORE the gate so an invalid write is rejected now, not at approve - time. Missing ``old_text`` is recoverable (it can't be schema-required — needs a - combinator the Codex backend rejects — and some clients omit it): return the - current inventory plus a retry instruction instead of a dead-end.""" + """Validate BEFORE the gate so an invalid write is rejected now, not at approve time. + Missing ``old_text`` is recoverable (it can't be schema-required — needs a combinator + the Codex backend rejects): return the inventory plus a retry instruction.""" if action == "add" and not content: return tool_error("Content is required for 'add' action.", success=False) if action in ("replace", "remove") and not old_text: @@ -154,26 +144,23 @@ def memory_tool(action: str = None, target: str = "memory", content: str = None, if operations: if not isinstance(operations, list): return tool_error("operations must be a list of {action, content?, old_text?} objects.", success=False) - gate_result = _apply_batch_write_gate(target, operations) + # Approval gate: stages (background/gateway) or prompts inline (CLI); off by default. + gate_result = _apply_write_gate("batch", target, None, None, operations) if gate_result is not None: return gate_result return json.dumps(store.apply_batch(target, operations), ensure_ascii=False) if action not in _STORE_ACTIONS: return tool_error(f"Unknown action '{action}'. Use: add, replace, remove", success=False) - invalid = _validate_single_op(store, action, target, content, old_text) + invalid = (_validate_single_op(store, action, target, content, old_text) + or _apply_write_gate(action, target, content, old_text)) if invalid is not None: return invalid - # Approval gate: stages (background/gateway) or prompts inline (CLI); off by default. - gate_result = _apply_write_gate(action, target, content, old_text) - if gate_result is not None: - return gate_result return json.dumps(_STORE_ACTIONS[action][0](store, target, content, old_text), ensure_ascii=False) def get_builtin_memory_config(config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]: - """Normalized ``memory`` config section ({} when missing/malformed → flags default - to enabled). ``agent_init`` reads the same section so availability and store - construction cannot diverge.""" + """Normalized ``memory`` config section ({} when missing/malformed → flags default to + enabled). ``agent_init`` reads the same section so availability and store cannot diverge.""" if config is None: try: from hermes_cli.config import load_config_readonly @@ -214,8 +201,7 @@ def _memory_target_error(store: "MemoryStore", target: str) -> Optional[Dict[str def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[str, Any]: """Replay a staged write against the store, bypassing the gate (/memory approve).""" - action = payload.get("action") - target = payload.get("target", "memory") + action, target = payload.get("action"), payload.get("target", "memory") target_error = _memory_target_error(store, target) if target_error is not None: return target_error @@ -226,8 +212,6 @@ def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[ return _STORE_ACTIONS[action][0](store, target, payload.get("content") or "", payload.get("old_text") or "") -# -- OpenAI Function-Calling Schema -- - MEMORY_SCHEMA = { "name": "memory", "description": ( @@ -318,21 +302,17 @@ def _build_memory_schema_overrides() -> Dict[str, Any]: _memory_surface_flags.set(None) targets = [t for t, on in zip(("memory", "user"), flags) if on] parameters = copy.deepcopy(MEMORY_SCHEMA["parameters"]) - target_schema = parameters["properties"]["target"] + target_schema, description = parameters["properties"]["target"], MEMORY_SCHEMA["description"] target_schema["enum"] = targets - description = MEMORY_SCHEMA["description"] - narrowed = _SINGLE_TARGET_TEXT.get(tuple(targets)) - if narrowed: + if narrowed := _SINGLE_TARGET_TEXT.get(tuple(targets)): target_schema["description"], replacement = narrowed description = description.replace( "TARGETS: 'user' = who the user is (name, role, preferences, style). 'memory' = your " - "notes (environment, conventions, tool quirks, lessons).", - replacement) + "notes (environment, conventions, tool quirks, lessons).", replacement) return {"description": description, "parameters": parameters} -# --- Registry --- -from tools.registry import registry, tool_error +from tools.registry import registry, tool_error # noqa: E402 (registration at import time) registry.register( name="memory", diff --git a/tools/memory_tool_store.py b/tools/memory_tool_store.py index 1ee1408358..6abd008dd0 100644 --- a/tools/memory_tool_store.py +++ b/tools/memory_tool_store.py @@ -5,7 +5,7 @@ in ``tools.memory_tool`` and is read lazily.""" import logging import time -from contextlib import contextmanager +from contextlib import contextmanager, suppress from pathlib import Path from typing import Any, Dict, List, Optional, Tuple @@ -22,37 +22,36 @@ MEMORY_BLOCK_HEADERS = { ENTRY_DELIMITER = "\n§\n" -def _memory_dir() -> Path: - from tools import memory_tool - return memory_tool.get_memory_dir() - - def _scan_memory_content(content: str) -> Optional[str]: """Error string if *content* matches injection/exfil patterns. Strict scope: memory enters the system prompt, so a poisoned entry persists across sessions.""" return _first_threat_message(content, scope="strict") -def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]: +def _error(message: str, **extra) -> Dict[str, Any]: + return {"success": False, "error": message, **extra} + + +def _drift_error(path: Path, bak_path: str) -> Dict[str, Any]: """External drift: the file wouldn't round-trip, so flushing would discard content.""" - return {"success": False, "error": ( + return _error(( f"Refusing to write {path.name}: file on disk has content that wouldn't round-trip " f"through the memory tool (likely added by the patch tool, a shell append, a manual edit, " f"or a concurrent session). A snapshot was saved to {bak_path}. Resolve the drift first — " f"either rewrite the file as a clean §-delimited list of entries, or move the extra " f"content out — then retry. This guard exists to prevent silent data loss (issue #26045)." - ), "drift_backup": bak_path, "remediation": ( + ), drift_backup=bak_path, remediation=( "Open the .bak file, integrate the missing entries into the memory tool one at a time via " - "memory(action=add, content=...), then remove or rewrite the original file to a clean state.")} + "memory(action=add, content=...), then remove or rewrite the original file to a clean state.")) -def _read_failed_error(path: "Path") -> Dict[str, Any]: +def _read_failed_error(path: Path) -> Dict[str, Any]: """Existing-but-unreadable file: saving from an assumed-empty view would wipe it.""" - return {"success": False, "error": ( + return _error( f"Refusing to write {path.name}: the file exists on disk but could not be read right now " f"(temporarily locked by another program, a permission change, invalid/corrupt text encoding, " f"or a filesystem error). Treating an unreadable file as empty and saving would wipe existing " - f"memory, so the write is refused. Nothing was changed — retry in a moment.")} + f"memory, so the write is refused. Nothing was changed — retry in a moment.") def _find_unique_match(entries: List[str], old_text: str) -> Tuple[Optional[int], bool]: @@ -78,19 +77,16 @@ class MemoryStore: memory_enabled: bool = True, user_profile_enabled: bool = True): self.memory_entries: List[str] = [] self.user_entries: List[str] = [] - self.memory_char_limit = memory_char_limit - self.user_char_limit = user_char_limit - self.memory_enabled = memory_enabled - self.user_profile_enabled = user_profile_enabled + self.memory_char_limit, self.user_char_limit = memory_char_limit, user_char_limit + self.memory_enabled, self.user_profile_enabled = memory_enabled, user_profile_enabled self._system_prompt_snapshot: Dict[str, str] = {"memory": "", "user": ""} self._consolidation_failures = 0 # per turn; reset by reset_consolidation_failures() def target_enabled(self, target: str) -> bool: - """Return whether this session's selected built-in store is writable.""" return self.user_profile_enabled if target == "user" else self.memory_enabled def reset_consolidation_failures(self) -> None: - """Reset the per-turn consolidation-failure counter (call at turn start).""" + """Call at turn start.""" self._consolidation_failures = 0 def _consolidation_failure(self, response: Dict[str, Any]) -> Dict[str, Any]: @@ -109,22 +105,10 @@ class MemoryStore: Threat hits are replaced by a ``[BLOCKED: …]`` placeholder in the SNAPSHOT only; live lists keep the raw text so the user can see and remove poisoned entries (dropping them silently would hide the attack).""" - mem_dir = _memory_dir() - mem_dir.mkdir(parents=True, exist_ok=True) - # Deduplicate (order-preserving, first occurrence wins). - self.memory_entries = list(dict.fromkeys(self._read_file(mem_dir / "MEMORY.md"))) - self.user_entries = list(dict.fromkeys(self._read_file(mem_dir / "USER.md"))) - self._system_prompt_snapshot = { - "memory": self._render_block("memory", self._sanitize_entries_for_snapshot(self.memory_entries, "MEMORY.md")), - "user": self._render_block("user", self._sanitize_entries_for_snapshot(self.user_entries, "USER.md"))} - - @staticmethod - def _sanitize_entries_for_snapshot(entries: List[str], filename: str) -> List[str]: - """*entries* with threat matches replaced by a ``[BLOCKED: …]`` placeholder - (strict scope, same as writes); empty / already-blocked entries pass through.""" from tools.threat_patterns import scan_for_threats - def _one(entry): + def _sanitize(entry, filename): + # Strict scope, same as writes; empty / already-blocked entries pass through. findings = scan_for_threats(entry, scope="strict") if entry and not entry.startswith("[BLOCKED:") else None if not findings: return entry @@ -132,7 +116,13 @@ class MemoryStore: return (f"[BLOCKED: {filename} entry contained threat pattern(s): {', '.join(findings)}. " f"Removed from system prompt; use memory(action=remove) to delete the original.]") - return [_one(e) for e in entries] + for target in ("memory", "user"): + path = self._path_for(target) + path.parent.mkdir(parents=True, exist_ok=True) + # Deduplicate (order-preserving, first occurrence wins). + entries = list(dict.fromkeys(self._read_file(path))) + self._set_entries(target, entries) + self._system_prompt_snapshot[target] = self._render_block(target, [_sanitize(e, path.name) for e in entries]) @staticmethod @contextmanager @@ -157,33 +147,13 @@ class MemoryStore: try: yield finally: - try: + with suppress(OSError): _flock(True) - except OSError: - pass @staticmethod def _path_for(target: str) -> Path: - return _memory_dir() / ("USER.md" if target == "user" else "MEMORY.md") - - def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]: - """Re-read entries from disk (under lock) before mutating; return the abort - error dict or None. Aborts on external drift (flushing would discard - un-roundtrippable content) and on an existing-but-unreadable file (even - append-only ``add`` rewrites the whole file). Drift check and parse use the - SAME raw snapshot — a failed second read used to count as "no drift".""" - path = self._path_for(target) - raw, read_ok = self._read_raw_checked(path) - if not read_ok: - return _read_failed_error(path) - bak = None if skip_drift else self._detect_external_drift(target, raw) - self._set_entries(target, list(dict.fromkeys(self._parse_entries(raw)))) - return _drift_error(path, bak) if bak else None - - def save_to_disk(self, target: str): - """Persist entries to the appropriate file. Called after every mutation.""" - _memory_dir().mkdir(parents=True, exist_ok=True) - self._write_file(self._path_for(target), self._entries_for(target)) + from tools import memory_tool # get_memory_dir is monkeypatched there + return memory_tool.get_memory_dir() / ("USER.md" if target == "user" else "MEMORY.md") def _entries_for(self, target: str) -> List[str]: return self.user_entries if target == "user" else self.memory_entries @@ -201,53 +171,45 @@ class MemoryStore: return f"{self._char_count(target):,}/{self._char_limit(target):,}" def _usage_pct(self, target: str, current: int) -> str: - """``"% — / chars"`` for the given target.""" limit = self._char_limit(target) - pct = min(100, int((current / limit) * 100)) if limit > 0 else 0 - return f"{pct}% — {current:,}/{limit:,} chars" + return f"{min(100, int((current / limit) * 100)) if limit > 0 else 0}% — {current:,}/{limit:,} chars" def _failure_with_entries(self, target: str, message: str) -> Dict[str, Any]: """Consolidation failure carrying the live entries so the model can consolidate.""" - return self._consolidation_failure({"success": False, "error": message, - "current_entries": self._entries_for(target), "usage": self._usage(target)}) - - def _locate(self, target: str, old_text: str, verb: str): - """Resolve *old_text* to a unique entry index, or an error dict.""" - entries = self._entries_for(target) - idx, ambiguous = _find_unique_match(entries, old_text) - if ambiguous: - return None, {"success": False, "error": f"Multiple entries matched '{old_text}'. Be more specific.", - "matches": [e[:80] + ("..." if len(e) > 80 else "") for e in entries if old_text in e]} - if idx is None: - return None, self._consolidation_failure({ - "success": False, - "error": f"No entry matched '{old_text}'. Check current_entries below and retry with the exact text of the entry you want to {verb}.", - "current_entries": entries}) - return idx, None + return self._consolidation_failure( + _error(message, current_entries=self._entries_for(target), usage=self._usage(target))) def _mutate(self, target: str, mutate, *, skip_drift: bool = False) -> Dict[str, Any]: - """Lock, reload, run ``mutate(entries, limit)`` -> ``(new_entries, message)`` or an - error dict, then persist and return the success response.""" - with self._file_lock(self._path_for(target)): - err = self._reload_or_error(target, skip_drift=skip_drift) - if err: - return err + """Lock, re-read from disk, run ``mutate(entries, limit)`` -> ``(new_entries, message)`` + or an error dict, then persist and return the success response. The reload aborts + on an existing-but-unreadable file (even append-only ``add`` rewrites the whole + file) and, unless *skip_drift*, on external drift (flushing would discard + un-roundtrippable content). Drift check and parse use the SAME raw snapshot — + a failed second read used to count as "no drift".""" + path = self._path_for(target) + with self._file_lock(path): + raw, read_ok = self._read_raw_checked(path) + if not read_ok: + return _read_failed_error(path) + bak = None if skip_drift else self._detect_external_drift(target, raw) + self._set_entries(target, list(dict.fromkeys(self._parse_entries(raw)))) + if bak: + return _drift_error(path, bak) result = mutate(self._entries_for(target), self._char_limit(target)) if isinstance(result, dict): return result - entries, message = result - self._set_entries(target, entries) - self.save_to_disk(target) - return self._success_response(target, message) + self._set_entries(target, result[0]) + path.parent.mkdir(parents=True, exist_ok=True) + self._write_file(path, result[0]) + return self._success_response(target, result[1]) def add(self, target: str, content: str) -> Dict[str, Any]: """Append a new entry. Returns error if it would exceed the char limit.""" content = content.strip() if not content: - return {"success": False, "error": "Content cannot be empty."} - scan_error = _scan_memory_content(content) - if scan_error: - return {"success": False, "error": scan_error} + return _error("Content cannot be empty.") + if scan_error := _scan_memory_content(content): + return _error(scan_error) def _add(entries, limit): if content in entries: @@ -259,7 +221,6 @@ class MemoryStore: f"overlapping entries into shorter ones or 'remove' stale or less important entries (see " f"current_entries below), then retry this add — all in this turn.")) return entries + [content], "Entry added." - # Append-only: skip the drift guard (appending never clobbers foreign # content) but still refuse a failed read — add rewrites the WHOLE file. return self._mutate(target, _add, skip_drift=True) @@ -268,29 +229,33 @@ class MemoryStore: """Find entry containing old_text substring, replace it with new_content.""" new_content = new_content.strip() if not old_text.strip(): - return {"success": False, "error": "old_text cannot be empty."} + return _error("old_text cannot be empty.") if not new_content: - return {"success": False, "error": "new_content cannot be empty. Use 'remove' to delete entries."} - scan_error = _scan_memory_content(new_content) - if scan_error: - return {"success": False, "error": scan_error} + return _error("new_content cannot be empty. Use 'remove' to delete entries.") + if scan_error := _scan_memory_content(new_content): + return _error(scan_error) return self._edit(target, old_text.strip(), new_content) def remove(self, target: str, old_text: str) -> Dict[str, Any]: """Remove the entry containing old_text substring.""" if not old_text.strip(): - return {"success": False, "error": "old_text cannot be empty."} + return _error("old_text cannot be empty.") return self._edit(target, old_text.strip(), None) def _edit(self, target: str, old_text: str, new_content: Optional[str]) -> Dict[str, Any]: - """Locked replace (``new_content`` set) or remove (None) of the unique entry matching *old_text*.""" + """Locked replace (``new_content`` set) or remove (None) of the entry matching *old_text*.""" def _apply(entries, limit): - idx, err = self._locate(target, old_text, "replace" if new_content else "remove") - if err: - return err + idx, ambiguous = _find_unique_match(entries, old_text) + if ambiguous: + return _error(f"Multiple entries matched '{old_text}'. Be more specific.", + matches=[e[:80] + ("..." if len(e) > 80 else "") for e in entries if old_text in e]) + if idx is None: + return self._consolidation_failure(_error( + f"No entry matched '{old_text}'. Check current_entries below and retry with the exact text " + f"of the entry you want to {'replace' if new_content else 'remove'}.", current_entries=entries)) + replaced = entries[:idx] + ([] if new_content is None else [new_content]) + entries[idx + 1:] if new_content is None: - return entries[:idx] + entries[idx + 1:], "Entry removed." - replaced = entries[:idx] + [new_content] + entries[idx + 1:] + return replaced, "Entry removed." new_total = len(ENTRY_DELIMITER.join(replaced)) if new_total > limit: return self._failure_with_entries(target, ( @@ -298,7 +263,6 @@ class MemoryStore: f"or 'remove' other stale or less important entries to make room (see current_entries " f"below), then retry — all in this turn.")) return replaced, "Entry replaced." - return self._mutate(target, _apply) @staticmethod @@ -321,79 +285,67 @@ class MemoryStore: return f"{pos}: '{old_text}' matched multiple distinct entries -- be more specific." if idx is None: return f"{pos}: no entry matched '{old_text}'." - if act == "replace": - working[idx] = content - else: - working.pop(idx) + working[idx:idx + 1] = [content] if act == "replace" else [] return None def apply_batch(self, target: str, operations: List[Dict[str, Any]]) -> Dict[str, Any]: - """Apply add/replace/remove ops to one target atomically against the FINAL - budget, so one call can free space and add entries. All-or-nothing: any - malformed / unmatched op or an over-limit result writes NOTHING and returns - the first failure plus live state.""" + """Apply add/replace/remove ops atomically against the FINAL budget, so one call + can free space and add entries. All-or-nothing: any malformed / unmatched op or + an over-limit result writes NOTHING and returns the first failure plus live state.""" if not operations: - return {"success": False, "error": "operations list is empty."} - + return _error("operations list is empty.") ops = [op or {} for op in operations] # Scan every add/replace content BEFORE touching disk -- one poisoned op rejects the batch. for i, op in enumerate(ops): scan_error = op.get("action") in {"add", "replace"} and op.get("content") and _scan_memory_content(op["content"]) if scan_error: - return {"success": False, "error": f"Operation {i + 1}: {scan_error}"} + return _error(f"Operation {i + 1}: {scan_error}") def _apply(entries, limit): working = list(entries) # only committed if the whole batch validates for i, op in enumerate(ops): act = op.get("action") - content = (op.get("content") or op.get("new_text") or "").strip() - old_text = (op.get("old_text") or "").strip() - pos = f"Operation {i + 1} ({act or 'unknown'})" - msg = self._apply_batch_op(working, act, content, old_text, pos) + msg = self._apply_batch_op(working, act, (op.get("content") or op.get("new_text") or "").strip(), + (op.get("old_text") or "").strip(), f"Operation {i + 1} ({act or 'unknown'})") if msg: - return self._failure_with_entries( - target, msg + " No operations were applied (batch is all-or-nothing).") - # Budget check against the FINAL state only. - new_total = len(ENTRY_DELIMITER.join(working)) + return self._failure_with_entries(target, msg + " No operations were applied (batch is all-or-nothing).") + new_total = len(ENTRY_DELIMITER.join(working)) # budget check against the FINAL state only if new_total > limit: return self._failure_with_entries(target, ( f"After applying all {len(operations)} operations, memory would be at " f"{new_total:,}/{limit:,} chars -- over the limit. Remove or shorten more " f"entries in the same batch (see current_entries below), then retry.")) return working, f"Applied {len(operations)} operation(s)." - return self._mutate(target, _apply) def format_for_system_prompt(self, target: str) -> Optional[str]: - """Frozen load-time snapshot for the system prompt (NOT live state — mid-session - writes don't touch it, preserving the prefix cache); None if empty.""" + """Frozen load-time snapshot (NOT live state — mid-session writes don't touch + it, preserving the prefix cache); None if empty.""" return self._system_prompt_snapshot.get(target, "") or None def _success_response(self, target: str, message: str = None) -> Dict[str, Any]: - # A successful write is progress: reset the per-turn (consecutive) failure budget. + """TERMINAL and WITHOUT the entries list: echoing entries invites the model to + "find more to fix" and re-issue the same ops. A successful write resets the + per-turn failure budget.""" self._consolidation_failures = 0 - # TERMINAL and WITHOUT the entries list: echoing entries invites the model to - # "find more to fix" and re-issue the same ops. Entries only appear on errors. return {"success": True, "done": True, "target": target, "usage": self._usage_pct(target, self._char_count(target)), "entry_count": len(self._entries_for(target)), **({"message": message} if message else {}), "note": "Write saved. This update is complete — do not repeat it."} def _render_block(self, target: str, entries: List[str]) -> str: - """Render a system prompt block with header and usage indicator.""" + """System prompt block: header + usage indicator + entries ("" when empty).""" if not entries: return "" - content = ENTRY_DELIMITER.join(entries) + content, sep = ENTRY_DELIMITER.join(entries), "═" * 46 title = MEMORY_BLOCK_HEADERS["user" if target == "user" else "memory"] - separator = "═" * 46 - return f"{separator}\n{title} [{self._usage_pct(target, len(content))}]\n{separator}\n{content}" + return f"{sep}\n{title} [{self._usage_pct(target, len(content))}]\n{sep}\n{content}" @staticmethod def _read_raw_checked(path: Path) -> Tuple[str, bool]: - """``(raw, read_ok)``; ``read_ok`` is False ONLY when the file EXISTS but can't - be read (absent → ``("", True)``). Decoding stays STRICT: ``errors="replace"`` - would hand callers a lossy view that a save then persists. ``utf-8-sig`` strips - a Notepad BOM that otherwise glues U+FEFF onto the first entry forever.""" + """``(raw, read_ok)``; ``read_ok`` is False ONLY when the file EXISTS but can't be + read. Decoding stays STRICT (``errors="replace"`` would hand callers a lossy view + a save then persists); ``utf-8-sig`` strips a Notepad BOM off the first entry.""" if not path.exists(): return "", True try: @@ -408,19 +360,26 @@ class MemoryStore: @staticmethod def _read_file(path: Path) -> List[str]: - """Entries of a memory file ([] on any error). Read-only callers only - (``load_from_disk``, learning_mutations); mutation paths must use - ``_read_raw_checked`` so they can refuse to overwrite an unreadable file.""" + """Entries of a memory file ([] on any error). Read-only callers only; mutation + paths use ``_read_raw_checked`` so they can refuse to overwrite an unreadable file.""" return MemoryStore._parse_entries(MemoryStore._read_raw_checked(path)[0]) + @staticmethod + def _write_file(path: Path, entries: List[str]): + """Atomic temp-file + rename: readers never see a truncated file. Also used by + agent/learning_mutations.py.""" + try: + atomic_write_text(path, ENTRY_DELIMITER.join(entries), tmp_prefix=".mem_") + except OSError as e: + raise RuntimeError(f"Failed to write memory file {path}: {e}") + def _detect_external_drift(self, target: str, raw: str) -> Optional[str]: - """Backup path if *raw* shows external drift, else None. Signals: round-trip - mismatch, or one entry over the whole-file limit (no tool-written entry can — - an external writer appended free-form text). Snapshots to ``.bak.``.""" - if not raw.strip(): - return None + """``.bak.`` snapshot path if *raw* shows external drift, else None. Signals: + round-trip mismatch, or one entry over the whole-file limit (no tool-written + entry can be — an external writer appended free-form text).""" parsed = self._parse_entries(raw) - if raw.strip() == ENTRY_DELIMITER.join(parsed) and max(map(len, parsed), default=0) <= self._char_limit(target): + if not raw.strip() or (raw.strip() == ENTRY_DELIMITER.join(parsed) + and max(map(len, parsed), default=0) <= self._char_limit(target)): return None path = self._path_for(target) bak_path = path.with_suffix(path.suffix + f".bak.{int(time.time())}") @@ -429,11 +388,3 @@ class MemoryStore: except OSError: return str(bak_path) + " (BACKUP FAILED — file unchanged on disk)" return str(bak_path) - - @staticmethod - def _write_file(path: Path, entries: List[str]): - """Atomic temp-file + rename: readers never see a truncated file.""" - try: - atomic_write_text(path, ENTRY_DELIMITER.join(entries), tmp_prefix=".mem_") - except OSError as e: - raise RuntimeError(f"Failed to write memory file {path}: {e}") diff --git a/tools/microsoft_graph_auth.py b/tools/microsoft_graph_auth.py index 5c8edc1160..538597dcf4 100644 --- a/tools/microsoft_graph_auth.py +++ b/tools/microsoft_graph_auth.py @@ -53,13 +53,10 @@ class GraphCredentials: @property def token_url(self) -> str: - tenant = self.tenant_id.strip().strip("/") - return f"{self.authority_url.rstrip('/')}/{tenant}/oauth2/v2.0/token" + return f"{self.authority_url.rstrip('/')}/{self.tenant_id.strip().strip('/')}/oauth2/v2.0/token" @classmethod - def from_env( - cls, environ: dict[str, str] | None = None, *, required: bool = True - ) -> "GraphCredentials | None": + def from_env(cls, environ: dict[str, str] | None = None, *, required: bool = True) -> "GraphCredentials | None": env = environ if environ is not None else os.environ values = [(env.get(name) or "").strip() for name in _REQUIRED_ENV] missing = [name for name, value in zip(_REQUIRED_ENV, values) if not value] @@ -90,10 +87,9 @@ class CachedAccessToken: class MicrosoftGraphTokenProvider: """Acquire and cache Microsoft Graph app-only access tokens.""" - def __init__( - self, credentials: GraphCredentials, *, timeout: float = 20.0, - skew_seconds: int = DEFAULT_TOKEN_SKEW_SECONDS, transport: httpx.AsyncBaseTransport | None = None, - ) -> None: + def __init__(self, credentials: GraphCredentials, *, timeout: float = 20.0, + skew_seconds: int = DEFAULT_TOKEN_SKEW_SECONDS, + transport: httpx.AsyncBaseTransport | None = None) -> None: self.credentials, self.timeout, self.skew_seconds = credentials, timeout, max(0, int(skew_seconds)) self._transport = transport self._cached_token: CachedAccessToken | None = None @@ -150,10 +146,8 @@ class MicrosoftGraphTokenProvider: except (TypeError, ValueError) as exc: raise MicrosoftGraphTokenError( "Microsoft Graph token response did not include a valid expires_in.") from exc - return CachedAccessToken( - access_token=access_token, - token_type=str(payload.get("token_type") or "Bearer").strip() or "Bearer", - expires_at=time.time() + max(0, expires_in_seconds)) + return CachedAccessToken(access_token, time.time() + max(0, expires_in_seconds), + str(payload.get("token_type") or "Bearer").strip() or "Bearer") def _extract_error_detail(response: httpx.Response) -> str: diff --git a/tools/microsoft_graph_client.py b/tools/microsoft_graph_client.py index 75c1038693..faf29d7acb 100644 --- a/tools/microsoft_graph_client.py +++ b/tools/microsoft_graph_client.py @@ -10,7 +10,7 @@ from typing import Any, Awaitable, Callable import httpx from agent.retry_utils import parse_retry_after_seconds -from tools.microsoft_graph_auth import GraphCredentials, MicrosoftGraphTokenProvider, format_graph_error +from tools.microsoft_graph_auth import MicrosoftGraphTokenProvider, format_graph_error DEFAULT_GRAPH_BASE_URL = "https://graph.microsoft.com/v1.0" @@ -26,9 +26,8 @@ class MicrosoftGraphClientError(RuntimeError): class MicrosoftGraphAPIError(MicrosoftGraphClientError): """Raised when a Graph API request fails.""" - def __init__( - self, status_code: int, method: str, url: str, message: str, *, - retry_after_seconds: float | None = None, payload: Any = None) -> None: + def __init__(self, status_code: int, method: str, url: str, message: str, *, + retry_after_seconds: float | None = None, payload: Any = None) -> None: self.status_code, self.method, self.url = status_code, method, url self.retry_after_seconds, self.payload = retry_after_seconds, payload super().__init__(f"Microsoft Graph API error {status_code} for {method} {url}: {message}") @@ -39,20 +38,15 @@ class MicrosoftGraphClient: alike): transport errors back off exponentially; 401 clears the token cache and refetches; 429/5xx honor ``Retry-After``. Each attempt uses a fresh ``AsyncClient``.""" - def __init__( - self, token_provider: MicrosoftGraphTokenProvider, *, - base_url: str = DEFAULT_GRAPH_BASE_URL, timeout: float = 60.0, max_retries: int = 3, - transport: httpx.AsyncBaseTransport | None = None, - sleep: Callable[[float], Awaitable[None]] | None = None, - user_agent: str = "Hermes-Agent/graph-client") -> None: + def __init__(self, token_provider: MicrosoftGraphTokenProvider, *, + base_url: str = DEFAULT_GRAPH_BASE_URL, timeout: float = 60.0, max_retries: int = 3, + transport: httpx.AsyncBaseTransport | None = None, + sleep: Callable[[float], Awaitable[None]] | None = None, + user_agent: str = "Hermes-Agent/graph-client") -> None: self.token_provider, self.base_url, self.timeout = token_provider, base_url.rstrip("/"), timeout self.max_retries, self.user_agent = max(0, int(max_retries)), user_agent self._transport, self._sleep = transport, sleep or asyncio.sleep - @classmethod - def from_env(cls, **kwargs: Any) -> "MicrosoftGraphClient": - return cls(MicrosoftGraphTokenProvider(GraphCredentials.from_env()), **kwargs) - async def get_json(self, path: str, *, params: Params = None, headers: Headers = None) -> Any: return self._decode_json(await self._request("GET", path, params=params, headers=headers)) @@ -60,27 +54,24 @@ class MicrosoftGraphClient: return self._decode_json(await self._request("POST", path, json_body=json_body, headers=headers)) async def patch_json(self, path: str, *, json_body: Any | None = None, headers: Headers = None) -> Any: + """Decoded body, or ``{}`` for a 204 / bodiless response.""" response = await self._request("PATCH", path, json_body=json_body, headers=headers) - return self._decode_json_or(response, {}) + return self._decode_json(response) if response.status_code != 204 and response.content else {} async def delete(self, path: str, *, headers: Headers = None) -> dict[str, Any]: + """Decoded body, or ``{"deleted": True, "status_code"}`` for a 204 / bodiless response.""" response = await self._request("DELETE", path, headers=headers) - return self._decode_json_or(response, {"deleted": True, "status_code": response.status_code}) + if response.status_code != 204 and response.content: + return self._decode_json(response) + return {"deleted": True, "status_code": response.status_code} - def _decode_json_or(self, response: httpx.Response, empty: Any) -> Any: - """*empty* for a 204 / bodiless response, else the decoded JSON body.""" - return empty if response.status_code == 204 or not response.content else self._decode_json(response) - - async def collect_paginated( - self, path: str, *, params: Params = None, headers: Headers = None) -> list[Any]: + async def collect_paginated(self, path: str, *, params: Params = None, headers: Headers = None) -> list[Any]: """Follow ``@odata.nextLink`` and concatenate every page's ``value`` list.""" items: list[Any] = [] # Query params go on the first request only; @odata.nextLink already embeds them. - next_url: str | None = self._resolve_url(path) - next_params = dict(params or {}) + next_url, next_params = self._resolve_url(path), dict(params or {}) while next_url: - response = await self._request("GET", next_url, params=next_params or None, headers=headers) - payload = self._decode_json(response) + payload = self._decode_json(await self._request("GET", next_url, params=next_params or None, headers=headers)) if not isinstance(payload, dict): raise MicrosoftGraphClientError( f"Expected paginated Graph response dict, got {type(payload).__name__}.") @@ -89,13 +80,11 @@ class MicrosoftGraphClient: next_url, next_params = payload.get("@odata.nextLink"), {} return items - async def download_to_file( - self, path: str, destination: str | Path, *, headers: Headers = None, chunk_size: int = 65536 - ) -> dict[str, Any]: + async def download_to_file(self, path: str, destination: str | Path, *, headers: Headers = None, + chunk_size: int = 65536) -> dict[str, Any]: """Stream a Graph resource to disk chunk-by-chunk (large recordings never fit in memory); written to ``.part`` and renamed into place only on success.""" - url = self._resolve_url(path) - target = Path(destination) + url, target = self._resolve_url(path), Path(destination) target.parent.mkdir(parents=True, exist_ok=True) tmp_target = target.with_suffix(target.suffix + ".part") @@ -103,8 +92,7 @@ class MicrosoftGraphClient: try: async with client.stream("GET", url, headers=request_headers) as response: if response.status_code >= 400: - # Materialize the (small) error body so the message is meaningful. - await response.aread() + await response.aread() # small error body -> meaningful message return response, None with tmp_target.open("wb") as handle: async for chunk in response.aiter_bytes(chunk_size=chunk_size): @@ -119,9 +107,8 @@ class MicrosoftGraphClient: os.replace(tmp_target, target) return {"path": str(target), "size_bytes": target.stat().st_size, "content_type": content_type} - async def _request( - self, method: str, path_or_url: str, *, - params: Params = None, json_body: Any | None = None, headers: Headers = None) -> httpx.Response: + async def _request(self, method: str, path_or_url: str, *, params: Params = None, + json_body: Any | None = None, headers: Headers = None) -> httpx.Response: url = self._resolve_url(path_or_url) async def perform(client: httpx.AsyncClient, request_headers: dict[str, str]): @@ -137,47 +124,37 @@ class MicrosoftGraphClient: """Run ``perform`` (-> ``(response, result)``) under the retry policy. ``kind`` only labels transport-failure messages. Raises ``MicrosoftGraphAPIError`` once retries are exhausted or the status is not retryable; only 401 forces a token refresh.""" - attempt = 0 last_error: Exception | None = None - - while attempt <= self.max_retries: + for attempt in range(self.max_retries + 1): token = await self.token_provider.get_access_token( - force_refresh=attempt > 0 - and isinstance(last_error, MicrosoftGraphAPIError) - and last_error.status_code == 401) - request_headers = {"Authorization": f"Bearer {token}", "Accept": accept, "User-Agent": self.user_agent} - if json_body is not None: - request_headers["Content-Type"] = "application/json" - if headers: - request_headers.update(headers) - + force_refresh=isinstance(last_error, MicrosoftGraphAPIError) and last_error.status_code == 401) + request_headers = {"Authorization": f"Bearer {token}", "Accept": accept, "User-Agent": self.user_agent, + **({"Content-Type": "application/json"} if json_body is not None else {}), + **(headers or {})} + exhausted = attempt >= self.max_retries try: async with httpx.AsyncClient(timeout=httpx.Timeout(self.timeout), transport=self._transport) as client: response, result = await perform(client, request_headers) except httpx.HTTPError as exc: last_error, response = exc, None - if attempt >= self.max_retries: + if exhausted: raise MicrosoftGraphClientError( f"Microsoft Graph {kind} failed for {method} {url}: {exc}") from exc else: if response.status_code < 400: return result - last_error = self._build_api_error(method, url, response) - status = response.status_code - if attempt >= self.max_retries or not (status in (401, 429) or 500 <= status < 600): + last_error, status = self._build_api_error(method, url, response), response.status_code + if exhausted or not (status in (401, 429) or 500 <= status < 600): raise last_error if status == 401: self.token_provider.clear_cache() await self._sleep(self._retry_delay(response, attempt)) - attempt += 1 - raise MicrosoftGraphClientError(f"Microsoft Graph {kind} exhausted retries for {method} {url}.") def _resolve_url(self, path_or_url: str) -> str: if path_or_url.startswith(("http://", "https://")): return path_or_url - path = path_or_url if path_or_url.startswith("/") else f"/{path_or_url}" - return f"{self.base_url}{path}" + return f"{self.base_url}{path_or_url if path_or_url.startswith('/') else '/' + path_or_url}" @staticmethod def _decode_json(response: httpx.Response) -> Any: @@ -201,5 +178,5 @@ class MicrosoftGraphClient: payload = None detail = format_graph_error(payload.get("error")) if isinstance(payload, dict) else None return MicrosoftGraphAPIError( - response.status_code, method, url, detail if detail is not None else (response.text.strip() or "unknown error"), + response.status_code, method, url, response.text.strip() or "unknown error" if detail is None else detail, retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload) 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/process_registry.py b/tools/process_registry.py index a32b0f67fd..a55b131d8d 100644 --- a/tools/process_registry.py +++ b/tools/process_registry.py @@ -363,8 +363,7 @@ class ProcessRegistry: the gateway asyncio loop (watchers, reset checks) and the cleanup thread.""" _SHELL_NOISE_SUBSTRINGS = ( - "no job control in this shell", - "cannot set terminal process group", + "no job control in this shell", "cannot set terminal process group", "tcsetattr: Inappropriate ioctl for device") def __init__(self): @@ -394,10 +393,8 @@ class ProcessRegistry: 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. @@ -432,16 +429,13 @@ class ProcessRegistry: # avoids stale notifications minutes after the process ended. if session.exited: return - matched_lines = [] - matched_pattern = None - for line in new_text.splitlines(): - pat = next((p for p in session.watch_patterns if p in line), None) # one match per line - if pat is not None: - matched_lines.append(line.rstrip()) - if matched_pattern is None: - matched_pattern = pat - 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() with session._lock: if session._watch_cooldown_until and now < session._watch_cooldown_until: @@ -535,45 +529,41 @@ class ProcessRegistry: 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 = trip_msg = None + 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: - 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) + 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 else: if now - self._global_watch_window_start >= WATCH_GLOBAL_WINDOW_SECONDS: - self._global_watch_window_start = now - self._global_watch_window_hits = 0 + 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_msg = self._global_watch_event( + 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.") - # Queue summary events outside the lock. - for msg in (release_msg, trip_msg): - if msg is not None: - self.completion_queue.put(msg) + f"{WATCH_GLOBAL_COOLDOWN_SECONDS}s.")) + for msg in events: + self.completion_queue.put(msg) return admit @staticmethod @@ -589,11 +579,9 @@ class ProcessRegistry: @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 @@ -604,9 +592,8 @@ class ProcessRegistry: 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 - return expected_start is None or cls._safe_host_start_time(pid) == expected_start + 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.""" @@ -619,9 +606,8 @@ class ProcessRegistry: with session._lock: if session.exited: return session - session.exited = True # No waitable handle survives recovery, so the real exit code is unknown. - session.exit_code = None + session.exited, session.exit_code = True, None self._move_to_finished(session) return session @@ -630,9 +616,7 @@ class ProcessRegistry: """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 @@ -921,8 +905,8 @@ class ProcessRegistry: 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 with - # stateless decoding; the incremental decoder holds the partial sequence. + # 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): @@ -1648,20 +1632,18 @@ class ProcessRegistry: 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: + 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. try: - sink(self.get(session_id), session_id) + self.on_close(self.get(session_id), session_id) except Exception as e: return {"status": "error", "error": str(e)} return { @@ -1836,10 +1818,9 @@ class ProcessRegistry: 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 env handle is gone logger.info( "Skipping recovery for non-host process: %s (pid=%s, scope=%s)", diff --git a/tools/process_registry_notifications.py b/tools/process_registry_notifications.py index 280a58814b..0f70f371fd 100644 --- a/tools/process_registry_notifications.py +++ b/tools/process_registry_notifications.py @@ -1,14 +1,13 @@ -"""Human-readable rendering of background-process notification events. - -Events come off ``ProcessRegistry.completion_queue`` (completion, watch_match, -watch_disabled, watch_overflow_*, async_delegation) and are turned into the -``[IMPORTANT: ...]`` / ``[ASYNC DELEGATION ...]`` text injected into the agent -conversation by the CLI drain loop, the gateway, and the TUI. -""" +"""Human-readable rendering of ``ProcessRegistry.completion_queue`` events (completion, +watch_match, watch_disabled, watch_overflow_*, async_delegation) into the +``[IMPORTANT: ...]`` / ``[ASYNC DELEGATION ...]`` text the CLI drain loop, gateway and +TUI inject into the agent conversation.""" import time from contextlib import suppress +_DONE = ("completed", "success") + def _format_age(seconds: float) -> str: """Human-friendly elapsed string ('18m', '2h3m', '45s').""" @@ -26,34 +25,28 @@ def _format_age(seconds: float) -> str: def _model_not_found_patterns() -> "list[str]": - """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.""" + """Model-not-found phrases from ``agent.error_classifier`` (the failover path's + own list, so nothing drifts); a minimal built-in set if the import fails.""" try: from agent.error_classifier import _MODEL_NOT_FOUND_PATTERNS - return list(_MODEL_NOT_FOUND_PATTERNS) except Exception: return ["is not a valid model", "model not found", "model_not_found"] def _delegation_config() -> dict: - """Active delegation config (model/provider/fallbacks); ``{}`` on any error. - Lazy ``tools.delegate_tool._load_config`` so the renderer sees the dispatcher's - model/provider without importing the heavy delegation module at import time.""" + """Active delegation config; ``{}`` on any error. Lazy: delegate_tool is heavy.""" try: from tools.delegate_tool import _load_config as _cfg - return _cfg() or {} except Exception: return {} 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.""" + """True when a result reflects a config-level model_not_found rejection: needs a + model-not-found phrase AND the currently-configured model name in the same text, + so a stale task failing on a removed model is not mis-attributed to the config.""" model = str((config or {}).get("model") or "").lower() if not model: return False @@ -74,11 +67,9 @@ 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)."] 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.") return lines @@ -94,71 +85,59 @@ def _is_truncated(entry: dict) -> bool: return bool(entry.get("truncated") or entry.get("exit_reason") == "max_iterations") -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.""" +def _notice_lines(results) -> "list[str]": + """Blank + model_not_found notice block, or [] when the notice does not apply.""" + notice = _delegation_model_not_found_notice(results) + return ["", *notice] if notice else [] + + +def _preamble(evt: dict, title: str, intro: str, completed_at: float, *, with_goal: bool) -> "list[str]": + """Shared preamble: title, intro, blank, dispatch time, [goal], context/toolsets, role+model.""" 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]": - lines = [] + if with_goal: + lines.append(f"Original goal: {evt.get('goal', '') or ''}") if evt.get("context"): lines.append(f"Context you provided: {evt['context']}") if evt.get("toolsets"): lines.append(f"Toolsets: {', '.join(evt['toolsets'])}") + lines.append(f"Role: {evt.get('role') or 'leaf'} Model: {evt.get('model') or '?'}") 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 [] - goals = evt.get("goals") or [] + results, goals = evt.get("results") or [], evt.get("goals") or [] n = len(results) if results else len(goals) - total_dur = evt.get("total_duration_seconds", evt.get("duration_seconds", "?")) - error = evt.get("error") - lines = _header_lines( + lines = _preamble( 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.", - completed_at) - lines.extend(_task_source_lines(evt)) - lines.append(f"{_role_model(evt)} Total duration: {total_dur}s") - if error and not results: - lines += ["--- ERROR ---", f"The batch did not complete successfully: {error}"] + completed_at, with_goal=False) + lines[-1] += f" Total duration: {evt.get('total_duration_seconds', evt.get('duration_seconds', '?'))}s" + if evt.get("error") and not results: + lines += ["--- ERROR ---", f"The batch did not complete successfully: {evt['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 += ["", *_notice] + lines += _notice_lines(results) for r in sorted(results, key=lambda x: x.get("task_index", 0)): - idx = r.get("task_index", 0) - r_status = r.get("status", "?") - r_summary = r.get("summary") - r_error = r.get("error") + idx, r_truncated = r.get("task_index", 0), _is_truncated(r) + r_status, r_summary, r_error = r.get("status", "?"), r.get("summary"), r.get("error") 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 "✗") - 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" + icon = "⚠" if r_truncated else ("✓" if r_status in _DONE else "✗") + header = (f"--- {icon} TASK {idx + 1}/{n}" + (f": {r_goal}" if r_goal else "") + f" (status={r_status}" + + (f", api_calls={r['api_calls']}" if r.get("api_calls") else "") + + (f", {r['duration_seconds']}s" if r.get("duration_seconds") is not None else "") + + (", TRUNCATED: hit max_iterations — work may be incomplete" if r_truncated else "")) lines += ["", header + ") ---"] - if r_status in ("completed", "success") and r_summary: + if r_status in _DONE and r_summary: if r_truncated: lines.append(_TRUNCATED_SUMMARY_NOTE) lines.append(r_summary) @@ -174,39 +153,27 @@ def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> s 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.""" + """Self-contained re-injection for an async-delegation completion: the FULL + original task source (goal, context, toolsets, role, model), dispatch time, status + and result, so an agent deep in unrelated context can act on it or re-dispatch.""" 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") + status, summary, error = evt.get("status") or "completed", evt.get("summary"), evt.get("error") truncated = _is_truncated(evt) - lines = _header_lines( + lines = _preamble( 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.", - completed_at) - lines.append(f"Original goal: {evt.get('goal', '') or ''}") - lines.extend(_task_source_lines(evt)) - lines.append(_role_model(evt)) - _notice = _delegation_model_not_found_notice([evt]) - if _notice: - lines += ["", *_notice] - _trunc = " [TRUNCATED: hit max_iterations — work may be incomplete]" if truncated else "" - lines += [ - f"Status: {status} API calls: {evt.get('api_calls', 0)} " - f"Duration: {evt.get('duration_seconds', '?')}s{_trunc}", - "--- RESULT ---", - ] - if status in ("completed", "success") and summary: + completed_at, with_goal=True) + lines += _notice_lines([evt]) + [ + f"Status: {status} API calls: {evt.get('api_calls', 0)} Duration: {evt.get('duration_seconds', '?')}s" + + (" [TRUNCATED: hit max_iterations — work may be incomplete]" if truncated else ""), + "--- RESULT ---"] + if status in _DONE and summary: if truncated: lines.append(_TRUNCATED_SUMMARY_NOTE) lines.append(summary) @@ -215,92 +182,71 @@ def _format_async_delegation(evt: dict) -> str: 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 += ["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. - 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.""" + """One-line provenance for a subagent-owned process event, else None. Such a process + outlives the child and lands in the PARENT conversation, which would otherwise see an + anonymous output wall. Keyed on ``owner_task_id`` — ``task_id`` may be 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 info = None with suppress(Exception): from tools.delegate_tool import get_subagent_attribution - info = get_subagent_attribution(task_id) if not info: # Registry entry aged out — still attribute generically, not anonymously. return f"Started by subagent {task_id} (delegate_task)." - goal = str(info.get("goal") or "").strip() - if len(goal) > 120: - goal = goal[:117] + "..." - deleg = info.get("delegation_id") - line = f"Started by subagent {task_id}" + (f" of delegation {deleg}" if deleg else "") + "." - if goal: - line += f' Task: "{goal}"' - return line + goal, deleg = str(info.get("goal") or "").strip(), info.get("delegation_id") + goal = goal[:117] + "..." if len(goal) > 120 else goal + return (f"Started by subagent {task_id}" + (f" of delegation {deleg}" if deleg else "") + "." + + (f' Task: "{goal}"' if goal else "")) -_REASON_STATUS = { - "lost": "marked lost because the process backend disappeared", - "failed_start": "failed to start", -} +_REASON_STATUS = {"lost": "marked lost because the process backend disappeared", "failed_start": "failed to start"} def _completion_status(evt: dict) -> str: reason = evt.get("completion_reason") or "exited" if reason == "killed": return f"terminated by {evt.get('termination_source') or 'Hermes'}" - if reason in _REASON_STATUS: - return _REASON_STATUS[reason] - return "completed normally" if evt.get("exit_code", "?") == 0 else "exited" + return _REASON_STATUS.get(reason) or ("completed normally" if evt.get("exit_code", "?") == 0 else "exited") def format_process_notification(evt: dict) -> "str | None": """Format a completion_queue event into an ``[IMPORTANT: ...]`` message.""" evt_type = evt.get("type", "completion") - _sid = evt.get("session_id", "unknown") - _cmd = evt.get("command", "unknown") - _attribution = _delegation_attribution_line(evt) - # 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', '')}]" - if evt_type == "watch_match": - _sup = evt.get("suppressed", 0) - text = f"[IMPORTANT: Background process {_sid} matched watch pattern \"{evt.get('pattern', '?')}\".\n" - if _attribution: - text += f"{_attribution}\n" - text += f"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) - + _sid, _cmd = evt.get("session_id", "unknown"), evt.get("command", "unknown") + _attribution = _delegation_attribution_line(evt) + attribution = f"{_attribution}\n" if _attribution else "" + if evt_type == "watch_match": + _sup = evt.get("suppressed", 0) + return ( + f"[IMPORTANT: Background process {_sid} matched watch pattern \"{evt.get('pattern', '?')}\".\n" + f"{attribution}Command: {_cmd}\nMatched output:\n{evt.get('output', '')}" + + (f"\n({_sup} earlier matches were suppressed by rate limit)" if _sup else "") + "]") _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)} (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}]") 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/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 b6962e4393..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,25 +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`` (null+multi-const unions land here, not in - ``strip_nullable_unions``). 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) @@ -241,53 +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); -# passed through (``required`` follows property renames). +# 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): @@ -295,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: @@ -361,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, @@ -379,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 '