diff --git a/src/agent_env/cli/eval/run.py b/src/agent_env/cli/eval/run.py index b47b20bb..89353d9a 100644 --- a/src/agent_env/cli/eval/run.py +++ b/src/agent_env/cli/eval/run.py @@ -11,9 +11,10 @@ from agent_env.cli.banner import print_banner from agent_env.cli.identity import get_agent_env_client_id +from agent_env.cli.teardown_output import echo_teardown from agent_env.store.ids import fs_safe from agent_env.task.interrupts import Interrupts -from agent_env.task.teardown import TeardownReport, kind, teardown_run +from agent_env.task.teardown import teardown_run from agent_env.task_step.context import TaskStepContext @@ -52,16 +53,6 @@ def _write_context(context, task_id, output_dir, prefix=""): click.echo(click.style(f"{prefix}Output written to: {output_path}", fg="blue")) -def _echo_teardown(tag: str, report: TeardownReport) -> None: - if report.terminated: - n = len(report.terminated) - click.echo(click.style(f"{tag} Tore down {n} sandbox{'es' if n != 1 else ''}", fg="blue")) - for sandbox, why in report.failed: - click.echo(click.style(f"{tag} Couldn't tear down {sandbox.sandbox_id}: {why}", fg="red")) - for sandbox in report.left: - click.echo(click.style(f"{tag} Still up: {sandbox.sandbox_id} ({kind(sandbox)})", fg="red")) - - async def _run_single(task, tag, output_dir, agent_model=None, agent_artifact_id=None, base_metadata=None): """Execute a single task run with logging callbacks, then tear down what it deployed, even when it raised or was cancelled.""" @@ -76,7 +67,7 @@ async def _run_single(task, tag, output_dir, agent_model=None, agent_artifact_id context=context, ) finally: - _echo_teardown(tag, await teardown_run(context)) + echo_teardown(tag, await teardown_run(context)) if output_dir: _write_context(context, task.id, output_dir, prefix=f"{tag} ") return task.id, tag, context diff --git a/src/agent_env/cli/task/run.py b/src/agent_env/cli/task/run.py index 057454c0..6a9eb702 100644 --- a/src/agent_env/cli/task/run.py +++ b/src/agent_env/cli/task/run.py @@ -6,13 +6,18 @@ import json import os import textwrap +import time from uuid import uuid4 import click from agent_env.cli.banner import print_banner from agent_env.cli.identity import get_agent_env_client_id +from agent_env.cli.teardown_output import echo_teardown +from agent_env.providers.sandbox_providers.local_sandbox import LocalSandbox from agent_env.store.ids import derive_id, fs_safe, is_local_id, validate_local_id +from agent_env.task.interrupts import Interrupts +from agent_env.task.teardown import TeardownReport, kind, teardown_run from agent_env.task_step.task_steps.collect_artifacts import CollectArtifactsTaskStep @@ -312,16 +317,93 @@ def _stamp_agent_env_client_metadata(context_metadata: dict, client_id: str | No context_metadata.setdefault("agent_env_hub", {}).setdefault("caller", client_id) -async def _run_single(task, run_index, task_id, output_dir, agent_model=None, agent_artifact_id=None, start_step=0, context=None): - """Execute a single parallel run.""" - tag = f"[run {run_index}] " - on_start, on_complete = _make_parallel_callbacks(run_index) - context = await task.run(on_step_start=on_start, on_step_complete=on_complete, agent_model=agent_model, agent_artifact_id=agent_artifact_id, start_step=start_step, context=context) +async def _run_and_settle(task, tag, context, *, keep, output_dir, output_name, on_start, on_complete, **run_kwargs): + """Run ``task`` in ``context``, then tear down what it deployed, however it ended. ``keep`` holds what a run + left up, whether it passed or failed, but never what Ctrl-C cancelled.""" + try: + await task.run(on_step_start=on_start, on_step_complete=on_complete, context=context, **run_kwargs) + except asyncio.CancelledError: + echo_teardown(tag, await teardown_run(context)) + raise + except Exception: + if not keep: + echo_teardown(tag, await teardown_run(context)) + elif output_dir: # what it left up is kept, so its context is what resumes or inspects it + try: + _write_context(context, output_name, output_dir, prefix=f"{tag} " if tag else "") + except OSError as e: # the run's own error is the one to report + click.echo(click.style(f"{f'{tag} ' if tag else ''}Couldn't write the run's context: {e}", fg="red")) + raise + if not keep: + echo_teardown(tag, await teardown_run(context)) if output_dir: - _write_context(context, task_id, output_dir, prefix=tag) + _write_context(context, output_name, output_dir, prefix=f"{tag} " if tag else "") return context +def _on_signal(count): + message = "Cancelling and tearing down (Ctrl-C again to stop now)" if count == 1 else "Stopping the teardown now" + click.echo(click.style(message, fg="yellow")) + + +def _outcomes(futures) -> list: + """Each run's context, or what it raised (a ``CancelledError`` for one Ctrl-C cancelled).""" + return [asyncio.CancelledError() if f.cancelled() else (f.exception() or f.result()) for f in futures] + + +def _kept(tagged_contexts) -> list: + """Each run that still has sandboxes up, with them: what ``--keep`` held, and nothing a teardown took down.""" + kept = [(tag, context, TeardownReport.skipped(context).left) for tag, context in tagged_contexts] + return [(tag, context, left) for tag, context, left in kept if left] + + +def _tear_down_kept(kept) -> None: + """Tear down what the runs kept; a Ctrl-C stops it. Counts signals afresh, so the one that got here doesn't.""" + count = sum(len(left) for _, _, left in kept) + click.echo(f"\nTearing down {count} {'sandbox' if count == 1 else 'sandboxes'} (Ctrl-C again to stop now)") + with Interrupts() as stopper: + try: + reports = stopper.run(stopper.stopping(teardown_run(context) for _, context, _ in kept)) + except asyncio.CancelledError: + click.echo(click.style("Stopped the teardown; what it hadn't reached is still up.", fg="yellow")) + return + for (tag, _, _), report in zip(kept, reports): + echo_teardown(tag, report) + + +def _settle_kept(tagged_contexts, interrupts: Interrupts, interrupted) -> None: + """``--keep``'s end: hold what the runs left up until Ctrl-C, unless Ctrl-C already ended the runs, which tears + down what the ones that finished first kept.""" + if interrupted is None: + _hold(tagged_contexts, interrupts) + elif kept := _kept(tagged_contexts): + _tear_down_kept(kept) + + +_HOLD_POLL_SECONDS = 0.2 # how often a hold looks for the Ctrl-C that ends it + + +def _hold(tagged_contexts, interrupts: Interrupts) -> None: + """Print what ``--keep`` left up, wait for Ctrl-C or SIGTERM, then tear it down; another one stops that.""" + kept = _kept(tagged_contexts) + count = sum(len(left) for _, _, left in kept) + if not count: + click.echo("\nNothing to keep up: no run left a sandbox.") + return + click.echo("\nKept up:") + for tag, _, left in kept: + for sandbox in left: + folder = LocalSandbox.find_work_dir(sandbox.sandbox_id) if sandbox.sandbox_type == "local" else None + click.echo(f" {f'{tag} ' if tag else ''}{sandbox.sandbox_id} {kind(sandbox)}{f' {folder}' if folder else ''}") + noun = "sandbox" if count == 1 else "sandboxes" + click.echo(f"\nHolding {count} {noun} up; Ctrl-C tears {'it' if count == 1 else 'them'} down.") + while not interrupts.count: + time.sleep(_HOLD_POLL_SECONDS) + if interrupts.count > 1: + return + _tear_down_kept(kept) + + @click.command() @click.option("--id", "task_id", required=True, help="Task id") @click.option("--version", "task_version", default=None, type=int, help="Task version (defaults to latest)") @@ -347,8 +429,16 @@ async def _run_single(task, run_index, task_id, output_dir, agent_model=None, ag help="Override the env state type for deploy_env steps") @click.option("--env-state-instance-id", default=None, type=str, help="Attach env to an existing EnvStateInstance") -def run(task_id: str, task_version: int | None, output_dir: str | None, k: int, agent_model: str | None, agent_artifact_id: str | None, start_step: int, context_json: str | None, litellm_api_key: str | None, judge_litellm_api_key: str | None, apply_trajectory_filter: bool | None, a2a_agent_id: str | None, agent_sandbox: str | None, env_sandbox: str | None, gateway_env_id: str | None, service_db_env_id: str | None, env_state_type: str | None, env_state_instance_id: str | None): - """Run a task by executing its steps sequentially.""" +@click.option("--keep", is_flag=True, + help="Keep what a run deployed up when it ends, print it, and tear it down on Ctrl-C; while it holds, " + "another terminal can resume the run with --start-step and --context-json. Without it, each run " + "is torn down as it ends, as Ctrl-C mid-run always does.") +def run(task_id: str, task_version: int | None, output_dir: str | None, k: int, agent_model: str | None, agent_artifact_id: str | None, start_step: int, context_json: str | None, litellm_api_key: str | None, judge_litellm_api_key: str | None, apply_trajectory_filter: bool | None, a2a_agent_id: str | None, agent_sandbox: str | None, env_sandbox: str | None, gateway_env_id: str | None, service_db_env_id: str | None, env_state_type: str | None, env_state_instance_id: str | None, keep: bool): + """Run a task by executing its steps sequentially. + + Each run's sandboxes are torn down as it ends, passed or failed; its instance and context JSON stay. + --keep holds them up until Ctrl-C instead. Ctrl-C or SIGTERM mid-run cancels the runs, tears them down + and exits 130 (143 for SIGTERM); a second one stops the teardown.""" if k < 1: raise click.BadParameter("must be at least 1", param_hint="'--k'") @@ -397,44 +487,51 @@ def run(task_id: str, task_version: int | None, output_dir: str | None, k: int, print_banner() + run_kwargs = dict(agent_model=agent_model, agent_artifact_id=agent_artifact_id, start_step=start_step) if k > 1: # Parallel runs: compact output with [run N] prefixes click.echo(click.style(f"Running {k} task runs in parallel...", fg="blue")) click.echo() - - async def _run_all(): - coros = [_run_single(task, i, task_id, output_dir, agent_model=agent_model, agent_artifact_id=agent_artifact_id, start_step=start_step, context=copy.deepcopy(initial_context) if initial_context else None) for i in range(1, k + 1)] - return await asyncio.gather(*coros, return_exceptions=True) - - results = asyncio.run(_run_all()) - - click.echo() - failures = [(i, r) for i, r in enumerate(results, 1) if isinstance(r, BaseException)] - if failures: - for run_index, exc in failures: - click.echo(click.style(f"[run {run_index}] FAILED: {exc}", fg="red")) - click.echo(click.style( - f"{len(results) - len(failures)}/{k} runs completed, {len(failures)}/{k} failed.", - fg="red", - )) - raise SystemExit(1) - else: - click.echo(click.style(f"All {k} runs completed!", fg="blue")) + tagged = [(f"[run {i}]", copy.deepcopy(initial_context)) for i in range(1, k + 1)] + callbacks = [_make_parallel_callbacks(i) for i in range(1, k + 1)] else: - # Single run: verbose output - context = asyncio.run(task.run( - on_step_start=_log_step_start, - on_step_complete=_log_step_complete, - agent_model=agent_model, - agent_artifact_id=agent_artifact_id, - start_step=start_step, - context=initial_context, - )) - + tagged = [("", initial_context)] + callbacks = [(_log_step_start, _log_step_complete)] + + async def _run_all(interrupts): + runs = [ + _run_and_settle(task, tag, context, keep=keep, output_dir=output_dir, output_name=task_id, + on_start=on_start, on_complete=on_complete, **run_kwargs) + for (tag, context), (on_start, on_complete) in zip(tagged, callbacks) + ] + return _outcomes(await interrupts.gather(runs, _on_signal)) + + with Interrupts() as interrupts: + results = interrupts.run(_run_all(interrupts)) + interrupted = interrupts.signum if interrupts.count else None click.echo() - click.echo(click.style("Task completed!", fg="blue")) - - _write_context(context, task_id, output_dir) + failures = [(tag, r) for (tag, _), r in zip(tagged, results) if isinstance(r, BaseException)] + if k > 1: + for tag, exc in failures: + outcome = "CANCELLED" if isinstance(exc, asyncio.CancelledError) else f"FAILED: {exc}" + click.echo(click.style(f"{tag} {outcome}", fg="red")) + if failures: + click.echo(click.style( + f"{len(results) - len(failures)}/{k} runs completed, {len(failures)}/{k} failed.", fg="red")) + else: + click.echo(click.style(f"All {k} runs completed!", fg="blue")) + elif failures and interrupted is not None: + click.echo(click.style("Task cancelled.", fg="red")) + elif not failures: + click.echo(click.style("Task completed!", fg="blue")) + if keep: + _settle_kept(tagged, interrupts, interrupted) + if interrupted is not None: + raise SystemExit(128 + interrupted) + if failures and k == 1: + raise failures[0][1] # torn down already; the CLI reports it as it always has + if failures: + raise SystemExit(1) def _seed_universe_id(task, seed: dict) -> str | None: @@ -463,9 +560,15 @@ def _seed_universe_id(task, seed: dict) -> str | None: help="Override the env state type for deploy_env steps") @click.option("--env-state-instance-id", default=None, type=str, help="Attach env to an existing EnvStateInstance") -def run_batch(task_id: str, task_version: int | None, seeds: str, concurrency: int, output_dir: str | None, agent_model: str | None, agent_artifact_id: str | None, litellm_api_key: str | None, judge_litellm_api_key: str | None, apply_trajectory_filter: bool | None, agent_sandbox: str | None, env_sandbox: str | None, env_state_type: str | None, env_state_instance_id: str | None): +@click.option("--keep", is_flag=True, + help="Keep what each run deployed up when it ends, print it, and tear it down on Ctrl-C. Without it, " + "each run is torn down as it ends, as Ctrl-C mid-batch always does.") +def run_batch(task_id: str, task_version: int | None, seeds: str, concurrency: int, output_dir: str | None, agent_model: str | None, agent_artifact_id: str | None, litellm_api_key: str | None, judge_litellm_api_key: str | None, apply_trajectory_filter: bool | None, agent_sandbox: str | None, env_sandbox: str | None, env_state_type: str | None, env_state_instance_id: str | None, keep: bool): """Run a task in batch against multiple seeds from a CSV file. + Each run's sandboxes are torn down as it ends; --keep holds them up until Ctrl-C instead. Ctrl-C or + SIGTERM mid-batch cancels the runs, tears them down and exits 130 (143 for SIGTERM). + Each row in the CSV becomes a seed dict passed to the task via context.metadata["seed"]. Prompt templates with variables are rendered using the seed values before being sent to the agent. @@ -506,6 +609,8 @@ def run_batch(task_id: str, task_version: int | None, seeds: str, concurrency: i batch_run_group_id = uuid4().hex client_id = get_agent_env_client_id() + tagged: list[tuple[str, TaskStepContext]] = [] + async def _run_seed(index: int, seed: dict): async with sem: seed_name = seed.get("name", seed.get("repo_url", f"seed-{index}")) @@ -513,6 +618,7 @@ async def _run_seed(index: int, seed: dict): on_start, on_complete = _make_parallel_callbacks(f"seed {index}") ctx = TaskStepContext() + tagged.append((tag, ctx)) _stamp_agent_env_client_metadata(ctx.metadata, client_id) ctx.metadata["run_group_id"] = batch_run_group_id ctx.metadata["seed"] = seed @@ -536,22 +642,21 @@ async def _run_seed(index: int, seed: dict): ctx.metadata["task_id"] = task_id click.echo(click.style(f"{tag} Starting...", fg="yellow")) - context = await task.run( - on_step_start=on_start, - on_step_complete=on_complete, - agent_model=agent_model, - agent_artifact_id=agent_artifact_id, - context=ctx, + context = await _run_and_settle( + task, tag, ctx, keep=keep, output_dir=output_dir, output_name=f"{task_id}-seed{index}", + on_start=on_start, on_complete=on_complete, agent_model=agent_model, agent_artifact_id=agent_artifact_id, ) - if output_dir: - _write_context(context, f"{task_id}-seed{index}", output_dir, prefix=f"{tag} ") return index, seed_name, context - async def _run_all(): + async def _run_all(interrupts): coros = [_run_seed(i, seed) for i, seed in enumerate(seed_rows, 1)] - return await asyncio.gather(*coros, return_exceptions=True) + return _outcomes(await interrupts.gather(coros, _on_signal)) - results = asyncio.run(_run_all()) + with Interrupts() as interrupts: + results = interrupts.run(_run_all(interrupts)) + interrupted = interrupts.signum if interrupts.count else None + if keep: + _settle_kept(tagged, interrupts, interrupted) click.echo() successes = 0 @@ -564,8 +669,11 @@ async def _run_all(): if failures: for exc in failures: - click.echo(click.style(f"FAILED: {exc}", fg="red")) + click.echo(click.style("CANCELLED" if isinstance(exc, asyncio.CancelledError) else f"FAILED: {exc}", fg="red")) click.echo(click.style(f"{successes}/{len(seed_rows)} succeeded, {len(failures)}/{len(seed_rows)} failed.", fg="red")) - raise SystemExit(1) else: click.echo(click.style(f"All {len(seed_rows)} seeds completed!", fg="blue")) + if interrupted is not None: + raise SystemExit(128 + interrupted) + if failures: + raise SystemExit(1) diff --git a/src/agent_env/cli/teardown_output.py b/src/agent_env/cli/teardown_output.py new file mode 100644 index 00000000..8bc2b4cb --- /dev/null +++ b/src/agent_env/cli/teardown_output.py @@ -0,0 +1,16 @@ +"""How a command reports what tearing a run down did.""" + +import click + +from agent_env.task.teardown import TeardownReport, kind + + +def echo_teardown(tag: str, report: TeardownReport) -> None: + prefix = f"{tag} " if tag else "" + if report.terminated: + n = len(report.terminated) + click.echo(click.style(f"{prefix}Tore down {n} sandbox{'es' if n != 1 else ''}", fg="blue")) + for sandbox, why in report.failed: + click.echo(click.style(f"{prefix}Couldn't tear down {sandbox.sandbox_id}: {why}", fg="red")) + for sandbox in report.left: + click.echo(click.style(f"{prefix}Still up: {sandbox.sandbox_id} ({kind(sandbox)})", fg="red")) diff --git a/src/agent_env/runner/local_runner.py b/src/agent_env/runner/local_runner.py index f9c734b8..434e1a70 100644 --- a/src/agent_env/runner/local_runner.py +++ b/src/agent_env/runner/local_runner.py @@ -18,6 +18,7 @@ from agent_env.runner import store as run_store from agent_env.runner.runner import RunHandle, RunRecord, Runner, RunStatus +from agent_env.task.teardown import TERMINATE_TIMEOUT_SECONDS, teardown_run logger = logging.getLogger(__name__) @@ -26,6 +27,13 @@ class LocalRunner(Runner): """In-process runner. The DocumentStore holds run records for status/listing only.""" type = "local" + # How long a run's teardown may take, from when it starts: each sandbox gets TERMINATE_TIMEOUT_SECONDS, + # and the folder of a local one is removed after that, outside it. + TEARDOWN_WAIT_SECONDS = 2 * TERMINATE_TIMEOUT_SECONDS + # How long stop() gives a cancelled run to unwind its steps and reach its teardown. + STOP_WAIT_SECONDS = 60 + # After those waits, how long stop() gives a run it cancels again before shutting down without it. + ABANDON_WAIT_SECONDS = 1 def __init__(self, workers: int = 2) -> None: if workers < 1: @@ -33,6 +41,9 @@ def __init__(self, workers: int = 2) -> None: self.workers = int(workers) self._sem = asyncio.Semaphore(self.workers) self._inflight: dict[str, asyncio.Task] = {} + self._tearing_down: set[str] = set() # runs past their task, removing what it deployed + self._contexts: dict = {} # each run's context, so stop() can still tear down a run it gives up on + self._abandoned: set[str] = set() # runs stop() gave up on and tore down itself self._stopping = False # --- lifecycle --------------------------------------------------------- @@ -49,13 +60,47 @@ async def start(self) -> None: logger.info("Failed %d run(s) left non-terminal by a previous process", orphaned) async def stop(self) -> None: + """Cancel the runs still working and wait for every run, letting a teardown under way finish: each bounds + its own teardown by TEARDOWN_WAIT_SECONDS. A run still going after a run's whole allowance, a step that + ignored its cancel or a teardown that never ended, is cancelled again and left behind.""" self._stopping = True - for task in list(self._inflight.values()): - task.cancel() + for run_id, task in list(self._inflight.items()): + if run_id not in self._tearing_down: + task.cancel() if self._inflight: - await asyncio.gather(*self._inflight.values(), return_exceptions=True) + _, stuck = await asyncio.wait( + list(self._inflight.values()), timeout=self.STOP_WAIT_SECONDS + self.TEARDOWN_WAIT_SECONDS) + if stuck: + for task in stuck: + task.cancel() + await asyncio.wait(stuck, timeout=self.ABANDON_WAIT_SECONDS) + # A run still going never reached its own teardown: remove what it has recorded so far. + left = [run_id for run_id, task in self._inflight.items() + if not task.done() and run_id not in self._tearing_down and run_id in self._contexts] + if left: + logger.warning("Shutting down without %d run(s) that didn't stop; tearing down what they recorded", + len(left)) + self._abandoned.update(left) # their own teardown waits while this one runs + try: + await self._tear_down_abandoned(left) + finally: # a run that ends later tears down whatever this one didn't + self._abandoned.difference_update(left) self._inflight.clear() + async def _tear_down_abandoned(self, run_ids: list[str]) -> None: + try: + async with asyncio.timeout(self.TEARDOWN_WAIT_SECONDS): + results = await asyncio.gather( + *(self._tear_down(run_id, self._contexts[run_id]) for run_id in run_ids), return_exceptions=True) + except TimeoutError: + logger.warning("Gave up tearing down %s after %ss; what they deployed may still be up", + ", ".join(run_ids), self.TEARDOWN_WAIT_SECONDS) + return + for run_id, result in zip(run_ids, results): + if isinstance(result, BaseException): + logger.warning("Run %s: tearing it down failed (%s: %s); what it deployed may still be up", + run_id, type(result).__name__, result) + # --- Runner API -------------------------------------------------------- async def submit( @@ -110,6 +155,7 @@ async def _run(self, record: RunRecord) -> None: from agent_env.task import Task run_task: Optional[asyncio.Task] = None + context = None try: async with self._sem: # bounded concurrency; the wait here IS the queue if run_store.get_run(record.run_id).status == RunStatus.CANCELED: @@ -121,6 +167,7 @@ async def _run(self, record: RunRecord) -> None: raise LookupError(f"Task {record.task_id} v{record.task_version} not found") context = self._seed_context(record) + self._contexts[record.run_id] = context # start_step rides in metadata but Task.run takes it as a keyword; # without lifting it out here every resume re-ran from zero. run_task = asyncio.ensure_future(task.run( @@ -154,7 +201,29 @@ async def _run(self, record: RunRecord) -> None: logger.exception("Run %s failed", record.run_id) run_store.mark_terminal(record.run_id, RunStatus.FAILED, error=f"{type(e).__name__}: {e}") finally: - self._inflight.pop(record.run_id, None) + try: + if context is not None and record.run_id not in self._abandoned: + self._tearing_down.add(record.run_id) + try: + async with asyncio.timeout(self.TEARDOWN_WAIT_SECONDS): + await self._tear_down(record.run_id, context) + except TimeoutError: + logger.warning("Run %s: stopped waiting for its teardown after %ss; what it hadn't " + "removed is still up", record.run_id, self.TEARDOWN_WAIT_SECONDS) + finally: + self._tearing_down.discard(record.run_id) + self._abandoned.discard(record.run_id) + self._contexts.pop(record.run_id, None) + self._inflight.pop(record.run_id, None) + + @staticmethod + async def _tear_down(run_id: str, context) -> None: + """Remove what the run deployed, however it ended: nothing resumes from a local run's sandboxes.""" + report = await teardown_run(context) + for sandbox, why in report.failed: + logger.warning("Run %s: couldn't tear down %s: %s", run_id, sandbox.sandbox_id, why) + for sandbox in report.left: + logger.warning("Run %s: %s is still up", run_id, sandbox.sandbox_id) def _finish(self, run_id: str, context, *, exc: Optional[BaseException] = None) -> None: """Persist the terminal state of a finished Task.run(): a raised step is FAILED, diff --git a/tst/unit/cli/task_run_teardown_test.py b/tst/unit/cli/task_run_teardown_test.py new file mode 100644 index 00000000..cde31098 --- /dev/null +++ b/tst/unit/cli/task_run_teardown_test.py @@ -0,0 +1,180 @@ +"""``agent-env task run`` and ``task run-batch`` tear down what each run deployed, however it ends; ``--keep`` holds +it up until Ctrl-C, and Ctrl-C mid-run tears it down and exits 130.""" + +import asyncio +import importlib +import logging +import os +import signal +import threading + +import pytest +from click.testing import CliRunner + +from agent_env.cli import cli +from agent_env.task import Task +from agent_env.task.teardown import TORN_DOWN_KEY, RecordedSandbox, TeardownReport +from agent_env.task_step.context import DeployedSandbox + +task_run = importlib.import_module("agent_env.cli.task.run") # the package's ``run`` is the click command +pytestmark = pytest.mark.usefixtures("sigint_handled") + + +class _Deploys: + """Records a sandbox on a fake backend, then naps, then ends as told.""" + + id, version, steps = "t", 1, [] + + def __init__(self, nap=0.0, ending=None): + self.nap, self.ending, self.contexts = nap, ending, [] + + async def run(self, context, **_): + self.contexts.append(context) + context.deployed_sandboxes.append(DeployedSandbox( + sandbox_name="box", sandbox_id=f"sb-{len(self.contexts)}", sandbox_mode="vm", sandbox_type="fake")) + await asyncio.sleep(self.nap) + if self.ending: + raise self.ending + return context + + +@pytest.fixture +def torn_down(monkeypatch): + torn = [] + + async def teardown_run(context): # records what it took down, as the real one does + torn.append(context) + context.metadata.setdefault(TORN_DOWN_KEY, []).extend(s.sandbox_id for s in context.deployed_sandboxes) + return TeardownReport(terminated=tuple( + RecordedSandbox(s.sandbox_id, s.sandbox_type, "sandbox") for s in context.deployed_sandboxes)) + + monkeypatch.setattr(task_run, "teardown_run", teardown_run) + return torn + + +def _invoke(monkeypatch, tmp_path, task, *args, signal_after=None): + """Invoke the CLI, sending Ctrl-C ``signal_after`` seconds in; a Ctrl-C not yet sent when it returns never is.""" + monkeypatch.setattr(Task, "get", classmethod(lambda cls, id, version=None: task)) + timer = threading.Timer(signal_after, lambda: os.kill(os.getpid(), signal.SIGINT)) if signal_after else None + if timer: + timer.start() + logging.disable(logging.CRITICAL) + try: + return CliRunner().invoke(cli, ["task", *args, "--id", "t", "--output-dir", str(tmp_path)]) + finally: + if timer: + timer.cancel() + logging.disable(logging.NOTSET) + + +def test_a_run_that_passes_is_torn_down(monkeypatch, tmp_path, torn_down): + task = _Deploys() + + result = _invoke(monkeypatch, tmp_path, task, "run") + + assert result.exit_code == 0, result.output + assert torn_down == task.contexts + assert "Tore down 1 sandbox" in result.output + + +def test_a_run_that_fails_is_torn_down_and_reported_as_before(monkeypatch, tmp_path, torn_down): + task = _Deploys(ending=RuntimeError("boom")) + + result = _invoke(monkeypatch, tmp_path, task, "run") + + assert result.exit_code != 0 + assert isinstance(result.exception, RuntimeError) + assert torn_down == task.contexts + + +def test_each_parallel_run_is_torn_down(monkeypatch, tmp_path, torn_down): + task = _Deploys() + + result = _invoke(monkeypatch, tmp_path, task, "run", "--k", "2") + + assert result.exit_code == 0, result.output + assert len(torn_down) == 2 and {id(c) for c in torn_down} == {id(c) for c in task.contexts} + + +def test_keep_holds_a_run_up_until_ctrl_c_then_tears_it_down(monkeypatch, tmp_path, torn_down): + task = _Deploys() + + result = _invoke(monkeypatch, tmp_path, task, "run", "--keep", signal_after=1.0) + + assert result.exit_code == 0, result.output + assert "Kept up:" in result.output and "sb-1" in result.output + assert "Holding 1 sandbox up; Ctrl-C tears it down." in result.output + assert torn_down == task.contexts + + +def test_keep_still_tears_down_a_run_ctrl_c_cancels(monkeypatch, tmp_path, torn_down): + task = _Deploys(nap=30) + + result = _invoke(monkeypatch, tmp_path, task, "run", "--keep", signal_after=1.0) + + assert result.exit_code == 130, result.output + assert torn_down == task.contexts + assert "Holding" not in result.output + + +def test_ctrl_c_mid_run_tears_down_every_run_and_exits_130(monkeypatch, tmp_path, torn_down): + task = _Deploys(nap=30) + + result = _invoke(monkeypatch, tmp_path, task, "run", "--k", "2", signal_after=1.0) + + assert result.exit_code == 130, result.output + assert len(torn_down) == 2 + assert "Cancelling and tearing down (Ctrl-C again to stop now)" in result.output + assert result.output.count("CANCELLED") == 2 + + +def test_each_seed_of_a_batch_is_torn_down(monkeypatch, tmp_path, torn_down): + seeds = tmp_path / "seeds.csv" + seeds.write_text("name\nfirst\nsecond\n") + task = _Deploys() + + result = _invoke(monkeypatch, tmp_path, task, "run-batch", "--seeds", str(seeds)) + + assert result.exit_code == 0, result.output + assert len(torn_down) == 2 + + +class _OneQuickOneSlow(_Deploys): + """Its first run ends at once; every later one naps.""" + + async def run(self, context, **kwargs): + self.nap = 0 if not self.contexts else 30 + return await super().run(context, **kwargs) + + +def test_ctrl_c_with_keep_also_tears_down_a_run_that_already_finished(monkeypatch, tmp_path, torn_down): + task = _OneQuickOneSlow() + + result = _invoke(monkeypatch, tmp_path, task, "run", "--k", "2", "--keep", signal_after=1.0) + + assert result.exit_code == 130, result.output + assert {id(c) for c in torn_down} == {id(c) for c in task.contexts} + assert "Holding" not in result.output + + +def test_keep_writes_the_context_of_a_run_that_failed(monkeypatch, tmp_path, torn_down): + task = _Deploys(ending=RuntimeError("boom")) + + result = _invoke(monkeypatch, tmp_path, task, "run", "--keep", signal_after=1.0) + + assert isinstance(result.exception, RuntimeError) + assert list(tmp_path.glob("t_*.json")), "no context file to resume the kept run from" + + +def test_a_kept_failed_run_reports_its_own_error_when_its_context_cant_be_written(monkeypatch, tmp_path, torn_down): + def disk_full(*_, **__): + raise OSError(28, "No space left on device") + + monkeypatch.setattr(task_run, "_write_context", disk_full) + task = _Deploys(ending=RuntimeError("boom")) + + result = _invoke(monkeypatch, tmp_path, task, "run", "--keep", signal_after=1.0) + + assert isinstance(result.exception, RuntimeError) and str(result.exception) == "boom" + assert "Couldn't write the run's context: [Errno 28] No space left on device" in result.output + assert torn_down == task.contexts # the Ctrl-C that ended the hold diff --git a/tst/unit/runner/local_runner_test.py b/tst/unit/runner/local_runner_test.py index 504b5974..f3a21204 100644 --- a/tst/unit/runner/local_runner_test.py +++ b/tst/unit/runner/local_runner_test.py @@ -279,3 +279,211 @@ async def test_start_reconciles_an_orphan_behind_many_terminal_runs(docs): await runner.start() await runner.stop() assert run_store.get_run("old-orphan").status is RunStatus.FAILED + + +async def _settled(runner, run_id, timeout=5.0): + """Wait for a run's task to finish, teardown included: its terminal state is recorded first.""" + deadline = asyncio.get_event_loop().time() + timeout + while run_id in runner._inflight: + if asyncio.get_event_loop().time() > deadline: + raise AssertionError(f"run {run_id} never finished") + await asyncio.sleep(0.02) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("ending", ["completes", "raises", "is cancelled"]) +async def test_what_a_run_deployed_is_torn_down_however_it_ends(docs, monkeypatch, ending): + from agent_env.runner import local_runner + from agent_env.task.teardown import TeardownReport + + torn_down = [] + + async def teardown_run(context): + torn_down.append(context.metadata["workflow_id"]) + return TeardownReport() + + monkeypatch.setattr(local_runner, "teardown_run", teardown_run) + started = asyncio.Event() + + async def on_run(_): + started.set() + if ending == "raises": + raise RuntimeError("boom") + if ending == "is cancelled": + await asyncio.sleep(30) + + _install_task(monkeypatch, _FakeTask(on_run=on_run)) + runner = LocalRunner(workers=1) + await runner.start() + handle = await runner.submit("t1", 1) + try: + if ending == "is cancelled": + await asyncio.wait_for(started.wait(), 5) + await runner.cancel(handle.run_id) + await _drain(runner, handle.run_id) + await _settled(runner, handle.run_id) + finally: + await runner.stop() + + assert torn_down == [handle.run_id] + + +@pytest.mark.asyncio +async def test_stop_lets_a_teardown_under_way_finish(docs, monkeypatch): + from agent_env.runner import local_runner + from agent_env.task.teardown import TeardownReport + + tearing, finished = asyncio.Event(), [] + + async def teardown_run(context): + tearing.set() + await asyncio.sleep(0.3) + finished.append(context.metadata["workflow_id"]) + return TeardownReport() + + monkeypatch.setattr(local_runner, "teardown_run", teardown_run) + _install_task(monkeypatch, _FakeTask()) + runner = LocalRunner(workers=1) + await runner.start() + handle = await runner.submit("t1", 1) + await asyncio.wait_for(tearing.wait(), 5) + + await runner.stop() + + assert finished == [handle.run_id] + + +@pytest.mark.asyncio +async def test_a_teardown_that_never_ends_is_given_up_on_so_stop_ends(docs, monkeypatch): + from agent_env.runner import local_runner + + tearing = asyncio.Event() + + async def teardown_run(context): + tearing.set() + await asyncio.sleep(3600) + + monkeypatch.setattr(local_runner, "teardown_run", teardown_run) + monkeypatch.setattr(LocalRunner, "TEARDOWN_WAIT_SECONDS", 0.2) + _install_task(monkeypatch, _FakeTask()) + runner = LocalRunner(workers=1) + await runner.start() + await runner.submit("t1", 1) + await asyncio.wait_for(tearing.wait(), 5) + + await asyncio.wait_for(runner.stop(), 5) + + assert not runner._inflight + + + +@pytest.mark.asyncio +async def test_a_teardown_gets_its_whole_wait_however_long_the_run_took_to_stop(docs, monkeypatch): + from agent_env.runner import local_runner + from agent_env.task.teardown import TeardownReport + + finished = [] + + async def teardown_run(context): + await asyncio.sleep(0.3) + finished.append(context.metadata["workflow_id"]) + return TeardownReport() + + async def slow_to_stop(_): + try: + await asyncio.sleep(3600) + except asyncio.CancelledError: + await asyncio.sleep(0.3) # a step that takes a while to unwind + raise + + monkeypatch.setattr(local_runner, "teardown_run", teardown_run) + monkeypatch.setattr(LocalRunner, "TEARDOWN_WAIT_SECONDS", 0.5) + _install_task(monkeypatch, _FakeTask(on_run=slow_to_stop)) + runner = LocalRunner(workers=1) + await runner.start() + handle = await runner.submit("t1", 1) + await asyncio.sleep(0.1) + + await asyncio.wait_for(runner.stop(), 5) + + assert finished == [handle.run_id] + + + +@pytest.mark.asyncio +async def test_stop_ends_even_when_a_step_ignores_its_cancel(docs, monkeypatch): + from agent_env.runner import local_runner + from agent_env.task.teardown import TeardownReport + + torn_down_at = [] + + async def teardown_run(context): + torn_down_at.append(asyncio.get_running_loop().time()) + return TeardownReport() + + started = asyncio.Event() + + loop = asyncio.get_running_loop() + gives_up_at = loop.time() + 1.5 + + async def stubborn(_): + started.set() + while loop.time() < gives_up_at: # swallows every cancel for 1.5 s, then ends so the loop can close + try: + await asyncio.sleep(gives_up_at - loop.time()) + except asyncio.CancelledError: + continue + + monkeypatch.setattr(local_runner, "teardown_run", teardown_run) + monkeypatch.setattr(LocalRunner, "STOP_WAIT_SECONDS", 0.2) + monkeypatch.setattr(LocalRunner, "TEARDOWN_WAIT_SECONDS", 0.1) + monkeypatch.setattr(LocalRunner, "ABANDON_WAIT_SECONDS", 0.1) + _install_task(monkeypatch, _FakeTask(on_run=stubborn)) + runner = LocalRunner(workers=1) + await runner.start() + await runner.submit("t1", 1) + await asyncio.wait_for(started.wait(), 5) + + stopping = loop.time() + await asyncio.wait_for(runner.stop(), 5) + + stopped = loop.time() + assert stopped - stopping < 1.0, "stop() waited for the step that ignored its cancel" + assert not runner._inflight + assert torn_down_at and torn_down_at[0] < stopped, "stop() left the abandoned run's sandboxes up" + await asyncio.sleep(1.6) # let the abandoned run end before the loop closes + assert all(t >= stopped for t in torn_down_at[1:]), "the run's own teardown overlapped stop()'s" + + +@pytest.mark.asyncio +async def test_a_shutdown_teardown_that_runs_out_of_time_is_reported(docs, monkeypatch, caplog): + from agent_env.runner import local_runner + + async def teardown_run(context): + await asyncio.sleep(10) + + loop = asyncio.get_running_loop() + gives_up_at = loop.time() + 1.0 + started = asyncio.Event() + + async def stubborn(_): + started.set() + while loop.time() < gives_up_at: + try: + await asyncio.sleep(gives_up_at - loop.time()) + except asyncio.CancelledError: + continue + + monkeypatch.setattr(local_runner, "teardown_run", teardown_run) + for name, value in (("STOP_WAIT_SECONDS", 0.1), ("TEARDOWN_WAIT_SECONDS", 0.1), ("ABANDON_WAIT_SECONDS", 0.1)): + monkeypatch.setattr(LocalRunner, name, value) + _install_task(monkeypatch, _FakeTask(on_run=stubborn)) + runner = LocalRunner(workers=1) + await runner.start() + handle = await runner.submit("t1", 1) + await asyncio.wait_for(started.wait(), 5) + + await asyncio.wait_for(runner.stop(), 5) + + assert f"Gave up tearing down {handle.run_id}" in caplog.text + await asyncio.sleep(1.1) # let the abandoned run end before the loop closes