use crate::{Error, Result, Record, Key, Item, SeqNo, Value, wal::Wal, sst::{SstWriter, SstReader}};
use crate::iterator::{QueryParams, QueryResult, ScanParams, ScanResult};
use crate::expression::{UpdateAction, UpdateExecutor, ExpressionContext, Expr, ExpressionEvaluator};
use crate::index::{TableSchema, encode_index_key, decode_index_key};
use crate::compaction::{CompactionManager, CompactionConfig, CompactionStatsAtomic};
use crate::config::DatabaseConfig;
use bytes::Bytes;
use parking_lot::RwLock;
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::fs;
const MEMTABLE_THRESHOLD: usize = 1000; const NUM_STRIPES: usize = 256;
pub struct LsmEngine {
inner: Arc<RwLock<LsmInner>>,
}
struct Stripe {
memtable: BTreeMap<Vec<u8>, Record>, memtable_size_bytes: usize, ssts: Vec<SstReader>, }
impl Stripe {
fn new() -> Self {
Self {
memtable: BTreeMap::new(),
memtable_size_bytes: 0,
ssts: Vec::new(),
}
}
fn estimate_record_size(key_enc: &[u8], record: &Record) -> usize {
let mut size = key_enc.len(); size += std::mem::size_of::<SeqNo>();
if let Some(item) = &record.value {
for (attr_name, value) in item {
size += attr_name.len();
size += match value {
Value::S(s) => s.len(),
Value::N(n) => n.len(),
Value::B(b) => b.len(),
Value::Bool(_) => 1,
Value::Null => 0,
Value::Ts(_) => 8,
Value::L(list) => {
list.len() * 32 }
Value::M(map) => {
map.len() * 64 }
Value::VecF32(vec) => {
vec.len() * 4
}
};
}
}
size
}
}
struct LsmInner {
dir: PathBuf,
wal: Wal,
stripes: Vec<Stripe>, next_seq: SeqNo, next_sst_id: u64, schema: TableSchema, stream_buffer: std::collections::VecDeque<crate::stream::StreamRecord>, compaction_config: CompactionConfig, compaction_stats: CompactionStatsAtomic, config: DatabaseConfig, }
#[derive(Debug, Clone)]
pub enum TransactWriteOperation {
Put {
item: Item,
condition: Option<Expr>,
},
Update {
actions: Vec<UpdateAction>,
condition: Option<Expr>,
},
Delete {
condition: Option<Expr>,
},
ConditionCheck {
condition: Expr,
},
}
impl TransactWriteOperation {
pub fn condition(&self) -> Option<&Expr> {
match self {
Self::Put { condition, .. } => condition.as_ref(),
Self::Update { condition, .. } => condition.as_ref(),
Self::Delete { condition } => condition.as_ref(),
Self::ConditionCheck { condition } => Some(condition),
}
}
}
impl LsmInner {
fn should_flush_stripe(&self, stripe_id: usize) -> bool {
let stripe = &self.stripes[stripe_id];
if stripe.memtable.len() >= self.config.max_memtable_records {
return true;
}
if let Some(max_bytes) = self.config.max_memtable_size_bytes {
if stripe.memtable_size_bytes >= max_bytes {
return true;
}
}
false
}
fn insert_into_memtable(&mut self, stripe_id: usize, key_enc: Vec<u8>, record: Record) {
let record_size = Stripe::estimate_record_size(&key_enc, &record);
if let Some(old_record) = self.stripes[stripe_id].memtable.get(&key_enc) {
let old_size = Stripe::estimate_record_size(&key_enc, old_record);
self.stripes[stripe_id].memtable_size_bytes =
self.stripes[stripe_id].memtable_size_bytes.saturating_sub(old_size);
}
self.stripes[stripe_id].memtable.insert(key_enc, record);
self.stripes[stripe_id].memtable_size_bytes += record_size;
}
}
impl LsmEngine {
pub fn create(dir: impl AsRef<Path>) -> Result<Self> {
Self::create_with_schema(dir, TableSchema::new())
}
pub fn create_with_schema(dir: impl AsRef<Path>, schema: TableSchema) -> Result<Self> {
Self::create_with_config(dir, DatabaseConfig::default(), schema)
}
pub fn create_with_config(
dir: impl AsRef<Path>,
config: DatabaseConfig,
schema: TableSchema,
) -> Result<Self> {
config.validate().map_err(|e| Error::InvalidArgument(e))?;
let dir = dir.as_ref();
fs::create_dir_all(dir)?;
let wal_path = dir.join("wal.log");
if wal_path.exists() {
return Err(Error::AlreadyExists(dir.display().to_string()));
}
let wal = Wal::create(&wal_path)?;
let stripes = (0..NUM_STRIPES).map(|_| Stripe::new()).collect();
Ok(Self {
inner: Arc::new(RwLock::new(LsmInner {
dir: dir.to_path_buf(),
wal,
stripes,
next_seq: 1,
next_sst_id: 1,
schema,
stream_buffer: std::collections::VecDeque::new(),
compaction_config: CompactionConfig::default(),
compaction_stats: CompactionStatsAtomic::new(),
config,
})),
})
}
pub fn open(dir: impl AsRef<Path>) -> Result<Self> {
let dir = dir.as_ref();
let wal_path = dir.join("wal.log");
let wal = Wal::open(&wal_path)?;
let mut stripes: Vec<Stripe> = (0..NUM_STRIPES).map(|_| Stripe::new()).collect();
let mut max_sst_id = 0u64;
for entry in fs::read_dir(dir)? {
let entry = entry?;
let path = entry.path();
if let Some(ext) = path.extension() {
if ext == "sst" {
if let Some(stem) = path.file_stem() {
if let Some(name) = stem.to_str() {
if let Some((stripe_str, id_str)) = name.split_once('-') {
if let (Ok(stripe), Ok(id)) = (stripe_str.parse::<usize>(), id_str.parse::<u64>()) {
if stripe < NUM_STRIPES {
max_sst_id = max_sst_id.max(id);
let reader = SstReader::open(&path)?;
stripes[stripe].ssts.push(reader);
}
}
} else {
if let Ok(id) = name.parse::<u64>() {
max_sst_id = max_sst_id.max(id);
let reader = SstReader::open(&path)?;
stripes[0].ssts.push(reader);
}
}
}
}
}
}
}
for stripe in &mut stripes {
stripe.ssts.reverse();
}
let records = wal.read_all()?;
let mut max_seq = 0;
for (_lsn, record) in records {
max_seq = max_seq.max(record.seq);
let key_enc = record.key.encode().to_vec();
let stripe_id = record.key.stripe() as usize;
stripes[stripe_id].memtable.insert(key_enc, record);
}
Ok(Self {
inner: Arc::new(RwLock::new(LsmInner {
dir: dir.to_path_buf(),
wal,
stripes,
next_seq: max_seq + 1,
next_sst_id: max_sst_id + 1,
schema: TableSchema::new(), stream_buffer: std::collections::VecDeque::new(),
compaction_config: CompactionConfig::default(),
compaction_stats: CompactionStatsAtomic::new(),
config: DatabaseConfig::default(), })),
})
}
pub fn put(&self, key: Key, item: Item) -> Result<()> {
let mut inner = self.inner.write();
let old_image = if inner.schema.stream_config.enabled {
let stripe_id = key.stripe() as usize;
let key_enc = key.encode().to_vec();
inner.stripes[stripe_id].memtable.get(&key_enc).and_then(|r| r.value.clone())
} else {
None
};
let seq = inner.next_seq;
inner.next_seq += 1;
let record = Record::put(key.clone(), item.clone(), seq);
inner.wal.append(record.clone())?;
inner.wal.flush()?;
let stripe_id = record.key.stripe() as usize;
let key_enc = record.key.encode().to_vec();
inner.insert_into_memtable(stripe_id, key_enc, record);
if !inner.schema.local_indexes.is_empty() {
self.materialize_lsi_entries(&mut inner, &key, &item)?;
}
if !inner.schema.global_indexes.is_empty() {
self.materialize_gsi_entries(&mut inner, &key, &item)?;
}
if inner.schema.stream_config.enabled {
let stream_record = if let Some(old) = old_image {
crate::stream::StreamRecord::modify(
seq,
key.clone(),
old,
item.clone(),
inner.schema.stream_config.view_type,
)
} else {
crate::stream::StreamRecord::insert(
seq,
key.clone(),
item.clone(),
inner.schema.stream_config.view_type,
)
};
self.emit_stream_record(&mut inner, stream_record);
}
if inner.should_flush_stripe(stripe_id) {
self.flush_stripe(&mut inner, stripe_id)?;
}
Ok(())
}
pub fn put_conditional(&self, key: Key, item: Item, condition: &Expr, context: &ExpressionContext) -> Result<()> {
let current_item = self.get(&key)?.unwrap_or_else(|| std::collections::HashMap::new());
let evaluator = ExpressionEvaluator::new(¤t_item, context);
let condition_passed = evaluator.evaluate(condition)?;
if !condition_passed {
return Err(Error::ConditionalCheckFailed("Put condition failed".into()));
}
self.put(key, item)
}
pub fn get(&self, key: &Key) -> Result<Option<Item>> {
let inner = self.inner.read();
let stripe_id = key.stripe() as usize;
let stripe = &inner.stripes[stripe_id];
let key_enc = key.encode().to_vec();
if let Some(record) = stripe.memtable.get(&key_enc) {
if let Some(item) = &record.value {
if inner.schema.is_expired(item) {
drop(inner); self.delete(key.clone())?;
return Ok(None);
}
}
return Ok(record.value.clone());
}
for sst in &stripe.ssts {
if let Some(record) = sst.get(key) {
if let Some(item) = &record.value {
if inner.schema.is_expired(item) {
drop(inner); self.delete(key.clone())?;
return Ok(None);
}
}
return Ok(record.value.clone());
}
}
Ok(None)
}
pub fn delete(&self, key: Key) -> Result<()> {
let mut inner = self.inner.write();
let old_image = if inner.schema.stream_config.enabled {
let stripe_id = key.stripe() as usize;
let key_enc = key.encode().to_vec();
inner.stripes[stripe_id].memtable.get(&key_enc).and_then(|r| r.value.clone())
} else {
None
};
let seq = inner.next_seq;
inner.next_seq += 1;
let record = Record::delete(key.clone(), seq);
inner.wal.append(record.clone())?;
inner.wal.flush()?;
let stripe_id = record.key.stripe() as usize;
let key_enc = record.key.encode().to_vec();
inner.stripes[stripe_id].memtable.insert(key_enc, record);
if inner.schema.stream_config.enabled {
if let Some(old) = old_image {
let stream_record = crate::stream::StreamRecord::remove(
seq,
key.clone(),
old,
inner.schema.stream_config.view_type,
);
self.emit_stream_record(&mut inner, stream_record);
}
}
if inner.should_flush_stripe(stripe_id) {
self.flush_stripe(&mut inner, stripe_id)?;
}
Ok(())
}
pub fn delete_conditional(&self, key: Key, condition: &Expr, context: &ExpressionContext) -> Result<()> {
let current_item = self.get(&key)?.unwrap_or_else(|| std::collections::HashMap::new());
let evaluator = ExpressionEvaluator::new(¤t_item, context);
let condition_passed = evaluator.evaluate(condition)?;
if !condition_passed {
return Err(Error::ConditionalCheckFailed("Delete condition failed".into()));
}
self.delete(key)
}
pub fn update(&self, key: &Key, actions: &[UpdateAction], context: &ExpressionContext) -> Result<Item> {
let current_item = self.get(key)?.unwrap_or_else(|| std::collections::HashMap::new());
let executor = UpdateExecutor::new(context);
let updated_item = executor.execute(¤t_item, actions)?;
self.put(key.clone(), updated_item.clone())?;
Ok(updated_item)
}
pub fn update_conditional(
&self,
key: &Key,
actions: &[UpdateAction],
condition: &Expr,
context: &ExpressionContext,
) -> Result<Item> {
let current_item = self.get(key)?.unwrap_or_else(|| std::collections::HashMap::new());
let evaluator = ExpressionEvaluator::new(¤t_item, context);
let condition_passed = evaluator.evaluate(condition)?;
if !condition_passed {
return Err(Error::ConditionalCheckFailed("Update condition failed".into()));
}
let executor = UpdateExecutor::new(context);
let updated_item = executor.execute(¤t_item, actions)?;
self.put(key.clone(), updated_item.clone())?;
Ok(updated_item)
}
pub fn query(&self, params: QueryParams) -> Result<QueryResult> {
let inner = self.inner.read();
let stripe_id = {
let temp_key = Key::new(params.pk.clone());
temp_key.stripe() as usize
};
let stripe = &inner.stripes[stripe_id];
let mut items = Vec::new();
let mut seen_keys: std::collections::HashSet<Vec<u8>> = std::collections::HashSet::new();
let mut scanned_count = 0;
let mut last_key = None;
let mut all_records: BTreeMap<Vec<u8>, Record> = BTreeMap::new();
let is_index_query = params.index_name.is_some();
for (key_enc, record) in &stripe.memtable {
if is_index_query {
if let Some(index_name) = ¶ms.index_name {
if let Some((idx_name, idx_pk, idx_sk)) = decode_index_key(key_enc) {
if idx_name != *index_name {
continue;
}
if idx_pk != params.pk {
continue;
}
if !params.matches_sk(&Some(idx_sk)) {
continue;
}
all_records.insert(key_enc.clone(), record.clone());
}
}
} else {
if record.key.pk != params.pk {
continue;
}
if !params.matches_sk(&record.key.sk) {
continue;
}
all_records.insert(key_enc.clone(), record.clone());
}
}
for _sst in &stripe.ssts {
}
let mut sorted_records: Vec<(Vec<u8>, Record)> = all_records.into_iter().collect();
if !params.forward {
sorted_records.reverse();
}
for (key_enc, record) in sorted_records {
if params.should_skip(&record.key) {
continue;
}
scanned_count += 1;
if seen_keys.contains(&key_enc) {
continue;
}
seen_keys.insert(key_enc);
if record.value.is_none() {
continue;
}
if let Some(ref item) = record.value {
if inner.schema.is_expired(item) {
continue; }
}
last_key = Some(record.key.clone());
if let Some(item) = record.value {
items.push(item);
if let Some(limit) = params.limit {
if items.len() >= limit {
break;
}
}
}
}
Ok(QueryResult::new(items, last_key, scanned_count))
}
pub fn batch_get(&self, keys: &[Key]) -> Result<std::collections::HashMap<Key, Option<Item>>> {
let mut results = std::collections::HashMap::new();
for key in keys {
let item = self.get(key)?;
results.insert(key.clone(), item);
}
Ok(results)
}
pub fn batch_write(&self, operations: &[(Key, Option<Item>)]) -> Result<usize> {
let mut processed = 0;
for (key, item_opt) in operations {
match item_opt {
Some(item) => {
self.put(key.clone(), item.clone())?;
processed += 1;
}
None => {
self.delete(key.clone())?;
processed += 1;
}
}
}
Ok(processed)
}
pub fn transact_get(&self, keys: &[Key]) -> Result<Vec<Option<Item>>> {
let _inner = self.inner.read();
let mut items = Vec::new();
for key in keys {
let item = self.get(key)?;
items.push(item);
}
Ok(items)
}
pub fn transact_write(
&self,
operations: &[(Key, TransactWriteOperation)],
context: &ExpressionContext,
) -> Result<usize> {
let mut inner = self.inner.write();
let mut current_items: Vec<Option<Item>> = Vec::new();
for (key, op) in operations {
let item = {
let stripe_id = key.stripe() as usize;
let stripe = &inner.stripes[stripe_id];
let key_enc = key.encode().to_vec();
if let Some(record) = stripe.memtable.get(&key_enc) {
record.value.clone()
} else {
let mut found = None;
for sst in &stripe.ssts {
if let Some(record) = sst.get(key) {
found = record.value.clone();
break;
}
}
found
}
};
current_items.push(item.clone());
if let Some(condition_expr) = op.condition() {
let current_item = item.unwrap_or_else(|| std::collections::HashMap::new());
let evaluator = ExpressionEvaluator::new(¤t_item, context);
let condition_passed = evaluator.evaluate(condition_expr)?;
if !condition_passed {
return Err(Error::TransactionCanceled(format!(
"Condition failed for key {:?}",
key
)));
}
}
}
let mut committed = 0;
for (i, (key, op)) in operations.iter().enumerate() {
match op {
TransactWriteOperation::Put { item, .. } => {
let seq = inner.next_seq;
inner.next_seq += 1;
let record = Record::put(key.clone(), item.clone(), seq);
inner.wal.append(record.clone())?;
inner.wal.flush()?;
let stripe_id = record.key.stripe() as usize;
let key_enc = record.key.encode().to_vec();
inner.stripes[stripe_id].memtable.insert(key_enc, record);
if inner.stripes[stripe_id].memtable.len() >= MEMTABLE_THRESHOLD {
self.flush_stripe(&mut inner, stripe_id)?;
}
committed += 1;
}
TransactWriteOperation::Delete { .. } => {
let seq = inner.next_seq;
inner.next_seq += 1;
let record = Record::delete(key.clone(), seq);
inner.wal.append(record.clone())?;
inner.wal.flush()?;
let stripe_id = record.key.stripe() as usize;
let key_enc = record.key.encode().to_vec();
inner.stripes[stripe_id].memtable.insert(key_enc, record);
if inner.stripes[stripe_id].memtable.len() >= MEMTABLE_THRESHOLD {
self.flush_stripe(&mut inner, stripe_id)?;
}
committed += 1;
}
TransactWriteOperation::Update { actions, .. } => {
let current_item = current_items[i].clone().unwrap_or_else(|| std::collections::HashMap::new());
let executor = UpdateExecutor::new(context);
let updated_item = executor.execute(¤t_item, actions)?;
let seq = inner.next_seq;
inner.next_seq += 1;
let record = Record::put(key.clone(), updated_item, seq);
inner.wal.append(record.clone())?;
inner.wal.flush()?;
let stripe_id = record.key.stripe() as usize;
let key_enc = record.key.encode().to_vec();
inner.stripes[stripe_id].memtable.insert(key_enc, record);
if inner.stripes[stripe_id].memtable.len() >= MEMTABLE_THRESHOLD {
self.flush_stripe(&mut inner, stripe_id)?;
}
committed += 1;
}
TransactWriteOperation::ConditionCheck { .. } => {
committed += 1;
}
}
}
Ok(committed)
}
pub fn scan(&self, params: ScanParams) -> Result<ScanResult> {
let inner = self.inner.read();
let mut all_records: BTreeMap<Vec<u8>, Record> = BTreeMap::new();
for stripe_id in 0..NUM_STRIPES {
if !params.should_scan_stripe(stripe_id) {
continue;
}
let stripe = &inner.stripes[stripe_id];
for (key_enc, record) in &stripe.memtable {
if record.value.is_none() {
continue;
}
all_records.insert(key_enc.clone(), record.clone());
}
}
let mut items = Vec::new();
let mut scanned_count = 0;
let mut last_key = None;
for (_, record) in all_records {
if params.should_skip(&record.key) {
continue;
}
scanned_count += 1;
if let Some(ref item) = record.value {
if inner.schema.is_expired(item) {
continue; }
}
last_key = Some(record.key.clone());
if let Some(item) = record.value {
items.push(item);
if let Some(limit) = params.limit {
if items.len() >= limit {
return Ok(ScanResult::new(items, last_key, scanned_count));
}
}
}
}
Ok(ScanResult::new(items, last_key, scanned_count))
}
fn materialize_lsi_entries(&self, inner: &mut LsmInner, key: &Key, item: &Item) -> Result<()> {
for lsi in &inner.schema.local_indexes {
if let Some(index_sk_value) = item.get(&lsi.sort_key_attribute) {
let index_sk_bytes = match index_sk_value {
Value::S(s) => Bytes::copy_from_slice(s.as_bytes()),
Value::N(n) => Bytes::copy_from_slice(n.as_bytes()),
Value::B(b) => b.clone(),
Value::Bool(b) => Bytes::copy_from_slice(if *b { b"true" } else { b"false" }),
Value::Ts(ts) => Bytes::copy_from_slice(&ts.to_le_bytes()),
_ => continue, };
let index_key_encoded = encode_index_key(&lsi.name, &key.pk, &index_sk_bytes);
let index_item = item.clone();
let index_key = Key::new(Bytes::copy_from_slice(&index_key_encoded));
let seq = inner.next_seq;
inner.next_seq += 1;
let index_record = Record::put(index_key, index_item, seq);
inner.wal.append(index_record.clone())?;
let stripe_id = key.stripe() as usize;
inner.stripes[stripe_id].memtable.insert(index_key_encoded, index_record);
}
}
Ok(())
}
fn materialize_gsi_entries(&self, inner: &mut LsmInner, base_key: &Key, item: &Item) -> Result<()> {
for gsi in &inner.schema.global_indexes {
if let Some(gsi_pk_value) = item.get(&gsi.partition_key_attribute) {
let gsi_pk_bytes = match gsi_pk_value {
Value::S(s) => Bytes::copy_from_slice(s.as_bytes()),
Value::N(n) => Bytes::copy_from_slice(n.as_bytes()),
Value::B(b) => b.clone(),
Value::Bool(b) => Bytes::copy_from_slice(if *b { b"true" } else { b"false" }),
Value::Ts(ts) => Bytes::copy_from_slice(&ts.to_le_bytes()),
_ => continue, };
let mut gsi_sk_bytes = if let Some(gsi_sk_attr) = &gsi.sort_key_attribute {
if let Some(gsi_sk_value) = item.get(gsi_sk_attr) {
match gsi_sk_value {
Value::S(s) => Bytes::copy_from_slice(s.as_bytes()),
Value::N(n) => Bytes::copy_from_slice(n.as_bytes()),
Value::B(b) => b.clone(),
Value::Bool(b) => Bytes::copy_from_slice(if *b { b"true" } else { b"false" }),
Value::Ts(ts) => Bytes::copy_from_slice(&ts.to_le_bytes()),
_ => Bytes::new(), }
} else {
continue; }
} else {
Bytes::new() };
let base_pk_encoded = base_key.encode();
let mut combined_sk = Vec::with_capacity(gsi_sk_bytes.len() + base_pk_encoded.len());
combined_sk.extend_from_slice(&gsi_sk_bytes);
combined_sk.extend_from_slice(&base_pk_encoded);
gsi_sk_bytes = Bytes::from(combined_sk);
let index_key_encoded = encode_index_key(&gsi.name, &gsi_pk_bytes, &gsi_sk_bytes);
let index_item = item.clone();
let index_key = Key::new(Bytes::copy_from_slice(&index_key_encoded));
let seq = inner.next_seq;
inner.next_seq += 1;
let index_record = Record::put(index_key, index_item, seq);
inner.wal.append(index_record.clone())?;
let gsi_stripe_key = Key::new(gsi_pk_bytes.clone());
let gsi_stripe_id = gsi_stripe_key.stripe() as usize;
inner.stripes[gsi_stripe_id].memtable.insert(index_key_encoded, index_record);
}
}
Ok(())
}
pub fn read_stream(&self, after_sequence_number: Option<u64>) -> Result<Vec<crate::stream::StreamRecord>> {
let inner = self.inner.read();
if !inner.schema.stream_config.enabled {
return Ok(Vec::new());
}
let records: Vec<crate::stream::StreamRecord> = inner.stream_buffer
.iter()
.filter(|record| {
if let Some(after) = after_sequence_number {
record.sequence_number > after
} else {
true
}
})
.cloned()
.collect();
Ok(records)
}
fn emit_stream_record(&self, inner: &mut LsmInner, record: crate::stream::StreamRecord) {
if !inner.schema.stream_config.enabled {
return;
}
inner.stream_buffer.push_back(record);
while inner.stream_buffer.len() > inner.schema.stream_config.buffer_size {
inner.stream_buffer.pop_front();
}
}
fn flush_stripe(&self, inner: &mut LsmInner, stripe_id: usize) -> Result<()> {
if inner.stripes[stripe_id].memtable.is_empty() {
return Ok(());
}
let sst_id = inner.next_sst_id;
inner.next_sst_id += 1;
let sst_path = inner.dir.join(format!("{:03}-{}.sst", stripe_id, sst_id));
let mut writer = SstWriter::new();
for record in inner.stripes[stripe_id].memtable.values() {
writer.add(record.clone());
}
writer.finish(&sst_path)?;
let reader = SstReader::open(&sst_path)?;
inner.stripes[stripe_id].ssts.insert(0, reader);
inner.stripes[stripe_id].memtable.clear();
inner.stripes[stripe_id].memtable_size_bytes = 0;
if inner.compaction_config.enabled && inner.stripes[stripe_id].ssts.len() >= inner.compaction_config.sst_threshold {
let _guard = inner.compaction_stats.start_compaction();
let compaction_mgr = CompactionManager::new(stripe_id, inner.dir.clone());
let ssts_to_compact = &inner.stripes[stripe_id].ssts;
let sst_count = ssts_to_compact.len();
let compacted_sst_id = inner.next_sst_id;
inner.next_sst_id += 1;
let (new_sst, old_paths) = compaction_mgr.compact(ssts_to_compact, compacted_sst_id)?;
inner.compaction_stats.record_ssts_merged(sst_count as u64);
inner.compaction_stats.record_ssts_created(1);
inner.stripes[stripe_id].ssts.clear();
inner.stripes[stripe_id].ssts.push(new_sst);
compaction_mgr.cleanup_old_ssts(old_paths)?;
}
Ok(())
}
pub fn flush(&self) -> Result<()> {
let mut inner = self.inner.write();
for stripe_id in 0..NUM_STRIPES {
if !inner.stripes[stripe_id].memtable.is_empty() {
self.flush_stripe(&mut inner, stripe_id)?;
}
}
Ok(())
}
pub fn set_compaction_config(&self, config: CompactionConfig) {
let mut inner = self.inner.write();
inner.compaction_config = config;
}
pub fn compaction_config(&self) -> CompactionConfig {
let inner = self.inner.read();
inner.compaction_config.clone()
}
pub fn compaction_stats(&self) -> crate::compaction::CompactionStats {
let inner = self.inner.read();
inner.compaction_stats.snapshot()
}
pub fn trigger_compaction(&self, stripe_id: usize) -> Result<()> {
if stripe_id >= NUM_STRIPES {
return Err(Error::InvalidArgument(format!(
"Invalid stripe_id: {}, must be < {}",
stripe_id, NUM_STRIPES
)));
}
let mut inner = self.inner.write();
if inner.stripes[stripe_id].ssts.len() >= inner.compaction_config.sst_threshold {
let _guard = inner.compaction_stats.start_compaction();
let compaction_mgr = CompactionManager::new(stripe_id, inner.dir.clone());
let sst_count = inner.stripes[stripe_id].ssts.len();
let compacted_sst_id = inner.next_sst_id;
inner.next_sst_id += 1;
let ssts_to_compact = &inner.stripes[stripe_id].ssts;
let (new_sst, old_paths) = compaction_mgr.compact(ssts_to_compact, compacted_sst_id)?;
inner.compaction_stats.record_ssts_merged(sst_count as u64);
inner.compaction_stats.record_ssts_created(1);
inner.stripes[stripe_id].ssts.clear();
inner.stripes[stripe_id].ssts.push(new_sst);
compaction_mgr.cleanup_old_ssts(old_paths)?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Value;
use tempfile::TempDir;
use std::collections::HashMap;
#[test]
fn test_lsm_create() {
let dir = TempDir::new().unwrap();
let _db = LsmEngine::create(dir.path()).unwrap();
}
#[test]
fn test_lsm_put_get() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let key = Key::new(b"user#123".to_vec());
let mut item = HashMap::new();
item.insert("name".to_string(), Value::string("Alice"));
item.insert("age".to_string(), Value::number(30));
db.put(key.clone(), item.clone()).unwrap();
let result = db.get(&key).unwrap();
assert_eq!(result, Some(item));
}
#[test]
fn test_lsm_delete() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let key = Key::new(b"user#123".to_vec());
let mut item = HashMap::new();
item.insert("name".to_string(), Value::string("Bob"));
db.put(key.clone(), item).unwrap();
assert!(db.get(&key).unwrap().is_some());
db.delete(key.clone()).unwrap();
assert!(db.get(&key).unwrap().is_none());
}
#[test]
fn test_lsm_reopen() {
let dir = TempDir::new().unwrap();
let path = dir.path().to_path_buf();
let key = Key::new(b"persistent".to_vec());
let mut item = HashMap::new();
item.insert("data".to_string(), Value::string("test"));
{
let db = LsmEngine::create(&path).unwrap();
db.put(key.clone(), item.clone()).unwrap();
}
let db = LsmEngine::open(&path).unwrap();
let result = db.get(&key).unwrap();
assert_eq!(result, Some(item));
}
#[test]
fn test_lsm_flush() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
for i in 0..MEMTABLE_THRESHOLD + 10 {
let key = Key::new(format!("key{}", i).into_bytes());
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
for i in 0..MEMTABLE_THRESHOLD + 10 {
let key = Key::new(format!("key{}", i).into_bytes());
let result = db.get(&key).unwrap();
assert!(result.is_some());
}
}
#[test]
fn test_lsm_overwrite() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let key = Key::new(b"counter".to_vec());
for i in 0..5 {
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i));
db.put(key.clone(), item).unwrap();
}
let result = db.get(&key).unwrap().unwrap();
match result.get("value").unwrap() {
Value::N(n) => assert_eq!(n, "4"),
_ => panic!("Expected number value"),
}
}
#[test]
fn test_lsm_striping() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let mut keys_by_stripe: HashMap<u8, Vec<Key>> = HashMap::new();
for i in 0..1000 {
let key = Key::new(format!("key{}", i).into_bytes());
let stripe = key.stripe();
let mut item = HashMap::new();
item.insert("id".to_string(), Value::number(i));
db.put(key.clone(), item).unwrap();
keys_by_stripe.entry(stripe).or_insert_with(Vec::new).push(key);
}
assert!(keys_by_stripe.len() > 1, "Expected keys to be distributed across multiple stripes");
for (stripe, keys) in keys_by_stripe {
for key in keys {
let result = db.get(&key).unwrap();
assert!(result.is_some(), "Key should exist in stripe {}", stripe);
}
}
}
#[test]
fn test_lsm_stripe_independent_flush() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let pk = b"user#123";
let base_key = Key::new(pk.to_vec());
let stripe = base_key.stripe();
for i in 0..MEMTABLE_THRESHOLD + 10 {
let key = Key::with_sk(pk.to_vec(), format!("item#{}", i).into_bytes());
assert_eq!(key.stripe(), stripe);
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
db.flush().unwrap();
let mut found_striped_sst = false;
for entry in fs::read_dir(dir.path()).unwrap() {
let entry = entry.unwrap();
if let Some(name) = entry.file_name().to_str() {
if name.ends_with(".sst") && name.starts_with(&format!("{:03}-", stripe)) {
found_striped_sst = true;
break;
}
}
}
assert!(found_striped_sst, "Expected SST file with stripe prefix");
}
#[test]
fn test_lsm_query_basic() {
use crate::iterator::QueryParams;
use bytes::Bytes;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let pk = b"user#123";
for i in 0..10 {
let key = Key::with_sk(pk.to_vec(), format!("item#{:03}", i).into_bytes());
let mut item = HashMap::new();
item.insert("id".to_string(), Value::number(i));
item.insert("name".to_string(), Value::string(format!("Item {}", i)));
db.put(key, item).unwrap();
}
let params = QueryParams::new(Bytes::from(pk.to_vec()));
let result = db.query(params).unwrap();
assert_eq!(result.items.len(), 10);
assert_eq!(result.scanned_count, 10);
}
#[test]
fn test_lsm_query_with_limit() {
use crate::iterator::QueryParams;
use bytes::Bytes;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let pk = b"user#456";
for i in 0..20 {
let key = Key::with_sk(pk.to_vec(), format!("item#{:03}", i).into_bytes());
let mut item = HashMap::new();
item.insert("id".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
let params = QueryParams::new(Bytes::from(pk.to_vec())).with_limit(5);
let result = db.query(params).unwrap();
assert_eq!(result.items.len(), 5);
assert!(result.last_key.is_some());
}
#[test]
fn test_lsm_query_with_sk_condition() {
use crate::iterator::{QueryParams, SortKeyCondition};
use bytes::Bytes;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let pk = b"user#789";
for i in 0..10 {
let key = Key::with_sk(pk.to_vec(), format!("item#{:03}", i).into_bytes());
let mut item = HashMap::new();
item.insert("id".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
let params = QueryParams::new(Bytes::from(pk.to_vec()))
.with_sk_condition(SortKeyCondition::BeginsWith, Bytes::from("item#00"), None);
let result = db.query(params).unwrap();
assert_eq!(result.items.len(), 10);
}
#[test]
fn test_lsm_query_reverse() {
use crate::iterator::QueryParams;
use bytes::Bytes;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let pk = b"user#999";
for i in 0..5 {
let key = Key::with_sk(pk.to_vec(), format!("item#{}", i).into_bytes());
let mut item = HashMap::new();
item.insert("id".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
let params = QueryParams::new(Bytes::from(pk.to_vec())).with_direction(false);
let result = db.query(params).unwrap();
assert_eq!(result.items.len(), 5);
if let Some(Value::N(n)) = result.items[0].get("id") {
assert_eq!(n, "4");
} else {
panic!("Expected number value");
}
}
#[test]
fn test_lsm_scan_basic() {
use crate::iterator::ScanParams;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
for i in 0..20 {
let pk = format!("user#{}", i);
let key = Key::new(pk.into_bytes());
let mut item = HashMap::new();
item.insert("id".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
let params = ScanParams::new();
let result = db.scan(params).unwrap();
assert_eq!(result.items.len(), 20);
assert_eq!(result.scanned_count, 20);
}
#[test]
fn test_lsm_scan_with_limit() {
use crate::iterator::ScanParams;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
for i in 0..50 {
let pk = format!("item#{:03}", i);
let key = Key::new(pk.into_bytes());
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
let params = ScanParams::new().with_limit(10);
let result = db.scan(params).unwrap();
assert_eq!(result.items.len(), 10);
assert!(result.last_key.is_some());
}
#[test]
fn test_lsm_scan_parallel() {
use crate::iterator::ScanParams;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
for i in 0..100 {
let pk = format!("key{}", i);
let key = Key::new(pk.into_bytes());
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
let mut total_items = 0;
for segment in 0..4 {
let params = ScanParams::new().with_segment(segment, 4);
let result = db.scan(params).unwrap();
total_items += result.items.len();
}
assert_eq!(total_items, 100);
}
#[test]
fn test_lsm_scan_pagination() {
use crate::iterator::ScanParams;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
for i in 0..30 {
let pk = format!("user#{:03}", i);
let key = Key::new(pk.into_bytes());
let mut item = HashMap::new();
item.insert("id".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
let params1 = ScanParams::new().with_limit(10);
let result1 = db.scan(params1).unwrap();
assert_eq!(result1.items.len(), 10);
assert!(result1.last_key.is_some());
let params2 = ScanParams::new()
.with_limit(10)
.with_start_key(result1.last_key.unwrap());
let result2 = db.scan(params2).unwrap();
assert_eq!(result2.items.len(), 10);
let params3 = ScanParams::new()
.with_limit(10)
.with_start_key(result2.last_key.unwrap());
let result3 = db.scan(params3).unwrap();
assert_eq!(result3.items.len(), 10);
}
#[test]
fn test_lsm_compaction_triggered() {
use crate::compaction::COMPACTION_THRESHOLD;
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
for batch in 0..COMPACTION_THRESHOLD {
for i in 0..MEMTABLE_THRESHOLD {
let key = Key::new(format!("batch{:02}_key{:04}", batch, i).into_bytes());
let mut item = HashMap::new();
item.insert("batch".to_string(), Value::number(batch as i64));
item.insert("seq".to_string(), Value::number(i as i64));
db.put(key, item).unwrap();
}
}
for i in 0..MEMTABLE_THRESHOLD {
let key = Key::new(format!("final_key{:04}", i).into_bytes());
let mut item = HashMap::new();
item.insert("final".to_string(), Value::number(1));
db.put(key, item).unwrap();
}
let key1 = Key::new(b"batch00_key0000".to_vec());
let result1 = db.get(&key1).unwrap();
assert!(result1.is_some());
let item1 = result1.unwrap();
assert_eq!(item1.get("batch").unwrap(), &Value::N("0".to_string()));
let key2 = Key::new(b"final_key0000".to_vec());
let result2 = db.get(&key2).unwrap();
assert!(result2.is_some());
}
#[test]
fn test_lsm_compaction_removes_tombstones() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
for i in 0..100 {
let key = Key::new(format!("key{:03}", i).into_bytes());
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
db.flush().unwrap();
for i in 0..50 {
let key = Key::new(format!("key{:03}", i).into_bytes());
db.delete(key).unwrap();
}
db.flush().unwrap();
for i in 0..50 {
let key = Key::new(format!("key{:03}", i).into_bytes());
assert!(db.get(&key).unwrap().is_none());
}
for i in 50..100 {
let key = Key::new(format!("key{:03}", i).into_bytes());
let result = db.get(&key).unwrap();
assert!(result.is_some());
}
}
#[test]
fn test_lsm_compaction_keeps_latest_version() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let key = Key::new(b"test_key".to_vec());
let mut item1 = HashMap::new();
item1.insert("version".to_string(), Value::number(1));
db.put(key.clone(), item1).unwrap();
db.flush().unwrap();
let mut item2 = HashMap::new();
item2.insert("version".to_string(), Value::number(2));
db.put(key.clone(), item2).unwrap();
db.flush().unwrap();
let mut item3 = HashMap::new();
item3.insert("version".to_string(), Value::number(3));
db.put(key.clone(), item3).unwrap();
db.flush().unwrap();
let result = db.get(&key).unwrap().unwrap();
assert_eq!(result.get("version").unwrap(), &Value::N("3".to_string()));
drop(db);
let db = LsmEngine::open(dir.path()).unwrap();
let result = db.get(&key).unwrap().unwrap();
assert_eq!(result.get("version").unwrap(), &Value::N("3".to_string()));
}
#[test]
fn test_compaction_configuration() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let config = db.compaction_config();
assert!(config.enabled);
assert_eq!(config.sst_threshold, 10);
db.set_compaction_config(CompactionConfig::disabled());
let config = db.compaction_config();
assert!(!config.enabled);
let custom_config = CompactionConfig::new()
.with_sst_threshold(5)
.with_check_interval(30);
db.set_compaction_config(custom_config);
let config = db.compaction_config();
assert!(config.enabled);
assert_eq!(config.sst_threshold, 5);
assert_eq!(config.check_interval_secs, 30);
}
#[test]
fn test_compaction_statistics() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
let stats = db.compaction_stats();
assert_eq!(stats.total_compactions, 0);
assert_eq!(stats.total_ssts_merged, 0);
assert_eq!(stats.active_compactions, 0);
let pk = b"testdata";
for batch in 0..12 {
for i in 0..MEMTABLE_THRESHOLD {
let key = Key::with_sk(pk.to_vec(), format!("batch{:02}_key{:04}", batch, i).into_bytes());
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i as i64));
db.put(key, item).unwrap();
}
}
let stats = db.compaction_stats();
assert!(stats.total_compactions > 0, "Expected at least one compaction");
assert!(stats.total_ssts_merged > 0, "Expected SSTs to be merged");
assert!(stats.total_ssts_created > 0, "Expected new SSTs to be created");
}
#[test]
fn test_manual_compaction_trigger() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
db.set_compaction_config(CompactionConfig::new().with_sst_threshold(3));
let mut count = 0;
for i in 0..50000 {
let key = Key::new(format!("key{:06}", i).into_bytes());
if key.stripe() == 0 {
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i));
db.put(key, item).unwrap();
count += 1;
if count >= MEMTABLE_THRESHOLD * 4 {
break;
}
}
}
db.flush().unwrap();
let stats_before = db.compaction_stats();
db.trigger_compaction(0).unwrap();
let stats_after = db.compaction_stats();
assert!(
stats_after.total_compactions >= stats_before.total_compactions,
"Compaction count should increase or stay the same"
);
}
#[test]
fn test_compaction_disabled() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
db.set_compaction_config(CompactionConfig::disabled());
for batch in 0..15 {
for i in 0..MEMTABLE_THRESHOLD {
let key = Key::new(format!("batch{:02}_key{:04}", batch, i).into_bytes());
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i as i64));
db.put(key, item).unwrap();
}
}
db.flush().unwrap();
let stats = db.compaction_stats();
assert_eq!(stats.total_compactions, 0, "No compactions should occur when disabled");
}
#[test]
fn test_compaction_with_deletes_reclaims_space() {
let dir = TempDir::new().unwrap();
let db = LsmEngine::create(dir.path()).unwrap();
for i in 0..200 {
let key = Key::new(format!("key{:03}", i).into_bytes());
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i));
db.put(key, item).unwrap();
}
db.flush().unwrap();
for i in 0..100 {
let key = Key::new(format!("key{:03}", i).into_bytes());
db.delete(key).unwrap();
}
db.flush().unwrap();
for batch in 0..12 {
for i in 200..(200 + MEMTABLE_THRESHOLD) {
let key = Key::new(format!("key{:06}_{:02}", i, batch).into_bytes());
let mut item = HashMap::new();
item.insert("value".to_string(), Value::number(i as i64));
db.put(key, item).unwrap();
}
}
for i in 0..100 {
let key = Key::new(format!("key{:03}", i).into_bytes());
let result = db.get(&key).unwrap();
assert!(result.is_none(), "Deleted key should not be found");
}
for i in 100..200 {
let key = Key::new(format!("key{:03}", i).into_bytes());
let result = db.get(&key).unwrap();
assert!(result.is_some(), "Non-deleted key should still exist");
}
}
}