1use std::collections::BTreeMap;
4use std::fmt;
5use std::ops::Bound;
6use std::sync::{Arc, Mutex};
7
8use super::{MemoryFault, lock, take_fault};
9use crate::rt::Clock;
10use crate::store::{
11 Batch, BatchOutcome, Cursor, Key, NamespaceStore, Partition, PartitionStats, Precondition,
12 ScanPage, StoreCapabilities, StoreError, Value, Write,
13};
14
15type Rows = BTreeMap<Key, Value>;
16
17struct RowEdit<'a> {
20 rows: &'a mut Rows,
21 undo: Vec<(Key, Option<Value>)>,
22}
23
24impl<'a> RowEdit<'a> {
25 fn new(rows: &'a mut Rows, writes: usize) -> Self {
26 Self {
27 rows,
28 undo: Vec::with_capacity(writes),
29 }
30 }
31
32 fn write(&mut self, write: Write) {
33 let key = match &write {
34 Write::Put(key, _) | Write::Delete(key) => key,
35 };
36 self.undo.push((key.clone(), self.rows.get(key).cloned()));
38 match write {
39 Write::Put(key, value) => self.rows.insert(key, value),
40 Write::Delete(key) => self.rows.remove(&key),
41 };
42 }
43
44 fn commit(mut self) {
45 self.undo.clear();
46 }
47}
48
49impl Drop for RowEdit<'_> {
50 fn drop(&mut self) {
51 for (key, prior) in self.undo.drain(..).rev() {
52 match prior {
53 Some(value) => self.rows.insert(key, value),
54 None => self.rows.remove(&key),
55 };
56 }
57 }
58}
59
60pub struct MemoryKv {
67 partitions: Mutex<BTreeMap<Partition, Rows>>,
68 clock: Arc<dyn Clock>,
69 caps: StoreCapabilities,
70 capacity: Option<u64>,
71 fault: Mutex<Option<MemoryFault>>,
72}
73
74impl fmt::Debug for MemoryKv {
75 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
76 f.debug_struct("MemoryKv")
77 .field("caps", &self.caps)
78 .field("capacity", &self.capacity)
79 .finish_non_exhaustive()
80 }
81}
82
83#[cfg(not(target_arch = "wasm32"))]
85impl Default for MemoryKv {
86 fn default() -> Self {
87 Self::with_clock(Arc::new(crate::rt::SystemClock))
88 }
89}
90
91impl MemoryKv {
92 #[must_use]
94 pub fn with_clock(clock: Arc<dyn Clock>) -> Self {
95 Self {
96 partitions: Mutex::default(),
97 clock,
98 caps: StoreCapabilities::full(),
99 capacity: None,
100 fault: Mutex::default(),
101 }
102 }
103
104 #[cfg(not(target_arch = "wasm32"))]
106 #[must_use]
107 pub fn new(caps: StoreCapabilities) -> Self {
108 Self::default().with_capabilities(caps)
109 }
110
111 #[must_use]
113 pub fn with_capabilities(mut self, caps: StoreCapabilities) -> Self {
114 self.caps = caps;
115 self
116 }
117
118 #[must_use]
120 pub fn with_fault(self, fault: MemoryFault) -> Self {
121 *lock(&self.fault) = Some(fault);
122 self
123 }
124
125 #[must_use]
129 pub fn with_capacity_limit(mut self, bytes: u64) -> Self {
130 self.capacity = Some(bytes);
131 self
132 }
133}
134
135fn size(rows: &Rows) -> u64 {
136 rows.iter()
137 .map(|(k, v)| (k.as_bytes().len() + v.as_bytes().len()) as u64)
138 .sum()
139}
140
141fn check(rows: &Rows, index: usize, pre: &Precondition, now: u64) -> Option<BatchOutcome> {
144 let (holds, observed) = match pre {
145 Precondition::Absent(k) => (!rows.contains_key(k), rows.get(k).cloned()),
146 Precondition::Present(k) => (rows.contains_key(k), None),
147 Precondition::Equals(k, v) => (rows.get(k) == Some(v), rows.get(k).cloned()),
148 Precondition::NotAfter(deadline) => {
149 return (now > *deadline).then_some(BatchOutcome::DeadlinePassed { backend_now: now });
150 }
151 };
152 (!holds).then_some(BatchOutcome::PreconditionFailed { index, observed })
153}
154
155impl NamespaceStore for MemoryKv {
156 fn capabilities(&self) -> StoreCapabilities {
157 self.caps
158 }
159
160 async fn get(&self, p: &Partition, key: &Key) -> Result<Option<Value>, StoreError> {
161 Ok(lock(&self.partitions)
162 .get(p)
163 .and_then(|rows| rows.get(key))
164 .cloned())
165 }
166
167 async fn scan(
168 &self,
169 p: &Partition,
170 start: &Key,
171 end: &Key,
172 after: Option<&Cursor>,
173 limit: u32,
174 ) -> Result<ScanPage, StoreError> {
175 if limit == 0 {
176 return Err(StoreError::Invalid("scan limit must be at least 1".into()));
177 }
178 let lower = match after.map(|c| Key::new(c.clone().into_bytes())) {
179 None => Bound::Included(start.clone()),
180 Some(cursor) if *start <= cursor && cursor < *end => Bound::Excluded(cursor),
182 Some(_) => {
183 return Err(StoreError::Invalid(
184 "scan cursor outside the scanned range".into(),
185 ));
186 }
187 };
188 let empty = match &lower {
189 Bound::Included(k) | Bound::Excluded(k) => k >= end,
190 Bound::Unbounded => false,
191 };
192 let partitions = lock(&self.partitions);
193 let (Some(rows), false) = (partitions.get(p), empty) else {
194 return Ok(ScanPage::default());
195 };
196 let want = usize::try_from(limit).unwrap_or(usize::MAX);
197 let mut entries: Vec<_> = rows
198 .range((lower, Bound::Excluded(end.clone())))
199 .take(want.saturating_add(1))
200 .map(|(k, v)| (k.clone(), v.clone()))
201 .collect();
202 let next = (entries.len() > want).then(|| {
203 entries.truncate(want);
204 Cursor::new(entries[want - 1].0.clone().into_bytes())
205 });
206 Ok(ScanPage { entries, next })
207 }
208
209 async fn apply(&self, p: &Partition, batch: Batch) -> Result<BatchOutcome, StoreError> {
210 batch.validate(&self.caps)?;
211 take_fault(&self.fault, MemoryFault::ApplyBefore)?;
212 let mut partitions = lock(&self.partitions);
213 let now = u64::try_from(self.clock.now_ms()).unwrap_or(u64::MAX);
216 let empty = Rows::new();
217 let rows = partitions.get(p).unwrap_or(&empty);
218 for (index, pre) in batch.preconditions.iter().enumerate() {
219 if let Some(failed) = check(rows, index, pre, now) {
220 return Ok(failed);
221 }
222 }
223 let adds = batch.has_put();
224 let mut edit = RowEdit::new(partitions.entry(p.clone()).or_default(), batch.writes.len());
225 for write in batch.writes {
226 edit.write(write);
227 }
228 if adds && self.capacity.is_some_and(|cap| size(edit.rows) > cap) {
229 return Err(StoreError::Full);
230 }
231 edit.commit();
232 drop(partitions);
233 take_fault(&self.fault, MemoryFault::ApplyAfterCommit)?;
234 Ok(BatchOutcome::Committed)
235 }
236
237 async fn stats(&self, p: &Partition) -> Result<PartitionStats, StoreError> {
238 let partitions = lock(&self.partitions);
239 let rows = partitions.get(p);
240 Ok(PartitionStats {
241 bytes: rows.map_or(0, size),
242 keys: Some(rows.map_or(0, |r| r.len() as u64)),
243 })
244 }
245
246 async fn probe(&self) -> Result<(), StoreError> {
247 Ok(())
248 }
249}
250
251#[cfg(test)]
252mod tests {
253 use std::future::Future;
254 use std::pin::pin;
255 use std::task::{Context, Poll, Waker};
256
257 use futures_executor::block_on;
258
259 use super::*;
260 use crate::repo::{NamespaceKey, RepoName};
261 use crate::rt::ManualClock;
262 use crate::store::{
263 KeyClasses, MAX_BATCH_BYTES, MAX_BATCH_OPS, MAX_KEY_BYTES, MAX_VALUE_BYTES, keys,
264 };
265 use BatchOutcome::Committed;
266 use Precondition::{Absent, Equals, NotAfter, Present};
267
268 fn ns() -> Partition {
269 Partition::Namespace(NamespaceKey::deployment_default())
270 }
271
272 fn k(s: &str) -> Key {
273 Key::new(s.as_bytes().to_vec())
274 }
275
276 fn v(s: &str) -> Value {
277 Value::new(s.as_bytes().to_vec())
278 }
279
280 fn apply(kv: &MemoryKv, batch: Batch) -> Result<BatchOutcome, StoreError> {
281 block_on(kv.apply(&ns(), batch))
282 }
283
284 fn ok(kv: &MemoryKv, batch: Batch) -> BatchOutcome {
285 apply(kv, batch).unwrap()
286 }
287
288 fn failed(index: usize, observed: Option<Value>) -> BatchOutcome {
289 BatchOutcome::PreconditionFailed { index, observed }
290 }
291
292 fn get(kv: &MemoryKv, key: &Key) -> Option<Value> {
293 block_on(kv.get(&ns(), key)).unwrap()
294 }
295
296 fn ab(kv: &MemoryKv) -> (Option<Value>, Option<Value>) {
297 (get(kv, &k("a")), get(kv, &k("b")))
298 }
299
300 #[test]
301 fn batch_journal_restores_rows_before_recovering_a_poisoned_lock() {
302 let kv = Arc::new(MemoryKv::default());
303 ok(
304 kv.as_ref(),
305 Batch::new().put(k("a"), v("1")).put(k("b"), v("2")),
306 );
307 let before = lock(&kv.partitions).clone();
308 let for_thread = kv.clone();
309 let panicked = std::thread::spawn(move || {
310 let mut partitions = lock(&for_thread.partitions);
311 let rows = partitions.get_mut(&ns()).expect("seeded partition");
312 let mut edit = RowEdit::new(rows, 5);
313 edit.write(Write::Put(k("a"), v("changed")));
314 edit.write(Write::Delete(k("b")));
315 edit.write(Write::Put(k("c"), v("new")));
316 edit.write(Write::Delete(k("a")));
317 edit.write(Write::Delete(k("c")));
318 panic!("interrupt an uncommitted batch");
319 })
320 .join();
321 assert!(panicked.is_err() && kv.partitions.is_poisoned());
322 assert_eq!(*lock(&kv.partitions), before);
323 assert_eq!(ab(kv.as_ref()), (Some(v("1")), Some(v("2"))));
324 assert_eq!(get(kv.as_ref(), &k("c")), None);
325 }
326
327 #[test]
328 fn apply_all_or_nothing_reports_first_failing_precondition_and_observed_value() {
329 let kv = MemoryKv::default();
330 assert_eq!(ok(&kv, Batch::new().put(k("a"), v("1"))), Committed);
331 let batch = |preconditions: Vec<Precondition>| Batch {
332 preconditions,
333 writes: vec![Write::Put(k("b"), v("2")), Write::Delete(k("a"))],
334 };
335 let cases = [
336 (vec![Absent(k("a"))], failed(0, Some(v("1")))),
337 (vec![Present(k("z"))], failed(0, None)),
338 (vec![Equals(k("z"), v("1"))], failed(0, None)),
339 (
340 vec![Present(k("a")), Equals(k("a"), v("9")), Absent(k("a"))],
341 failed(1, Some(v("1"))),
342 ),
343 ];
344 for (pres, outcome) in cases {
345 assert_eq!(ok(&kv, batch(pres)), outcome);
346 assert_eq!(ab(&kv), (Some(v("1")), None));
347 }
348 let pres = vec![Equals(k("a"), v("1")), Absent(k("b"))];
349 assert_eq!(ok(&kv, batch(pres)), Committed);
350 assert_eq!(ab(&kv), (None, Some(v("2"))));
351 assert!(block_on(kv.has(&ns(), &k("b"))).unwrap());
352 let other = Partition::ContentShard(0);
353 assert_eq!(block_on(kv.get(&other, &k("b"))).unwrap(), None);
354 }
355
356 #[test]
357 fn apply_rejects_oversize_key_and_value_writing_nothing() {
358 let kv = MemoryKv::default();
359 let long_key = Key::new(vec![b'k'; MAX_KEY_BYTES + 1]);
360 let long_value = Value::new(vec![0; MAX_VALUE_BYTES + 1]);
361 let a = || Batch::new().put(k("a"), v("1"));
362 for batch in [
363 a().put(long_key.clone(), v("1")),
364 a().require(Absent(long_key)),
365 a().put(k("b"), long_value.clone()),
366 a().require(Equals(k("b"), long_value)),
367 ] {
368 assert!(matches!(apply(&kv, batch), Err(StoreError::Invalid(_))));
369 }
370 assert_eq!(ab(&kv), (None, None));
371 let max_key = Key::new(vec![1; MAX_KEY_BYTES]);
372 let max = Batch::new().put(max_key, Value::new(vec![0; MAX_VALUE_BYTES]));
373 assert_eq!(ok(&kv, max), Committed);
374 }
375
376 #[test]
377 fn apply_rejects_batches_over_the_op_and_byte_caps() {
378 let kv = MemoryKv::default();
379 let puts = |n: usize, len: usize| {
380 (0..n).fold(Batch::new(), |b, i| {
381 b.put(Key::new(i.to_be_bytes().to_vec()), Value::new(vec![0; len]))
382 })
383 };
384 assert_eq!(ok(&kv, puts(MAX_BATCH_OPS, 1)), Committed);
385 let too_many = puts(MAX_BATCH_OPS, 1).require(Absent(k("z")));
386 assert!(matches!(apply(&kv, too_many), Err(StoreError::Invalid(_))));
387 let too_big = puts(MAX_BATCH_BYTES / MAX_VALUE_BYTES + 1, MAX_VALUE_BYTES);
389 assert!(matches!(apply(&kv, too_big), Err(StoreError::Invalid(_))));
390 }
391
392 #[test]
393 fn not_after_uses_store_clock_at_apply() {
394 let clock = Arc::new(ManualClock::new(1_000));
395 let kv = MemoryKv::with_clock(clock.clone());
396 let late = Batch::new()
399 .require(NotAfter(1_500))
400 .require(Equals(k("x"), v("never")))
401 .put(k("a"), v("1"));
402 clock.set(1_501);
403 let passed = BatchOutcome::DeadlinePassed { backend_now: 1_501 };
404 assert_eq!(ok(&kv, late), passed);
405 assert_eq!(ab(&kv), (None, None));
406 clock.set(1_500);
407 let at_deadline = Batch::new().require(NotAfter(1_500)).put(k("a"), v("1"));
408 assert_eq!(ok(&kv, at_deadline), Committed);
409 clock.set(-5);
411 let never = Batch::new()
412 .require(NotAfter(u64::MAX - 1))
413 .put(k("b"), v("1"));
414 let invalid = BatchOutcome::DeadlinePassed {
415 backend_now: u64::MAX,
416 };
417 assert_eq!(ok(&kv, never), invalid);
418 assert_eq!(get(&kv, &k("b")), None);
419 }
420
421 #[test]
422 fn not_after_accepted_by_refs_only_non_atomic_store() {
423 let clock = Arc::new(ManualClock::new(10));
424 let caps = StoreCapabilities::refs_only();
425 let kv = MemoryKv::with_clock(clock).with_capabilities(caps);
426 assert_eq!(kv.capabilities().key_classes, KeyClasses::RefsOnly);
427 assert_eq!(kv.capabilities().implicit_layout_version, Some(1));
428 let repo = RepoName::new("r").unwrap();
429 let main = keys::ref_key(&repo, "refs/heads/main");
430 let dev = keys::ref_key(&repo, "refs/heads/dev");
431 let one = Batch::new()
432 .require(NotAfter(10))
433 .require(Absent(main.clone()))
434 .put(main.clone(), v("id"));
435 assert_eq!(ok(&kv, one), Committed);
436 for bad in [
437 Batch::new()
438 .put(main.clone(), v("1"))
439 .put(dev.clone(), v("2")),
440 Batch::new()
441 .require(Absent(dev.clone()))
442 .put(main.clone(), v("1")),
443 Batch::new().put(keys::layout_version(), v("1")),
444 Batch::new().require(Absent(keys::grant_epoch())),
445 ] {
446 assert!(matches!(apply(&kv, bad), Err(StoreError::Unsupported(_))));
447 }
448 assert_eq!((get(&kv, &main), get(&kv, &dev)), (Some(v("id")), None));
449 }
450
451 #[test]
452 fn capacity_limit_returns_full_but_reads_and_deletes_work() {
453 let kv = MemoryKv::default().with_capacity_limit(4);
454 assert_eq!(ok(&kv, Batch::new().put(k("a"), v("12"))), Committed);
455 let b = || Batch::new().put(k("b"), v("1"));
456 assert!(matches!(apply(&kv, b()), Err(StoreError::Full)));
457 assert_eq!(ab(&kv), (Some(v("12")), None));
458 let prune = Batch::new()
461 .require(Equals(k("a"), v("12")))
462 .require(Absent(k("b")))
463 .delete(k("a"));
464 assert_eq!(ok(&kv, prune), Committed);
465 assert_eq!(block_on(kv.stats(&ns())).unwrap().bytes, 0);
466 assert_eq!(ok(&kv, b()), Committed);
467 assert_eq!(block_on(kv.stats(&ns())).unwrap().keys, Some(1));
468 }
469
470 #[test]
471 fn scan_orders_by_bytes_and_cursor_resumes_strictly_after() {
472 let kv = MemoryKv::default();
473 let keys = [&b"a"[..], b"a\0", b"ab", b"a\xff", b"b", b"c"].map(|b| Key::new(b.to_vec()));
474 let batch = keys
475 .iter()
476 .rev()
477 .fold(Batch::new(), |b, key| b.put(key.clone(), v("x")));
478 assert_eq!(ok(&kv, batch), Committed);
479 let scan = |start: &str, after: Option<&Cursor>, limit| {
480 block_on(kv.scan(&ns(), &k(start), &k("c"), after, limit))
481 };
482 let names = |page: &ScanPage| page.entries.iter().map(|e| e.0.clone()).collect::<Vec<_>>();
483 let first = scan("a", None, 2).unwrap();
484 assert_eq!(names(&first), keys[..2]);
485 let rest = scan("a", first.next.as_ref(), 10).unwrap();
486 assert_eq!((names(&rest), rest.next), (keys[2..5].to_vec(), None));
487 assert_eq!(scan("a", None, 5).unwrap().next, None);
489 assert_eq!(scan("d", None, 1).unwrap(), ScanPage::default());
490 let foreign = Cursor::new(&b"0"[..]);
492 assert!(matches!(
493 scan("a", Some(&foreign), 1),
494 Err(StoreError::Invalid(_))
495 ));
496 assert!(matches!(scan("a", None, 0), Err(StoreError::Invalid(_))));
497 }
498
499 #[test]
500 fn get_many_preserves_order() {
501 let kv = MemoryKv::default();
502 let batch = Batch::new().put(k("a"), v("1")).put(k("b"), v("2"));
503 assert_eq!(ok(&kv, batch), Committed);
504 let got = block_on(kv.get_many(&ns(), &[k("b"), k("z"), k("a")])).unwrap();
505 assert_eq!(got, vec![Some(v("2")), None, Some(v("1"))]);
506 }
507
508 #[test]
509 fn dropped_apply_future_leaves_store_consistent() {
510 let kv = MemoryKv::default();
511 let p = ns();
512 let batch = || Batch::new().put(k("a"), v("1")).put(k("b"), v("2"));
513 drop(kv.apply(&p, batch()));
515 assert_eq!(ab(&kv), (None, None));
516 let mut fut = pin!(kv.apply(&p, batch()));
518 let poll = fut.as_mut().poll(&mut Context::from_waker(Waker::noop()));
519 assert!(matches!(poll, Poll::Ready(Ok(Committed))));
520 assert_eq!(ab(&kv), (Some(v("1")), Some(v("2"))));
521 }
522
523 #[test]
524 fn poisoned_lock_recovers() {
525 let kv = Arc::new(MemoryKv::default());
526 assert_eq!(ok(&kv, Batch::new().put(k("a"), v("1"))), Committed);
527 let held = kv.clone();
528 let panicked = std::thread::spawn(move || {
529 let _guard = held.partitions.lock().unwrap();
530 panic!("poison the store lock");
531 })
532 .join();
533 assert!(panicked.is_err() && kv.partitions.is_poisoned());
534 assert_eq!(ok(&kv, Batch::new().put(k("b"), v("2"))), Committed);
535 assert_eq!(ab(&kv), (Some(v("1")), Some(v("2"))));
536 }
537
538 #[test]
539 fn injected_faults_fire_once() {
540 let a = || Batch::new().put(k("a"), v("1"));
541 let kv = MemoryKv::default().with_fault(MemoryFault::ApplyBefore);
542 assert!(matches!(apply(&kv, a()), Err(StoreError::Unavailable(_))));
543 assert_eq!(ab(&kv), (None, None));
544 assert_eq!(ok(&kv, a()), Committed);
545 let kv = MemoryKv::default().with_fault(MemoryFault::ApplyAfterCommit);
546 assert!(matches!(apply(&kv, a()), Err(StoreError::Unavailable(_))));
547 assert_eq!(ab(&kv), (Some(v("1")), None));
548 }
549}