test_orchestrator.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151
  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.orchestrator import complete_step
  9. def _session() -> Session:
  10. engine = create_engine("sqlite+pysqlite:///:memory:")
  11. PipelineRun.__table__.create(engine)
  12. PipelineStepRun.__table__.create(engine)
  13. return Session(engine, expire_on_commit=False)
  14. def _seed(session: Session) -> tuple[PipelineRun, PipelineStepRun, PipelineStepRun]:
  15. run = PipelineRun(
  16. run_id="run-1",
  17. dedupe_key="supply_pipeline:20260727:full",
  18. pipeline_key="supply_pipeline",
  19. biz_dt="20260727",
  20. trigger_type="test",
  21. run_mode="full",
  22. dry_run=True,
  23. status="running",
  24. deadline_at=datetime.now(timezone.utc).replace(tzinfo=None) + timedelta(days=1),
  25. config_snapshot_json={},
  26. date_snapshot_json={},
  27. lease_owner="worker-1",
  28. )
  29. first = PipelineStepRun(
  30. step_run_id="step-1",
  31. run_id=run.run_id,
  32. step_key="global_tree_sync",
  33. step_order=1,
  34. attempt=1,
  35. status="running",
  36. critical=True,
  37. dependency_snapshot_json=[],
  38. input_snapshot_json={},
  39. timeout_seconds=10,
  40. max_attempts=1,
  41. retryable=False,
  42. lease_owner="worker-1",
  43. )
  44. second = PipelineStepRun(
  45. step_run_id="step-2",
  46. run_id=run.run_id,
  47. step_key="demand_pool_source_sync",
  48. step_order=2,
  49. attempt=1,
  50. status="pending",
  51. critical=True,
  52. dependency_snapshot_json=["global_tree_sync"],
  53. input_snapshot_json={},
  54. timeout_seconds=10,
  55. max_attempts=1,
  56. retryable=False,
  57. )
  58. session.add_all([run, first, second])
  59. session.flush()
  60. return run, first, second
  61. def test_failure_blocks_all_downstream_steps() -> None:
  62. with _session() as session:
  63. run, first, second = _seed(session)
  64. state = complete_step(
  65. session,
  66. step_run_id=first.step_run_id,
  67. owner="worker-1",
  68. result=StepResult(
  69. success=False,
  70. payload={"success": False},
  71. error_code="test_failure",
  72. error_message="boom",
  73. exit_code=1,
  74. ),
  75. )
  76. assert state == "failed"
  77. assert run.status == "failed"
  78. assert first.status == "failed"
  79. assert second.status == "blocked"
  80. def test_success_only_releases_immediate_next_step() -> None:
  81. with _session() as session:
  82. run, first, second = _seed(session)
  83. state = complete_step(
  84. session,
  85. step_run_id=first.step_run_id,
  86. owner="worker-1",
  87. result=StepResult(
  88. success=True,
  89. payload={"success": True},
  90. exit_code=0,
  91. ),
  92. )
  93. assert state == "queued"
  94. assert run.status == "queued"
  95. assert first.status == "succeeded"
  96. assert second.status == "ready"
  97. def test_completion_cannot_revive_a_cancelled_run() -> None:
  98. with _session() as session:
  99. run, first, second = _seed(session)
  100. run.status = "cancelling"
  101. state = complete_step(
  102. session,
  103. step_run_id=first.step_run_id,
  104. owner="worker-1",
  105. result=StepResult(
  106. success=True,
  107. payload={"success": True},
  108. exit_code=0,
  109. ),
  110. )
  111. assert state == "cancelled"
  112. assert run.status == "cancelled"
  113. assert first.status == "cancelled"
  114. assert second.status == "pending"
  115. def test_completion_cannot_revive_a_deadline_exceeded_run() -> None:
  116. with _session() as session:
  117. run, first, second = _seed(session)
  118. run.status = "deadline_exceeded"
  119. state = complete_step(
  120. session,
  121. step_run_id=first.step_run_id,
  122. owner="worker-1",
  123. result=StepResult(
  124. success=True,
  125. payload={"success": True},
  126. exit_code=0,
  127. ),
  128. )
  129. assert state == "deadline_exceeded"
  130. assert run.status == "deadline_exceeded"
  131. assert first.status == "timed_out"
  132. assert second.status == "pending"