refactor(tools): extract process checkpoint persistence
This commit is contained in:
committed by
Teknium
parent
3bc839508e
commit
5695ebc40f
@@ -32,6 +32,7 @@ from hermes_cli.config import get_hermes_home
|
||||
|
||||
from agent.redact import redact_sensitive_text
|
||||
from tools.process_registry_notifications import format_process_notification
|
||||
from tools.process_registry_checkpoint import ProcessCheckpointMixin
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -384,7 +385,7 @@ _CHECKPOINT_DEFAULTS = {
|
||||
}
|
||||
|
||||
|
||||
class ProcessRegistry:
|
||||
class ProcessRegistry(ProcessCheckpointMixin):
|
||||
"""In-memory registry of running and finished background processes.
|
||||
Thread-safe: accessed from executor threads (terminal_tool, process handlers),
|
||||
the gateway asyncio loop (watchers, reset checks) and the cleanup thread."""
|
||||
@@ -1893,97 +1894,6 @@ class ProcessRegistry:
|
||||
self._completion_consumed &= tracked
|
||||
self._poll_observed &= tracked
|
||||
|
||||
# ----- Checkpoint (crash recovery) -----
|
||||
|
||||
def _write_checkpoint(self, extra_entries: Optional[List[Dict[str, Any]]] = None):
|
||||
"""Write running process metadata to the checkpoint file atomically."""
|
||||
try:
|
||||
with self._lock:
|
||||
entries = []
|
||||
for s in self._running.values():
|
||||
if s.exited:
|
||||
continue
|
||||
# Backfill the start time so recovery can detect PID recycling
|
||||
# even for sessions spawned before this field existed.
|
||||
if s.host_start_time is None and s.pid_scope == "host" and s.pid:
|
||||
s.host_start_time = self._safe_host_start_time(s.pid)
|
||||
entry = {"session_id": s.id, **{f: getattr(s, f) for f in _CHECKPOINT_FIELDS}}
|
||||
# Redact inline credentials before persisting (~/.hermes/processes.json).
|
||||
# Recovery uses command only for display (adoption re-validates the
|
||||
# PID, never re-runs it), so masking is lossless.
|
||||
# See #77484.
|
||||
entry["command"] = redact_sensitive_text(s.command, code_file=True)
|
||||
entry["owner_task_id"] = s.owner_task_id or s.task_id
|
||||
entries.append(entry)
|
||||
if extra_entries:
|
||||
tracked_ids = {item.get("session_id") for item in entries}
|
||||
entries.extend(item for item in extra_entries if item.get("session_id") not in tracked_ids)
|
||||
from utils import atomic_json_write
|
||||
atomic_json_write(CHECKPOINT_PATH, entries)
|
||||
except Exception as e:
|
||||
logger.debug("Failed to write checkpoint file: %s", e, exc_info=True)
|
||||
|
||||
def recover_from_checkpoint(self) -> int:
|
||||
"""On gateway startup, probe PIDs from the checkpoint file; returns how many
|
||||
were recovered as detached sessions."""
|
||||
if not CHECKPOINT_PATH.exists():
|
||||
return 0
|
||||
try:
|
||||
entries = json.loads(CHECKPOINT_PATH.read_text(encoding="utf-8"))
|
||||
except Exception:
|
||||
return 0
|
||||
recovered = 0
|
||||
unresolved_scope_entries: List[Dict[str, Any]] = []
|
||||
for entry in entries:
|
||||
pid, pid_scope = entry.get("pid"), entry.get("pid_scope", "host")
|
||||
if not pid:
|
||||
continue
|
||||
if pid_scope != "host": # in-sandbox PIDs mean nothing once the env handle is gone
|
||||
logger.info(
|
||||
"Skipping recovery for non-host process: %s (pid=%s, scope=%s)",
|
||||
entry.get("command", "unknown")[:60], pid, pid_scope)
|
||||
continue
|
||||
# Alive AND the same process: across a restart the kernel may have
|
||||
# recycled the PID onto a stranger, and adopting it would let a later
|
||||
# kill tree-kill e.g. a browser.
|
||||
if not self._host_pid_is_ours(pid, entry.get("host_start_time")):
|
||||
if self._is_host_pid_alive(pid):
|
||||
logger.info(
|
||||
"Not recovering session %s: pid %d is alive but its "
|
||||
"start time no longer matches — PID was recycled onto "
|
||||
"an unrelated process; refusing to adopt it.",
|
||||
entry.get("session_id", "?"), pid)
|
||||
systemd_unit = entry.get("systemd_unit", "")
|
||||
if systemd_unit and not _stop_systemd_unit(systemd_unit):
|
||||
logger.warning(
|
||||
"Could not reap persisted scope %s for dead wrapper pid %s; "
|
||||
"retaining checkpoint entry for the next startup",
|
||||
systemd_unit, pid)
|
||||
unresolved_scope_entries.append(entry)
|
||||
continue
|
||||
fields = {f: entry.get(f, _CHECKPOINT_DEFAULTS[f]) for f in _CHECKPOINT_FIELDS}
|
||||
fields.update(
|
||||
command=entry.get("command", "unknown"),
|
||||
owner_task_id=entry.get("owner_task_id", "") or entry.get("task_id", ""),
|
||||
started_at=entry.get("started_at", time.time()))
|
||||
# detached: can't read output, but can report status + kill
|
||||
session = ProcessSession(id=entry["session_id"], detached=True, **fields)
|
||||
with self._lock:
|
||||
self._running[session.id] = session
|
||||
recovered += 1
|
||||
logger.info("Recovered detached process: %s (pid=%d)", session.command[:60], pid)
|
||||
# Re-enqueue watcher so gateway can resume notifications
|
||||
if session.watcher_interval > 0:
|
||||
self.pending_watchers.append({
|
||||
"session_id": session.id,
|
||||
"check_interval": session.watcher_interval,
|
||||
"session_key": session.session_key,
|
||||
**{key: getattr(session, f"watcher_{key}") for key in _WATCHER_ROUTE_KEYS},
|
||||
"notify_on_complete": session.notify_on_complete,
|
||||
"parent_session_id": session.parent_session_id,
|
||||
})
|
||||
self._write_checkpoint(extra_entries=unresolved_scope_entries)
|
||||
return recovered
|
||||
|
||||
|
||||
process_registry = ProcessRegistry()
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
"""Running-process checkpoint persistence and PID-safe recovery."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from agent.redact import redact_sensitive_text
|
||||
|
||||
logger = logging.getLogger("tools.process_registry")
|
||||
|
||||
|
||||
class ProcessCheckpointMixin:
|
||||
# ----- Checkpoint (crash recovery) -----
|
||||
|
||||
def _write_checkpoint(self, extra_entries: Optional[List[Dict[str, Any]]] = None):
|
||||
"""Write running process metadata to the checkpoint file atomically."""
|
||||
from tools.process_registry import CHECKPOINT_PATH, _CHECKPOINT_FIELDS
|
||||
|
||||
try:
|
||||
with self._lock:
|
||||
entries = []
|
||||
for s in self._running.values():
|
||||
if s.exited:
|
||||
continue
|
||||
# Backfill the start time so recovery can detect PID recycling
|
||||
# even for sessions spawned before this field existed.
|
||||
if s.host_start_time is None and s.pid_scope == "host" and s.pid:
|
||||
s.host_start_time = self._safe_host_start_time(s.pid)
|
||||
entry = {"session_id": s.id, **{f: getattr(s, f) for f in _CHECKPOINT_FIELDS}}
|
||||
# Redact inline credentials before persisting (~/.hermes/processes.json).
|
||||
# Recovery uses command only for display (adoption re-validates the
|
||||
# PID, never re-runs it), so masking is lossless.
|
||||
# See #77484.
|
||||
entry["command"] = redact_sensitive_text(s.command, code_file=True)
|
||||
entry["owner_task_id"] = s.owner_task_id or s.task_id
|
||||
entries.append(entry)
|
||||
if extra_entries:
|
||||
tracked_ids = {item.get("session_id") for item in entries}
|
||||
entries.extend(item for item in extra_entries if item.get("session_id") not in tracked_ids)
|
||||
from utils import atomic_json_write
|
||||
atomic_json_write(CHECKPOINT_PATH, entries)
|
||||
except Exception as e:
|
||||
logger.debug("Failed to write checkpoint file: %s", e, exc_info=True)
|
||||
|
||||
def recover_from_checkpoint(self) -> int:
|
||||
"""On gateway startup, probe PIDs from the checkpoint file; returns how many
|
||||
were recovered as detached sessions."""
|
||||
from tools.process_registry import (
|
||||
CHECKPOINT_PATH, ProcessSession, _CHECKPOINT_FIELDS,
|
||||
_CHECKPOINT_DEFAULTS, _WATCHER_ROUTE_KEYS, _stop_systemd_unit,
|
||||
)
|
||||
|
||||
if not CHECKPOINT_PATH.exists():
|
||||
return 0
|
||||
try:
|
||||
entries = json.loads(CHECKPOINT_PATH.read_text(encoding="utf-8"))
|
||||
except Exception:
|
||||
return 0
|
||||
recovered = 0
|
||||
unresolved_scope_entries: List[Dict[str, Any]] = []
|
||||
for entry in entries:
|
||||
pid, pid_scope = entry.get("pid"), entry.get("pid_scope", "host")
|
||||
if not pid:
|
||||
continue
|
||||
if pid_scope != "host": # in-sandbox PIDs mean nothing once the env handle is gone
|
||||
logger.info(
|
||||
"Skipping recovery for non-host process: %s (pid=%s, scope=%s)",
|
||||
entry.get("command", "unknown")[:60], pid, pid_scope)
|
||||
continue
|
||||
# Alive AND the same process: across a restart the kernel may have
|
||||
# recycled the PID onto a stranger, and adopting it would let a later
|
||||
# kill tree-kill e.g. a browser.
|
||||
if not self._host_pid_is_ours(pid, entry.get("host_start_time")):
|
||||
if self._is_host_pid_alive(pid):
|
||||
logger.info(
|
||||
"Not recovering session %s: pid %d is alive but its "
|
||||
"start time no longer matches — PID was recycled onto "
|
||||
"an unrelated process; refusing to adopt it.",
|
||||
entry.get("session_id", "?"), pid)
|
||||
systemd_unit = entry.get("systemd_unit", "")
|
||||
if systemd_unit and not _stop_systemd_unit(systemd_unit):
|
||||
logger.warning(
|
||||
"Could not reap persisted scope %s for dead wrapper pid %s; "
|
||||
"retaining checkpoint entry for the next startup",
|
||||
systemd_unit, pid)
|
||||
unresolved_scope_entries.append(entry)
|
||||
continue
|
||||
fields = {f: entry.get(f, _CHECKPOINT_DEFAULTS[f]) for f in _CHECKPOINT_FIELDS}
|
||||
fields.update(
|
||||
command=entry.get("command", "unknown"),
|
||||
owner_task_id=entry.get("owner_task_id", "") or entry.get("task_id", ""),
|
||||
started_at=entry.get("started_at", time.time()))
|
||||
# detached: can't read output, but can report status + kill
|
||||
session = ProcessSession(id=entry["session_id"], detached=True, **fields)
|
||||
with self._lock:
|
||||
self._running[session.id] = session
|
||||
recovered += 1
|
||||
logger.info("Recovered detached process: %s (pid=%d)", session.command[:60], pid)
|
||||
# Re-enqueue watcher so gateway can resume notifications
|
||||
if session.watcher_interval > 0:
|
||||
self.pending_watchers.append({
|
||||
"session_id": session.id,
|
||||
"check_interval": session.watcher_interval,
|
||||
"session_key": session.session_key,
|
||||
**{key: getattr(session, f"watcher_{key}") for key in _WATCHER_ROUTE_KEYS},
|
||||
"notify_on_complete": session.notify_on_complete,
|
||||
"parent_session_id": session.parent_session_id,
|
||||
})
|
||||
self._write_checkpoint(extra_entries=unresolved_scope_entries)
|
||||
return recovered
|
||||
Reference in New Issue
Block a user