Coverage for core / src / sensorkit / backend / fake.py: 97%

231 statements  

« prev     ^ index     » next       coverage.py v7.13.5, created at 2026-09-02 00:03 +0000

1# SPDX-License-Identifier: Apache-2.0 

2import asyncio 

3import bisect 

4from collections.abc import AsyncIterator, Iterable, MutableMapping 

5from dataclasses import dataclass, field 

6from datetime import UTC, datetime 

7from typing import override 

8 

9import pygtrie 

10from loguru import logger 

11 

12from sensorkit.backend.base import ( 

13 BackendImpl, 

14 Entity, 

15 KeyNotFound, 

16 KVEntry, 

17 RequestCallback, 

18 RevisionError, 

19 SpecialProperty, 

20 StreamMessage, 

21 Subject, 

22 UnregisteredResponder, 

23) 

24 

25type Trie[K, V] = pygtrie.Trie | MutableMapping[K, V] 

26 

27 

28@dataclass 

29class FakeBrokerElement[K, V]: 

30 """Storage node for a single key in the FakeBroker trie, holding value history and subscriber queues.""" 

31 

32 key: K 

33 history: list[V] = field(default_factory=list) 

34 direct: set[asyncio.Queue[tuple[K, V]]] = field(default_factory=set) 

35 shallow_wildcard: set[asyncio.Queue[tuple[K, V]]] = field(default_factory=set) 

36 deep_wildcard: set[asyncio.Queue[tuple[K, V]]] = field(default_factory=set) 

37 

38 @classmethod 

39 async def _observe(cls, queue: asyncio.Queue[tuple[K, V]], queue_set: set[asyncio.Queue]): 

40 try: 

41 while True: 

42 yield await queue.get() 

43 finally: 

44 queue_set.discard(queue) 

45 

46 def subscribe(self): 

47 """Return an async generator that yields (key, value) tuples published directly to this key.""" 

48 queue = asyncio.Queue() 

49 self.direct.add(queue) 

50 return self._observe(queue, self.direct) 

51 

52 def subscribe_children(self): 

53 """Return an async generator that yields (key, value) tuples for immediate children.""" 

54 queue = asyncio.Queue() 

55 self.shallow_wildcard.add(queue) 

56 return self._observe(queue, self.shallow_wildcard) 

57 

58 def subscribe_descendants(self): 

59 """Return an async generator that yields (key, value) tuples for all descendants.""" 

60 queue = asyncio.Queue() 

61 self.deep_wildcard.add(queue) 

62 return self._observe(queue, self.deep_wildcard) 

63 

64 def direct_set(self, value: V): 

65 """Append value to history and notify direct subscribers.""" 

66 self.history.append(value) 

67 data = (self.key, value) 

68 

69 for queue in self.direct: 

70 queue.put_nowait(data) 

71 

72 def descendant_set(self, value: V, immediate=False): 

73 """Notify wildcard subscribers of a descendant value update.""" 

74 data = (self.key, value) 

75 

76 if immediate: 

77 for queue in self.shallow_wildcard: 

78 queue.put_nowait(data) 

79 

80 for queue in self.deep_wildcard: 

81 queue.put_nowait(data) 

82 

83 

84class FakeBroker[K, V]: 

85 """In-memory pub/sub broker with hierarchical subscription support. 

86 

87 Manages a trie-based storage system for publish-subscribe messaging with support 

88 for direct subscriptions, shallow wildcard subscriptions (immediate children), 

89 and deep wildcard subscriptions (all descendants). 

90 

91 Type Parameters: 

92 K: Key type 

93 V: Value type for messages published to the broker 

94 """ 

95 

96 def __init__(self): 

97 self._trie: Trie[K, FakeBrokerElement[K, V]] = pygtrie.Trie() 

98 

99 def element(self, key: K) -> FakeBrokerElement[K, V]: 

100 """Return the broker element for the given key, creating it if it does not yet exist.""" 

101 if key not in self._trie: 

102 self._trie[key] = FakeBrokerElement(key=key) 

103 

104 return self._trie[key] 

105 

106 def children_of(self, key: K): 

107 """Return the immediate child keys of the given key in the trie.""" 

108 def node_factory(path_conv, path, children, *value): 

109 current_key = path_conv(path) 

110 if current_key == key: 

111 return [c for c in children if c is not None] 

112 return current_key if value else None 

113 

114 try: 

115 return self._trie.traverse(node_factory, prefix=key) 

116 except KeyError: 

117 return [] 

118 

119 def descendants_of(self, key: K): 

120 """Return all descendant keys (at any depth) of the given key in the trie.""" 

121 try: 

122 return (k for k in self._trie.keys(prefix=key, shallow=False) if k != key) 

123 except KeyError: 

124 return [] 

125 

126 def publish(self, key: K, value: V): 

127 """Publish a value to a key, notifying all direct and ancestor wildcard subscribers.""" 

128 # Get the element corresponding to the given key and set its new value. 

129 elem = self.element(key) 

130 elem.direct_set(value) 

131 

132 # Use our trie to find all prefixes of the key. 

133 ancestors: list[FakeBrokerElement[K, V]] = list( 

134 step.value for step in self._trie.prefixes(key) 

135 ) 

136 

137 # The last prefix should always correspond to our full key. 

138 _elem = ancestors.pop() 

139 assert elem is _elem 

140 

141 # If there is a parent, propagate to the parent to handle shallow wildcards. 

142 parent = ancestors.pop() if ancestors else None 

143 

144 if parent: 

145 parent.descendant_set(value, immediate=True) 

146 

147 # If there are more ancestors, propagate similarly to handle deep wildcards. 

148 for ancestor in reversed(ancestors): 

149 ancestor.descendant_set(value) 

150 

151 def defined(self, key: K): 

152 """Return True if the key exists in the trie and has at least one history entry.""" 

153 return key in self._trie and len(self._trie[key].history) > 0 

154 

155 

156class FakeBackendImpl(BackendImpl): 

157 """In-memory BackendImpl used for testing, backed by FakeBroker instances.""" 

158 

159 @override 

160 @classmethod 

161 async def create(cls, *args, **kwargs) -> BackendImpl: 

162 return cls() 

163 

164 def __init__(self): 

165 self._req: dict[Subject, RequestCallback] = {} 

166 self._kv: FakeBroker[tuple[str, ...], KVEntry] = FakeBroker() 

167 self._kv_ttls: dict[Subject, float] = {} 

168 self._kv_expirers: dict[Subject, asyncio.Task] = {} 

169 self._streams: FakeBroker[tuple[str, ...], StreamMessage] = FakeBroker() 

170 self._sequence = 0 

171 

172 @override 

173 async def request_invoke(self, target: Subject, payload: bytes) -> bytes: 

174 if target not in self._req: 

175 raise UnregisteredResponder(target) 

176 

177 return await asyncio.create_task(self._req[target](payload)) 

178 

179 @override 

180 async def request_listen(self, target: Subject, func: RequestCallback): 

181 self._req[target] = func 

182 

183 @override 

184 async def request_purge_listeners(self): 

185 self._req.clear() 

186 

187 @override 

188 async def stream_list(self, entity: Entity) -> list[Subject]: 

189 return [ 

190 Subject(path=p[:-1], prop=p[-1]) for p in self._streams.descendants_of(entity.path) 

191 ] 

192 

193 @override 

194 async def stream_publish(self, target: Subject, payload: bytes): 

195 stream = self._streams.element(target.full_path()) 

196 was_empty = len(stream.history) == 0 

197 now = datetime.now(UTC) 

198 self._sequence += 1 

199 

200 assert was_empty or stream.history[-1].timestamp < now 

201 

202 msg = StreamMessage( 

203 subject=target, 

204 sequence=self._sequence, 

205 timestamp=now, 

206 data=payload, 

207 ) 

208 

209 self._streams.publish(target.full_path(), msg) 

210 

211 def _stream_filtered_history( 

212 self, 

213 keys: Iterable[tuple[str, ...]], 

214 include_latest: bool, 

215 start_at: int | datetime | None, 

216 ): 

217 histories = [self._streams.element(key).history for key in keys] 

218 

219 for history in histories: 

220 if not history: 

221 continue 

222 

223 last = len(history) - 1 

224 

225 if include_latest: 

226 assert start_at is None 

227 

228 yield history[last] 

229 else: 

230 match start_at: 

231 case int(): 

232 start_idx = bisect.bisect_left(history, start_at, key=lambda x: x.sequence) 

233 case datetime(): 

234 start_idx = bisect.bisect_left( 

235 history, start_at, key=lambda x: x.timestamp 

236 ) 

237 case None: 

238 start_idx = len(history) - 1 

239 

240 for i in range(start_idx, last + 1): 

241 yield history[i] 

242 

243 @override 

244 async def stream_consume( 

245 self, 

246 target: Subject, 

247 *, 

248 durable_name: str | None = None, 

249 start_at: int | datetime | None = None, 

250 include_latest: bool = False, 

251 ) -> AsyncIterator[StreamMessage]: 

252 match target.prop: 

253 case SpecialProperty.ALL_PROPERTIES: 

254 elem = self._streams.element(target.path) 

255 subscription = elem.subscribe_children() 

256 history = self._stream_filtered_history( 

257 self._streams.children_of(target.path), 

258 include_latest, 

259 start_at, 

260 ) 

261 case SpecialProperty.ALL_DESCENDANTS: 

262 elem = self._streams.element(target.path) 

263 subscription = elem.subscribe_descendants() 

264 history = self._stream_filtered_history( 

265 self._streams.descendants_of(target.path), 

266 include_latest, 

267 start_at, 

268 ) 

269 case _: 

270 elem = self._streams.element(target.full_path()) 

271 subscription = elem.subscribe() 

272 history = self._stream_filtered_history( 

273 (target.full_path(),), include_latest, start_at 

274 ) 

275 

276 async def _stream_consume(): 

277 for msg in history: 

278 yield msg 

279 

280 async for _, msg in subscription: 

281 yield msg 

282 

283 return _stream_consume() 

284 

285 @override 

286 async def kv_monitor(self, target: Subject) -> AsyncIterator[list[KVEntry]]: 

287 match target.prop: 

288 case SpecialProperty.ALL_PROPERTIES: 

289 elem = self._kv.element(target.path) 

290 subscription = elem.subscribe_children() 

291 history = [ 

292 self._kv.element(key).history[-1] for key in self._kv.children_of(target.path) 

293 ] 

294 case SpecialProperty.ALL_DESCENDANTS: 

295 elem = self._kv.element(target.path) 

296 subscription = elem.subscribe_descendants() 

297 history = [ 

298 self._kv.element(key).history[-1] 

299 for key in self._kv.descendants_of(target.path) 

300 ] 

301 case _: 

302 elem = self._kv.element(target.full_path()) 

303 subscription = elem.subscribe() 

304 history = [elem.history[-1]] if elem.history else [] 

305 

306 async def _kv_monitor(): 

307 yield history 

308 

309 async for _, entry in subscription: 

310 yield [entry] 

311 

312 return _kv_monitor() 

313 

314 def _kv_defined_not_deleted(self, key: tuple[str, ...]): 

315 return self._kv.defined(key) and not self._kv.element(key).history[-1].deleted() 

316 

317 @override 

318 async def kv_get(self, target: Subject) -> KVEntry: 

319 key = target.full_path() 

320 

321 if not self._kv_defined_not_deleted(key): 

322 raise KeyNotFound(subject=target, deleted=self._kv.defined(key)) 

323 

324 return self._kv.element(key).history[-1] 

325 

326 NOT_SET = object() 

327 

328 def _kv_do_update( 

329 self, 

330 target: Subject, 

331 payload: bytes, 

332 revision: int | None, 

333 ttl: float | None = NOT_SET, 

334 ): 

335 key = target.full_path() 

336 

337 if revision is None: 

338 revision = 0 

339 elif self._kv.defined(key): 

340 entry = self._kv.element(key).history[-1] 

341 

342 if revision != entry.revision: 

343 raise RevisionError("revision does not match") 

344 

345 revision += 1 

346 

347 # Determine whether there is an applicable TTL, and if so, enforce it by creating tasks 

348 # to send a `kv_delete` at expiry. 

349 

350 if ttl is self.NOT_SET: 

351 ttl = self._kv_ttls.get(key, None) 

352 else: 

353 self._kv_ttls[key] = ttl 

354 

355 if old_expirer := self._kv_expirers.get(target): 

356 old_expirer.cancel() 

357 

358 if ttl: 

359 async def _expire_ttl(): 

360 await asyncio.sleep(ttl) 

361 await self.kv_delete(target, revision) 

362 

363 def _expire_cleanup(t): 

364 del self._kv_expirers[target] 

365 

366 if not t.cancelled(): 

367 t.exception() 

368 

369 if payload is not KVEntry.DELETE_MARKER: 

370 task = asyncio.create_task(_expire_ttl()) 

371 self._kv_expirers[target] = task 

372 task.add_done_callback(_expire_cleanup) 

373 

374 entry = KVEntry(key=target, value=payload, revision=revision) 

375 self._kv.publish(key, entry) 

376 logger.debug(f"KV: {entry}") 

377 return entry 

378 

379 @override 

380 async def kv_create( 

381 self, 

382 target: Subject, 

383 payload: bytes, 

384 ttl: float | None = None, 

385 ) -> KVEntry: 

386 try: 

387 return self._kv_do_update(target, payload, 0, ttl) 

388 except RevisionError as err: 

389 entry = self._kv.element(target.full_path()).history[-1] 

390 

391 if entry.deleted(): 

392 return self._kv_do_update(target, payload, entry.revision, ttl) 

393 

394 raise err 

395 

396 @override 

397 async def kv_update( 

398 self, 

399 target: Subject, 

400 payload: bytes, 

401 revision: int | None = None, 

402 ttl: float | None = None, 

403 ) -> KVEntry: 

404 return self._kv_do_update(target, payload, revision) 

405 

406 @override 

407 async def kv_delete(self, target: Subject, revision: int | None = None): 

408 if self._kv_defined_not_deleted(target.full_path()): 

409 self._kv_do_update(target, KVEntry.DELETE_MARKER, revision)