refactor(tools): compact registry/schema_sanitizer/self_repo_guard/plugin_guard/project_tools/path_security

This commit is contained in:
Teknium
2026-09-02 22:15:46 -07:00
parent 113f04616b
commit b494260f14
6 changed files with 310 additions and 634 deletions

View File

@@ -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/``..`` normalised), else None."""
try:
path.resolve().relative_to(root.resolve())
except (ValueError, OSError) as exc:

View File

@@ -1,20 +1,12 @@
#!/usr/bin/env python3
"""Plugin Guard — security scanner for externally-installed plugins.
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
Reuses the ``tools/skills_guard.py`` engine for ``hermes plugins install``/``update``.
Plugins run Python in-process but are *expected* to read their own env keys, call
provider HTTP 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 is
exempt on *code* files; plugin-sized structural limits; VCS/venv noise skipped.
Verdict policy: ``safe`` installs; ``caution`` needs confirmation; ``dangerous`` is
blocked and ``--force`` does NOT override.
"""
@@ -41,14 +33,12 @@ EXCLUDED_DIRS = {
".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 the normal
# documented plugin pattern (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.
# skills_guard pattern ids exempt on code files (every legitimate provider plugin
# exhibits them); they still apply in full to docs/config files.
CODE_EXEMPT_PATTERN_IDS = {
"python_environ_get_secret",
"python_getenv_secret",
@@ -60,24 +50,22 @@ CODE_EXEMPT_PATTERN_IDS = {
"env_exfil_fetch",
"env_exfil_curl",
"env_exfil_wget",
# Agent-facing instruction patterns are meaningless inside code
# (docstrings/comments about prompts trip them constantly).
# Agent-facing instruction patterns are meaningless inside code (docstrings
# 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.
# Plugins legitimately write their settings into config.yaml during post_setup
# and base64-encode credentials (HTTP 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.
# Plugin severity remaps: a bundled binary is warn-tier (repos occasionally vendor
# one legitimately); a mere ``~/.hermes/.env`` mention is how READMEs tell users where
# keys go (READING it still trips ``read_secrets_file``, critical); ``curl | sh``
# install instructions are common in READMEs — caution, not an unoverridable block.
SEVERITY_REMAP = {
"binary_file": "high",
"hermes_env_access": "medium",
@@ -169,11 +157,7 @@ def _check_plugin_structure(plugin_dir: Path) -> List[Finding]:
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():
@@ -205,15 +189,8 @@ def scan_plugin(plugin_dir: Path, source: str = "") -> ScanResult:
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``.
"""
def should_allow_plugin_install(result: ScanResult, force: bool = False) -> Tuple[Optional[bool], str]:
"""Map a verdict to ``(allowed, reason)``: True installs, None needs confirmation, False is blocked."""
n = len(result.findings)
if result.verdict == "safe":
return True, "Allowed (clean scan)"

View File

@@ -1,15 +1,11 @@
#!/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.
Projects (per-profile ``projects.db``) are the named workspaces the desktop sidebar
groups sessions into; creating/switching one is an explicit tool call, never a side
effect of a terminal ``cd``. GUI-only: the `project` toolset stays off
``_HERMES_CORE_TOOLS`` and is folded in by the desktop/TUI gateway, which also wires
``set_project_workspace_callback`` so the live session's cwd and sidebar follow.
"""
import json
@@ -18,10 +14,9 @@ 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 at session wiring: ``(task_id, primary_path, project_name)``
# re-anchors that session's workspace. ``None`` in CLI/messaging contexts — the DB
# write still happens; there is just no live GUI session to move.
_workspace_callback: Optional[Callable[[str, str, str], None]] = None
@@ -66,6 +61,12 @@ 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
@@ -76,13 +77,7 @@ def project_list(task_id: Optional[str] = None) -> str:
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,
}
{"id": p.id, "slug": p.slug, "name": p.name, "primary_path": _primary_path(p), "active": p.id == active}
for p in projects
],
})
@@ -103,9 +98,8 @@ def project_create(name: str, path: Optional[str] = None, task_id: Optional[str]
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: re-activating the folder's project beats minting a
# duplicate (duplicates render N identical sidebar subtrees).
pdb.set_active(conn, existing.id)
proj = existing
else:
@@ -117,11 +111,7 @@ def project_create(name: str, path: Optional[str] = None, task_id: Optional[str]
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:
@@ -132,28 +122,24 @@ def project_switch(project: str, task_id: Optional[str] = None) -> str:
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",

View File

@@ -1,17 +1,10 @@
"""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.
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 keeping parallel data structures. Import chain (cycle-safe):
tools/registry.py imports nothing from model_tools or tool files; tools/*.py import
tools.registry at module level; model_tools.py imports both; run_agent/cli import that.
"""
import ast
@@ -42,9 +35,7 @@ def _bound_error_text(text: str) -> str:
if len(text) <= _MAX_TOOL_ERROR_CHARS:
return text
logger.debug(
"tool error body truncated for context (%d chars): %s",
len(text),
text[:_MAX_LOGGED_ERROR_CHARS],
"tool error body truncated for context (%d chars): %s", len(text), text[:_MAX_LOGGED_ERROR_CHARS],
)
return text[:_MAX_TOOL_ERROR_CHARS] + _TOOL_ERROR_TRUNCATION_MARKER
@@ -52,10 +43,8 @@ def _bound_error_text(text: str) -> str:
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.
Handlers that ``json.dumps({"error": str(exc)})`` directly bypass ``tool_error``'s
cap; applied at the dispatch boundary so no tool can stack unbounded errors across retries.
"""
if len(result) <= _MAX_TOOL_ERROR_CHARS or '"error"' not in result:
return result
@@ -73,27 +62,21 @@ 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.
"""True when the module body (or a module-level ``for``) calls ``registry.register(...)``.
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.
Only module-body statements count, so helpers that register 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")
@@ -105,26 +88,20 @@ def _module_registers_tools(module_path: Path) -> bool:
tree = ast.parse(source, filename=str(module_path))
except 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.
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 cache 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
@@ -143,11 +120,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 (cached[0], cached[1]) == stat_key:
registers = bool(cached[2])
else:
registers = _module_registers_tools(path)
@@ -173,8 +146,7 @@ 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"
@@ -211,8 +183,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 with ``is``)."""
name: str
toolset: str
@@ -224,10 +195,9 @@ 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 returning schema overrides merged (shallow) onto the base schema
# at every get_definitions() call — for fields tracking runtime config (e.g.
# delegate_task's description reflects delegation.max_concurrent_children).
dynamic_schema_overrides: Optional[Callable] = None
@@ -249,18 +219,16 @@ _OVERRIDE_DENIED_MSG = (
# ---------------------------------------------------------------------------
# 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.
# 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
# Transient-failure suppression: probes can flap (a ``docker version`` timing 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 each check's last success is remembered and a fresh failure within a short
# grace window serves 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.
# ---------------------------------------------------------------------------
@@ -287,10 +255,7 @@ 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``.
"""
"""Expire stale entries and cap profile-dimensional cache growth. Caller holds the lock."""
for key, (timestamp, _) in list(_check_fn_cache.items()):
if now - timestamp >= _CHECK_FN_TTL_SECONDS:
_check_fn_cache.pop(key, None)
@@ -307,13 +272,11 @@ 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.
attach/detach, so a fully bound browser-control request bypasses both this cache
and model_tools' outer definition cache (same sentinel) — one Browser session's
live tools must not leak into another. Single-profile processes keep the
process-wide cache; a multiplex gateway installs a Hermes-home override per
profile turn, so the canonical profile key is the isolation boundary.
"""
try:
from gateway.session_context import get_session_env
@@ -353,11 +316,9 @@ def _run_check_fn_uncached(fn: Callable, *, unresolved_scope: bool = False) -> b
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: with multiplexing on, boot-time check_fns run
# before any profile secret scope exists, so get_secret raises by design.
# No traceback so this cannot be 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 "
@@ -365,8 +326,7 @@ def _run_check_fn_uncached(fn: Callable, *, unresolved_scope: bool = False) -> b
_fn_label(fn),
)
return False
# The scope resolved but the read still failed closed: a genuinely
# lost scope. Keep the loud crash-style report.
# 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",
@@ -378,9 +338,7 @@ def _run_check_fn_uncached(fn: Callable, *, unresolved_scope: bool = False) -> b
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,
_fn_label(fn), detail, exc_info=True,
)
return False
@@ -416,9 +374,8 @@ def _check_fn_cached(fn: Callable) -> bool:
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 and do NOT cache the
# failure, so the next call re-probes instead of pinning a stale verdict.
logger.warning(
"check_fn %s failed (%s) within %.0fs of last success; "
"treating as transient and keeping tool(s) available",
@@ -426,26 +383,31 @@ def _check_fn_cached(fn: Callable) -> bool:
)
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 (or grace expired) — honor the failure. Logged so silent
# tool loss in quiet mode (subagents) 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 such as ``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.
"""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.
@@ -468,29 +430,21 @@ class ToolRegistry:
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.
# Plugin registrations are overlays keyed by resolved HERMES_HOME: a profile
# sees its own overlay first, then the global built-ins.
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 module namespace -> operator opt-in for built-in override. Policies
# are lifecycle-managed; scope attribution below 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 dynamic refresh can mutate the registry while other threads read tool
# metadata: mutations are serialized and readers get stable snapshots.
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; external callers (get_tool_definitions) memoize
# against it — an entry keyed on the generation is valid until it changes.
self._generation: int = 0
@staticmethod
@@ -508,22 +462,16 @@ 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, {}))
merged.update(self._scoped_tools.get(scope or self.current_scope_key(), {}))
return merged
def _snapshot_state(
self,
scope: Optional[str] = None,
) -> tuple[List[ToolEntry], Dict[str, Callable]]:
def _snapshot_state(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())
@@ -537,46 +485,27 @@ class ToolRegistry:
"""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_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``.
Mirrors :meth:`get_definitions` per-tool filtering so doctor, banners and other
toolset-level surfaces agree with runtime exposure: mixed toolsets (``terminal``
plus desktop-only ``read_terminal``) must not be gated by the first ``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]:
if not entry.check_fn or _memo_check(entry.check_fn, check_results):
return True
return False
def get_entry(
self,
name: str,
*,
scope: Optional[str] = None,
) -> Optional[ToolEntry]:
def get_entry(self, name: str, *, scope: Optional[str] = None) -> Optional[ToolEntry]:
"""Return the 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]:
def snapshot_registration(self, name: str, *, scope: Optional[str] = None) -> Optional[ToolEntry]:
"""Return the local slot state without following global fallback."""
with self._lock:
return self._slot(scope).get(name)
@@ -591,10 +520,7 @@ class ToolRegistry:
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(entry.name for entry in self._snapshot_entries() if entry.toolset == toolset)
def register_toolset_alias(self, alias: str, toolset: str) -> None:
"""Register an explicit alias for a canonical toolset name."""
@@ -602,8 +528,7 @@ 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
@@ -623,11 +548,7 @@ class ToolRegistry:
# ------------------------------------------------------------------
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.
@@ -641,10 +562,7 @@ 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:
@@ -669,11 +587,7 @@ 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))
@@ -682,9 +596,9 @@ class ToolRegistry:
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.
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.
"""
mod = self._callable_module(handler)
return self._plugin_namespace_of_module(mod) if mod else None
@@ -718,17 +632,12 @@ class ToolRegistry:
return str(module_name)
return str(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)
@@ -773,8 +682,7 @@ class ToolRegistry:
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 ""
@@ -794,13 +702,12 @@ class ToolRegistry:
override: bool = False,
scope: Optional[str] = None,
):
"""Register a tool. Called at module-import time by each tool file.
"""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.
``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.
"""
handler_owner = self._plugin_owner_of(handler)
caller_owner = self._plugin_namespace_of_module(self._caller_module())
@@ -809,27 +716,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)
)
plugin_override_denied = (
owner is not None and not self._plugin_override_allowed(scope, owner)
)
existing = self._tools.get(name) 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)
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,
owner, name,
)
return
if plugin_override_denied:
@@ -846,17 +743,16 @@ class ToolRegistry:
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): replace the tool; INFO so
# the override is auditable in agent.log.
logger.info(
"Tool '%s': toolset '%s' overriding existing toolset '%s' "
"(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, including MCP-to-MCP collisions.
# MCP reconnect/refresh re-registers within the same toolset: allowed.
logger.error(
"Tool registration REJECTED: '%s' (toolset '%s') would "
"shadow existing tool from toolset '%s'. Pass "
@@ -878,22 +774,20 @@ class ToolRegistry:
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.
# 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 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``).
"""Remove a tool; drops the toolset check/aliases if it was the last tool in its toolset.
``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.
``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,
@@ -903,11 +797,7 @@ class ToolRegistry:
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 "
@@ -934,9 +824,7 @@ class ToolRegistry:
if (
caller_owner is not None
and not same_plugin
and not self._plugin_override_allowed(
caller_scope, caller_owner
)
and not self._plugin_override_allowed(caller_scope, caller_owner)
):
logger.error(
"Tool deregistration REJECTED: plugin %r attempted to "
@@ -954,33 +842,20 @@ class ToolRegistry:
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 any(e.toolset == entry.toolset for e in self._merged_tools(scope).values()):
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,
self, name: str, current: ToolEntry, previous: Optional[ToolEntry], *, scope: Optional[str] = None,
) -> bool:
"""Restore a host-owned registration if it is still current.
"""Restore a host-owned registration if it is still current (plugin ownership ledger).
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.
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 one untouched.
"""
with self._lock:
target = self._slot(scope, create=True)
@@ -994,22 +869,15 @@ 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 the affected toolset checks from the surviving entries: a plugin
# may have replaced an entry in the same toolset, so leaving the current
# check_fn behind would retain stale plugin state after restoration.
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 = [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)
if scope is None:
if check_fn is None:
self._toolset_checks.pop(toolset, None)
@@ -1030,29 +898,25 @@ class ToolRegistry:
# ------------------------------------------------------------------
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 go through the ~30 s TTL cache so ``hermes tools enable`` still 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
# (model_tools.get_tool_definitions) 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()
@@ -1060,9 +924,7 @@ class ToolRegistry:
schema_with_name.update(overrides)
except Exception as exc:
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,
)
result.append({"type": "function", "function": schema_with_name})
return result
@@ -1073,9 +935,9 @@ class ToolRegistry:
@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 (
@@ -1086,11 +948,7 @@ class ToolRegistry:
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",
@@ -1098,17 +956,9 @@ class ToolRegistry:
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": ...}``."""
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": ...}``."""
entry = self.get_entry(name, scope=scope)
if not entry:
return tool_error(f"Unknown tool: {name}")
@@ -1121,9 +971,7 @@ 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))
)
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.
raw = f"Tool execution failed: {type(e).__name__}: {e}"
@@ -1135,7 +983,7 @@ class ToolRegistry:
return tool_error(sanitized)
# ------------------------------------------------------------------
# Query helpers (replace redundant dicts in model_tools.py)
# Query helpers
# ------------------------------------------------------------------
def get_max_result_size(self, name: str, default: int | float | None = None) -> int | float:
@@ -1153,11 +1001,7 @@ class ToolRegistry:
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.
"""
"""A tool's raw schema dict, bypassing check_fn filtering (token estimation, introspection)."""
entry = self.get_entry(name)
return entry.schema if entry else None

View File

@@ -1,21 +1,11 @@
"""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 silently accept: llama.cpp's
``json-schema-to-grammar`` 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``; OpenAI's
Codex backend rejects top-level combinators. This module walks the final tool
schema tree on a deep copy and fixes only those shapes.
"""
from __future__ import annotations
@@ -28,9 +18,9 @@ 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 tool input schemas whose
# property keys don't match this; one bad key anywhere in the tools array 400s the
# entire request (Cloudflare's MCP ships 61 such keys).
_PROP_KEY_RE = re.compile(r"^[a-zA-Z0-9_.-]{1,64}$")
_PROP_KEY_BAD_CHARS = re.compile(r"[^a-zA-Z0-9_.-]")
@@ -49,11 +39,10 @@ def sanitize_property_key(key: str) -> str:
def _rename_property_keys(props: dict, path: str) -> dict[str, str]:
"""Return {original_key: conforming_key} for one properties dict.
"""Return {original_key: conforming_key} for one properties dict (identity entries omitted).
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.
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)}
@@ -78,8 +67,8 @@ def _rename_property_keys(props: dict, path: str) -> dict[str, str]:
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.
``params_schema`` is the ORIGINAL (unsanitized) registry schema. Recurses into
object values and array items; unknown keys pass through untouched.
"""
if not isinstance(params_schema, dict) or not isinstance(args, dict):
return args
@@ -96,8 +85,7 @@ 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
unrename_tool_args(subschema["items"], item) if isinstance(item, dict) else item
for item in value
]
out[orig] = value
@@ -105,15 +93,13 @@ def unrename_tool_args(params_schema: Any, args: Any) -> Any:
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 each parameter schema sanitized; callers may 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):
@@ -134,10 +120,9 @@ def _sanitize_single_tool(tool: dict) -> dict:
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``.
# 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"`` to None.
top = strip_nullable_unions(top, keep_nullable_hint=True)
top = _strip_top_level_combinators(top, path=name)
fn["parameters"] = _strip_ref_siblings(top)
@@ -149,8 +134,7 @@ _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``)."""
"""Recursively drop forbidden siblings from ``$ref`` nodes (Fireworks rejects ``default`` there)."""
if isinstance(node, list):
return [_strip_ref_siblings(item) for item in node]
if not isinstance(node, dict):
@@ -166,12 +150,10 @@ _TOP_LEVEL_FORBIDDEN_KEYS = ("allOf", "anyOf", "oneOf", "enum", "not")
def _strip_top_level_combinators(params: dict, *, path: str = "<tool>") -> dict:
"""Drop combinator keywords from the TOP level of a parameters schema only.
"""Drop combinators from the TOP level only (Codex rejects them there).
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.
They are usually conditional-required hints; dropping them does not change which
argument values are valid (handlers re-validate). Nested combinators are preserved.
"""
if not isinstance(params, dict):
return params
@@ -201,30 +183,21 @@ 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:
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).
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`` for 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()
}
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):
@@ -239,28 +212,21 @@ def strip_nullable_unions(
return stripped
_CONST_PRIMITIVE_TYPES: dict[type, str] = {
bool: "boolean",
int: "integer",
float: "number",
str: "string",
}
_CONST_PRIMITIVE_TYPES: dict[type, str] = {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.
Qualifies when the dict carries a primitive ``const`` and any declared ``type``
matches it; ``title``/``description`` are allowed, any other keyword disqualifies.
"""
if not isinstance(branch, dict) or "const" not in branch:
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))
json_type = _CONST_PRIMITIVE_TYPES.get(type(branch["const"]))
if json_type is None:
return None
declared = branch.get("type")
@@ -272,17 +238,14 @@ def _const_branch_type(branch: Any) -> str | 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.
Ported from block/goose ``tool_schema_normalize.rs`` (Apache-2.0). Rust/TS-generated
MCP servers 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`` (null+multi-const unions land here,
not in ``strip_nullable_unions``). 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]
@@ -294,19 +257,14 @@ def collapse_const_unions(schema: Any) -> Any:
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
]
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],
}
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)
@@ -315,20 +273,17 @@ def collapse_const_unions(schema: Any) -> Any:
_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).
# Sibling keywords whose values are NOT schemas: recursing would mistake literal strings
# like "path" for bare-string schemas. Passed through (``required`` follows property renames).
_NON_SCHEMA_LIST_KEYS = frozenset({"required", "enum", "examples", "dependentRequired"})
def _normalize_type_array(value: list, out: dict) -> None:
"""Normalize a ``type: [...]`` array into *out*.
"""Normalize a ``type: [...]`` array into *out* (llama.cpp and Gemini-via-OpenAI reject arrays).
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.
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.
"""
has_null = "null" in value
non_null = [t for t in value if isinstance(t, str) and t != "null"]
@@ -346,16 +301,11 @@ 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": <value>}`` (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).
Bare-string schemas become ``{"type": <value>}`` (unknown strings → permissive
object); object nodes gain ``properties: {}``; ``type`` arrays are normalized;
recursion covers ``properties``/``items``/``additionalProperties``/combinators/
``$defs``/``definitions``; 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:
@@ -420,19 +370,19 @@ def _sanitize_node(node: Any, path: str) -> Any:
return out
# =============================================================================
# ---------------------------------------------------------------------------
# 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."""
"""Apply *strip_node* (returns keywords removed) to every dict node of every tool's parameters, in place.
Handles OpenAI format (``{"function": {"parameters": ...}}``) and Responses format
(``{"name": ..., "parameters": ...}``). Returns ``(tools, stripped_count)`` — same list.
"""
if not tools:
return tools, 0
stripped = 0
@@ -453,8 +403,7 @@ def _reactive_strip(tools: list[dict], strip_node: Callable[[dict], int], log_ms
fn = tool.get("function")
if isinstance(fn, dict) and isinstance(fn.get("parameters"), dict):
_walk(fn["parameters"])
continue
if isinstance(tool.get("parameters"), dict):
elif isinstance(tool.get("parameters"), dict):
_walk(tool["parameters"])
if stripped:
@@ -465,14 +414,11 @@ def _reactive_strip(tools: list[dict], strip_node: Callable[[dict], int], log_ms
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``.
Reactive: only after llama.cpp's grammar converter rejected a schema (HTTP 400); its
regex engine supports a small ECMAScript subset and few ``format`` values, while
cloud providers use these as prompting hints, so they stay in the default schema.
Only strips beside ``type``/combinators (schema nodes), so a property literally
*named* ``pattern`` (``search_files``) inside ``properties`` is untouched.
"""
def _strip(node: dict) -> int:
if not ("type" in node or "anyOf" in node or "oneOf" in node or "allOf" in node):
@@ -492,10 +438,9 @@ def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]:
def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]:
"""Strip ``enum`` keywords whose string values contain ``/``, in place.
xAI'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.
xAI's Responses/chat endpoints compile 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; the model still sees the description.
"""
def _strip(node: dict) -> int:
enum_val = node.get("enum")

View File

@@ -17,7 +17,6 @@ from tools.approval import (
_read_shell_word,
)
# bisect is included: it drives repeated checkouts of the running root — the
# exact module-version-skew hazard this guard exists for.
_WORKTREE_MUTATIONS = frozenset({
@@ -44,16 +43,13 @@ _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")
_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",
}),
"env": frozenset({"-a", "--argv0", "-C", "--chdir", "-S", "--split-string", "-u", "--unset"}),
"command": _NO_OPTIONS,
"builtin": _NO_OPTIONS,
"exec": frozenset({"-a"}),
@@ -63,9 +59,7 @@ _WRAPPER_OPTIONS_WITH_ARG: dict[str, frozenset[str]] = {
}
_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,20 +112,14 @@ 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
return words
def _consume_options(
words: list[str],
start: int,
options_with_arg: frozenset[str] = _NO_OPTIONS,
) -> int:
def _consume_options(words: list[str], start: int, options_with_arg: frozenset[str] = _NO_OPTIONS) -> int:
"""Index of the first positional at/after ``start`` (``--`` ends options)."""
index = start
while index < len(words):
@@ -140,11 +128,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 +136,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 +150,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, []
@@ -233,12 +214,11 @@ def _operator_before(command: str, start: int) -> str | None:
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]
if index >= 0:
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
@@ -256,21 +236,19 @@ 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 '<script>'`` hides ``-c`` behind an operand). When it
finds no ``-c``, fall back to a permissive positional scan because
zsh/dash/ksh option letters (``zsh -yc``) fall outside bash's alphabet and
would otherwise make this block-guard fail open.
(``-o pipefail -c '<script>'`` hides ``-c`` behind an operand). When it finds
no ``-c``, fall back to a permissive positional scan: zsh/dash/ksh option
letters (``zsh -yc``) fall outside bash's alphabet and would otherwise make
this block-guard fail open.
"""
has_c, payload = _bash_exec_payload(args)
if has_c:
return payload
for index, arg in enumerate(args):
if arg == "--":
if arg == "--" or not arg.startswith("-"):
break
if arg.startswith("-") and "c" in arg[1:]:
if "c" in arg[1:]:
return args[index + 1] if index + 1 < len(args) else None
if not arg.startswith("-"):
break
return None
@@ -318,9 +296,7 @@ def _heredoc_specs(line: str) -> list[_Heredoc]:
index = end + 1
else:
end = index
while (
end < len(line) and not line[end].isspace() and line[end] not in ";|&<>"
):
while end < len(line) and not line[end].isspace() and line[end] not in ";|&<>":
end += 1
delimiter = line[index:end]
index = end
@@ -342,10 +318,6 @@ def _heredoc_specs(line: str) -> list[_Heredoc]:
return specs
def _masked_line(line: str) -> str:
return "".join(char if char in {"\r", "\n"} else " " for char in line)
def _mask_heredocs(command: str) -> tuple[str, list[str]]:
"""Blank heredoc bodies; return (masked command, bodies a bare shell would execute).
@@ -365,9 +337,8 @@ def _mask_heredocs(command: str) -> tuple[str, list[str]]:
finished.append(pending.pop(0))
else:
current.body.append(line)
output.append(_masked_line(line))
output.append("".join(char if char in {"\r", "\n"} else " " for char in line))
continue
output.append(line)
pending.extend(_heredoc_specs(line))
@@ -383,9 +354,7 @@ def _record_alias(config: str, aliases: dict[str, str]) -> None:
def _git_target_and_subcommand(
args: list[str],
current_dir: Path,
env: dict[str, str],
args: list[str], current_dir: Path, env: dict[str, str],
) -> tuple[Path, str | None, list[str], dict[str, str]]:
"""Parse git's global options -> (target dir, subcommand, sub args, inline aliases)."""
target = current_dir
@@ -441,8 +410,7 @@ def _stash_mutates(args: list[str]) -> bool:
def _clean_mutates(args: list[str]) -> bool:
return not any(
arg == "--dry-run" or (not arg.startswith("--") and _has_short_flag(arg, "n"))
for arg in args
arg == "--dry-run" or (not arg.startswith("--") and _has_short_flag(arg, "n")) for arg in args
)
@@ -463,9 +431,7 @@ _CONDITIONAL_MUTATIONS: dict[str, Callable[[list[str]], bool]] = {
def _mutates_worktree(subcommand: str, args: list[str]) -> bool:
check = _CONDITIONAL_MUTATIONS.get(subcommand)
if check is not None:
return check(args)
return subcommand in _WORKTREE_MUTATIONS
return check(args) if check is not None else subcommand in _WORKTREE_MUTATIONS
def _inspect_git_worktree(args: list[str], cwd: Path, root: Path) -> str | None:
@@ -477,9 +443,7 @@ def _inspect_git_worktree(args: list[str], cwd: Path, root: Path) -> str | None:
if action not in _WORKTREE_TARGET_ACTIONS:
return None
target_index = _consume_options(args, action_index + 1)
if target_index >= len(args):
return None
if _resolve(args[target_index], cwd) == root:
if target_index < len(args) and _resolve(args[target_index], cwd) == root:
return f"git worktree {action}"
return None
@@ -488,10 +452,7 @@ def _read_git_alias(executable: str, target: Path, alias: str) -> str | None:
try:
result = subprocess.run(
[executable, "-C", str(target), "config", "--get", f"alias.{alias}"],
capture_output=True,
text=True,
timeout=1,
check=False,
capture_output=True, text=True, timeout=1, check=False,
)
except (OSError, subprocess.SubprocessError):
return None
@@ -500,16 +461,9 @@ def _read_git_alias(executable: str, target: Path, alias: str) -> str | None:
def _inspect_git(
executable: str,
args: list[str],
current_dir: Path,
env: dict[str, str],
root: Path,
depth: int,
executable: str, args: list[str], current_dir: Path, env: dict[str, str], root: Path, depth: int,
) -> str | None:
target, subcommand, sub_args, inline_aliases = _git_target_and_subcommand(
args, current_dir, env
)
target, subcommand, sub_args, inline_aliases = _git_target_and_subcommand(args, current_dir, env)
if subcommand is None:
return None
# `worktree` names its victim as an argument, so the cwd check does not apply.
@@ -533,45 +487,25 @@ def _inspect_git(
alias_args = shlex.split(alias, posix=True)
except ValueError:
return None
return _inspect_git(
executable,
[*alias_args, *sub_args],
target,
{},
root,
depth + 1,
)
return _inspect_git(executable, [*alias_args, *sub_args], target, {}, root, depth + 1)
def _inspect_github_cli(
executable: str,
args: list[str],
current_dir: Path,
env: dict[str, str],
root: Path,
depth: int,
executable: str, args: list[str], current_dir: Path, env: dict[str, str], root: Path, depth: int,
) -> str | None:
if not _is_within(current_dir, root):
return None
name = _executable_name(executable)
index = _consume_options(args, 0, frozenset({"-R", "--repo", "--hostname"}))
if args[index : index + 2] == ["pr", "checkout"]:
return f"{name} pr checkout"
return f"{_executable_name(executable)} pr checkout"
return None
def _inspect_shell(
executable: str,
args: list[str],
current_dir: Path,
env: dict[str, str],
root: Path,
depth: int,
executable: str, args: list[str], current_dir: Path, env: dict[str, str], root: Path, depth: int,
) -> str | None:
script = _shell_script_arg(args)
if script:
return _find_mutation(script, current_dir, root, depth + 1)
return None
return _find_mutation(script, current_dir, root, depth + 1) if script else None
# executable name -> inspector(executable, args, current_dir, env, root, depth)
@@ -611,8 +545,7 @@ def _find_mutation(command: str, cwd: Path, root: Path, depth: int = 0) -> str |
if pending is not None and operator in {"&&", ";", "\n"}:
cwd_by_scope[scope] = pending
words = _shell_words_at(masked_command, start)
env, executable, args = _command_parts(words)
env, executable, args = _command_parts(_shell_words_at(masked_command, start))
if executable is None:
continue
@@ -634,19 +567,16 @@ def _find_mutation(command: str, cwd: Path, root: Path, depth: int = 0) -> str |
def guard_active() -> bool:
"""Whether the self-repo git guard applies on this platform.
Windows-only: NTFS locks loaded .py/.pyd files, so an in-place overwrite of
the live checkout can corrupt the running process. On POSIX, open handles
keep the old inode alive, so already-imported modules keep running; the
mixed-module hazard is limited to later lazy imports — not worth blocking
Windows-only: NTFS locks loaded .py/.pyd files, so overwriting the live checkout
can corrupt the running process. On POSIX, open handles keep the old inode alive;
the mixed-module hazard is limited to later lazy imports — not worth blocking
every git workflow for.
"""
return os.name == "nt"
def detect_self_repo_git_mutation(
command: str,
cwd: str | None,
source_root: Path | None = None,
command: str, cwd: str | None, source_root: Path | None = None,
) -> tuple[bool, str | None]:
"""Return whether a command would rewrite the live source checkout."""
root = source_root if source_root is not None else get_running_source_root()