pipeline_run_repo.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170
  1. from __future__ import annotations
  2. from datetime import datetime
  3. from typing import Any
  4. from sqlalchemy import exists, select
  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.db.repositories.base import BaseRepository
  8. class PipelineRunRepository(BaseRepository[PipelineRun]):
  9. model = PipelineRun
  10. def get(self, run_id: str, *, for_update: bool = False) -> PipelineRun | None:
  11. stmt = select(PipelineRun).where(PipelineRun.run_id == run_id)
  12. if for_update:
  13. stmt = stmt.with_for_update()
  14. return self.session.scalar(stmt)
  15. def get_by_dedupe_key(self, dedupe_key: str) -> PipelineRun | None:
  16. stmt = select(PipelineRun).where(PipelineRun.dedupe_key == dedupe_key)
  17. return self.session.scalar(stmt)
  18. def get_for_business_day(
  19. self,
  20. *,
  21. pipeline_key: str,
  22. biz_dt: str,
  23. ) -> PipelineRun | None:
  24. stmt = (
  25. select(PipelineRun)
  26. .where(PipelineRun.pipeline_key == pipeline_key)
  27. .where(PipelineRun.biz_dt == biz_dt)
  28. .order_by(PipelineRun.created_at.desc())
  29. .limit(1)
  30. )
  31. return self.session.scalar(stmt)
  32. def create(self, values: dict[str, Any]) -> PipelineRun:
  33. return self.add(PipelineRun(**values))
  34. def list_recent(
  35. self,
  36. *,
  37. limit: int = 50,
  38. status: str | None = None,
  39. biz_dt: str | None = None,
  40. ) -> list[PipelineRun]:
  41. stmt = select(PipelineRun)
  42. if status:
  43. stmt = stmt.where(PipelineRun.status == status)
  44. if biz_dt:
  45. stmt = stmt.where(PipelineRun.biz_dt == biz_dt)
  46. stmt = stmt.order_by(PipelineRun.created_at.desc()).limit(limit)
  47. return list(self.session.scalars(stmt).all())
  48. def get_latest_nonterminal(self) -> PipelineRun | None:
  49. has_running_step = exists(
  50. select(PipelineStepRun.step_run_id)
  51. .where(PipelineStepRun.run_id == PipelineRun.run_id)
  52. .where(PipelineStepRun.status == "running")
  53. )
  54. stmt = (
  55. select(PipelineRun)
  56. .where(
  57. PipelineRun.status.in_(
  58. ("queued", "running", "cancelling", "blocked_previous_run")
  59. )
  60. | has_running_step
  61. )
  62. .order_by(PipelineRun.created_at.desc())
  63. .limit(1)
  64. )
  65. return self.session.scalar(stmt)
  66. def list_nonterminal(self, *, limit: int = 100) -> list[PipelineRun]:
  67. stmt = (
  68. select(PipelineRun)
  69. .where(
  70. PipelineRun.status.in_(
  71. ("queued", "running", "cancelling", "blocked_previous_run")
  72. )
  73. )
  74. .order_by(PipelineRun.created_at)
  75. .limit(limit)
  76. )
  77. return list(self.session.scalars(stmt).all())
  78. def list_past_deadline(
  79. self,
  80. now: datetime,
  81. *,
  82. limit: int = 100,
  83. ) -> list[PipelineRun]:
  84. stmt = (
  85. select(PipelineRun)
  86. .where(PipelineRun.status.in_(("queued", "running")))
  87. .where(PipelineRun.deadline_at <= now)
  88. .order_by(PipelineRun.deadline_at)
  89. .with_for_update(skip_locked=True)
  90. .limit(limit)
  91. )
  92. return list(self.session.scalars(stmt).all())
  93. def list_blocked_previous(self, *, limit: int = 100) -> list[PipelineRun]:
  94. stmt = (
  95. select(PipelineRun)
  96. .where(PipelineRun.status == "blocked_previous_run")
  97. .order_by(PipelineRun.biz_dt, PipelineRun.created_at)
  98. .with_for_update(skip_locked=True)
  99. .limit(limit)
  100. )
  101. return list(self.session.scalars(stmt).all())
  102. def mark_running(
  103. self,
  104. run: PipelineRun,
  105. *,
  106. step_key: str,
  107. owner: str,
  108. now: datetime,
  109. lease_until: datetime,
  110. ) -> None:
  111. if run.started_at is None:
  112. run.started_at = now
  113. run.status = "running"
  114. run.current_step = step_key
  115. run.heartbeat_at = now
  116. run.lease_owner = owner
  117. run.lease_until = lease_until
  118. self.session.flush()
  119. def finish(
  120. self,
  121. run: PipelineRun,
  122. *,
  123. status: str,
  124. now: datetime,
  125. summary: dict[str, Any] | None = None,
  126. error_code: str | None = None,
  127. error_message: str | None = None,
  128. ) -> None:
  129. run.status = status
  130. run.finished_at = now
  131. run.heartbeat_at = now
  132. run.lease_owner = None
  133. run.lease_until = None
  134. run.current_step = None
  135. run.summary_json = summary
  136. run.error_code = error_code
  137. run.error_message = error_message
  138. self.session.flush()
  139. def heartbeat(
  140. self,
  141. run_id: str,
  142. *,
  143. owner: str,
  144. now: datetime,
  145. lease_until: datetime,
  146. ) -> bool:
  147. run = self.get(run_id, for_update=True)
  148. if run is None or run.lease_owner != owner or run.status != "running":
  149. return False
  150. run.heartbeat_at = now
  151. run.lease_until = lease_until
  152. self.session.flush()
  153. return True