From 8d6232878e1902dbbc8d1bc6030b59b4ae61b08b Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:44:46 -0700 Subject: [PATCH] refactor(tools): compact file_tools companions; one eviction helper in file_state --- tools/file_state.py | 154 ++++++++----------- tools/file_tools_paths.py | 45 ++---- tools/file_tools_read_tracking.py | 149 ++++++------------- tools/file_tools_write_guards.py | 236 ++++++++++-------------------- 4 files changed, 192 insertions(+), 392 deletions(-) diff --git a/tools/file_state.py b/tools/file_state.py index 77dd0577e0..99f3e2d069 100644 --- a/tools/file_state.py +++ b/tools/file_state.py @@ -1,24 +1,11 @@ """Cross-agent file state coordination. Prevents mangled edits when concurrent subagents (same process, same -filesystem) touch the same file: subagent B writes a file that subagent A -already read, so A's next write would clobber B's changes with stale content. -Complements the single-agent path-overlap check in -``run_agent._should_parallelize_tool_batch``. - -A process-wide ``FileStateRegistry`` tracks, per resolved path: - * per-agent read stamps {task_id: {path: (mtime, read_ts, partial)}} - * last writer globally {path: (task_id, write_ts)} - * a per-path ``threading.Lock`` for read->modify->write sections - -Hooks used by the file tools: ``record_read`` (read_file), ``note_write`` -(after write_file/patch), ``check_stale`` (BEFORE write_file/patch), -``lock_path`` (wrap the whole read->modify->write block) and ``writes_since`` -(delegate_tool's subagent-completion reminder). - -All methods are no-ops when ``HERMES_DISABLE_FILE_STATE_GUARD=1``. This is -separate from ``file_tools._read_tracker``, which handles per-task -consecutive-read loop detection. +filesystem) touch the same file: B writes a file A already read, so A's next +write would clobber B's changes. Complements the single-agent path-overlap +check in ``run_agent._should_parallelize_tool_batch``. A process-wide +``FileStateRegistry`` tracks per-agent read stamps, the global last writer and +a per-path lock; every method is a no-op under ``HERMES_DISABLE_FILE_STATE_GUARD=1``. """ from __future__ import annotations @@ -35,12 +22,45 @@ from typing import Dict, Iterable, List, Optional, Tuple # so the model re-reads in full. ReadStamp = Tuple[float, float, bool] -# Bounded so long sessions don't accumulate unbounded state; oldest by -# insertion order are dropped on overflow. +# Bounded so long sessions don't accumulate unbounded state. _MAX_PATHS_PER_AGENT = 4096 _MAX_GLOBAL_WRITERS = 4096 +def _disabled() -> bool: + # Re-read each call so tests can toggle via monkeypatch.setenv. + return os.environ.get("HERMES_DISABLE_FILE_STATE_GUARD", "").strip() == "1" + + +def _mtime_or_none(resolved: str) -> Optional[float]: + try: + return os.path.getmtime(resolved) + except OSError: + return None + + +def _fmt_ts(ts: float) -> str: + # Short wall-clock for warnings; avoids datetime formatting on the hot path. + return time.strftime("%H:%M:%S", time.localtime(ts)) + + +def _evict_oldest(container, cap: int) -> None: + """Pop entries until *container* is within *cap*. + + Sets pop arbitrary entries (they only feed diagnostic summaries); dicts pop + oldest by insertion order. An evicted entry costs one redundant re-send + (dedup) or one non-mtime staleness check — graceful degradation, not a bug. + """ + for _ in range(len(container) - cap): + try: + if isinstance(container, set): + container.pop() + else: + container.pop(next(iter(container))) + except (StopIteration, KeyError): + break + + class FileStateRegistry: """Process-wide coordinator for cross-agent file edits.""" @@ -51,50 +71,31 @@ class FileStateRegistry: self._meta_lock = threading.Lock() # guards _path_locks self._state_lock = threading.Lock() # guards _reads + _last_writer - def _lock_for(self, resolved: str) -> threading.Lock: - with self._meta_lock: - lock = self._path_locks.get(resolved) - if lock is None: - lock = threading.Lock() - self._path_locks[resolved] = lock - return lock - @contextmanager def lock_path(self, resolved: str): """Per-path lock: threads on the same path serialize, different paths proceed.""" - lock = self._lock_for(resolved) - lock.acquire() - try: + with self._meta_lock: + lock = self._path_locks.setdefault(resolved, threading.Lock()) + with lock: yield - finally: - lock.release() - def record_read( - self, - task_id: str, - resolved: str, - *, - partial: bool = False, - mtime: Optional[float] = None, - ) -> None: + def _stamp(self, task_id: str, resolved: str, mtime: float, now: float, partial: bool) -> None: + """Caller holds ``_state_lock``.""" + agent_reads = self._reads[task_id] + agent_reads[resolved] = (float(mtime), now, bool(partial)) + _evict_oldest(agent_reads, _MAX_PATHS_PER_AGENT) + + def record_read(self, task_id: str, resolved: str, *, partial: bool = False, + mtime: Optional[float] = None) -> None: if _disabled(): return mtime = _mtime_or_none(resolved) if mtime is None else mtime if mtime is None: return - now = time.time() with self._state_lock: - agent_reads = self._reads[task_id] - agent_reads[resolved] = (float(mtime), now, bool(partial)) - _cap_dict(agent_reads, _MAX_PATHS_PER_AGENT) + self._stamp(task_id, resolved, mtime, time.time(), partial) - def note_write( - self, - task_id: str, - resolved: str, - *, - mtime: Optional[float] = None, - ) -> None: + def note_write(self, task_id: str, resolved: str, *, mtime: Optional[float] = None) -> None: """Record a successful write: global last-writer AND this agent's own read stamp (a write is an implicit read of the current content).""" if _disabled(): @@ -105,9 +106,8 @@ class FileStateRegistry: now = time.time() with self._state_lock: self._last_writer[resolved] = (task_id, now) - _cap_dict(self._last_writer, _MAX_GLOBAL_WRITERS) - self._reads[task_id][resolved] = (float(mtime), now, False) - _cap_dict(self._reads[task_id], _MAX_PATHS_PER_AGENT) + _evict_oldest(self._last_writer, _MAX_GLOBAL_WRITERS) + self._stamp(task_id, resolved, mtime, now, False) def check_stale(self, task_id: str, resolved: str) -> Optional[str]: """Model-facing warning if this write would be stale, else ``None``. @@ -171,12 +171,8 @@ class FileStateRegistry: "Read the file first so you can write an informed edit." ) - def writes_since( - self, - exclude_task_id: str, - since_ts: float, - paths: Iterable[str], - ) -> Dict[str, List[str]]: + def writes_since(self, exclude_task_id: str, since_ts: float, + paths: Iterable[str]) -> Dict[str, List[str]]: """``{writer_task_id: [paths]}`` for writes after ``since_ts`` by agents other than ``exclude_task_id`` (delegate_task's "subagent modified files you previously read" reminder).""" @@ -213,36 +209,6 @@ def get_registry() -> FileStateRegistry: return _registry -def _disabled() -> bool: - # Re-read each call so tests can toggle via monkeypatch.setenv. - return os.environ.get("HERMES_DISABLE_FILE_STATE_GUARD", "").strip() == "1" - - -def _mtime_or_none(resolved: str) -> Optional[float]: - try: - return os.path.getmtime(resolved) - except OSError: - return None - - -def _fmt_ts(ts: float) -> str: - # Short wall-clock for warnings; avoids datetime formatting on the hot path. - return time.strftime("%H:%M:%S", time.localtime(ts)) - - -def _cap_dict(d: dict, limit: int) -> None: - """Trim ``d`` to ``limit`` entries by dropping the insertion-order oldest.""" - over = len(d) - limit - if over <= 0: - return - it = iter(d) - for _ in range(over): - try: - d.pop(next(it)) - except (StopIteration, KeyError): - break - - # Convenience wrappers (short names used at call sites). def record_read(task_id: str, resolved_or_path: str | Path, *, partial: bool = False) -> None: _registry.record_read(task_id, str(resolved_or_path), partial=partial) @@ -260,11 +226,7 @@ def lock_path(resolved_or_path: str | Path): return _registry.lock_path(str(resolved_or_path)) -def writes_since( - exclude_task_id: str, - since_ts: float, - paths: Iterable[str | Path], -) -> Dict[str, List[str]]: +def writes_since(exclude_task_id: str, since_ts: float, paths: Iterable[str | Path]) -> Dict[str, List[str]]: return _registry.writes_since(exclude_task_id, since_ts, [str(p) for p in paths]) diff --git a/tools/file_tools_paths.py b/tools/file_tools_paths.py index bb79f5c8bd..d25b0f6247 100644 --- a/tools/file_tools_paths.py +++ b/tools/file_tools_paths.py @@ -47,10 +47,7 @@ def _terminal_env_type_for_task(task_id: str = "default") -> str: """Best-effort terminal backend type for path-resolution decisions.""" try: from tools.terminal_tool import ( - _active_environments, - _env_lock, - _get_env_config, - _resolve_container_task_id, + _active_environments, _env_lock, _get_env_config, _resolve_container_task_id, ) try: @@ -61,14 +58,11 @@ def _terminal_env_type_for_task(task_id: str = "default") -> str: env = _active_environments.get(container_key) or _active_environments.get(task_id) if env is not None: name = env.__class__.__name__.lower() - for hint in _ENV_CLASS_NAME_HINTS: - if hint in name: - return hint + hint = next((h for h in _ENV_CLASS_NAME_HINTS if h in name), None) stamped = getattr(env, "_hermes_backend_name", None) - if isinstance(stamped, str) and stamped: - return stamped - cfg = _get_env_config() - return str(cfg.get("env_type") or os.getenv("TERMINAL_ENV") or "local").lower() + if hint or (isinstance(stamped, str) and stamped): + return hint or stamped + return str(_get_env_config().get("env_type") or os.getenv("TERMINAL_ENV") or "local").lower() except Exception: return str(os.getenv("TERMINAL_ENV") or "local").lower() @@ -103,9 +97,7 @@ def _sentinel_free_abs_cwd(raw: str | None) -> str | None: if raw.lower() in _TERMINAL_CWD_SENTINELS: return None expanded = _expand_tilde(raw) - if not os.path.isabs(expanded): - return None - return expanded + return expanded if os.path.isabs(expanded) else None def _configured_terminal_cwd() -> str | None: @@ -151,12 +143,7 @@ def _authoritative_workspace_root(task_id: str = "default") -> str | None: recorded = get_session_cwd(task_id) except Exception: recorded = None - if recorded: - return recorded - registered = _registered_task_cwd_override(task_id) - if registered: - return registered - return _configured_terminal_cwd() + return recorded or _registered_task_cwd_override(task_id) or _configured_terminal_cwd() def _resolve_base_dir( @@ -255,16 +242,14 @@ def _path_resolution_warning(filepath: str, resolved: Path, task_id: str = "defa root = _normalize_without_host_deref(Path(_expand_tilde(workspace_root))) else: root = Path(_expand_tilde(workspace_root)).resolve() - try: - resolved.relative_to(root) + if resolved.is_relative_to(root): return None - except ValueError: - return ( - f"Relative path {filepath!r} resolved to {str(resolved)!r}, which is " - f"OUTSIDE the active workspace ({str(root)!r}). The edit will land in " - f"a different directory than the terminal's cwd. If this is not " - f"intended (e.g. a git-worktree session writing into the main " - f"checkout), pass an absolute path under the workspace instead." - ) + return ( + f"Relative path {filepath!r} resolved to {str(resolved)!r}, which is " + f"OUTSIDE the active workspace ({str(root)!r}). The edit will land in " + f"a different directory than the terminal's cwd. If this is not " + f"intended (e.g. a git-worktree session writing into the main " + f"checkout), pass an absolute path under the workspace instead." + ) except Exception: return None diff --git a/tools/file_tools_read_tracking.py b/tools/file_tools_read_tracking.py index 1789c76bdc..de50abe7af 100644 --- a/tools/file_tools_read_tracking.py +++ b/tools/file_tools_read_tracking.py @@ -1,23 +1,12 @@ """Per-task read/search bookkeeping for the file tools. -Process-lifetime state behind ``read_file`` / ``search_files`` / ``write_file`` / -``patch``; ``tools.file_tools`` re-imports every name here. Per task_id -``_read_tracker`` stores: - - last_key / consecutive most recent read-or-search key and its repeat count - (loop detection; reset by any OTHER tool call). - read_history set of (path, offset, limit) — diagnostic summaries only. - dedup (resolved_path, offset, limit) -> mtime; skip identical - re-reads of unchanged files. Cleared on context - compression (the original content was summarised away). - dedup_hits per-key count of stub returns, to break stub loops. - read_timestamps resolved_path -> mtime at last read/write by this task; - write/patch warn when the file changed underneath. - not_found (op, resolved_path) -> (monotonic, cached error JSON); - short-TTL negative cache for retried missing paths. - -Every container is hard-capped (``_cap_read_tracker_data``) so a long CLI -session accretes a few hundred KB at most instead of ~1.5MB per 10k reads. +Process-lifetime state behind read_file/search_files/write_file/patch; +``tools.file_tools`` re-imports every name here. Per task_id ``_read_tracker`` +stores: ``last_key``/``consecutive`` (loop detection; reset by any OTHER tool +call), ``read_history`` (diagnostics), ``dedup`` (key -> mtime; cleared on +context compression), ``dedup_hits`` (stub-loop breaker), ``read_timestamps`` +(staleness warnings) and ``not_found`` (short-TTL negative cache). Every +container is hard-capped (``_cap_read_tracker_data``) so long sessions stay small. """ import logging @@ -25,6 +14,7 @@ import os import threading import time +from tools.file_state import _evict_oldest from tools.file_tools_paths import _authoritative_workspace_root, _resolve_path_for_task logger = logging.getLogger("tools.file_tools") @@ -33,8 +23,7 @@ _read_tracker_lock = threading.Lock() _read_tracker: dict = {} # Consecutive patch failures per (task_id, resolved_path); escalates the hint -# when the model keeps failing the same file (stale view, ambiguous old_string). -# Reset on a successful patch to that path. +# when the model keeps failing the same file. Reset on a successful patch. _patch_failure_lock = threading.Lock() _patch_failure_tracker: dict = {} # {task_id: {resolved_path: count}} _PATCH_FAILURE_PATHS_CAP = 64 @@ -48,24 +37,17 @@ _NOT_FOUND_CAP = 500 _NOT_FOUND_TTL_SECONDS = 60.0 # a path that didn't exist may be created soon -def _new_task_data() -> dict: - return { - "last_key": None, "consecutive": 0, - "read_history": set(), "dedup": {}, - "dedup_hits": {}, "read_timestamps": {}, - } - - def _task_data(task_id: str) -> dict: """Get-or-create the tracker entry for *task_id*, back-filling any missing keys. Must be called with ``_read_tracker_lock`` held. Entries created by older code paths (or injected by tests) may lack the newer containers. """ - task_data = _read_tracker.setdefault(task_id, _new_task_data()) - for key, factory in (("dedup", dict), ("dedup_hits", dict), ("read_timestamps", dict)): - if key not in task_data: - task_data[key] = factory() + task_data = _read_tracker.setdefault(task_id, { + "last_key": None, "consecutive": 0, "read_history": set(), + }) + for key in ("dedup", "dedup_hits", "read_timestamps"): + task_data.setdefault(key, {}) return task_data @@ -74,11 +56,8 @@ def _record_patch_failure(task_id: str, resolved_path: str) -> int: with _patch_failure_lock: task_failures = _patch_failure_tracker.setdefault(task_id, {}) # Evict the oldest entry once a task has failed on many distinct files. - if len(task_failures) >= _PATCH_FAILURE_PATHS_CAP and resolved_path not in task_failures: - try: - del task_failures[next(iter(task_failures))] - except StopIteration: - pass + if resolved_path not in task_failures: + _evict_oldest(task_failures, _PATCH_FAILURE_PATHS_CAP - 1) task_failures[resolved_path] = task_failures.get(resolved_path, 0) + 1 return task_failures[resolved_path] @@ -89,29 +68,10 @@ def _reset_patch_failures(task_id: str, resolved_paths: list) -> None: return with _patch_failure_lock: task_failures = _patch_failure_tracker.get(task_id) - if not task_failures: - return - for rp in resolved_paths: + for rp in resolved_paths if task_failures else (): task_failures.pop(rp, None) -def _evict_oldest(container, cap: int) -> None: - """Pop entries until *container* is within *cap*. - - Sets pop arbitrary entries (they only feed diagnostic summaries); dicts pop - oldest by insertion order. An evicted entry costs one redundant re-send - (dedup) or one non-mtime staleness check — graceful degradation, not a bug. - """ - for _ in range(len(container) - cap): - try: - if isinstance(container, set): - container.pop() - else: - container.pop(next(iter(container))) - except (StopIteration, KeyError): - break - - def _cap_read_tracker_data(task_data: dict) -> None: """Enforce size caps on the per-task sub-containers. Call with ``_read_tracker_lock`` held.""" # Caps are read at call time so tests can monkeypatch the module constants. @@ -127,6 +87,14 @@ def _cap_read_tracker_data(task_data: dict) -> None: _evict_oldest(container, cap) +def _pop_not_found(op: str, resolved_str: str, task_id: str) -> None: + """Drop the negative-cache entry for *(op, resolved_str)*. Lock must be held.""" + task_data = _read_tracker.get(task_id) + nf = task_data.get("not_found") if task_data else None + if nf: + nf.pop((op, resolved_str), None) + + def _check_not_found_cache(op: str, resolved_str: str, task_id: str) -> str | None: """Return cached not-found JSON for *(op, resolved_str)* if still fresh. @@ -136,28 +104,20 @@ def _check_not_found_cache(op: str, resolved_str: str, task_id: str) -> str | No """ with _read_tracker_lock: task_data = _read_tracker.get(task_id) - if not task_data: - return None - nf = task_data.get("not_found") - if not nf: - return None - entry = nf.get((op, resolved_str)) + entry = (task_data.get("not_found") or {}).get((op, resolved_str)) if task_data else None if entry is None: return None ts, cached_json = entry if time.monotonic() - ts > _NOT_FOUND_TTL_SECONDS: - nf.pop((op, resolved_str), None) + _pop_not_found(op, resolved_str, task_id) return None # The path may have been created since the miss was cached (terminal, - # another agent, ...) — the "check → create → read" pattern is common, so - # serving a stale miss breaks it. The stat runs OUTSIDE the global tracker - # lock: a hung stat on a dead network mount must not stall every task. + # another agent, ...) — "check → create → read" is common, so serving a + # stale miss breaks it. The stat runs OUTSIDE the global tracker lock: a + # hung stat on a dead network mount must not stall every task. if os.path.exists(resolved_str): with _read_tracker_lock: - task_data = _read_tracker.get(task_id) - nf = task_data.get("not_found") if task_data else None - if nf: - nf.pop((op, resolved_str), None) + _pop_not_found(op, resolved_str, task_id) return None return cached_json @@ -166,8 +126,7 @@ def _record_not_found(op: str, resolved_str: str, task_id: str, error_json: str) """Cache a not-found error so the next *op* call for *resolved_str* skips I/O.""" with _read_tracker_lock: task_data = _task_data(task_id) - nf = task_data.setdefault("not_found", {}) - nf[(op, resolved_str)] = (time.monotonic(), error_json) + task_data.setdefault("not_found", {})[(op, resolved_str)] = (time.monotonic(), error_json) _cap_read_tracker_data(task_data) @@ -212,11 +171,9 @@ def notify_other_tool_call(task_id: str = "default"): if task_data: task_data["last_key"] = None task_data["consecutive"] = 0 - if "dedup_hits" in task_data: - task_data["dedup_hits"].clear() - nf = task_data.get("not_found") - if nf: - nf.clear() + for key in ("dedup_hits", "not_found"): + if task_data.get(key): + task_data[key].clear() def _invalidate_dedup_for_path(filepath: str, task_id: str) -> None: @@ -237,10 +194,8 @@ def _invalidate_dedup_for_path(filepath: str, task_id: str) -> None: if dedup: for k in [k for k in dedup if k[0] == resolved]: del dedup[k] - nf = task_data.get("not_found") - if nf: - nf.pop(("read", resolved), None) - nf.pop(("search", resolved), None) + _pop_not_found("read", resolved, task_id) + _pop_not_found("search", resolved, task_id) def _update_read_timestamp(filepath: str, task_id: str) -> None: @@ -271,9 +226,7 @@ def _check_file_staleness(filepath: str, task_id: str) -> str | None: return None with _read_tracker_lock: task_data = _read_tracker.get(task_id) - if not task_data: - return None - read_mtime = task_data.get("read_timestamps", {}).get(resolved) + read_mtime = task_data.get("read_timestamps", {}).get(resolved) if task_data else None if read_mtime is None: return None try: @@ -289,11 +242,8 @@ def _check_file_staleness(filepath: str, task_id: str) -> str | None: return None -def _mark_verification_stale( - task_id: str, - resolved_paths: list[str], - session_id: str | None = None, -) -> None: +def _mark_verification_stale(task_id: str, resolved_paths: list[str], + session_id: str | None = None) -> None: """Best-effort note that successful edits made prior verification stale. The workspace cwd is the first edited path's project root when one is @@ -308,22 +258,9 @@ def _mark_verification_stale( from agent.coding_context import project_facts_for from agent.verification_evidence import mark_workspace_edited - cwd = None - for path in paths: - try: - candidate = str(Path(path).parent) - except Exception: - continue - if project_facts_for(candidate): - cwd = candidate - break - if cwd is None: - cwd = _authoritative_workspace_root(task_id) - if cwd is None: - try: - cwd = str(Path(paths[0]).parent) - except Exception: - cwd = None + parents = [str(Path(p).parent) for p in paths] + cwd = (next((c for c in parents if project_facts_for(c)), None) + or _authoritative_workspace_root(task_id) or parents[0]) mark_workspace_edited(session_id=session_id or task_id, cwd=cwd, paths=paths) except Exception: logger.debug("verification stale marker failed", exc_info=True) diff --git a/tools/file_tools_write_guards.py b/tools/file_tools_write_guards.py index 294ce52715..f17db4ee23 100644 --- a/tools/file_tools_write_guards.py +++ b/tools/file_tools_write_guards.py @@ -2,14 +2,10 @@ Every guard returns ``None`` when the write may proceed, else an error string the tool returns verbatim. ``tools.file_tools`` re-imports every name here. - -Guards (in the order the tools apply them): - * ``_check_sensitive_path`` — system paths + the Hermes config.yaml (hard deny). - * ``_check_binary_document_write`` — text write would corrupt an Office/PDF container. - * ``_check_protected_instruction_write`` — AGENTS.md-style files: ALWAYS ask, no yolo bypass. - * ``_check_approval_required_write`` — ~/.ssh/config-style files: normal approval gate. - * ``_check_cross_profile_path`` — sandbox-mirror writes the host never reads (lost-work guard). - * ``_is_internal_file_tool_content`` — refuse to persist read_file display text as a file. +Guards, in the order the tools apply them: ``_check_sensitive_path`` (hard +deny), ``_check_binary_document_write``, ``_check_protected_instruction_write`` +(ALWAYS ask), ``_check_approval_required_write`` (normal gate), +``_check_cross_profile_path`` (sandbox-mirror lost-work), ``_is_internal_file_tool_content``. """ import fnmatch @@ -31,25 +27,42 @@ _SENSITIVE_EXACT_PATHS = {"/var/run/docker.sock", "/run/docker.sock"} _hermes_config_resolved: str | None = None _hermes_config_resolved_loaded = False +_real_hermes_home_cached: str | None = None +_real_hermes_home_loaded = False def _get_hermes_config_resolved() -> str | None: """Return the resolved absolute path of the Hermes config file (cached).""" global _hermes_config_resolved, _hermes_config_resolved_loaded - if _hermes_config_resolved_loaded: - return _hermes_config_resolved - _hermes_config_resolved_loaded = True - try: - from hermes_cli.config import get_config_path - _hermes_config_resolved = str(get_config_path().resolve()) - except Exception: + if not _hermes_config_resolved_loaded: + _hermes_config_resolved_loaded = True try: - _hermes_config_resolved = str(Path(_expand_tilde("~/.hermes/config.yaml")).resolve()) + from hermes_cli.config import get_config_path + _hermes_config_resolved = str(get_config_path().resolve()) except Exception: - _hermes_config_resolved = None + try: + _hermes_config_resolved = str(Path(_expand_tilde("~/.hermes/config.yaml")).resolve()) + except Exception: + _hermes_config_resolved = None return _hermes_config_resolved +def _get_real_hermes_home() -> str | None: + """Return the realpath of the authoritative Hermes home (cached).""" + global _real_hermes_home_cached, _real_hermes_home_loaded + if not _real_hermes_home_loaded: + _real_hermes_home_loaded = True + try: + from hermes_constants import get_hermes_home + _real_hermes_home_cached = os.path.realpath(str(get_hermes_home())) + except Exception: + try: + _real_hermes_home_cached = os.path.realpath(_expand_tilde("~/.hermes")) + except Exception: + _real_hermes_home_cached = None + return _real_hermes_home_cached + + def _resolved_or_raw(filepath: str, task_id: str) -> str: """Task-resolved path string, falling back to the raw input on resolution failure.""" try: @@ -60,21 +73,16 @@ def _resolved_or_raw(filepath: str, task_id: str) -> str: def _check_sensitive_path(filepath: str, task_id: str = "default") -> str | None: """Return an error message if the path targets a sensitive system location.""" - resolved = _resolved_or_raw(filepath, task_id) - normalized = os.path.normpath(_expand_tilde(filepath)) - _err = ( - f"Refusing to write to sensitive system path: {filepath}\n" - "Use the terminal tool with sudo if you need to modify system files." - ) - for prefix in _SENSITIVE_PATH_PREFIXES: - if resolved.startswith(prefix) or normalized.startswith(prefix): - return _err - if resolved in _SENSITIVE_EXACT_PATHS or normalized in _SENSITIVE_EXACT_PATHS: - return _err + candidates = (_resolved_or_raw(filepath, task_id), os.path.normpath(_expand_tilde(filepath))) + if any(c.startswith(_SENSITIVE_PATH_PREFIXES) or c in _SENSITIVE_EXACT_PATHS for c in candidates): + return ( + f"Refusing to write to sensitive system path: {filepath}\n" + "Use the terminal tool with sudo if you need to modify system files." + ) # approvals.mode and other security settings live in config.yaml; a # prompt-injected agent could silently disable exec approval by editing it. hermes_config = _get_hermes_config_resolved() - if hermes_config and (resolved == hermes_config or normalized == hermes_config): + if hermes_config and hermes_config in candidates: return ( f"Refusing to write to Hermes config file: {filepath}\n" "Agent cannot modify security-sensitive configuration. " @@ -88,37 +96,13 @@ def _check_sensitive_path(filepath: str, task_id: str = "default") -> str | None # --------------------------------------------------------------------------- # Files that steer FUTURE agent behavior are a prompt-injection persistence # vector: an injected edit to AGENTS.md / CLAUDE.md / SOUL.md / .cursorrules -# (or a project-local .hermes tree) outlives the turn and poisons every later -# session. Writes ALWAYS require human approval — even under --yolo — and fail -# closed when no human channel exists. (Ported from Roo-Code's -# RooProtectedController; the terminal-tool vector is gated separately.) -# -# Basenames match in ANY directory (instruction files load from cwd trees) and -# case-insensitively (case-insensitive filesystems; loaders probe variants). +# (or a project-local .hermes tree) outlives the turn. Writes ALWAYS require +# human approval — even under --yolo — and fail closed without a human channel. +# Basenames match in ANY directory and case-insensitively (loaders probe variants). _PROTECTED_INSTRUCTION_BASENAMES = frozenset({ "agents.md", "claude.md", "soul.md", ".cursorrules", }) -_real_hermes_home_cached: str | None = None -_real_hermes_home_loaded = False - - -def _get_real_hermes_home() -> str | None: - """Return the realpath of the authoritative Hermes home (cached).""" - global _real_hermes_home_cached, _real_hermes_home_loaded - if _real_hermes_home_loaded: - return _real_hermes_home_cached - _real_hermes_home_loaded = True - try: - from hermes_constants import get_hermes_home - _real_hermes_home_cached = os.path.realpath(str(get_hermes_home())) - except Exception: - try: - _real_hermes_home_cached = os.path.realpath(_expand_tilde("~/.hermes")) - except Exception: - _real_hermes_home_cached = None - return _real_hermes_home_cached - def _protected_instruction_config() -> tuple[bool, list[str]]: """Return ``(enabled, extra_patterns)`` from ``security.protected_instruction_files`` / @@ -129,10 +113,8 @@ def _protected_instruction_config() -> tuple[bool, list[str]]: try: from hermes_cli.config import load_config, cfg_get cfg = load_config() - enabled = cfg_get(cfg, "security", "protected_instruction_files", - default=True) - extra = cfg_get(cfg, "security", "protected_instruction_extra_patterns", - default=[]) + enabled = cfg_get(cfg, "security", "protected_instruction_files", default=True) + extra = cfg_get(cfg, "security", "protected_instruction_extra_patterns", default=[]) except Exception: return True, [] if not isinstance(enabled, bool): @@ -166,18 +148,15 @@ def _protected_instruction_reason(filepath: str, task_id: str = "default", # mirror guard, write_approval); this gate targets PROJECT-LOCAL files only. # Must run before the ``.hermes`` component rule, which would match the home. real_home = _get_real_hermes_home() - if real_home and (resolved == real_home - or resolved.startswith(real_home + os.sep)): + if real_home and (resolved == real_home or resolved.startswith(real_home + os.sep)): return None for candidate in (normalized, resolved): base = os.path.basename(candidate) base_lower = base.lower() - if base_lower in _PROTECTED_INSTRUCTION_BASENAMES: + if base_lower in _PROTECTED_INSTRUCTION_BASENAMES or any( + fnmatch.fnmatch(base_lower, pattern.lower()) for pattern in extra_patterns): return base - for pattern in extra_patterns: - if fnmatch.fnmatch(base_lower, pattern.lower()): - return base # Project-local .hermes config dirs (/.hermes/config.yaml) steer # behavior too. Only the IMMEDIATE parent counts — matching any ancestor # would gate every write inside a checkout living under ~/.hermes. @@ -187,8 +166,11 @@ def _protected_instruction_reason(filepath: str, task_id: str = "default", return None -def _request_protected_instruction_approval( - reasons: list[str], task_id: str = "default") -> str | None: +_APPROVAL_UNAVAILABLE = "requires approval but the approval subsystem is unavailable." +_NO_HUMAN = "requires approval but no interactive user or gateway is present to approve it." + + +def _request_protected_instruction_approval(reasons: list[str], task_id: str = "default") -> str | None: """Ask the human to approve a write to protected instruction file(s). Returns ``None`` when approved, else a BLOCKED error string. Deliberately @@ -209,21 +191,17 @@ def _request_protected_instruction_approval( "attempt the same edit via another path (terminal, execute_code, " "etc.)." ) - timed_out = blocked.format( - why="approval prompt timed out without a user response. " - "Silence is not consent.") + timed_out = blocked.format(why="approval prompt timed out without a user response. Silence is not consent.") denied = blocked.format(why="was denied by the user.") try: import tools.approval as _approval except Exception: - return blocked.format(why="requires approval but the approval " - "subsystem is unavailable.") + return blocked.format(why=_APPROVAL_UNAVAILABLE) # Gateway surface: block on the button round-trip when a notify callback # is registered for this session. One-operation only — no scope buttons. session_key = _approval.get_current_session_key() - notify_cb = None try: with _approval._lock: notify_cb = _approval._gateway_notify_cbs.get(session_key) @@ -239,23 +217,15 @@ def _request_protected_instruction_approval( "allow_permanent": False, "allow_session": False, } - decision = _approval._await_gateway_decision( - session_key, notify_cb, approval_data, surface="gateway", - ) + decision = _approval._await_gateway_decision(session_key, notify_cb, approval_data, surface="gateway") if decision.get("notify_failed"): - return blocked.format( - why="requires approval but the approval request could not " - "be delivered.") - choice = decision.get("choice") + return blocked.format(why="requires approval but the approval request could not be delivered.") # Any tapped scope is a one-operation grant; nothing is persisted. - if decision.get("resolved") and choice in {"once", "session", "always"}: + if decision.get("resolved") and decision.get("choice") in {"once", "session", "always"}: return None - if not decision.get("resolved"): - return timed_out - return denied + return timed_out if not decision.get("resolved") else denied # CLI surface: per-thread approval callback (prompt_toolkit panel). - callback = None try: from tools.terminal_tool import _get_approval_callback callback = _get_approval_callback() @@ -264,26 +234,17 @@ def _request_protected_instruction_approval( if callback is not None: choice = _approval.prompt_dangerous_approval( - display, description, - allow_permanent=False, - allow_session=False, - approval_callback=callback, - ) + display, description, allow_permanent=False, allow_session=False, approval_callback=callback) if choice in {"once", "session", "always"}: return None - if choice == "timeout": - return timed_out - return denied + return timed_out if choice == "timeout" else denied # No human channel (script, cron, background thread): fail closed — # auto-approving here would recreate the persistence vector. - return blocked.format( - why="requires approval but no interactive user or gateway is " - "present to approve it.") + return blocked.format(why=_NO_HUMAN) -def _check_protected_instruction_write(paths: list[str], - task_id: str = "default") -> str | None: +def _check_protected_instruction_write(paths: list[str], task_id: str = "default") -> str | None: """Gate a write/patch touching protected instruction files. ONE protected file gates the ENTIRE multi-file patch: a single prompt lists @@ -293,19 +254,14 @@ def _check_protected_instruction_write(paths: list[str], enabled, extra = _protected_instruction_config() if not enabled: return None - reasons = [ - r for r in ( - _protected_instruction_reason(p, task_id, enabled=enabled, extra_patterns=extra) - for p in paths - ) if r - ] + reasons = [r for r in (_protected_instruction_reason(p, task_id, enabled=enabled, extra_patterns=extra) + for p in paths) if r] if not reasons: return None return _request_protected_instruction_approval(reasons, task_id) -def _check_approval_required_write(paths: list[str], - task_id: str = "default") -> str | None: +def _check_approval_required_write(paths: list[str], task_id: str = "default") -> str | None: """Gate a write/patch touching an approval-required path (``~/.ssh/config``). Not credentials and not hard-denied, but they can steer process execution @@ -337,15 +293,13 @@ def _check_approval_required_write(paths: list[str], try: import tools.approval as _approval except Exception: - return blocked.format(why="requires approval but the approval " - "subsystem is unavailable.") + return blocked.format(why=_APPROVAL_UNAVAILABLE) result = _approval._run_approval_gate( pattern_key="ssh_config_write", description=description, display_target=f"", - cron_deny_message=blocked.format( - why="requires approval but this cron session denies it."), + cron_deny_message=blocked.format(why="requires approval but this cron session denies it."), single_query_deny_message=blocked.format( why="requires approval but single-query (-q) sessions run " "without a user present to approve it. To allow flagged " @@ -353,9 +307,7 @@ def _check_approval_required_write(paths: list[str], "approve in config.yaml."), autoapprove_log_prefix="ssh_config_write", fail_closed_when_no_human=True, - no_human_block_message=blocked.format( - why="requires approval but no interactive user or gateway is " - "present to approve it."), + no_human_block_message=blocked.format(why=_NO_HUMAN), ) if result.get("approved"): return None @@ -366,31 +318,18 @@ def _get_container_mirror_prefix_for_task(task_id: str = "default") -> str | Non """Return the container-side Hermes mirror prefix for persistent Docker file tools.""" try: from tools.terminal_tool import ( - _active_environments, - _env_lock, - _get_env_config, - _resolve_container_task_id, + _active_environments, _env_lock, _get_env_config, _resolve_container_task_id, ) - container_key = _resolve_container_task_id(task_id) - except Exception: - return None - - try: with _env_lock: env = _active_environments.get(container_key) or _active_environments.get(task_id) - if env is not None: - if env.__class__.__name__ == "DockerEnvironment" and bool( - getattr(env, "_persistent", False) - ): - return "/root/.hermes" - return None - + persistent_docker = (env.__class__.__name__ == "DockerEnvironment" + and bool(getattr(env, "_persistent", False))) + return "/root/.hermes" if persistent_docker else None config = _get_env_config() except Exception: return None - if config.get("env_type") == "docker" and config.get("container_persistent", True): return "/root/.hermes" return None @@ -406,23 +345,14 @@ def _check_cross_profile_path(filepath: str, task_id: str = "default") -> str | Fails open on import error — the sensitive-path guard and denylist still apply. """ try: - from agent.file_safety import ( - get_container_mirror_warning, - get_sandbox_mirror_warning, - ) + from agent.file_safety import get_container_mirror_warning, get_sandbox_mirror_warning except Exception: return None - resolved = _resolved_or_raw(filepath, task_id) - warning = get_sandbox_mirror_warning(resolved) if warning is not None: return warning - - return get_container_mirror_warning( - resolved, - mirror_prefix=_get_container_mirror_prefix_for_task(task_id), - ) + return get_container_mirror_warning(resolved, mirror_prefix=_get_container_mirror_prefix_for_task(task_id)) def _check_binary_document_write(filepath: str, task_id: str = "default") -> str | None: @@ -484,12 +414,8 @@ def _is_internal_file_status_text(content: str) -> bool: if not isinstance(content, str): return False stripped = content.strip() - if not stripped: - return False - if stripped == _READ_DEDUP_STATUS_MESSAGE: - return True - return (_READ_DEDUP_STATUS_MESSAGE in stripped - and len(stripped) <= 2 * len(_READ_DEDUP_STATUS_MESSAGE)) + return bool(stripped) and _READ_DEDUP_STATUS_MESSAGE in stripped and ( + len(stripped) <= 2 * len(_READ_DEDUP_STATUS_MESSAGE)) def _looks_like_read_file_line_numbered_content(content: str) -> bool: @@ -501,30 +427,20 @@ def _looks_like_read_file_line_numbered_content(content: str) -> bool: """ if not isinstance(content, str): return False - lines = [line for line in content.splitlines() if line.strip()] if len(lines) < 2: return False - numbered: list[int] = [] for line in lines: prefix, sep, _rest = line.lstrip().partition("|") if sep and prefix.isdigit(): numbered.append(int(prefix)) - if len(numbered) < 2 or len(numbered) / len(lines) < 0.6: return False - - consecutive_pairs = sum( - 1 for prev, current in zip(numbered, numbered[1:]) - if current == prev + 1 - ) + consecutive_pairs = sum(1 for prev, current in zip(numbered, numbered[1:]) if current == prev + 1) return consecutive_pairs >= len(numbered) - 1 def _is_internal_file_tool_content(content: str) -> bool: """Return True when content is file-tool display text, not intended file bytes.""" - return ( - _is_internal_file_status_text(content) - or _looks_like_read_file_line_numbered_content(content) - ) + return _is_internal_file_status_text(content) or _looks_like_read_file_line_numbered_content(content)