Skip to content
Merged
15 changes: 3 additions & 12 deletions src/agent_env/cli/eval/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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."""
Expand All @@ -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
Expand Down
216 changes: 162 additions & 54 deletions src/agent_env/cli/task/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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
Comment thread
earakely-scale marked this conversation as resolved.
Comment thread
earakely-scale marked this conversation as resolved.
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)")
Expand All @@ -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'")

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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 <placeholder> variables
are rendered using the seed values before being sent to the agent.
Expand Down Expand Up @@ -506,13 +609,16 @@ 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}"))
tag = f"[seed {index}: {seed_name[:50]}]"
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
Expand All @@ -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
Expand All @@ -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)
16 changes: 16 additions & 0 deletions src/agent_env/cli/teardown_output.py
Original file line number Diff line number Diff line change
@@ -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"))
Loading
Loading