from __future__ import annotations from datetime import datetime, timedelta, timezone from sqlalchemy import create_engine from sqlalchemy.orm import Session from supply_infra.db.models.pipeline_run import PipelineRun from supply_infra.db.models.pipeline_step_run import PipelineStepRun from supply_infra.pipeline.contracts import StepResult from supply_infra.pipeline.dag import PIPELINE_STEPS from supply_infra.pipeline.orchestrator import complete_step from supply_infra.pipeline.registry import STEP_REGISTRY def _session() -> Session: engine = create_engine("sqlite+pysqlite:///:memory:") PipelineRun.__table__.create(engine) PipelineStepRun.__table__.create(engine) return Session(engine, expire_on_commit=False) def _seed_full_run(session: Session) -> tuple[PipelineRun, list[PipelineStepRun]]: run = PipelineRun( run_id="full-run", dedupe_key="supply_pipeline:20260731:full", pipeline_key="supply_pipeline", biz_dt="20260731", trigger_type="test", run_mode="full", dry_run=True, status="queued", deadline_at=( datetime.now(timezone.utc).replace(tzinfo=None) + timedelta(days=1) ), config_snapshot_json={}, date_snapshot_json={}, ) steps = [ PipelineStepRun( step_run_id=f"step-{definition.order}", run_id=run.run_id, step_key=definition.key, step_order=definition.order, attempt=1, status="ready" if definition.order == 1 else "pending", critical=definition.critical, dependency_snapshot_json=list(definition.dependencies), input_snapshot_json={}, timeout_seconds=definition.timeout_seconds, max_attempts=definition.max_attempts, retryable=definition.retryable, ) for definition in PIPELINE_STEPS ] session.add_all([run, *steps]) session.flush() return run, steps def test_all_fifteen_steps_advance_to_terminal_success() -> None: with _session() as session: run, steps = _seed_full_run(session) assert {step.step_key for step in steps} == set(STEP_REGISTRY) for index, step in enumerate(steps): step.status = "running" step.lease_owner = "worker" run.status = "running" state = complete_step( session, step_run_id=step.step_run_id, owner="worker", result=StepResult( success=True, payload={"success": True}, exit_code=0, ), ) if index < len(steps) - 1: assert state == "queued" assert steps[index + 1].status == "ready" else: assert state == "succeeded" assert run.status == "succeeded" def test_business_failure_blocks_every_remaining_step() -> None: with _session() as session: run, steps = _seed_full_run(session) target = next( step for step in steps if step.step_key == "video_discovery" ) for step in steps[: target.step_order - 1]: step.status = "succeeded" target.status = "running" target.lease_owner = "worker" run.status = "running" state = complete_step( session, step_run_id=target.step_run_id, owner="worker", result=StepResult( success=True, payload={"success": True, "failed": 1}, exit_code=0, ), ) assert state == "failed" assert run.status == "failed" assert all( step.status == "blocked" for step in steps if step.step_order > target.step_order )