test_orchestrator.py 5.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180
  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_category_match_item_failures_create_retry_attempt() -> None:
  98. with _session() as session:
  99. run, first, second = _seed(session)
  100. first.step_key = "demand_classify"
  101. first.retryable = True
  102. first.max_attempts = 2
  103. state = complete_step(
  104. session,
  105. step_run_id=first.step_run_id,
  106. owner="worker-1",
  107. result=StepResult(
  108. success=False,
  109. payload={
  110. "success": False,
  111. "failed_items": [{"term": "牺牲", "error": "timeout"}],
  112. },
  113. error_code="category_match_failed_items",
  114. error_message="1 category match item(s) failed",
  115. exit_code=1,
  116. ),
  117. )
  118. assert state == "retry_wait"
  119. assert run.status == "queued"
  120. assert first.status == "failed"
  121. assert second.status == "pending"
  122. def test_completion_cannot_revive_a_cancelled_run() -> None:
  123. with _session() as session:
  124. run, first, second = _seed(session)
  125. run.status = "cancelling"
  126. state = complete_step(
  127. session,
  128. step_run_id=first.step_run_id,
  129. owner="worker-1",
  130. result=StepResult(
  131. success=True,
  132. payload={"success": True},
  133. exit_code=0,
  134. ),
  135. )
  136. assert state == "cancelled"
  137. assert run.status == "cancelled"
  138. assert first.status == "cancelled"
  139. assert second.status == "pending"
  140. def test_completion_cannot_revive_a_deadline_exceeded_run() -> None:
  141. with _session() as session:
  142. run, first, second = _seed(session)
  143. run.status = "deadline_exceeded"
  144. state = complete_step(
  145. session,
  146. step_run_id=first.step_run_id,
  147. owner="worker-1",
  148. result=StepResult(
  149. success=True,
  150. payload={"success": True},
  151. exit_code=0,
  152. ),
  153. )
  154. assert state == "deadline_exceeded"
  155. assert run.status == "deadline_exceeded"
  156. assert first.status == "timed_out"
  157. assert second.status == "pending"