pipeline_lock_repo.py 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138
  1. from __future__ import annotations
  2. from datetime import datetime, timedelta
  3. from sqlalchemy import func, select
  4. from supply_infra.db.models.pipeline_lock import PipelineLock
  5. from supply_infra.db.repositories.base import BaseRepository
  6. CONTROL_PLANE_GUARD_KEY = "pipeline:claim_guard"
  7. class PipelineLockRepository(BaseRepository[PipelineLock]):
  8. model = PipelineLock
  9. def get(self, lock_key: str) -> PipelineLock | None:
  10. return self.session.get(PipelineLock, lock_key)
  11. def count_active_with_prefix(
  12. self,
  13. lock_key_prefix: str,
  14. *,
  15. now: datetime,
  16. ) -> int:
  17. count = self.session.scalar(
  18. select(func.count())
  19. .select_from(PipelineLock)
  20. .where(
  21. PipelineLock.lock_key.startswith(lock_key_prefix),
  22. PipelineLock.lease_until > now,
  23. )
  24. )
  25. return int(count or 0)
  26. def lock_control_plane(self, *, now: datetime) -> PipelineLock:
  27. """Lock the singleton row used for atomic scheduling decisions."""
  28. lock = self.session.scalar(
  29. select(PipelineLock)
  30. .where(PipelineLock.lock_key == CONTROL_PLANE_GUARD_KEY)
  31. .with_for_update()
  32. )
  33. if lock is None:
  34. # Production gets this row from Alembic. This fallback supports
  35. # local databases created from SQLAlchemy metadata.
  36. lock = PipelineLock(
  37. lock_key=CONTROL_PLANE_GUARD_KEY,
  38. owner_run_id="control-plane",
  39. owner_instance="claim-coordinator",
  40. lease_until=now,
  41. heartbeat_at=now,
  42. version=1,
  43. )
  44. self.session.add(lock)
  45. self.session.flush()
  46. return lock
  47. def touch(
  48. self,
  49. *,
  50. lock_key: str,
  51. owner: str,
  52. lease_seconds: int,
  53. now: datetime,
  54. ) -> PipelineLock:
  55. lock = self.session.scalar(
  56. select(PipelineLock)
  57. .where(PipelineLock.lock_key == lock_key)
  58. .with_for_update()
  59. )
  60. lease_until = now + timedelta(seconds=lease_seconds)
  61. if lock is None:
  62. lock = PipelineLock(
  63. lock_key=lock_key,
  64. owner_run_id="control-plane",
  65. owner_instance=owner,
  66. lease_until=lease_until,
  67. heartbeat_at=now,
  68. version=1,
  69. )
  70. self.session.add(lock)
  71. else:
  72. lock.owner_instance = owner
  73. lock.lease_until = lease_until
  74. lock.heartbeat_at = now
  75. lock.version += 1
  76. self.session.flush()
  77. return lock
  78. def acquire(
  79. self,
  80. *,
  81. lock_key: str,
  82. run_id: str,
  83. owner: str,
  84. lease_seconds: int,
  85. now: datetime,
  86. ) -> bool:
  87. stmt = (
  88. select(PipelineLock)
  89. .where(PipelineLock.lock_key == lock_key)
  90. .with_for_update()
  91. )
  92. lock = self.session.scalar(stmt)
  93. lease_until = now + timedelta(seconds=lease_seconds)
  94. if lock is None:
  95. self.add(
  96. PipelineLock(
  97. lock_key=lock_key,
  98. owner_run_id=run_id,
  99. owner_instance=owner,
  100. lease_until=lease_until,
  101. heartbeat_at=now,
  102. version=1,
  103. )
  104. )
  105. return True
  106. if lock.lease_until >= now and lock.owner_run_id != run_id:
  107. return False
  108. lock.owner_run_id = run_id
  109. lock.owner_instance = owner
  110. lock.lease_until = lease_until
  111. lock.heartbeat_at = now
  112. lock.version += 1
  113. self.session.flush()
  114. return True
  115. def release(self, lock_key: str, *, run_id: str) -> bool:
  116. stmt = (
  117. select(PipelineLock)
  118. .where(PipelineLock.lock_key == lock_key)
  119. .with_for_update()
  120. )
  121. lock = self.session.scalar(stmt)
  122. if lock is None or lock.owner_run_id != run_id:
  123. return False
  124. self.session.delete(lock)
  125. self.session.flush()
  126. return True