Skip to main content

dynamo_runtime/storage/kv/
mem.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;
5use std::collections::hash_map::Entry;
6use std::pin::Pin;
7use std::sync::Arc;
8use std::time::Duration;
9
10use async_trait::async_trait;
11use rand::Rng as _;
12use tokio::sync::broadcast;
13
14use super::{Bucket, Key, KeyValue, Store, StoreError, StoreOutcome, WatchEvent};
15
16const MEMORY_EVENT_BUFFER_CAPACITY: usize = 16_384;
17
18#[derive(Clone, Debug)]
19enum MemoryEvent {
20    Put {
21        bucket: String,
22        key: String,
23        value: bytes::Bytes,
24    },
25    Delete {
26        bucket: String,
27        key: String,
28    },
29}
30
31#[derive(Clone)]
32pub struct MemoryStore {
33    inner: Arc<MemoryStoreInner>,
34    connection_id: u64,
35}
36
37impl Default for MemoryStore {
38    fn default() -> Self {
39        Self::new()
40    }
41}
42
43struct MemoryStoreInner {
44    data: parking_lot::Mutex<HashMap<String, MemoryBucket>>,
45    change_sender: broadcast::Sender<MemoryEvent>,
46}
47
48pub struct MemoryBucketRef {
49    name: String,
50    inner: Arc<MemoryStoreInner>,
51}
52
53struct MemoryBucket {
54    data: HashMap<String, (u64, bytes::Bytes)>,
55}
56
57impl MemoryBucket {
58    fn new() -> Self {
59        MemoryBucket {
60            data: HashMap::new(),
61        }
62    }
63}
64
65impl MemoryStore {
66    pub(super) fn new() -> Self {
67        let (tx, _) = broadcast::channel(MEMORY_EVENT_BUFFER_CAPACITY);
68        MemoryStore {
69            inner: Arc::new(MemoryStoreInner {
70                data: parking_lot::Mutex::new(HashMap::new()),
71                change_sender: tx,
72            }),
73            connection_id: rand::rng().random(),
74        }
75    }
76}
77
78#[async_trait]
79impl Store for MemoryStore {
80    type Bucket = MemoryBucketRef;
81
82    async fn get_or_create_bucket(
83        &self,
84        bucket_name: &str,
85        // MemoryStore doesn't respect TTL yet
86        _ttl: Option<Duration>,
87    ) -> Result<Self::Bucket, StoreError> {
88        let mut locked_data = self.inner.data.lock();
89        // Ensure the bucket exists
90        locked_data
91            .entry(bucket_name.to_string())
92            .or_insert_with(MemoryBucket::new);
93        // Return an object able to access it
94        Ok(MemoryBucketRef {
95            name: bucket_name.to_string(),
96            inner: self.inner.clone(),
97        })
98    }
99
100    /// This operation cannot fail on MemoryStore. Always returns Ok.
101    async fn get_bucket(&self, bucket_name: &str) -> Result<Option<Self::Bucket>, StoreError> {
102        let locked_data = self.inner.data.lock();
103        match locked_data.get(bucket_name) {
104            Some(_) => Ok(Some(MemoryBucketRef {
105                name: bucket_name.to_string(),
106                inner: self.inner.clone(),
107            })),
108            None => Ok(None),
109        }
110    }
111
112    fn connection_id(&self) -> u64 {
113        self.connection_id
114    }
115
116    fn shutdown(&self) {}
117}
118
119#[async_trait]
120impl Bucket for MemoryBucketRef {
121    async fn insert(
122        &self,
123        key: &Key,
124        value: bytes::Bytes,
125        revision: u64,
126    ) -> Result<StoreOutcome, StoreError> {
127        let mut locked_data = self.inner.data.lock();
128        let mut b = locked_data.get_mut(&self.name);
129        let Some(bucket) = b.as_mut() else {
130            return Err(StoreError::MissingBucket(self.name.to_string()));
131        };
132        let outcome = match bucket.data.entry(key.to_string()) {
133            Entry::Vacant(e) => {
134                e.insert((revision, value.clone()));
135                let _ = self.inner.change_sender.send(MemoryEvent::Put {
136                    bucket: self.name.clone(),
137                    key: key.to_string(),
138                    value,
139                });
140                StoreOutcome::Created(revision)
141            }
142            Entry::Occupied(mut entry) => {
143                let (rev, _v) = entry.get();
144                if *rev == revision {
145                    StoreOutcome::Exists(revision)
146                } else {
147                    entry.insert((revision, value.clone()));
148                    let _ = self.inner.change_sender.send(MemoryEvent::Put {
149                        bucket: self.name.clone(),
150                        key: key.to_string(),
151                        value,
152                    });
153                    StoreOutcome::Created(revision)
154                }
155            }
156        };
157        Ok(outcome)
158    }
159
160    async fn get(&self, key: &Key) -> Result<Option<bytes::Bytes>, StoreError> {
161        let locked_data = self.inner.data.lock();
162        let Some(bucket) = locked_data.get(&self.name) else {
163            return Ok(None);
164        };
165        Ok(bucket.data.get(&key.0).map(|(_, v)| v.clone()))
166    }
167
168    async fn delete(&self, key: &Key) -> Result<(), StoreError> {
169        let mut locked_data = self.inner.data.lock();
170        let Some(bucket) = locked_data.get_mut(&self.name) else {
171            return Err(StoreError::MissingBucket(self.name.to_string()));
172        };
173        if bucket.data.remove(&key.0).is_some() {
174            let _ = self.inner.change_sender.send(MemoryEvent::Delete {
175                bucket: self.name.clone(),
176                key: key.to_string(),
177            });
178        }
179        Ok(())
180    }
181
182    /// All current values in the bucket first, then block waiting for new
183    /// values to be published.
184    async fn watch(
185        &self,
186    ) -> Result<Pin<Box<dyn futures::Stream<Item = WatchEvent> + Send + 'life0>>, StoreError> {
187        // Subscribe while holding the data lock so the snapshot and incremental stream have no
188        // race: mutations take the same lock before broadcasting their update.
189        let data_lock = self.inner.data.lock();
190        let Some(bucket) = data_lock.get(&self.name) else {
191            return Err(StoreError::MissingBucket(self.name.to_string()));
192        };
193        let mut changes = self.inner.change_sender.subscribe();
194        let existing_items: Vec<_> = bucket
195            .data
196            .iter()
197            .map(|(key, (_revision, value))| {
198                WatchEvent::Put(KeyValue::new(Key::new(key.clone()), value.clone()))
199            })
200            .collect();
201        drop(data_lock);
202        let bucket_name = self.name.clone();
203        let inner = self.inner.clone();
204
205        Ok(Box::pin(async_stream::stream! {
206            for event in existing_items {
207                yield event;
208            }
209            loop {
210                match changes.recv().await {
211                    Ok(MemoryEvent::Put { bucket, key, value }) => {
212                        if bucket != bucket_name {
213                            continue;
214                        }
215                        let item = KeyValue::new(Key::new(key), value);
216                        yield WatchEvent::Put(item);
217                    },
218                    Ok(MemoryEvent::Delete { bucket, key }) => {
219                        if bucket != bucket_name {
220                            continue;
221                        }
222                        yield WatchEvent::Delete(Key::new(key));
223                    },
224                    Err(broadcast::error::RecvError::Lagged(_)) => {
225                        let snapshot = {
226                            let data = inner.data.lock();
227                            // Discard retained events that predate this authoritative snapshot.
228                            // Mutations take the same lock before broadcasting, so the replacement
229                            // receiver observes every update that follows the snapshot.
230                            changes = inner.change_sender.subscribe();
231                            data.get(&bucket_name)
232                                .map(|bucket| {
233                                    bucket.data
234                                        .iter()
235                                        .map(|(key, (_revision, value))| {
236                                            (Key::new(key.clone()), value.clone())
237                                        })
238                                        .collect()
239                                })
240                                .unwrap_or_default()
241                        };
242                        yield WatchEvent::Resync(snapshot);
243                    },
244                    Err(broadcast::error::RecvError::Closed) => break,
245                }
246            }
247        }))
248    }
249
250    async fn entries(&self) -> Result<HashMap<Key, bytes::Bytes>, StoreError> {
251        let locked_data = self.inner.data.lock();
252        match locked_data.get(&self.name) {
253            Some(bucket) => {
254                let mut out = HashMap::new();
255                for (k, (_rev, v)) in bucket.data.iter() {
256                    let key = Key::new([self.name.clone(), k.to_string()].join("/"));
257                    let value = v.clone();
258                    out.insert(key, value);
259                }
260                Ok(out)
261            }
262            None => Err(StoreError::MissingBucket(self.name.clone())),
263        }
264    }
265}
266
267#[cfg(test)]
268mod tests {
269    use super::MEMORY_EVENT_BUFFER_CAPACITY;
270    use crate::storage::kv::{Bucket as _, Key, MemoryStore, Store as _, WatchEvent};
271    use futures::StreamExt;
272    use std::collections::HashSet;
273    use std::time::Duration;
274
275    #[tokio::test]
276    async fn multiple_watchers_receive_updates_without_cross_bucket_stealing() {
277        let store = MemoryStore::new();
278        let bucket_a = store.get_or_create_bucket("bucket-a", None).await.unwrap();
279        let bucket_b = store.get_or_create_bucket("bucket-b", None).await.unwrap();
280        let mut a_first = bucket_a.watch().await.unwrap();
281        let mut a_second = bucket_a.watch().await.unwrap();
282        let mut b = bucket_b.watch().await.unwrap();
283
284        bucket_a
285            .insert(&Key::new("shared-key".to_string()), "a".into(), 1)
286            .await
287            .unwrap();
288        bucket_b
289            .insert(&Key::new("shared-key".to_string()), "b".into(), 1)
290            .await
291            .unwrap();
292
293        for watcher in [&mut a_first, &mut a_second] {
294            let WatchEvent::Put(item) = watcher.next().await.unwrap() else {
295                panic!("expected bucket-a put");
296            };
297            assert_eq!(item.value(), b"a");
298        }
299        let WatchEvent::Put(item) = b.next().await.unwrap() else {
300            panic!("expected bucket-b put");
301        };
302        assert_eq!(item.value(), b"b");
303    }
304
305    #[tokio::test]
306    async fn watcher_observes_updates_to_an_existing_key() {
307        let store = MemoryStore::new();
308        let bucket = store.get_or_create_bucket("bucket", None).await.unwrap();
309        bucket
310            .insert(&Key::new("key".to_string()), "old".into(), 1)
311            .await
312            .unwrap();
313        let mut watcher = bucket.watch().await.unwrap();
314        assert!(matches!(watcher.next().await, Some(WatchEvent::Put(_))));
315
316        bucket
317            .insert(&Key::new("key".to_string()), "new".into(), 2)
318            .await
319            .unwrap();
320        let WatchEvent::Put(item) = watcher.next().await.unwrap() else {
321            panic!("expected updated put");
322        };
323        assert_eq!(item.value(), b"new");
324    }
325
326    #[tokio::test]
327    async fn lag_resync_discards_retained_pre_snapshot_events() {
328        let store = MemoryStore::new();
329        let bucket = store.get_or_create_bucket("bucket", None).await.unwrap();
330        let mut watcher = bucket.watch().await.unwrap();
331        let key = Key::new("key".to_string());
332
333        for revision in 1..=MEMORY_EVENT_BUFFER_CAPACITY + 1 {
334            bucket
335                .insert(&key, format!("value-{revision}").into(), revision as u64)
336                .await
337                .unwrap();
338        }
339
340        let WatchEvent::Resync(snapshot) = watcher.next().await.unwrap() else {
341            panic!("expected authoritative resync after lag");
342        };
343        assert_eq!(snapshot.len(), 1);
344        let latest_value = format!("value-{}", MEMORY_EVENT_BUFFER_CAPACITY + 1);
345        assert_eq!(snapshot.get(&key).unwrap(), latest_value.as_bytes());
346
347        bucket
348            .insert(
349                &key,
350                "post-resync".into(),
351                (MEMORY_EVENT_BUFFER_CAPACITY + 2) as u64,
352            )
353            .await
354            .unwrap();
355        let next = tokio::time::timeout(Duration::from_secs(1), watcher.next())
356            .await
357            .expect("post-resync update timed out")
358            .unwrap();
359        let WatchEvent::Put(item) = next else {
360            panic!("expected post-resync put");
361        };
362        assert_eq!(item.value(), b"post-resync");
363    }
364
365    #[tokio::test]
366    async fn test_entries_full_path() {
367        let m = MemoryStore::new();
368        let bucket = m.get_or_create_bucket("bucket1", None).await.unwrap();
369        let _ = bucket
370            .insert(&Key::new("key1".to_string()), "value1".into(), 0)
371            .await
372            .unwrap();
373        let _ = bucket
374            .insert(&Key::new("key2".to_string()), "value2".into(), 0)
375            .await
376            .unwrap();
377        let entries = bucket.entries().await.unwrap();
378        let keys: HashSet<Key> = entries.into_keys().collect();
379        assert!(keys.contains(&Key::new("bucket1/key1".to_string())));
380        assert!(keys.contains(&Key::new("bucket1/key2".to_string())));
381    }
382}