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
« 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
9import pygtrie
10from loguru import logger
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)
25type Trie[K, V] = pygtrie.Trie | MutableMapping[K, V]
28@dataclass
29class FakeBrokerElement[K, V]:
30 """Storage node for a single key in the FakeBroker trie, holding value history and subscriber queues."""
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)
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)
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)
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)
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)
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)
69 for queue in self.direct:
70 queue.put_nowait(data)
72 def descendant_set(self, value: V, immediate=False):
73 """Notify wildcard subscribers of a descendant value update."""
74 data = (self.key, value)
76 if immediate:
77 for queue in self.shallow_wildcard:
78 queue.put_nowait(data)
80 for queue in self.deep_wildcard:
81 queue.put_nowait(data)
84class FakeBroker[K, V]:
85 """In-memory pub/sub broker with hierarchical subscription support.
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).
91 Type Parameters:
92 K: Key type
93 V: Value type for messages published to the broker
94 """
96 def __init__(self):
97 self._trie: Trie[K, FakeBrokerElement[K, V]] = pygtrie.Trie()
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)
104 return self._trie[key]
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
114 try:
115 return self._trie.traverse(node_factory, prefix=key)
116 except KeyError:
117 return []
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 []
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)
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 )
137 # The last prefix should always correspond to our full key.
138 _elem = ancestors.pop()
139 assert elem is _elem
141 # If there is a parent, propagate to the parent to handle shallow wildcards.
142 parent = ancestors.pop() if ancestors else None
144 if parent:
145 parent.descendant_set(value, immediate=True)
147 # If there are more ancestors, propagate similarly to handle deep wildcards.
148 for ancestor in reversed(ancestors):
149 ancestor.descendant_set(value)
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
156class FakeBackendImpl(BackendImpl):
157 """In-memory BackendImpl used for testing, backed by FakeBroker instances."""
159 @override
160 @classmethod
161 async def create(cls, *args, **kwargs) -> BackendImpl:
162 return cls()
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
172 @override
173 async def request_invoke(self, target: Subject, payload: bytes) -> bytes:
174 if target not in self._req:
175 raise UnregisteredResponder(target)
177 return await asyncio.create_task(self._req[target](payload))
179 @override
180 async def request_listen(self, target: Subject, func: RequestCallback):
181 self._req[target] = func
183 @override
184 async def request_purge_listeners(self):
185 self._req.clear()
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 ]
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
200 assert was_empty or stream.history[-1].timestamp < now
202 msg = StreamMessage(
203 subject=target,
204 sequence=self._sequence,
205 timestamp=now,
206 data=payload,
207 )
209 self._streams.publish(target.full_path(), msg)
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]
219 for history in histories:
220 if not history:
221 continue
223 last = len(history) - 1
225 if include_latest:
226 assert start_at is None
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
240 for i in range(start_idx, last + 1):
241 yield history[i]
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 )
276 async def _stream_consume():
277 for msg in history:
278 yield msg
280 async for _, msg in subscription:
281 yield msg
283 return _stream_consume()
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 []
306 async def _kv_monitor():
307 yield history
309 async for _, entry in subscription:
310 yield [entry]
312 return _kv_monitor()
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()
317 @override
318 async def kv_get(self, target: Subject) -> KVEntry:
319 key = target.full_path()
321 if not self._kv_defined_not_deleted(key):
322 raise KeyNotFound(subject=target, deleted=self._kv.defined(key))
324 return self._kv.element(key).history[-1]
326 NOT_SET = object()
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()
337 if revision is None:
338 revision = 0
339 elif self._kv.defined(key):
340 entry = self._kv.element(key).history[-1]
342 if revision != entry.revision:
343 raise RevisionError("revision does not match")
345 revision += 1
347 # Determine whether there is an applicable TTL, and if so, enforce it by creating tasks
348 # to send a `kv_delete` at expiry.
350 if ttl is self.NOT_SET:
351 ttl = self._kv_ttls.get(key, None)
352 else:
353 self._kv_ttls[key] = ttl
355 if old_expirer := self._kv_expirers.get(target):
356 old_expirer.cancel()
358 if ttl:
359 async def _expire_ttl():
360 await asyncio.sleep(ttl)
361 await self.kv_delete(target, revision)
363 def _expire_cleanup(t):
364 del self._kv_expirers[target]
366 if not t.cancelled():
367 t.exception()
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)
374 entry = KVEntry(key=target, value=payload, revision=revision)
375 self._kv.publish(key, entry)
376 logger.debug(f"KV: {entry}")
377 return entry
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]
391 if entry.deleted():
392 return self._kv_do_update(target, payload, entry.revision, ttl)
394 raise err
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)
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)