diff --git a/src/agent_env/env/gateway/triggers.py b/src/agent_env/env/gateway/triggers.py index a2b07f77..823dba7f 100644 --- a/src/agent_env/env/gateway/triggers.py +++ b/src/agent_env/env/gateway/triggers.py @@ -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") diff --git a/src/agent_env/providers/sandbox_providers/sandbox.py b/src/agent_env/providers/sandbox_providers/sandbox.py index 86985046..e143efad 100644 --- a/src/agent_env/providers/sandbox_providers/sandbox.py +++ b/src/agent_env/providers/sandbox_providers/sandbox.py @@ -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" @@ -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.""" @@ -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. @@ -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.""" @@ -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: @@ -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.""" @@ -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.""" @@ -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]: diff --git a/tst/unit/env/gateway/triggers_test.py b/tst/unit/env/gateway/triggers_test.py index 7f16d956..41ebe35f 100644 --- a/tst/unit/env/gateway/triggers_test.py +++ b/tst/unit/env/gateway/triggers_test.py @@ -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"), diff --git a/tst/unit/providers/sandbox_providers/sandbox_method_compatibility_test.py b/tst/unit/providers/sandbox_providers/sandbox_method_compatibility_test.py new file mode 100644 index 00000000..daffb84d --- /dev/null +++ b/tst/unit/providers/sandbox_providers/sandbox_method_compatibility_test.py @@ -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",), + ] diff --git a/tst/unit/providers/sandbox_providers/vm_sandbox_test.py b/tst/unit/providers/sandbox_providers/vm_sandbox_test.py index 511ea945..04017b4c 100644 --- a/tst/unit/providers/sandbox_providers/vm_sandbox_test.py +++ b/tst/unit/providers/sandbox_providers/vm_sandbox_test.py @@ -2,9 +2,12 @@ import hashlib import io import logging +import os +import subprocess import threading import time from concurrent.futures import ThreadPoolExecutor +from pathlib import Path from types import SimpleNamespace import pytest @@ -37,10 +40,9 @@ def signing_store(): class _RecordingVmSandbox(VmSandbox): """Concrete VmSandbox that records executed scripts instead of touching a VM.""" - def __init__(self, images_stdout: str = ""): + def __init__(self): self.sandbox_id = "vm-test" self.scripts: list[str] = [] - self._images_stdout = images_stdout async def terminate(self) -> None: # pragma: no cover - not exercised pass @@ -49,10 +51,9 @@ async def exec(self, *command): # pragma: no cover - not exercised return None async def exec_with_output(self, *args): - # load_docker_images runs `sudo docker images` to verify; everything - # else flows through exec_script as `sudo bash -c