1use std::collections::BTreeMap;
12use std::collections::VecDeque;
13use std::ops::Bound;
14use std::ops::Bound::{Excluded, Included, Unbounded};
15use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
16use std::sync::{Arc, RwLock};
17
18use crate::types::{Key, Value};
19
20#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct MemTableEntry {
23 pub value: Option<Value>,
25 pub timestamp: u64,
27 pub sequence: u64,
29}
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub struct MemTableConfig {
34 pub flush_threshold: usize,
36 pub max_immutable_count: usize,
38}
39
40impl Default for MemTableConfig {
41 fn default() -> Self {
42 Self {
43 flush_threshold: 64 * 1024 * 1024,
44 max_immutable_count: 4,
45 }
46 }
47}
48
49fn encode_be_u64(v: u64) -> [u8; 8] {
50 v.to_be_bytes()
51}
52
53fn invert_u64(v: u64) -> u64 {
54 u64::MAX - v
55}
56
57fn internal_key_prefix(user_key: &[u8]) -> Vec<u8> {
58 let mut out = Vec::with_capacity(user_key.len() + 1);
59 out.extend_from_slice(user_key);
60 out.push(0);
61 out
62}
63
64fn internal_key(user_key: &[u8], timestamp: u64, sequence: u64) -> Vec<u8> {
65 let mut out = Vec::with_capacity(user_key.len() + 1 + 16);
66 out.extend_from_slice(user_key);
67 out.push(0);
68 out.extend_from_slice(&encode_be_u64(invert_u64(timestamp)));
69 out.extend_from_slice(&encode_be_u64(invert_u64(sequence)));
70 out
71}
72
73fn decode_user_key(internal_key: &[u8]) -> &[u8] {
74 let user_len = internal_key
77 .len()
78 .checked_sub(1 + 16)
79 .expect("internal key has fixed trailer");
80 &internal_key[..user_len]
81}
82
83fn next_prefix(prefix: &[u8]) -> Option<Vec<u8>> {
84 if prefix.is_empty() {
85 return None;
86 }
87 let mut out = prefix.to_vec();
88 for i in (0..out.len()).rev() {
89 if out[i] != 0xFF {
90 out[i] = out[i].wrapping_add(1);
91 out.truncate(i + 1);
92 return Some(out);
93 }
94 }
95 None
96}
97
98fn update_min(atom: &AtomicU64, v: u64) {
99 let mut cur = atom.load(Ordering::Relaxed);
100 while v < cur {
101 match atom.compare_exchange_weak(cur, v, Ordering::Relaxed, Ordering::Relaxed) {
102 Ok(_) => return,
103 Err(next) => cur = next,
104 }
105 }
106}
107
108fn update_max(atom: &AtomicU64, v: u64) {
109 let mut cur = atom.load(Ordering::Relaxed);
110 while v > cur {
111 match atom.compare_exchange_weak(cur, v, Ordering::Relaxed, Ordering::Relaxed) {
112 Ok(_) => return,
113 Err(next) => cur = next,
114 }
115 }
116}
117
118#[derive(Debug)]
120pub struct MemTable {
121 data: RwLock<BTreeMap<Vec<u8>, MemTableEntry>>,
123 memory_usage: AtomicUsize,
125 min_timestamp: AtomicU64,
127 max_timestamp: AtomicU64,
129}
130
131impl Default for MemTable {
132 fn default() -> Self {
133 Self::new()
134 }
135}
136
137impl MemTable {
138 pub fn new() -> Self {
140 Self {
141 data: RwLock::new(BTreeMap::new()),
142 memory_usage: AtomicUsize::new(0),
143 min_timestamp: AtomicU64::new(u64::MAX),
144 max_timestamp: AtomicU64::new(0),
145 }
146 }
147
148 pub fn memory_usage_bytes(&self) -> usize {
152 self.memory_usage.load(Ordering::Relaxed)
153 }
154
155 pub fn min_timestamp(&self) -> Option<u64> {
157 let v = self.min_timestamp.load(Ordering::Relaxed);
158 if v == u64::MAX {
159 None
160 } else {
161 Some(v)
162 }
163 }
164
165 pub fn max_timestamp(&self) -> Option<u64> {
167 let v = self.max_timestamp.load(Ordering::Relaxed);
168 if self.memory_usage_bytes() == 0 {
169 None
170 } else {
171 Some(v)
172 }
173 }
174
175 fn insert_entry(&self, user_key: &[u8], entry: MemTableEntry) {
176 let ikey = internal_key(user_key, entry.timestamp, entry.sequence);
177 let value_len = entry.value.as_ref().map(|v| v.len()).unwrap_or(0);
178 let approx_bytes = ikey.len().saturating_add(value_len);
179
180 let mut data = self.data.write().expect("memtable lock poisoned");
181 if let Some(old) = data.insert(ikey, entry.clone()) {
182 let old_value_len = old.value.as_ref().map(|v| v.len()).unwrap_or(0);
183 let old_key_len = internal_key(user_key, old.timestamp, old.sequence).len();
184 let old_bytes = old_key_len.saturating_add(old_value_len);
185 self.memory_usage
186 .fetch_sub(old_bytes.min(self.memory_usage_bytes()), Ordering::Relaxed);
187 }
188 self.memory_usage.fetch_add(approx_bytes, Ordering::Relaxed);
189 drop(data);
190
191 update_min(&self.min_timestamp, entry.timestamp);
192 update_max(&self.max_timestamp, entry.timestamp);
193 }
194
195 pub fn put(&self, key: Key, value: Value, timestamp: u64, sequence: u64) {
197 self.insert_entry(
198 &key,
199 MemTableEntry {
200 value: Some(value),
201 timestamp,
202 sequence,
203 },
204 );
205 }
206
207 pub fn delete(&self, key: Key, timestamp: u64, sequence: u64) {
209 self.insert_entry(
210 &key,
211 MemTableEntry {
212 value: None,
213 timestamp,
214 sequence,
215 },
216 );
217 }
218
219 pub fn get(&self, key: &[u8], read_timestamp: u64) -> Option<MemTableEntry> {
221 let prefix = internal_key_prefix(key);
222 let start = internal_key(key, read_timestamp, u64::MAX);
223 let end = next_prefix(&prefix);
224
225 let data = self.data.read().expect("memtable lock poisoned");
226 let range = match end {
227 Some(end_key) => data.range((Included(start), Excluded(end_key))),
228 None => data.range((Included(start), Unbounded)),
229 };
230 for (k, entry) in range {
231 if decode_user_key(k) != key {
232 break;
233 }
234 if entry.timestamp <= read_timestamp {
235 return Some(entry.clone());
236 }
237 }
238 None
239 }
240
241 fn collect_scan(
242 &self,
243 start: Bound<Vec<u8>>,
244 end: Bound<Vec<u8>>,
245 read_timestamp: u64,
246 ) -> Vec<(Key, MemTableEntry)> {
247 let data = self.data.read().expect("memtable lock poisoned");
248 let mut out = Vec::new();
249 let mut last_user_key: Option<Vec<u8>> = None;
250
251 for (k, entry) in data.range((start, end)) {
252 let user_key = decode_user_key(k);
253 if last_user_key.as_deref() == Some(user_key) {
254 continue;
255 }
256 if entry.timestamp > read_timestamp {
257 continue;
260 }
261 last_user_key = Some(user_key.to_vec());
262 out.push((user_key.to_vec(), entry.clone()));
263 }
264 out
265 }
266
267 pub fn scan_prefix(&self, prefix: &[u8], read_timestamp: u64) -> Vec<(Key, MemTableEntry)> {
269 let start = Included(internal_key_prefix(prefix));
272 let end = next_prefix(prefix).map(Excluded).unwrap_or(Unbounded);
273 self.collect_scan(start, end, read_timestamp)
274 }
275
276 pub fn scan_range(
278 &self,
279 start: &[u8],
280 end: &[u8],
281 read_timestamp: u64,
282 ) -> Vec<(Key, MemTableEntry)> {
283 self.collect_scan(
284 Included(start.to_vec()),
285 Excluded(end.to_vec()),
286 read_timestamp,
287 )
288 }
289
290 pub fn freeze(self) -> ImmutableMemTable {
292 let min_timestamp = self.min_timestamp();
293 let max_timestamp = self.max_timestamp();
294 let memory_usage = self.memory_usage.load(Ordering::Relaxed);
295 let data = self.data.into_inner().expect("memtable lock poisoned");
296 ImmutableMemTable {
297 data: Arc::new(data),
298 memory_usage,
299 min_timestamp,
300 max_timestamp,
301 }
302 }
303}
304
305#[derive(Debug, Clone)]
307pub struct ImmutableMemTable {
308 data: Arc<BTreeMap<Vec<u8>, MemTableEntry>>,
309 memory_usage: usize,
310 min_timestamp: Option<u64>,
311 max_timestamp: Option<u64>,
312}
313
314#[derive(Debug)]
321pub struct ImmutableMemTableCache {
322 max_immutable_count: usize,
323 next_id: u64,
324 entries: VecDeque<ImmutableMemTableCacheEntry>,
325}
326
327#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
329pub struct ImmutableMemTableId(u64);
330
331#[derive(Debug)]
332struct ImmutableMemTableCacheEntry {
333 id: ImmutableMemTableId,
334 table: Arc<ImmutableMemTable>,
335 flushing: bool,
336}
337
338#[derive(Debug)]
340pub struct ImmutableMemTablePushOutcome {
341 pub id: ImmutableMemTableId,
343 pub evicted: Vec<ImmutableMemTableEvicted>,
345}
346
347#[derive(Debug)]
349pub struct ImmutableMemTableEvicted {
350 pub id: ImmutableMemTableId,
352 pub table: Arc<ImmutableMemTable>,
354}
355
356impl ImmutableMemTableCache {
357 pub fn new(max_immutable_count: usize) -> Self {
359 Self {
360 max_immutable_count,
361 next_id: 1,
362 entries: VecDeque::new(),
363 }
364 }
365
366 pub fn len(&self) -> usize {
368 self.entries.len()
369 }
370
371 pub fn is_empty(&self) -> bool {
373 self.entries.is_empty()
374 }
375
376 pub fn max_immutable_count(&self) -> usize {
378 self.max_immutable_count
379 }
380
381 pub fn try_push(
386 &mut self,
387 table: Arc<ImmutableMemTable>,
388 ) -> Option<ImmutableMemTablePushOutcome> {
389 if self.max_immutable_count == 0 {
390 return None;
391 }
392
393 let mut evicted = Vec::new();
394 while self.entries.len() >= self.max_immutable_count {
395 let victim_index = self.select_eviction_candidate_index()?;
396 let victim = self
397 .entries
398 .remove(victim_index)
399 .expect("candidate index is in range");
400 evicted.push(ImmutableMemTableEvicted {
401 id: victim.id,
402 table: victim.table,
403 });
404 }
405
406 let id = ImmutableMemTableId(self.next_id);
407 self.next_id = self.next_id.wrapping_add(1).max(1);
408 self.entries.push_back(ImmutableMemTableCacheEntry {
409 id,
410 table,
411 flushing: false,
412 });
413 Some(ImmutableMemTablePushOutcome { id, evicted })
414 }
415
416 pub fn get(&self, id: ImmutableMemTableId) -> Option<Arc<ImmutableMemTable>> {
418 self.entries
419 .iter()
420 .find(|e| e.id == id)
421 .map(|e| Arc::clone(&e.table))
422 }
423
424 pub fn set_flushing(&mut self, id: ImmutableMemTableId, flushing: bool) -> bool {
426 let Some(entry) = self.entries.iter_mut().find(|e| e.id == id) else {
427 return false;
428 };
429 entry.flushing = flushing;
430 true
431 }
432
433 pub fn remove(&mut self, id: ImmutableMemTableId) -> Option<Arc<ImmutableMemTable>> {
435 let index = self.entries.iter().position(|e| e.id == id)?;
436 let entry = self.entries.remove(index)?;
437 Some(entry.table)
438 }
439
440 fn select_eviction_candidate_index(&self) -> Option<usize> {
441 let mut best: Option<(usize, u64, usize, u64)> = None;
447 for (idx, entry) in self.entries.iter().enumerate() {
448 if entry.flushing {
449 continue;
450 }
451 let ts = entry.table.max_timestamp().unwrap_or(0);
452 let mem = entry.table.memory_usage_bytes();
453 let id = entry.id.0;
454
455 let key = (idx, ts, mem, id);
456 best = match best {
457 None => Some(key),
458 Some((best_idx, best_ts, best_mem, best_id)) => {
459 let better = (ts < best_ts)
460 || (ts == best_ts && mem > best_mem)
461 || (ts == best_ts && mem == best_mem && id < best_id);
462 if better {
463 Some((idx, ts, mem, id))
464 } else {
465 Some((best_idx, best_ts, best_mem, best_id))
466 }
467 }
468 };
469 }
470 best.map(|(idx, _, _, _)| idx)
471 }
472}
473
474impl ImmutableMemTable {
475 pub fn memory_usage_bytes(&self) -> usize {
479 self.memory_usage
480 }
481
482 pub fn min_timestamp(&self) -> Option<u64> {
484 self.min_timestamp
485 }
486
487 pub fn max_timestamp(&self) -> Option<u64> {
489 self.max_timestamp
490 }
491
492 pub fn get(&self, key: &[u8], read_timestamp: u64) -> Option<MemTableEntry> {
494 let prefix = internal_key_prefix(key);
495 let start = internal_key(key, read_timestamp, u64::MAX);
496 let end = next_prefix(&prefix);
497
498 let range = match end {
499 Some(end_key) => self.data.range((Included(start), Excluded(end_key))),
500 None => self.data.range((Included(start), Unbounded)),
501 };
502 for (k, entry) in range {
503 if decode_user_key(k) != key {
504 break;
505 }
506 if entry.timestamp <= read_timestamp {
507 return Some(entry.clone());
508 }
509 }
510 None
511 }
512
513 fn collect_scan(
514 &self,
515 start: Bound<Vec<u8>>,
516 end: Bound<Vec<u8>>,
517 read_timestamp: u64,
518 ) -> Vec<(Key, MemTableEntry)> {
519 let mut out = Vec::new();
520 let mut last_user_key: Option<Vec<u8>> = None;
521 for (k, entry) in self.data.range((start, end)) {
522 let user_key = decode_user_key(k);
523 if last_user_key.as_deref() == Some(user_key) {
524 continue;
525 }
526 if entry.timestamp > read_timestamp {
527 continue;
528 }
529 last_user_key = Some(user_key.to_vec());
530 out.push((user_key.to_vec(), entry.clone()));
531 }
532 out
533 }
534
535 pub fn scan_prefix(&self, prefix: &[u8], read_timestamp: u64) -> Vec<(Key, MemTableEntry)> {
537 let start = Included(internal_key_prefix(prefix));
538 let end = next_prefix(prefix).map(Excluded).unwrap_or(Unbounded);
539 self.collect_scan(start, end, read_timestamp)
540 }
541
542 pub fn scan_range(
544 &self,
545 start: &[u8],
546 end: &[u8],
547 read_timestamp: u64,
548 ) -> Vec<(Key, MemTableEntry)> {
549 self.collect_scan(
550 Included(start.to_vec()),
551 Excluded(end.to_vec()),
552 read_timestamp,
553 )
554 }
555}
556
557#[cfg(all(test, not(target_arch = "wasm32")))]
558mod tests {
559 use super::*;
560
561 #[test]
562 fn get_obeys_read_timestamp_and_sequence() {
563 let mem = MemTable::new();
564 mem.put(b"k".to_vec(), b"v1".to_vec(), 10, 1);
565 mem.put(b"k".to_vec(), b"v2".to_vec(), 20, 1);
566 mem.put(b"k".to_vec(), b"v2b".to_vec(), 20, 2);
567
568 assert_eq!(mem.get(b"k", 9), None);
569 assert_eq!(mem.get(b"k", 10).unwrap().value.unwrap(), b"v1".to_vec());
570 assert_eq!(mem.get(b"k", 20).unwrap().value.unwrap(), b"v2b".to_vec());
571 assert_eq!(mem.get(b"k", 999).unwrap().value.unwrap(), b"v2b".to_vec());
572 }
573
574 #[test]
575 fn tombstone_is_visible() {
576 let mem = MemTable::new();
577 mem.put(b"k".to_vec(), b"v".to_vec(), 10, 1);
578 mem.delete(b"k".to_vec(), 20, 1);
579
580 let e = mem.get(b"k", 20).unwrap();
581 assert_eq!(e.value, None);
582 }
583
584 #[test]
585 fn scan_prefix_returns_latest_visible_per_key() {
586 let mem = MemTable::new();
587 mem.put(b"p:a".to_vec(), b"v1".to_vec(), 10, 1);
588 mem.put(b"p:a".to_vec(), b"v2".to_vec(), 20, 1);
589 mem.put(b"p:b".to_vec(), b"x".to_vec(), 15, 1);
590 mem.delete(b"p:c".to_vec(), 12, 1);
591 mem.put(b"q:z".to_vec(), b"no".to_vec(), 99, 1);
592
593 let got = mem.scan_prefix(b"p:", 20);
594 assert_eq!(got.len(), 3);
595 assert_eq!(got[0].0, b"p:a".to_vec());
596 assert_eq!(got[0].1.value.as_deref(), Some(b"v2".as_slice()));
597 assert_eq!(got[1].0, b"p:b".to_vec());
598 assert_eq!(got[2].0, b"p:c".to_vec());
599 assert!(got[2].1.value.is_none());
600 }
601
602 #[test]
603 fn scan_range_is_end_exclusive_and_obeys_read_timestamp() {
604 let mem = MemTable::new();
605 mem.put(b"a".to_vec(), b"1".to_vec(), 10, 1);
606 mem.put(b"b".to_vec(), b"2_old".to_vec(), 10, 1);
607 mem.put(b"b".to_vec(), b"2_new".to_vec(), 20, 1);
608 mem.delete(b"c".to_vec(), 12, 1);
609 mem.put(b"d".to_vec(), b"4".to_vec(), 40, 1);
610
611 let got = mem.scan_range(b"b", b"d", 15);
613 assert_eq!(got.len(), 2);
614 assert_eq!(got[0].0, b"b".to_vec());
615 assert_eq!(got[0].1.value.as_deref(), Some(b"2_old".as_slice()));
617 assert_eq!(got[1].0, b"c".to_vec());
618 assert!(got[1].1.value.is_none());
620 }
621
622 #[test]
623 fn freeze_produces_read_only_snapshot() {
624 let mem = MemTable::new();
625 mem.put(b"k".to_vec(), b"v".to_vec(), 10, 1);
626 let imm = mem.freeze();
627 assert_eq!(imm.get(b"k", 10).unwrap().value.unwrap(), b"v".to_vec());
628 }
629}
630
631#[cfg(all(test, not(target_arch = "wasm32")))]
632mod cache {
633 use super::*;
634
635 fn frozen_with_ts(ts: u64, mem: usize) -> Arc<ImmutableMemTable> {
636 let memtable = MemTable::new();
637 memtable.put(b"k".to_vec(), vec![0u8; mem], ts, 1);
638 Arc::new(memtable.freeze())
639 }
640
641 #[test]
642 fn evicts_oldest_non_flushing_when_full() {
643 let mut cache = ImmutableMemTableCache::new(2);
644
645 let a = cache.try_push(frozen_with_ts(10, 10)).unwrap().id;
646 let _b = cache.try_push(frozen_with_ts(20, 10)).unwrap().id;
647 let outcome = cache.try_push(frozen_with_ts(30, 10)).unwrap();
648
649 assert_eq!(cache.len(), 2);
650 assert_eq!(outcome.evicted.len(), 1);
651 assert_eq!(outcome.evicted[0].id, a);
652 assert!(cache.get(a).is_none());
653 }
654
655 #[test]
656 fn does_not_evict_flushing_entries() {
657 let mut cache = ImmutableMemTableCache::new(2);
658
659 let a = cache.try_push(frozen_with_ts(10, 10)).unwrap().id;
660 let b = cache.try_push(frozen_with_ts(20, 10)).unwrap().id;
661 assert!(cache.set_flushing(a, true));
662 assert!(cache.set_flushing(b, true));
663
664 assert!(cache.try_push(frozen_with_ts(30, 10)).is_none());
666 assert_eq!(cache.len(), 2);
667 assert!(cache.get(a).is_some());
668 assert!(cache.get(b).is_some());
669 }
670
671 #[test]
672 fn eviction_prefers_older_then_larger_memory() {
673 let mut cache = ImmutableMemTableCache::new(2);
674
675 let a = cache.try_push(frozen_with_ts(10, 5)).unwrap().id;
677 let b = cache.try_push(frozen_with_ts(10, 50)).unwrap().id;
678 let outcome = cache.try_push(frozen_with_ts(11, 5)).unwrap();
679
680 assert_eq!(outcome.evicted.len(), 1);
681 assert_eq!(outcome.evicted[0].id, b);
682 assert!(cache.get(a).is_some());
683 assert!(cache.get(b).is_none());
684 }
685}