"""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"