use std::collections::BTreeMap;
use std::ops::Bound;
use std::sync::{RwLock, RwLockReadGuard, RwLockWriteGuard};
use crate::engine::{Durability, KvEngine, ReadTx, Readable, WriteTx};
use crate::error::{StoreError, StoreResult};
type Table = BTreeMap<Vec<u8>, Vec<u8>>;
type Tables = BTreeMap<String, Table>;
#[derive(Default)]
pub struct MemEngine {
tables: RwLock<Tables>,
}
impl MemEngine {
pub fn new() -> Self {
Self::default()
}
}
fn poisoned() -> StoreError {
StoreError::engine("in-memory lock poisoned")
}
fn range_table(table: &Table, lo: &[u8], hi: &[u8]) -> Vec<(Vec<u8>, Vec<u8>)> {
table
.range::<[u8], _>((Bound::Included(lo), Bound::Excluded(hi)))
.map(|(k, v)| (k.clone(), v.clone()))
.collect()
}
pub struct MemReadTx<'a> {
guard: RwLockReadGuard<'a, Tables>,
}
impl Readable for MemReadTx<'_> {
fn get_raw(&self, table: &str, key: &[u8]) -> StoreResult<Option<Vec<u8>>> {
Ok(self.guard.get(table).and_then(|t| t.get(key).cloned()))
}
fn range_raw(&self, table: &str, lo: &[u8], hi: &[u8]) -> StoreResult<Vec<(Vec<u8>, Vec<u8>)>> {
Ok(self
.guard
.get(table)
.map(|t| range_table(t, lo, hi))
.unwrap_or_default())
}
}
impl ReadTx for MemReadTx<'_> {}
type Overlay = BTreeMap<String, BTreeMap<Vec<u8>, Option<Vec<u8>>>>;
pub struct MemWriteTx<'a> {
base: RwLockWriteGuard<'a, Tables>,
overlay: Overlay,
_durability: Durability,
}
impl Readable for MemWriteTx<'_> {
fn get_raw(&self, table: &str, key: &[u8]) -> StoreResult<Option<Vec<u8>>> {
if let Some(t) = self.overlay.get(table) {
if let Some(slot) = t.get(key) {
return Ok(slot.clone());
}
}
Ok(self.base.get(table).and_then(|t| t.get(key).cloned()))
}
fn range_raw(&self, table: &str, lo: &[u8], hi: &[u8]) -> StoreResult<Vec<(Vec<u8>, Vec<u8>)>> {
let mut merged: BTreeMap<Vec<u8>, Vec<u8>> = self
.base
.get(table)
.map(|t| range_table(t, lo, hi).into_iter().collect())
.unwrap_or_default();
if let Some(t) = self.overlay.get(table) {
for (k, slot) in t.range::<[u8], _>((Bound::Included(lo), Bound::Excluded(hi))) {
match slot {
Some(v) => {
merged.insert(k.clone(), v.clone());
}
None => {
merged.remove(k);
}
}
}
}
Ok(merged.into_iter().collect())
}
}
impl WriteTx for MemWriteTx<'_> {
fn put_raw(&mut self, table: &str, key: &[u8], val: &[u8]) -> StoreResult<()> {
self.overlay
.entry(table.to_string())
.or_default()
.insert(key.to_vec(), Some(val.to_vec()));
Ok(())
}
fn delete_raw(&mut self, table: &str, key: &[u8]) -> StoreResult<bool> {
let existed = self.get_raw(table, key)?.is_some();
self.overlay
.entry(table.to_string())
.or_default()
.insert(key.to_vec(), None);
Ok(existed)
}
fn commit(mut self) -> StoreResult<()> {
let overlay = std::mem::take(&mut self.overlay);
for (table, edits) in overlay {
let t = self.base.entry(table).or_default();
for (key, slot) in edits {
match slot {
Some(v) => {
t.insert(key, v);
}
None => {
t.remove(&key);
}
}
}
}
Ok(())
}
}
impl KvEngine for MemEngine {
type RTx<'a> = MemReadTx<'a>;
type WTx<'a> = MemWriteTx<'a>;
fn begin_read(&self) -> StoreResult<Self::RTx<'_>> {
let guard = self.tables.read().map_err(|_| poisoned())?;
Ok(MemReadTx { guard })
}
fn begin_write(&self, durability: Durability) -> StoreResult<Self::WTx<'_>> {
let base = self.tables.write().map_err(|_| poisoned())?;
Ok(MemWriteTx {
base,
overlay: Overlay::new(),
_durability: durability,
})
}
}