permanence-training / tests /test_pipeline_orchestration.py
chane35's picture
PERMANENCE: reversibility-aware RL environment for training LLM agents
796da7c verified
Raw
History Blame
5.78 kB
"""Tests for the pipeline orchestrator's wiring and control flow.
These tests replace each stage's ``run_*`` function with a fake so we can
verify:
* Artifact paths are passed correctly between stages
* A failing gate aborts the pipeline (bail_on_failure=True)
* ``--from`` and ``--only`` flags skip the right stages
* ``pipeline_summary.json`` is written with the right shape
Run on CPU only.
"""
from __future__ import annotations
import json
import sys
from pathlib import Path
from unittest.mock import patch
_ROOT = Path(__file__).resolve().parent.parent
if str(_ROOT) not in sys.path:
sys.path.insert(0, str(_ROOT))
from training.config import TrainingConfig
from training.pipeline import STAGES, run_pipeline
def _fake_stage(ok: bool = True, extra: dict | None = None):
def fake(config, *args, **kwargs):
return {"ok": ok, **(extra or {})}
return fake
def test_stages_list_is_ordered():
"""Pipeline stages run in this exact order: sft → gate → grpo → eval."""
assert STAGES == ["sft", "gate", "grpo", "eval"]
def test_pipeline_runs_all_stages_when_all_pass():
"""Happy path: every stage returns ok=True, pipeline completes."""
cfg = TrainingConfig()
with patch("training.stages.stage_1_sft.run_sft", _fake_stage(True)), \
patch("training.stages.stage_2_gate.run_gate", _fake_stage(True, {"coverage": 1.0})), \
patch("training.stages.stage_3_grpo.run_grpo", _fake_stage(True, {"mean_reward": 0.8})), \
patch("training.stages.stage_4_eval.run_eval", _fake_stage(True)):
summary = run_pipeline(cfg, list(STAGES), bail_on_failure=True)
assert summary["final_status"] == "completed"
assert set(summary["stages"].keys()) == set(STAGES)
for stage in STAGES:
assert summary["stages"][stage]["ok"] is True
def test_pipeline_bails_when_gate_fails():
"""If the gate fails, GRPO and eval must NOT run — this is the whole
point of the gate: fail fast, don't burn GPU on a broken SFT."""
cfg = TrainingConfig()
grpo_called = [False]
eval_called = [False]
def track_grpo(*args, **kwargs):
grpo_called[0] = True
return {"ok": True}
def track_eval(*args, **kwargs):
eval_called[0] = True
return {"ok": True}
with patch("training.stages.stage_1_sft.run_sft", _fake_stage(True)), \
patch("training.stages.stage_2_gate.run_gate", _fake_stage(False, {"coverage": 0.5})), \
patch("training.stages.stage_3_grpo.run_grpo", track_grpo), \
patch("training.stages.stage_4_eval.run_eval", track_eval):
summary = run_pipeline(cfg, list(STAGES), bail_on_failure=True)
assert summary["final_status"] == "failed_at_gate"
assert grpo_called[0] is False, "GRPO ran even though gate failed!"
assert eval_called[0] is False, "Eval ran even though gate failed!"
def test_pipeline_bails_when_sft_fails():
"""Even earlier: if SFT fails (loss too high), nothing downstream runs."""
cfg = TrainingConfig()
gate_called = [False]
with patch("training.stages.stage_1_sft.run_sft", _fake_stage(False, {"final_training_loss": 2.5})), \
patch("training.stages.stage_2_gate.run_gate", lambda *a, **k: gate_called.__setitem__(0, True) or {"ok": True}):
summary = run_pipeline(cfg, list(STAGES), bail_on_failure=True)
assert summary["final_status"] == "failed_at_sft"
assert gate_called[0] is False
def test_pipeline_no_bail_runs_all_stages_even_on_failure():
"""With bail_on_failure=False, each stage runs regardless of prior
failures. Used for post-mortem runs where we want partial artifacts."""
cfg = TrainingConfig()
with patch("training.stages.stage_1_sft.run_sft", _fake_stage(False)), \
patch("training.stages.stage_2_gate.run_gate", _fake_stage(False)), \
patch("training.stages.stage_3_grpo.run_grpo", _fake_stage(False)), \
patch("training.stages.stage_4_eval.run_eval", _fake_stage(True)):
summary = run_pipeline(cfg, list(STAGES), bail_on_failure=False)
assert summary["final_status"] == "completed"
assert all(stage in summary["stages"] for stage in STAGES)
def test_pipeline_with_subset_of_stages():
"""``--only grpo`` or ``--from gate`` narrows the stage list. Pipeline
runs exactly those stages."""
cfg = TrainingConfig()
with patch("training.stages.stage_3_grpo.run_grpo", _fake_stage(True)):
summary = run_pipeline(cfg, ["grpo"], bail_on_failure=True)
assert list(summary["stages"].keys()) == ["grpo"]
assert summary["final_status"] == "completed"
def test_exception_in_stage_surfaces_cleanly():
"""If a stage's run function raises (not returns ok=False), the
orchestrator must catch it and record ``final_status=fatal``."""
cfg = TrainingConfig()
def raiser(*args, **kwargs):
raise RuntimeError("simulated stage crash")
with patch("training.stages.stage_1_sft.run_sft", raiser):
summary = run_pipeline(cfg, ["sft"], bail_on_failure=True)
assert summary["final_status"] == "fatal"
assert "error" in summary["stages"]["sft"]
def test_pipeline_summary_is_json_serializable():
"""The final summary must round-trip through JSON so it can be written
to artifacts/pipeline_summary.json."""
cfg = TrainingConfig()
with patch("training.stages.stage_1_sft.run_sft", _fake_stage(True, {"custom_metric": 0.42})):
summary = run_pipeline(cfg, ["sft"], bail_on_failure=True)
# This serialization is what pipeline.py main() does; if it fails,
# the artifact won't be written.
s = json.dumps(summary, default=str)
assert len(s) > 10
# And re-parses
parsed = json.loads(s)
assert parsed["final_status"] == "completed"