run_service.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380
  1. from __future__ import annotations
  2. import os
  3. import uuid
  4. from dataclasses import dataclass
  5. from datetime import datetime
  6. from typing import Any
  7. from sqlalchemy.exc import IntegrityError
  8. from supply_infra.config import InfraSettings, get_infra_settings
  9. from supply_infra.db.models.pipeline_run import PipelineRun
  10. from supply_infra.db.models.pipeline_step_run import PipelineStepRun
  11. from supply_infra.db.repositories.pipeline_lock_repo import PipelineLockRepository
  12. from supply_infra.db.repositories.pipeline_outbox_repo import PipelineOutboxRepository
  13. from supply_infra.db.repositories.pipeline_run_repo import PipelineRunRepository
  14. from supply_infra.db.repositories.pipeline_step_run_repo import PipelineStepRunRepository
  15. from supply_infra.db.session import get_session
  16. from supply_infra.pipeline.dag import PIPELINE_STEPS
  17. from supply_infra.pipeline.dates import (
  18. build_date_snapshot,
  19. CHINA_TIMEZONE,
  20. china_now,
  21. deadline_for_trigger,
  22. next_schedule_china,
  23. resolve_biz_dt,
  24. )
  25. from supply_infra.pipeline.enums import RunStatus, StepStatus
  26. PIPELINE_KEY = "supply_pipeline"
  27. @dataclass(frozen=True)
  28. class RunSubmission:
  29. run_id: str
  30. created: bool
  31. status: str
  32. biz_dt: str
  33. dry_run: bool = False
  34. def to_dict(self) -> dict[str, Any]:
  35. return {
  36. "accepted": True,
  37. "run_id": self.run_id,
  38. "created": self.created,
  39. "status": self.status,
  40. "biz_dt": self.biz_dt,
  41. "dry_run": self.dry_run,
  42. }
  43. def _config_snapshot(settings: InfraSettings) -> dict[str, Any]:
  44. return {
  45. "scheduler_timezone": settings.scheduler_timezone,
  46. "scheduler_cron_hour": settings.scheduler_cron_hour,
  47. "scheduler_cron_minute": settings.scheduler_cron_minute,
  48. "lease_seconds": settings.pipeline_lease_seconds,
  49. "heartbeat_seconds": settings.pipeline_heartbeat_seconds,
  50. "worker_processes": settings.pipeline_worker_processes,
  51. "max_active_steps": settings.pipeline_max_active_steps,
  52. "mysql_connection_budget": settings.mysql_connection_budget,
  53. }
  54. def _code_version() -> str | None:
  55. return os.getenv("BUILD_TIMESTAMP") or os.getenv("GIT_COMMIT")
  56. def submit_pipeline_run(
  57. *,
  58. biz_dt: str | None = None,
  59. trigger_type: str = "api",
  60. trigger_source: str | None = None,
  61. trigger_reason: str | None = None,
  62. settings: InfraSettings | None = None,
  63. ) -> RunSubmission:
  64. active = settings or get_infra_settings()
  65. resolved_biz_dt = resolve_biz_dt(biz_dt, settings=active)
  66. dedupe_key = f"{PIPELINE_KEY}:{resolved_biz_dt}:full"
  67. run_id = str(uuid.uuid4())
  68. date_snapshot = build_date_snapshot(resolved_biz_dt)
  69. try:
  70. with get_session() as session:
  71. run_repo = PipelineRunRepository(session)
  72. PipelineLockRepository(session).lock_control_plane(now=china_now())
  73. existing = run_repo.get_by_dedupe_key(dedupe_key)
  74. if existing is not None:
  75. return RunSubmission(
  76. run_id=existing.run_id,
  77. created=False,
  78. status=existing.status,
  79. biz_dt=existing.biz_dt,
  80. dry_run=False,
  81. )
  82. # All entry points share the same no-overlap rule. Linking each new
  83. # run to the latest nonterminal run also creates a safe FIFO chain
  84. # when several future batches are submitted in advance.
  85. previous = run_repo.get_latest_nonterminal()
  86. initial_status = (
  87. RunStatus.BLOCKED_PREVIOUS_RUN.value
  88. if previous is not None
  89. else RunStatus.QUEUED.value
  90. )
  91. run = run_repo.create(
  92. {
  93. "run_id": run_id,
  94. "dedupe_key": dedupe_key,
  95. "pipeline_key": PIPELINE_KEY,
  96. "biz_dt": resolved_biz_dt,
  97. "trigger_type": trigger_type,
  98. "trigger_source": trigger_source or trigger_type,
  99. "trigger_reason": trigger_reason,
  100. "run_mode": "full",
  101. "dry_run": False,
  102. "status": initial_status,
  103. "deadline_at": deadline_for_trigger(
  104. trigger_type,
  105. settings=active,
  106. ),
  107. "code_version": _code_version(),
  108. "config_snapshot_json": _config_snapshot(active),
  109. "date_snapshot_json": date_snapshot,
  110. "summary_json": (
  111. {"blocked_by_run_id": previous.run_id}
  112. if previous is not None
  113. else None
  114. ),
  115. }
  116. )
  117. rows: list[dict[str, Any]] = []
  118. for index, step in enumerate(PIPELINE_STEPS):
  119. if previous is not None:
  120. status = StepStatus.BLOCKED.value
  121. else:
  122. status = (
  123. StepStatus.READY.value if index == 0 else StepStatus.PENDING.value
  124. )
  125. rows.append(
  126. {
  127. "step_run_id": str(uuid.uuid4()),
  128. "run_id": run.run_id,
  129. "step_key": step.key,
  130. "step_order": step.order,
  131. "attempt": 1,
  132. "status": status,
  133. "critical": step.critical,
  134. "dependency_snapshot_json": list(step.dependencies),
  135. "input_snapshot_json": {},
  136. "timeout_seconds": step.timeout_seconds,
  137. "max_attempts": step.max_attempts,
  138. "retryable": step.retryable,
  139. "error_code": (
  140. "previous_run_active" if previous is not None else None
  141. ),
  142. "error_message": (
  143. f"Blocked by previous run {previous.run_id}"
  144. if previous is not None
  145. else None
  146. ),
  147. }
  148. )
  149. PipelineStepRunRepository(session).create_steps(rows)
  150. return RunSubmission(
  151. run_id=run.run_id,
  152. created=True,
  153. status=run.status,
  154. biz_dt=run.biz_dt,
  155. dry_run=False,
  156. )
  157. except IntegrityError:
  158. with get_session() as session:
  159. existing = PipelineRunRepository(session).get_by_dedupe_key(dedupe_key)
  160. if existing is None:
  161. raise
  162. return RunSubmission(
  163. run_id=existing.run_id,
  164. created=False,
  165. status=existing.status,
  166. biz_dt=existing.biz_dt,
  167. dry_run=False,
  168. )
  169. def get_pipeline_run(run_id: str) -> dict[str, Any] | None:
  170. with get_session() as session:
  171. run = PipelineRunRepository(session).get(run_id)
  172. if run is None:
  173. return None
  174. steps = PipelineStepRunRepository(session).list_for_run(run_id)
  175. outbox = PipelineOutboxRepository(session).list_for_run(run_id)
  176. return {
  177. **serialize_run(run),
  178. "steps": [serialize_step(item) for item in steps],
  179. "effects": [
  180. {
  181. "outbox_id": item.outbox_id,
  182. "step_run_id": item.step_run_id,
  183. "effect_type": item.effect_type,
  184. "payload_hash": item.payload_hash,
  185. "payload_uri": item.payload_uri,
  186. "status": item.status,
  187. "dry_run": item.dry_run,
  188. "created_at": _iso(item.created_at),
  189. }
  190. for item in outbox
  191. ],
  192. }
  193. def list_pipeline_runs(
  194. *,
  195. limit: int = 50,
  196. status: str | None = None,
  197. biz_dt: str | None = None,
  198. ) -> list[dict[str, Any]]:
  199. with get_session() as session:
  200. rows = PipelineRunRepository(session).list_recent(
  201. limit=min(max(limit, 1), 200),
  202. status=status,
  203. biz_dt=biz_dt,
  204. )
  205. return [serialize_run(item) for item in rows]
  206. def cancel_pipeline_run(run_id: str) -> bool:
  207. now = china_now()
  208. with get_session() as session:
  209. run_repo = PipelineRunRepository(session)
  210. step_repo = PipelineStepRunRepository(session)
  211. run = run_repo.get(run_id, for_update=True)
  212. if run is None:
  213. return False
  214. if run.status in {
  215. RunStatus.SUCCEEDED.value,
  216. RunStatus.FAILED.value,
  217. RunStatus.CANCELLED.value,
  218. RunStatus.DEADLINE_EXCEEDED.value,
  219. }:
  220. return True
  221. has_running_step = False
  222. for step in step_repo.list_for_run(run_id):
  223. if step.status == StepStatus.RUNNING.value:
  224. has_running_step = True
  225. continue
  226. if step.status in {
  227. StepStatus.PENDING.value,
  228. StepStatus.READY.value,
  229. StepStatus.RETRY_WAIT.value,
  230. StepStatus.BLOCKED.value,
  231. }:
  232. step.status = StepStatus.CANCELLED.value
  233. step.finished_at = now
  234. if has_running_step:
  235. run.status = RunStatus.CANCELLING.value
  236. run.error_code = "cancellation_requested"
  237. run.error_message = "Waiting for the running step to stop"
  238. run.heartbeat_at = now
  239. run.lease_owner = None
  240. run.lease_until = None
  241. else:
  242. run_repo.finish(run, status=RunStatus.CANCELLED.value, now=now)
  243. return True
  244. def resume_pipeline_run(run_id: str, *, step_key: str | None = None) -> bool:
  245. with get_session() as session:
  246. run_repo = PipelineRunRepository(session)
  247. step_repo = PipelineStepRunRepository(session)
  248. run = run_repo.get(run_id, for_update=True)
  249. if run is None:
  250. return False
  251. if run.status not in {
  252. RunStatus.FAILED.value,
  253. RunStatus.CANCELLED.value,
  254. RunStatus.DEADLINE_EXCEEDED.value,
  255. }:
  256. return False
  257. latest = step_repo.latest_attempts(run_id)
  258. resumable_statuses = (
  259. {StepStatus.CANCELLED.value}
  260. if run.status == RunStatus.CANCELLED.value
  261. else {
  262. StepStatus.FAILED.value,
  263. StepStatus.TIMED_OUT.value,
  264. StepStatus.BLOCKED.value,
  265. }
  266. )
  267. resume_from = next(
  268. (
  269. item
  270. for item in sorted(latest.values(), key=lambda row: row.step_order)
  271. if item.status in resumable_statuses
  272. ),
  273. None,
  274. )
  275. if resume_from is None:
  276. return False
  277. if step_key is not None and resume_from.step_key != step_key:
  278. return False
  279. retry = step_repo.create_retry_attempt(
  280. resume_from,
  281. status=StepStatus.READY.value,
  282. )
  283. for item in latest.values():
  284. if item.step_order > resume_from.step_order and item.status in {
  285. StepStatus.BLOCKED.value,
  286. StepStatus.CANCELLED.value,
  287. }:
  288. item.status = StepStatus.PENDING.value
  289. item.error_code = None
  290. item.error_message = None
  291. run.status = RunStatus.QUEUED.value
  292. run.current_step = retry.step_key
  293. if run.deadline_at is not None and run.deadline_at <= china_now():
  294. run.deadline_at = next_schedule_china()
  295. run.finished_at = None
  296. run.summary_json = None
  297. run.error_code = None
  298. run.error_message = None
  299. return True
  300. def retry_pipeline_step(run_id: str, step_key: str) -> bool:
  301. return resume_pipeline_run(run_id, step_key=step_key)
  302. def serialize_run(run: PipelineRun) -> dict[str, Any]:
  303. return {
  304. "run_id": run.run_id,
  305. "pipeline_key": run.pipeline_key,
  306. "biz_dt": run.biz_dt,
  307. "trigger_type": run.trigger_type,
  308. "trigger_source": run.trigger_source,
  309. "run_mode": run.run_mode,
  310. "dry_run": run.dry_run,
  311. "status": run.status,
  312. "current_step": run.current_step,
  313. "deadline_at": _iso(run.deadline_at),
  314. "started_at": _iso(run.started_at),
  315. "finished_at": _iso(run.finished_at),
  316. "heartbeat_at": _iso(run.heartbeat_at),
  317. "summary": run.summary_json,
  318. "error_code": run.error_code,
  319. "error_message": run.error_message,
  320. "created_at": _iso(run.created_at),
  321. "updated_at": _iso(run.updated_at),
  322. }
  323. def serialize_step(step: PipelineStepRun) -> dict[str, Any]:
  324. return {
  325. "step_run_id": step.step_run_id,
  326. "run_id": step.run_id,
  327. "step_key": step.step_key,
  328. "step_order": step.step_order,
  329. "attempt": step.attempt,
  330. "status": step.status,
  331. "critical": step.critical,
  332. "timeout_seconds": step.timeout_seconds,
  333. "started_at": _iso(step.started_at),
  334. "finished_at": _iso(step.finished_at),
  335. "heartbeat_at": _iso(step.heartbeat_at),
  336. "result": step.result_summary_json,
  337. "error_code": step.error_code,
  338. "error_message": step.error_message,
  339. "log_uri": step.log_uri,
  340. }
  341. def _iso(value: datetime | None) -> str | None:
  342. if value is None:
  343. return None
  344. aware = (
  345. value.replace(tzinfo=CHINA_TIMEZONE)
  346. if value.tzinfo is None
  347. else value.astimezone(CHINA_TIMEZONE)
  348. )
  349. return aware.isoformat()