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

1# SPDX-License-Identifier: Apache-2.0 

2from __future__ import annotations 

3 

4import asyncio 

5import contextlib 

6import heapq 

7import os 

8import random 

9from datetime import datetime, timedelta 

10from typing import Callable 

11 

12from loguru import logger 

13from pydantic import BaseModel 

14 

15from sensorkit.backend.base import BackendError, KeyValueContext, RevisionError 

16from sensorkit.common.aio import scoped_waiter 

17 

18 

19class LeaseModel[M: BaseModel | None](BaseModel): 

20 """Serialisable metadata stored in the KV entry for an active lease.""" 

21 

22 acquired_at: datetime 

23 refreshed_at: datetime 

24 record: M 

25 

26 

27class LeaseUnavailableError(BackendError): 

28 """Raised when a lease cannot be acquired because the key is already held.""" 

29 

30 

31class LeaseAbandonedError(BackendError): 

32 """Raised when a lease is lost unexpectedly (TTL expiry or external deletion).""" 

33 

34 

35class Lease[M: BaseModel | None]: 

36 """A distributed, TTL-backed lock stored in the KV backend. 

37 

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 """ 

42 

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. 

52 

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 ) 

61 

62 if os.environ.get("SENSORKIT_NO_ENFORCE_LEASES"): 

63 await kv.delete(key) 

64 

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 

75 

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()) 

93 

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) 

98 

99 async def _expiry_monitor(self): 

100 if os.environ.get("SENSORKIT_NO_ENFORCE_LEASES"): 

101 return 

102 

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) 

108 

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) 

114 

115 if entry.deleted(): 

116 # We have lost the lease. Fallthrough to ensure we are marked as expired. 

117 break 

118 

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() 

125 

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 

130 

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 

139 

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)}") 

146 

147 async with asyncio.timeout(1.0): 

148 await self._kv.delete( 

149 self._key, 

150 revision=self._revision, 

151 ) 

152 

153 self._expired.set() 

154 self._task.cancel() 

155 

156 with contextlib.suppress(asyncio.CancelledError): 

157 await self._task 

158 

159 @property 

160 def expired(self): 

161 """Return True if the lease is expired.""" 

162 return self._expired.is_set() 

163 

164 def wait_expired(self): 

165 """Wait until lease expiry.""" 

166 return self._expired.wait() 

167 

168 async def __aenter__(self): 

169 return self 

170 

171 async def __aexit__(self, exc_type, exc_val, exc_tb): 

172 await self.expire() 

173 

174 

175class LeaseGroup: 

176 """Manage a group of Leases with shared expiration semantics.""" 

177 

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() 

182 

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 

201 

202 async def refresh_loop(self, record_callback: Callable[[Lease], None] | None = None): 

203 """Yield Leases as their scheduled refresh time becomes current. 

204 

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() 

216 

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 

220 

221 # Refresh and reschedule the lease. 

222 await lease.refresh(record) 

223 self._schedule_add(lease) 

224 

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() 

237 

238 if abandoned: 

239 keys = ", ".join(lease._key for lease in abandoned) 

240 raise LeaseAbandonedError(f"Lease(s) lost unexpectedly: {keys}") 

241 

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 ) 

249 

250 for task in self._lease_waiters.values(): 

251 task.cancel() 

252 

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 ) 

261 

262 if errors: 

263 raise ExceptionGroup("One or more leases failed to expire", errors) 

264 

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()) 

268 

269 # Wait until the next scheduled refresh, or until a lease is acquired or lost. 

270 max_wait = self._time_until_next_refresh() 

271 

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) 

277 

278 self._lease_added.clear() 

279 return not expired 

280 

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 

285 

286 return None 

287 

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 ) 

294 

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))