Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion src/agent_env/env/gateway/triggers.py
Original file line number Diff line number Diff line change
Expand Up @@ -538,12 +538,17 @@ async def _drive_time(self) -> None:
logger.exception("clock time-driver tick failed")

async def stop_driver(self) -> None:
"""Cancel the background poller on gateway teardown (idempotent)."""
"""Cancel the poller and drain trigger work before gateway clients close."""
task, self._driver_task = self._driver_task, None
if task is not None:
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
while self._tasks:
tasks = tuple(self._tasks)
for pending in tasks:
pending.cancel()
await asyncio.gather(*tasks, return_exceptions=True)

def _resolve_watch_roles(self, body: dict) -> set[str]:
raw = body.get("watch_roles")
Expand Down
92 changes: 73 additions & 19 deletions src/agent_env/providers/sandbox_providers/sandbox.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,8 @@
# the partial bytes already consumed, corrupting the stream. Pipe consumers must
# download to a temp file first, then read the file (see load_docker_images).
CURL_RETRY_FLAGS = "--retry 5 --retry-all-errors --retry-delay 1"
_DOCKER_IMAGE_INSPECT_FORMAT = "{{.Id}}"
_DOCKER_IMAGE_LOAD_EXEC_RETRIES = 2
# Docker label on what a step starts on a sandbox's Docker host (containers, images, networks), valued with the
# sandbox id, so a sandbox that shares its host (the local one) can remove its own when it terminates.
SANDBOX_LABEL = "agentenv.sandbox"
Expand Down Expand Up @@ -134,6 +136,21 @@ def from_dict(cls, data: dict) -> "NetworkPolicy":
)


def _object_file_method(legacy: Callable) -> Callable:
async def method(self, object_url: str, destination_path: str) -> None:
await legacy(self, object_url, destination_path)

return method


def _legacy_file_method(canonical: Callable, old_symbol: str, new_name: str) -> Callable:
async def method(self, s3_url: str, destination_path: str) -> None:
warn_deprecated(old_symbol, new_name, kind="method")
await canonical(self, s3_url, destination_path)

return method


class Sandbox(ABC):
"""Universal sandbox contract — anything that can host a process and expose ports."""

Expand All @@ -150,6 +167,20 @@ class Sandbox(ABC):
_VM_READY_TIMEOUT = 1200 # wait_for_vm wall-clock budget (s)
_VM_READY_POLL_INTERVAL = 30 # sparse polling (s)

def __init_subclass__(cls, **kwargs) -> None:
"""Keep overrides of either file-method spelling in the dispatch path, including super() calls."""
super().__init_subclass__(**kwargs)
for legacy, neutral, owner in (
("write_file_from_s3", "write_file_from_object", "Sandbox"),
("load_s3_file", "load_object_file", "VmSandbox"),
):
legacy_impl = cls.__dict__.get(legacy)
neutral_impl = cls.__dict__.get(neutral)
if legacy_impl is not None and neutral_impl is None:
setattr(cls, neutral, _object_file_method(legacy_impl))
elif neutral_impl is not None and legacy_impl is None:
setattr(cls, legacy, _legacy_file_method(neutral_impl, f"{owner}.{legacy}", neutral))

def host_port(self, port: int) -> int:
"""The host-side port a published container port is reachable on.

Expand Down Expand Up @@ -182,7 +213,7 @@ async def write_file_from_object(self, object_url: str, destination_path: str) -
async def write_file_from_s3(self, s3_url: str, destination_path: str) -> None:
"""Deprecated: ``write_file_from_object``."""
warn_deprecated("Sandbox.write_file_from_s3", "write_file_from_object", kind="method")
await self.write_file_from_object(s3_url, destination_path)
await Sandbox.write_file_from_object(self, s3_url, destination_path)

async def write_file_from_url(self, url: str, destination_path: str) -> None:
"""Download an HTTP(S) URL into the agent process's filesystem at destination_path."""
Expand Down Expand Up @@ -307,31 +338,51 @@ async def sign(artifact) -> str | None:

async def _load_docker_images(self, artifacts: list, signed_urls: list[str | None]) -> None:
load_commands = []
staging_paths = [f"/tmp/_docker_image_{self.sandbox_id}_{idx}.tar.gz" for idx in range(len(artifacts))]
for idx, (artifact, signed) in enumerate(zip(artifacts, signed_urls, strict=True)):
tmp_tar = f"/tmp/_docker_image_{self.sandbox_id}_{idx}.tar.gz"
tmp_tar = staging_paths[idx]
if signed is not None:
# Download to a file first (retry-safe with -o); a `curl | ... docker load`
# pipe can't be retried without corrupting the stream (curl won't rewind).
load_commands.append(
f'(curl -fsSL {CURL_RETRY_FLAGS} "{signed}" -o {shlex.quote(tmp_tar)} '
f"&& gunzip -c {shlex.quote(tmp_tar)} | docker load && rm -f {shlex.quote(tmp_tar)})"
f'(curl -fsSL {CURL_RETRY_FLAGS} {shlex.quote(signed)} -o {shlex.quote(tmp_tar)} '
f"&& gunzip -c {shlex.quote(tmp_tar)} | docker load)"
)
else:
await self._download_object_to_vm(artifact.tar_gz_object_url, tmp_tar)
load_commands.append(
f"(gunzip -c {shlex.quote(tmp_tar)} | docker load && rm -f {shlex.quote(tmp_tar)})"
f"(gunzip -c {shlex.quote(tmp_tar)} | docker load)"
)
logger.info(f" Queued: {artifact.image_name}")
await self.exec_script(" & ".join(load_commands) + " & wait", max_retries=2)
try:
for idx, (artifact, signed) in enumerate(zip(artifacts, signed_urls, strict=True)):
if signed is None:
await self._download_object_to_vm(artifact.tar_gz_object_url, staging_paths[idx])
# exec_script explicitly invokes `bash -c`; pipefail and wait-status
# collection therefore use the shell the provider actually guarantees.
workers = "\n".join(
f"{command} & pid_{idx}=$!" for idx, command in enumerate(load_commands)
)
waits = "\n".join(f"wait $pid_{idx} || status=1" for idx in range(len(load_commands)))
script = (
"set -o pipefail\n"
f"{workers}\n"
"status=0\n"
f"{waits}\n"
"exit $status"
)
# Signed archives are downloaded per attempt; keep unsigned archives staged until retries finish.
await self.exec_script(script, max_retries=_DOCKER_IMAGE_LOAD_EXEC_RETRIES)
finally:
await self._remove_vm_temp_file(*staging_paths)

logger.info("Verifying Docker images...")
exit_code, stdout, stderr = await self.exec_with_output("sudo", "docker", "images")
image_refs = [artifact.image_name for artifact in artifacts]
exit_code, _, stderr = await self.exec_with_output(
"sudo", "docker", "image", "inspect", "--format", _DOCKER_IMAGE_INSPECT_FORMAT,
*image_refs,
)
if exit_code != 0:
raise RuntimeError(f"docker images failed: {stderr}")
for artifact in artifacts:
base_name = artifact.image_name.split(":")[0]
if base_name not in stdout:
raise RuntimeError(f"{artifact.image_name} image not found. stdout: {stdout}")
raise RuntimeError(f"Docker image reference verification failed for {image_refs}: {stderr}")
logger.info(" All images loaded successfully")

async def load_object_file(self, object_url: str, destination_path: str) -> None:
Expand All @@ -341,7 +392,7 @@ async def load_object_file(self, object_url: str, destination_path: str) -> None
async def load_s3_file(self, s3_url: str, destination_path: str) -> None:
"""Deprecated: ``load_object_file``."""
warn_deprecated("VmSandbox.load_s3_file", "load_object_file", kind="method")
await self.load_object_file(s3_url, destination_path)
await VmSandbox.load_object_file(self, s3_url, destination_path)

async def _download_object_to_vm(self, object_url: str, vm_path: str) -> None:
"""Place object_url onto the VM host at vm_path, backend-agnostically."""
Expand Down Expand Up @@ -422,10 +473,13 @@ async def _write_bytes_to_vm_path(self, data: bytes, vm_path: str) -> None:
return
# Too big for one heredoc arg: append the (shell-safe) base64 in bounded chunks.
vm_b64 = f"{vm_path}.b64"
await self.exec_script(f": > {shlex.quote(vm_b64)}")
for i in range(0, len(encoded), self._WFT_CHUNK_BYTES):
await self.exec_script(f"printf '%s' {shlex.quote(encoded[i:i + self._WFT_CHUNK_BYTES])} >> {shlex.quote(vm_b64)}")
await self.exec_script(f"base64 -d {shlex.quote(vm_b64)} > {shlex.quote(vm_path)} && rm -f {shlex.quote(vm_b64)}")
try:
await self.exec_script(f": > {shlex.quote(vm_b64)}")
for i in range(0, len(encoded), self._WFT_CHUNK_BYTES):
await self.exec_script(f"printf '%s' {shlex.quote(encoded[i:i + self._WFT_CHUNK_BYTES])} >> {shlex.quote(vm_b64)}")
await self.exec_script(f"base64 -d {shlex.quote(vm_b64)} > {shlex.quote(vm_path)}")
finally:
await self._remove_vm_temp_file(vm_b64)

async def write_host_file(self, data: bytes, vm_path: str) -> None:
"""Write bytes to vm_path on the VM host itself, not into the agent container."""
Expand All @@ -440,7 +494,7 @@ async def write_file_from_text(self, content: str, destination_path: str) -> Non
await self._write_bytes_to_vm_path(content.encode(), vm_path)
await self._copy_into_container(vm_path, destination_path)
finally:
await self._remove_vm_temp_file(vm_path, f"{vm_path}.b64")
await self._remove_vm_temp_file(vm_path)


def port_bindings(host_ips: Iterable[str], host_port: int, container_port: int) -> list[str]:
Expand Down
33 changes: 33 additions & 0 deletions tst/unit/env/gateway/triggers_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,39 @@ def _registration(**overrides):
return base


@pytest.mark.asyncio
async def test_stop_driver_cancels_and_awaits_pending_action_before_client_teardown(engine):
started = asyncio.Event()
cancelled = asyncio.Event()

async def slow_call(_tool_name, _arguments):
started.set()
try:
await asyncio.Event().wait()
except asyncio.CancelledError:
cancelled.set()
raise

engine._internal_call = slow_call
engine.register({"watch_roles": ["default"], "triggers": [
{"id": "shutdown", "when": {"type": "action", "tool": "provoking_tool"},
"actions": [{"type": "tool", "tool": "slack_send_message", "args": {}}]}
]})
engine.start_driver()
engine.on_tool_call("default", "provoking_tool", {}, _result())
await asyncio.wait_for(started.wait(), 1)
tracked = tuple(engine._tasks)
assert tracked

await engine.stop_driver()

assert cancelled.is_set()
assert all(task.done() for task in tracked)
assert not engine._tasks
await engine.stop_driver()
assert not engine._tasks


@pytest.mark.parametrize("bad,fragment", [
({"triggers": "nope"}, "triggers must be a list"),
({"triggers": [{"id": "x", "when": {"type": "step"}, "actions": []}]}, "when.type"),
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
from pathlib import Path
from types import SimpleNamespace

import pytest

from agent_env.providers.sandbox_providers.sandbox import Sandbox, VmSandbox, stage_files_into_container


_METHODS = (
(Sandbox, "write_file_from_s3", "write_file_from_object"),
(VmSandbox, "write_file_from_s3", "write_file_from_object"),
(VmSandbox, "load_s3_file", "load_object_file"),
)


async def _terminate(self):
pass


@pytest.mark.asyncio
@pytest.mark.parametrize("base, legacy, neutral", _METHODS)
async def test_new_calls_reach_inherited_legacy_provider_overrides(base, legacy, neutral):
calls = []

async def implementation(self, s3_url, destination_path):
calls.append((s3_url, destination_path))

provider = type("LegacySandbox", (base,), {legacy: implementation, "terminate": _terminate})
child = type("InheritedSandbox", (provider,), {})

await getattr(child(), neutral)(object_url="file:///objects/input", destination_path="/app/input")

assert calls == [("file:///objects/input", "/app/input")]


@pytest.mark.asyncio
@pytest.mark.parametrize("base, legacy, neutral", _METHODS)
async def test_old_calls_reach_modern_overrides_of_legacy_providers(base, legacy, neutral):
calls = []

async def old(self, s3_url, destination_path):
raise AssertionError("the parent provider was bypassed by the modern override")

async def new(self, object_url, destination_path):
calls.append((object_url, destination_path))

parent = type("LegacySandbox", (base,), {legacy: old, "terminate": _terminate})
child = type("ModernSandbox", (parent,), {neutral: new})

with pytest.warns(DeprecationWarning, match=f"use {neutral}"):
await getattr(child(), legacy)(s3_url="file:///objects/input", destination_path="/app/input")

assert calls == [("file:///objects/input", "/app/input")]


@pytest.mark.asyncio
async def test_artifact_staging_uses_an_older_container_provider(tmp_path):
class LegacySandbox(Sandbox):
async def terminate(self):
pass

async def exec(self, *command):
Path(command[-1]).mkdir(parents=True, exist_ok=True)

async def write_file_from_s3(self, s3_url, destination_path):
Path(destination_path).write_bytes(b"artifact payload")

artifacts = {"input.json": SimpleNamespace(object_url="file:///objects/input")}

loaded = await stage_files_into_container(LegacySandbox(), artifacts, str(tmp_path))

assert loaded == {"input.json": str(tmp_path / "input.json")}
assert (tmp_path / "input.json").read_bytes() == b"artifact payload"


@pytest.mark.asyncio
async def test_legacy_vm_overrides_can_delegate_to_super(monkeypatch):
calls = []

class LegacyVm(VmSandbox):
async def terminate(self):
pass

async def load_s3_file(self, s3_url, destination_path):
calls.append("load override")
await super().load_s3_file(s3_url, destination_path)

async def write_file_from_s3(self, s3_url, destination_path):
calls.append("write override")
await super().write_file_from_s3(s3_url, destination_path)

async def _download_object_to_vm(self, object_url, vm_path):
calls.append((object_url, vm_path))

async def _copy_into_container(self, vm_path, destination_path):
calls.append((vm_path, destination_path))

async def _remove_vm_temp_file(self, *vm_paths):
calls.append(vm_paths)

monkeypatch.setattr(LegacyVm, "_staging_path", staticmethod(lambda kind, dest: "/tmp/staged"))

with pytest.warns(DeprecationWarning):
await LegacyVm().write_file_from_object("file:///objects/input", "/app/input")

assert calls == [
"write override", "load override", ("file:///objects/input", "/tmp/staged"),
("/tmp/staged", "/app/input"), ("/tmp/staged",),
]
Loading
Loading