test_full_flow_contract.py 3.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118
  1. from __future__ import annotations
  2. from datetime import datetime, timedelta, timezone
  3. from sqlalchemy import create_engine
  4. from sqlalchemy.orm import Session
  5. from supply_infra.db.models.pipeline_run import PipelineRun
  6. from supply_infra.db.models.pipeline_step_run import PipelineStepRun
  7. from supply_infra.pipeline.contracts import StepResult
  8. from supply_infra.pipeline.dag import PIPELINE_STEPS
  9. from supply_infra.pipeline.orchestrator import complete_step
  10. from supply_infra.pipeline.registry import STEP_REGISTRY
  11. def _session() -> Session:
  12. engine = create_engine("sqlite+pysqlite:///:memory:")
  13. PipelineRun.__table__.create(engine)
  14. PipelineStepRun.__table__.create(engine)
  15. return Session(engine, expire_on_commit=False)
  16. def _seed_full_run(session: Session) -> tuple[PipelineRun, list[PipelineStepRun]]:
  17. run = PipelineRun(
  18. run_id="full-run",
  19. dedupe_key="supply_pipeline:20260731:full",
  20. pipeline_key="supply_pipeline",
  21. biz_dt="20260731",
  22. trigger_type="test",
  23. run_mode="full",
  24. dry_run=True,
  25. status="queued",
  26. deadline_at=(
  27. datetime.now(timezone.utc).replace(tzinfo=None)
  28. + timedelta(days=1)
  29. ),
  30. config_snapshot_json={},
  31. date_snapshot_json={},
  32. )
  33. steps = [
  34. PipelineStepRun(
  35. step_run_id=f"step-{definition.order}",
  36. run_id=run.run_id,
  37. step_key=definition.key,
  38. step_order=definition.order,
  39. attempt=1,
  40. status="ready" if definition.order == 1 else "pending",
  41. critical=definition.critical,
  42. dependency_snapshot_json=list(definition.dependencies),
  43. input_snapshot_json={},
  44. timeout_seconds=definition.timeout_seconds,
  45. max_attempts=definition.max_attempts,
  46. retryable=definition.retryable,
  47. )
  48. for definition in PIPELINE_STEPS
  49. ]
  50. session.add_all([run, *steps])
  51. session.flush()
  52. return run, steps
  53. def test_all_fifteen_steps_advance_to_terminal_success() -> None:
  54. with _session() as session:
  55. run, steps = _seed_full_run(session)
  56. assert {step.step_key for step in steps} == set(STEP_REGISTRY)
  57. for index, step in enumerate(steps):
  58. step.status = "running"
  59. step.lease_owner = "worker"
  60. run.status = "running"
  61. state = complete_step(
  62. session,
  63. step_run_id=step.step_run_id,
  64. owner="worker",
  65. result=StepResult(
  66. success=True,
  67. payload={"success": True},
  68. exit_code=0,
  69. ),
  70. )
  71. if index < len(steps) - 1:
  72. assert state == "queued"
  73. assert steps[index + 1].status == "ready"
  74. else:
  75. assert state == "succeeded"
  76. assert run.status == "succeeded"
  77. def test_business_failure_blocks_every_remaining_step() -> None:
  78. with _session() as session:
  79. run, steps = _seed_full_run(session)
  80. target = next(
  81. step for step in steps if step.step_key == "video_discovery"
  82. )
  83. for step in steps[: target.step_order - 1]:
  84. step.status = "succeeded"
  85. target.status = "running"
  86. target.lease_owner = "worker"
  87. run.status = "running"
  88. state = complete_step(
  89. session,
  90. step_run_id=target.step_run_id,
  91. owner="worker",
  92. result=StepResult(
  93. success=True,
  94. payload={"success": True, "failed": 1},
  95. exit_code=0,
  96. ),
  97. )
  98. assert state == "failed"
  99. assert run.status == "failed"
  100. assert all(
  101. step.status == "blocked"
  102. for step in steps
  103. if step.step_order > target.step_order
  104. )