Spaces:
Sleeping
Sleeping
| """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" | |