dynamo_runtime/storage/kv/
mem.rs1use 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 _ttl: Option<Duration>,
87 ) -> Result<Self::Bucket, StoreError> {
88 let mut locked_data = self.inner.data.lock();
89 locked_data
91 .entry(bucket_name.to_string())
92 .or_insert_with(MemoryBucket::new);
93 Ok(MemoryBucketRef {
95 name: bucket_name.to_string(),
96 inner: self.inner.clone(),
97 })
98 }
99
100 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 async fn watch(
212 &self,
213 ) -> Result<Pin<Box<dyn futures::Stream<Item = WatchEvent> + Send + 'life0>>, StoreError> {
214 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 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}