Files
hermes-agent/tools/file_state.py

285 lines
9.8 KiB
Python

"""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.
"""
from __future__ import annotations
import os
import threading
import time
from collections import defaultdict
from contextlib import contextmanager
from pathlib import Path
from typing import Dict, Iterable, List, Optional, Tuple
# (mtime, read_ts, partial). partial=True when read_file returned a windowed
# view (offset > 1 or limit < total_lines) — a later write should still warn
# 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.
_MAX_PATHS_PER_AGENT = 4096
_MAX_GLOBAL_WRITERS = 4096
class FileStateRegistry:
"""Process-wide coordinator for cross-agent file edits."""
def __init__(self) -> None:
self._reads: Dict[str, Dict[str, ReadStamp]] = defaultdict(dict)
self._last_writer: Dict[str, Tuple[str, float]] = {}
self._path_locks: Dict[str, threading.Lock] = {}
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:
yield
finally:
lock.release()
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)
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():
return
mtime = _mtime_or_none(resolved) if mtime is None else mtime
if mtime is None:
return
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)
def check_stale(self, task_id: str, resolved: str) -> Optional[str]:
"""Model-facing warning if this write would be stale, else ``None``.
Checked in severity order: (1) a sibling subagent wrote after this
agent's last read; (2) mtime drifted since our read (external edit) or
the read was partial; (3) this agent never read the file. Never raises
— callers decide whether to block or warn.
"""
if _disabled():
return None
with self._state_lock:
stamp = self._reads.get(task_id, {}).get(resolved)
last_writer = self._last_writer.get(resolved)
# Never read and no write record: net-new file or first touch —
# existing sensitive-path / file-exists logic handles it.
if stamp is None and last_writer is None:
return None
current_mtime = _mtime_or_none(resolved)
if current_mtime is None:
return None # file doesn't exist — write creates it; not stale
if last_writer is not None:
writer_tid, writer_ts = last_writer
if writer_tid != task_id:
if stamp is None:
return (
f"{resolved} was modified by sibling subagent "
f"{writer_tid!r} but this agent never read it. "
"Read the file before writing to avoid overwriting "
"the sibling's changes."
)
read_ts = stamp[1]
if writer_ts > read_ts:
return (
f"{resolved} was modified by sibling subagent "
f"{writer_tid!r} at {_fmt_ts(writer_ts)} — after "
f"this agent's last read at {_fmt_ts(read_ts)}. "
"Re-read the file before writing."
)
if stamp is not None:
read_mtime, _read_ts, partial = stamp
if current_mtime != read_mtime:
return (
f"{resolved} was modified since you last read it "
"on disk (external edit or unrecorded writer). "
"Re-read the file before writing."
)
if partial:
return (
f"{resolved} was last read with offset/limit pagination "
"(partial view). Re-read the whole file before "
"overwriting it."
)
return None
return (
f"{resolved} was not read by this agent. "
"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]]:
"""``{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)."""
if _disabled():
return {}
paths_set = set(paths)
out: Dict[str, List[str]] = defaultdict(list)
with self._state_lock:
for p, (writer_tid, ts) in self._last_writer.items():
if writer_tid != exclude_task_id and ts >= since_ts and p in paths_set:
out[writer_tid].append(p)
return dict(out)
def known_reads(self, task_id: str) -> List[str]:
"""Resolved paths this agent has read."""
if _disabled():
return []
with self._state_lock:
return list(self._reads.get(task_id, {}).keys())
def clear(self) -> None:
"""Reset all state. Intended for tests only."""
with self._state_lock:
self._reads.clear()
self._last_writer.clear()
with self._meta_lock:
self._path_locks.clear()
_registry = FileStateRegistry()
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)
def note_write(task_id: str, resolved_or_path: str | Path) -> None:
_registry.note_write(task_id, str(resolved_or_path))
def check_stale(task_id: str, resolved_or_path: str | Path) -> Optional[str]:
return _registry.check_stale(task_id, str(resolved_or_path))
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]]:
return _registry.writes_since(exclude_task_id, since_ts, [str(p) for p in paths])
def known_reads(task_id: str) -> List[str]:
return _registry.known_reads(task_id)
__all__ = [
"FileStateRegistry",
"get_registry",
"record_read",
"note_write",
"check_stale",
"lock_path",
"writes_since",
"known_reads",
]