run_service.py 14 KB

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