Skip to content
Merged
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
38 changes: 18 additions & 20 deletions src/agent_env/task/task.py
Original file line number Diff line number Diff line change
Expand Up @@ -279,12 +279,14 @@ def _audit(before: dict, after: dict) -> ContextUpdateOps:
class _SchedulerState:
"""Initial state the DAG scheduler needs to drive execution.

- `dependencies[step_id]`: steps `step_id` waits on; a retry's rearm recounts `pending` from it.
- `dependents[step_id]`: steps released when `step_id` completes.
- `pending[step_id]`: count of deps that still need to complete before `step_id` runs.
- `completed`: seeded with steps skipped by `start_step`.
- `ready`: steps whose deps are all satisfied and can launch immediately.
"""
step_by_id: dict[str, "TaskStep"]
dependencies: dict[str, set[str]]
Comment thread
morluto marked this conversation as resolved.
dependents: dict[str, list[str]]
pending: dict[str, int]
completed: set[str]
Expand Down Expand Up @@ -313,33 +315,30 @@ def __init__(

@classmethod
def _validate_dag(cls, steps: list[TaskStep]) -> None:
ids = {s.id for s in steps}
ids = {s.id for s in steps} if any(s.depends_on is not None for s in steps) else set()
seen: set[str] = set()
for i, step in enumerate(steps):
if step.id in seen:
raise ValueError(
f"Duplicate step id '{step.id}' at position {i}; step ids must be unique within a task"
)
seen.add(step.id)
prior = {s.id for s in steps[:i]}
retry_config = getattr(step, "retry_config", None)
if retry_config is not None:
if retry_config is not None and retry_config.retry_from_step_id != step.id:
# Trust the task author's resume point, but catch an obviously
# invalid one at creation: it must be the step itself (re-run just
# this step) or an earlier dependency ancestor. We do not judge
# which spanned steps are replay-safe.
ancestors = _ancestor_ids(step.id, _dependency_ids(steps[: i + 1]))
if (
retry_config.retry_from_step_id != step.id
and retry_config.retry_from_step_id not in ancestors
):
if retry_config.retry_from_step_id not in ancestors:
raise ValueError(
f"Step '{step.id}' at position {i} declares retry_config.retry_from_step_id "
f"'{retry_config.retry_from_step_id}' which is not in the step's dependency "
f"ancestry (ancestor ids: {sorted(ancestors)})"
)
if step.depends_on is None:
continue
prior = {s.id for s in steps[:i]}
for dep in step.depends_on:
if dep.task_step_id in prior:
continue
Expand Down Expand Up @@ -527,17 +526,7 @@ async def _drive_dag(
from agent_env.task import store as _store

index_of = {s.id: i for i, s in enumerate(steps)}
# Dependency ids per active step (mirrors _build_scheduler_state), used to
# re-arm after a rollback.
active_ids = set(state.step_by_id)
dep_ids: dict[str, set[str]] = {}
for i, s in enumerate(steps):
if s.id not in active_ids:
continue
if s.depends_on is None:
dep_ids[s.id] = {p.id for p in steps[:i] if p.id in active_ids}
else:
dep_ids[s.id] = {d.task_step_id for d in s.depends_on}

in_flight: dict[str, asyncio.Task] = {}
task_to_step_id: dict[asyncio.Task, str] = {}
Expand Down Expand Up @@ -577,7 +566,7 @@ def _rearm() -> None:
re-dispatched span (and any other uncompleted, dep-satisfied step) runs.
Nothing is in flight here: the drain preceding a rollback emptied it."""
for sid in state.step_by_id:
state.pending[sid] = sum(1 for d in dep_ids[sid] if d not in state.completed)
state.pending[sid] = sum(1 for d in state.dependencies[sid] if d not in state.completed)
state.ready = [
state.step_by_id[s.id] for s in steps
if s.id in active_ids and s.id not in state.completed and state.pending[s.id] == 0
Expand Down Expand Up @@ -733,17 +722,25 @@ def _build_scheduler_state(self, steps: list, start_step: int, end_step: int | N
# earlier steps, so active steps never reference truncated ones.
active_steps = steps[:end_step] if end_step is not None else steps
step_by_id = {s.id: s for s in active_steps}
dependencies: dict[str, set[str]] = {}
dependents: dict[str, list[str]] = {step_id: [] for step_id in step_by_id}
pending: dict[str, int] = {}
implicit_dependencies = all(step.depends_on is None for step in active_steps)
previous_step_id: str | None = None
for i, step in enumerate(active_steps):
# When depends_on is None, we assume all prior steps are dependencies.
if step.depends_on is None:
if implicit_dependencies:
# One predecessor has the same reachability as every earlier step here.
dependency_step_ids = {previous_step_id} if previous_step_id is not None else set()
elif step.depends_on is None:
# When depends_on is None, we assume all prior steps are dependencies.
dependency_step_ids = {s.id for s in active_steps[:i]}
else:
dependency_step_ids = {d.task_step_id for d in step.depends_on}
dependencies[step.id] = dependency_step_ids
for dependency_step_id in dependency_step_ids:
dependents[dependency_step_id].append(step.id)
pending[step.id] = len(dependency_step_ids)
previous_step_id = step.id

completed: set[str] = set()
for s in active_steps[:start_step]:
Expand All @@ -757,6 +754,7 @@ def _build_scheduler_state(self, steps: list, start_step: int, end_step: int | N
]
return _SchedulerState(
step_by_id=step_by_id,
dependencies=dependencies,
dependents=dependents,
pending=pending,
completed=completed,
Expand Down
125 changes: 125 additions & 0 deletions tst/benchmarks/task_graphs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
"""Compare task validation and scheduler graph construction with a Git ref.

From the repository root, run:

PYTHONPATH=src:packages/agentenv-protocol/src .venv/bin/python tst/benchmarks/task_graphs.py --base origin/main

Baseline methods are extracted from ``git show <base>:src/agent_env/task/task.py``;
the current methods are imported from the working tree. Timing and allocation
measurements run separately. Synthetic steps keep the workload focused on graph work.
"""

from __future__ import annotations

import argparse
import ast
import statistics
import subprocess
import time
import tracemalloc
from dataclasses import dataclass
from types import SimpleNamespace

from agent_env.task.task import Task


def _baseline_methods(base_ref: str):
source = subprocess.check_output(
["git", "show", f"{base_ref}:src/agent_env/task/task.py"], text=True,
)
module = ast.parse(source)
dependencies = {"_SchedulerState", "_ancestor_ids", "_dependency_ids"}
baseline_dependencies = [
node for node in module.body
if isinstance(node, (ast.ClassDef, ast.FunctionDef)) and node.name in dependencies
]
if {node.name for node in baseline_dependencies} != dependencies:
raise RuntimeError(f"{base_ref} does not contain the expected scheduler dependencies")
task_class = next(
node for node in module.body if isinstance(node, ast.ClassDef) and node.name == "Task"
)
methods = {
node.name: node for node in task_class.body
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
and node.name in {"_validate_dag", "_build_scheduler_state"}
}
if set(methods) != {"_validate_dag", "_build_scheduler_state"}:
raise RuntimeError(f"{base_ref} does not contain the expected Task methods")
baseline_class = ast.ClassDef(
name="BaselineTask", bases=[], keywords=[],
body=[methods["_validate_dag"], methods["_build_scheduler_state"]],
decorator_list=[],
)
future_annotations = ast.ImportFrom(
module="__future__", names=[ast.alias(name="annotations")], level=0,
)
extracted = ast.fix_missing_locations(
ast.Module(body=[future_annotations, *baseline_dependencies, baseline_class], type_ignores=[])
)
namespace = {"dataclass": dataclass, "__name__": __name__}
exec(compile(extracted, f"{base_ref}:src/agent_env/task/task.py", "exec"), namespace)
return namespace["BaselineTask"]


def _steps(count: int, self_retry: bool = False):
return [
SimpleNamespace(
id=f"s{i}",
depends_on=None,
retry_config=(SimpleNamespace(retry_from_step_id=f"s{i}") if self_retry else None),
)
for i in range(count)
]


def _measure(fn, repetitions: int) -> tuple[float, int]:
samples = []
for _ in range(repetitions):
started = time.perf_counter()
fn()
samples.append(time.perf_counter() - started)

tracemalloc.start()
fn()
_, peak_bytes = tracemalloc.get_traced_memory()
tracemalloc.stop()
return statistics.median(samples), peak_bytes


def _row(name: str, base, count: int, repetitions: int) -> str:
steps = _steps(count, self_retry=(name == "self-retry validation"))
baseline = getattr(base, "_validate_dag") if name.endswith("validation") else None
if name == "scheduler graph":
baseline_fn = lambda: base._build_scheduler_state(object.__new__(base), steps, 0)
current_fn = lambda: Task._build_scheduler_state(object.__new__(Task), steps, 0)
else:
baseline_fn = lambda: baseline(steps)
current_fn = lambda: Task._validate_dag(steps)
base_s, base_peak = _measure(baseline_fn, repetitions)
current_s, current_peak = _measure(current_fn, repetitions)
return (
f"{name:<23} {count:>6} {base_s * 1000:>10.3f} {current_s * 1000:>10.3f} "
f"{base_peak:>12,} {current_peak:>12,} {base_s / current_s:>7.1f}x"
)


def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--base", default="origin/main", help="Git ref to compare against")
parser.add_argument("--steps", nargs="+", type=int, default=[300, 1000])
parser.add_argument("--retry-steps", type=int, default=300)
parser.add_argument("--repeats", type=int, default=5)
args = parser.parse_args()
base = _baseline_methods(args.base)

print(f"Base ref: {args.base}; timing: median of {args.repeats}; allocation: one separate traced run")
print(f"{'workload':<23} {'steps':>6} {'base ms':>10} {'current ms':>10} {'base peak B':>12} {'current B':>12} {'speedup':>8}")
for count in args.steps:
print(_row("implicit validation", base, count, args.repeats))
print(_row("scheduler graph", base, count, args.repeats))
print(_row("self-retry validation", base, args.retry_steps, args.repeats))
print("Baseline methods, dependency helpers and state class: extracted together from the selected Git ref.")


if __name__ == "__main__":
main()
31 changes: 31 additions & 0 deletions tst/unit/task/benchmark_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
"""A benchmark baseline must carry its own helper functions and state schema."""

from tst.benchmarks import task_graphs


def test_task_benchmark_loads_helpers_and_state_from_the_selected_revision(monkeypatch):
source = '''
@dataclass
class _SchedulerState:
baseline_marker: str

def _dependency_ids(steps):
return {"baseline": set()}

def _ancestor_ids(step_id, dependencies):
return set(dependencies)

class Task:
@classmethod
def _validate_dag(cls, steps):
return _ancestor_ids("step", _dependency_ids(steps))

def _build_scheduler_state(self, steps, start_step):
return _SchedulerState(baseline_marker="baseline")
'''
monkeypatch.setattr(task_graphs.subprocess, "check_output", lambda *args, **kwargs: source)

baseline = task_graphs._baseline_methods("baseline-ref")

assert baseline._validate_dag([]) == {"baseline"}
assert baseline._build_scheduler_state(object.__new__(baseline), [], 0).baseline_marker == "baseline"
64 changes: 64 additions & 0 deletions tst/unit/task/depends_on_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,3 +107,67 @@ async def test_bare_and_object_entries_gate_the_scheduler_alike(local_stores):
assert sorted(_Step.log[2:4]) == ["start b", "start c"]
assert sorted(_Step.log[4:6]) == ["end b", "end c"]
assert _Step.log[6:] == ["start d", "end d"]


def test_implicit_scheduler_graph_is_a_predecessor_chain():
task = Task(id="t", version=None, steps=[_step(str(i)) for i in range(5)])

state = task._build_scheduler_state(task.steps, start_step=0)

assert state.dependencies == {"0": set(), "1": {"0"}, "2": {"1"}, "3": {"2"}, "4": {"3"}}
assert state.dependents == {"0": ["1"], "1": ["2"], "2": ["3"], "3": ["4"], "4": []}
assert state.pending == {"0": 0, "1": 1, "2": 1, "3": 1, "4": 1}


def test_mixed_scheduler_graph_keeps_implicit_all_prior_edges():
task = Task(id="t", version=None, steps=[
_step("a"), _step("b", []), _step("c"),
])

state = task._build_scheduler_state(task.steps, start_step=0)

assert state.dependencies == {"a": set(), "b": set(), "c": {"a", "b"}}
assert state.dependents == {"a": ["c"], "b": ["c"], "c": []}
assert state.pending == {"a": 0, "b": 0, "c": 2}


@pytest.mark.asyncio
async def test_implicit_run_resume_end_boundary_and_callbacks(local_stores):
task = Task(id="t", version=None, steps=[_step(str(i)) for i in range(4)])
starts = []
completes = []

await task.run(
start_step=1,
end_step=3,
on_step_start=lambda idx, total, step, context: starts.append((idx, total, step.id)),
on_step_complete=lambda idx, total, step, context, duration: completes.append((idx, total, step.id)),
)

assert starts == [(1, 4, "1"), (2, 4, "2")]
assert [entry[2] for entry in completes] == ["1", "2"]
assert _Step.log == ["start 1", "end 1", "start 2", "end 2"]


@pytest.mark.asyncio
async def test_implicit_run_cancellation_drains_the_active_step(local_stores):
entered = asyncio.Event()
cancelled = asyncio.Event()

class _BlockingStep(_Step):
async def execute(self, context):
entered.set()
try:
await asyncio.Event().wait()
finally:
cancelled.set()

task = Task(id="t", version=None, steps=[_BlockingStep(id="a", version=None), _step("b")])
running = asyncio.create_task(task.run())
await entered.wait()
running.cancel()

with pytest.raises(asyncio.CancelledError):
await running
assert cancelled.is_set()
assert _Step.log == []
Loading
Loading