Coverage for core / src / sensorkit / backend / nats.py: 77%
225 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 logging
4import os
5from datetime import datetime
6from enum import StrEnum
7from typing import override
9import nats
10import nats.errors
11import nats.js.errors
12import nats.micro
13from nats.aio.msg import Msg
14from nats.aio.subscription import Subscription
15from nats.js.api import ConsumerConfig, DeliverPolicy, StreamConfig
16from nats.js.kv import KV_DEL, KV_PURGE, KeyValue
18from sensorkit.backend.base import (
19 BackendError,
20 BackendImpl,
21 Entity,
22 KeyNotFound,
23 KVEntry,
24 KVError,
25 RequestCallback,
26 RevisionError,
27 ServiceInfo,
28 SpecialProperty,
29 StreamMessage,
30 Subject,
31 UnregisteredResponder,
32)
33from sensorkit.common.aio import cleanup_future
35SUBJECT_PREFIX = "sensorkit"
36DELETE_OPS = (KV_DEL, KV_PURGE)
38logging.getLogger("nats").setLevel(
39 logging.WARNING if os.environ.get("SENSORKIT_DEBUG") else logging.CRITICAL
40)
43class Method(StrEnum):
44 """NATS subject namespace identifiers for the three backend communication channels."""
46 REQ = "request"
47 STR = "stream"
48 KV = "kv"
50 def with_prefix(self):
51 """Return the fully-qualified NATS subject prefix for this method (e.g. 'sensorkit.request')."""
52 return SUBJECT_PREFIX + "." + str(self)
55def _subject_to_nats(method: Method, target: Subject):
56 tokens = [method.with_prefix(), *target.path]
58 match target.prop:
59 case SpecialProperty.ALL_DESCENDANTS:
60 tokens.append(">")
61 case SpecialProperty.ALL_PROPERTIES:
62 tokens.append("*")
63 case SpecialProperty.EVENTS:
64 tokens.append("_EVENTS")
65 case SpecialProperty.NONE:
66 pass
67 case _:
68 tokens.append(target.prop)
70 return ".".join(tokens)
73def _subject_from_nats(method: Method, subject: str):
74 tokens = subject.split(".")
76 if tokens[0] != SUBJECT_PREFIX or tokens[1] != str(method):
77 raise ValueError()
79 if tokens[-1] == "_EVENTS":
80 tokens[-1] = SpecialProperty.EVENTS
82 return Subject(tuple[str](tokens[2:-1]), tokens[-1])
85def _kv_entry(entry: KeyValue.Entry, *, key: Subject | None = None):
86 return KVEntry(
87 key=key if key else _subject_from_nats(Method.KV, entry.key),
88 value=entry.value if entry.operation not in DELETE_OPS else KVEntry.DELETE_MARKER,
89 revision=entry.revision,
90 )
93class NATSBackendImpl(BackendImpl):
94 """BackendImpl backed by a live NATS JetStream cluster."""
96 @override
97 @classmethod
98 async def create(cls, servers: str | list[str] | None = None, **kwargs):
99 if servers is None:
100 servers = os.environ.get("NATS_URL", os.environ.get("SENSORKIT_BACKEND_ARG"))
102 if isinstance(servers, str):
103 servers = servers.split(",")
105 if servers is not None:
106 kwargs["servers"] = servers
108 # Connect to the NATS broker.
109 nc = await nats.connect(**kwargs)
111 # Create JetStream resources.
112 js = nc.jetstream()
113 info = await js.add_stream(
114 name="sensorkit",
115 subjects=[f"{Method.STR.with_prefix()}.>"],
116 )
117 # TODO: When limit_marker_ttl support is added to nats-py, watchers should be able to
118 # detect key expiration.
119 #kv = await js.create_key_value(bucket="sensorkit", limit_marker_ttl=1)
120 kv = await js.create_key_value(bucket="sensorkit")
122 return cls(
123 client=nc,
124 stream=info.config,
125 key_value=kv,
126 )
128 def __init__(self, client: nats.NATS, stream: StreamConfig, key_value: KeyValue):
129 self._nc = client
130 self._kv = key_value
131 self._stream = stream
132 self._js = client.jetstream()
133 self._subs: dict[Subject, Subscription] = {}
135 @override
136 async def register_service(self, info: ServiceInfo):
137 try:
138 await nats.micro.add_service(
139 self._nc,
140 name=info.name,
141 version=info.version,
142 )
143 except nats.errors.Error as e:
144 raise BackendError() from e
146 @override
147 async def request_invoke(self, target: Subject, payload: bytes):
148 try:
149 msg = await self._nc.request(
150 _subject_to_nats(Method.REQ, target),
151 payload,
152 )
153 except nats.errors.NoRespondersError as e:
154 raise UnregisteredResponder(target) from e
155 except nats.errors.Error as e:
156 raise BackendError() from e
158 return msg.data
160 @override
161 async def request_listen(self, target: Subject, coro: RequestCallback):
162 async def _handle_request(msg: Msg):
163 resp = await coro(msg.data)
165 if resp is None:
166 resp = b""
168 await msg.respond(resp)
170 async def _receive_request(msg: Msg):
171 t = asyncio.create_task(_handle_request(msg))
172 t.add_done_callback(cleanup_future)
174 try:
175 self._subs[target] = await self._nc.subscribe(
176 _subject_to_nats(Method.REQ, target),
177 cb=_receive_request,
178 )
179 except nats.errors.Error as e:
180 raise BackendError() from e
182 @override
183 async def request_purge_listeners(self):
184 req_prefix = Method.REQ.with_prefix()
185 subjects = tuple(
186 key for key, sub in self._subs.items() if sub.subject.startswith(req_prefix)
187 )
189 subscriptions = (self._subs.pop(subject) for subject in subjects)
190 unsubscribes = (subscription.unsubscribe() for subscription in subscriptions)
192 await asyncio.gather(*unsubscribes, return_exceptions=True)
194 @override
195 async def stream_list(self, entity: Entity):
196 try:
197 info = await self._js.stream_info(
198 self._stream.name,
199 subjects_filter=_subject_to_nats(
200 Method.STR,
201 entity.subject(SpecialProperty.ALL_PROPERTIES),
202 ),
203 )
204 return [entity.subject(key) for key in info.state.subjects.keys()]
205 except nats.errors.Error as e:
206 raise BackendError() from e
208 @override
209 async def stream_publish(self, target: Subject, payload: bytes):
210 try:
211 await self._js.publish(
212 subject=_subject_to_nats(Method.STR, target),
213 payload=payload,
214 stream=self._stream.name,
215 )
216 except nats.errors.Error as e:
217 raise BackendError() from e
219 @staticmethod
220 async def _stream_iterator(sub: Subscription):
221 try:
222 async for msg in sub.messages:
223 try:
224 yield StreamMessage(
225 subject=_subject_from_nats(Method.STR, msg.subject),
226 sequence=msg.metadata.sequence.stream,
227 timestamp=msg.metadata.timestamp,
228 data=msg.data,
229 )
230 finally:
231 await msg.ack()
232 except nats.errors.Error as e:
233 raise BackendError() from e
234 finally:
235 await sub.unsubscribe()
237 @override
238 async def stream_consume(
239 self,
240 target: Subject,
241 *,
242 durable_name: str | None = None,
243 start_at: int | datetime | None = None,
244 include_latest: bool = False,
245 ):
246 config = ConsumerConfig()
247 config.stream = self._stream.name
248 config.durable_name = durable_name
250 if start_at is not None:
251 if include_latest:
252 raise RuntimeError("Cannot specify both `start_at` and `include_latest`")
254 match start_at:
255 case datetime():
256 config.deliver_policy = DeliverPolicy.BY_START_TIME
257 config.opt_start_time = start_at.isoformat() # type: ignore
258 case int():
259 if start_at <= 0:
260 config.deliver_policy = DeliverPolicy.ALL
261 else:
262 config.deliver_policy = DeliverPolicy.BY_START_SEQUENCE
263 config.opt_start_seq = start_at
264 elif include_latest:
265 config.deliver_policy = DeliverPolicy.LAST_PER_SUBJECT
266 else:
267 config.deliver_policy = DeliverPolicy.NEW
269 try:
270 sub = await self._js.subscribe(
271 _subject_to_nats(Method.STR, target),
272 config=config,
273 )
274 except nats.errors.Error as e:
275 raise BackendError() from e
277 return self._stream_iterator(sub)
279 @staticmethod
280 async def _kv_iterator(watcher: KeyValue.KeyWatcher):
281 try:
282 # Collect initial entries until the watcher yields None (end-of-initial sentinel).
283 initial = []
284 async for entry in watcher:
285 if entry is None:
286 break
287 initial.append(_kv_entry(entry))
288 yield initial
290 # Stream live updates.
291 async for entry in watcher:
292 if entry is not None:
293 yield [_kv_entry(entry)]
294 finally:
295 await watcher.stop()
297 @override
298 async def kv_monitor(self, target: Subject):
299 try:
300 watcher = await self._kv.watch(
301 keys=_subject_to_nats(Method.KV, target),
302 include_history=True,
303 )
304 except nats.js.errors.KeyValueError as e:
305 raise KVError() from e
306 except nats.errors.Error as e:
307 raise BackendError() from e
309 return self._kv_iterator(watcher)
311 @override
312 async def kv_get(self, target: Subject, revision: int | None = None) -> KVEntry:
313 try:
314 entry = await self._kv.get(key=_subject_to_nats(Method.KV, target), revision=revision)
315 except nats.js.errors.KeyNotFoundError as e:
316 raise KeyNotFound(target, deleted=e.op in DELETE_OPS) from e
317 except nats.js.errors.KeyValueError as e:
318 raise KVError() from e
319 except nats.errors.Error as e:
320 raise BackendError() from e
322 return _kv_entry(entry, key=target)
324 @override
325 async def kv_create(self, target: Subject, payload: bytes, ttl: float | None = None):
326 # Presently NATS (or nats-py, as of v2.13) supports only whole second TTLs.
327 if ttl:
328 ttl = max(round(ttl), 1)
330 try:
331 rev = await self._kv.create(
332 key=_subject_to_nats(Method.KV, target),
333 value=payload,
334 msg_ttl=ttl,
335 )
336 except nats.js.errors.KeyWrongLastSequenceError as e:
337 raise RevisionError() from e
338 except nats.js.errors.KeyValueError as e:
339 raise KVError() from e
340 except nats.errors.Error as e:
341 raise BackendError() from e
343 return KVEntry(target, payload, rev)
345 @override
346 async def kv_update(
347 self,
348 target: Subject,
349 payload: bytes,
350 revision: int | None = None,
351 ttl: float | None = None,
352 ):
353 key = _subject_to_nats(Method.KV, target)
355 # Presently NATS (or nats-py, as of v2.13) supports only whole second TTLs.
356 if ttl:
357 ttl = max(round(ttl), 1)
359 try:
360 if revision is None:
361 if ttl is not None:
362 raise BackendError("Revision required for update with TTL")
364 rev = await self._kv.put(key, payload)
365 else:
366 rev = await self._kv.update(key, payload, revision, msg_ttl=ttl)
367 except nats.js.errors.KeyWrongLastSequenceError as e:
368 raise RevisionError() from e
369 except nats.js.errors.KeyValueError as e:
370 raise KVError() from e
371 except nats.errors.Error as e:
372 raise BackendError() from e
374 return KVEntry(target, payload, rev)
376 @override
377 async def kv_delete(self, target: Subject, revision: int | None = None):
378 try:
379 await self._kv.delete(
380 key=_subject_to_nats(Method.KV, target),
381 last=revision,
382 )
383 except nats.js.errors.KeyWrongLastSequenceError as e:
384 raise RevisionError() from e
385 except nats.errors.Error as e:
386 raise BackendError() from e