refactor(tools): second compaction pass on registry/sanitizer/guard; wrap >100-col lines

This commit is contained in:
Teknium
2026-09-02 22:38:41 -07:00
parent 502e68b92e
commit 1155d44dcb
6 changed files with 174 additions and 172 deletions

View File

@@ -5,7 +5,7 @@ from typing import Optional
def validate_within_dir(path: Path, root: Path) -> Optional[str]:
"""Error message if *path* does not resolve inside *root* (symlinks/``..`` normalised), else None."""
"""Error message if *path* does not resolve inside *root* (symlinks and ``..`` followed)."""
try:
path.resolve().relative_to(root.resolve())
except (ValueError, OSError) as exc:

View File

@@ -185,8 +185,10 @@ 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 verdict to ``(allowed, reason)``: True installs, None needs confirmation, False is blocked."""
def should_allow_plugin_install(
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)"

View File

@@ -64,7 +64,9 @@ def _resolve(conn, token: str):
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})
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:
@@ -77,7 +79,10 @@ 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
],
})
@@ -127,7 +132,9 @@ def project_switch(project: str, task_id: Optional[str] = None) -> str:
_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),
"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),
}

View File

@@ -35,7 +35,8 @@ 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,9 +53,7 @@ def _bound_json_error_result(result: str) -> str:
payload = json.loads(result)
except ValueError:
return result
if not isinstance(payload, dict):
return result
error = payload.get("error")
error = payload.get("error") if isinstance(payload, dict) else None
if not isinstance(error, str) or len(error) <= _MAX_TOOL_ERROR_CHARS:
return result
payload["error"] = _bound_error_text(error)
@@ -68,8 +67,7 @@ def _is_registry_register_call(node: ast.AST) -> bool:
func = node.value.func
return (
isinstance(func, ast.Attribute) and func.attr == "register"
and isinstance(func.value, ast.Name) and func.value.id == "registry"
)
and isinstance(func.value, ast.Name) and func.value.id == "registry")
def _module_registers_tools(module_path: Path) -> bool:
@@ -92,8 +90,7 @@ def _module_registers_tools(module_path: Path) -> bool:
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
)
for stmt in tree.body)
def discover_builtin_tools(tools_dir: Optional[Path] = None) -> List[str]:
@@ -120,7 +117,7 @@ def discover_builtin_tools(tools_dir: Optional[Path] = None) -> List[str]:
except OSError:
continue
cached = cache.get(abs_path)
if isinstance(cached, (list, tuple)) and len(cached) == 3 and (cached[0], cached[1]) == stat_key:
if isinstance(cached, (list, tuple)) and len(cached) == 3 and tuple(cached[:2]) == stat_key:
registers = bool(cached[2])
else:
registers = _module_registers_tools(path)
@@ -183,7 +180,7 @@ def _save_discovery_cache(cache: Dict[str, list]) -> None:
@dataclass(eq=False, slots=True)
class ToolEntry:
"""Metadata for one registered tool (identity semantics: restore/CAS paths compare with ``is``)."""
"""Metadata for one registered tool (identity semantics: restore/CAS paths compare ``is``)."""
name: str
toolset: str
@@ -212,8 +209,7 @@ class _PluginOverridePolicy:
_OVERRIDE_DENIED_MSG = (
"Plugin module {owner!r} cannot override built-in tool {name!r} "
"without operator opt-in (allow_tool_override)."
)
"without operator opt-in (allow_tool_override).")
# ---------------------------------------------------------------------------
@@ -242,6 +238,11 @@ _check_fn_last_good: Dict[tuple[Callable, Optional[str]], float] = {}
_check_fn_cache_lock = threading.Lock()
CHECK_FN_CACHE_BYPASS = ""
_NO_CACHE_CHECK_FNS: Set[Callable] = set()
_BROWSER_IDENTITY_KEYS = (
"HERMES_SESSION_ID",
"HERMES_BROWSER_CONTROL_PRINCIPAL",
"HERMES_BROWSER_CONTROL_TRANSPORT_FAMILY",
)
def no_cache_check_fn(fn: Callable) -> Callable:
@@ -281,11 +282,7 @@ def check_fn_cache_scope() -> Optional[str]:
try:
from gateway.session_context import get_session_env
browser_identity = (
get_session_env("HERMES_SESSION_ID", ""),
get_session_env("HERMES_BROWSER_CONTROL_PRINCIPAL", ""),
get_session_env("HERMES_BROWSER_CONTROL_TRANSPORT_FAMILY", ""),
)
browser_identity = tuple(get_session_env(key, "") for key in _BROWSER_IDENTITY_KEYS)
if all(str(value or "").strip() for value in browser_identity):
return CHECK_FN_CACHE_BYPASS
except Exception:
@@ -299,9 +296,7 @@ def check_fn_cache_scope() -> Optional[str]:
from hermes_constants import get_hermes_home_override
override = get_hermes_home_override()
if not override:
return CHECK_FN_CACHE_BYPASS
return str(Path(override).expanduser().resolve())
return str(Path(override).expanduser().resolve()) if override else CHECK_FN_CACHE_BYPASS
except Exception:
# Fail closed: bypass both cache layers rather than aliasing requests
# whose multiplex profile identity could not be resolved.
@@ -323,24 +318,19 @@ def _run_check_fn_uncached(fn: Callable, *, unresolved_scope: bool = False) -> b
"check_fn %s hit the multiplex fail-closed path with no "
"profile secret scope active; dependent tools re-probe on "
"the first scoped turn",
_fn_label(fn),
)
return False
# The scope resolved but the read still failed closed: a genuinely lost scope.
logger.warning(
"check_fn %s raised UnscopedSecretError while the profile cache "
"scope was resolved; dependent tools will be unavailable this turn",
_fn_label(fn),
exc_info=True,
)
return False
_fn_label(fn))
else:
# The scope resolved but the read still failed closed: a genuinely lost scope.
logger.warning(
"check_fn %s raised UnscopedSecretError while the profile cache "
"scope was resolved; dependent tools will be unavailable this turn",
_fn_label(fn), exc_info=True)
except Exception:
detail = " while profile cache scope was unresolved" if unresolved_scope else ""
logger.warning(
"check_fn %s raised%s; dependent tools will be unavailable this turn",
_fn_label(fn), detail, exc_info=True,
)
return False
_fn_label(fn), detail, exc_info=True)
return False
def _check_fn_cached(fn: Callable) -> bool:
@@ -379,15 +369,13 @@ def _check_fn_cached(fn: Callable) -> bool:
logger.warning(
"check_fn %s failed (%s) within %.0fs of last success; "
"treating as transient and keeping tool(s) available",
_fn_label(fn), outcome, _CHECK_FN_FAILURE_GRACE_SECONDS,
)
_fn_label(fn), outcome, _CHECK_FN_FAILURE_GRACE_SECONDS)
return True
# No recent success (or grace expired) — honor the failure. 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
@@ -400,7 +388,7 @@ def _memo_check(fn: Callable, memo: Dict[Callable, bool]) -> bool:
def invalidate_check_fn_cache() -> None:
"""Drop all cached ``check_fn`` results (after config changes such as ``hermes tools enable``)."""
"""Drop all cached ``check_fn`` results (after config changes like ``hermes tools enable``)."""
with _check_fn_cache_lock:
_check_fn_cache.clear()
_check_fn_last_good.clear()
@@ -462,16 +450,18 @@ 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."""
merged = dict(self._tools)
merged.update(self._scoped_tools.get(scope or self.current_scope_key(), {}))
return merged
return {**self._tools, **self._scoped_tools.get(scope or self.current_scope_key(), {})}
def _snapshot_state(self, scope: Optional[str] = None) -> tuple[List[ToolEntry], Dict[str, Callable]]:
def _toolset_entries(self, toolset: str, scope: Optional[str]) -> List[ToolEntry]:
return [entry for entry in self._merged_tools(scope).values() if entry.toolset == toolset]
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())
@@ -493,19 +483,18 @@ class ToolRegistry:
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 or _memo_check(entry.check_fn, check_results):
return True
return False
return any(
not entry.check_fn or _memo_check(entry.check_fn, check_results)
for entry in entries if entry.toolset == toolset)
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)
@@ -528,7 +517,8 @@ 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
@@ -574,8 +564,7 @@ class ToolRegistry:
current: _PluginOverridePolicy,
previous: Optional[_PluginOverridePolicy],
*,
scope: Optional[str] = None,
) -> bool:
scope: Optional[str] = None) -> bool:
"""CAS-restore policy state while retaining durable scope attribution."""
with self._lock:
key = (scope, module_namespace)
@@ -618,15 +607,12 @@ class ToolRegistry:
current = func
continue
globals_dict = getattr(current, "__globals__", None)
if isinstance(globals_dict, dict):
module_name = globals_dict.get("__name__", "")
if module_name:
return str(module_name)
if isinstance(globals_dict, dict) and globals_dict.get("__name__", ""):
return str(globals_dict["__name__"])
wrapped = getattr(current, "__wrapped__", None)
if wrapped is not None:
current = wrapped
continue
break
if wrapped is None:
break
current = wrapped
module_name = getattr(current, "__module__", "")
if module_name:
return str(module_name)
@@ -637,8 +623,7 @@ class ToolRegistry:
with self._lock:
matches = [
namespace for namespace in self._plugin_module_scopes
if module_namespace == namespace or module_namespace.startswith(f"{namespace}.")
]
if module_namespace == namespace or module_namespace.startswith(f"{namespace}.")]
if matches:
return max(matches, key=len)
# Also gate plugin modules currently loading but not yet policy-recorded
@@ -660,8 +645,7 @@ class ToolRegistry:
return next(iter(scopes))
raise PermissionError(
f"Plugin module {module_namespace!r} is active in multiple "
"profiles and cannot register outside one of those scopes."
)
"profiles and cannot register outside one of those scopes.")
def plugin_scope_for_module(self, module_namespace: str) -> Optional[str]:
"""Public host lookup for a loaded plugin module's immutable scope."""
@@ -700,8 +684,7 @@ class ToolRegistry:
max_result_size_chars: int | float | None = None,
dynamic_schema_overrides: Callable = None,
override: bool = False,
scope: Optional[str] = None,
):
scope: Optional[str] = None):
"""Register a tool. Called at module-import time by each tool file.
``override=True`` is an explicit opt-in for plugins that intend to replace an
@@ -716,18 +699,20 @@ 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 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:
raise PermissionError(_OVERRIDE_DENIED_MSG.format(owner=owner, name=name))
@@ -740,16 +725,14 @@ class ToolRegistry:
"operator opt-in. Set "
"plugins.entries.<plugin_id>.allow_tool_override: true "
"in config.yaml to allow it.",
owner, name, existing.toolset,
)
owner, name, existing.toolset)
raise PermissionError(_OVERRIDE_DENIED_MSG.format(owner=owner, name=name))
# 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,
)
name, toolset, existing.toolset)
else:
# Reject every cross-toolset shadow, including MCP-to-MCP collisions.
# MCP reconnect/refresh re-registers within the same toolset: allowed.
@@ -758,8 +741,7 @@ class ToolRegistry:
"shadow existing tool from toolset '%s'. Pass "
"override=True to register() if the replacement is "
"intentional, or deregister the existing tool first.",
name, toolset, existing.toolset,
)
name, toolset, existing.toolset)
return
target[name] = ToolEntry(
name=name,
@@ -772,8 +754,7 @@ class ToolRegistry:
description=description or schema.get("description", ""),
emoji=emoji,
max_result_size_chars=max_result_size_chars,
dynamic_schema_overrides=dynamic_schema_overrides,
)
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,
@@ -801,8 +782,7 @@ class ToolRegistry:
if caller_owner is not None and scope is not None and scope != caller_scope:
raise PermissionError(
f"Plugin module {caller_mod!r} cannot deregister tools "
"outside its own profile scope."
)
"outside its own profile scope.")
if scope is None:
scope = caller_scope
target = self._slot(scope)
@@ -812,8 +792,7 @@ class ToolRegistry:
raise PermissionError(
f"Scoped plugin module {caller_mod!r} cannot deregister "
f"process-global tool {name!r}; register a scoped "
"override instead."
)
"override instead.")
return
if not entry.toolset.startswith("mcp-"):
owner = self._plugin_owner_of(entry.handler)
@@ -824,32 +803,34 @@ 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 "
"remove tool %r (toolset %r) it does not own, without "
"operator opt-in. Set "
"plugins.entries.%s.allow_tool_override: true in "
"config.yaml to allow it.",
caller_mod, name, entry.toolset, caller_mod,
)
caller_mod, name, entry.toolset, caller_mod)
raise PermissionError(
f"Plugin module {caller_mod!r} cannot deregister tool "
f"{name!r} (toolset {entry.toolset!r}) without operator "
f"opt-in (allow_tool_override)."
)
f"opt-in (allow_tool_override).")
del target[name]
if scope is not None and not target:
self._scoped_tools.pop(scope, None)
if not any(e.toolset == entry.toolset for e in self._merged_tools(scope).values()):
if not self._toolset_entries(entry.toolset, scope):
self._toolset_checks.pop(entry.toolset, None)
self._drop_toolset_aliases(entry.toolset)
self._generation += 1
logger.debug("Deregistered tool: %s", name)
def restore_registration(
self, name: str, current: ToolEntry, previous: Optional[ToolEntry], *, scope: Optional[str] = None,
self,
name: str,
current: ToolEntry,
previous: Optional[ToolEntry],
*,
scope: Optional[str] = None,
) -> bool:
"""Restore a host-owned registration if it is still current (plugin ownership ledger).
@@ -876,7 +857,7 @@ class ToolRegistry:
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]
surviving = self._toolset_entries(toolset, scope)
check_fn = next((entry.check_fn for entry in surviving if entry.check_fn), None)
if scope is None:
if check_fn is None:
@@ -886,8 +867,7 @@ class ToolRegistry:
if not surviving and not any(
entry.toolset == toolset
for entries in self._scoped_tools.values()
for entry in entries.values()
):
for entry in entries.values()):
self._drop_toolset_aliases(toolset)
self._generation += 1
logger.debug("Restored tool registration: %s", name)
@@ -924,7 +904,8 @@ 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
@@ -943,8 +924,7 @@ class ToolRegistry:
if (
isinstance(result, dict)
and result.get("_multimodal") is True
and isinstance(result.get("content"), list)
):
and isinstance(result.get("content"), list)):
return result
result_type = type(result).__name__
@@ -953,10 +933,11 @@ class ToolRegistry:
f"Tool handler returned unsupported result type: {result_type}",
error_type="tool_result_contract",
tool=name,
result_type=result_type,
)
result_type=result_type)
def dispatch(self, name: str, args: dict, *, scope: Optional[str] = None, **kwargs) -> str | dict:
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)
@@ -1001,7 +982,7 @@ class ToolRegistry:
return sorted(entry.name for entry in self._snapshot_entries())
def get_schema(self, name: str) -> Optional[dict]:
"""A tool's raw schema dict, bypassing check_fn filtering (token estimation, introspection)."""
"""A tool's raw schema dict, bypassing check_fn filtering (token estimates, introspection)."""
entry = self.get_entry(name)
return entry.schema if entry else None
@@ -1028,8 +1009,7 @@ class ToolRegistry:
entries = self._snapshot_entries()
return {
toolset: self._toolset_has_exposable_tools(toolset, entries)
for toolset in sorted({entry.toolset for entry in entries})
}
for toolset in sorted({entry.toolset for entry in entries})}
def get_available_toolsets(self) -> Dict[str, dict]:
"""Return toolset metadata for UI display."""
@@ -1042,8 +1022,7 @@ class ToolRegistry:
"available": self._toolset_has_exposable_tools(entry.toolset, entries),
"tools": [],
"description": "",
"requirements": [],
}
"requirements": []}
info["tools"].append(entry.name)
_extend_unique(info["requirements"], entry.requires_env or [])
return toolsets
@@ -1058,8 +1037,7 @@ class ToolRegistry:
"env_vars": [],
"check_fn": toolset_checks.get(entry.toolset),
"setup_url": None,
"tools": [],
})
"tools": []})
_extend_unique(info["tools"], [entry.name])
_extend_unique(info["env_vars"], entry.requires_env)
return result
@@ -1077,8 +1055,7 @@ class ToolRegistry:
unavailable.append({
"name": ts,
"env_vars": ts_entries[0].requires_env if ts_entries else [],
"tools": [entry.name for entry in ts_entries],
})
"tools": [entry.name for entry in ts_entries]})
return available, unavailable

View File

@@ -59,8 +59,7 @@ 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
@@ -86,14 +85,13 @@ def unrename_tool_args(params_schema: Any, args: Any) -> Any:
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
]
for item in value]
out[orig] = value
return out
def sanitize_tool_schemas(tools: list[dict]) -> list[dict]:
"""Deep-copied ``tools`` (OpenAI format) with each parameter schema sanitized; callers may mutate."""
"""Deep-copied ``tools`` (OpenAI format) with sanitized parameter schemas; callers may mutate."""
if not tools:
return tools
return [_sanitize_single_tool(tool) for tool in tools]
@@ -134,7 +132,7 @@ _REF_FORBIDDEN_SIBLINGS = frozenset({"default"})
def _strip_ref_siblings(node: Any) -> Any:
"""Recursively drop forbidden siblings from ``$ref`` nodes (Fireworks rejects ``default`` there)."""
"""Recursively drop forbidden siblings of ``$ref`` (Fireworks rejects ``default`` there)."""
if isinstance(node, list):
return [_strip_ref_siblings(item) for item in node]
if not isinstance(node, dict):
@@ -163,8 +161,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
@@ -193,11 +190,16 @@ def strip_nullable_unions(schema: Any, *, keep_nullable_hint: bool = True) -> An
``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]
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):
@@ -212,7 +214,9 @@ def strip_nullable_unions(schema: Any, *, keep_nullable_hint: bool = True) -> An
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:
@@ -264,7 +268,10 @@ def collapse_const_unions(schema: Any) -> Any:
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)
@@ -312,13 +319,11 @@ def _sanitize_node(node: Any, path: str) -> Any:
logger.debug(
"schema_sanitizer[%s]: replacing bare-string schema %r "
"with {'type': %r}",
path, node, node,
)
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):
@@ -341,8 +346,7 @@ 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.
@@ -377,8 +381,10 @@ def _sanitize_node(node: Any, path: str) -> Any:
_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]:
"""Apply *strip_node* (returns keywords removed) to every dict node of every tool's parameters, in place.
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.
Handles OpenAI format (``{"function": {"parameters": ...}}``) and Responses format
(``{"name": ..., "parameters": ...}``). Returns ``(tools, stripped_count)`` — same list.
@@ -431,8 +437,7 @@ def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]:
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]:
@@ -452,5 +457,4 @@ def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]:
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)")

View File

@@ -14,15 +14,13 @@ 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.
_WORKTREE_MUTATIONS = frozenset({
"checkout", "switch", "rebase", "merge", "pull", "restore", "clean",
"cherry-pick", "revert", "bisect",
})
"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"})
@@ -37,8 +35,7 @@ _KNOWN_GIT_BUILTINS = frozenset({
"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",
})
"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")
@@ -47,19 +44,19 @@ _NO_OPTIONS: frozenset[str] = frozenset()
_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",
}),
"-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,
"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
@@ -119,7 +116,9 @@ def _shell_words_at(command: str, start: int) -> list[str]:
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):
@@ -311,8 +310,7 @@ 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
@@ -410,7 +408,8 @@ 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
)
@@ -425,8 +424,7 @@ _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:
@@ -452,8 +450,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
value = result.stdout.strip()
@@ -461,9 +458,16 @@ 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.
@@ -491,7 +495,12 @@ def _inspect_git(
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
@@ -502,7 +511,12 @@ def _inspect_github_cli(
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)
return _find_mutation(script, current_dir, root, depth + 1) if script else None
@@ -513,8 +527,7 @@ _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:
@@ -576,8 +589,7 @@ def guard_active() -> bool:
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:
@@ -594,7 +606,8 @@ def detect_self_repo_git_mutation(
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. "
@@ -604,5 +617,4 @@ def _block_message(operation: str, root: Path) -> str:
"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."
)
"Hermes.")