From 1545afb892bbbea38fe34d5e39bddad4600aae69 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 23:56:07 -0700 Subject: [PATCH] refactor(tools): simplify registry, schema_sanitizer, self_repo_guard, plugin_guard, project_tools, path_security (-25% LOC) Zero behavior change; tool schemas byte-identical; golden corpus (376 guard commands, 93 heredoc forms, sanitizer/unrename cases) identical old vs new. registry.py (1263 -> 924): _memo_check per-pass check_fn memo (2 sites), _grouped/_toolset_entries toolset grouping (6 sites), _unique_env replaces _extend_unique, _attr accessor for get_schema/get_toolset_for_tool/get_emoji/ get_max_result_size, fold cache-prune loops, flatten _callable_module walk, merge guard try-blocks, docstring/comment compaction (all WHY kept). schema_sanitizer.py (511 -> 392): _rewrite bottom-up tree map shared by _strip_ref_siblings/strip_nullable_unions/collapse_const_unions, _dict_nodes generator replaces nested _walk closure, collapsed _const_branch_type guards. self_repo_guard.py (678 -> 517): _scope_keys state machine as one if/elif ladder, heredoc opener parsed by one regex (_HEREDOC_OPENER_RE), _masked_line inlined, _operator_before via rstrip, folded tail expressions. plugin_guard.py (235 -> 163): positional Finding ctor, packed tables, docs. project_tools.py (181 -> 156): _activated shared by create/switch, _ACTIONS dict dispatch replaces the if-chain in _handle_project. path_security.py (24 -> 18): unused logger/logging import dropped. --- tools/path_security.py | 8 +- tools/plugin_guard.py | 140 ++---- tools/project_tools.py | 91 ++-- tools/registry.py | 919 ++++++++++++-------------------------- tools/schema_sanitizer.py | 407 ++++++----------- tools/self_repo_guard.py | 357 ++++----------- 6 files changed, 600 insertions(+), 1322 deletions(-) diff --git a/tools/path_security.py b/tools/path_security.py index 79051e53db..222df08fd0 100644 --- a/tools/path_security.py +++ b/tools/path_security.py @@ -1,17 +1,11 @@ """Shared path validation helpers for tool implementations (skills, cron, credential files).""" -import logging from pathlib import Path from typing import Optional -logger = logging.getLogger(__name__) - def validate_within_dir(path: Path, root: Path) -> Optional[str]: - """Return an error message if *path* does not resolve inside *root*, else None. - - ``Path.resolve()`` follows symlinks and normalises ``..`` before the check. - """ + """Error message if *path* does not resolve inside *root* (symlinks and ``..`` followed).""" try: path.resolve().relative_to(root.resolve()) except (ValueError, OSError) as exc: diff --git a/tools/plugin_guard.py b/tools/plugin_guard.py index 26880b3140..b3814c336b 100644 --- a/tools/plugin_guard.py +++ b/tools/plugin_guard.py @@ -1,21 +1,11 @@ #!/usr/bin/env python3 -"""Plugin Guard — security scanner for externally-installed plugins. +"""Plugin Guard — ``skills_guard`` engine applied to ``hermes plugins install``/``update``. -Reuses the ``tools/skills_guard.py`` static-analysis engine for -``hermes plugins install`` / ``update``, which otherwise clone and execute -arbitrary Git repositories unscanned. - -Plugins run Python in-process (more dangerous than skills) but are *expected* -to read their own API keys from env vars, call provider HTTP APIs and spawn -subprocesses, so the raw skill patterns would flag every legitimate provider -plugin. Hence: full pattern set on docs/config files (where prompt-injection -lives); the "reads own env secret" / "HTTP call with key" family is exempt on -*code* files while genuinely malicious signals stay; plugin-sized structural -limits; VCS/venv noise skipped. - -Verdict → install policy: ``safe`` installs; ``caution`` requires explicit -confirmation (prompt, ``--force``, or caller callback); ``dangerous`` is -blocked and ``--force`` does NOT override. +Plugins run in-process but are *expected* to read their own env keys, call provider APIs +and spawn subprocesses, so: full pattern set on docs/config files (where prompt-injection +lives); the "reads own secret"/"HTTP call with key" family exempt on *code* files; +plugin-sized structural limits; VCS/venv noise skipped. ``safe`` installs, ``caution`` +needs confirmation, ``dangerous`` is blocked and ``--force`` does NOT override. """ from __future__ import annotations @@ -25,64 +15,35 @@ from pathlib import Path from typing import Iterator, List, Optional, Tuple from tools.skills_guard import ( - Finding, - ScanResult, - SUSPICIOUS_BINARY_EXTENSIONS, - _determine_verdict, - format_scan_report, - scan_file, -) + Finding, ScanResult, SUSPICIOUS_BINARY_EXTENSIONS, _determine_verdict, format_scan_report, + scan_file) PLUGIN_SCANNER_VERSION = "plugin-guard-v1" # Never scanned: VCS internals, caches, vendored envs. EXCLUDED_DIRS = { ".git", "__pycache__", "node_modules", ".venv", "venv", - ".mypy_cache", ".pytest_cache", ".ruff_cache", ".tox", -} + ".mypy_cache", ".pytest_cache", ".ruff_cache", ".tox"} -# Code files, where "reads an env secret" / "HTTP call with a key variable" -# is the NORMAL, documented plugin pattern (requires_env). -CODE_FILE_EXTENSIONS = { - ".py", ".js", ".ts", ".sh", ".bash", ".rb", ".pl", ".php", -} +# Code files, where "reads an env secret" / "HTTP call with a key" is normal (requires_env). +CODE_FILE_EXTENSIONS = {".py", ".js", ".ts", ".sh", ".bash", ".rb", ".pl", ".php"} -# skills_guard pattern ids exempt on code files (every legitimate provider -# plugin exhibits them); they still apply in full to docs/config files. +# Pattern ids exempt on code files (every legitimate provider plugin trips them); still +# applied in full to docs/config files. CODE_EXEMPT_PATTERN_IDS = { - "python_environ_get_secret", - "python_getenv_secret", - "python_os_environ", - "node_process_env", - "ruby_env_secret", - "env_exfil_httpx", - "env_exfil_requests", - "env_exfil_fetch", - "env_exfil_curl", - "env_exfil_wget", - # Agent-facing instruction patterns are meaningless inside code - # (docstrings/comments about prompts trip them constantly). - "context_exfil", - "send_to_url", - "fake_policy", - # Plugins legitimately write their own settings into config.yaml during - # post_setup, and encode credentials (e.g. HTTP Basic auth) with base64. - "agent_config_mod", - "agent_config_contract", - "encoded_exfil", -} + "python_environ_get_secret", "python_getenv_secret", "python_os_environ", "node_process_env", + "ruby_env_secret", "env_exfil_httpx", "env_exfil_requests", "env_exfil_fetch", + "env_exfil_curl", "env_exfil_wget", + # Agent-facing instruction patterns are meaningless inside code (prompt docstrings trip them). + "context_exfil", "send_to_url", "fake_policy", + # Plugins legitimately write config.yaml in post_setup and base64 credentials (Basic auth). + "agent_config_mod", "agent_config_contract", "encoded_exfil"} -# Severity remaps for plugins. A bundled binary is warn-tier (plugin repos -# occasionally vendor one legitimately; skills never should). A mere -# ``~/.hermes/.env`` reference is the DOCUMENTED way plugin READMEs tell users -# where keys go — informational; actually READING it still trips -# ``read_secrets_file`` (critical). ``curl | sh`` install instructions are -# common in READMEs: caution, not an unoverridable block. +# Severity remaps: a bundled binary is warn-tier (repos occasionally vendor one); a mere +# ``~/.hermes/.env`` mention is how READMEs say where keys go (READING it still trips +# ``read_secrets_file``, critical); ``curl | sh`` in READMEs is caution, not a hard block. SEVERITY_REMAP = { - "binary_file": "high", - "hermes_env_access": "medium", - "curl_pipe_shell": "high", -} + "binary_file": "high", "hermes_env_access": "medium", "curl_pipe_shell": "high"} # Structural limits — plugins are real codebases, far larger than skills. MAX_PLUGIN_FILE_COUNT = 400 @@ -102,8 +63,7 @@ def _walk(plugin_dir: Path) -> Iterator[Tuple[Path, str]]: def _finding(pattern_id: str, severity: str, category: str, file: str, match: str, description: str) -> Finding: - return Finding(pattern_id=pattern_id, severity=severity, category=category, - file=file, line=0, match=match, description=description) + return Finding(pattern_id, severity, category, file, 0, match, description) def _filter_findings(findings: List[Finding], rel_path: str) -> List[Finding]: @@ -124,7 +84,6 @@ def _check_plugin_structure(plugin_dir: Path) -> List[Finding]: file_count = 0 total_size = 0 resolved_root = plugin_dir.resolve() - for f, rel in _walk(plugin_dir): if f.is_symlink(): file_count += 1 @@ -138,50 +97,38 @@ def _check_plugin_structure(plugin_dir: Path) -> List[Finding]: findings.append(_finding("symlink_escape", "critical", "traversal", rel, f"symlink -> {resolved}", "symlink points outside the plugin directory")) continue - if not f.is_file(): continue file_count += 1 - try: size = f.stat().st_size except OSError: continue total_size += size - if size > MAX_PLUGIN_SINGLE_FILE_KB * 1024: findings.append(_finding("oversized_file", "medium", "structural", rel, f"{size // 1024}KB", f"file is {size // 1024}KB (limit: {MAX_PLUGIN_SINGLE_FILE_KB}KB)")) - ext = f.suffix.lower() if ext in SUSPICIOUS_BINARY_EXTENSIONS: findings.append(_finding("binary_file", SEVERITY_REMAP["binary_file"], "structural", rel, f"binary: {ext}", f"binary/executable file ({ext}) bundled in plugin (cannot be scanned)")) - if file_count > MAX_PLUGIN_FILE_COUNT: findings.append(_finding("too_many_files", "medium", "structural", "(directory)", f"{file_count} files", f"plugin has {file_count} files (limit: {MAX_PLUGIN_FILE_COUNT})")) if total_size > MAX_PLUGIN_TOTAL_SIZE_KB * 1024: findings.append(_finding("oversized_bundle", "medium", "structural", "(directory)", f"{total_size // 1024}KB", f"plugin is {total_size // 1024}KB total (limit: {MAX_PLUGIN_TOTAL_SIZE_KB}KB)")) - return findings def scan_plugin(plugin_dir: Path, source: str = "") -> ScanResult: - """Scan a plugin directory (typically the temp clone) for security threats. - - Returns a ScanResult with verdict ``safe`` | ``caution`` | ``dangerous``; - every externally installed plugin is ``community`` trust. - """ + """Scan a plugin directory (typically the temp clone); every external plugin is ``community`` trust.""" all_findings: List[Finding] = [] - if plugin_dir.is_dir(): all_findings.extend(_check_plugin_structure(plugin_dir)) for f, rel in sorted(_walk(plugin_dir)): if f.is_file() and not f.is_symlink(): all_findings.extend(_filter_findings(scan_file(f, rel_path=rel), rel)) - verdict = _determine_verdict(all_findings) if all_findings: categories = sorted({f.category for f in all_findings}) @@ -189,31 +136,17 @@ def scan_plugin(plugin_dir: Path, source: str = "") -> ScanResult: else: summary = f"{plugin_dir.name}: clean scan, no threats detected" result = ScanResult( - skill_name=plugin_dir.name, - source=source or plugin_dir.name, - trust_level="community", - verdict=verdict, - findings=all_findings, - scanned_at=datetime.now(timezone.utc).isoformat(), - summary=summary, - ) + skill_name=plugin_dir.name, source=source or plugin_dir.name, trust_level="community", + verdict=verdict, findings=all_findings, scanned_at=datetime.now(timezone.utc).isoformat(), + summary=summary) result.scan_provenance = { - "scanner_version": PLUGIN_SCANNER_VERSION, - "verdict": verdict, - "source": result.source, - } + "scanner_version": PLUGIN_SCANNER_VERSION, "verdict": verdict, "source": result.source} return result def should_allow_plugin_install( - result: ScanResult, - force: bool = False, -) -> Tuple[Optional[bool], str]: - """Map a plugin scan verdict to ``(allowed, reason)``. - - ``True`` installs, ``None`` needs explicit confirmation (caution), ``False`` - is blocked — ``force`` never overrides ``dangerous``. - """ + result: ScanResult, force: bool = False) -> Tuple[Optional[bool], str]: + """Map a verdict to ``(allowed, reason)``: True installs, None asks to confirm, False blocks.""" n = len(result.findings) if result.verdict == "safe": return True, "Allowed (clean scan)" @@ -223,13 +156,8 @@ def should_allow_plugin_install( return None, f"Requires confirmation (caution verdict, {n} findings)" return False, ( f"Blocked (dangerous verdict, {n} findings). " - f"--force does not override a dangerous verdict." - ) + f"--force does not override a dangerous verdict.") __all__ = [ - "scan_plugin", - "should_allow_plugin_install", - "format_scan_report", - "PLUGIN_SCANNER_VERSION", -] + "scan_plugin", "should_allow_plugin_install", "format_scan_report", "PLUGIN_SCANNER_VERSION"] diff --git a/tools/project_tools.py b/tools/project_tools.py index e7ff1fc90b..c89c7d28a2 100644 --- a/tools/project_tools.py +++ b/tools/project_tools.py @@ -1,16 +1,9 @@ #!/usr/bin/env python3 -"""Project tools — the agent's INTENTIONAL handle on first-class Projects. - -Projects (per-profile ``projects.db``) are the named workspaces the desktop -sidebar groups sessions into. Creating / switching a project is a deliberate act -expressed as explicit tools — never a side effect of a terminal ``cd``. - -Exposed only on GUI sessions: the tools live in the `project` toolset (kept off -``_HERMES_CORE_TOOLS``) which the desktop/TUI gateway folds into its resolved -toolsets, so no CLI/messaging/cron schema carries them. The GUI also wires -``set_project_workspace_callback`` so a create/switch re-anchors the live -session's cwd and the sidebar follows the move; the DB write is the durable part. -""" +"""Project tools — the agent's INTENTIONAL handle on first-class Projects (per-profile +``projects.db``, the desktop sidebar's named workspaces). Creating/switching is an explicit +tool call, never a side effect of ``cd``. GUI-only: the `project` toolset stays off +``_HERMES_CORE_TOOLS``; the desktop/TUI gateway folds it in and wires +``set_project_workspace_callback`` so the live session's cwd and sidebar follow.""" import json import os @@ -18,10 +11,8 @@ from typing import Callable, Optional from tools.registry import registry -# Set by the GUI gateway (tui_gateway) at session wiring. Receives -# ``(task_id, primary_path, project_name)`` and re-anchors that session's -# workspace + refreshes the sidebar. ``None`` in CLI / messaging contexts — the -# DB write still happens; there's just no live GUI session to move. +# Set by the GUI gateway: ``(task_id, primary_path, project_name)`` re-anchors that session's +# workspace. ``None`` in CLI/messaging — the DB write still happens, nothing to move. _workspace_callback: Optional[Callable[[str, str, str], None]] = None @@ -50,7 +41,6 @@ def _apply_workspace(task_id: Optional[str], path: Optional[str], name: str) -> def _resolve(conn, token: str): from hermes_cli import projects_db as pdb - token = (token or "").strip() if not token: return None @@ -66,46 +56,41 @@ def _resolve(conn, token: str): return None +def _activated(proj, task_id: Optional[str]) -> str: + primary = _primary_path(proj) + _apply_workspace(task_id, primary, proj.name) + return json.dumps({ + "success": True, "id": proj.id, "slug": proj.slug, "name": proj.name, + "primary_path": primary}) + + def project_list(task_id: Optional[str] = None) -> str: from hermes_cli import projects_db as pdb - with pdb.connect_closing() as conn: active = pdb.get_active_id(conn) projects = pdb.list_projects(conn) - return json.dumps({ "active_id": active, "projects": [ { - "id": p.id, - "slug": p.slug, - "name": p.name, - "primary_path": _primary_path(p), - "active": p.id == active, - } - for p in projects - ], - }) + "id": p.id, "slug": p.slug, "name": p.name, + "primary_path": _primary_path(p), "active": p.id == active} + for p in projects]}) def project_create(name: str, path: Optional[str] = None, task_id: Optional[str] = None) -> str: name = (name or "").strip() if not name: return json.dumps({"success": False, "error": "name is required"}) - from hermes_cli import projects_db as pdb - folder = (path or "").strip() if folder: folder = os.path.abspath(os.path.expanduser(folder)) - try: with pdb.connect_closing() as conn: existing = pdb.find_by_primary_path(conn, folder) if folder else None if existing is not None: - # Idempotent create: the folder already belongs to a project. - # Re-activating it beats minting a duplicate — duplicated - # projects render N identical sidebar subtrees (#75820). + # Idempotent create: duplicates would render N identical sidebar subtrees. pdb.set_active(conn, existing.id) proj = existing else: @@ -114,46 +99,37 @@ def project_create(name: str, path: Optional[str] = None, task_id: Optional[str] proj = pdb.get_project(conn, pid) except ValueError as exc: return json.dumps({"success": False, "error": str(exc)}) - if proj is None: return json.dumps({"success": False, "error": "project vanished after create"}) - - primary = _primary_path(proj) - _apply_workspace(task_id, primary, proj.name) - - return json.dumps({"success": True, "id": proj.id, "slug": proj.slug, "name": proj.name, "primary_path": primary}) + return _activated(proj, task_id) def project_switch(project: str, task_id: Optional[str] = None) -> str: from hermes_cli import projects_db as pdb - with pdb.connect_closing() as conn: proj = _resolve(conn, project) if proj is None: return json.dumps({"success": False, "error": f"no project matching '{project}'"}) pdb.set_active(conn, proj.id) + return _activated(proj, task_id) - primary = _primary_path(proj) - _apply_workspace(task_id, primary, proj.name) - return json.dumps({"success": True, "id": proj.id, "slug": proj.slug, "name": proj.name, "primary_path": primary}) +_ACTIONS = { + "list": lambda args, tid: project_list(task_id=tid), + "create": lambda args, tid: project_create( + name=args.get("name", ""), path=args.get("path"), task_id=tid), + "switch": lambda args, tid: project_switch(project=args.get("name", ""), task_id=tid)} def _handle_project(args, **kw): - action = (args.get("action") or "").strip() - tid = kw.get("task_id") - if action == "list": - return project_list(task_id=tid) - if action == "create": - return project_create(name=args.get("name", ""), path=args.get("path"), task_id=tid) - if action == "switch": - return project_switch(project=args.get("name", ""), task_id=tid) - return json.dumps({"success": False, "error": "action must be one of: create, switch, list."}) + action = _ACTIONS.get((args.get("action") or "").strip()) + if action is None: + return json.dumps( + {"success": False, "error": "action must be one of: create, switch, list."}) + return action(args, kw.get("task_id")) -# Consolidated (#95681, maintainer-directed): project_list/create/switch each -# re-taught "desktop Projects (named workspaces)"; one action enum says it -# once (244 -> ~145 tok). +# One action enum instead of three tools (each re-taught "desktop Projects"; 244 -> ~145 tok). registry.register( name="desktop_project", toolset="project", @@ -164,8 +140,7 @@ registry.register( "this chat into it — pass path to anchor it to a repo/folder (the " "chat's workspace moves there, the sidebar follows). switch: move " "this chat into an existing project by name/slug/id — the " - "intentional way to move the session, not `cd`. list: all " - "projects + which is active." + "intentional way to move the session, not `cd`. list: all projects + which is active." ), "parameters": { "type": "object", diff --git a/tools/registry.py b/tools/registry.py index ae829360bc..1bca3621d9 100644 --- a/tools/registry.py +++ b/tools/registry.py @@ -1,18 +1,7 @@ -"""Central registry for all hermes-agent tools. - -Each tool file calls ``registry.register()`` at module level to declare its -schema, handler, toolset membership, and availability check. ``model_tools.py`` -queries the registry instead of maintaining its own parallel data structures. - -Import chain (circular-import safe): - tools/registry.py (no imports from model_tools or tool files) - ^ - tools/*.py (import from tools.registry at module level) - ^ - model_tools.py (imports tools.registry + all tool modules) - ^ - run_agent.py, cli.py, batch_runner.py, etc. -""" +"""Central registry for all hermes-agent tools: each tool file calls ``registry.register()`` +at import to declare schema, handler, toolset and availability check; ``model_tools.py`` +queries it. Cycle-safe chain: this module imports nothing from model_tools or tool files; +tools/*.py import it; model_tools.py imports both; run_agent/cli import model_tools.""" import ast import functools @@ -43,29 +32,21 @@ def _bound_error_text(text: str) -> str: return text logger.debug( "tool error body truncated for context (%d chars): %s", - len(text), - text[:_MAX_LOGGED_ERROR_CHARS], - ) + len(text), text[:_MAX_LOGGED_ERROR_CHARS]) return text[:_MAX_TOOL_ERROR_CHARS] + _TOOL_ERROR_TRUNCATION_MARKER def _bound_json_error_result(result: str) -> str: - """Trim an oversized ``error`` field in a JSON string result. - - Handlers that serialize exceptions directly — ``json.dumps({"error": - str(exc), ...})`` instead of ``tool_error()`` — bypass the cap in - ``tool_error``. Applied at the dispatch boundary so no registered tool - can return an unbounded error body that stacks across retries. - """ + """Trim an oversized ``error`` field in a JSON string result: handlers that + ``json.dumps({"error": str(exc)})`` directly bypass ``tool_error``'s cap, so this runs + at the dispatch boundary to stop unbounded errors stacking across retries.""" if len(result) <= _MAX_TOOL_ERROR_CHARS or '"error"' not in result: return result try: payload = json.loads(result) except ValueError: - return result - if not isinstance(payload, dict): - return result - error = payload.get("error") + payload = None + error = payload.get("error") if isinstance(payload, dict) else None if not isinstance(error, str) or len(error) <= _MAX_TOOL_ERROR_CHARS: return result payload["error"] = _bound_error_text(error) @@ -73,65 +54,42 @@ def _bound_json_error_result(result: str) -> str: def _is_registry_register_call(node: ast.AST) -> bool: - """Return True when *node* is a ``registry.register(...)`` call expression.""" + """True when *node* is a ``registry.register(...)`` call expression.""" if not isinstance(node, ast.Expr) or not isinstance(node.value, ast.Call): return False func = node.value.func return ( - isinstance(func, ast.Attribute) - and func.attr == "register" - and isinstance(func.value, ast.Name) - and func.value.id == "registry" - ) + isinstance(func, ast.Attribute) and func.attr == "register" + and isinstance(func.value, ast.Name) and func.value.id == "registry") def _module_registers_tools(module_path: Path) -> bool: - """Return True when the module contains a top-level ``registry.register(...)`` call. - - Only inspects module-body statements so that helper modules which happen - to call ``registry.register()`` inside a function are not picked up. - - A cheap text prefilter avoids the ``ast.parse`` cost for files that do not - mention both ``registry`` and ``register`` — a necessary condition for a - top-level ``registry.register()`` call to exist. - """ + """True when the module body (or a module-level ``for``) calls ``registry.register(...)``. + Only module-body statements count, so helpers registering inside a function are skipped; + a text prefilter avoids ``ast.parse`` for files lacking both words.""" try: source = module_path.read_text(encoding="utf-8") - except OSError: - return False - if "registry" not in source or "register" not in source: - return False - try: + if "registry" not in source or "register" not in source: + return False tree = ast.parse(source, filename=str(module_path)) - except SyntaxError: + except (OSError, SyntaxError): return False - - # Module-level ``for`` loops count too: table-driven modules register - # several tools from one loop, which still runs at import time. - for stmt in tree.body: - if _is_registry_register_call(stmt): - return True - if isinstance(stmt, ast.For) and any(_is_registry_register_call(s) for s in stmt.body): - return True - return False + # Table-driven modules register several tools from one loop, still at import time. + return any( + _is_registry_register_call(stmt) + or (isinstance(stmt, ast.For) and any(_is_registry_register_call(s) for s in stmt.body)) + for stmt in tree.body) def discover_builtin_tools(tools_dir: Optional[Path] = None) -> List[str]: - """Import built-in self-registering tool modules and return their module names. - - The per-file AST scan (:func:`_module_registers_tools`) costs ~145 ms over - ~100 files on a warm cache, so verdicts are memoized on disk keyed by - ``(mtime_ns, size)``. A file whose mtime_ns+size match the cached entry is - trusted without re-reading; any mismatch (or a corrupt/missing cache file) - falls back to a fresh scan for that file. The cache write is best-effort - and atomic, so concurrent processes can race harmlessly. - """ + """Import built-in self-registering tool modules and return their module names. The + per-file AST scan costs ~145 ms over ~100 files, so verdicts are memoized on disk keyed + by ``(mtime_ns, size)``; a mismatch or corrupt cache re-scans that file. The write is + best-effort and atomic, so concurrent processes race harmlessly.""" tools_path = Path(tools_dir) if tools_dir is not None else Path(__file__).resolve().parent - cache = _load_discovery_cache() fresh_cache: Dict[str, list] = {} cache_dirty = False - module_names: List[str] = [] for path in sorted(tools_path.glob("*.py")): if path.name in {"__init__.py", "registry.py", "mcp_tool.py"}: @@ -143,11 +101,7 @@ def discover_builtin_tools(tools_dir: Optional[Path] = None) -> List[str]: except OSError: continue cached = cache.get(abs_path) - if ( - isinstance(cached, (list, tuple)) - and len(cached) == 3 - and (cached[0], cached[1]) == stat_key - ): + if isinstance(cached, (list, tuple)) and len(cached) == 3 and tuple(cached[:2]) == stat_key: registers = bool(cached[2]) else: registers = _module_registers_tools(path) @@ -159,7 +113,6 @@ def discover_builtin_tools(tools_dir: Optional[Path] = None) -> List[str]: # Drop entries for files that no longer exist; rewrite only when changed. if cache_dirty or set(fresh_cache) != set(cache): _save_discovery_cache(fresh_cache) - imported: List[str] = [] for mod_name in module_names: try: @@ -173,10 +126,8 @@ def discover_builtin_tools(tools_dir: Optional[Path] = None) -> List[str]: def _discovery_cache_path() -> Optional[Path]: """Path of the tool-discovery verdict cache, or None if unresolvable.""" try: - # Deferred import keeps tools/registry.py a no-deps leaf at module - # import time (hermes_constants itself is stdlib-only, so no cycle). + # Deferred import keeps tools/registry.py a no-deps leaf at import time. from hermes_constants import get_hermes_home - return Path(get_hermes_home()) / "cache" / "tool_discovery_cache.json" except Exception: return None @@ -202,7 +153,6 @@ def _save_discovery_cache(cache: Dict[str, list]) -> None: return try: from utils import atomic_json_write # stdlib+yaml only; no cycle - path.parent.mkdir(parents=True, exist_ok=True) atomic_json_write(path, cache, indent=0) except Exception as e: @@ -211,8 +161,7 @@ def _save_discovery_cache(cache: Dict[str, list]) -> None: @dataclass(eq=False, slots=True) class ToolEntry: - """Metadata for a single registered tool (identity semantics: registry - restore/CAS paths compare entries with ``is``).""" + """Metadata for one registered tool (identity semantics: restore/CAS paths compare ``is``).""" name: str toolset: str @@ -224,10 +173,8 @@ class ToolEntry: description: str emoji: str max_result_size_chars: int | float | None = None - # Zero-arg callable returning schema overrides merged (shallow) on top of - # the base schema at every get_definitions() call — for fields that track - # runtime config (e.g. delegate_task's description must reflect the current - # delegation.max_concurrent_children / max_spawn_depth). + # Zero-arg callable whose dict is shallow-merged onto the schema at every get_definitions() + # — for fields tracking runtime config (delegate_task's description reflects limits). dynamic_schema_overrides: Optional[Callable] = None @@ -242,27 +189,17 @@ class _PluginOverridePolicy: _OVERRIDE_DENIED_MSG = ( "Plugin module {owner!r} cannot override built-in tool {name!r} " - "without operator opt-in (allow_tool_override)." -) + "without operator opt-in (allow_tool_override).") -# --------------------------------------------------------------------------- -# check_fn TTL cache -# -# check_fns probe external state (Docker daemon, Modal SDK, playwright binary) -# that changes on human timescales, so results are cached ~30 s: env-var flips -# via ``hermes tools`` still propagate within a turn or two with no explicit -# invalidation. -# -# Transient-failure suppression: probes can flap (a ``docker version`` that -# times out under load), which would silently strip a whole toolset from the -# agent being built at that instant — most visibly a delegate_task subagent -# reporting "Tool read_file does not exist". So we remember each check's last -# success and, when a fresh probe fails within a short grace window of it, -# serve the last-good True WITHOUT caching the failure. A failure persisting -# past the window is honored, so a backend that really went down stops -# advertising its tools. -# --------------------------------------------------------------------------- +# ---- check_fn TTL cache ---------------------------------------------------- +# check_fns probe external state (Docker, Modal SDK, playwright) that changes on human +# timescales, so results are cached ~30 s: env-var flips via ``hermes tools`` still land +# within a turn or two. Transient-failure suppression: a flapping probe (``docker version`` +# timing out under load) would silently strip a whole toolset from the agent being built — +# most visibly a subagent reporting "Tool read_file does not exist" — so a failure within a +# short grace window of the last success serves the last-good True WITHOUT caching it; a +# failure persisting past the window is honored so a dead backend stops advertising tools. _CHECK_FN_TTL_SECONDS = 30.0 # Grace window after a success in which a failure counts as a flake; kept short @@ -274,6 +211,10 @@ _check_fn_last_good: Dict[tuple[Callable, Optional[str]], float] = {} _check_fn_cache_lock = threading.Lock() CHECK_FN_CACHE_BYPASS = "" _NO_CACHE_CHECK_FNS: Set[Callable] = set() +_BROWSER_IDENTITY_KEYS = ( + "HERMES_SESSION_ID", + "HERMES_BROWSER_CONTROL_PRINCIPAL", + "HERMES_BROWSER_CONTROL_TRANSPORT_FAMILY") def no_cache_check_fn(fn: Callable) -> Callable: @@ -287,58 +228,36 @@ def _fn_label(fn: Callable) -> object: def _prune_check_fn_caches(now: float) -> None: - """Expire stale entries and cap profile-dimensional cache growth. - - Caller must hold ``_check_fn_cache_lock``. - """ - for key, (timestamp, _) in list(_check_fn_cache.items()): - if now - timestamp >= _CHECK_FN_TTL_SECONDS: - _check_fn_cache.pop(key, None) - for key, timestamp in list(_check_fn_last_good.items()): - if now - timestamp >= _CHECK_FN_FAILURE_GRACE_SECONDS: - _check_fn_last_good.pop(key, None) - while len(_check_fn_cache) >= _CHECK_FN_CACHE_MAX: - _check_fn_cache.pop(next(iter(_check_fn_cache))) - while len(_check_fn_last_good) >= _CHECK_FN_CACHE_MAX: - _check_fn_last_good.pop(next(iter(_check_fn_last_good))) + """Expire stale entries and cap profile-dimensional cache growth. Caller holds the lock.""" + for cache, ttl, stamp in ( + (_check_fn_cache, _CHECK_FN_TTL_SECONDS, lambda v: v[0]), + (_check_fn_last_good, _CHECK_FN_FAILURE_GRACE_SECONDS, lambda v: v)): + for key, value in list(cache.items()): + if now - stamp(value) >= ttl: + cache.pop(key, None) + while len(cache) >= _CHECK_FN_CACHE_MAX: + cache.pop(next(iter(cache))) def check_fn_cache_scope() -> Optional[str]: - """Return the active profile key when availability is profile-scoped. - - Browser-controller availability is request-bound and can change on every - attach/detach, so a fully bound browser-control request bypasses both this - cache and model_tools' outer definition cache (same sentinel for both - layers) — one Browser session's live tools must not leak into another. - - Single-profile processes keep the historical process-wide cache. A - multiplex gateway installs a Hermes-home override per profile turn, so the - canonical profile key is the stable isolation boundary. - """ + """Return the active profile key when availability is profile-scoped. Browser-controller + availability is request-bound (changes on every attach/detach), so a fully bound + browser-control request bypasses this cache AND model_tools' outer definition cache (same + sentinel) — one Browser session's live tools must not leak into another. A multiplex + gateway installs a Hermes-home override per profile turn, so that key is the boundary.""" try: from gateway.session_context import get_session_env - - browser_identity = ( - get_session_env("HERMES_SESSION_ID", ""), - get_session_env("HERMES_BROWSER_CONTROL_PRINCIPAL", ""), - get_session_env("HERMES_BROWSER_CONTROL_TRANSPORT_FAMILY", ""), - ) - if all(str(value or "").strip() for value in browser_identity): + if all(str(get_session_env(k, "") or "").strip() for k in _BROWSER_IDENTITY_KEYS): return CHECK_FN_CACHE_BYPASS except Exception: pass - try: from agent.secret_scope import is_multiplex_active - + from hermes_constants import get_hermes_home_override if not is_multiplex_active(): return None - from hermes_constants import get_hermes_home_override - override = get_hermes_home_override() - if not override: - return CHECK_FN_CACHE_BYPASS - return str(Path(override).expanduser().resolve()) + return str(Path(override).expanduser().resolve()) if override else CHECK_FN_CACHE_BYPASS except Exception: # Fail closed: bypass both cache layers rather than aliasing requests # whose multiplex profile identity could not be resolved. @@ -348,41 +267,28 @@ def check_fn_cache_scope() -> Optional[str]: def _run_check_fn_uncached(fn: Callable, *, unresolved_scope: bool = False) -> bool: """Run an availability check without cache/grace handling.""" from agent.secret_scope import UnscopedSecretError - try: return bool(fn()) except UnscopedSecretError: if unresolved_scope: - # Expected fail-closed probe: with multiplexing on, boot-time - # check_fns run before any profile secret scope exists, so - # get_secret raises by design. The tool re-probes on the first - # scoped turn — log without a traceback so this cannot be - # mistaken for a crashed check_fn. + # Expected fail-closed probe: boot-time check_fns run before any multiplex profile + # secret scope exists. No traceback, so it isn't mistaken for a crashed check_fn. logger.debug( "check_fn %s hit the multiplex fail-closed path with no " - "profile secret scope active; dependent tools re-probe on " - "the first scoped turn", - _fn_label(fn), - ) - return False - # The scope resolved but the read still failed closed: a genuinely - # lost scope. Keep the loud crash-style report. - logger.warning( - "check_fn %s raised UnscopedSecretError while the profile cache " - "scope was resolved; dependent tools will be unavailable this turn", - _fn_label(fn), - exc_info=True, - ) - return False + "profile secret scope active; dependent tools re-probe on the first scoped turn", + _fn_label(fn)) + else: + # The scope resolved but the read still failed closed: a genuinely lost scope. + logger.warning( + "check_fn %s raised UnscopedSecretError while the profile cache " + "scope was resolved; dependent tools will be unavailable this turn", + _fn_label(fn), exc_info=True) except Exception: detail = " while profile cache scope was unresolved" if unresolved_scope else "" logger.warning( "check_fn %s raised%s; dependent tools will be unavailable this turn", - _fn_label(fn), - detail, - exc_info=True, - ) - return False + _fn_label(fn), detail, exc_info=True) + return False def _check_fn_cached(fn: Callable) -> bool: @@ -399,57 +305,49 @@ def _check_fn_cached(fn: Callable) -> bool: cached = _check_fn_cache.get(cache_key) if cached is not None: return cached[1] - try: - value = bool(fn()) - outcome = "returned False" + value, outcome = bool(fn()), "returned False" except Exception: - value = False - outcome = "raised" - + value, outcome = False, "raised" with _check_fn_cache_lock: _prune_check_fn_caches(now) if value: _check_fn_last_good[cache_key] = now _check_fn_cache[cache_key] = (now, True) return True - last_good = _check_fn_last_good.get(cache_key) if last_good is not None and now - last_good < _CHECK_FN_FAILURE_GRACE_SECONDS: - # Recent success → treat this failure as a flake. Serve last-good - # True and do NOT cache the failure, so the next call re-probes - # rather than pinning a stale verdict for the full TTL. + # Recent success → flake: serve last-good True, do NOT cache (next call re-probes). logger.warning( "check_fn %s failed (%s) within %.0fs of last success; " "treating as transient and keeping tool(s) available", - _fn_label(fn), outcome, _CHECK_FN_FAILURE_GRACE_SECONDS, - ) + _fn_label(fn), outcome, _CHECK_FN_FAILURE_GRACE_SECONDS) return True - # No recent success (or grace expired) — honor the failure. Log it so - # silent tool loss in quiet mode (subagents) is diagnosable. + # No recent success — honor the failure; logged so silent tool loss is diagnosable. logger.warning( - "check_fn %s %s; dependent tools will be unavailable this turn", - _fn_label(fn), outcome, - ) + "check_fn %s %s; dependent tools will be unavailable this turn", _fn_label(fn), outcome) _check_fn_cache[cache_key] = (now, False) return False +def _memo_check(fn: Callable, memo: Dict[Callable, bool]) -> bool: + """Per-pass memo on top of the TTL cache: one probe per distinct check_fn.""" + if fn not in memo: + memo[fn] = _check_fn_cached(fn) + return memo[fn] + + def invalidate_check_fn_cache() -> None: - """Drop all cached ``check_fn`` results. Call after config changes that - affect tool availability (e.g. ``hermes tools enable``).""" + """Drop all cached ``check_fn`` results (after config changes like ``hermes tools enable``).""" with _check_fn_cache_lock: _check_fn_cache.clear() _check_fn_last_good.clear() def get_cached_check_fn_result(fn: Callable) -> Optional[bool]: - """Return the cached verdict for *fn* if its TTL is still valid, else None. - - NEVER executes the probe: for read-only surfaces (dashboard status panels) - that must not trigger network / auth / SDK work inside a request path. - """ + """Cached verdict for *fn* if its TTL is still valid, else None. NEVER runs the probe: + for read-only surfaces (dashboard panels) that must not do network/auth/SDK work.""" now = time.monotonic() scope = check_fn_cache_scope() if scope == CHECK_FN_CACHE_BYPASS: @@ -457,47 +355,40 @@ def get_cached_check_fn_result(fn: Callable) -> Optional[bool]: return None with _check_fn_cache_lock: cached = _check_fn_cache.get((fn, scope)) - if cached is not None and now - cached[0] < _CHECK_FN_TTL_SECONDS: - return cached[1] - return None + return cached[1] if cached is not None and now - cached[0] < _CHECK_FN_TTL_SECONDS else None class ToolRegistry: """Singleton registry that collects tool schemas + handlers from tool files.""" def __init__(self): - # Built-in and other process-global registrations. - self._tools: Dict[str, ToolEntry] = {} - # Plugin registrations are overlays keyed by resolved HERMES_HOME. A - # profile sees its own overlay first and then the global built-ins. + self._tools: Dict[str, ToolEntry] = {} # built-in / process-global registrations + # Plugin overlays keyed by resolved HERMES_HOME; a profile sees its overlay first. self._scoped_tools: Dict[str, Dict[str, ToolEntry]] = {} - # Plugin module namespace -> operator opt-in for built-in override. - # Authorization records are lifecycle-managed; the separate scope map - # remains durable so delayed callbacks stay profile-confined. - self._plugin_override_policy: Dict[ - tuple[Optional[str], str], _PluginOverridePolicy - ] = {} - # Scope attribution stays durable after policy removal so delayed code - # remains confined to the profile where its module was loaded. + # Plugin namespace -> operator opt-in for built-in override (lifecycle-managed); + # scope attribution stays durable after policy removal so delayed callbacks + # remain confined to the profile that loaded them. + self._plugin_override_policy: Dict[tuple[Optional[str], str], _PluginOverridePolicy] = {} self._plugin_module_scopes: Dict[str, Set[Optional[str]]] = {} self._toolset_checks: Dict[str, Callable] = {} self._toolset_aliases: Dict[str, str] = {} - # MCP dynamic refresh can mutate the registry while other threads are - # reading tool metadata, so keep mutations serialized and readers on - # stable snapshots. + # MCP refresh mutates while other threads read: serialize writes, snapshot reads. self._lock = threading.RLock() - # Monotonically-increasing generation counter. Bumped on every - # mutation (register / deregister / register_toolset_alias / MCP - # refresh). External callers (e.g. get_tool_definitions) can memoize - # against it: a cache entry keyed on the generation is valid for as - # long as the generation hasn't changed. + # Bumped on every mutation; get_tool_definitions memoizes against it. self._generation: int = 0 @staticmethod def current_scope_key() -> str: - """Return the active profile's canonical registry scope.""" return hermes_home_key() + @staticmethod + def _grouped(entries) -> Dict[str, List[ToolEntry]]: + """``{toolset: entries}`` in first-appearance order.""" + groups: Dict[str, List[ToolEntry]] = {} + for entry in entries: + groups.setdefault(entry.toolset, []).append(entry) + return groups + def _slot(self, scope: Optional[str], *, create: bool = False) -> Dict[str, ToolEntry]: """The registration map for *scope*: global when None, else that profile's overlay.""" if scope is None: @@ -508,93 +399,55 @@ class ToolRegistry: def _drop_toolset_aliases(self, toolset: str) -> None: self._toolset_aliases = { - alias: target - for alias, target in self._toolset_aliases.items() - if target != toolset - } + alias: target for alias, target in self._toolset_aliases.items() if target != toolset} def _merged_tools(self, scope: Optional[str] = None) -> Dict[str, ToolEntry]: """Return global tools overlaid with one profile's plugin tools.""" - active_scope = scope or self.current_scope_key() - merged = dict(self._tools) - merged.update(self._scoped_tools.get(active_scope, {})) - return merged + return {**self._tools, **self._scoped_tools.get(scope or self.current_scope_key(), {})} + + def _toolset_entries(self, toolset: str, scope: Optional[str]) -> List[ToolEntry]: + return self._grouped(self._merged_tools(scope).values()).get(toolset, []) def _snapshot_state( - self, - scope: Optional[str] = None, - ) -> tuple[List[ToolEntry], Dict[str, Callable]]: + self, scope: Optional[str] = None) -> tuple[List[ToolEntry], Dict[str, Callable]]: """Return a coherent snapshot of registry entries and toolset checks.""" with self._lock: entries = list(self._merged_tools(scope).values()) checks = dict(self._toolset_checks) - for entry in entries: - if entry.check_fn is not None: - checks[entry.toolset] = entry.check_fn + checks.update({e.toolset: e.check_fn for e in entries if e.check_fn is not None}) return entries, checks def _snapshot_entries(self) -> List[ToolEntry]: - """Return a stable snapshot of registered tool entries.""" return self._snapshot_state()[0] - def _toolset_has_exposable_tools( - self, - toolset: str, - entries: List[ToolEntry], - ) -> bool: - """Return True when at least one tool in *toolset* would be exposed. + def _toolset_has_exposable_tools(self, toolset: str, entries: List[ToolEntry]) -> bool: + """True when at least one tool in *toolset* would be exposed. Mirrors + :meth:`get_definitions` per-tool filtering so doctor/banners agree with runtime: + mixed toolsets (``terminal`` + desktop-only ``read_terminal``) must not be gated + by the first ``check_fn``.""" + memo: Dict[Callable, bool] = {} + members = (e for e in entries if e.toolset == toolset) + return any(not e.check_fn or _memo_check(e.check_fn, memo) for e in members) - Mirrors :meth:`get_tool_definitions` per-tool filtering so doctor, - banners, and other toolset-level surfaces agree with runtime exposure. - Mixed toolsets (e.g. ``terminal`` plus desktop-only ``read_terminal``) - must not be gated solely by the first registered ``check_fn``. - """ - check_results: Dict[Callable, bool] = {} - for entry in entries: - if entry.toolset != toolset: - continue - if not entry.check_fn: - return True - if entry.check_fn not in check_results: - check_results[entry.check_fn] = _check_fn_cached(entry.check_fn) - if check_results[entry.check_fn]: - return True - return False - - def get_entry( - self, - name: str, - *, - scope: Optional[str] = None, - ) -> Optional[ToolEntry]: - """Return the active profile's entry by name, falling back to global.""" + def get_entry(self, name: str, *, scope: Optional[str] = None) -> Optional[ToolEntry]: + """Active profile's entry by name, falling back to global.""" with self._lock: return self._merged_tools(scope).get(name) def snapshot_registration( - self, - name: str, - *, - scope: Optional[str] = None, - ) -> Optional[ToolEntry]: - """Return the local slot state without following global fallback.""" + self, name: str, *, scope: Optional[str] = None) -> Optional[ToolEntry]: + """Local slot state only — no global fallback.""" with self._lock: return self._slot(scope).get(name) def get_registered_toolset_names(self) -> List[str]: - """Return sorted unique toolset names present in the registry.""" - return sorted({entry.toolset for entry in self._snapshot_entries()}) + return sorted(self._grouped(self._snapshot_entries())) def get_all_entries(self) -> List[ToolEntry]: - """Return the active profile's merged tool entries.""" return self._snapshot_entries() def get_tool_names_for_toolset(self, toolset: str) -> List[str]: - """Return sorted tool names registered under a given toolset.""" - return sorted( - entry.name for entry in self._snapshot_entries() - if entry.toolset == toolset - ) + return sorted(e.name for e in self._grouped(self._snapshot_entries()).get(toolset, [])) def register_toolset_alias(self, alias: str, toolset: str) -> None: """Register an explicit alias for a canonical toolset name.""" @@ -602,38 +455,26 @@ class ToolRegistry: existing = self._toolset_aliases.get(alias) if existing and existing != toolset: logger.warning( - "Toolset alias collision: '%s' (%s) overwritten by %s", - alias, existing, toolset, + "Toolset alias collision: '%s' (%s) overwritten by %s", alias, existing, toolset ) self._toolset_aliases[alias] = toolset self._generation += 1 def get_registered_toolset_aliases(self) -> Dict[str, str]: - """Return a snapshot of ``{alias: canonical_toolset}`` mappings.""" with self._lock: return dict(self._toolset_aliases) def get_toolset_alias_target(self, alias: str) -> Optional[str]: - """Return the canonical toolset name for an alias, or None.""" with self._lock: return self._toolset_aliases.get(alias) - # ------------------------------------------------------------------ - # Registration - # ------------------------------------------------------------------ + # ---- Registration ------------------------------------------------ def register_plugin_override_policy( - self, - module_namespace: str, - allowed: bool, - *, - scope: Optional[str] = None, + self, module_namespace: str, allowed: bool, *, scope: Optional[str] = None, ) -> _PluginOverridePolicy: - """Bind a plugin module namespace to its current operator opt-in. - - The identity-bearing result lets plugin unload/reload revoke a stale - authorization without losing durable module-to-profile attribution. - """ + """Bind a plugin module namespace to its current operator opt-in. The identity-bearing + result lets unload/reload revoke a stale authorization without losing attribution.""" with self._lock: policy = _PluginOverridePolicy(allowed) self._plugin_override_policy[(scope, module_namespace)] = policy @@ -641,23 +482,15 @@ class ToolRegistry: return policy def snapshot_plugin_override_policy( - self, - module_namespace: str, - *, - scope: Optional[str] = None, + self, module_namespace: str, *, scope: Optional[str] = None, ) -> Optional[_PluginOverridePolicy]: """Return one local authorization generation without fallback.""" with self._lock: return self._plugin_override_policy.get((scope, module_namespace)) def restore_plugin_override_policy( - self, - module_namespace: str, - current: _PluginOverridePolicy, - previous: Optional[_PluginOverridePolicy], - *, - scope: Optional[str] = None, - ) -> bool: + self, module_namespace: str, current: _PluginOverridePolicy, + previous: Optional[_PluginOverridePolicy], *, scope: Optional[str] = None) -> bool: """CAS-restore policy state while retaining durable scope attribution.""" with self._lock: key = (scope, module_namespace) @@ -669,23 +502,17 @@ class ToolRegistry: self._plugin_override_policy[key] = previous return True - def _plugin_override_allowed( - self, - scope: Optional[str], - module_namespace: str, - ) -> bool: + def _plugin_override_allowed(self, scope: Optional[str], module_namespace: str) -> bool: policy = self._plugin_override_policy.get((scope, module_namespace)) if policy is None and scope is not None: policy = self._plugin_override_policy.get((None, module_namespace)) return bool(policy and policy.allowed) def _plugin_owner_of(self, handler: Callable) -> Optional[str]: - """Plugin namespace that DEFINED *handler* (None for built-in/MCP handlers). - - Bound to ``handler.__globals__["__name__"]``, fixed at definition time so - it cannot drift with call site, thread, or timing; lambdas and nested - functions inherit it, so a plugin cannot launder an override via a callback. - """ + """Plugin namespace that DEFINED *handler* (None for built-in/MCP handlers). Bound to + ``handler.__globals__["__name__"]``, fixed at definition time so it cannot drift with + call site/thread/timing; lambdas and nested functions inherit it, so a plugin cannot + launder an override via a callback.""" mod = self._callable_module(handler) return self._plugin_namespace_of_module(mod) if mod else None @@ -696,40 +523,26 @@ class ToolRegistry: seen: Set[int] = set() while id(current) not in seen: seen.add(id(current)) + globals_dict = getattr(current, "__globals__", None) if isinstance(current, functools.partial): current = current.func - continue - func = getattr(current, "__func__", None) - if func is not None: - current = func - continue - globals_dict = getattr(current, "__globals__", None) - if isinstance(globals_dict, dict): - module_name = globals_dict.get("__name__", "") - if module_name: - return str(module_name) - wrapped = getattr(current, "__wrapped__", None) - if wrapped is not None: - current = wrapped - continue - break + elif getattr(current, "__func__", None) is not None: + current = current.__func__ + elif isinstance(globals_dict, dict) and globals_dict.get("__name__", ""): + return str(globals_dict["__name__"]) + elif getattr(current, "__wrapped__", None) is not None: + current = current.__wrapped__ + else: + break module_name = getattr(current, "__module__", "") - if module_name: - return str(module_name) - return str(getattr(type(current), "__module__", "") or "") + return str(module_name or getattr(type(current), "__module__", "") or "") - def _plugin_namespace_of_module( - self, - module_namespace: str, - ) -> Optional[str]: + def _plugin_namespace_of_module(self, module_namespace: str) -> Optional[str]: """Resolve a module/submodule to its durable plugin namespace.""" with self._lock: matches = [ - namespace - for namespace in self._plugin_module_scopes - if module_namespace == namespace - or module_namespace.startswith(f"{namespace}.") - ] + namespace for namespace in self._plugin_module_scopes + if module_namespace == namespace or module_namespace.startswith(f"{namespace}.")] if matches: return max(matches, key=len) # Also gate plugin modules currently loading but not yet policy-recorded @@ -750,9 +563,8 @@ class ToolRegistry: if len(scopes) == 1: return next(iter(scopes)) raise PermissionError( - f"Plugin module {module_namespace!r} is active in multiple " - "profiles and cannot register outside one of those scopes." - ) + f"Plugin module {module_namespace!r} is active in multiple profiles and cannot " + "register outside one of those scopes.") def plugin_scope_for_module(self, module_namespace: str) -> Optional[str]: """Public host lookup for a loaded plugin module's immutable scope.""" @@ -767,41 +579,22 @@ class ToolRegistry: @staticmethod def _caller_module() -> str: """Best-effort module name of the registry method's caller (two frames up). - - ``deregister()`` takes only a tool name — no handler to bind authorization - to via ``_plugin_owner_of`` — so frame inspection is the only way to know - who is asking. - """ + ``deregister()`` takes only a tool name — no handler for ``_plugin_owner_of`` — + so frame inspection is the only way to know who is asking.""" try: - frame = sys._getframe(2) - return frame.f_globals.get("__name__", "") or "" + return sys._getframe(2).f_globals.get("__name__", "") or "" except Exception: return "" def register( - self, - name: str, - toolset: str, - schema: dict, - handler: Callable, - check_fn: Callable = None, - requires_env: list = None, - is_async: bool = False, - description: str = "", - emoji: str = "", - max_result_size_chars: int | float | None = None, - dynamic_schema_overrides: Callable = None, - override: bool = False, - scope: Optional[str] = None, - ): - """Register a tool. Called at module-import time by each tool file. - - ``override=True`` is an explicit opt-in for plugins that intend to - replace an existing built-in tool implementation (e.g. swap the - default browser tool for a headed-Chrome CDP backend). Without it, - registrations that would shadow an existing tool from a different - toolset are rejected to prevent accidental overwrites. - """ + self, name: str, toolset: str, schema: dict, handler: Callable, + check_fn: Callable = None, requires_env: list = None, is_async: bool = False, + description: str = "", emoji: str = "", max_result_size_chars: int | float | None = None, + dynamic_schema_overrides: Callable = None, override: bool = False, + scope: Optional[str] = None): + """Register a tool (called at import time by each tool file). ``override=True`` is an + explicit opt-in for plugins replacing a built-in implementation (e.g. a headed-Chrome + browser backend); without it, cross-toolset shadowing is rejected.""" handler_owner = self._plugin_owner_of(handler) caller_owner = self._plugin_namespace_of_module(self._caller_module()) owner = caller_owner or handler_owner @@ -809,28 +602,17 @@ class ToolRegistry: scope = self._plugin_scope_of(owner) with self._lock: target = self._slot(scope, create=True) - existing = ( - self._tools.get(name) - if scope is None - else self._merged_tools(scope).get(name) - ) + existing = (self._tools if scope is None else self._merged_tools(scope)).get(name) plugin_override_denied = ( - owner is not None and not self._plugin_override_allowed(scope, owner) - ) + owner is not None and not self._plugin_override_allowed(scope, owner)) shadows_global = ( - owner is not None - and scope is not None - and name not in target - and name in self._tools - ) + owner is not None and scope is not None + and name not in target and name in self._tools) if shadows_global: if not override: logger.error( - "Tool registration REJECTED: plugin %r attempted to " - "shadow global tool %r without override=True", - owner, - name, - ) + "Tool registration REJECTED: plugin %r attempted to shadow global tool %r " + "without override=True", owner, name) return if plugin_override_denied: raise PermissionError(_OVERRIDE_DENIED_MSG.format(owner=owner, name=name)) @@ -838,81 +620,53 @@ class ToolRegistry: if override: if plugin_override_denied: logger.error( - "Tool registration REJECTED: plugin %r attempted to " - "override built-in tool %r (existing toolset %r) without " - "operator opt-in. Set " - "plugins.entries..allow_tool_override: true " - "in config.yaml to allow it.", - owner, name, existing.toolset, - ) + "Tool registration REJECTED: plugin %r attempted to override built-in " + "tool %r (existing toolset %r) without operator opt-in. Set " + "plugins.entries..allow_tool_override: true in config.yaml " + "to allow it.", + owner, name, existing.toolset) raise PermissionError(_OVERRIDE_DENIED_MSG.format(owner=owner, name=name)) - # Explicit opt-in (or non-plugin caller): replace the tool. - # Logged at INFO so the override is auditable in agent.log. + # Explicit opt-in (or non-plugin caller): INFO so the override is auditable. logger.info( "Tool '%s': toolset '%s' overriding existing toolset '%s' " - "(override=True opt-in)", - name, toolset, existing.toolset, - ) + "(override=True opt-in)", name, toolset, existing.toolset) else: - # Reject every cross-toolset shadow, including MCP-to-MCP - # collisions. Legitimate MCP reconnect/refresh re-registers - # within the same canonical toolset and remains allowed. + # Reject every cross-toolset shadow (incl. MCP-to-MCP); same-toolset + # re-registration (MCP reconnect/refresh) stays allowed. logger.error( - "Tool registration REJECTED: '%s' (toolset '%s') would " - "shadow existing tool from toolset '%s'. Pass " - "override=True to register() if the replacement is " - "intentional, or deregister the existing tool first.", - name, toolset, existing.toolset, - ) + "Tool registration REJECTED: '%s' (toolset '%s') would shadow existing " + "tool from toolset '%s'. Pass override=True to register() if the " + "replacement is intentional, or deregister the existing tool first.", + name, toolset, existing.toolset) return target[name] = ToolEntry( - name=name, - toolset=toolset, - schema=schema, - handler=handler, - check_fn=check_fn, - requires_env=requires_env or [], - is_async=is_async, - description=description or schema.get("description", ""), - emoji=emoji, + name=name, toolset=toolset, schema=schema, handler=handler, check_fn=check_fn, + requires_env=requires_env or [], is_async=is_async, + description=description or schema.get("description", ""), emoji=emoji, max_result_size_chars=max_result_size_chars, - dynamic_schema_overrides=dynamic_schema_overrides, - ) - # Availability is derived per-tool (_toolset_has_exposable_tools), so - # this map no longer gates a toolset. It still feeds - # get_toolset_requirements -> TOOLSET_REQUIREMENTS["check_fn"], which - # banner.py reads (presence only, never called) to classify an - # already-unavailable toolset as lazy-init vs disabled. + dynamic_schema_overrides=dynamic_schema_overrides) + # Availability is derived per-tool, so this map no longer gates a toolset; it still + # feeds TOOLSET_REQUIREMENTS["check_fn"], which banner.py reads (presence only, + # never called) to classify an unavailable toolset as lazy-init vs disabled. if scope is None and check_fn and toolset not in self._toolset_checks: self._toolset_checks[toolset] = check_fn self._generation += 1 def deregister(self, name: str, *, scope: Optional[str] = None) -> None: - """Remove a tool; also drops the toolset check/aliases if it was the last - tool in its toolset (MCP nuke-and-repave on ``tools/list_changed``). - - ``scope`` selects a profile overlay explicitly (multiplexed MCP tools live - in the owning profile's overlay). Plugin callers may not name another - scope; non-plugin callers without ``scope`` target the process-global map. - - Gated by the same opt-in as ``register(override=True)``: otherwise a plugin - could deregister a tool it doesn't own and re-register over the empty slot, - skipping the override check (which only runs when an entry exists). - ``mcp-*`` toolsets are exempt — discovery repaves its own tools per refresh. - """ + """Remove a tool; drops the toolset check/aliases if it was the last in its toolset. + ``scope`` selects a profile overlay explicitly (multiplexed MCP tools live there); + plugin callers may not name another scope, non-plugin callers default to the global + map. Gated by the same opt-in as ``register(override=True)``, else a plugin could + deregister a tool it doesn't own and re-register over the empty slot (the override + check only runs when an entry exists). ``mcp-*`` is exempt — discovery repaves.""" with self._lock: caller_mod = self._caller_module() caller_owner = self._plugin_namespace_of_module(caller_mod) - caller_scope = ( - self._plugin_scope_of(caller_owner) - if caller_owner is not None - else None - ) + caller_scope = self._plugin_scope_of(caller_owner) if caller_owner is not None else None if caller_owner is not None and scope is not None and scope != caller_scope: raise PermissionError( f"Plugin module {caller_mod!r} cannot deregister tools " - "outside its own profile scope." - ) + "outside its own profile scope.") if scope is None: scope = caller_scope target = self._slot(scope) @@ -920,73 +674,44 @@ class ToolRegistry: if entry is None: if scope is not None and caller_owner is not None and name in self._tools: raise PermissionError( - f"Scoped plugin module {caller_mod!r} cannot deregister " - f"process-global tool {name!r}; register a scoped " - "override instead." - ) + f"Scoped plugin module {caller_mod!r} cannot deregister process-global " + f"tool {name!r}; register a scoped override instead.") return if not entry.toolset.startswith("mcp-"): owner = self._plugin_owner_of(entry.handler) - # Ownership binds to the plugin package root (``hermes_plugins.{name}``), - # not the exact module string: a handler defined in a submodule is - # still owned by the package, so root-module cleanup may remove it. + # Ownership binds to the plugin package root (``hermes_plugins.{name}``), not + # the exact module: a submodule's handler is still the package's to remove. same_plugin = bool(owner and caller_owner == owner) - if ( - caller_owner is not None - and not same_plugin - and not self._plugin_override_allowed( - caller_scope, caller_owner - ) - ): + if caller_owner is not None and not same_plugin \ + and not self._plugin_override_allowed(caller_scope, caller_owner): logger.error( "Tool deregistration REJECTED: plugin %r attempted to " - "remove tool %r (toolset %r) it does not own, without " - "operator opt-in. Set " - "plugins.entries.%s.allow_tool_override: true in " - "config.yaml to allow it.", - caller_mod, name, entry.toolset, caller_mod, - ) + "remove tool %r (toolset %r) it does not own, without operator opt-in. Set " + "plugins.entries.%s.allow_tool_override: true in config.yaml to allow it.", + caller_mod, name, entry.toolset, caller_mod) raise PermissionError( - f"Plugin module {caller_mod!r} cannot deregister tool " - f"{name!r} (toolset {entry.toolset!r}) without operator " - f"opt-in (allow_tool_override)." - ) + f"Plugin module {caller_mod!r} cannot deregister tool {name!r} (toolset " + f"{entry.toolset!r}) without operator opt-in (allow_tool_override).") del target[name] if scope is not None and not target: self._scoped_tools.pop(scope, None) - # Drop the toolset check and aliases if this was the last tool in - # that toolset. - toolset_still_exists = any( - e.toolset == entry.toolset - for e in self._merged_tools(scope).values() - ) - if not toolset_still_exists: + if not self._toolset_entries(entry.toolset, scope): self._toolset_checks.pop(entry.toolset, None) self._drop_toolset_aliases(entry.toolset) self._generation += 1 logger.debug("Deregistered tool: %s", name) def restore_registration( - self, - name: str, - current: ToolEntry, - previous: Optional[ToolEntry], - *, - scope: Optional[str] = None, - ) -> bool: - """Restore a host-owned registration if it is still current. - - This is the narrow inverse used by the plugin ownership ledger. The - identity check is deliberate: another plugin (or another - ``PluginManager`` in a multi-profile process) may have registered a - newer entry under the same name, in which case unloading this entry - must leave the newer entry untouched. - """ + self, name: str, current: ToolEntry, previous: Optional[ToolEntry], *, + scope: Optional[str] = None) -> bool: + """Restore a host-owned registration if it is still current (plugin ownership ledger). + The identity check is deliberate: another plugin (or ``PluginManager`` in a + multi-profile process) may have registered a newer entry under the same name, and + unloading this entry must leave that newer one untouched.""" with self._lock: target = self._slot(scope, create=True) if target.get(name) is not current: return False - if previous is None: target.pop(name, None) else: @@ -994,121 +719,80 @@ class ToolRegistry: if scope is not None and not target: self._scoped_tools.pop(scope, None) - # Rebuild the affected toolset checks from the surviving entries. - # A plugin may have replaced an entry in the same toolset, so - # simply leaving the current check_fn behind would retain stale - # plugin state after restoration. + # Rebuild affected toolset checks from survivors: a plugin may have replaced an + # entry in the same toolset, so its check_fn would otherwise linger after restore. affected_toolsets = {current.toolset} if previous is not None: affected_toolsets.add(previous.toolset) for toolset in affected_toolsets: - surviving = [ - entry for entry in self._merged_tools(scope).values() - if entry.toolset == toolset - ] - check_fn = next( - (entry.check_fn for entry in surviving if entry.check_fn), - None, - ) + surviving = self._toolset_entries(toolset, scope) + check_fn = next((entry.check_fn for entry in surviving if entry.check_fn), None) if scope is None: if check_fn is None: self._toolset_checks.pop(toolset, None) else: self._toolset_checks[toolset] = check_fn - if not surviving and not any( - entry.toolset == toolset - for entries in self._scoped_tools.values() - for entry in entries.values() - ): + in_overlays = (e for m in self._scoped_tools.values() for e in m.values()) + if not surviving and not any(e.toolset == toolset for e in in_overlays): self._drop_toolset_aliases(toolset) self._generation += 1 logger.debug("Restored tool registration: %s", name) return True - # ------------------------------------------------------------------ - # Schema retrieval - # ------------------------------------------------------------------ + # ---- Schema retrieval -------------------------------------------- def get_definitions(self, tool_names: Set[str], quiet: bool = False) -> List[dict]: - """Return OpenAI-format schemas for the requested tools whose ``check_fn`` - passes (or is absent). Probes go through the ~30 s TTL cache - (:func:`_check_fn_cached`) so ``hermes tools enable`` still lands quickly. - """ + """OpenAI-format schemas for the requested tools whose ``check_fn`` passes (or is + absent). Probes use the ~30 s TTL cache so ``hermes tools enable`` lands quickly.""" result = [] - # Per-call memo on top of the TTL: one probe per distinct check_fn per pass. check_results: Dict[Callable, bool] = {} entries_by_name = {entry.name: entry for entry in self._snapshot_entries()} for name in sorted(tool_names): entry = entries_by_name.get(name) if not entry: continue - if entry.check_fn: - if entry.check_fn not in check_results: - check_results[entry.check_fn] = _check_fn_cached(entry.check_fn) - if not check_results[entry.check_fn]: - if not quiet: - logger.debug("Tool %s unavailable (check failed)", name) - continue + if entry.check_fn and not _memo_check(entry.check_fn, check_results): + if not quiet: + logger.debug("Tool %s unavailable (check failed)", name) + continue schema_with_name = {**entry.schema, "name": entry.name} - # Runtime-dynamic overrides (e.g. delegate_task limits). The caller's - # memo (model_tools.get_tool_definitions) is keyed on config.yaml - # mtime+size, so config changes invalidate it automatically. + # Runtime-dynamic overrides (e.g. delegate_task limits); the caller's memo is + # keyed on config.yaml mtime+size, so config changes invalidate it automatically. if entry.dynamic_schema_overrides is not None: try: overrides = entry.dynamic_schema_overrides() - if isinstance(overrides, dict): - schema_with_name.update(overrides) except Exception as exc: + overrides = None logger.warning( - "dynamic_schema_overrides for tool %s raised %s; " - "using static schema", - name, exc, - ) + "dynamic_schema_overrides for tool %s raised %s; using static schema", + name, exc) + if isinstance(overrides, dict): + schema_with_name.update(overrides) result.append({"type": "function", "function": schema_with_name}) return result - # ------------------------------------------------------------------ - # Dispatch - # ------------------------------------------------------------------ + # ---- Dispatch ---------------------------------------------------- @staticmethod def _normalize_handler_result(name: str, result): - """Results must be a string or the multimodal envelope; anything else - becomes a string error so logging/hooks/budgeting/persistence never - receive values they cannot slice or size.""" + """Results must be a string or the multimodal envelope; anything else becomes a + string error so logging/hooks/budgeting/persistence never receive values they + cannot slice or size.""" if isinstance(result, str): return _bound_json_error_result(result) - if ( - isinstance(result, dict) - and result.get("_multimodal") is True - and isinstance(result.get("content"), list) - ): + if isinstance(result, dict) and result.get("_multimodal") is True \ + and isinstance(result.get("content"), list): return result - result_type = type(result).__name__ - logger.error( - "Tool %s handler returned unsupported result type: %s", - name, - result_type, - ) + logger.error("Tool %s handler returned unsupported result type: %s", name, result_type) return tool_error( f"Tool handler returned unsupported result type: {result_type}", - error_type="tool_result_contract", - tool=name, - result_type=result_type, - ) + error_type="tool_result_contract", tool=name, result_type=result_type) def dispatch( - self, - name: str, - args: dict, - *, - scope: Optional[str] = None, - **kwargs, - ) -> str | dict: - """Execute a tool handler by name: async handlers bridged via - ``_run_async()``, results normalized, every exception returned as - ``{"error": ...}``.""" + self, name: str, args: dict, *, scope: Optional[str] = None, **kwargs) -> str | dict: + """Execute a tool handler by name: async handlers bridged via ``_run_async()``, + results normalized, every exception returned as ``{"error": ...}``.""" entry = self.get_entry(name, scope=scope) if not entry: return tool_error(f"Unknown tool: {name}") @@ -1121,11 +805,8 @@ class ToolRegistry: return self._normalize_handler_result(name, result) except Exception as e: # exc_info already renders the exception, so keep the message copy bounded. - logger.exception( - "Tool %s dispatch error: %s", name, _bound_error_text(str(e)) - ) - # Sanitize so framing tokens / CDATA / fences in exception strings - # don't reach the model as structural noise. + logger.exception("Tool %s dispatch error: %s", name, _bound_error_text(str(e))) + # Sanitize so framing tokens/CDATA/fences in exception text aren't structural noise. raw = f"Tool execution failed: {type(e).__name__}: {e}" try: from model_tools import _sanitize_tool_error @@ -1134,45 +815,36 @@ class ToolRegistry: sanitized = raw # defensive: never let the sanitizer block error propagation return tool_error(sanitized) - # ------------------------------------------------------------------ - # Query helpers (replace redundant dicts in model_tools.py) - # ------------------------------------------------------------------ + # ---- Query helpers ----------------------------------------------- + + def _attr(self, name: str, attr: str): + return getattr(self.get_entry(name), attr, None) def get_max_result_size(self, name: str, default: int | float | None = None) -> int | float: """Return per-tool max result size, or *default* (or global default).""" - entry = self.get_entry(name) - if entry and entry.max_result_size_chars is not None: - return entry.max_result_size_chars + size = self._attr(name, "max_result_size_chars") + if size is not None: + return size if default is not None: return default from tools.budget_config import DEFAULT_RESULT_SIZE_CHARS return DEFAULT_RESULT_SIZE_CHARS def get_all_tool_names(self) -> List[str]: - """Return sorted list of all registered tool names.""" return sorted(entry.name for entry in self._snapshot_entries()) def get_schema(self, name: str) -> Optional[dict]: - """Return a tool's raw schema dict, bypassing check_fn filtering. - - Useful for token estimation and introspection where availability - doesn't matter — only the schema content does. - """ - entry = self.get_entry(name) - return entry.schema if entry else None + """Raw schema dict, bypassing check_fn filtering (token estimates, introspection).""" + return self._attr(name, "schema") def get_toolset_for_tool(self, name: str) -> Optional[str]: - """Return the toolset a tool belongs to, or None.""" - entry = self.get_entry(name) - return entry.toolset if entry else None + return self._attr(name, "toolset") def get_emoji(self, name: str, default: str = "⚡") -> str: """Return the emoji for a tool, or *default* if unset.""" - entry = self.get_entry(name) - return (entry.emoji if entry and entry.emoji else default) + return self._attr(name, "emoji") or default def get_tool_to_toolset_map(self) -> Dict[str, str]: - """Return ``{tool_name: toolset_name}`` for every registered tool.""" return {entry.name: entry.toolset for entry in self._snapshot_entries()} def is_toolset_available(self, toolset: str) -> bool: @@ -1180,68 +852,57 @@ class ToolRegistry: return self._toolset_has_exposable_tools(toolset, self._snapshot_entries()) def check_toolset_requirements(self) -> Dict[str, bool]: - """Return ``{toolset: available_bool}`` for every toolset.""" entries = self._snapshot_entries() return { toolset: self._toolset_has_exposable_tools(toolset, entries) - for toolset in sorted({entry.toolset for entry in entries}) - } + for toolset in sorted(self._grouped(entries))} def get_available_toolsets(self) -> Dict[str, dict]: """Return toolset metadata for UI display.""" - toolsets: Dict[str, dict] = {} entries = self._snapshot_entries() - for entry in entries: - info = toolsets.get(entry.toolset) - if info is None: - info = toolsets[entry.toolset] = { - "available": self._toolset_has_exposable_tools(entry.toolset, entries), - "tools": [], - "description": "", - "requirements": [], - } - info["tools"].append(entry.name) - _extend_unique(info["requirements"], entry.requires_env or []) + toolsets: Dict[str, dict] = {} + for toolset, members in self._grouped(entries).items(): + toolsets[toolset] = { + "available": self._toolset_has_exposable_tools(toolset, entries), + "tools": [entry.name for entry in members], + "description": "", + "requirements": _unique_env(members)} return toolsets def get_toolset_requirements(self) -> Dict[str, dict]: """Build a TOOLSET_REQUIREMENTS-compatible dict for backward compat.""" - result: Dict[str, dict] = {} entries, toolset_checks = self._snapshot_state() - for entry in entries: - info = result.setdefault(entry.toolset, { - "name": entry.toolset, - "env_vars": [], - "check_fn": toolset_checks.get(entry.toolset), + result: Dict[str, dict] = {} + for toolset, members in self._grouped(entries).items(): + result[toolset] = { + "name": toolset, + "env_vars": _unique_env(members), + "check_fn": toolset_checks.get(toolset), "setup_url": None, - "tools": [], - }) - _extend_unique(info["tools"], [entry.name]) - _extend_unique(info["env_vars"], entry.requires_env) + "tools": [entry.name for entry in members]} return result def check_tool_availability(self, quiet: bool = False): """Return (available_toolsets, unavailable_info) like the old function.""" - available = [] - unavailable = [] + available, unavailable = [], [] entries = self._snapshot_entries() - for ts in sorted({entry.toolset for entry in entries}): - ts_entries = [entry for entry in entries if entry.toolset == ts] + groups = self._grouped(entries) + for ts in sorted(groups): if self._toolset_has_exposable_tools(ts, entries): available.append(ts) else: unavailable.append({ - "name": ts, - "env_vars": ts_entries[0].requires_env if ts_entries else [], - "tools": [entry.name for entry in ts_entries], - }) + "name": ts, "env_vars": groups[ts][0].requires_env, + "tools": [entry.name for entry in groups[ts]]}) return available, unavailable -def _extend_unique(target: list, items) -> None: - for item in items: - if item not in target: - target.append(item) +def _unique_env(entries: List[ToolEntry]) -> list: + """Union of ``requires_env`` across *entries*, first-seen order, no duplicates.""" + out: list = [] + for entry in entries: + out.extend(v for v in (entry.requires_env or []) if v not in out) + return out # Module-level singleton diff --git a/tools/schema_sanitizer.py b/tools/schema_sanitizer.py index c6a0f30665..39a3d6c937 100644 --- a/tools/schema_sanitizer.py +++ b/tools/schema_sanitizer.py @@ -1,21 +1,10 @@ """Sanitize tool JSON schemas for broad LLM-backend compatibility. -Some backends are strict about JSON Schema shapes that OpenAI/Anthropic/most -cloud providers silently accept — llama.cpp's ``json-schema-to-grammar`` fails -the whole request (``Unrecognized schema: "object"``), Anthropic rejects -nullable ``anyOf`` at the top of ``input_schema``, Fireworks rejects ``default`` -beside ``$ref``, OpenAI's Codex backend rejects top-level combinators. Known -hostile constructs: - -* ``{"type": "object"}`` with no ``properties``. -* A bare string (``"object"``) where a schema dict belongs (malformed MCP output). -* ``"type": ["string", "null"]`` array types. -* ``anyOf``/``oneOf`` unions whose only purpose is to permit ``null``. -* ``default`` (etc.) alongside ``$ref`` — e.g. ``{"$ref": "#/$defs/Foo", "default": null}``. - -This module walks the final tool schema tree (after MCP normalization and any -per-tool dynamic rebuilds) and fixes those in place on a deep copy. It is -deliberately conservative: it only modifies shapes the backend couldn't use. +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. """ from __future__ import annotations @@ -28,9 +17,8 @@ from typing import Any, Callable logger = logging.getLogger(__name__) -# Anthropic (and Bedrock/Vertex/Azure fronting it) reject tool input schemas -# whose property keys don't match this pattern; one bad key anywhere in the -# tools array 400s the entire request (Cloudflare's MCP ships 61 such keys). +# 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_.-]") @@ -43,18 +31,24 @@ def _empty_object() -> dict: return {"type": "object", "properties": {}} +def _rewrite(schema: Any, fn: Callable[[dict], Any]) -> Any: + """Bottom-up map over a schema tree: lists/dicts recurse, then *fn* sees each dict.""" + if isinstance(schema, list): + return [_rewrite(item, fn) for item in schema] + if not isinstance(schema, dict): + return schema + return fn({k: _rewrite(v, fn) for k, v in schema.items()}) + + def sanitize_property_key(key: str) -> str: """Deterministically map an arbitrary property key to a conforming one.""" return _PROP_KEY_BAD_CHARS.sub("_", key)[:64] or "param" def _rename_property_keys(props: dict, path: str) -> dict[str, str]: - """Return {original_key: conforming_key} for one properties dict. - - Identity entries are omitted. Deterministic (insertion order, numeric - suffixes on collision) so the model-visible schema and the dispatch-time - reverse map computed from the registry's original schema always agree. - """ + """{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.""" renames: dict[str, str] = {} taken = {k for k in props if _PROP_KEY_RE.match(k)} for key in props: @@ -70,17 +64,13 @@ def _rename_property_keys(props: dict, path: str) -> dict[str, str]: renames[key] = candidate logger.debug( "schema_sanitizer[%s]: renamed property key %r -> %r " - "(provider key-pattern compat)", path, key, candidate, - ) + "(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 (unsanitized) registry schema. Recurses - into object values and array items; unknown keys pass through untouched. - """ + """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") @@ -96,48 +86,38 @@ def unrename_tool_args(params_schema: Any, args: Any) -> Any: 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 - ] + unrename_tool_args(subschema["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]: - """Return a deep-copied ``tools`` list (OpenAI format) with each tool's - parameter schema sanitized; callers may mutate the result freely.""" + """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] def _sanitize_single_tool(tool: dict) -> dict: - """Deep-copy and sanitize a single OpenAI-format tool entry.""" out = copy.deepcopy(tool) fn = out.get("function") if isinstance(out, dict) else None if not isinstance(fn, dict): return out - params = fn.get("parameters") if not isinstance(params, dict): # missing / non-dict → minimal valid shape fn["parameters"] = _empty_object() 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 = _empty_object() - else: - if top.get("type") != "object": - top["type"] = "object" - if not isinstance(top.get("properties"), dict): - top["properties"] = {} - # Collapse nullable unions the recursive pass leaves intact (it only - # handles the array-form ``type: [X, "null"]``); keep ``nullable: true`` so - # runtime coercion (``model_tools._schema_allows_null``) still maps a - # model-emitted ``"null"`` string to Python ``None``. + top = {} + top["type"] = "object" + if not isinstance(top.get("properties"), dict): + top["properties"] = {} + # The recursive pass only handles array-form ``type: [X, "null"]``; collapse anyOf unions + # here, keeping ``nullable: true`` so ``model_tools._schema_allows_null`` still coerces. top = strip_nullable_unions(top, keep_nullable_hint=True) top = _strip_top_level_combinators(top, path=name) fn["parameters"] = _strip_ref_siblings(top) @@ -149,30 +129,22 @@ _REF_FORBIDDEN_SIBLINGS = frozenset({"default"}) def _strip_ref_siblings(node: Any) -> Any: - """Recursively drop forbidden sibling keywords from nodes carrying ``$ref`` - (Fireworks: ``keyword(s) ['default'] not allowed at the same level as $ref``).""" - if isinstance(node, list): - return [_strip_ref_siblings(item) for item in node] - if not isinstance(node, dict): - return node - out = {key: _strip_ref_siblings(value) for key, value in node.items()} - if "$ref" in out: - for key in _REF_FORBIDDEN_SIBLINGS: - out.pop(key, None) - return out + """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) + return out + return _rewrite(node, strip) _TOP_LEVEL_FORBIDDEN_KEYS = ("allOf", "anyOf", "oneOf", "enum", "not") def _strip_top_level_combinators(params: dict, *, path: str = "") -> dict: - """Drop combinator keywords from the TOP level of a parameters schema only. - - OpenAI's Codex backend rejects ``oneOf/anyOf/allOf/enum/not`` at the top - level. They are usually conditional-required hints; dropping them does not - change which argument values are valid (handlers re-validate). Nested - combinators are preserved. - """ + """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.""" if not isinstance(params, dict): return params out = dict(params) @@ -181,8 +153,7 @@ def _strip_top_level_combinators(params: dict, *, path: str = "") -> dict: logger.debug( "schema_sanitizer[%s]: stripped top-level %r combinator " "from tool parameters (strict-backend compat)", - path, key, - ) + path, key) out.pop(key, None) return out @@ -201,135 +172,83 @@ def _carry_union_meta(outer: dict, replacement: dict, *, skip_default_on_ref: bo 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"}], "default": null}``; - Anthropic rejects the null branch, and optionality is already expressed by - the parent's ``required``. Only collapses when a null branch was dropped - AND exactly one non-null branch survives. Outer metadata is carried over. - ``keep_nullable_hint`` sets ``nullable: true`` on the replacement for - downstream consumers (runtime ``"null"`` → ``None`` coercion). - """ - if isinstance(schema, list): - return [strip_nullable_unions(item, keep_nullable_hint=keep_nullable_hint) for item in schema] - if not isinstance(schema, dict): - return schema - - stripped = { - k: strip_nullable_unions(v, keep_nullable_hint=keep_nullable_hint) - for k, v in schema.items() - } - for key in _UNION_KEYS: - variants = stripped.get(key) - if not isinstance(variants, list): - continue - non_null = [item for item in variants if not _is_null_branch(item)] - if len(non_null) == 1 and len(non_null) != len(variants): - replacement = dict(non_null[0]) if isinstance(non_null[0], dict) else {} - if keep_nullable_hint: - replacement.setdefault("nullable", True) - _carry_union_meta(stripped, replacement, skip_default_on_ref=True) - return strip_nullable_unions(replacement, keep_nullable_hint=keep_nullable_hint) - return stripped +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. + ``keep_nullable_hint`` sets ``nullable: true`` for runtime ``"null"`` → ``None`` coercion.""" + def collapse(stripped: dict) -> Any: + for key in _UNION_KEYS: + variants = stripped.get(key) + if not isinstance(variants, list): + continue + non_null = [item for item in variants if not _is_null_branch(item)] + if len(non_null) == 1 and len(non_null) != len(variants): + replacement = dict(non_null[0]) if isinstance(non_null[0], dict) else {} + if keep_nullable_hint: + replacement.setdefault("nullable", True) + _carry_union_meta(stripped, replacement, skip_default_on_ref=True) + return _rewrite(replacement, collapse) + return stripped + return _rewrite(schema, collapse) _CONST_PRIMITIVE_TYPES: dict[type, str] = { - bool: "boolean", - int: "integer", - float: "number", - str: "string", -} + bool: "boolean", int: "integer", float: "number", str: "string"} def _const_branch_type(branch: Any) -> str | None: - """JSON-Schema primitive type of a pure ``const`` branch, else None. - - Qualifies when the dict carries a primitive ``const`` and any declared - ``type`` matches it; ``title``/``description`` are allowed, any other - constraining keyword disqualifies. - """ - if not isinstance(branch, dict) or "const" not in branch: + """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.""" + if not isinstance(branch, dict) or "const" not in branch \ + or set(branch) - {"const", "type", "title", "description"}: return None - if set(branch) - {"const", "type", "title", "description"}: - return None - value = branch["const"] - # ``type(value) is`` (not isinstance): bool is a subclass of int. - json_type = _CONST_PRIMITIVE_TYPES.get(type(value)) - if json_type is None: - return None - declared = branch.get("type") - if declared is not None and declared != json_type: - return None - return json_type + # ``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 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). MCP - servers generated from Rust/TS union types emit - ``{"anyOf": [{"const": "red"}, {"const": "green"}]}``; strict backends - mishandle these while ``{"type": "string", "enum": [...]}`` is universal. - - 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`` - (``strip_nullable_unions`` only handles single-non-null unions, so - null+multi-const unions land here). Enum order preserves branch order; - outer metadata is carried over; input is never mutated. - """ - if isinstance(schema, list): - return [collapse_const_unions(item) for item in schema] - if not isinstance(schema, dict): - return schema - - out = {k: collapse_const_unions(v) for k, v in schema.items()} - for key in _UNION_KEYS: - variants = out.get(key) - if not isinstance(variants, list) or not variants: - continue - null_branches = [ - item for item in variants if _is_null_branch(item) and "const" not in item - ] - const_branches = [item for item in variants if item not in null_branches] - if len(null_branches) > 1 or not const_branches: - continue - branch_types = {_const_branch_type(item) for item in const_branches} - if len(branch_types) != 1 or None in branch_types: - continue - replacement: dict = { - "type": branch_types.pop(), - "enum": [item["const"] for item in const_branches], - } - if null_branches: - replacement["nullable"] = True - _carry_union_meta(out, replacement, skip_default_on_ref=False) - return replacement - return out + """Collapse ``anyOf``/``oneOf`` unions of same-typed consts to ``enum`` (ported from + block/goose ``tool_schema_normalize.rs``, Apache-2.0). Rust/TS MCP servers emit + ``{"anyOf": [{"const": "red"}, {"const": "green"}]}``, which strict backends mishandle. + Applies only when EVERY non-null branch is a pure ``const`` of one primitive type + (``bool`` never merges with ``integer``); one ``{"type": "null"}`` branch is tolerated + and recorded as ``nullable: true``. Branch order is kept; outer metadata carried over; + input never mutated.""" + def collapse(out: dict) -> Any: + for key in _UNION_KEYS: + variants = out.get(key) + if not isinstance(variants, list) or not variants: + continue + null_branches = [i for i in variants if _is_null_branch(i) and "const" not in i] + const_branches = [item for item in variants if item not in null_branches] + if len(null_branches) > 1 or not const_branches: + continue + branch_types = {_const_branch_type(item) for item in const_branches} + if len(branch_types) != 1 or None in branch_types: + continue + replacement: dict = { + "type": branch_types.pop(), "enum": [item["const"] for item in const_branches]} + if null_branches: + replacement["nullable"] = True + _carry_union_meta(out, replacement, skip_default_on_ref=False) + return replacement + return out + return _rewrite(schema, collapse) _BARE_TYPE_NAMES = frozenset({"object", "string", "number", "integer", "boolean", "array", "null"}) -# Sibling keywords whose values are NOT schemas: recursing would mistake literal -# strings like "path" for bare-string schemas. Passed through unchanged -# (``required`` remapped through property renames). +# Keys whose values are NOT schemas (recursing would treat "path" as a bare-string 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*. - - Several backends reject array types (llama.cpp's grammar generator; Gemini - via OpenAI-compatible transports 400s). Per the AI-SDK behavior: 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`` or + the object fallback. Ported from anomalyco/opencode#31877.""" has_null = "null" in value non_null = [t for t in value if isinstance(t, str) and t != "null"] if len(non_null) == 1: @@ -344,45 +263,29 @@ def _normalize_type_array(value: list, out: dict) -> None: def _sanitize_node(node: Any, path: str) -> Any: - """Recursively sanitize a JSON-Schema fragment. - - - Bare-string schema values become ``{"type": }`` (unknown strings - become a permissive object schema rather than something backends reject). - - Object-typed nodes gain ``properties: {}`` (llama.cpp can't constrain a - free-form object). - - ``type`` arrays are normalized (see ``_normalize_type_array``). - - Recurses into ``properties``, ``items``, ``additionalProperties``, - ``anyOf``/``oneOf``/``allOf`` and ``$defs``/``definitions``; property - keys are renamed to the provider-safe pattern and ``required`` follows. - - ``required`` entries that don't exist in ``properties`` are pruned - (malformed MCP schemas; built-in/plugin tools skip the MCP-level check). - """ + """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.""" 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, - ) + "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, - ) + "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 are computed up front so ``required`` can be remapped even when - # it precedes ``properties`` in the source dict. + # 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") - out: dict = {} for key, value in node.items(): if key == "type" and isinstance(value, list): @@ -391,11 +294,9 @@ def _sanitize_node(node: Any, path: str) -> Any: 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() - } + for sub_k, sub_v in value.items()} elif key in {"items", "additionalProperties"}: - # Bool ``additionalProperties`` is valid and widely accepted; - # ``items: true/false`` is non-standard but preserved rather than dropped. + # 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)] @@ -406,7 +307,6 @@ def _sanitize_node(node: Any, path: str) -> Any: out[key] = copy.deepcopy(value) if isinstance(value, (list, dict)) else value else: 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"] = {} @@ -420,60 +320,49 @@ def _sanitize_node(node: Any, path: str) -> Any: 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"}) -def _reactive_strip(tools: list[dict], strip_node: Callable[[dict], int], log_msg: str) -> tuple[list[dict], int]: - """Walk every tool's parameters in place, applying *strip_node* to each dict - node (it returns how many keywords it removed). Handles OpenAI format - (``{"function": {"parameters": ...}}``) and Responses format - (``{"name": ..., "parameters": ...}`` — codex_responses mode, xAI, etc.). - Returns ``(tools, stripped_count)`` — the same list reference.""" +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.""" + 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) + + +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 stripped = 0 - - def _walk(node: Any) -> None: - nonlocal stripped - if isinstance(node, dict): - stripped += strip_node(node) - for v in node.values(): - _walk(v) - elif isinstance(node, list): - for item in node: - _walk(item) - for tool in tools: if not isinstance(tool, dict): continue fn = tool.get("function") - if isinstance(fn, dict) and isinstance(fn.get("parameters"), dict): - _walk(fn["parameters"]) - continue - if isinstance(tool.get("parameters"), dict): - _walk(tool["parameters"]) - + params = fn.get("parameters") if isinstance(fn, dict) else None + if not isinstance(params, dict): + params = tool.get("parameters") + if isinstance(params, dict): + stripped += sum(strip_node(node) for node in _dict_nodes(params)) if stripped: logger.info(log_msg, stripped) return tools, stripped def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]: - """Strip ``pattern``/``format`` keywords from tool schemas, in place. - - Reactive: invoked only after llama.cpp's grammar converter rejected a - schema with HTTP 400. Its regex engine supports a small ECMAScript subset - (no ``\\d``/``\\w``/``\\s``) and most ``format`` values; cloud providers rely - on these as prompting hints, so they stay in the default schema. - - Only strips as a sibling of ``type``/combinators (i.e. on schema nodes), so - a property literally *named* ``pattern`` (``search_files``) is untouched — - property names live inside ``properties``, not beside ``type``. - """ + """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.""" def _strip(node: dict) -> int: if not ("type" in node or "anyOf" in node or "oneOf" in node or "allOf" in node): return 0 @@ -481,31 +370,23 @@ def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]: for k in hits: node.pop(k, None) return len(hits) - return _reactive_strip( tools, _strip, "schema_sanitizer: stripped %d pattern/format keyword(s) from " - "tool schemas (llama.cpp grammar-parse recovery)", - ) + "tool schemas (llama.cpp grammar-parse recovery)") def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]: - """Strip ``enum`` keywords whose string values contain ``/``, in place. - - xAI's ``/v1/responses`` and ``/v1/chat/completions`` compile schemas to a - grammar that rejects ``/`` in enum values (HTTP 400 before any token) — - typically MCP enums of HuggingFace model IDs or owner/name env IDs. The - constraint is a prompting hint only; the model still sees the description. - """ + """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.""" 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) return 1 return 0 - return _reactive_strip( tools, _strip, "schema_sanitizer: stripped %d enum keyword(s) containing '/' " - "from tool schemas (xAI Responses grammar-compile recovery)", - ) + "from tool schemas (xAI Responses grammar-compile recovery)") diff --git a/tools/self_repo_guard.py b/tools/self_repo_guard.py index 4f3a3c9bd7..32b72a1fba 100644 --- a/tools/self_repo_guard.py +++ b/tools/self_repo_guard.py @@ -14,58 +14,45 @@ from tools.approval import ( _bash_exec_payload, _deobfuscate_shell_word_for_detection, _iter_shell_command_starts, - _read_shell_word, -) + _read_shell_word) - -# bisect is included: it drives repeated checkouts of the running root — the -# exact module-version-skew hazard this guard exists for. +# bisect drives repeated checkouts of the running root — the exact skew hazard guarded here. _WORKTREE_MUTATIONS = frozenset({ - "checkout", "switch", "rebase", "merge", "pull", "restore", "clean", - "cherry-pick", "revert", "bisect", -}) + "checkout", "switch", "rebase", "merge", "pull", "restore", "clean", "cherry-pick", "revert", + "bisect"}) _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 in _inspect_git); -# listing them only avoids a pointless `git config --get alias.` -# subprocess for `stash list`, `reset --soft`, `clean -n`, `restore --staged`. +# `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`. _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", "ls-files", "ls-remote", "ls-tree", - "maintenance", "merge-base", "mv", "notes", "push", "range-diff", "reflog", - "remote", "repack", "replace", "reset", "restore", "rev-list", "rev-parse", - "rm", "shortlog", "show", "show-ref", "stash", "status", "submodule", "tag", - "worktree", -}) + "add", "am", "apply", "blame", "branch", "bundle", "cat-file", "clean", "clone", "commit", + "config", "describe", "diff", "fetch", "format-patch", "grep", "help", "init", "log", + "ls-files", "ls-remote", "ls-tree", "maintenance", "merge-base", "mv", "notes", "push", + "range-diff", "reflog", "remote", "repack", "replace", "reset", "restore", "rev-list", + "rev-parse", "rm", "shortlog", "show", "show-ref", "stash", "status", "submodule", "tag", + "worktree"}) _SHELL_EXECUTABLES = frozenset({"bash", "dash", "ksh", "sh", "zsh"}) _ASSIGNMENT_RE = re.compile(r"[A-Za-z_][A-Za-z0-9_]*=(.*)", re.DOTALL) _RESET_HARD_RE = re.compile(r"--h(?:a(?:r(?:d)?)?)?\Z") +# ``<<-`` opener + optional blanks + a quoted delimiter (closing quote required) or a bare word. +_HEREDOC_OPENER_RE = re.compile( + r"<<(?P-?)[ \t]*(?:(?P['\"])(?P.*?)(?P=q)|(?!['\"])(?P[^\s;|&<>]*))") _NO_OPTIONS: frozenset[str] = frozenset() -# Wrapper executables that are skipped to reach the real command, mapped to -# the options that consume a following argument. +# Wrapper executables skipped to reach the real command -> options that consume an argument. _WRAPPER_OPTIONS_WITH_ARG: dict[str, frozenset[str]] = { "sudo": frozenset({ "-C", "--chdir", "-c", "--close-from", "-g", "--group", "-h", "--host", - "-p", "--prompt", "-R", "--chroot", "-T", "--command-timeout", "-u", "--user", - }), - "env": frozenset({ - "-a", "--argv0", "-C", "--chdir", "-S", "--split-string", "-u", "--unset", - }), - "command": _NO_OPTIONS, - "builtin": _NO_OPTIONS, + "-p", "--prompt", "-R", "--chroot", "-T", "--command-timeout", "-u", "--user"}), + "env": frozenset({"-a", "--argv0", "-C", "--chdir", "-S", "--split-string", "-u", "--unset"}), + "command": _NO_OPTIONS, "builtin": _NO_OPTIONS, "nohup": _NO_OPTIONS, "setsid": _NO_OPTIONS, "exec": frozenset({"-a"}), - "nohup": _NO_OPTIONS, - "setsid": _NO_OPTIONS, - "time": frozenset({"-f", "--format", "-o", "--output"}), -} + "time": frozenset({"-f", "--format", "-o", "--output"})} _MAX_RECURSION = 4 # git global options that consume the next argument (-C/--work-tree/-c are acted on). -_GIT_GLOBAL_OPTIONS_WITH_ARG = frozenset({ - "-C", "-c", "--work-tree", "--git-dir", "--namespace", "--exec-path", -}) +_GIT_GLOBAL_OPTIONS_WITH_ARG = frozenset( + {"-C", "-c", "--work-tree", "--git-dir", "--namespace", "--exec-path"}) @dataclass @@ -118,9 +105,7 @@ def _shell_words_at(command: str, start: int) -> list[str]: cursor = start for _ in range(64): word_start, word_end, raw_word = _read_shell_word(command, cursor) - if word_start == word_end: - break - if words and "\n" in command[cursor:word_start]: + if word_start == word_end or (words and "\n" in command[cursor:word_start]): break words.append(_deobfuscate_shell_word_for_detection(raw_word)) cursor = word_end @@ -128,10 +113,7 @@ def _shell_words_at(command: str, start: int) -> list[str]: def _consume_options( - words: list[str], - start: int, - options_with_arg: frozenset[str] = _NO_OPTIONS, -) -> int: + 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): @@ -140,11 +122,7 @@ def _consume_options( return index + 1 if not option.startswith("-") or option == "-": break - option_name = option.split("=", 1)[0] - if "=" not in option and option_name in options_with_arg: - index += 2 - else: - index += 1 + index += 2 if "=" not in option and option in options_with_arg else 1 return index @@ -152,14 +130,12 @@ def _command_parts(words: list[str]) -> tuple[dict[str, str], str | None, list[s """Split leading VAR=value assignments and wrappers off -> (env, executable, args).""" env: dict[str, str] = {} index = 0 - while index < len(words): if _ASSIGNMENT_RE.fullmatch(words[index]): name, value = words[index].split("=", 1) env[name] = value index += 1 continue - executable = _executable_name(words[index]) wrapper_options = _WRAPPER_OPTIONS_WITH_ARG.get(executable) if wrapper_options is None: @@ -168,7 +144,6 @@ def _command_parts(words: list[str]) -> tuple[dict[str, str], str | None, list[s if executable == "command" and words[index + 1 : index + 2] in (["-v"], ["-V"]): return env, None, [] index = _consume_options(words, index + 1, wrapper_options) - return env, None, [] @@ -177,69 +152,46 @@ def _scope_keys(command: str, starts: list[int]) -> dict[int, tuple[int, ...]]: contexts = [_ShellContext("root", -1)] scopes: dict[int, tuple[int, ...]] = {} cursor = 0 - for start in sorted(set(starts)): while cursor < start: context = contexts[-1] quote = context.quote char = command[cursor] - + nested = len(contexts) > 1 if quote == "'": if char == "'": context.quote = None + elif char == "\\" and cursor + 1 < start: cursor += 1 - continue - # Unquoted or inside double quotes: substitutions still open scopes. - if char == "\\" and cursor + 1 < start: - cursor += 2 - continue - if quote == '"': - if char == '"': - context.quote = None - cursor += 1 - continue - elif char in {"'", '"'}: + elif quote == '"' and char == '"': + context.quote = None + elif quote is None and char in {"'", '"'}: context.quote = char - cursor += 1 - continue - if command.startswith("$(", cursor): + # Unquoted or inside double quotes: substitutions still open scopes. + elif command.startswith("$(", cursor): contexts.append(_ShellContext("$(", cursor)) - cursor += 2 - continue - if quote is None: - if char == "(": - contexts.append(_ShellContext("(", cursor)) - cursor += 1 - continue - if char == ")" and len(contexts) > 1 and contexts[-1].kind in {"(", "$("}: - contexts.pop() - cursor += 1 - continue - if char == "`": - if quote is None and len(contexts) > 1 and contexts[-1].kind == "`": + 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 {"(", "$("}: + contexts.pop() + elif char == "`": + if quote is None and nested and contexts[-1].kind == "`": contexts.pop() else: 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: - index = start - 1 - saw_newline = False - while index >= 0 and command[index].isspace(): - saw_newline = saw_newline or command[index] == "\n" - index -= 1 - if index < 0: - return "\n" if saw_newline else None - if index > 0 and command[index - 1 : index + 1] in {"&&", "||"}: - return command[index - 1 : index + 1] - if command[index] in {";", "|", "&", "(", "{"}: - return command[index] - return "\n" if saw_newline else None + head = command[:start].rstrip() + if head[-2:] in {"&&", "||"}: + return head[-2:] + if head[-1:] in {";", "|", "&", "(", "{"}: + return head[-1:] + return "\n" if "\n" in command[len(head):start] else None def _cd_target(executable: str, args: list[str], cwd: Path) -> Path | None: @@ -253,24 +205,19 @@ 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 '