| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118 |
- 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
- )
|