From 0f261269983d9872f4524c743cbba6ca72dc9676 Mon Sep 17 00:00:00 2001
From: morluto <76467478+morluto@users.noreply.github.com>
Date: Wed, 7 Oct 2026 03:47:32 +0800
Subject: [PATCH 1/4] perf(task): skip unused validation work
---
src/agent_env/task/task.py | 11 +++-----
tst/unit/task_step/retry_test.py | 48 ++++++++++++++++++++++++++++++++
2 files changed, 52 insertions(+), 7 deletions(-)
diff --git a/src/agent_env/task/task.py b/src/agent_env/task/task.py
index 5626789d..eb503ede 100644
--- a/src/agent_env/task/task.py
+++ b/src/agent_env/task/task.py
@@ -313,7 +313,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 +321,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 +336,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
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"])
From ffb49cf2a66d8040ad4c40bb704562ac22635e9a Mon Sep 17 00:00:00 2001
From: morluto <76467478+morluto@users.noreply.github.com>
Date: Wed, 7 Oct 2026 03:47:32 +0800
Subject: [PATCH 2/4] perf(task): reduce implicit scheduler graphs to chains
---
src/agent_env/task/task.py | 20 ++++-
tst/benchmarks/task_graphs.py | 121 +++++++++++++++++++++++++++++++
tst/unit/task/depends_on_test.py | 62 ++++++++++++++++
3 files changed, 200 insertions(+), 3 deletions(-)
create mode 100644 tst/benchmarks/task_graphs.py
diff --git a/src/agent_env/task/task.py b/src/agent_env/task/task.py
index eb503ede..18465bf3 100644
--- a/src/agent_env/task/task.py
+++ b/src/agent_env/task/task.py
@@ -528,13 +528,21 @@ async def _drive_dag(
# re-arm after a rollback.
active_ids = set(state.step_by_id)
dep_ids: dict[str, set[str]] = {}
+ implicit_dependencies = all(
+ s.depends_on is None for s in steps if s.id in active_ids
+ )
+ previous_active_id: str | None = None
for i, s in enumerate(steps):
if s.id not in active_ids:
continue
- if s.depends_on is None:
+ if implicit_dependencies:
+ # Match the scheduler's chain so rollback rearming preserves reachability.
+ dep_ids[s.id] = {previous_active_id} if previous_active_id is not None else set()
+ elif 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}
+ previous_active_id = s.id
in_flight: dict[str, asyncio.Task] = {}
task_to_step_id: dict[asyncio.Task, str] = {}
@@ -732,15 +740,21 @@ def _build_scheduler_state(self, steps: list, start_step: int, end_step: int | N
step_by_id = {s.id: s for s in active_steps}
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 ()
+ 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}
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]:
diff --git a/tst/benchmarks/task_graphs.py b/tst/benchmarks/task_graphs.py
new file mode 100644
index 00000000..30f69612
--- /dev/null
+++ b/tst/benchmarks/task_graphs.py
@@ -0,0 +1,121 @@
+"""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 types import SimpleNamespace
+
+from agent_env.task.task import Task, _SchedulerState, _ancestor_ids, _dependency_ids
+
+
+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)
+ 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_class], type_ignores=[])
+ )
+ namespace = {
+ "_SchedulerState": _SchedulerState,
+ "_ancestor_ids": _ancestor_ids,
+ "_dependency_ids": _dependency_ids,
+ }
+ 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 method source: AST-extracted from the selected Git ref; unchanged dependency helpers are shared.")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/tst/unit/task/depends_on_test.py b/tst/unit/task/depends_on_test.py
index 2e58178c..9213a827 100644
--- a/tst/unit/task/depends_on_test.py
+++ b/tst/unit/task/depends_on_test.py
@@ -107,3 +107,65 @@ 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.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.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 == []
From 2e0b997cbdfb758146f05d2badd2973fe0047481 Mon Sep 17 00:00:00 2001
From: Edgar Arakelyan
Date: Tue, 6 Oct 2026 17:13:19 -0700
Subject: [PATCH 3/4] refactor(task): build the scheduler's dependency map once
_drive_dag rebuilt each active step's dependency ids to re-arm after a
retry's rollback, mirroring _build_scheduler_state's chain-or-dense
choice. _build_scheduler_state now records them on _SchedulerState as
`dependencies`, and _rearm reads that, so the choice lives in one place.
Co-Authored-By: Claude Opus 5.5
---
src/agent_env/task/task.py | 27 +++++++--------------------
tst/unit/task/depends_on_test.py | 2 ++
2 files changed, 9 insertions(+), 20 deletions(-)
diff --git a/src/agent_env/task/task.py b/src/agent_env/task/task.py
index 18465bf3..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]
@@ -524,25 +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]] = {}
- implicit_dependencies = all(
- s.depends_on is None for s in steps if s.id in active_ids
- )
- previous_active_id: str | None = None
- for i, s in enumerate(steps):
- if s.id not in active_ids:
- continue
- if implicit_dependencies:
- # Match the scheduler's chain so rollback rearming preserves reachability.
- dep_ids[s.id] = {previous_active_id} if previous_active_id is not None else set()
- elif 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}
- previous_active_id = s.id
in_flight: dict[str, asyncio.Task] = {}
task_to_step_id: dict[asyncio.Task, str] = {}
@@ -582,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
@@ -738,6 +722,7 @@ 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)
@@ -745,12 +730,13 @@ def _build_scheduler_state(self, steps: list, start_step: int, end_step: int | N
for i, step in enumerate(active_steps):
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 ()
+ 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)
@@ -768,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/unit/task/depends_on_test.py b/tst/unit/task/depends_on_test.py
index 9213a827..faf4d314 100644
--- a/tst/unit/task/depends_on_test.py
+++ b/tst/unit/task/depends_on_test.py
@@ -114,6 +114,7 @@ def test_implicit_scheduler_graph_is_a_predecessor_chain():
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}
@@ -125,6 +126,7 @@ def test_mixed_scheduler_graph_keeps_implicit_all_prior_edges():
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}
From eacf8211df584987c79a687a5c978844aae7fbb0 Mon Sep 17 00:00:00 2001
From: morluto <76467478+morluto@users.noreply.github.com>
Date: Wed, 7 Oct 2026 08:56:52 +0800
Subject: [PATCH 4/4] fix(benchmark): load task baselines from one revision
---
tst/benchmarks/task_graphs.py | 20 ++++++++++++--------
tst/unit/task/benchmark_test.py | 31 +++++++++++++++++++++++++++++++
2 files changed, 43 insertions(+), 8 deletions(-)
create mode 100644 tst/unit/task/benchmark_test.py
diff --git a/tst/benchmarks/task_graphs.py b/tst/benchmarks/task_graphs.py
index 30f69612..22b57a26 100644
--- a/tst/benchmarks/task_graphs.py
+++ b/tst/benchmarks/task_graphs.py
@@ -17,9 +17,10 @@
import subprocess
import time
import tracemalloc
+from dataclasses import dataclass
from types import SimpleNamespace
-from agent_env.task.task import Task, _SchedulerState, _ancestor_ids, _dependency_ids
+from agent_env.task.task import Task
def _baseline_methods(base_ref: str):
@@ -27,6 +28,13 @@ def _baseline_methods(base_ref: str):
["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"
)
@@ -46,13 +54,9 @@ def _baseline_methods(base_ref: str):
module="__future__", names=[ast.alias(name="annotations")], level=0,
)
extracted = ast.fix_missing_locations(
- ast.Module(body=[future_annotations, baseline_class], type_ignores=[])
+ ast.Module(body=[future_annotations, *baseline_dependencies, baseline_class], type_ignores=[])
)
- namespace = {
- "_SchedulerState": _SchedulerState,
- "_ancestor_ids": _ancestor_ids,
- "_dependency_ids": _dependency_ids,
- }
+ namespace = {"dataclass": dataclass, "__name__": __name__}
exec(compile(extracted, f"{base_ref}:src/agent_env/task/task.py", "exec"), namespace)
return namespace["BaselineTask"]
@@ -114,7 +118,7 @@ def main() -> None:
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 method source: AST-extracted from the selected Git ref; unchanged dependency helpers are shared.")
+ print("Baseline methods, dependency helpers and state class: extracted together from the selected Git ref.")
if __name__ == "__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"