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

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 

8 

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 

17 

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 

34 

35SUBJECT_PREFIX = "sensorkit" 

36DELETE_OPS = (KV_DEL, KV_PURGE) 

37 

38logging.getLogger("nats").setLevel( 

39 logging.WARNING if os.environ.get("SENSORKIT_DEBUG") else logging.CRITICAL 

40) 

41 

42 

43class Method(StrEnum): 

44 """NATS subject namespace identifiers for the three backend communication channels.""" 

45 

46 REQ = "request" 

47 STR = "stream" 

48 KV = "kv" 

49 

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) 

53 

54 

55def _subject_to_nats(method: Method, target: Subject): 

56 tokens = [method.with_prefix(), *target.path] 

57 

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) 

69 

70 return ".".join(tokens) 

71 

72 

73def _subject_from_nats(method: Method, subject: str): 

74 tokens = subject.split(".") 

75 

76 if tokens[0] != SUBJECT_PREFIX or tokens[1] != str(method): 

77 raise ValueError() 

78 

79 if tokens[-1] == "_EVENTS": 

80 tokens[-1] = SpecialProperty.EVENTS 

81 

82 return Subject(tuple[str](tokens[2:-1]), tokens[-1]) 

83 

84 

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 ) 

91 

92 

93class NATSBackendImpl(BackendImpl): 

94 """BackendImpl backed by a live NATS JetStream cluster.""" 

95 

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

101 

102 if isinstance(servers, str): 

103 servers = servers.split(",") 

104 

105 if servers is not None: 

106 kwargs["servers"] = servers 

107 

108 # Connect to the NATS broker. 

109 nc = await nats.connect(**kwargs) 

110 

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

121 

122 return cls( 

123 client=nc, 

124 stream=info.config, 

125 key_value=kv, 

126 ) 

127 

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] = {} 

134 

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 

145 

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 

157 

158 return msg.data 

159 

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) 

164 

165 if resp is None: 

166 resp = b"" 

167 

168 await msg.respond(resp) 

169 

170 async def _receive_request(msg: Msg): 

171 t = asyncio.create_task(_handle_request(msg)) 

172 t.add_done_callback(cleanup_future) 

173 

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 

181 

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 ) 

188 

189 subscriptions = (self._subs.pop(subject) for subject in subjects) 

190 unsubscribes = (subscription.unsubscribe() for subscription in subscriptions) 

191 

192 await asyncio.gather(*unsubscribes, return_exceptions=True) 

193 

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 

207 

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 

218 

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

236 

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 

249 

250 if start_at is not None: 

251 if include_latest: 

252 raise RuntimeError("Cannot specify both `start_at` and `include_latest`") 

253 

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 

268 

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 

276 

277 return self._stream_iterator(sub) 

278 

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 

289 

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

296 

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 

308 

309 return self._kv_iterator(watcher) 

310 

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 

321 

322 return _kv_entry(entry, key=target) 

323 

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) 

329 

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 

342 

343 return KVEntry(target, payload, rev) 

344 

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) 

354 

355 # Presently NATS (or nats-py, as of v2.13) supports only whole second TTLs. 

356 if ttl: 

357 ttl = max(round(ttl), 1) 

358 

359 try: 

360 if revision is None: 

361 if ttl is not None: 

362 raise BackendError("Revision required for update with TTL") 

363 

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 

373 

374 return KVEntry(target, payload, rev) 

375 

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