Skip to main content

dynamo_runtime/storage/kv/
nats.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use std::{collections::HashMap, pin::Pin, time::Duration};
5
6use crate::{protocols::EndpointId, slug::Slug, storage::kv, transports::nats::Client};
7use async_nats::jetstream::kv::Operation;
8use async_trait::async_trait;
9use futures::StreamExt;
10
11use super::{Bucket, Store, StoreError, StoreOutcome};
12
13#[derive(Clone)]
14pub struct NATSStore {
15    client: Client,
16    endpoint: EndpointId,
17}
18
19pub struct NATSBucket {
20    nats_store: async_nats::jetstream::kv::Store,
21}
22
23#[async_trait]
24impl Store for NATSStore {
25    type Bucket = NATSBucket;
26
27    async fn get_or_create_bucket(
28        &self,
29        bucket_name: &str,
30        ttl: Option<Duration>,
31    ) -> Result<Self::Bucket, StoreError> {
32        let name = Slug::slugify(bucket_name);
33        let nats_store = self
34            .get_or_create_key_value(&self.endpoint.namespace, &name, ttl)
35            .await?;
36        Ok(NATSBucket { nats_store })
37    }
38
39    async fn get_bucket(&self, bucket_name: &str) -> Result<Option<Self::Bucket>, StoreError> {
40        let name = Slug::slugify(bucket_name);
41        match self.get_key_value(&self.endpoint.namespace, &name).await? {
42            Some(nats_store) => Ok(Some(NATSBucket { nats_store })),
43            None => Ok(None),
44        }
45    }
46
47    fn connection_id(&self) -> u64 {
48        self.client.client().server_info().client_id
49    }
50
51    fn shutdown(&self) {
52        // TODO: Track and delete any owned keys
53        // The TTL should ensure NATS does it, but best we do it immediately
54    }
55}
56
57impl NATSStore {
58    pub fn new(client: Client, endpoint: EndpointId) -> Self {
59        NATSStore { client, endpoint }
60    }
61
62    /// Get or create a key-value store (aka bucket) in NATS.
63    ///
64    /// ttl is only used if we are creating the bucket, so if that has
65    /// changed first delete the bucket.
66    async fn get_or_create_key_value(
67        &self,
68        namespace: &str,
69        bucket_name: &Slug,
70        // Delete entries older than this
71        ttl: Option<Duration>,
72    ) -> Result<async_nats::jetstream::kv::Store, StoreError> {
73        if let Ok(Some(kv)) = self.get_key_value(namespace, bucket_name).await {
74            return Ok(kv);
75        }
76
77        // It doesn't exist, create it
78
79        let bucket_name = single_name(namespace, bucket_name);
80        let js = self.client.jetstream();
81        let create_result = js
82            .create_key_value(
83                // TODO: configure the bucket, probably need to pass some of these values in
84                async_nats::jetstream::kv::Config {
85                    bucket: bucket_name.clone(),
86                    max_age: ttl.unwrap_or_default(),
87                    ..Default::default()
88                },
89            )
90            .await;
91        let nats_store = create_result
92            .map_err(|err| StoreError::KeyValueError(err.to_string(), bucket_name.clone()))?;
93        tracing::debug!("Created bucket {bucket_name}");
94        Ok(nats_store)
95    }
96
97    async fn get_key_value(
98        &self,
99        namespace: &str,
100        bucket_name: &Slug,
101    ) -> Result<Option<async_nats::jetstream::kv::Store>, StoreError> {
102        let bucket_name = single_name(namespace, bucket_name);
103        let js = self.client.jetstream();
104
105        use async_nats::jetstream::context::KeyValueErrorKind;
106        match js.get_key_value(&bucket_name).await {
107            Ok(store) => Ok(Some(store)),
108            Err(err) if err.kind() == KeyValueErrorKind::GetBucket => {
109                // bucket doesn't exist
110                Ok(None)
111            }
112            Err(err) => Err(StoreError::KeyValueError(err.to_string(), bucket_name)),
113        }
114    }
115}
116
117#[async_trait]
118impl Bucket for NATSBucket {
119    async fn insert(
120        &self,
121        key: &kv::Key,
122        value: bytes::Bytes,
123        revision: u64,
124    ) -> Result<StoreOutcome, StoreError> {
125        if revision == 0 {
126            self.create(key, value).await
127        } else {
128            self.update(key, value, revision).await
129        }
130    }
131
132    async fn compare_and_replace(
133        &self,
134        key: &kv::Key,
135        expected: bytes::Bytes,
136        value: bytes::Bytes,
137    ) -> Result<StoreOutcome, StoreError> {
138        let entry = self
139            .nats_store
140            .entry(key)
141            .await
142            .map_err(|error| StoreError::NATSError(error.to_string()))?
143            .ok_or_else(|| StoreError::MissingKey(key.to_string()))?;
144        if matches!(entry.operation, Operation::Delete | Operation::Purge) {
145            return Err(StoreError::MissingKey(key.to_string()));
146        }
147        if entry.value != expected {
148            return Err(StoreError::Retry);
149        }
150
151        match self.nats_store.update(key, value, entry.revision).await {
152            Ok(revision) => Ok(StoreOutcome::Created(revision)),
153            Err(error)
154                if error.kind()
155                    == async_nats::jetstream::kv::UpdateErrorKind::WrongLastRevision =>
156            {
157                match self.nats_store.entry(key).await {
158                    Ok(None) => Err(StoreError::MissingKey(key.to_string())),
159                    Ok(Some(entry))
160                        if matches!(entry.operation, Operation::Delete | Operation::Purge) =>
161                    {
162                        Err(StoreError::MissingKey(key.to_string()))
163                    }
164                    Ok(Some(_)) => Err(StoreError::Retry),
165                    Err(error) => Err(StoreError::NATSError(error.to_string())),
166                }
167            }
168            Err(error) => Err(StoreError::NATSError(error.to_string())),
169        }
170    }
171
172    async fn get(&self, key: &kv::Key) -> Result<Option<bytes::Bytes>, StoreError> {
173        self.nats_store
174            .get(key)
175            .await
176            .map_err(|e| StoreError::NATSError(e.to_string()))
177    }
178
179    async fn delete(&self, key: &kv::Key) -> Result<(), StoreError> {
180        self.nats_store
181            .delete(key)
182            .await
183            .map_err(|e| StoreError::NATSError(e.to_string()))
184    }
185
186    async fn watch(
187        &self,
188    ) -> Result<Pin<Box<dyn futures::Stream<Item = kv::WatchEvent> + Send + 'life0>>, StoreError>
189    {
190        let watch_stream = self
191            .nats_store
192            .watch_with_history(">")
193            .await
194            .map_err(|e| StoreError::NATSError(e.to_string()))?;
195        // Map the `Entry` to `Entry.value` which is Bytes of the stored value.
196        Ok(Box::pin(
197            watch_stream.filter_map(
198                |maybe_entry: Result<
199                    async_nats::jetstream::kv::Entry,
200                    async_nats::error::Error<_>,
201                >| async move {
202                    match maybe_entry {
203                        Ok(entry) => {
204                            let key = kv::Key::new(entry.key);
205                            Some(match entry.operation {
206                                Operation::Put => {
207                                    let item = kv::KeyValue::new(key, entry.value);
208                                    kv::WatchEvent::Put(item)
209                                }
210                                Operation::Delete => kv::WatchEvent::Delete(key),
211                                // TODO: What is Purge? Not urgent, NATS impl not used
212                                Operation::Purge => kv::WatchEvent::Delete(key),
213                            })
214                        }
215                        Err(e) => {
216                            tracing::error!(error=%e, "watch fatal err");
217                            None
218                        }
219                    }
220                },
221            ),
222        ))
223    }
224
225    async fn entries(&self) -> Result<HashMap<kv::Key, bytes::Bytes>, StoreError> {
226        let mut key_stream = self
227            .nats_store
228            .keys()
229            .await
230            .map_err(|e| StoreError::NATSError(e.to_string()))?;
231        let mut out = HashMap::new();
232        while let Some(Ok(key)) = key_stream.next().await {
233            if let Ok(Some(entry)) = self.nats_store.entry(&key).await {
234                out.insert(kv::Key::new(key), entry.value);
235            }
236        }
237        Ok(out)
238    }
239}
240
241impl NATSBucket {
242    async fn create(&self, key: &kv::Key, value: bytes::Bytes) -> Result<StoreOutcome, StoreError> {
243        match self.nats_store.create(&key, value).await {
244            Ok(revision) => Ok(StoreOutcome::Created(revision)),
245            Err(err) if err.kind() == async_nats::jetstream::kv::CreateErrorKind::AlreadyExists => {
246                // key exists, get the revsion
247                match self.nats_store.entry(key).await {
248                    Ok(Some(entry)) => Ok(StoreOutcome::Exists(entry.revision)),
249                    Ok(None) => {
250                        tracing::error!(
251                            %key,
252                            "Race condition, key deleted between create and fetch. Retry."
253                        );
254                        Err(StoreError::Retry)
255                    }
256                    Err(err) => Err(StoreError::NATSError(err.to_string())),
257                }
258            }
259            Err(err) => Err(StoreError::NATSError(err.to_string())),
260        }
261    }
262
263    async fn update(
264        &self,
265        key: &kv::Key,
266        value: bytes::Bytes,
267        revision: u64,
268    ) -> Result<StoreOutcome, StoreError> {
269        match self.nats_store.update(key, value.clone(), revision).await {
270            Ok(revision) => Ok(StoreOutcome::Created(revision)),
271            Err(err)
272                if err.kind() == async_nats::jetstream::kv::UpdateErrorKind::WrongLastRevision =>
273            {
274                tracing::warn!(revision, %key, "Update WrongLastRevision, resync");
275                self.resync_update(key, value).await
276            }
277            Err(err) => Err(StoreError::NATSError(err.to_string())),
278        }
279    }
280
281    /// We have the wrong revision for a key. Fetch it's entry to get the correct revision,
282    /// and try the update again.
283    async fn resync_update(
284        &self,
285        key: &kv::Key,
286        value: bytes::Bytes,
287    ) -> Result<StoreOutcome, StoreError> {
288        match self.nats_store.entry(key).await {
289            Ok(Some(entry)) => {
290                // Re-try the update with new version number
291                let next_rev = entry.revision + 1;
292                match self.nats_store.update(key, value, next_rev).await {
293                    Ok(correct_revision) => Ok(StoreOutcome::Created(correct_revision)),
294                    Err(err) => Err(StoreError::NATSError(format!(
295                        "Error during update of key {key} after resync: {err}"
296                    ))),
297                }
298            }
299            Ok(None) => {
300                tracing::warn!(%key, "Entry does not exist during resync, creating.");
301                self.create(key, value).await
302            }
303            Err(err) => {
304                tracing::error!(%key, %err, "Failed fetching entry during resync");
305                Err(StoreError::NATSError(err.to_string()))
306            }
307        }
308    }
309}
310
311/// async-nats won't let us use a multi-part subject to create KV buckets (and probably many other
312/// things).
313fn single_name(namespace: &str, name: &Slug) -> String {
314    format!("{namespace}_{name}")
315}
316
317#[cfg(feature = "integration")]
318#[cfg(test)]
319mod compare_and_replace_tests {
320    use std::sync::Arc;
321
322    use tokio::sync::Barrier;
323
324    use super::*;
325    use crate::storage::kv::{Bucket as _, Key, Store as _};
326
327    #[tokio::test]
328    async fn delete_wins_race_with_compare_and_replace() {
329        let client = crate::transports::nats::ClientOptions::default()
330            .connect()
331            .await
332            .unwrap();
333        let endpoint = EndpointId {
334            namespace: "test".to_string(),
335            component: "storage".to_string(),
336            name: "compare-and-replace".to_string(),
337        };
338        let store = NATSStore::new(client, endpoint);
339        let bucket_name = format!("compare_and_replace_{}", uuid::Uuid::new_v4());
340        let jetstream_bucket_name = single_name("test", &Slug::slugify(&bucket_name));
341        let bucket = Arc::new(
342            store
343                .get_or_create_bucket(&bucket_name, None)
344                .await
345                .unwrap(),
346        );
347        let key = Key::new("model".to_string());
348        bucket.insert(&key, "old".into(), 0).await.unwrap();
349
350        let barrier = Arc::new(Barrier::new(3));
351        let task_bucket = bucket.clone();
352        let task_key = key.clone();
353        let task_barrier = barrier.clone();
354        let update = tokio::spawn(async move {
355            task_barrier.wait().await;
356            task_bucket
357                .compare_and_replace(&task_key, "old".into(), "new".into())
358                .await
359        });
360        let task_bucket = bucket.clone();
361        let task_key = key.clone();
362        let task_barrier = barrier.clone();
363        let delete = tokio::spawn(async move {
364            task_barrier.wait().await;
365            task_bucket.delete(&task_key).await
366        });
367
368        barrier.wait().await;
369        let update_result = update.await.unwrap();
370        delete.await.unwrap().unwrap();
371        assert!(update_result.is_ok() || matches!(update_result, Err(StoreError::MissingKey(_))));
372        assert_eq!(bucket.get(&key).await.unwrap(), None);
373
374        drop(bucket);
375        store
376            .client
377            .jetstream()
378            .delete_key_value(&jetstream_bucket_name)
379            .await
380            .unwrap();
381    }
382}