refactor(tools): simplify remote env backends (modal/managed_modal/daytona/ssh/vercel) — dead code, shared helpers, defensive collapse
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user