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 revision == 0 || *rev == revision {
145                    StoreOutcome::Exists(*rev)
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 compare_and_replace(
161        &self,
162        key: &Key,
163        expected: bytes::Bytes,
164        value: bytes::Bytes,
165    ) -> Result<StoreOutcome, StoreError> {
166        let mut locked_data = self.inner.data.lock();
167        let bucket = locked_data
168            .get_mut(&self.name)
169            .ok_or_else(|| StoreError::MissingBucket(self.name.to_string()))?;
170        let entry = bucket
171            .data
172            .get_mut(&key.0)
173            .ok_or_else(|| StoreError::MissingKey(key.to_string()))?;
174        if entry.1 != expected {
175            return Err(StoreError::Retry);
176        }
177        let next_revision = entry.0.saturating_add(1).max(1);
178        *entry = (next_revision, value.clone());
179        let _ = self.inner.change_sender.send(MemoryEvent::Put {
180            bucket: self.name.clone(),
181            key: key.to_string(),
182            value,
183        });
184        Ok(StoreOutcome::Created(next_revision))
185    }
186
187    async fn get(&self, key: &Key) -> Result<Option<bytes::Bytes>, StoreError> {
188        let locked_data = self.inner.data.lock();
189        let Some(bucket) = locked_data.get(&self.name) else {
190            return Ok(None);
191        };
192        Ok(bucket.data.get(&key.0).map(|(_, v)| v.clone()))
193    }
194
195    async fn delete(&self, key: &Key) -> Result<(), StoreError> {
196        let mut locked_data = self.inner.data.lock();
197        let Some(bucket) = locked_data.get_mut(&self.name) else {
198            return Err(StoreError::MissingBucket(self.name.to_string()));
199        };
200        if bucket.data.remove(&key.0).is_some() {
201            let _ = self.inner.change_sender.send(MemoryEvent::Delete {
202                bucket: self.name.clone(),
203                key: key.to_string(),
204            });
205        }
206        Ok(())
207    }
208
209    /// All current values in the bucket first, then block waiting for new
210    /// values to be published.
211    async fn watch(
212        &self,
213    ) -> Result<Pin<Box<dyn futures::Stream<Item = WatchEvent> + Send + 'life0>>, StoreError> {
214        // Subscribe while holding the data lock so the snapshot and incremental stream have no
215        // race: mutations take the same lock before broadcasting their update.
216        let data_lock = self.inner.data.lock();
217        let Some(bucket) = data_lock.get(&self.name) else {
218            return Err(StoreError::MissingBucket(self.name.to_string()));
219        };
220        let mut changes = self.inner.change_sender.subscribe();
221        let existing_items: Vec<_> = bucket
222            .data
223            .iter()
224            .map(|(key, (_revision, value))| {
225                WatchEvent::Put(KeyValue::new(Key::new(key.clone()), value.clone()))
226            })
227            .collect();
228        drop(data_lock);
229        let bucket_name = self.name.clone();
230        let inner = self.inner.clone();
231
232        Ok(Box::pin(async_stream::stream! {
233            for event in existing_items {
234                yield event;
235            }
236            loop {
237                match changes.recv().await {
238                    Ok(MemoryEvent::Put { bucket, key, value }) => {
239                        if bucket != bucket_name {
240                            continue;
241                        }
242                        let item = KeyValue::new(Key::new(key), value);
243                        yield WatchEvent::Put(item);
244                    },
245                    Ok(MemoryEvent::Delete { bucket, key }) => {
246                        if bucket != bucket_name {
247                            continue;
248                        }
249                        yield WatchEvent::Delete(Key::new(key));
250                    },
251                    Err(broadcast::error::RecvError::Lagged(_)) => {
252                        let snapshot = {
253                            let data = inner.data.lock();
254                            // Discard retained events that predate this authoritative snapshot.
255                            // Mutations take the same lock before broadcasting, so the replacement
256                            // receiver observes every update that follows the snapshot.
257                            changes = inner.change_sender.subscribe();
258                            data.get(&bucket_name)
259                                .map(|bucket| {
260                                    bucket.data
261                                        .iter()
262                                        .map(|(key, (_revision, value))| {
263                                            (Key::new(key.clone()), value.clone())
264                                        })
265                                        .collect()
266                                })
267                                .unwrap_or_default()
268                        };
269                        yield WatchEvent::Resync(snapshot);
270                    },
271                    Err(broadcast::error::RecvError::Closed) => break,
272                }
273            }
274        }))
275    }
276
277    async fn entries(&self) -> Result<HashMap<Key, bytes::Bytes>, StoreError> {
278        let locked_data = self.inner.data.lock();
279        match locked_data.get(&self.name) {
280            Some(bucket) => {
281                let mut out = HashMap::new();
282                for (k, (_rev, v)) in bucket.data.iter() {
283                    let key = Key::new([self.name.clone(), k.to_string()].join("/"));
284                    let value = v.clone();
285                    out.insert(key, value);
286                }
287                Ok(out)
288            }
289            None => Err(StoreError::MissingBucket(self.name.clone())),
290        }
291    }
292}
293
294#[cfg(test)]
295mod tests {
296    use super::MEMORY_EVENT_BUFFER_CAPACITY;
297    use crate::storage::kv::{
298        Bucket as _, Key, MemoryStore, Store as _, StoreError, StoreOutcome, WatchEvent,
299    };
300    use futures::StreamExt;
301    use std::collections::HashSet;
302    use std::sync::Arc;
303    use std::time::Duration;
304    use tokio::sync::Barrier;
305
306    #[tokio::test]
307    async fn delete_wins_race_with_compare_and_replace() {
308        let store = MemoryStore::new();
309        let bucket = Arc::new(store.get_or_create_bucket("bucket", None).await.unwrap());
310        let key = Key::new("model".to_string());
311        bucket.insert(&key, "old".into(), 0).await.unwrap();
312
313        let barrier = Arc::new(Barrier::new(3));
314        let update_bucket = bucket.clone();
315        let update_key = key.clone();
316        let update_barrier = barrier.clone();
317        let update = tokio::spawn(async move {
318            update_barrier.wait().await;
319            update_bucket
320                .compare_and_replace(&update_key, "old".into(), "new".into())
321                .await
322        });
323        let delete_bucket = bucket.clone();
324        let delete_key = key.clone();
325        let delete_barrier = barrier.clone();
326        let delete = tokio::spawn(async move {
327            delete_barrier.wait().await;
328            delete_bucket.delete(&delete_key).await
329        });
330
331        barrier.wait().await;
332        let update_result = update.await.unwrap();
333        delete.await.unwrap().unwrap();
334        assert!(update_result.is_ok() || matches!(update_result, Err(StoreError::MissingKey(_))));
335        assert_eq!(bucket.get(&key).await.unwrap(), None);
336    }
337
338    #[tokio::test]
339    async fn revision_zero_replay_preserves_cas_update() {
340        let store = MemoryStore::new();
341        let bucket = store.get_or_create_bucket("bucket", None).await.unwrap();
342        let key = Key::new("model".to_string());
343
344        assert_eq!(
345            bucket.insert(&key, "initial".into(), 0).await.unwrap(),
346            StoreOutcome::Created(0)
347        );
348        assert_eq!(
349            bucket
350                .compare_and_replace(&key, "initial".into(), "updated".into())
351                .await
352                .unwrap(),
353            StoreOutcome::Created(1)
354        );
355        assert_eq!(
356            bucket.insert(&key, "initial".into(), 0).await.unwrap(),
357            StoreOutcome::Exists(1)
358        );
359        assert_eq!(
360            bucket.get(&key).await.unwrap().unwrap().as_ref(),
361            b"updated"
362        );
363    }
364
365    #[tokio::test]
366    async fn multiple_watchers_receive_updates_without_cross_bucket_stealing() {
367        let store = MemoryStore::new();
368        let bucket_a = store.get_or_create_bucket("bucket-a", None).await.unwrap();
369        let bucket_b = store.get_or_create_bucket("bucket-b", None).await.unwrap();
370        let mut a_first = bucket_a.watch().await.unwrap();
371        let mut a_second = bucket_a.watch().await.unwrap();
372        let mut b = bucket_b.watch().await.unwrap();
373
374        bucket_a
375            .insert(&Key::new("shared-key".to_string()), "a".into(), 1)
376            .await
377            .unwrap();
378        bucket_b
379            .insert(&Key::new("shared-key".to_string()), "b".into(), 1)
380            .await
381            .unwrap();
382
383        for watcher in [&mut a_first, &mut a_second] {
384            let WatchEvent::Put(item) = watcher.next().await.unwrap() else {
385                panic!("expected bucket-a put");
386            };
387            assert_eq!(item.value(), b"a");
388        }
389        let WatchEvent::Put(item) = b.next().await.unwrap() else {
390            panic!("expected bucket-b put");
391        };
392        assert_eq!(item.value(), b"b");
393    }
394
395    #[tokio::test]
396    async fn watcher_observes_updates_to_an_existing_key() {
397        let store = MemoryStore::new();
398        let bucket = store.get_or_create_bucket("bucket", None).await.unwrap();
399        bucket
400            .insert(&Key::new("key".to_string()), "old".into(), 1)
401            .await
402            .unwrap();
403        let mut watcher = bucket.watch().await.unwrap();
404        assert!(matches!(watcher.next().await, Some(WatchEvent::Put(_))));
405
406        bucket
407            .insert(&Key::new("key".to_string()), "new".into(), 2)
408            .await
409            .unwrap();
410        let WatchEvent::Put(item) = watcher.next().await.unwrap() else {
411            panic!("expected updated put");
412        };
413        assert_eq!(item.value(), b"new");
414    }
415
416    #[tokio::test]
417    async fn lag_resync_discards_retained_pre_snapshot_events() {
418        let store = MemoryStore::new();
419        let bucket = store.get_or_create_bucket("bucket", None).await.unwrap();
420        let mut watcher = bucket.watch().await.unwrap();
421        let key = Key::new("key".to_string());
422
423        for revision in 1..=MEMORY_EVENT_BUFFER_CAPACITY + 1 {
424            bucket
425                .insert(&key, format!("value-{revision}").into(), revision as u64)
426                .await
427                .unwrap();
428        }
429
430        let WatchEvent::Resync(snapshot) = watcher.next().await.unwrap() else {
431            panic!("expected authoritative resync after lag");
432        };
433        assert_eq!(snapshot.len(), 1);
434        let latest_value = format!("value-{}", MEMORY_EVENT_BUFFER_CAPACITY + 1);
435        assert_eq!(snapshot.get(&key).unwrap(), latest_value.as_bytes());
436
437        bucket
438            .insert(
439                &key,
440                "post-resync".into(),
441                (MEMORY_EVENT_BUFFER_CAPACITY + 2) as u64,
442            )
443            .await
444            .unwrap();
445        let next = tokio::time::timeout(Duration::from_secs(1), watcher.next())
446            .await
447            .expect("post-resync update timed out")
448            .unwrap();
449        let WatchEvent::Put(item) = next else {
450            panic!("expected post-resync put");
451        };
452        assert_eq!(item.value(), b"post-resync");
453    }
454
455    #[tokio::test]
456    async fn test_entries_full_path() {
457        let m = MemoryStore::new();
458        let bucket = m.get_or_create_bucket("bucket1", None).await.unwrap();
459        let _ = bucket
460            .insert(&Key::new("key1".to_string()), "value1".into(), 0)
461            .await
462            .unwrap();
463        let _ = bucket
464            .insert(&Key::new("key2".to_string()), "value2".into(), 0)
465            .await
466            .unwrap();
467        let entries = bucket.entries().await.unwrap();
468        let keys: HashSet<Key> = entries.into_keys().collect();
469        assert!(keys.contains(&Key::new("bucket1/key1".to_string())));
470        assert!(keys.contains(&Key::new("bucket1/key2".to_string())));
471    }
472}