Files
hermes-agent/tools/code_kernel_remote.py

397 lines
16 KiB
Python

"""Session-persistent kernels for REMOTE terminal backends (docker/ssh/modal).
Remote backends offer one primitive — ``env.execute(cmd)``, run-to-completion
— so the three things the local kernel gets from owning a child are rebuilt:
1. **A process outliving one env.execute():** the runner starts detached
(``nohup ... &``) with its PID recorded; each cell first probes ``kill -0``.
2. **A conversation channel:** a file-based CELL protocol in the kernel dir
(``cell_req_NNNNNN.json`` / ``cell_res_NNNNNN.json``), sibling to the
file-based TOOL-RPC protocol (req_/res_) reused unchanged — the host-side
``_rpc_poll_loop`` starts per cell with the calling thread's context, which
is what gives per-cell tool authority.
3. **Death detection:** a failed liveness probe (transport drop, container
restart, OOM-killed runner) reads as *kernel died: state lost*; the next
call respawns and says so — never a hung poll, every wait is bounded by
the cell timeout.
Same invariants as local: owner = approval session key with the ``::child::``
qualifier (one resolver in tools.code_kernel, cannot drift), same generated
tool stubs, same output post-processing in the caller. ``reset=true`` kills
and respawns. Spawn failure fails OPEN to the per-call path with a note.
"""
from __future__ import annotations
import atexit
import json
import logging
import shlex
import threading
import time
import uuid
from dataclasses import dataclass, field
from typing import Any, Dict, Optional, Tuple
from tools.code_kernel import RUNNER_CELL_SOURCE, KernelRegistry
logger = logging.getLogger(__name__)
# How often the host polls the remote for a cell result file. Each poll is
# one env.execute round-trip (typically 0.1-0.4s on ssh/docker), so this is
# a floor, not a rate.
_CELL_POLL_INTERVAL = 0.5
# The remote runner: a tiny forever-loop that polls for cell request files,
# execs them in one persistent namespace, and writes response files. It is
# deliberately transport-agnostic (pure files) and stdlib-only. Cells and
# tool-RPC share the kernel dir but use distinct prefixes.
REMOTE_KERNEL_RUNNER_SOURCE = '''\
"""Auto-generated Hermes REMOTE session-kernel runner (file cell protocol)."""
import contextlib
import io
import json
import os
import sys
import time
import traceback
KDIR = os.environ["HERMES_KERNEL_DIR"]
CELLS = os.path.join(KDIR, "cells")
_CAPTURE_LIMIT = {capture_limit}
IDLE_EXIT_SECONDS = {idle_exit}
{cell_source}
def main():
execution_count = 0
last_activity = time.time()
while True:
pending = sorted(
f for f in os.listdir(CELLS)
if f.startswith("cell_req_") and f.endswith(".json")
)
if not pending:
if time.time() - last_activity > IDLE_EXIT_SECONDS:
return # self-reap: nobody is talking to us anymore
time.sleep(0.2)
continue
for name in pending:
req_path = os.path.join(CELLS, name)
try:
with open(req_path, "r", encoding="utf-8") as f:
request = json.load(f)
except Exception:
# Partially-written request (ship in progress): retry next tick.
continue
os.remove(req_path)
last_activity = time.time()
execution_count += 1
payload, _ = run_cell(request, execution_count)
res_name = name.replace("cell_req_", "cell_res_")
tmp = os.path.join(CELLS, res_name + ".tmp")
with open(tmp, "w", encoding="utf-8") as f:
json.dump(payload, f, ensure_ascii=False)
os.replace(tmp, os.path.join(CELLS, res_name))
if payload["status"] == "exit":
return
if __name__ == "__main__":
main()
'''
@dataclass
class RemoteKernel:
"""Host-side record of one detached remote kernel process."""
env: Any
env_type: str
kernel_dir: str
pid: str
rpc_token: str
owner: str
created: float = field(default_factory=time.monotonic)
last_used: float = field(default_factory=time.monotonic)
execution_count: int = 0
cell_seq: int = 0
def sh(self, cmd: str, timeout: int = 15) -> str:
return _sh(self.env, cmd, timeout)
def _sh(env, cmd: str, timeout: int = 15) -> str:
"""Run *cmd* on the remote from ``/`` and return its output text."""
result = env.execute(cmd, cwd="/", timeout=timeout)
return (result.get("output", "") if isinstance(result, dict) else "") or ""
def _kernel_key(owner: str, env_type: str, task_env_id: str) -> Tuple:
return (owner, "remote", env_type, task_env_id)
def _is_alive(kernel: RemoteKernel) -> bool:
"""Bounded liveness probe: kill -0 through the transport.
Any transport failure counts as dead — the caller respawns. A dropped ssh
connection and a dead runner are indistinguishable from here, and both
have the same correct answer.
"""
try:
return "ALIVE" in kernel.sh(f"kill -0 {shlex.quote(kernel.pid)} 2>/dev/null && echo ALIVE")
except Exception:
return False
def _kill(kernel: RemoteKernel) -> None:
"""Best-effort kill of the runner and its subprocesses, then rm -rf."""
q_pid = shlex.quote(kernel.pid)
steps = (
# Kill the runner's children if the shell gave it a group, then the PID itself.
(f"pkill -TERM -P {q_pid} 2>/dev/null; kill {q_pid} 2>/dev/null; true",
"remote kernel kill failed (transport?)"),
(f"rm -rf {shlex.quote(kernel.kernel_dir)}", "remote kernel dir cleanup failed"),
)
for cmd, failure in steps:
try:
kernel.sh(cmd)
except Exception:
logger.debug(failure, exc_info=True)
# Registry + lock shared-shape with code_kernel; teardown runs outside the lock.
_REGISTRY = KernelRegistry(lambda kernel: _kill(kernel))
_REMOTE_KERNELS: Dict[Tuple, RemoteKernel] = _REGISTRY.kernels
_REMOTE_KERNELS_LOCK = _REGISTRY.lock
def shutdown_all_remote_kernels() -> None:
_REGISTRY.shutdown()
def shutdown_remote_kernels_for_owner(owner: str) -> None:
"""Session-boundary disposal — wired to the same clear_session hook as
local kernels, so /new and session close reap both kinds."""
if owner:
_REGISTRY.shutdown(owner)
atexit.register(shutdown_all_remote_kernels)
def _spawn_remote_kernel(env, env_type: str, owner: str, task_env_id: str,
sandbox_tools: frozenset, *, idle_exit: int) -> Optional[RemoteKernel]:
"""Start a detached kernel runner on the remote. None on failure."""
from tools.code_execution_tool import (
MAX_STDOUT_BYTES, _ship_file_to_remote, _env_temp_dir, generate_hermes_tools_module,
)
import secrets as _secrets
kernel_dir = f"{_env_temp_dir(env)}/hermes_rkernel_{uuid.uuid4().hex[:12]}"
q_dir = shlex.quote(kernel_dir)
def sh(cmd: str, timeout: int = 15) -> str:
return _sh(env, cmd, timeout)
def start() -> Optional[RemoteKernel]:
sh(f"mkdir -p {q_dir}/cells {q_dir}/rpc")
rpc_token = _secrets.token_urlsafe(32)
runner_src = REMOTE_KERNEL_RUNNER_SOURCE.format(
cell_source=RUNNER_CELL_SOURCE, capture_limit=MAX_STDOUT_BYTES, idle_exit=idle_exit)
_ship_file_to_remote(env, f"{kernel_dir}/kernel_runner.py", runner_src)
_ship_file_to_remote(env, f"{kernel_dir}/hermes_tools.py",
generate_hermes_tools_module(list(sandbox_tools), transport="file"))
env_prefix = (
f"HERMES_KERNEL_DIR={q_dir} "
f"HERMES_RPC_DIR={shlex.quote(kernel_dir + '/rpc')} "
f"HERMES_RPC_TOKEN={shlex.quote(rpc_token)} "
f"PYTHONDONTWRITEBYTECODE=1 PYTHONPATH={q_dir}"
)
started = sh(f"cd {q_dir} && nohup env {env_prefix} python3 kernel_runner.py "
f"> {q_dir}/runner.log 2>&1 & echo PID:$!", timeout=20)
pid = next((line.strip()[4:].strip() for line in started.splitlines()
if line.strip().startswith("PID:")), "")
if not pid.isdigit():
logger.warning("remote kernel spawn returned no PID: %r", started)
return None
kernel = RemoteKernel(env=env, env_type=env_type, kernel_dir=kernel_dir,
pid=pid, rpc_token=rpc_token, owner=owner)
if not _is_alive(kernel):
# Died instantly (missing python3 was pre-checked by the caller,
# so this is unexpected) — surface the runner log.
try:
logger.warning("remote kernel died at spawn: %s",
sh(f"cat {q_dir}/runner.log", timeout=10)[:500])
except Exception:
pass
return None
return kernel
kernel = None
try:
kernel = start()
except Exception:
logger.warning("remote kernel spawn failed", exc_info=True)
if kernel is None:
try:
sh(f"rm -rf {q_dir}")
except Exception:
pass
return kernel
def _acquire_remote_kernel(env, env_type: str, owner: str, task_env_id: str,
sandbox_tools: frozenset, *, reset: bool,
idle_exit: int) -> Tuple[Optional[RemoteKernel], bool, bool, bool]:
"""Find/respawn the owner's kernel: (kernel|None, reused, state_reset, state_lost)."""
key = _kernel_key(owner, env_type, task_env_id)
state_lost = state_reset = False
with _REMOTE_KERNELS_LOCK:
kernel = _REMOTE_KERNELS.get(key)
if kernel is not None and reset:
_REGISTRY.discard(key, kernel)
kernel, state_reset = None, True
if kernel is not None and not _is_alive(kernel):
# Transport drop, container restart, self-reaped on idle, OOM — all
# the same answer: report the loss, respawn fresh (_kill is then only
# best-effort dir cleanup; the process is already gone).
_REGISTRY.discard(key, kernel)
kernel, state_lost = None, True
reused = kernel is not None
if kernel is None:
kernel = _spawn_remote_kernel(env, env_type, owner, task_env_id, sandbox_tools,
idle_exit=idle_exit)
if kernel is not None:
with _REMOTE_KERNELS_LOCK:
_REMOTE_KERNELS[key] = kernel
return kernel, reused, state_reset, state_lost
def _run_remote_cell(kernel: RemoteKernel, code: str, timeout: int) -> Tuple[str, Dict[str, Any]]:
"""Ship one cell request and poll for its result: (cell status, payload)."""
from tools.code_execution_tool import _ship_file_to_remote
kernel.cell_seq += 1
seq = f"{kernel.cell_seq:06d}"
q_cells = shlex.quote(f"{kernel.kernel_dir}/cells")
res_name = f"cell_res_{seq}.json"
request = json.dumps({"id": seq, "code": code}, ensure_ascii=False)
_ship_file_to_remote(kernel.env, f"{kernel.kernel_dir}/cells/cell_req_{seq}.json.tmp", request)
kernel.sh(f"mv {q_cells}/cell_req_{seq}.json.tmp {q_cells}/cell_req_{seq}.json", timeout=10)
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
try:
body = kernel.sh(f"cat {q_cells}/{shlex.quote(res_name)} 2>/dev/null", timeout=20).strip()
except Exception:
# One flaky round-trip is not kernel death; liveness decides.
time.sleep(_CELL_POLL_INTERVAL)
continue
if body:
try:
payload = json.loads(body)
status = payload.get("status", "error")
except ValueError:
payload, status = {}, "protocol-error"
kernel.sh(f"rm -f {q_cells}/{shlex.quote(res_name)}", timeout=10)
return status, payload
time.sleep(_CELL_POLL_INTERVAL)
return "timeout", {}
def execute_in_remote_kernel(
code: str, *, env, env_type: str, task_env_id: str, sandbox_tools: frozenset,
timeout: int, max_tool_calls: int, reset: bool, idle_exit: int = 1800,
) -> Optional[Dict[str, Any]]:
"""Run one cell in the owner's remote kernel.
Returns the raw cell result dict (caller does output post-processing),
or ``None`` when no kernel could be spawned — the caller falls open to
the per-call path. ``state_lost`` / ``state_reset`` / ``reused`` ride in
the ``kernel`` sub-dict, matching the local kernel's result shape.
"""
from tools.code_kernel import _resolve_owner
from tools.code_execution_tool import _rpc_poll_loop
from tools.thread_context import propagate_context_to_thread
owner = _resolve_owner(task_env_id)
kernel, reused, state_reset, state_lost = _acquire_remote_kernel(
env, env_type, owner, task_env_id, sandbox_tools, reset=reset, idle_exit=idle_exit)
if kernel is None:
return None # fail open to per-call
key = _kernel_key(owner, env_type, task_env_id)
kernel.last_used = time.monotonic()
# Clean stale tool-RPC requests from a previous cell before arming this
# cell's poll loop, so a background thread the last cell leaked cannot
# smuggle a call into this cell's authority window.
q_rpc = shlex.quote(kernel.kernel_dir + '/rpc')
try:
kernel.sh(f"rm -f {q_rpc}/req_* {q_rpc}/res_*", timeout=10)
except Exception:
pass
tool_call_log: list = []
tool_call_counter, stop_event = [0], threading.Event()
# Per-cell RPC thread carrying THIS call's approval/session context —
# the remote analogue of CellAuthority: authority lives exactly as long
# as the cell's poll loop.
rpc_thread = threading.Thread(
target=propagate_context_to_thread(_rpc_poll_loop),
args=(env, f"{kernel.kernel_dir}/rpc", task_env_id, tool_call_log, tool_call_counter,
max_tool_calls, sandbox_tools, stop_event, kernel.rpc_token),
daemon=True,
)
rpc_thread.start()
cell_status, cell_payload = "no-result", {}
try:
cell_status, cell_payload = _run_remote_cell(kernel, code, timeout)
finally:
stop_event.set()
rpc_thread.join(timeout=5)
kernel_info: Dict[str, Any] = {"reused": reused, "remote": True}
result: Dict[str, Any] = {
"status": "error", "stdout": cell_payload.get("stdout", ""),
"stderr": cell_payload.get("stderr", ""), "traceback": cell_payload.get("traceback", ""),
"tool_calls_made": tool_call_counter[0], "kernel": kernel_info,
}
if cell_status in ("timeout", "protocol-error", "no-result"):
# No safe way to interrupt one cell in place (same contract as
# local): kill the kernel, report the loss, respawn next call.
_REGISTRY.discard(key, kernel)
if cell_status == "timeout":
result["status"] = "timeout"
note = ("Cell timed out; the remote session kernel was killed and "
"its state was lost. The next call starts a fresh kernel.")
else:
note = "Remote kernel protocol failure; kernel killed, state lost."
kernel_info.update(ended=True, state_lost=True, note=note)
return result
if cell_status == "exit":
_REGISTRY.discard(key, kernel)
kernel_info["ended"] = True
kernel.execution_count = int(cell_payload.get("execution_count", 0) or 0)
kernel_info["execution_count"] = kernel.execution_count
if cell_status in ("ok", "exit"):
result["status"] = "success"
result["stdout_clipped"] = bool(cell_payload.get("stdout_clipped"))
result["stderr_clipped"] = bool(cell_payload.get("stderr_clipped"))
if state_reset:
kernel_info["state_reset"] = True
if state_lost:
kernel_info.update(state_lost=True, note=(
"The previous remote kernel was gone (transport drop, container "
"restart, or idle self-exit); state from earlier calls was lost "
"and a fresh kernel was started."))
if cell_status == "error" and result["traceback"]:
result["error"] = result["traceback"].strip().splitlines()[-1]
return result