| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138 |
- from __future__ import annotations
- from datetime import datetime, timedelta
- from sqlalchemy import func, select
- from supply_infra.db.models.pipeline_lock import PipelineLock
- from supply_infra.db.repositories.base import BaseRepository
- CONTROL_PLANE_GUARD_KEY = "pipeline:claim_guard"
- class PipelineLockRepository(BaseRepository[PipelineLock]):
- model = PipelineLock
- def get(self, lock_key: str) -> PipelineLock | None:
- return self.session.get(PipelineLock, lock_key)
- def count_active_with_prefix(
- self,
- lock_key_prefix: str,
- *,
- now: datetime,
- ) -> int:
- count = self.session.scalar(
- select(func.count())
- .select_from(PipelineLock)
- .where(
- PipelineLock.lock_key.startswith(lock_key_prefix),
- PipelineLock.lease_until > now,
- )
- )
- return int(count or 0)
- def lock_control_plane(self, *, now: datetime) -> PipelineLock:
- """Lock the singleton row used for atomic scheduling decisions."""
- lock = self.session.scalar(
- select(PipelineLock)
- .where(PipelineLock.lock_key == CONTROL_PLANE_GUARD_KEY)
- .with_for_update()
- )
- if lock is None:
- # Production gets this row from Alembic. This fallback supports
- # local databases created from SQLAlchemy metadata.
- lock = PipelineLock(
- lock_key=CONTROL_PLANE_GUARD_KEY,
- owner_run_id="control-plane",
- owner_instance="claim-coordinator",
- lease_until=now,
- heartbeat_at=now,
- version=1,
- )
- self.session.add(lock)
- self.session.flush()
- return lock
- def touch(
- self,
- *,
- lock_key: str,
- owner: str,
- lease_seconds: int,
- now: datetime,
- ) -> PipelineLock:
- lock = self.session.scalar(
- select(PipelineLock)
- .where(PipelineLock.lock_key == lock_key)
- .with_for_update()
- )
- lease_until = now + timedelta(seconds=lease_seconds)
- if lock is None:
- lock = PipelineLock(
- lock_key=lock_key,
- owner_run_id="control-plane",
- owner_instance=owner,
- lease_until=lease_until,
- heartbeat_at=now,
- version=1,
- )
- self.session.add(lock)
- else:
- lock.owner_instance = owner
- lock.lease_until = lease_until
- lock.heartbeat_at = now
- lock.version += 1
- self.session.flush()
- return lock
- def acquire(
- self,
- *,
- lock_key: str,
- run_id: str,
- owner: str,
- lease_seconds: int,
- now: datetime,
- ) -> bool:
- stmt = (
- select(PipelineLock)
- .where(PipelineLock.lock_key == lock_key)
- .with_for_update()
- )
- lock = self.session.scalar(stmt)
- lease_until = now + timedelta(seconds=lease_seconds)
- if lock is None:
- self.add(
- PipelineLock(
- lock_key=lock_key,
- owner_run_id=run_id,
- owner_instance=owner,
- lease_until=lease_until,
- heartbeat_at=now,
- version=1,
- )
- )
- return True
- if lock.lease_until >= now and lock.owner_run_id != run_id:
- return False
- lock.owner_run_id = run_id
- lock.owner_instance = owner
- lock.lease_until = lease_until
- lock.heartbeat_at = now
- lock.version += 1
- self.session.flush()
- return True
- def release(self, lock_key: str, *, run_id: str) -> bool:
- stmt = (
- select(PipelineLock)
- .where(PipelineLock.lock_key == lock_key)
- .with_for_update()
- )
- lock = self.session.scalar(stmt)
- if lock is None or lock.owner_run_id != run_id:
- return False
- self.session.delete(lock)
- self.session.flush()
- return True
|