| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170 |
- from __future__ import annotations
- import logging
- import os
- import signal
- import socket
- import threading
- import uuid
- from dataclasses import dataclass
- from datetime import timedelta
- from supply_infra.config import InfraSettings, get_infra_settings
- from supply_infra.db.repositories.pipeline_run_repo import PipelineRunRepository
- from supply_infra.db.repositories.pipeline_step_run_repo import PipelineStepRunRepository
- from supply_infra.db.session import dispose_engine, get_session
- from supply_infra.pipeline.dates import china_now
- from supply_infra.pipeline.orchestrator import complete_step
- from supply_infra.pipeline.contracts import StepResult
- from supply_infra.pipeline.step_runner import StepExecutionRequest, run_step_subprocess
- logger = logging.getLogger(__name__)
- @dataclass(frozen=True)
- class ClaimedStep:
- run_id: str
- step_run_id: str
- step_key: str
- timeout_seconds: int
- class PipelineWorker:
- def __init__(
- self,
- *,
- settings: InfraSettings | None = None,
- worker_id: str | None = None,
- ) -> None:
- self.settings = settings or get_infra_settings()
- self.worker_id = worker_id or (
- f"{socket.gethostname()}:{os.getpid()}:{uuid.uuid4().hex[:8]}"
- )
- self._stop = threading.Event()
- def stop(self, *_args: object) -> None:
- logger.info("Worker stop requested: worker_id=%s", self.worker_id)
- self._stop.set()
- def claim(self) -> ClaimedStep | None:
- now = china_now()
- with get_session() as session:
- step = PipelineStepRunRepository(session).claim_next_ready(
- owner=self.worker_id,
- lease_seconds=self.settings.pipeline_lease_seconds,
- max_active_steps=self.settings.pipeline_max_active_steps,
- now=now,
- )
- if step is None:
- return None
- run = PipelineRunRepository(session).get(step.run_id, for_update=True)
- if run is None:
- raise ValueError(f"pipeline run not found: {step.run_id}")
- PipelineRunRepository(session).mark_running(
- run,
- step_key=step.step_key,
- owner=self.worker_id,
- now=now,
- lease_until=now + timedelta(seconds=self.settings.pipeline_lease_seconds),
- )
- return ClaimedStep(
- run_id=step.run_id,
- step_run_id=step.step_run_id,
- step_key=step.step_key,
- timeout_seconds=step.timeout_seconds,
- )
- def _heartbeat_loop(self, claimed: ClaimedStep, stopped: threading.Event) -> None:
- while not stopped.wait(self.settings.pipeline_heartbeat_seconds):
- now = china_now()
- lease_until = now + timedelta(seconds=self.settings.pipeline_lease_seconds)
- try:
- with get_session() as session:
- step_ok = PipelineStepRunRepository(session).heartbeat(
- claimed.step_run_id,
- owner=self.worker_id,
- now=now,
- lease_until=lease_until,
- )
- run_ok = PipelineRunRepository(session).heartbeat(
- claimed.run_id,
- owner=self.worker_id,
- now=now,
- lease_until=lease_until,
- )
- if not step_ok or not run_ok:
- logger.error(
- "Pipeline lease lost: worker=%s step=%s",
- self.worker_id,
- claimed.step_run_id,
- )
- stopped.set()
- except Exception:
- logger.exception("Pipeline heartbeat failed: step=%s", claimed.step_run_id)
- def run_once(self) -> bool:
- claimed = self.claim()
- if claimed is None:
- return False
- heartbeat_stop = threading.Event()
- heartbeat = threading.Thread(
- target=self._heartbeat_loop,
- args=(claimed, heartbeat_stop),
- name=f"pipeline-heartbeat-{claimed.step_run_id[:8]}",
- daemon=True,
- )
- heartbeat.start()
- try:
- try:
- result = run_step_subprocess(
- StepExecutionRequest(
- run_id=claimed.run_id,
- step_run_id=claimed.step_run_id,
- step_key=claimed.step_key,
- timeout_seconds=claimed.timeout_seconds,
- ),
- settings=self.settings,
- )
- except Exception as exc:
- logger.exception("Unexpected step runner failure: %s", claimed.step_run_id)
- result = StepResult(
- success=False,
- payload={},
- error_code="step_runner_exception",
- error_message=str(exc),
- )
- finally:
- heartbeat_stop.set()
- heartbeat.join(timeout=2)
- try:
- with get_session() as session:
- state = complete_step(
- session,
- step_run_id=claimed.step_run_id,
- owner=self.worker_id,
- result=result,
- )
- except RuntimeError:
- logger.exception(
- "Discarding step result after lease loss: step=%s",
- claimed.step_run_id,
- )
- return True
- logger.info(
- "Pipeline step completed: step=%s state=%s",
- claimed.step_run_id,
- state,
- )
- return True
- def run_forever(self) -> None:
- signal.signal(signal.SIGTERM, self.stop)
- signal.signal(signal.SIGINT, self.stop)
- logger.info("Pipeline worker started: worker_id=%s", self.worker_id)
- try:
- while not self._stop.is_set():
- if not self.run_once():
- self._stop.wait(self.settings.pipeline_worker_poll_seconds)
- finally:
- dispose_engine()
- logger.info("Pipeline worker stopped: worker_id=%s", self.worker_id)
|