| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170 |
- from __future__ import annotations
- from datetime import datetime
- from typing import Any
- from sqlalchemy import exists, 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
- class PipelineRunRepository(BaseRepository[PipelineRun]):
- model = PipelineRun
- def get(self, run_id: str, *, for_update: bool = False) -> PipelineRun | None:
- stmt = select(PipelineRun).where(PipelineRun.run_id == run_id)
- if for_update:
- stmt = stmt.with_for_update()
- return self.session.scalar(stmt)
- def get_by_dedupe_key(self, dedupe_key: str) -> PipelineRun | None:
- stmt = select(PipelineRun).where(PipelineRun.dedupe_key == dedupe_key)
- return self.session.scalar(stmt)
- def get_for_business_day(
- self,
- *,
- pipeline_key: str,
- biz_dt: str,
- ) -> PipelineRun | None:
- stmt = (
- select(PipelineRun)
- .where(PipelineRun.pipeline_key == pipeline_key)
- .where(PipelineRun.biz_dt == biz_dt)
- .order_by(PipelineRun.created_at.desc())
- .limit(1)
- )
- return self.session.scalar(stmt)
- def create(self, values: dict[str, Any]) -> PipelineRun:
- return self.add(PipelineRun(**values))
- def list_recent(
- self,
- *,
- limit: int = 50,
- status: str | None = None,
- biz_dt: str | None = None,
- ) -> list[PipelineRun]:
- stmt = select(PipelineRun)
- if status:
- stmt = stmt.where(PipelineRun.status == status)
- if biz_dt:
- stmt = stmt.where(PipelineRun.biz_dt == biz_dt)
- stmt = stmt.order_by(PipelineRun.created_at.desc()).limit(limit)
- return list(self.session.scalars(stmt).all())
- def get_latest_nonterminal(self) -> PipelineRun | None:
- has_running_step = exists(
- select(PipelineStepRun.step_run_id)
- .where(PipelineStepRun.run_id == PipelineRun.run_id)
- .where(PipelineStepRun.status == "running")
- )
- stmt = (
- select(PipelineRun)
- .where(
- PipelineRun.status.in_(
- ("queued", "running", "cancelling", "blocked_previous_run")
- )
- | has_running_step
- )
- .order_by(PipelineRun.created_at.desc())
- .limit(1)
- )
- return self.session.scalar(stmt)
- def list_nonterminal(self, *, limit: int = 100) -> list[PipelineRun]:
- stmt = (
- select(PipelineRun)
- .where(
- PipelineRun.status.in_(
- ("queued", "running", "cancelling", "blocked_previous_run")
- )
- )
- .order_by(PipelineRun.created_at)
- .limit(limit)
- )
- return list(self.session.scalars(stmt).all())
- def list_past_deadline(
- self,
- now: datetime,
- *,
- limit: int = 100,
- ) -> list[PipelineRun]:
- stmt = (
- select(PipelineRun)
- .where(PipelineRun.status.in_(("queued", "running")))
- .where(PipelineRun.deadline_at <= now)
- .order_by(PipelineRun.deadline_at)
- .with_for_update(skip_locked=True)
- .limit(limit)
- )
- return list(self.session.scalars(stmt).all())
- def list_blocked_previous(self, *, limit: int = 100) -> list[PipelineRun]:
- stmt = (
- select(PipelineRun)
- .where(PipelineRun.status == "blocked_previous_run")
- .order_by(PipelineRun.biz_dt, PipelineRun.created_at)
- .with_for_update(skip_locked=True)
- .limit(limit)
- )
- return list(self.session.scalars(stmt).all())
- def mark_running(
- self,
- run: PipelineRun,
- *,
- step_key: str,
- owner: str,
- now: datetime,
- lease_until: datetime,
- ) -> None:
- if run.started_at is None:
- run.started_at = now
- run.status = "running"
- run.current_step = step_key
- run.heartbeat_at = now
- run.lease_owner = owner
- run.lease_until = lease_until
- self.session.flush()
- def finish(
- self,
- run: PipelineRun,
- *,
- status: str,
- now: datetime,
- summary: dict[str, Any] | None = None,
- error_code: str | None = None,
- error_message: str | None = None,
- ) -> None:
- run.status = status
- run.finished_at = now
- run.heartbeat_at = now
- run.lease_owner = None
- run.lease_until = None
- run.current_step = None
- run.summary_json = summary
- run.error_code = error_code
- run.error_message = error_message
- self.session.flush()
- def heartbeat(
- self,
- run_id: str,
- *,
- owner: str,
- now: datetime,
- lease_until: datetime,
- ) -> bool:
- run = self.get(run_id, for_update=True)
- if run is None or run.lease_owner != owner or run.status != "running":
- return False
- run.heartbeat_at = now
- run.lease_until = lease_until
- self.session.flush()
- return True
|