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)