refactor(tools): simplify remote env backends (modal/managed_modal/daytona/ssh/vercel) — dead code, shared helpers, defensive collapse

This commit is contained in:
Teknium
2026-09-02 21:59:50 -07:00
parent 113f04616b
commit 876790dcc0
6 changed files with 267 additions and 483 deletions

View File

@@ -4,6 +4,7 @@ Runs commands in Daytona cloud sandboxes via the Python SDK. Persistent mode sto
the sandbox on cleanup and resumes it next time, preserving the filesystem.
"""
import contextlib
import logging
import math
import os
@@ -32,9 +33,7 @@ class DaytonaEnvironment(BaseEnvironment):
def __init__(self, image: str, cwd: str = "/home/daytona", timeout: int = 60, cpu: int = 1,
memory: int = 5120, disk: int = 10240, persistent_filesystem: bool = True,
task_id: str = "default"):
requested_cwd = cwd
super().__init__(cwd=cwd, timeout=timeout)
ensure_lazy_dep("terminal.daytona")
from daytona import Daytona, CreateSandboxFromImageParams, DaytonaError, Resources, SandboxState
@@ -52,7 +51,6 @@ class DaytonaEnvironment(BaseEnvironment):
"Capping to 10GB.", disk_gib)
disk_gib = 10
resources = Resources(cpu=cpu, memory=memory_gib, disk=disk_gib)
labels = {"hermes_task_id": task_id}
sandbox_name = f"hermes-{task_id}"
@@ -66,7 +64,6 @@ class DaytonaEnvironment(BaseEnvironment):
except Exception as e:
logger.warning("Daytona: failed to resume sandbox for task %s: %s", task_id, e)
self._sandbox = None
if self._sandbox is None:
try:
# SDK list() is a cursor-paginated iterator (offset pagination is gone).
@@ -78,38 +75,30 @@ class DaytonaEnvironment(BaseEnvironment):
except Exception as e:
logger.debug("Daytona: no legacy sandbox found for task %s: %s", task_id, e)
self._sandbox = None
if self._sandbox is None:
self._sandbox = self._daytona.create(CreateSandboxFromImageParams(
image=image, name=sandbox_name, labels=labels, auto_stop_interval=0, resources=resources,
))
image=image, name=sandbox_name, labels=labels, auto_stop_interval=0, resources=resources))
logger.info("Daytona: created sandbox %s for task %s", self._sandbox.id, task_id)
self._remote_home = "/root"
try:
with contextlib.suppress(Exception):
home = self._sandbox.process.exec("echo $HOME").result.strip()
if home:
self._remote_home = home
if requested_cwd in {"~", "/home/daytona"}:
if cwd in {"~", "/home/daytona"}:
self.cwd = home
except Exception:
pass
logger.info("Daytona: resolved home to %s, cwd to %s", self._remote_home, self.cwd)
self._sync_manager = FileSyncManager(
get_files_fn=lambda: iter_sync_files(f"{self._remote_home}/.hermes"),
upload_fn=self._daytona_upload,
delete_fn=self._daytona_delete,
bulk_upload_fn=self._daytona_bulk_upload,
bulk_download_fn=self._daytona_bulk_download,
)
upload_fn=self._daytona_upload, delete_fn=self._daytona_delete,
bulk_upload_fn=self._daytona_bulk_upload, bulk_download_fn=self._daytona_bulk_download)
self._sync_manager.sync(force=True)
self.init_session()
def _daytona_upload(self, host_path: str, remote_path: str) -> None:
"""Upload a single file via Daytona SDK."""
parent = str(Path(remote_path).parent)
self._sandbox.process.exec(quoted_mkdir_command([parent]))
self._sandbox.process.exec(quoted_mkdir_command([str(Path(remote_path).parent)]))
self._sandbox.fs.upload_file(host_path, remote_path)
def _daytona_bulk_upload(self, files: list[tuple[str, str]]) -> None:
@@ -122,8 +111,7 @@ class DaytonaEnvironment(BaseEnvironment):
if parents:
self._sandbox.process.exec(quoted_mkdir_command(parents))
self._sandbox.fs.upload_files(
[FileUpload(source=host_path, destination=remote_path) for host_path, remote_path in files]
)
[FileUpload(source=host_path, destination=remote_path) for host_path, remote_path in files])
def _daytona_bulk_download(self, dest: Path) -> None:
"""Download remote .hermes/ as a tar archive."""
@@ -132,17 +120,13 @@ class DaytonaEnvironment(BaseEnvironment):
remote_tar = f"/tmp/.hermes_sync.{os.getpid()}.tar"
self._sandbox.process.exec(f"tar cf {shlex.quote(remote_tar)} -C / {shlex.quote(rel_base)}")
self._sandbox.fs.download_file(remote_tar, str(dest))
try:
with contextlib.suppress(Exception): # best-effort cleanup
self._sandbox.process.exec(f"rm -f {shlex.quote(remote_tar)}")
except Exception:
pass # best-effort cleanup
def _daytona_delete(self, remote_paths: list[str]) -> None:
"""Batch-delete remote files via SDK exec."""
self._sandbox.process.exec(quoted_rm_command(remote_paths))
# -- Sandbox lifecycle ----------------------------------------------
def _ensure_sandbox_ready(self) -> None:
"""Restart sandbox if it was stopped (e.g., by a previous interrupt)."""
self._sandbox.refresh_data()
@@ -163,11 +147,8 @@ class DaytonaEnvironment(BaseEnvironment):
lock = self._lock
def cancel():
with lock:
try:
sandbox.stop()
except Exception:
pass
with lock, contextlib.suppress(Exception):
sandbox.stop()
shell_cmd = f"bash {'-l ' if login else ''}-c {shlex.quote(cmd_string)}"
@@ -181,7 +162,6 @@ class DaytonaEnvironment(BaseEnvironment):
with self._lock:
if self._sandbox is None:
return
# sync_back runs inside the lock and after the None guard so an
# already-cleaned-up env can't trigger a 3-attempt retry storm on a nil sandbox.
if self._sync_manager:
@@ -190,7 +170,6 @@ class DaytonaEnvironment(BaseEnvironment):
self._sync_manager.sync_back()
except Exception as e:
logger.warning("Daytona: sync_back failed: %s", e)
try:
if self._persistent:
self._sandbox.stop()

View File

@@ -7,7 +7,6 @@ import logging
import os
import requests
import uuid
from dataclasses import dataclass
from typing import Any, Dict, Optional
from tools.environments.modal_utils import BaseModalExecutionEnvironment, ModalExecStart, PreparedModalExec
@@ -26,13 +25,11 @@ def _request_timeout_env(name: str, default: float) -> float:
return default
@dataclass(frozen=True)
class _ManagedModalExecHandle:
exec_id: str
class ManagedModalEnvironment(BaseModalExecutionEnvironment):
"""Gateway-owned Modal sandbox with Hermes-compatible execute/cleanup."""
"""Gateway-owned Modal sandbox with Hermes-compatible execute/cleanup.
The exec handle passed between ``_start_modal_exec`` / ``_poll_modal_exec`` /
``_cancel_modal_exec`` is the gateway exec id string."""
_CONNECT_TIMEOUT_SECONDS = _request_timeout_env("TERMINAL_MANAGED_MODAL_CONNECT_TIMEOUT_SECONDS", 1.0)
_POLL_READ_TIMEOUT_SECONDS = _request_timeout_env("TERMINAL_MANAGED_MODAL_POLL_READ_TIMEOUT_SECONDS", 5.0)
@@ -45,7 +42,17 @@ class ManagedModalEnvironment(BaseModalExecutionEnvironment):
modal_sandbox_kwargs: Optional[Dict[str, Any]] = None,
persistent_filesystem: bool = True, task_id: str = "default"):
super().__init__(cwd=cwd, timeout=timeout)
self._guard_unsupported_credential_passthrough()
# Managed Modal does not sync or mount host credential files.
try:
from tools.credential_files import get_credential_file_mounts
except Exception:
get_credential_file_mounts = None
if get_credential_file_mounts is not None and get_credential_file_mounts():
raise ValueError(
"Managed Modal does not support host credential-file passthrough. "
"Use TERMINAL_MODAL_MODE=direct when skills or config require "
"credential files inside the sandbox."
)
gateway = resolve_managed_tool_gateway("modal")
if gateway is None:
raise ValueError("Managed Modal requires a configured tool gateway and Nous user token")
@@ -84,12 +91,12 @@ class ManagedModalEnvironment(BaseModalExecutionEnvironment):
if body.get("execId") != exec_id:
return ModalExecStart(immediate_result=self._error_result(
"Managed Modal exec start did not return the expected exec id"))
return ModalExecStart(handle=_ManagedModalExecHandle(exec_id=exec_id))
return ModalExecStart(handle=exec_id)
def _poll_modal_exec(self, handle: _ManagedModalExecHandle) -> dict | None:
def _poll_modal_exec(self, handle: str) -> dict | None:
try:
status_response = self._request(
"GET", f"/v1/sandboxes/{self._sandbox_id}/execs/{handle.exec_id}",
"GET", f"/v1/sandboxes/{self._sandbox_id}/execs/{handle}",
timeout=(self._CONNECT_TIMEOUT_SECONDS, self._POLL_READ_TIMEOUT_SECONDS))
except Exception as exc:
return self._error_result(f"Managed Modal exec poll failed: {exc}")
@@ -99,8 +106,12 @@ class ManagedModalEnvironment(BaseModalExecutionEnvironment):
return self._error_result(self._format_error("Managed Modal exec poll failed", status_response))
return self._result_from_body(status_response.json())
def _cancel_modal_exec(self, handle: _ManagedModalExecHandle) -> None:
self._cancel_exec(handle.exec_id)
def _cancel_modal_exec(self, handle: str) -> None:
try:
self._request("POST", f"/v1/sandboxes/{self._sandbox_id}/execs/{handle}/cancel",
timeout=(self._CONNECT_TIMEOUT_SECONDS, self._CANCEL_READ_TIMEOUT_SECONDS))
except Exception as exc:
logger.warning("Managed Modal exec cancel failed: %s", exc)
def _timeout_result_for_modal(self, timeout: int) -> dict:
return self._result(f"Managed Modal exec timed out after {timeout}s", 124)
@@ -137,39 +148,16 @@ class ManagedModalEnvironment(BaseModalExecutionEnvironment):
raise RuntimeError("Managed Modal create did not return a sandbox id")
return sandbox_id
def _guard_unsupported_credential_passthrough(self) -> None:
"""Managed Modal does not sync or mount host credential files."""
try:
from tools.credential_files import get_credential_file_mounts
except Exception:
return
if get_credential_file_mounts():
raise ValueError(
"Managed Modal does not support host credential-file passthrough. "
"Use TERMINAL_MODAL_MODE=direct when skills or config require "
"credential files inside the sandbox."
)
def _request(self, method: str, path: str, *, json: Dict[str, Any] | None = None, timeout: int = 30,
extra_headers: Dict[str, str] | None = None) -> requests.Response:
headers = {"Authorization": f"Bearer {self._nous_user_token}", "Content-Type": "application/json"}
if extra_headers:
headers.update(extra_headers)
headers = {"Authorization": f"Bearer {self._nous_user_token}", "Content-Type": "application/json",
**(extra_headers or {})}
return requests.request(method, f"{self._gateway_origin}{path}", headers=headers, json=json, timeout=timeout)
def _cancel_exec(self, exec_id: str) -> None:
try:
self._request("POST", f"/v1/sandboxes/{self._sandbox_id}/execs/{exec_id}/cancel",
timeout=(self._CONNECT_TIMEOUT_SECONDS, self._CANCEL_READ_TIMEOUT_SECONDS))
except Exception as exc:
logger.warning("Managed Modal exec cancel failed: %s", exc)
@staticmethod
def _coerce_number(value: Any, default: float) -> float:
try:
if value is None:
return default
return float(value)
return default if value is None else float(value)
except (TypeError, ValueError):
return default

View File

@@ -22,7 +22,6 @@ from tools.environments.remote_common import bash_argv, ensure_lazy_dep
logger = logging.getLogger(__name__)
_SNAPSHOT_STORE = get_hermes_home() / "modal_snapshots.json"
_DIRECT_SNAPSHOT_NAMESPACE = "direct"
def _load_snapshots() -> dict:
@@ -33,14 +32,10 @@ def _save_snapshots(data: dict) -> None:
_save_json_store(_SNAPSHOT_STORE, data)
def _direct_snapshot_key(task_id: str) -> str:
return f"{_DIRECT_SNAPSHOT_NAMESPACE}:{task_id}"
def _get_snapshot_restore_candidate(task_id: str) -> tuple[str | None, bool]:
"""Return (snapshot_id, from_legacy_key); the namespaced key wins over the legacy bare task id."""
snapshots = _load_snapshots()
for key, legacy in ((_direct_snapshot_key(task_id), False), (task_id, True)):
for key, legacy in ((f"direct:{task_id}", False), (task_id, True)):
snapshot_id = snapshots.get(key)
if isinstance(snapshot_id, str) and snapshot_id:
return snapshot_id, legacy
@@ -49,61 +44,54 @@ def _get_snapshot_restore_candidate(task_id: str) -> tuple[str | None, bool]:
def _store_direct_snapshot(task_id: str, snapshot_id: str) -> None:
snapshots = _load_snapshots()
snapshots[_direct_snapshot_key(task_id)] = snapshot_id
snapshots[f"direct:{task_id}"] = snapshot_id
snapshots.pop(task_id, None)
_save_snapshots(snapshots)
def _delete_direct_snapshot(task_id: str, snapshot_id: str | None = None) -> None:
snapshots = _load_snapshots()
updated = False
for key in (_direct_snapshot_key(task_id), task_id):
value = snapshots.get(key)
if value is not None and (snapshot_id is None or value == snapshot_id):
stale = [k for k in (f"direct:{task_id}", task_id)
if snapshots.get(k) is not None and (snapshot_id is None or snapshots[k] == snapshot_id)]
if stale:
for key in stale:
snapshots.pop(key, None)
updated = True
if updated:
_save_snapshots(snapshots)
def _ensure_modal_sdk() -> None:
"""Lazy-install modal on demand. Idempotent — fast no-op once installed."""
ensure_lazy_dep("terminal.modal")
def _resolve_modal_image(image_spec: Any) -> Any:
"""Convert registry references or snapshot ids into Modal image objects. Registry images
get pip repaired (ensurepip) before Modal's bootstrap; ubuntu/debian also get python3."""
_ensure_modal_sdk()
ensure_lazy_dep("terminal.modal")
import modal as _modal
if not isinstance(image_spec, str):
return image_spec
if image_spec.startswith("im-"):
return _modal.Image.from_id(image_spec)
setup_commands = [
"RUN rm -rf /usr/local/lib/python*/site-packages/pip* 2>/dev/null; "
"python -m ensurepip --upgrade --default-pip 2>/dev/null || true",
]
if any(base in image_spec.lower() for base in ("ubuntu", "debian")):
setup_commands.insert(0,
"RUN apt-get update -qq && apt-get install -y -qq python3 python3-venv > /dev/null 2>&1 || true"
)
"RUN apt-get update -qq && apt-get install -y -qq python3 python3-venv > /dev/null 2>&1 || true")
return _modal.Image.from_registry(image_spec, setup_dockerfile_commands=setup_commands)
async def _stream_stdin(proc, payload: str, chunk_size: int) -> None:
"""Write ``payload`` to ``proc.stdin`` in ``chunk_size`` pieces, draining after each, then EOF."""
offset = 0
while offset < len(payload):
for offset in range(0, len(payload), chunk_size):
proc.stdin.write(payload[offset:offset + chunk_size])
await proc.stdin.drain.aio()
offset += chunk_size
proc.stdin.write_eof()
await proc.stdin.drain.aio()
def _as_text(value) -> str:
return value.decode("utf-8", errors="replace") if isinstance(value, bytes) else value
class _AsyncWorker:
"""Background thread with its own event loop for async-safe Modal calls."""
@@ -147,10 +135,9 @@ class ModalEnvironment(BaseEnvironment):
_stdin_mode = "heredoc"
_snapshot_timeout = 60 # Modal cold starts can be slow
# Modal SDK stdin buffer limit: the command-router path allows 16 MB but the
# legacy server path caps at 2 MB, so chunks stay under 2 MB and each is
# flushed individually via drain().
_STDIN_CHUNK_SIZE = 1 * 1024 * 1024 # 1 MB — safe for both transport paths
# Modal SDK stdin buffer limit: the command-router path allows 16 MB but the legacy server
# path caps at 2 MB, so chunks stay under 2 MB and each is flushed individually via drain().
_STDIN_CHUNK_SIZE = 1 * 1024 * 1024
def __init__(self, image: str, cwd: str = "/root", timeout: int = 60,
modal_sandbox_kwargs: Optional[dict[str, Any]] = None,
@@ -163,25 +150,20 @@ class ModalEnvironment(BaseEnvironment):
self._worker = _AsyncWorker()
self._sync_manager: FileSyncManager | None = None # initialized after sandbox creation
sandbox_kwargs = dict(modal_sandbox_kwargs or {})
restored_snapshot_id = None
restored_from_legacy_key = False
restored_snapshot_id, restored_from_legacy_key = None, False
if self._persistent:
restored_snapshot_id, restored_from_legacy_key = _get_snapshot_restore_candidate(self._task_id)
if restored_snapshot_id:
logger.info("Modal: restoring from snapshot %s", restored_snapshot_id[:20])
_ensure_modal_sdk()
ensure_lazy_dep("terminal.modal")
import modal as _modal
cred_mounts = []
try:
from tools.credential_files import get_credential_file_mounts, iter_skills_files, iter_cache_files
# from_iterable keeps each source lazy so a failure mid-way leaves the earlier mounts in place
for entry in itertools.chain.from_iterable(
fn() for fn in (get_credential_file_mounts, iter_skills_files, iter_cache_files)
):
cred_mounts.append(
_modal.Mount.from_local_file(entry["host_path"], remote_path=entry["container_path"])
)
fn() for fn in (get_credential_file_mounts, iter_skills_files, iter_cache_files)):
cred_mounts.append(_modal.Mount.from_local_file(entry["host_path"], remote_path=entry["container_path"]))
except Exception as e:
logger.debug("Modal: could not load credential file mounts: %s", e)
self._worker.start()
@@ -193,8 +175,7 @@ class ModalEnvironment(BaseEnvironment):
create_kwargs["mounts"] = list(create_kwargs.pop("mounts", [])) + cred_mounts
sandbox = await _modal.Sandbox.create.aio(
"sleep", "infinity", image=image_spec, app=app,
timeout=int(create_kwargs.pop("timeout", 3600)), **create_kwargs,
)
timeout=int(create_kwargs.pop("timeout", 3600)), **create_kwargs)
return app, sandbox
try:
try:
@@ -203,10 +184,8 @@ class ModalEnvironment(BaseEnvironment):
except Exception as exc:
if not restored_snapshot_id:
raise
logger.warning(
"Modal: failed to restore snapshot %s, retrying with base image: %s",
restored_snapshot_id[:20], exc,
)
logger.warning("Modal: failed to restore snapshot %s, retrying with base image: %s",
restored_snapshot_id[:20], exc)
_delete_direct_snapshot(self._task_id, restored_snapshot_id)
self._app, self._sandbox = self._worker.run_coroutine(
_create_sandbox(_resolve_modal_image(image)), timeout=300)
@@ -220,22 +199,28 @@ class ModalEnvironment(BaseEnvironment):
self._sync_manager = FileSyncManager(
get_files_fn=lambda: iter_sync_files("/root/.hermes"),
upload_fn=self._modal_upload, delete_fn=self._modal_delete,
bulk_upload_fn=self._modal_bulk_upload, bulk_download_fn=self._modal_bulk_download,
)
bulk_upload_fn=self._modal_bulk_upload, bulk_download_fn=self._modal_bulk_download)
self._sync_manager.sync(force=True)
self.init_session()
def _exec_with_stdin(self, cmd: str, payload: str, *, timeout: int, fail_label: str | None = None) -> None:
"""Run ``bash -c cmd`` in the sandbox feeding ``payload`` on stdin; when ``fail_label`` is
given a non-zero exit raises with the remote stderr."""
async def _run():
proc = await self._sandbox.exec.aio("bash", "-c", cmd)
await _stream_stdin(proc, payload, self._STDIN_CHUNK_SIZE)
exit_code = await proc.wait.aio()
if fail_label and exit_code != 0:
stderr_text = await proc.stderr.read.aio()
raise RuntimeError(f"Modal {fail_label} failed (exit {exit_code}): {stderr_text}")
self._worker.run_coroutine(_run(), timeout=timeout)
def _modal_upload(self, host_path: str, remote_path: str) -> None:
"""Upload a single file via base64 piped through stdin."""
b64 = base64.b64encode(Path(host_path).read_bytes()).decode("ascii")
container_dir = str(Path(remote_path).parent)
cmd = f"mkdir -p {shlex.quote(container_dir)} && base64 -d > {shlex.quote(remote_path)}"
async def _write():
proc = await self._sandbox.exec.aio("bash", "-c", cmd)
await _stream_stdin(proc, b64, self._STDIN_CHUNK_SIZE)
await proc.wait.aio()
self._worker.run_coroutine(_write(), timeout=30)
self._exec_with_stdin(cmd, b64, timeout=30)
def _modal_bulk_upload(self, files: list[tuple[str, str]]) -> None:
"""Upload many files as one in-memory gzipped tar streamed through stdin
@@ -248,15 +233,7 @@ class ModalEnvironment(BaseEnvironment):
tar.add(host_path, arcname=remote_path.lstrip("/"))
payload = base64.b64encode(buf.getvalue()).decode("ascii")
cmd = f"{quoted_mkdir_command(unique_parent_dirs(files))} && base64 -d | tar xzf - -C /"
async def _bulk():
proc = await self._sandbox.exec.aio("bash", "-c", cmd)
await _stream_stdin(proc, payload, self._STDIN_CHUNK_SIZE)
exit_code = await proc.wait.aio()
if exit_code != 0:
stderr_text = await proc.stderr.read.aio()
raise RuntimeError(f"Modal bulk upload failed (exit {exit_code}): {stderr_text}")
self._worker.run_coroutine(_bulk(), timeout=120)
self._exec_with_stdin(cmd, payload, timeout=120, fail_label="bulk upload")
def _modal_bulk_download(self, dest: Path) -> None:
"""Download remote .hermes/ as a tar archive (sandboxes run as root, so /root/.hermes)."""
@@ -268,9 +245,7 @@ class ModalEnvironment(BaseEnvironment):
raise RuntimeError(f"Modal bulk download failed (exit {exit_code})")
return data
tar_bytes = self._worker.run_coroutine(_download(), timeout=120)
if isinstance(tar_bytes, str):
tar_bytes = tar_bytes.encode()
dest.write_bytes(tar_bytes)
dest.write_bytes(tar_bytes.encode() if isinstance(tar_bytes, str) else tar_bytes)
def _modal_delete(self, remote_paths: list[str]) -> None:
"""Batch-delete remote files via exec."""
@@ -296,13 +271,9 @@ class ModalEnvironment(BaseEnvironment):
def exec_fn() -> tuple[str, int]:
async def _do():
process = await sandbox.exec.aio(*bash_argv(cmd_string, login), timeout=timeout)
stdout = await process.stdout.read.aio()
stderr = await process.stderr.read.aio()
stdout = _as_text(await process.stdout.read.aio())
stderr = _as_text(await process.stderr.read.aio())
exit_code = await process.wait.aio()
if isinstance(stdout, bytes):
stdout = stdout.decode("utf-8", errors="replace")
if isinstance(stderr, bytes):
stderr = stderr.decode("utf-8", errors="replace")
output = stdout
if stderr:
output = f"{stdout}\n{stderr}" if stdout else stderr
@@ -318,19 +289,19 @@ class ModalEnvironment(BaseEnvironment):
logger.info("Modal: syncing files from sandbox...")
self._sync_manager.sync_back()
if self._persistent:
async def _snapshot():
img = await self._sandbox.snapshot_filesystem.aio()
return img.object_id
try:
async def _snapshot():
img = await self._sandbox.snapshot_filesystem.aio()
return img.object_id
snapshot_id = self._worker.run_coroutine(_snapshot(), timeout=60)
except Exception:
snapshot_id = None
if snapshot_id:
try:
snapshot_id = self._worker.run_coroutine(_snapshot(), timeout=60)
except Exception:
snapshot_id = None
if snapshot_id:
_store_direct_snapshot(self._task_id, snapshot_id)
logger.info("Modal: saved filesystem snapshot %s for task %s", snapshot_id[:20], self._task_id)
except Exception as e:
logger.warning("Modal: filesystem snapshot failed: %s", e)
except Exception as e:
logger.warning("Modal: filesystem snapshot failed: %s", e)
try:
self._worker.run_coroutine(self._sandbox.terminate.aio(), timeout=15)
except Exception:

View File

@@ -1,16 +1,14 @@
"""Shared Hermes-side execution flow for Modal transports.
Stops at the Hermes boundary: command preparation, cwd/timeout normalization,
stdin/sudo shell wrapping, common result shape, interrupt/cancel polling.
Direct and managed Modal keep transport, persistence and trust-boundary logic
in their own modules.
sudo shell wrapping, common result shape, interrupt/cancel polling. The managed
transport keeps HTTP, persistence and trust-boundary logic in its own module.
"""
from __future__ import annotations
import shlex
import time
import uuid
from abc import abstractmethod
from dataclasses import dataclass
from typing import Any
@@ -37,19 +35,6 @@ class ModalExecStart:
immediate_result: dict | None = None
def wrap_modal_stdin_heredoc(command: str, stdin_data: str) -> str:
"""Append stdin as a shell heredoc for transports without stdin piping."""
marker = f"HERMES_EOF_{uuid.uuid4().hex[:8]}"
while marker in stdin_data:
marker = f"HERMES_EOF_{uuid.uuid4().hex[:8]}"
return f"{command} << '{marker}'\n{stdin_data}\n{marker}"
def wrap_modal_sudo_pipe(command: str, sudo_stdin: str) -> str:
"""Feed sudo via a shell pipe for transports without direct stdin piping."""
return f"printf '%s\\n' {shlex.quote(sudo_stdin.rstrip())} | {command}"
class BaseModalExecutionEnvironment(BaseEnvironment):
"""Execution flow for the *managed* Modal transport (gateway-owned sandbox).
@@ -66,10 +51,10 @@ class BaseModalExecutionEnvironment(BaseEnvironment):
def execute(self, command: str, cwd: str = "", *, timeout: int | None = None, stdin_data: str | None = None,
rewrite_compound_background: bool = True, bounded_capture: bool = False) -> dict:
# Signature parity with BaseEnvironment.execute only: the transport runs
# commands explicitly (no shell background rewriting) and returns the
# remote result in one payload, so streaming-time bounding does not apply
# (the terminal tool's final truncation still caps it).
# Signature parity with BaseEnvironment.execute only: the transport runs commands
# explicitly (no shell background rewriting) and returns the remote result in one
# payload, so streaming-time bounding does not apply (the terminal tool's final
# truncation still caps it).
del rewrite_compound_background, bounded_capture
self._before_execute()
prepared = self._prepare_modal_exec(command, cwd=cwd, timeout=timeout, stdin_data=stdin_data)
@@ -99,7 +84,6 @@ class BaseModalExecutionEnvironment(BaseEnvironment):
if deadline is not None and time.monotonic() >= deadline:
self._cancel_quietly(start.handle)
return self._timeout_result_for_modal(prepared.timeout)
# Periodic activity touch so the gateway knows we're alive (lazy import:
# tests stub tools.environments.base with only BaseEnvironment)
try:
@@ -120,15 +104,12 @@ class BaseModalExecutionEnvironment(BaseEnvironment):
def _prepare_modal_exec(self, command: str, *, cwd: str = "", timeout: int | None = None,
stdin_data: str | None = None) -> PreparedModalExec:
exec_command = command
exec_stdin = stdin_data if self._stdin_mode == "payload" else None
if stdin_data is not None and self._stdin_mode == "heredoc":
exec_command = wrap_modal_stdin_heredoc(exec_command, stdin_data)
exec_command, sudo_stdin = self._prepare_command(exec_command)
exec_command, sudo_stdin = self._prepare_command(command)
if sudo_stdin is not None:
exec_command = wrap_modal_sudo_pipe(exec_command, sudo_stdin)
# Feed sudo via a shell pipe: the transport has no direct stdin piping.
exec_command = f"printf '%s\\n' {shlex.quote(sudo_stdin.rstrip())} | {exec_command}"
return PreparedModalExec(command=exec_command, cwd=cwd or self.cwd, timeout=timeout or self.timeout,
stdin_data=exec_stdin)
stdin_data=stdin_data)
def _result(self, output: str, returncode: int) -> dict:
return {"output": output, "returncode": returncode}

View File

@@ -1,5 +1,6 @@
"""SSH remote execution environment with ControlMaster connection persistence."""
import contextlib
import hashlib
import logging
import os
@@ -24,10 +25,15 @@ _SSH_MULTIPLEX = os.name != "nt"
def _ensure_ssh_available() -> None:
"""Fail fast with a clear error when the SSH client is unavailable."""
if not shutil.which("ssh"):
raise RuntimeError("SSH is not installed or not in PATH. Install OpenSSH client: apt install openssh-client")
if not shutil.which("scp"):
raise RuntimeError("SCP is not installed or not in PATH. Install OpenSSH client: apt install openssh-client")
for tool in ("ssh", "scp"):
if not shutil.which(tool):
raise RuntimeError(f"{tool.upper()} is not installed or not in PATH. "
"Install OpenSSH client: apt install openssh-client")
def _sync_error(reason: str, subject: str, what: str = "the SSH connection") -> EnvironmentConnectionError:
return EnvironmentConnectionError(
reason, retry_hint=f"{subject} failed — verify {what} is healthy, then retry.")
class SSHEnvironment(BaseEnvironment):
@@ -57,68 +63,63 @@ class SSHEnvironment(BaseEnvironment):
_ensure_ssh_available()
self._establish_connection()
self._remote_home = self._detect_remote_home()
self._ensure_remote_dirs()
self._sync_manager = FileSyncManager(
get_files_fn=lambda: iter_sync_files(f"{self._remote_home}/.hermes"),
upload_fn=self._scp_upload,
delete_fn=self._ssh_delete,
bulk_upload_fn=self._ssh_bulk_upload,
bulk_download_fn=self._ssh_bulk_download,
)
upload_fn=self._scp_upload, delete_fn=self._ssh_delete,
bulk_upload_fn=self._ssh_bulk_upload, bulk_download_fn=self._ssh_bulk_download)
self._sync_manager.sync(force=True)
self.init_session()
def _target_flags(self, port_flag: str) -> list:
"""Port/key flags shared by ssh (``-p``) and scp (``-P``)."""
flags = [port_flag, str(self.port)] if self.port != 22 else []
return flags + (["-i", self.key_path] if self.key_path else [])
def _build_ssh_command(self, extra_args: list | None = None) -> list:
cmd = ["ssh"]
if _SSH_MULTIPLEX:
cmd.extend(["-o", f"ControlPath={self.control_socket}",
"-o", "ControlMaster=auto", "-o", "ControlPersist=300"])
cmd.extend(["-o", "BatchMode=yes", "-o", "StrictHostKeyChecking=accept-new",
"-o", "ConnectTimeout=10"])
if self.port != 22:
cmd.extend(["-p", str(self.port)])
if self.key_path:
cmd.extend(["-i", self.key_path])
if extra_args:
cmd.extend(extra_args)
cmd.extend(["-o", "BatchMode=yes", "-o", "StrictHostKeyChecking=accept-new", "-o", "ConnectTimeout=10"])
cmd.extend(self._target_flags("-p"))
cmd.extend(extra_args or [])
cmd.append(f"{self.user}@{self.host}")
return cmd
def _run_ssh(self, remote_cmd: str, timeout: float) -> subprocess.CompletedProcess:
"""Run one remote shell command over the multiplexed connection, capturing output."""
cmd = self._build_ssh_command()
cmd.append(remote_cmd)
return run_capture(cmd, timeout=timeout)
return run_capture(self._build_ssh_command() + [remote_cmd], timeout=timeout)
def _run_ssh_checked(self, remote_cmd: str, timeout: float, reason: str, subject: str) -> None:
result = self._run_ssh(remote_cmd, timeout=timeout)
if result.returncode != 0:
raise _sync_error(f"{reason}: {result.stderr.strip()}", subject)
def _establish_connection(self):
try:
result = self._run_ssh("echo 'SSH connection established'", timeout=15)
if result.returncode != 0:
error_msg = result.stderr.strip() or result.stdout.strip()
raise EnvironmentConnectionError(
f"SSH connection failed: {error_msg}",
retry_hint=(f"Verify {self.user}@{self.host}:{self.port} is reachable "
"(host up, sshd running, key/agent auth working), then "
"retry — the connection is re-established automatically."),
)
except subprocess.TimeoutExpired:
raise EnvironmentConnectionError(
f"SSH connection to {self.user}@{self.host} timed out",
retry_hint=(f"Check network connectivity to {self.host}:{self.port} "
"and that sshd is accepting connections, then retry."),
)
"and that sshd is accepting connections, then retry."))
if result.returncode != 0:
error_msg = result.stderr.strip() or result.stdout.strip()
raise EnvironmentConnectionError(
f"SSH connection failed: {error_msg}",
retry_hint=(f"Verify {self.user}@{self.host}:{self.port} is reachable "
"(host up, sshd running, key/agent auth working), then "
"retry — the connection is re-established automatically."))
def _detect_remote_home(self) -> str:
"""Detect the remote user's home directory."""
try:
with contextlib.suppress(Exception):
result = self._run_ssh("echo $HOME", timeout=10)
home = result.stdout.strip()
if home and result.returncode == 0:
logger.debug("SSH: remote home = %s", home)
return home
except Exception:
pass
return "/root" if self.user == "root" else f"/home/{self.user}"
# -- File sync (via FileSyncManager) --------------------------------
@@ -126,29 +127,17 @@ class SSHEnvironment(BaseEnvironment):
def _ensure_remote_dirs(self) -> None:
"""Create base ~/.hermes directory tree on remote in one SSH call."""
base = f"{self._remote_home}/.hermes"
dirs = [base, f"{base}/skills", f"{base}/credentials", f"{base}/cache"]
self._run_ssh(quoted_mkdir_command(dirs), timeout=10)
self._run_ssh(quoted_mkdir_command([base, f"{base}/skills", f"{base}/credentials", f"{base}/cache"]),
timeout=10)
def _scp_upload(self, host_path: str, remote_path: str) -> None:
"""Upload a single file via scp over ControlMaster."""
parent = str(Path(remote_path).parent)
self._run_ssh(f"mkdir -p {shlex.quote(parent)}", timeout=10)
scp_cmd = ["scp"]
if _SSH_MULTIPLEX:
scp_cmd.extend(["-o", f"ControlPath={self.control_socket}"])
if self.port != 22:
scp_cmd.extend(["-P", str(self.port)])
if self.key_path:
scp_cmd.extend(["-i", self.key_path])
scp_cmd.extend([host_path, f"{self.user}@{self.host}:{remote_path}"])
self._run_ssh(f"mkdir -p {shlex.quote(str(Path(remote_path).parent))}", timeout=10)
scp_cmd = ["scp"] + (["-o", f"ControlPath={self.control_socket}"] if _SSH_MULTIPLEX else [])
scp_cmd += self._target_flags("-P") + [host_path, f"{self.user}@{self.host}:{remote_path}"]
result = run_capture(scp_cmd, timeout=30)
if result.returncode != 0:
raise EnvironmentConnectionError(
f"scp failed: {result.stderr.strip()}",
retry_hint=(f"File sync to {self.user}@{self.host} failed — verify the "
"SSH connection is healthy, then retry."),
)
raise _sync_error(f"scp failed: {result.stderr.strip()}", f"File sync to {self.user}@{self.host}")
def _ssh_bulk_upload(self, files: list[tuple[str, str]]) -> None:
"""Upload many files in a single tar-over-SSH stream.
@@ -158,17 +147,11 @@ class SSHEnvironment(BaseEnvironment):
"""
if not files:
return
base = f"{self._remote_home}/.hermes"
parents = unique_parent_dirs(files)
if parents:
result = self._run_ssh(quoted_mkdir_command(parents), timeout=30)
if result.returncode != 0:
raise EnvironmentConnectionError(
f"remote mkdir failed: {result.stderr.strip()}",
retry_hint=(f"Remote directory setup on {self.host} failed — verify "
"the SSH connection is healthy, then retry."),
)
self._run_ssh_checked(quoted_mkdir_command(parents), 30, "remote mkdir failed",
f"Remote directory setup on {self.host}")
# Symlink staging avoids fragile GNU tar --transform rules. On Windows
# without Developer Mode symlink creation raises OSError winerror 1314;
@@ -181,24 +164,19 @@ class SSHEnvironment(BaseEnvironment):
raise RuntimeError(f"remote path {remote_path!r} is not under sync base {base!r}") from exc
if rel_remote == "." or rel_remote.startswith("../"):
raise RuntimeError(f"remote path {remote_path!r} escapes sync base {base!r}")
staged = os.path.join(staging, rel_remote)
os.makedirs(os.path.dirname(staged), exist_ok=True)
try:
os.symlink(os.path.abspath(host_path), staged)
except OSError as e:
if getattr(e, "winerror", None) == 1314:
shutil.copy2(host_path, staged)
else:
if getattr(e, "winerror", None) != 1314:
raise
shutil.copy2(host_path, staged)
tar_cmd = ["tar", "-chf", "-", "-C", staging, "."]
ssh_cmd = self._build_ssh_command()
# --no-overwrite-dir keeps tar from stamping the staging dir's mode onto
# existing dirs (e.g. /home/<user>); a umask-002 0775 home breaks sshd StrictModes.
ssh_cmd.append(f"tar xf - --no-overwrite-dir -C {shlex.quote(base)}")
tar_proc = subprocess.Popen(tar_cmd, stdin=subprocess.DEVNULL,
ssh_cmd = self._build_ssh_command() + [f"tar xf - --no-overwrite-dir -C {shlex.quote(base)}"]
tar_proc = subprocess.Popen(["tar", "-chf", "-", "-C", staging, "."], stdin=subprocess.DEVNULL,
stdout=subprocess.PIPE, stderr=subprocess.PIPE)
try:
ssh_proc = subprocess.Popen(ssh_cmd, stdin=tar_proc.stdout,
@@ -207,39 +185,29 @@ class SSHEnvironment(BaseEnvironment):
tar_proc.kill()
tar_proc.wait()
raise
# Allow tar_proc to receive SIGPIPE if ssh_proc exits early
tar_proc.stdout.close()
tar_proc.stdout.close() # let tar_proc receive SIGPIPE if ssh_proc exits early
try:
_, ssh_stderr = ssh_proc.communicate(timeout=120)
# communicate() (not wait()) drains stderr so tar can't deadlock on >PIPE_BUF errors.
tar_stderr_raw = b""
if tar_proc.poll() is None:
_, tar_stderr_raw = tar_proc.communicate(timeout=10)
else:
tar_stderr_raw = tar_proc.stderr.read() if tar_proc.stderr else b""
except subprocess.TimeoutExpired:
tar_proc.kill()
ssh_proc.kill()
tar_proc.wait()
ssh_proc.wait()
for proc in (tar_proc, ssh_proc):
proc.kill()
for proc in (tar_proc, ssh_proc):
proc.wait()
raise EnvironmentConnectionError(
"SSH bulk upload timed out",
retry_hint=f"Bulk file sync to {self.host} timed out — check the connection and retry.",
)
retry_hint=f"Bulk file sync to {self.host} timed out — check the connection and retry.")
if tar_proc.returncode != 0:
raise RuntimeError(f"tar create failed (rc={tar_proc.returncode}): "
f"{tar_stderr_raw.decode(errors='replace').strip()}")
if ssh_proc.returncode != 0:
raise EnvironmentConnectionError(
f"tar extract over SSH failed (rc={ssh_proc.returncode}): "
f"{ssh_stderr.decode(errors='replace').strip()}",
retry_hint=(f"File sync over SSH to {self.host} failed — verify the "
"connection is healthy, then retry."),
)
raise _sync_error(f"tar extract over SSH failed (rc={ssh_proc.returncode}): "
f"{ssh_stderr.decode(errors='replace').strip()}",
f"File sync over SSH to {self.host}", what="the connection")
logger.debug("SSH: bulk-uploaded %d file(s) via tar pipe", len(files))
def _ssh_bulk_download(self, dest: Path) -> None:
@@ -247,27 +215,17 @@ class SSHEnvironment(BaseEnvironment):
# Tar from / with the full path so archive entries keep absolute paths
# (home/user/.hermes/skills/f.py), matching _pushed_hashes keys.
rel_base = f"{self._remote_home}/.hermes".lstrip("/")
ssh_cmd = self._build_ssh_command()
ssh_cmd.append(f"tar cf - -C / {shlex.quote(rel_base)}")
ssh_cmd = self._build_ssh_command() + [f"tar cf - -C / {shlex.quote(rel_base)}"]
with open(dest, "wb") as f:
result = subprocess.run(ssh_cmd, stdin=subprocess.DEVNULL, stdout=f,
stderr=subprocess.PIPE, timeout=120)
result = subprocess.run(ssh_cmd, stdin=subprocess.DEVNULL, stdout=f, stderr=subprocess.PIPE, timeout=120)
if result.returncode != 0:
raise EnvironmentConnectionError(
f"SSH bulk download failed: {result.stderr.decode(errors='replace').strip()}",
retry_hint=(f"File sync from {self.host} failed — verify the SSH "
"connection is healthy, then retry."),
)
raise _sync_error(f"SSH bulk download failed: {result.stderr.decode(errors='replace').strip()}",
f"File sync from {self.host}")
def _ssh_delete(self, remote_paths: list[str]) -> None:
"""Batch-delete remote files in one SSH call."""
result = self._run_ssh(quoted_rm_command(remote_paths), timeout=10)
if result.returncode != 0:
raise EnvironmentConnectionError(
f"remote rm failed: {result.stderr.strip()}",
retry_hint=(f"Remote file cleanup on {self.host} failed — verify the "
"SSH connection is healthy, then retry."),
)
self._run_ssh_checked(quoted_rm_command(remote_paths), 10, "remote rm failed",
f"Remote file cleanup on {self.host}")
def _before_execute(self) -> None:
"""Sync files to remote via FileSyncManager (rate-limited internally)."""
@@ -278,22 +236,15 @@ class SSHEnvironment(BaseEnvironment):
def _run_bash(self, cmd_string: str, *, login: bool = False, timeout: int = 120,
stdin_data: str | None = None) -> subprocess.Popen:
"""Spawn an SSH process that runs bash on the remote host."""
cmd = self._build_ssh_command()
cmd.extend(bash_argv(shlex.quote(cmd_string), login))
return _popen_bash(cmd, stdin_data)
return _popen_bash(self._build_ssh_command() + bash_argv(shlex.quote(cmd_string), login), stdin_data)
def cleanup(self):
if self._sync_manager:
logger.info("SSH: syncing files from sandbox...")
self._sync_manager.sync_back()
if self.control_socket.exists():
try:
with contextlib.suppress(OSError, subprocess.SubprocessError):
cmd = ["ssh", "-o", f"ControlPath={self.control_socket}", "-O", "exit", f"{self.user}@{self.host}"]
subprocess.run(cmd, capture_output=True, timeout=5, stdin=subprocess.DEVNULL)
except (OSError, subprocess.SubprocessError):
pass
try:
with contextlib.suppress(OSError):
self.control_socket.unlink()
except OSError:
pass

View File

@@ -7,15 +7,15 @@ under ``HERMES_HOME`` and new sandboxes are restored from them on task reuse.
from __future__ import annotations
from functools import cache
from dataclasses import dataclass
from datetime import timedelta
import contextlib
import logging
import math
import os
import shlex
import threading
import time
from datetime import timedelta
from functools import cache
from pathlib import Path
from typing import TYPE_CHECKING, Any
@@ -29,68 +29,43 @@ from tools.environments.remote_common import ensure_lazy_dep
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from vercel.sandbox import Resources, Sandbox, SandboxStatus, WriteFile
from vercel.sandbox import Sandbox, SandboxStatus, WriteFile
DEFAULT_VERCEL_CWD = "/vercel/sandbox"
_DEFAULT_CONTAINER_DISK_MB = 51200
_CREATE_RETRY_ATTEMPTS = 3
_TRANSIENT_STATUS_CODES = frozenset({408, 425, 429, 500, 502, 503, 504})
_RUNNING_WAIT_TIMEOUT = timedelta(seconds=30)
_SNAPSHOT_STORE_NAME = "vercel_sandbox_snapshots.json"
def _ensure_vercel_sdk() -> None:
"""Lazy-install vercel SDK on demand. Idempotent."""
# The vercel SDK (>=0.7) ships default-on usage telemetry posting to
# telemetry.vercel.com. Hermes policy is no outbound telemetry without
# explicit opt-in, so disable it before the SDK is ever imported. Only the
# default is set — an explicit user value (e.g. "0") is never overridden.
# The vercel SDK (>=0.7) ships default-on usage telemetry. Hermes policy is no
# outbound telemetry without explicit opt-in, so disable it before the SDK is
# ever imported. Only the default is set — an explicit user value is never overridden.
os.environ.setdefault("VERCEL_TELEMETRY_DISABLED", "1")
ensure_lazy_dep("terminal.vercel")
_CREATE_RETRY_ATTEMPTS = 3
_WRITE_RETRY_ATTEMPTS = 3
_TRANSIENT_STATUS_CODES = frozenset({408, 425, 429, 500, 502, 503, 504})
_RETRY_BACKOFF_STEP = timedelta(milliseconds=100)
_MIN_SANDBOX_TIMEOUT = timedelta(minutes=5)
_MIN_RUNNING_WAIT = timedelta(seconds=1)
_RUNNING_WAIT_TIMEOUT = timedelta(seconds=30)
_RUNNING_WAIT_POLL_INTERVAL = timedelta(milliseconds=250)
_STOP_TIMEOUT = timedelta(seconds=15)
_STOP_POLL_INTERVAL = timedelta(milliseconds=500)
_SNAPSHOT_STORE_NAME = "vercel_sandbox_snapshots.json"
_SNAPSHOT_ID_KEYS = ("snapshot_id", "snapshotId", "id")
_MISSING = object()
def _exception_chain(exc: BaseException) -> list[BaseException]:
chain: list[BaseException] = []
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
chain.append(current)
seen.add(id(current))
current = current.__cause__ or current.__context__
return chain
def _extract_status_code(exc: BaseException) -> int | None:
response = getattr(exc, "response", None)
for value in (getattr(exc, "status_code", None), getattr(response, "status_code", None)):
if isinstance(value, int):
return value
return None
def _is_transient_vercel_error(exc: BaseException) -> bool:
for error in _exception_chain(exc):
error_name = type(error).__name__.lower()
if (_extract_status_code(error) in _TRANSIENT_STATUS_CODES
"""True when any exception in the cause/context chain looks retryable."""
seen: set[int] = set()
error: BaseException | None = exc
while error is not None and id(error) not in seen:
seen.add(id(error))
codes = (getattr(error, "status_code", None), getattr(getattr(error, "response", None), "status_code", None))
status = next((c for c in codes if isinstance(c, int)), None)
name = type(error).__name__.lower()
if (status in _TRANSIENT_STATUS_CODES
or isinstance(error, (httpx.NetworkError, httpx.ProtocolError, httpx.ReadError))
or "ratelimit" in error_name or "servererror" in error_name):
or "ratelimit" in name or "servererror" in name):
return True
error = error.__cause__ or error.__context__
return False
def _retry_vercel_call(label: str, callback, *, attempts: int):
backoff_seconds = _RETRY_BACKOFF_STEP.total_seconds()
for attempt in range(1, attempts + 1):
try:
return callback()
@@ -98,7 +73,7 @@ def _retry_vercel_call(label: str, callback, *, attempts: int):
if attempt >= attempts or not _is_transient_vercel_error(exc):
raise
logger.warning("Vercel: %s failed (%s); retrying %d/%d", label, exc, attempt, attempts)
time.sleep(backoff_seconds * attempt)
time.sleep(0.1 * attempt)
def _coerce_text(value: Any) -> str:
@@ -115,9 +90,8 @@ def _extract_result_output(result: Any) -> str:
def _extract_result_returncode(result: Any) -> int:
exit_code = getattr(result, "exit_code", _MISSING)
if exit_code is _MISSING:
exit_code = getattr(result, "returncode", None)
attr = "exit_code" if hasattr(result, "exit_code") else "returncode"
exit_code = getattr(result, attr, None)
return exit_code if isinstance(exit_code, int) else 1
@@ -130,38 +104,28 @@ def _save_snapshots(data: dict) -> None:
def _get_snapshot_id(task_id: str) -> str | None:
if not task_id:
return None
snapshot_id = _load_snapshots().get(task_id)
snapshot_id = _load_snapshots().get(task_id) if task_id else None
return snapshot_id if isinstance(snapshot_id, str) and snapshot_id else None
def _store_snapshot(task_id: str, snapshot_id: str) -> None:
if not task_id or not snapshot_id:
return
snapshots = _load_snapshots()
snapshots[task_id] = snapshot_id
_save_snapshots(snapshots)
if task_id and snapshot_id:
_save_snapshots({**_load_snapshots(), task_id: snapshot_id})
def _delete_snapshot(task_id: str, snapshot_id: str | None = None) -> None:
if not task_id:
return
snapshots = _load_snapshots()
existing = snapshots.get(task_id)
if existing is None or (snapshot_id is not None and existing != snapshot_id):
return
snapshots.pop(task_id, None)
_save_snapshots(snapshots)
existing = snapshots.get(task_id) if task_id else None
if existing is not None and (snapshot_id is None or existing == snapshot_id):
snapshots.pop(task_id, None)
_save_snapshots(snapshots)
def _extract_snapshot_id(snapshot: Any) -> str | None:
"""Accept SDK objects or raw dicts; attribute lookup first, then dict keys."""
getters = [lambda k: getattr(snapshot, k, None)]
if isinstance(snapshot, dict):
getters.append(snapshot.get)
getters = [lambda k: getattr(snapshot, k, None)] + ([snapshot.get] if isinstance(snapshot, dict) else [])
for get in getters:
for key in _SNAPSHOT_ID_KEYS:
for key in ("snapshot_id", "snapshotId", "id"):
value = get(key)
if isinstance(value, str) and value:
return value
@@ -175,17 +139,9 @@ def _sandbox_status_type() -> type[SandboxStatus]:
return SandboxStatus
@cache
def _terminal_sandbox_states() -> frozenset[SandboxStatus]:
SandboxStatus = _sandbox_status_type()
return frozenset({SandboxStatus.ABORTED, SandboxStatus.FAILED, SandboxStatus.STOPPED})
@dataclass(frozen=True, slots=True)
class _SandboxCreateParams:
timeout: timedelta
runtime: str | None = None
resources: Resources | None = None
def _is_terminal(status: Any) -> bool:
S = _sandbox_status_type()
return status in {S.ABORTED, S.FAILED, S.STOPPED}
class VercelSandboxEnvironment(BaseEnvironment):
@@ -197,7 +153,11 @@ class VercelSandboxEnvironment(BaseEnvironment):
cpu: float = 1, memory: int = 5120, disk: int = _DEFAULT_CONTAINER_DISK_MB,
persistent_filesystem: bool = True, task_id: str = "default"):
super().__init__(cwd=cwd, timeout=timeout)
self._runtime = runtime or None
if disk not in {0, _DEFAULT_CONTAINER_DISK_MB}:
raise ValueError(
"Vercel Sandbox does not support configurable container_disk. "
"Use the default shared setting."
)
self._persistent = persistent_filesystem
self._task_id = task_id
self._requested_cwd = cwd
@@ -206,77 +166,57 @@ class VercelSandboxEnvironment(BaseEnvironment):
self._workspace_root = DEFAULT_VERCEL_CWD
self._remote_home = DEFAULT_VERCEL_CWD
self._sync_manager: FileSyncManager | None = None
self._create_params = self._build_create_params(cpu=cpu, memory=memory, disk=disk)
_ensure_vercel_sdk()
from vercel.sandbox import Resources
vcpus = math.floor(cpu) if cpu > 0 else None
memory_mb = memory if memory > 0 else None
resources = Resources(vcpus=vcpus, memory=memory_mb) if vcpus is not None or memory_mb is not None else None
self._create_kwargs = {
"timeout": max(timedelta(seconds=max(self.timeout, 0)), timedelta(minutes=5)),
"runtime": runtime or None, "resources": resources,
}
self._attach_fresh_sandbox(cwd)
self._sync_manager.sync(force=True)
self.init_session()
def _require_sandbox(self) -> Sandbox:
sandbox = self._sandbox
if sandbox is None:
if self._sandbox is None:
raise RuntimeError("Vercel sandbox is not attached")
return sandbox
def _attach_fresh_sandbox(self, requested_cwd: str) -> None:
self._sandbox = self._create_sandbox()
self._configure_attached_sandbox(requested_cwd=requested_cwd)
return self._sandbox
def _remote_hermes_dir(self) -> str:
home = self._remote_home
return "/.hermes" if home == "/" else f"{home.rstrip('/')}/.hermes"
def _build_create_params(self, *, cpu: float, memory: int, disk: int) -> _SandboxCreateParams:
if disk not in {0, _DEFAULT_CONTAINER_DISK_MB}:
raise ValueError(
"Vercel Sandbox does not support configurable container_disk. "
"Use the default shared setting."
)
_ensure_vercel_sdk()
from vercel.sandbox import Resources
sandbox_timeout = max(timedelta(seconds=max(self.timeout, 0)), _MIN_SANDBOX_TIMEOUT)
vcpus = math.floor(cpu) if cpu > 0 else None
memory_mb = memory if memory > 0 else None
resources = Resources(vcpus=vcpus, memory=memory_mb) if vcpus is not None or memory_mb is not None else None
return _SandboxCreateParams(timeout=sandbox_timeout, runtime=self._runtime, resources=resources)
def _create_sandbox(self) -> Sandbox:
_ensure_vercel_sdk()
from vercel.sandbox import Sandbox
params = self._create_params
snapshot_id = _get_snapshot_id(self._task_id) if self._persistent else None
if snapshot_id:
try:
return _retry_vercel_call(
"sandbox restore",
lambda: Sandbox.create(timeout=params.timeout, runtime=params.runtime, resources=params.resources,
source={"type": "snapshot", "snapshot_id": snapshot_id}),
attempts=_CREATE_RETRY_ATTEMPTS,
)
lambda: Sandbox.create(**self._create_kwargs, source={"type": "snapshot", "snapshot_id": snapshot_id}),
attempts=_CREATE_RETRY_ATTEMPTS)
except Exception as exc:
logger.warning(
"Vercel: failed to restore snapshot %s for task %s; "
"falling back to a fresh sandbox: %s",
snapshot_id, self._task_id, exc,
)
logger.warning("Vercel: failed to restore snapshot %s for task %s; falling back to a fresh sandbox: %s",
snapshot_id, self._task_id, exc)
_delete_snapshot(self._task_id, snapshot_id)
return _retry_vercel_call(
"sandbox create",
lambda: Sandbox.create(timeout=params.timeout, runtime=params.runtime, resources=params.resources),
attempts=_CREATE_RETRY_ATTEMPTS,
)
return _retry_vercel_call("sandbox create", lambda: Sandbox.create(**self._create_kwargs),
attempts=_CREATE_RETRY_ATTEMPTS)
def _configure_attached_sandbox(self, *, requested_cwd: str) -> None:
def _attach_fresh_sandbox(self, requested_cwd: str) -> None:
"""Create a sandbox, wait until it runs, then wire cwd/home and the file sync manager."""
self._sandbox = self._create_sandbox()
self._wait_for_running()
self._workspace_root = self._detect_workspace_root()
cwd = self._require_sandbox().sandbox.cwd
self._workspace_root = cwd if cwd.startswith("/") else DEFAULT_VERCEL_CWD
self._remote_home = self._detect_remote_home()
container_base = self._remote_hermes_dir()
self._sync_manager = FileSyncManager(
get_files_fn=lambda: iter_sync_files(container_base),
upload_fn=self._vercel_upload,
delete_fn=self._vercel_delete,
bulk_upload_fn=self._vercel_bulk_upload,
bulk_download_fn=self._vercel_bulk_download,
)
upload_fn=self._vercel_upload, delete_fn=self._vercel_delete,
bulk_upload_fn=self._vercel_bulk_upload, bulk_download_fn=self._vercel_bulk_download)
if requested_cwd == "~":
self.cwd = self._remote_home
elif requested_cwd in {"", DEFAULT_VERCEL_CWD}:
@@ -284,14 +224,9 @@ class VercelSandboxEnvironment(BaseEnvironment):
else:
self.cwd = requested_cwd
def _detect_workspace_root(self) -> str:
cwd = self._require_sandbox().sandbox.cwd
return cwd if cwd.startswith("/") else DEFAULT_VERCEL_CWD
def _detect_remote_home(self) -> str:
sandbox = self._require_sandbox()
try:
result = sandbox.run_command("sh", ["-lc", 'printf %s "$HOME"'], cwd=self._workspace_root)
result = self._require_sandbox().run_command("sh", ["-lc", 'printf %s "$HOME"'], cwd=self._workspace_root)
except Exception as exc:
logger.debug("Vercel: home detection failed for task %s: %s", self._task_id, exc)
return self._workspace_root
@@ -300,51 +235,42 @@ class VercelSandboxEnvironment(BaseEnvironment):
def _wait_for_running(self, timeout: timedelta = _RUNNING_WAIT_TIMEOUT) -> None:
sandbox = self._require_sandbox()
SandboxStatus = _sandbox_status_type()
status = sandbox.status
if status is None or status == SandboxStatus.RUNNING:
if status is None or status == _sandbox_status_type().RUNNING:
return
if status in _terminal_sandbox_states():
if _is_terminal(status):
raise RuntimeError(f"Sandbox entered terminal state: {status}")
try:
sandbox.wait_for_status(SandboxStatus.RUNNING, timeout=max(timeout, _MIN_RUNNING_WAIT),
poll_interval=_RUNNING_WAIT_POLL_INTERVAL)
sandbox.wait_for_status(_sandbox_status_type().RUNNING, timeout=max(timeout, timedelta(seconds=1)),
poll_interval=timedelta(milliseconds=250))
except TimeoutError as exc:
status = sandbox.status
if status in _terminal_sandbox_states():
if _is_terminal(status):
raise RuntimeError(f"Sandbox entered terminal state: {status}") from exc
raise RuntimeError(f"Sandbox did not reach running state (last status: {status})") from exc
def _close_sandbox_client(self, sandbox: Sandbox | None) -> None:
if sandbox is None:
return
try:
sandbox.client.close()
except Exception:
pass
if sandbox is not None:
with contextlib.suppress(Exception):
sandbox.client.close()
def _stop_sandbox(self, sandbox: Sandbox | None) -> None:
if sandbox is None:
return
try:
sandbox.stop(blocking=True, timeout=_STOP_TIMEOUT, poll_interval=_STOP_POLL_INTERVAL)
except TypeError:
with contextlib.suppress(Exception):
try:
sandbox.stop(blocking=True, timeout=timedelta(seconds=15), poll_interval=timedelta(milliseconds=500))
except TypeError: # older SDKs: stop() takes no arguments
sandbox.stop()
except Exception:
pass
except Exception:
pass
def _snapshot_sandbox(self, sandbox: Sandbox) -> str | None:
if not self._persistent or not self._task_id:
return None
try:
snapshot = sandbox.snapshot()
snapshot_id = _extract_snapshot_id(sandbox.snapshot())
except Exception as exc:
logger.warning("Vercel: filesystem snapshot failed for task %s: %s", self._task_id, exc)
return None
snapshot_id = _extract_snapshot_id(snapshot)
if not snapshot_id:
logger.warning("Vercel: filesystem snapshot for task %s did not return a snapshot id", self._task_id)
return None
@@ -353,25 +279,27 @@ class VercelSandboxEnvironment(BaseEnvironment):
return snapshot_id
def _ensure_sandbox_ready(self) -> None:
"""Reuse a healthy sandbox; recreate when refresh fails or it hit a terminal state."""
sandbox = self._sandbox
requested_cwd = self.cwd or self._requested_cwd or DEFAULT_VERCEL_CWD
if sandbox is None:
self._attach_fresh_sandbox(requested_cwd)
return
try:
sandbox.refresh()
except Exception as exc:
logger.warning("Vercel: sandbox refresh failed for task %s: %s; recreating", self._task_id, exc)
if sandbox is not None:
try:
sandbox.refresh()
except Exception as exc:
logger.warning("Vercel: sandbox refresh failed for task %s: %s; recreating", self._task_id, exc)
else:
status = sandbox.status
if not _is_terminal(status):
self._wait_for_running()
return
logger.warning("Vercel: sandbox entered state %s for task %s; recreating", status, self._task_id)
self._close_sandbox_client(sandbox)
self._attach_fresh_sandbox(requested_cwd)
return
status = sandbox.status
if status in _terminal_sandbox_states():
logger.warning("Vercel: sandbox entered state %s for task %s; recreating", status, self._task_id)
self._close_sandbox_client(sandbox)
self._attach_fresh_sandbox(requested_cwd)
return
self._wait_for_running()
self._attach_fresh_sandbox(requested_cwd)
def _run_checked(self, script: str, label: str) -> None:
result = self._require_sandbox().run_command("bash", ["-lc", script], cwd=self._workspace_root)
if _extract_result_returncode(result) != 0:
raise RuntimeError(f"Vercel {label} failed: {_extract_result_output(result).strip()}")
def _vercel_upload(self, host_path: str, remote_path: str) -> None:
self._vercel_bulk_upload([(host_path, remote_path)])
@@ -380,37 +308,24 @@ class VercelSandboxEnvironment(BaseEnvironment):
if not files:
return
payload: list[WriteFile] = [
{"path": remote_path, "content": Path(host_path).read_bytes()} for host_path, remote_path in files
]
{"path": remote_path, "content": Path(host_path).read_bytes()} for host_path, remote_path in files]
sandbox = self._require_sandbox()
_retry_vercel_call("write_files", lambda: sandbox.write_files(payload), attempts=_WRITE_RETRY_ATTEMPTS)
_retry_vercel_call("write_files", lambda: sandbox.write_files(payload), attempts=3)
def _vercel_delete(self, remote_paths: list[str]) -> None:
if not remote_paths:
return
result = self._require_sandbox().run_command(
"bash", ["-lc", quoted_rm_command(remote_paths)], cwd=self._workspace_root,
)
if _extract_result_returncode(result) != 0:
raise RuntimeError(f"Vercel delete failed: {_extract_result_output(result).strip()}")
if remote_paths:
self._run_checked(quoted_rm_command(remote_paths), "delete")
def _vercel_bulk_download(self, dest_tar_path: Path) -> None:
archive_member = self._remote_hermes_dir().lstrip("/")
remote_tar = f"/tmp/.hermes_sync.{os.getpid()}.tar"
sandbox = self._require_sandbox()
try:
result = sandbox.run_command(
"bash", ["-lc", f"tar cf {shlex.quote(remote_tar)} -C / {shlex.quote(archive_member)}"],
cwd=self._workspace_root,
)
if _extract_result_returncode(result) != 0:
raise RuntimeError(f"Vercel bulk download failed: {_extract_result_output(result).strip()}")
self._run_checked(f"tar cf {shlex.quote(remote_tar)} -C / {shlex.quote(archive_member)}", "bulk download")
sandbox.download_file(remote_tar, dest_tar_path)
finally:
try:
with contextlib.suppress(Exception):
sandbox.run_command("bash", ["-lc", f"rm -f {shlex.quote(remote_tar)}"], cwd=self._workspace_root)
except Exception:
pass
def _before_execute(self) -> None:
with self._lock:
@@ -439,8 +354,7 @@ class VercelSandboxEnvironment(BaseEnvironment):
def cleanup(self):
with self._lock:
sandbox = self._sandbox
sync_manager = self._sync_manager
sandbox, sync_manager = self._sandbox, self._sync_manager
if sandbox is not None and sync_manager is not None:
try:
sync_manager.sync_back()