diff --git a/src/agent_env/task/task.py b/src/agent_env/task/task.py index 5626789d..30c54cd2 100644 --- a/src/agent_env/task/task.py +++ b/src/agent_env/task/task.py @@ -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]] dependents: dict[str, list[str]] pending: dict[str, int] completed: set[str] @@ -313,7 +315,7 @@ 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: @@ -321,18 +323,14 @@ def _validate_dag(cls, steps: list[TaskStep]) -> None: 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 " @@ -340,6 +338,7 @@ def _validate_dag(cls, steps: list[TaskStep]) -> None: ) 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 @@ -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] = {} @@ -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 @@ -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]: @@ -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, diff --git a/tst/benchmarks/task_graphs.py b/tst/benchmarks/task_graphs.py new file mode 100644 index 00000000..22b57a26 --- /dev/null +++ b/tst/benchmarks/task_graphs.py @@ -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 :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() diff --git a/tst/unit/task/benchmark_test.py b/tst/unit/task/benchmark_test.py new file mode 100644 index 00000000..9de36c0a --- /dev/null +++ b/tst/unit/task/benchmark_test.py @@ -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" diff --git a/tst/unit/task/depends_on_test.py b/tst/unit/task/depends_on_test.py index 2e58178c..faf4d314 100644 --- a/tst/unit/task/depends_on_test.py +++ b/tst/unit/task/depends_on_test.py @@ -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 == [] diff --git a/tst/unit/task_step/retry_test.py b/tst/unit/task_step/retry_test.py index 067ca6ac..d67f4914 100644 --- a/tst/unit/task_step/retry_test.py +++ b/tst/unit/task_step/retry_test.py @@ -21,6 +21,7 @@ import pytest +import agent_env.task.task as task_module import agent_env.task.store as store_mod from agent_env.env.env import DeployedEnv, DeployedGatewayEnv from agent_env.store import LocalSqliteDocumentStore @@ -932,6 +933,53 @@ def test_validate_dag_rejects_non_ancestor_resume_point(): Task(id="t", version=1, steps=[a, b, c]) +def test_validate_dag_rejects_unknown_resume_point(): + step = _FlakyPrompt("step", retry_config=RetryConfig(retry_from_step_id="missing")) + + with pytest.raises(ValueError, match="not in the step's dependency ancestry"): + Task(id="t", version=1, steps=[step]) + + +def test_validate_dag_rejects_duplicate_ids(): + with pytest.raises(ValueError, match="Duplicate step id 'same' at position 1"): + Task(id="t", version=1, steps=[_FakeDeployEnv("same"), _FakeDeployEnv("same")]) + + +@pytest.mark.parametrize("self_retry", [False, True]) +def test_implicit_validation_skips_prefix_slices_and_self_retry_ancestry(monkeypatch, self_retry): + class _SliceCountingSteps(list): + slice_count = 0 + + def __getitem__(self, key): + if isinstance(key, slice): + self.slice_count += 1 + return super().__getitem__(key) + + class _Step: + def __init__(self, index): + self.id = f"s{index}" + self.depends_on = None + self.retry_config = ( + RetryConfig(retry_from_step_id=self.id) if self_retry else None + ) + + ancestry_calls = 0 + original_ancestor_ids = task_module._ancestor_ids + + def count_ancestry(*args): + nonlocal ancestry_calls + ancestry_calls += 1 + return original_ancestor_ids(*args) + + monkeypatch.setattr(task_module, "_ancestor_ids", count_ancestry) + steps = _SliceCountingSteps(_Step(i) for i in range(100)) + + Task(id="t", version=1, steps=steps) + + assert steps.slice_count == 0 + assert ancestry_calls == 0 + + def test_validate_dag_accepts_ancestor_resume_point(): a = _FakeDeployEnv("a", depends_on=[]) c = _FlakyPrompt("c", retry_config=RetryConfig(retry_from_step_id="a"), depends_on=["a"])