pipeline_step_run_repo.py 8.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242
  1. from __future__ import annotations
  2. import uuid
  3. from datetime import datetime, timedelta
  4. from typing import Any
  5. from sqlalchemy import func, select
  6. from supply_infra.db.models.pipeline_run import PipelineRun
  7. from supply_infra.db.models.pipeline_step_run import PipelineStepRun
  8. from supply_infra.db.repositories.base import BaseRepository
  9. from supply_infra.db.repositories.pipeline_lock_repo import PipelineLockRepository
  10. class PipelineStepRunRepository(BaseRepository[PipelineStepRun]):
  11. model = PipelineStepRun
  12. def get(
  13. self,
  14. step_run_id: str,
  15. *,
  16. for_update: bool = False,
  17. ) -> PipelineStepRun | None:
  18. stmt = select(PipelineStepRun).where(
  19. PipelineStepRun.step_run_id == step_run_id
  20. )
  21. if for_update:
  22. stmt = stmt.with_for_update()
  23. return self.session.scalar(stmt)
  24. def create_steps(self, rows: list[dict[str, Any]]) -> list[PipelineStepRun]:
  25. entities = [PipelineStepRun(**row) for row in rows]
  26. self.session.add_all(entities)
  27. self.session.flush()
  28. return entities
  29. def list_for_run(self, run_id: str) -> list[PipelineStepRun]:
  30. stmt = (
  31. select(PipelineStepRun)
  32. .where(PipelineStepRun.run_id == run_id)
  33. .order_by(PipelineStepRun.step_order, PipelineStepRun.attempt)
  34. )
  35. return list(self.session.scalars(stmt).all())
  36. def latest_attempts(self, run_id: str) -> dict[str, PipelineStepRun]:
  37. latest: dict[str, PipelineStepRun] = {}
  38. for item in self.list_for_run(run_id):
  39. latest[item.step_key] = item
  40. return latest
  41. def claim_next_ready(
  42. self,
  43. *,
  44. owner: str,
  45. lease_seconds: int,
  46. max_active_steps: int,
  47. now: datetime,
  48. ) -> PipelineStepRun | None:
  49. # Serialize count-and-claim so concurrent workers cannot oversubscribe
  50. # the configured global active-step cap.
  51. PipelineLockRepository(self.session).lock_control_plane(now=now)
  52. active_steps = self.session.scalar(
  53. select(func.count())
  54. .select_from(PipelineStepRun)
  55. .where(PipelineStepRun.status == "running")
  56. )
  57. if int(active_steps or 0) >= max_active_steps:
  58. return None
  59. running_run_ids = list(
  60. self.session.scalars(
  61. select(PipelineStepRun.run_id)
  62. .where(PipelineStepRun.status == "running")
  63. .distinct()
  64. ).all()
  65. )
  66. stmt = (
  67. select(PipelineStepRun)
  68. .where(PipelineStepRun.status == "ready")
  69. .where(
  70. (PipelineStepRun.next_retry_at.is_(None))
  71. | (PipelineStepRun.next_retry_at <= now)
  72. )
  73. .order_by(PipelineStepRun.created_at, PipelineStepRun.step_order)
  74. .with_for_update(skip_locked=True)
  75. .limit(1)
  76. )
  77. if running_run_ids:
  78. # A daily run may contain parallel steps in the future, but a
  79. # different run must never overlap while an old process is alive.
  80. stmt = stmt.where(PipelineStepRun.run_id.in_(running_run_ids))
  81. step = self.session.scalar(stmt)
  82. if step is None:
  83. return None
  84. step.status = "running"
  85. step.started_at = step.started_at or now
  86. step.heartbeat_at = now
  87. step.lease_owner = owner
  88. step.lease_until = now + timedelta(seconds=lease_seconds)
  89. self.session.flush()
  90. return step
  91. def has_running_for_run(self, run_id: str) -> bool:
  92. count = self.session.scalar(
  93. select(func.count())
  94. .select_from(PipelineStepRun)
  95. .where(PipelineStepRun.run_id == run_id)
  96. .where(PipelineStepRun.status == "running")
  97. )
  98. return bool(count)
  99. def heartbeat(
  100. self,
  101. step_run_id: str,
  102. *,
  103. owner: str,
  104. now: datetime,
  105. lease_until: datetime,
  106. ) -> bool:
  107. step = self.get(step_run_id, for_update=True)
  108. if step is None or step.status != "running" or step.lease_owner != owner:
  109. return False
  110. step.heartbeat_at = now
  111. step.lease_until = lease_until
  112. self.session.flush()
  113. return True
  114. def finish(
  115. self,
  116. step: PipelineStepRun,
  117. *,
  118. status: str,
  119. now: datetime,
  120. exit_code: int | None,
  121. result: dict[str, Any] | None,
  122. error_code: str | None,
  123. error_message: str | None,
  124. log_uri: str | None,
  125. ) -> None:
  126. step.status = status
  127. step.finished_at = now
  128. step.heartbeat_at = now
  129. step.lease_owner = None
  130. step.lease_until = None
  131. step.exit_code = exit_code
  132. step.result_summary_json = result
  133. step.error_code = error_code
  134. step.error_message = error_message
  135. step.log_uri = log_uri
  136. self.session.flush()
  137. def create_retry_attempt(
  138. self,
  139. previous: PipelineStepRun,
  140. *,
  141. status: str = "ready",
  142. next_retry_at: datetime | None = None,
  143. ) -> PipelineStepRun:
  144. entity = PipelineStepRun(
  145. step_run_id=str(uuid.uuid4()),
  146. run_id=previous.run_id,
  147. step_key=previous.step_key,
  148. step_order=previous.step_order,
  149. attempt=previous.attempt + 1,
  150. status=status,
  151. critical=previous.critical,
  152. dependency_snapshot_json=list(previous.dependency_snapshot_json or []),
  153. input_snapshot_json=dict(previous.input_snapshot_json or {}),
  154. timeout_seconds=previous.timeout_seconds,
  155. max_attempts=previous.max_attempts,
  156. retryable=previous.retryable,
  157. next_retry_at=next_retry_at,
  158. )
  159. return self.add(entity)
  160. def list_expired_running(self, now: datetime, *, limit: int = 100) -> list[PipelineStepRun]:
  161. stmt = (
  162. select(PipelineStepRun)
  163. .where(PipelineStepRun.status == "running")
  164. .where(PipelineStepRun.lease_until.is_not(None))
  165. .where(PipelineStepRun.lease_until < now)
  166. .order_by(PipelineStepRun.lease_until)
  167. .with_for_update(skip_locked=True)
  168. .limit(limit)
  169. )
  170. return list(self.session.scalars(stmt).all())
  171. def promote_due_retries(self, now: datetime, *, limit: int = 100) -> int:
  172. stmt = (
  173. select(PipelineStepRun)
  174. .where(PipelineStepRun.status == "retry_wait")
  175. .where(PipelineStepRun.next_retry_at <= now)
  176. .order_by(PipelineStepRun.next_retry_at)
  177. .with_for_update(skip_locked=True)
  178. .limit(limit)
  179. )
  180. count = 0
  181. for step in self.session.scalars(stmt).all():
  182. step.status = "ready"
  183. count += 1
  184. self.session.flush()
  185. return count
  186. def block_unfinished_downstream(
  187. self,
  188. *,
  189. run_id: str,
  190. after_order: int,
  191. reason: str,
  192. ) -> int:
  193. count = 0
  194. for step in self.list_for_run(run_id):
  195. if step.step_order <= after_order or step.status in {
  196. "succeeded",
  197. "failed",
  198. "timed_out",
  199. "cancelled",
  200. }:
  201. continue
  202. step.status = "blocked"
  203. step.error_code = "upstream_failed"
  204. step.error_message = reason
  205. count += 1
  206. self.session.flush()
  207. return count
  208. def get_latest_succeeded_biz_dt(self, step_key: str) -> str | None:
  209. """返回指定步骤最近一次成功完成时所属 pipeline_run 的 biz_dt。"""
  210. stmt = (
  211. select(PipelineRun.biz_dt)
  212. .join(PipelineStepRun, PipelineStepRun.run_id == PipelineRun.run_id)
  213. .where(
  214. PipelineStepRun.step_key == step_key,
  215. PipelineStepRun.status == "succeeded",
  216. PipelineStepRun.finished_at.is_not(None),
  217. )
  218. .order_by(PipelineStepRun.finished_at.desc())
  219. .limit(1)
  220. )
  221. value = self.session.scalar(stmt)
  222. return str(value) if value else None