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.
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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",
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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", "<tool>")
|
||||
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 = "<tool>") -> 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 = "<tool>") -> 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": <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).
|
||||
"""
|
||||
"""Recursively sanitize a JSON-Schema fragment: bare-string schemas become ``{"type":
|
||||
<value>}`` (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)")
|
||||
|
||||
@@ -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.<sub>`
|
||||
# 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.<sub>` 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<dash>-?)[ \t]*(?:(?P<q>['\"])(?P<quoted>.*?)(?P=q)|(?!['\"])(?P<bare>[^\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 '<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.
|
||||
"""
|
||||
"""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, since 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
|
||||
|
||||
|
||||
@@ -279,54 +226,26 @@ def _heredoc_specs(line: str) -> list[_Heredoc]:
|
||||
specs: list[_Heredoc] = []
|
||||
quote: str | None = None
|
||||
index = 0
|
||||
|
||||
while index < len(line):
|
||||
char = line[index]
|
||||
if quote:
|
||||
if char == "\\" and quote == '"' and index + 1 < len(line):
|
||||
index += 2
|
||||
continue
|
||||
if char == quote:
|
||||
index += 1 # skip the escaped character too
|
||||
elif char == quote:
|
||||
quote = None
|
||||
index += 1
|
||||
continue
|
||||
if char in {"'", '"'}:
|
||||
elif char in {"'", '"'}:
|
||||
quote = char
|
||||
if quote or not line.startswith("<<", index) or line.startswith("<<<", index):
|
||||
index += 1
|
||||
continue
|
||||
if not line.startswith("<<", index) or line.startswith("<<<", index):
|
||||
index += 1
|
||||
continue
|
||||
|
||||
operator_at = index
|
||||
index += 2
|
||||
strip_tabs = index < len(line) and line[index] == "-"
|
||||
if strip_tabs:
|
||||
index += 1
|
||||
while index < len(line) and line[index] in {" ", "\t"}:
|
||||
index += 1
|
||||
if index >= len(line):
|
||||
opener = _HEREDOC_OPENER_RE.match(line, index)
|
||||
if opener is None: # unterminated quoted delimiter: give up on this line
|
||||
break
|
||||
|
||||
delimiter_quote = line[index] if line[index] in {"'", '"'} else None
|
||||
if delimiter_quote:
|
||||
index += 1
|
||||
end = line.find(delimiter_quote, index)
|
||||
if end == -1:
|
||||
break
|
||||
delimiter = line[index:end]
|
||||
index = end + 1
|
||||
else:
|
||||
end = index
|
||||
while (
|
||||
end < len(line) and not line[end].isspace() and line[end] not in ";|&<>"
|
||||
):
|
||||
end += 1
|
||||
delimiter = line[index:end]
|
||||
index = end
|
||||
operator_at, index = index, opener.end()
|
||||
strip_tabs = bool(opener.group("dash"))
|
||||
delimiter = opener.group("quoted") if opener.group("q") else opener.group("bare")
|
||||
if not delimiter:
|
||||
continue
|
||||
|
||||
header = line[:operator_at]
|
||||
starts = list(_iter_shell_command_starts(header))
|
||||
words = _shell_words_at(header, starts[-1]) if starts else []
|
||||
@@ -335,26 +254,17 @@ def _heredoc_specs(line: str) -> list[_Heredoc]:
|
||||
executable
|
||||
and _executable_name(executable) in _SHELL_EXECUTABLES
|
||||
and _shell_script_arg(args) is None
|
||||
and not any(arg and not arg.startswith("-") for arg in args)
|
||||
)
|
||||
and not any(arg and not arg.startswith("-") for arg in args))
|
||||
specs.append(_Heredoc(delimiter, strip_tabs, execute_as_shell))
|
||||
|
||||
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).
|
||||
|
||||
Unterminated heredocs run to end of input and are still reported.
|
||||
"""
|
||||
Unterminated heredocs run to end of input and are still reported."""
|
||||
output: list[str] = []
|
||||
pending: list[_Heredoc] = []
|
||||
finished: list[_Heredoc] = []
|
||||
|
||||
for line in command.splitlines(keepends=True):
|
||||
if pending:
|
||||
current = pending[0]
|
||||
@@ -365,12 +275,10 @@ 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))
|
||||
|
||||
shell_scripts = ["".join(spec.body) for spec in finished + pending if spec.execute_as_shell]
|
||||
return "".join(output), shell_scripts
|
||||
|
||||
@@ -383,16 +291,13 @@ 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
|
||||
work_tree: str | None = None
|
||||
aliases: dict[str, str] = {}
|
||||
index = 0
|
||||
|
||||
while index < len(args):
|
||||
arg = args[index]
|
||||
if arg == "--":
|
||||
@@ -418,12 +323,10 @@ def _git_target_and_subcommand(
|
||||
elif arg.startswith("-calias."):
|
||||
_record_alias(arg[2:], aliases)
|
||||
index += 1
|
||||
|
||||
explicit_work_tree = work_tree or env.get("GIT_WORK_TREE")
|
||||
if explicit_work_tree:
|
||||
target = _resolve(explicit_work_tree, target)
|
||||
subcommand = args[index].lower() if index < len(args) else None
|
||||
return target, subcommand, args[index + 1 :], aliases
|
||||
return target, args[index].lower() if index < len(args) else None, args[index + 1 :], aliases
|
||||
|
||||
|
||||
def _has_short_flag(arg: str, letter: str) -> bool:
|
||||
@@ -442,8 +345,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
|
||||
)
|
||||
for arg in args)
|
||||
|
||||
|
||||
def _restore_mutates(args: list[str]) -> bool:
|
||||
@@ -457,29 +359,22 @@ _CONDITIONAL_MUTATIONS: dict[str, Callable[[list[str]], bool]] = {
|
||||
"reset": _reset_mutates,
|
||||
"stash": _stash_mutates,
|
||||
"clean": _clean_mutates,
|
||||
"restore": _restore_mutates,
|
||||
}
|
||||
"restore": _restore_mutates}
|
||||
|
||||
|
||||
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:
|
||||
"""Block `worktree remove|move` aimed at the running root, from any directory."""
|
||||
action_index = _consume_options(args, 0)
|
||||
if action_index >= len(args):
|
||||
return None
|
||||
action = args[action_index].lower()
|
||||
action = args[action_index].lower() if action_index < len(args) else 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,28 +383,17 @@ 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
|
||||
value = result.stdout.strip()
|
||||
return value if result.returncode == 0 and value else None
|
||||
return (result.stdout.strip() or None) if result.returncode == 0 else 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
|
||||
)
|
||||
args, current_dir, env)
|
||||
if subcommand is None:
|
||||
return None
|
||||
# `worktree` names its victim as an argument, so the cwd check does not apply.
|
||||
@@ -521,7 +405,6 @@ def _inspect_git(
|
||||
return f"git {subcommand}"
|
||||
if subcommand in _KNOWN_GIT_BUILTINS or depth >= _MAX_RECURSION:
|
||||
return None
|
||||
|
||||
alias = inline_aliases.get(subcommand)
|
||||
if alias is None:
|
||||
alias = _read_git_alias(executable, target, subcommand)
|
||||
@@ -533,45 +416,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)
|
||||
@@ -579,100 +442,76 @@ _INSPECTORS: dict[str, Callable[..., str | None]] = {
|
||||
"git": _inspect_git,
|
||||
"gh": _inspect_github_cli,
|
||||
"hub": _inspect_github_cli,
|
||||
**{shell: _inspect_shell for shell in _SHELL_EXECUTABLES},
|
||||
}
|
||||
**{shell: _inspect_shell for shell in _SHELL_EXECUTABLES}}
|
||||
|
||||
|
||||
def _find_mutation(command: str, cwd: Path, root: Path, depth: int = 0) -> str | None:
|
||||
"""Name of the first command in ``command`` that would rewrite ``root``, else None."""
|
||||
if depth > _MAX_RECURSION:
|
||||
return None
|
||||
|
||||
masked_command, heredoc_scripts = _mask_heredocs(command)
|
||||
for script in heredoc_scripts:
|
||||
operation = _find_mutation(script, cwd, root, depth + 1)
|
||||
if operation:
|
||||
return operation
|
||||
|
||||
starts = sorted(set(_iter_shell_command_starts(masked_command)))
|
||||
scopes = _scope_keys(masked_command, starts)
|
||||
# cwd is tracked per subshell scope: `cd` only takes effect for the NEXT
|
||||
# command when joined by `&&`, `;` or a newline (not `||` / `|`).
|
||||
# cwd per subshell scope; `cd` applies to the NEXT command only via `&&`, `;`, newline.
|
||||
cwd_by_scope: dict[tuple[int, ...], Path] = {(): cwd}
|
||||
pending_cd: dict[tuple[int, ...], Path] = {}
|
||||
|
||||
for start in starts:
|
||||
scope = scopes[start]
|
||||
if scope not in cwd_by_scope:
|
||||
cwd_by_scope[scope] = cwd_by_scope.get(scope[:-1], cwd)
|
||||
|
||||
cwd_by_scope.setdefault(scope, cwd_by_scope.get(scope[:-1], cwd))
|
||||
operator = _operator_before(masked_command, start)
|
||||
pending = pending_cd.pop(scope, None)
|
||||
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
|
||||
|
||||
current_dir = cwd_by_scope[scope]
|
||||
cd_target = _cd_target(executable, args, current_dir)
|
||||
if cd_target is not None:
|
||||
pending_cd[scope] = cd_target
|
||||
continue
|
||||
|
||||
inspect = _INSPECTORS.get(_executable_name(executable))
|
||||
if inspect is not None:
|
||||
operation = inspect(executable, args, current_dir, env, root, depth)
|
||||
if operation:
|
||||
return operation
|
||||
|
||||
return None
|
||||
|
||||
|
||||
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
|
||||
every git workflow for.
|
||||
"""
|
||||
"""Whether the self-repo git guard applies on this platform. 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,
|
||||
) -> tuple[bool, str | 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()
|
||||
if root is None or not command:
|
||||
return False, None
|
||||
|
||||
root = _resolve(str(root), Path("/"))
|
||||
base = _resolve(cwd, Path("/")) if cwd else Path("/")
|
||||
operation = _find_mutation(command, base, root)
|
||||
if operation is None:
|
||||
return False, None
|
||||
return True, _block_message(operation, root)
|
||||
operation = _find_mutation(command, _resolve(cwd, Path("/")) if cwd else Path("/"), root)
|
||||
return (True, _block_message(operation, root)) if operation is not None else (False, None)
|
||||
|
||||
|
||||
def _block_message(operation: str, root: Path) -> str:
|
||||
# Suggest a disk-backed scratch dir: /tmp is usually tmpfs (see message).
|
||||
hermes_home = os.environ.get("HERMES_HOME", "").strip()
|
||||
scratch = (Path(hermes_home).expanduser() if hermes_home else Path.home() / ".hermes") / "scratch"
|
||||
home = Path(hermes_home).expanduser() if hermes_home else Path.home() / ".hermes"
|
||||
scratch = home / "scratch"
|
||||
return (
|
||||
f"Blocked: `{operation}` would rewrite Hermes's live source checkout "
|
||||
f"({root}) and can mix module versions in this running process. "
|
||||
f"Use a separate worktree or a shared clone on real disk, e.g. "
|
||||
f"`git clone --shared {root} {scratch}/<task>` — avoid /tmp for "
|
||||
"clones that install node/python deps: /tmp is usually RAM-backed "
|
||||
"tmpfs and a few dependency installs can fill it and ENOSPC other "
|
||||
"work. Delete the clone when the branch is pushed. To change this "
|
||||
"checkout, stop Hermes, run the command externally, then restart "
|
||||
"Hermes."
|
||||
)
|
||||
"clones that install node/python deps: /tmp is usually RAM-backed tmpfs and a few "
|
||||
"dependency installs can fill it and ENOSPC other work. Delete the clone when the branch "
|
||||
"is pushed. To change this checkout, stop Hermes, run the command externally, then restart "
|
||||
"Hermes.")
|
||||
|
||||
Reference in New Issue
Block a user