worker.py 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170
  1. from __future__ import annotations
  2. import logging
  3. import os
  4. import signal
  5. import socket
  6. import threading
  7. import uuid
  8. from dataclasses import dataclass
  9. from datetime import timedelta
  10. from supply_infra.config import InfraSettings, get_infra_settings
  11. from supply_infra.db.repositories.pipeline_run_repo import PipelineRunRepository
  12. from supply_infra.db.repositories.pipeline_step_run_repo import PipelineStepRunRepository
  13. from supply_infra.db.session import dispose_engine, get_session
  14. from supply_infra.pipeline.dates import china_now
  15. from supply_infra.pipeline.orchestrator import complete_step
  16. from supply_infra.pipeline.contracts import StepResult
  17. from supply_infra.pipeline.step_runner import StepExecutionRequest, run_step_subprocess
  18. logger = logging.getLogger(__name__)
  19. @dataclass(frozen=True)
  20. class ClaimedStep:
  21. run_id: str
  22. step_run_id: str
  23. step_key: str
  24. timeout_seconds: int
  25. class PipelineWorker:
  26. def __init__(
  27. self,
  28. *,
  29. settings: InfraSettings | None = None,
  30. worker_id: str | None = None,
  31. ) -> None:
  32. self.settings = settings or get_infra_settings()
  33. self.worker_id = worker_id or (
  34. f"{socket.gethostname()}:{os.getpid()}:{uuid.uuid4().hex[:8]}"
  35. )
  36. self._stop = threading.Event()
  37. def stop(self, *_args: object) -> None:
  38. logger.info("Worker stop requested: worker_id=%s", self.worker_id)
  39. self._stop.set()
  40. def claim(self) -> ClaimedStep | None:
  41. now = china_now()
  42. with get_session() as session:
  43. step = PipelineStepRunRepository(session).claim_next_ready(
  44. owner=self.worker_id,
  45. lease_seconds=self.settings.pipeline_lease_seconds,
  46. max_active_steps=self.settings.pipeline_max_active_steps,
  47. now=now,
  48. )
  49. if step is None:
  50. return None
  51. run = PipelineRunRepository(session).get(step.run_id, for_update=True)
  52. if run is None:
  53. raise ValueError(f"pipeline run not found: {step.run_id}")
  54. PipelineRunRepository(session).mark_running(
  55. run,
  56. step_key=step.step_key,
  57. owner=self.worker_id,
  58. now=now,
  59. lease_until=now + timedelta(seconds=self.settings.pipeline_lease_seconds),
  60. )
  61. return ClaimedStep(
  62. run_id=step.run_id,
  63. step_run_id=step.step_run_id,
  64. step_key=step.step_key,
  65. timeout_seconds=step.timeout_seconds,
  66. )
  67. def _heartbeat_loop(self, claimed: ClaimedStep, stopped: threading.Event) -> None:
  68. while not stopped.wait(self.settings.pipeline_heartbeat_seconds):
  69. now = china_now()
  70. lease_until = now + timedelta(seconds=self.settings.pipeline_lease_seconds)
  71. try:
  72. with get_session() as session:
  73. step_ok = PipelineStepRunRepository(session).heartbeat(
  74. claimed.step_run_id,
  75. owner=self.worker_id,
  76. now=now,
  77. lease_until=lease_until,
  78. )
  79. run_ok = PipelineRunRepository(session).heartbeat(
  80. claimed.run_id,
  81. owner=self.worker_id,
  82. now=now,
  83. lease_until=lease_until,
  84. )
  85. if not step_ok or not run_ok:
  86. logger.error(
  87. "Pipeline lease lost: worker=%s step=%s",
  88. self.worker_id,
  89. claimed.step_run_id,
  90. )
  91. stopped.set()
  92. except Exception:
  93. logger.exception("Pipeline heartbeat failed: step=%s", claimed.step_run_id)
  94. def run_once(self) -> bool:
  95. claimed = self.claim()
  96. if claimed is None:
  97. return False
  98. heartbeat_stop = threading.Event()
  99. heartbeat = threading.Thread(
  100. target=self._heartbeat_loop,
  101. args=(claimed, heartbeat_stop),
  102. name=f"pipeline-heartbeat-{claimed.step_run_id[:8]}",
  103. daemon=True,
  104. )
  105. heartbeat.start()
  106. try:
  107. try:
  108. result = run_step_subprocess(
  109. StepExecutionRequest(
  110. run_id=claimed.run_id,
  111. step_run_id=claimed.step_run_id,
  112. step_key=claimed.step_key,
  113. timeout_seconds=claimed.timeout_seconds,
  114. ),
  115. settings=self.settings,
  116. )
  117. except Exception as exc:
  118. logger.exception("Unexpected step runner failure: %s", claimed.step_run_id)
  119. result = StepResult(
  120. success=False,
  121. payload={},
  122. error_code="step_runner_exception",
  123. error_message=str(exc),
  124. )
  125. finally:
  126. heartbeat_stop.set()
  127. heartbeat.join(timeout=2)
  128. try:
  129. with get_session() as session:
  130. state = complete_step(
  131. session,
  132. step_run_id=claimed.step_run_id,
  133. owner=self.worker_id,
  134. result=result,
  135. )
  136. except RuntimeError:
  137. logger.exception(
  138. "Discarding step result after lease loss: step=%s",
  139. claimed.step_run_id,
  140. )
  141. return True
  142. logger.info(
  143. "Pipeline step completed: step=%s state=%s",
  144. claimed.step_run_id,
  145. state,
  146. )
  147. return True
  148. def run_forever(self) -> None:
  149. signal.signal(signal.SIGTERM, self.stop)
  150. signal.signal(signal.SIGINT, self.stop)
  151. logger.info("Pipeline worker started: worker_id=%s", self.worker_id)
  152. try:
  153. while not self._stop.is_set():
  154. if not self.run_once():
  155. self._stop.wait(self.settings.pipeline_worker_poll_seconds)
  156. finally:
  157. dispose_engine()
  158. logger.info("Pipeline worker stopped: worker_id=%s", self.worker_id)