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.orchestrator import complete_step 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(session: Session) -> tuple[PipelineRun, PipelineStepRun, PipelineStepRun]: run = PipelineRun( run_id="run-1", dedupe_key="supply_pipeline:20260727:full", pipeline_key="supply_pipeline", biz_dt="20260727", trigger_type="test", run_mode="full", dry_run=True, status="running", deadline_at=datetime.now(timezone.utc).replace(tzinfo=None) + timedelta(days=1), config_snapshot_json={}, date_snapshot_json={}, lease_owner="worker-1", ) first = PipelineStepRun( step_run_id="step-1", run_id=run.run_id, step_key="global_tree_sync", step_order=1, attempt=1, status="running", critical=True, dependency_snapshot_json=[], input_snapshot_json={}, timeout_seconds=10, max_attempts=1, retryable=False, lease_owner="worker-1", ) second = PipelineStepRun( step_run_id="step-2", run_id=run.run_id, step_key="demand_pool_source_sync", step_order=2, attempt=1, status="pending", critical=True, dependency_snapshot_json=["global_tree_sync"], input_snapshot_json={}, timeout_seconds=10, max_attempts=1, retryable=False, ) session.add_all([run, first, second]) session.flush() return run, first, second def test_failure_blocks_all_downstream_steps() -> None: with _session() as session: run, first, second = _seed(session) state = complete_step( session, step_run_id=first.step_run_id, owner="worker-1", result=StepResult( success=False, payload={"success": False}, error_code="test_failure", error_message="boom", exit_code=1, ), ) assert state == "failed" assert run.status == "failed" assert first.status == "failed" assert second.status == "blocked" def test_success_only_releases_immediate_next_step() -> None: with _session() as session: run, first, second = _seed(session) state = complete_step( session, step_run_id=first.step_run_id, owner="worker-1", result=StepResult( success=True, payload={"success": True}, exit_code=0, ), ) assert state == "queued" assert run.status == "queued" assert first.status == "succeeded" assert second.status == "ready" def test_completion_cannot_revive_a_cancelled_run() -> None: with _session() as session: run, first, second = _seed(session) run.status = "cancelling" state = complete_step( session, step_run_id=first.step_run_id, owner="worker-1", result=StepResult( success=True, payload={"success": True}, exit_code=0, ), ) assert state == "cancelled" assert run.status == "cancelled" assert first.status == "cancelled" assert second.status == "pending" def test_completion_cannot_revive_a_deadline_exceeded_run() -> None: with _session() as session: run, first, second = _seed(session) run.status = "deadline_exceeded" state = complete_step( session, step_run_id=first.step_run_id, owner="worker-1", result=StepResult( success=True, payload={"success": True}, exit_code=0, ), ) assert state == "deadline_exceeded" assert run.status == "deadline_exceeded" assert first.status == "timed_out" assert second.status == "pending"