From c8492e2adeb073ababb07fe03fce66dca95a1e8b Mon Sep 17 00:00:00 2001 From: morluto <76467478+morluto@users.noreply.github.com> Date: Wed, 7 Oct 2026 09:29:03 +0800 Subject: [PATCH 1/4] fix(lifecycle): drain gateway work and verify VM image loads --- src/agent_env/env/gateway/triggers.py | 7 +- .../providers/sandbox_providers/sandbox.py | 50 +++-- tst/unit/env/gateway/triggers_test.py | 33 ++++ .../sandbox_providers/vm_sandbox_test.py | 172 ++++++++++++++++-- 4 files changed, 238 insertions(+), 24 deletions(-) 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 d89f8b09..eff555e5 100644 --- a/src/agent_env/providers/sandbox_providers/sandbox.py +++ b/src/agent_env/providers/sandbox_providers/sandbox.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import contextlib import logging import os import posixpath @@ -40,6 +41,7 @@ # 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 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" @@ -281,31 +283,55 @@ 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))) + cleanup_command = "rm -f " + " ".join(shlex.quote(path) for path in staging_paths) + script = ( + f"trap {shlex.quote(cleanup_command)} EXIT\n" + "set -o pipefail\n" + f"{workers}\n" + "status=0\n" + f"{waits}\n" + "exit $status" + ) + await self.exec_script(script) + except BaseException: + # Download failures happen before the shell trap can be installed. + with contextlib.suppress(Exception): + await self._remove_vm_temp_file(*staging_paths) + raise 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_s3_file(self, s3_url: str, destination_path: str) -> None: 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/vm_sandbox_test.py b/tst/unit/providers/sandbox_providers/vm_sandbox_test.py index 75ae6d81..7b3d4e9f 100644 --- a/tst/unit/providers/sandbox_providers/vm_sandbox_test.py +++ b/tst/unit/providers/sandbox_providers/vm_sandbox_test.py @@ -1,4 +1,7 @@ import asyncio +import os +import subprocess +from pathlib import Path import threading import time from concurrent.futures import ThreadPoolExecutor @@ -33,10 +36,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 @@ -45,10 +47,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