refactor(tools): fold modal_utils abstract seam into managed_modal (its only consumer); layout hug

This commit is contained in:
Teknium
2026-09-02 22:12:25 -07:00
parent 876790dcc0
commit e1cacc050e
7 changed files with 100 additions and 182 deletions
@@ -109,7 +109,6 @@ class _FakeResponse:
def test_managed_modal_execute_polls_until_completed(monkeypatch):
_install_fake_tools_package()
managed_modal = _load_tool_module("tools.environments.managed_modal", "environments/managed_modal.py")
modal_common = sys.modules["tools.environments.modal_utils"]
calls = []
poll_count = {"value": 0}
@@ -135,7 +134,7 @@ def test_managed_modal_execute_polls_until_completed(monkeypatch):
raise AssertionError(f"Unexpected request: {method} {url}")
monkeypatch.setattr(managed_modal.requests, "request", fake_request)
monkeypatch.setattr(modal_common.time, "sleep", lambda _: None)
monkeypatch.setattr(managed_modal.time, "sleep", lambda _: None)
env = managed_modal.ManagedModalEnvironment(image="python:3.11")
result = env.execute("echo hello")
@@ -161,7 +160,6 @@ def test_managed_modal_rejects_host_credential_passthrough():
def test_managed_modal_execute_times_out_and_cancels(monkeypatch):
_install_fake_tools_package()
managed_modal = _load_tool_module("tools.environments.managed_modal", "environments/managed_modal.py")
modal_common = sys.modules["tools.environments.modal_utils"]
calls = []
monotonic_values = iter([0.0, 0.0, 0.0, 12.5, 12.5])
@@ -181,8 +179,8 @@ def test_managed_modal_execute_times_out_and_cancels(monkeypatch):
raise AssertionError(f"Unexpected request: {method} {url}")
monkeypatch.setattr(managed_modal.requests, "request", fake_request)
monkeypatch.setattr(modal_common.time, "monotonic", lambda: next(monotonic_values))
monkeypatch.setattr(modal_common.time, "sleep", lambda _: None)
monkeypatch.setattr(managed_modal.time, "monotonic", lambda: next(monotonic_values))
monkeypatch.setattr(managed_modal.time, "sleep", lambda _: None)
env = managed_modal.ManagedModalEnvironment(image="python:3.11")
result = env.execute("sleep 30", timeout=2)
+1 -2
View File
@@ -14,8 +14,7 @@ from pathlib import Path
from tools.environments.base import BaseEnvironment, _ThreadedProcessHandle
from tools.environments.file_sync import (
FileSyncManager, iter_sync_files, quoted_mkdir_command, quoted_rm_command, unique_parent_dirs,
)
FileSyncManager, iter_sync_files, quoted_mkdir_command, quoted_rm_command, unique_parent_dirs)
from tools.environments.remote_common import ensure_lazy_dep
logger = logging.getLogger(__name__)
+91 -32
View File
@@ -1,4 +1,9 @@
"""Managed Modal environment backed by tool-gateway."""
"""Managed Modal environment backed by tool-gateway.
Deliberately overrides :meth:`BaseEnvironment.execute`: the tool-gateway does command
preparation, CWD tracking and env-snapshot management server-side, so the base
``_wrap_command`` / ``_wait_for_process`` / snapshot machinery does not apply.
"""
from __future__ import annotations
@@ -6,15 +11,22 @@ import json
import logging
import os
import requests
import shlex
import time
import uuid
from typing import Any, Dict, Optional
from tools.environments.modal_utils import BaseModalExecutionEnvironment, ModalExecStart, PreparedModalExec
from tools.environments.base import BaseEnvironment
from tools.interrupt import is_interrupted
from tools.managed_tool_gateway import resolve_managed_tool_gateway
logger = logging.getLogger(__name__)
_TERMINAL_EXEC_STATUSES = frozenset({"completed", "failed", "cancelled", "timeout"})
_POLL_INTERVAL_SECONDS = 0.25
_CLIENT_TIMEOUT_GRACE_SECONDS = 10.0
_INTERRUPT_OUTPUT = "[Command interrupted - Modal sandbox exec cancelled]"
_ERROR_PREFIX = "Managed Modal exec failed"
def _request_timeout_env(name: str, default: float) -> float:
@@ -25,18 +37,21 @@ def _request_timeout_env(name: str, default: float) -> float:
return default
class ManagedModalEnvironment(BaseModalExecutionEnvironment):
"""Gateway-owned Modal sandbox with Hermes-compatible execute/cleanup.
def _result(output: str, returncode: int) -> dict:
return {"output": output, "returncode": returncode}
The exec handle passed between ``_start_modal_exec`` / ``_poll_modal_exec`` /
``_cancel_modal_exec`` is the gateway exec id string."""
def _error_result(output: str) -> dict:
return _result(output, 1)
class ManagedModalEnvironment(BaseEnvironment):
"""Gateway-owned Modal sandbox with Hermes-compatible execute/cleanup."""
_stdin_mode = "payload"
_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)
_CANCEL_READ_TIMEOUT_SECONDS = _request_timeout_env("TERMINAL_MANAGED_MODAL_CANCEL_READ_TIMEOUT_SECONDS", 5.0)
_client_timeout_grace_seconds = 10.0
_interrupt_output = "[Command interrupted - Modal sandbox exec cancelled]"
_unexpected_error_prefix = "Managed Modal exec failed"
def __init__(self, image: str, cwd: str = "/root", timeout: int = 60,
modal_sandbox_kwargs: Optional[Dict[str, Any]] = None,
@@ -65,57 +80,99 @@ class ManagedModalEnvironment(BaseModalExecutionEnvironment):
self._create_idempotency_key = str(uuid.uuid4())
self._sandbox_id = self._create_sandbox()
# -- Execution ------------------------------------------------------
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 gateway 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
exec_command, sudo_stdin = self._prepare_command(command)
if sudo_stdin is not None:
# 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}"
timeout = timeout or self.timeout
try:
exec_id, immediate = self._start_exec(exec_command, cwd or self.cwd, timeout, stdin_data)
except Exception as exc:
return _error_result(f"{_ERROR_PREFIX}: {exc}")
if immediate is not None:
return immediate
deadline = time.monotonic() + timeout + _CLIENT_TIMEOUT_GRACE_SECONDS
_now = time.monotonic()
_activity_state = {"last_touch": _now, "start": _now}
while True:
if is_interrupted():
self._cancel_exec(exec_id)
return _result(_INTERRUPT_OUTPUT, 130)
try:
result = self._poll_exec(exec_id)
except Exception as exc:
return _error_result(f"{_ERROR_PREFIX}: {exc}")
if result is not None:
return result
if time.monotonic() >= deadline:
self._cancel_exec(exec_id)
return _result(f"Managed Modal exec timed out after {timeout}s", 124)
# Periodic activity touch so the gateway knows we're alive (lazy import:
# tests stub tools.environments.base with only BaseEnvironment)
try:
from tools.environments.base import touch_activity_if_due
touch_activity_if_due(_activity_state, "modal command running")
except Exception:
pass
time.sleep(_POLL_INTERVAL_SECONDS)
def _result_from_body(self, body: dict) -> dict | None:
"""Final result dict if the exec body reports a terminal status, else ``None``."""
if body.get("status") in _TERMINAL_EXEC_STATUSES:
return self._result(body.get("output", ""), body.get("returncode", 1))
return _result(body.get("output", ""), body.get("returncode", 1))
return None
def _start_modal_exec(self, prepared: PreparedModalExec) -> ModalExecStart:
def _start_exec(self, command: str, cwd: str, timeout: int,
stdin_data: str | None) -> tuple[str, dict | None]:
"""POST the exec; return (exec_id, immediate_result). A non-None result ends execute()."""
exec_id = str(uuid.uuid4())
payload: Dict[str, Any] = {"execId": exec_id, "command": prepared.command, "cwd": prepared.cwd,
"timeoutMs": int(prepared.timeout * 1000)}
if prepared.stdin_data is not None:
payload["stdinData"] = prepared.stdin_data
payload: Dict[str, Any] = {"execId": exec_id, "command": command, "cwd": cwd,
"timeoutMs": int(timeout * 1000)}
if stdin_data is not None:
payload["stdinData"] = stdin_data
try:
response = self._request("POST", f"/v1/sandboxes/{self._sandbox_id}/execs", json=payload, timeout=10)
except Exception as exc:
return ModalExecStart(immediate_result=self._error_result(f"Managed Modal exec failed: {exc}"))
return exec_id, _error_result(f"Managed Modal exec failed: {exc}")
if response.status_code >= 400:
return ModalExecStart(immediate_result=self._error_result(
self._format_error("Managed Modal exec failed", response)))
return exec_id, _error_result(self._format_error("Managed Modal exec failed", response))
body = response.json()
final = self._result_from_body(body)
if final is not None:
return ModalExecStart(immediate_result=final)
return exec_id, final
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=exec_id)
return exec_id, _error_result("Managed Modal exec start did not return the expected exec id")
return exec_id, None
def _poll_modal_exec(self, handle: str) -> dict | None:
def _poll_exec(self, exec_id: str) -> dict | None:
try:
status_response = self._request(
"GET", f"/v1/sandboxes/{self._sandbox_id}/execs/{handle}",
"GET", f"/v1/sandboxes/{self._sandbox_id}/execs/{exec_id}",
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}")
return _error_result(f"Managed Modal exec poll failed: {exc}")
if status_response.status_code == 404:
return self._error_result("Managed Modal exec not found")
return _error_result("Managed Modal exec not found")
if status_response.status_code >= 400:
return self._error_result(self._format_error("Managed Modal exec poll failed", status_response))
return _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: str) -> None:
def _cancel_exec(self, exec_id: str) -> None:
try:
self._request("POST", f"/v1/sandboxes/{self._sandbox_id}/execs/{handle}/cancel",
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)
def _timeout_result_for_modal(self, timeout: int) -> dict:
return self._result(f"Managed Modal exec timed out after {timeout}s", 124)
def cleanup(self):
if not getattr(self, "_sandbox_id", None):
return
@@ -127,6 +184,8 @@ class ManagedModalEnvironment(BaseModalExecutionEnvironment):
finally:
self._sandbox_id = None
# -- Gateway HTTP ---------------------------------------------------
def _create_sandbox(self) -> str:
kw = self._sandbox_kwargs
cpu = self._coerce_number(kw.get("cpu"), 1)
+2 -4
View File
@@ -15,8 +15,7 @@ from typing import Any, Optional
from hermes_constants import get_hermes_home
from tools.environments.base import BaseEnvironment, _ThreadedProcessHandle, _load_json_store, _save_json_store
from tools.environments.file_sync import (
FileSyncManager, iter_sync_files, quoted_mkdir_command, quoted_rm_command, unique_parent_dirs,
)
FileSyncManager, iter_sync_files, quoted_mkdir_command, quoted_rm_command, unique_parent_dirs)
from tools.environments.remote_common import bash_argv, ensure_lazy_dep
logger = logging.getLogger(__name__)
@@ -71,8 +70,7 @@ def _resolve_modal_image(image_spec: Any) -> Any:
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",
]
"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")
-133
View File
@@ -1,133 +0,0 @@
"""Shared Hermes-side execution flow for Modal transports.
Stops at the Hermes boundary: command preparation, cwd/timeout normalization,
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
from abc import abstractmethod
from dataclasses import dataclass
from typing import Any
from tools.environments.base import BaseEnvironment
from tools.interrupt import is_interrupted
@dataclass(frozen=True)
class PreparedModalExec:
"""Normalized command data passed to a transport-specific exec runner."""
command: str
cwd: str
timeout: int
stdin_data: str | None = None
@dataclass(frozen=True)
class ModalExecStart:
"""Transport response after starting an exec."""
handle: Any | None = None
immediate_result: dict | None = None
class BaseModalExecutionEnvironment(BaseEnvironment):
"""Execution flow for the *managed* Modal transport (gateway-owned sandbox).
Deliberately overrides :meth:`BaseEnvironment.execute`: the tool-gateway does command
preparation, CWD tracking and env-snapshot management server-side, so the base
``_wrap_command`` / ``_wait_for_process`` / snapshot machinery does not apply.
See ``ManagedModalEnvironment`` for the concrete subclass."""
_stdin_mode = "payload"
_poll_interval_seconds = 0.25
_client_timeout_grace_seconds: float | None = None
_interrupt_output = "[Command interrupted]"
_unexpected_error_prefix = "Modal execution error"
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).
del rewrite_compound_background, bounded_capture
self._before_execute()
prepared = self._prepare_modal_exec(command, cwd=cwd, timeout=timeout, stdin_data=stdin_data)
try:
start = self._start_modal_exec(prepared)
except Exception as exc:
return self._error_result(f"{self._unexpected_error_prefix}: {exc}")
if start.immediate_result is not None:
return start.immediate_result
if start.handle is None:
return self._error_result(f"{self._unexpected_error_prefix}: transport did not return an exec handle")
deadline = None
if self._client_timeout_grace_seconds is not None:
deadline = time.monotonic() + prepared.timeout + self._client_timeout_grace_seconds
_now = time.monotonic()
_activity_state = {"last_touch": _now, "start": _now}
while True:
if is_interrupted():
self._cancel_quietly(start.handle)
return self._result(self._interrupt_output, 130)
try:
result = self._poll_modal_exec(start.handle)
except Exception as exc:
return self._error_result(f"{self._unexpected_error_prefix}: {exc}")
if result is not None:
return result
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:
from tools.environments.base import touch_activity_if_due
touch_activity_if_due(_activity_state, "modal command running")
except Exception:
pass
time.sleep(self._poll_interval_seconds)
def _cancel_quietly(self, handle: Any) -> None:
try:
self._cancel_modal_exec(handle)
except Exception:
pass
def _before_execute(self) -> None:
"""Hook for backends that need pre-exec sync or validation."""
def _prepare_modal_exec(self, command: str, *, cwd: str = "", timeout: int | None = None,
stdin_data: str | None = None) -> PreparedModalExec:
exec_command, sudo_stdin = self._prepare_command(command)
if sudo_stdin is not None:
# 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=stdin_data)
def _result(self, output: str, returncode: int) -> dict:
return {"output": output, "returncode": returncode}
def _error_result(self, output: str) -> dict:
return self._result(output, 1)
def _timeout_result_for_modal(self, timeout: int) -> dict:
return self._result(f"Command timed out after {timeout}s", 124)
@abstractmethod
def _start_modal_exec(self, prepared: PreparedModalExec) -> ModalExecStart:
"""Begin a transport-specific exec."""
@abstractmethod
def _poll_modal_exec(self, handle: Any) -> dict | None:
"""Return a final result dict when complete, else ``None``."""
@abstractmethod
def _cancel_modal_exec(self, handle: Any) -> None:
"""Cancel or terminate the active transport exec."""
+1 -2
View File
@@ -12,8 +12,7 @@ from pathlib import Path
from tools.environments.base import BaseEnvironment, EnvironmentConnectionError, _popen_bash
from tools.environments.file_sync import (
FileSyncManager, iter_sync_files, quoted_mkdir_command, quoted_rm_command, unique_parent_dirs,
)
FileSyncManager, iter_sync_files, quoted_mkdir_command, quoted_rm_command, unique_parent_dirs)
from tools.environments.remote_common import bash_argv, run_capture
logger = logging.getLogger(__name__)
+2 -4
View File
@@ -156,8 +156,7 @@ class VercelSandboxEnvironment(BaseEnvironment):
if disk not in {0, _DEFAULT_CONTAINER_DISK_MB}:
raise ValueError(
"Vercel Sandbox does not support configurable container_disk. "
"Use the default shared setting."
)
"Use the default shared setting.")
self._persistent = persistent_filesystem
self._task_id = task_id
self._requested_cwd = cwd
@@ -173,8 +172,7 @@ class VercelSandboxEnvironment(BaseEnvironment):
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,
}
"runtime": runtime or None, "resources": resources}
self._attach_fresh_sandbox(cwd)
self._sync_manager.sync(force=True)
self.init_session()