| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242 |
- from __future__ import annotations
- import uuid
- from datetime import datetime, timedelta
- from typing import Any
- from sqlalchemy import func, select
- from supply_infra.db.models.pipeline_run import PipelineRun
- from supply_infra.db.models.pipeline_step_run import PipelineStepRun
- from supply_infra.db.repositories.base import BaseRepository
- from supply_infra.db.repositories.pipeline_lock_repo import PipelineLockRepository
- class PipelineStepRunRepository(BaseRepository[PipelineStepRun]):
- model = PipelineStepRun
- def get(
- self,
- step_run_id: str,
- *,
- for_update: bool = False,
- ) -> PipelineStepRun | None:
- stmt = select(PipelineStepRun).where(
- PipelineStepRun.step_run_id == step_run_id
- )
- if for_update:
- stmt = stmt.with_for_update()
- return self.session.scalar(stmt)
- def create_steps(self, rows: list[dict[str, Any]]) -> list[PipelineStepRun]:
- entities = [PipelineStepRun(**row) for row in rows]
- self.session.add_all(entities)
- self.session.flush()
- return entities
- def list_for_run(self, run_id: str) -> list[PipelineStepRun]:
- stmt = (
- select(PipelineStepRun)
- .where(PipelineStepRun.run_id == run_id)
- .order_by(PipelineStepRun.step_order, PipelineStepRun.attempt)
- )
- return list(self.session.scalars(stmt).all())
- def latest_attempts(self, run_id: str) -> dict[str, PipelineStepRun]:
- latest: dict[str, PipelineStepRun] = {}
- for item in self.list_for_run(run_id):
- latest[item.step_key] = item
- return latest
- def claim_next_ready(
- self,
- *,
- owner: str,
- lease_seconds: int,
- max_active_steps: int,
- now: datetime,
- ) -> PipelineStepRun | None:
- # Serialize count-and-claim so concurrent workers cannot oversubscribe
- # the configured global active-step cap.
- PipelineLockRepository(self.session).lock_control_plane(now=now)
- active_steps = self.session.scalar(
- select(func.count())
- .select_from(PipelineStepRun)
- .where(PipelineStepRun.status == "running")
- )
- if int(active_steps or 0) >= max_active_steps:
- return None
- running_run_ids = list(
- self.session.scalars(
- select(PipelineStepRun.run_id)
- .where(PipelineStepRun.status == "running")
- .distinct()
- ).all()
- )
- stmt = (
- select(PipelineStepRun)
- .where(PipelineStepRun.status == "ready")
- .where(
- (PipelineStepRun.next_retry_at.is_(None))
- | (PipelineStepRun.next_retry_at <= now)
- )
- .order_by(PipelineStepRun.created_at, PipelineStepRun.step_order)
- .with_for_update(skip_locked=True)
- .limit(1)
- )
- if running_run_ids:
- # A daily run may contain parallel steps in the future, but a
- # different run must never overlap while an old process is alive.
- stmt = stmt.where(PipelineStepRun.run_id.in_(running_run_ids))
- step = self.session.scalar(stmt)
- if step is None:
- return None
- step.status = "running"
- step.started_at = step.started_at or now
- step.heartbeat_at = now
- step.lease_owner = owner
- step.lease_until = now + timedelta(seconds=lease_seconds)
- self.session.flush()
- return step
- def has_running_for_run(self, run_id: str) -> bool:
- count = self.session.scalar(
- select(func.count())
- .select_from(PipelineStepRun)
- .where(PipelineStepRun.run_id == run_id)
- .where(PipelineStepRun.status == "running")
- )
- return bool(count)
- def heartbeat(
- self,
- step_run_id: str,
- *,
- owner: str,
- now: datetime,
- lease_until: datetime,
- ) -> bool:
- step = self.get(step_run_id, for_update=True)
- if step is None or step.status != "running" or step.lease_owner != owner:
- return False
- step.heartbeat_at = now
- step.lease_until = lease_until
- self.session.flush()
- return True
- def finish(
- self,
- step: PipelineStepRun,
- *,
- status: str,
- now: datetime,
- exit_code: int | None,
- result: dict[str, Any] | None,
- error_code: str | None,
- error_message: str | None,
- log_uri: str | None,
- ) -> None:
- step.status = status
- step.finished_at = now
- step.heartbeat_at = now
- step.lease_owner = None
- step.lease_until = None
- step.exit_code = exit_code
- step.result_summary_json = result
- step.error_code = error_code
- step.error_message = error_message
- step.log_uri = log_uri
- self.session.flush()
- def create_retry_attempt(
- self,
- previous: PipelineStepRun,
- *,
- status: str = "ready",
- next_retry_at: datetime | None = None,
- ) -> PipelineStepRun:
- entity = PipelineStepRun(
- step_run_id=str(uuid.uuid4()),
- run_id=previous.run_id,
- step_key=previous.step_key,
- step_order=previous.step_order,
- attempt=previous.attempt + 1,
- status=status,
- critical=previous.critical,
- dependency_snapshot_json=list(previous.dependency_snapshot_json or []),
- input_snapshot_json=dict(previous.input_snapshot_json or {}),
- timeout_seconds=previous.timeout_seconds,
- max_attempts=previous.max_attempts,
- retryable=previous.retryable,
- next_retry_at=next_retry_at,
- )
- return self.add(entity)
- def list_expired_running(self, now: datetime, *, limit: int = 100) -> list[PipelineStepRun]:
- stmt = (
- select(PipelineStepRun)
- .where(PipelineStepRun.status == "running")
- .where(PipelineStepRun.lease_until.is_not(None))
- .where(PipelineStepRun.lease_until < now)
- .order_by(PipelineStepRun.lease_until)
- .with_for_update(skip_locked=True)
- .limit(limit)
- )
- return list(self.session.scalars(stmt).all())
- def promote_due_retries(self, now: datetime, *, limit: int = 100) -> int:
- stmt = (
- select(PipelineStepRun)
- .where(PipelineStepRun.status == "retry_wait")
- .where(PipelineStepRun.next_retry_at <= now)
- .order_by(PipelineStepRun.next_retry_at)
- .with_for_update(skip_locked=True)
- .limit(limit)
- )
- count = 0
- for step in self.session.scalars(stmt).all():
- step.status = "ready"
- count += 1
- self.session.flush()
- return count
- def block_unfinished_downstream(
- self,
- *,
- run_id: str,
- after_order: int,
- reason: str,
- ) -> int:
- count = 0
- for step in self.list_for_run(run_id):
- if step.step_order <= after_order or step.status in {
- "succeeded",
- "failed",
- "timed_out",
- "cancelled",
- }:
- continue
- step.status = "blocked"
- step.error_code = "upstream_failed"
- step.error_message = reason
- count += 1
- self.session.flush()
- return count
- def get_latest_succeeded_biz_dt(self, step_key: str) -> str | None:
- """返回指定步骤最近一次成功完成时所属 pipeline_run 的 biz_dt。"""
- stmt = (
- select(PipelineRun.biz_dt)
- .join(PipelineStepRun, PipelineStepRun.run_id == PipelineRun.run_id)
- .where(
- PipelineStepRun.step_key == step_key,
- PipelineStepRun.status == "succeeded",
- PipelineStepRun.finished_at.is_not(None),
- )
- .order_by(PipelineStepRun.finished_at.desc())
- .limit(1)
- )
- value = self.session.scalar(stmt)
- return str(value) if value else None
|