Coverage for core / src / sensorkit / backend / lease.py: 95%
140 statements
« prev ^ index » next coverage.py v7.13.5, created at 2026-09-02 00:03 +0000
« prev ^ index » next coverage.py v7.13.5, created at 2026-09-02 00:03 +0000
1# SPDX-License-Identifier: Apache-2.0
2from __future__ import annotations
4import asyncio
5import contextlib
6import heapq
7import os
8import random
9from datetime import datetime, timedelta
10from typing import Callable
12from loguru import logger
13from pydantic import BaseModel
15from sensorkit.backend.base import BackendError, KeyValueContext, RevisionError
16from sensorkit.common.aio import scoped_waiter
19class LeaseModel[M: BaseModel | None](BaseModel):
20 """Serialisable metadata stored in the KV entry for an active lease."""
22 acquired_at: datetime
23 refreshed_at: datetime
24 record: M
27class LeaseUnavailableError(BackendError):
28 """Raised when a lease cannot be acquired because the key is already held."""
31class LeaseAbandonedError(BackendError):
32 """Raised when a lease is lost unexpectedly (TTL expiry or external deletion)."""
35class Lease[M: BaseModel | None]:
36 """A distributed, TTL-backed lock stored in the KV backend.
38 A Lease is obtained via Lease.acquire() and must be periodically refreshed (or managed
39 via a LeaseGroup) to prevent expiry. It can be used as an async context manager: the lease
40 is expired automatically on exit.
41 """
43 @classmethod
44 async def acquire(
45 cls,
46 kv: KeyValueContext,
47 key: str,
48 ttl: float,
49 record: M,
50 ):
51 """Attempt to create a new lease at the given KV key with the specified TTL.
53 Raises LeaseUnavailableError if the key is already held.
54 """
55 now = datetime.now()
56 model = LeaseModel.model_construct(
57 acquired_at=now,
58 refreshed_at=now,
59 record=record,
60 )
62 if os.environ.get("SENSORKIT_NO_ENFORCE_LEASES"):
63 await kv.delete(key)
65 try:
66 logger.debug(f"acquiring lease {kv.entity.subject(key)} {record=}")
67 entry = await kv.create(
68 key=key,
69 value=model.model_dump_json().encode(),
70 ttl=ttl,
71 )
72 return cls(kv, key, ttl, entry.revision, model)
73 except RevisionError as e:
74 raise LeaseUnavailableError() from e
76 def __init__(
77 self,
78 kv: KeyValueContext,
79 key: str,
80 ttl: float,
81 revision: int,
82 model: LeaseModel[M],
83 ):
84 self._kv = kv
85 self._key = key
86 self.ttl = ttl
87 self._revision = revision
88 self.model = model
89 self.expire_called = False
90 self._expired = asyncio.Event()
91 self._lock = asyncio.Lock()
92 self._task = asyncio.create_task(self._expiry_monitor())
94 async def _temp_monitor_lease(self, key: str, *, poll_wait=2.0):
95 while True:
96 yield await self._kv.get(key)
97 await asyncio.sleep(poll_wait)
99 async def _expiry_monitor(self):
100 if os.environ.get("SENSORKIT_NO_ENFORCE_LEASES"):
101 return
103 # FIXME: The NATS TTL saga continues: while nats-py v2.13 has added basic support for KV
104 # TTLs, this does not include the ability for watchers to detect KV expiry. So, here
105 # we employ yet another stopgap that polls the lease key instead.
106 monitor = self._temp_monitor_lease(self._key)
107 # monitor = await self._kv.monitor(self._key)
109 try:
110 async with asyncio.timeout(self.ttl) as timeout:
111 while not self.expired:
112 # Wait for a change to our KV entry.
113 entry = await anext(monitor)
115 if entry.deleted():
116 # We have lost the lease. Fallthrough to ensure we are marked as expired.
117 break
119 # Extend local timeout.
120 timeout.reschedule(asyncio.get_running_loop().time() + self.ttl)
121 except TimeoutError:
122 logger.warning(f"Abandoning lease for {self._kv.entity.subject(self._key)}!")
123 finally:
124 self._expired.set()
126 async def refresh(self, record: M | None = None):
127 """Extend the lease, optionally updating the associated record."""
128 if record:
129 self.model.record = record
131 self.model.refreshed_at = datetime.now()
132 entry = await self._kv.update(
133 self._key,
134 self.model.model_dump_json().encode(),
135 ttl=self.ttl,
136 revision=self._revision,
137 )
138 self._revision = entry.revision
140 async def expire(self):
141 """Expire the lease if it isn't already expired."""
142 async with self._lock:
143 if not self._expired.is_set():
144 self.expire_called = True
145 logger.debug(f"expiring lease for {self._kv.entity.subject(self._key)}")
147 async with asyncio.timeout(1.0):
148 await self._kv.delete(
149 self._key,
150 revision=self._revision,
151 )
153 self._expired.set()
154 self._task.cancel()
156 with contextlib.suppress(asyncio.CancelledError):
157 await self._task
159 @property
160 def expired(self):
161 """Return True if the lease is expired."""
162 return self._expired.is_set()
164 def wait_expired(self):
165 """Wait until lease expiry."""
166 return self._expired.wait()
168 async def __aenter__(self):
169 return self
171 async def __aexit__(self, exc_type, exc_val, exc_tb):
172 await self.expire()
175class LeaseGroup:
176 """Manage a group of Leases with shared expiration semantics."""
178 def __init__(self):
179 self._lease_waiters: dict[Lease, asyncio.Task] = {}
180 self._refresh_queue: list[tuple[float, Lease]] = []
181 self._lease_added = asyncio.Event()
183 async def acquire[M: BaseModel | None](
184 self,
185 kv: KeyValueContext,
186 key: str,
187 ttl: float,
188 record: M,
189 ) -> Lease[M]:
190 """Acquire a Lease and add it to the maintenance schedule."""
191 lease = await Lease.acquire(
192 kv,
193 key,
194 ttl,
195 record,
196 )
197 self._lease_waiters[lease] = asyncio.create_task(lease.wait_expired())
198 self._schedule_add(lease)
199 self._lease_added.set()
200 return lease
202 async def refresh_loop(self, record_callback: Callable[[Lease], None] | None = None):
203 """Yield Leases as their scheduled refresh time becomes current.
205 NOTE: This implementation requires that lease refreshes are performed serially. There are
206 potential performance pitfalls with this if using a high-latency backend (which, it
207 should be noted, is not a supported use case) or if servicing a large number of
208 entities such that a burst of coincident refreshes causes a delay beyond a lease TTL.
209 Lease refreshes are scheduled with a random element, which reduces the potential for
210 this happening, and TTLs can always be tuned longer. However, a callback-based
211 implementation where refreshes can be performed concurrently could be considered.
212 """
213 try:
214 while True:
215 lease = self._schedule_pop()
217 if lease and not lease.expired:
218 # Generate the user data if a callback was provided.
219 record = record_callback(lease) if record_callback else None
221 # Refresh and reschedule the lease.
222 await lease.refresh(record)
223 self._schedule_add(lease)
225 # Wait for a lease to expire, a new lease to be added, or until our next scheduled
226 # refresh. If the return value is False, an expiration occurred.
227 if not await self._wait_for_event():
228 abandoned = [
229 lease
230 for lease in self._lease_waiters
231 if lease.expired and not lease.expire_called
232 ]
233 break
234 finally:
235 # Expire all leases in the group on expiration of any lease or on error.
236 await self.expire()
238 if abandoned:
239 keys = ", ".join(lease._key for lease in abandoned)
240 raise LeaseAbandonedError(f"Lease(s) lost unexpectedly: {keys}")
242 async def expire(self):
243 """Expire all leases in the group."""
244 leases = tuple(lease for lease in self._lease_waiters.keys() if not lease.expired)
245 results = await asyncio.gather(
246 *(lease.expire() for lease in leases),
247 return_exceptions=True,
248 )
250 for task in self._lease_waiters.values():
251 task.cancel()
253 try:
254 await asyncio.gather(*self._lease_waiters.values(), return_exceptions=True)
255 finally:
256 errors = tuple(
257 result if isinstance(result, Exception) else RuntimeError("Lease failed to expire")
258 for lease, result in zip(leases, results, strict=True)
259 if not lease.expired
260 )
262 if errors:
263 raise ExceptionGroup("One or more leases failed to expire", errors)
265 async def _wait_for_event(self):
266 # Make a copy of the current set of tasks waiting for each lease expiry.
267 waiters = list(self._lease_waiters.values())
269 # Wait until the next scheduled refresh, or until a lease is acquired or lost.
270 max_wait = self._time_until_next_refresh()
272 async with scoped_waiter(self._lease_added.wait()) as wait_added:
273 expired, _ = await asyncio.wait(
274 [wait_added, *waiters], timeout=max_wait, return_when=asyncio.FIRST_COMPLETED
275 )
276 expired.discard(wait_added)
278 self._lease_added.clear()
279 return not expired
281 def _schedule_pop(self):
282 if self._refresh_queue and self._refresh_queue[0][0] <= datetime.now().timestamp():
283 _, lease = heapq.heappop(self._refresh_queue)
284 return lease
286 return None
288 def _time_until_next_refresh(self):
289 return (
290 max(0.0, self._refresh_queue[0][0] - datetime.now().timestamp())
291 if self._refresh_queue
292 else None
293 )
295 def _schedule_add(self, lease: Lease):
296 delay = lease.ttl * 0.3 + random.random() * lease.ttl * 0.25
297 dt = lease.model.refreshed_at + timedelta(seconds=delay)
298 heapq.heappush(self._refresh_queue, (dt.timestamp(), lease))