quatzal-storage 0.1.0

Sharded LSM row-storage engine for Quatzal: WAL, snapshots, and crash recovery on io_uring (Linux only).
// SPDX-License-Identifier: Apache-2.0
//! A single, self-contained, single-core LSM instance: WAL + MemTable + L0/L1 SSTables +
//! L0->L1 compaction. Every `Shard` runs on exactly one glommio executor thread and is never
//! shared across threads directly (see `crate::engine::Engine` for the cross-shard router) —
//! this is the "per-core, shared-nothing LSM instance" mitigation from `CLAUDE.md`.

use std::collections::BTreeSet;
use std::path::PathBuf;

use quatzal_schema::{Row, UaceError, UaceResult, Value};

use crate::memtable::MemTable;
use crate::pax::{PaxRecord, VectorField};
use crate::sstable::SsTable;
use crate::wal::Wal;

#[derive(Clone, Debug)]
pub struct ShardConfig {
    pub dir: PathBuf,
    pub memtable_flush_threshold_bytes: usize,
    pub l0_compaction_trigger: usize,
}

impl ShardConfig {
    pub fn new(dir: PathBuf) -> Self {
        ShardConfig {
            dir,
            memtable_flush_threshold_bytes: 1 << 20, // 1 MiB
            l0_compaction_trigger: 4,
        }
    }
}

#[derive(Default, Clone, Copy)]
pub struct ShardStats {
    pub memtable_records: usize,
    pub l0_tables: usize,
    pub l1_present: bool,
}

pub struct Shard {
    cfg: ShardConfig,
    wal_path: PathBuf,
    wal: Wal,
    memtable: MemTable,
    l0: Vec<SsTable>,
    l1: Option<SsTable>,
    next_seq: u64,
    next_sst_id: u64,
}

impl Shard {
    pub async fn open(cfg: ShardConfig) -> UaceResult<Self> {
        std::fs::create_dir_all(&cfg.dir)
            .map_err(|e| UaceError::Io(format!("create shard dir {:?}: {e}", cfg.dir)))?;

        let mut max_seq: u64 = 0;
        let mut max_sst_id: u64 = 0;
        let mut l0 = Vec::new();
        let mut l1 = None;

        let entries = std::fs::read_dir(&cfg.dir)
            .map_err(|e| UaceError::Io(format!("read shard dir {:?}: {e}", cfg.dir)))?;
        let mut sst_paths: Vec<(u32, u64, PathBuf)> = Vec::new();
        for entry in entries {
            let entry = entry.map_err(|e| UaceError::Io(format!("dir entry: {e}")))?;
            let path = entry.path();
            if let Some(name) = path.file_name().and_then(|n| n.to_str())
                && let Some((level, id)) = parse_sst_name(name)
            {
                sst_paths.push((level, id, path));
            }
        }
        for (level, id, path) in sst_paths {
            let table = SsTable::open(path, id, level).await?;
            max_sst_id = max_sst_id.max(id);
            match level {
                0 => l0.push(table),
                1 => l1 = Some(table),
                _ => return Err(UaceError::Io(format!("unexpected sstable level {level}"))),
            }
        }
        // l0 files sort by id ascending (oldest first) so iteration-newest-first is `.rev()`.
        l0.sort_by_key(|t| t.id);
        for t in &l0 {
            max_seq = max_seq.max(max_seq_in(t));
        }
        if let Some(t) = &l1 {
            max_seq = max_seq.max(max_seq_in(t));
        }

        let wal_path = cfg.dir.join("wal.log");
        let replayed = Wal::replay(&wal_path).await?;
        let mut memtable = MemTable::new();
        for record in replayed {
            max_seq = max_seq.max(record.seq);
            memtable.put(record);
        }
        let wal = Wal::open_for_append(&wal_path).await?;

        Ok(Shard {
            cfg,
            wal_path,
            wal,
            memtable,
            l0,
            l1,
            next_seq: max_seq + 1,
            next_sst_id: max_sst_id + 1,
        })
    }

    fn take_seq(&mut self) -> u64 {
        let s = self.next_seq;
        self.next_seq += 1;
        s
    }

    fn take_sst_id(&mut self) -> u64 {
        let id = self.next_sst_id;
        self.next_sst_id += 1;
        id
    }

    /// Full row insert/replace. `UACE-FR-1.1`, `UACE-FR-1.2`.
    pub async fn put(&mut self, row: Row) -> UaceResult<()> {
        let seq = self.take_seq();
        let record = PaxRecord::encode_full(seq, &row)?;
        self.wal.append(&record).await?;
        self.memtable.put(record);
        self.maybe_flush().await
    }

    /// `UACE-FR-1.9`: group-committed batch insert -- every row is encoded, the whole set
    /// is WAL-appended under one fsync, then applied to the MemTable in sequence order.
    /// Sequence numbers are still assigned per row, so newest-wins resolution is identical
    /// to the single-row path; only the fsync is shared.
    pub async fn put_batch(&mut self, rows: Vec<Row>) -> UaceResult<()> {
        if rows.is_empty() {
            return Ok(());
        }
        let mut records = Vec::with_capacity(rows.len());
        for row in &rows {
            records.push(PaxRecord::encode_full(self.take_seq(), row)?);
        }
        self.wal.append_batch(&records).await?;
        for record in records {
            self.memtable.put(record);
        }
        // One flush check for the batch rather than per row -- the MemTable threshold is a
        // size bound, and checking it once after a bounded batch keeps the same guarantee.
        self.maybe_flush().await
    }

    /// `UACE-FR-1.10`: many point reads in one call. Purely an amortization of the
    /// cross-thread request round trip -- each lookup is the same in-memory
    /// MemTable/L0/L1 resolution `get` performs, so the two can never disagree.
    pub fn get_batch(&self, keys: &[Vec<u8>]) -> UaceResult<Vec<Option<Row>>> {
        keys.iter().map(|key| self.get(key)).collect()
    }

    /// Scalar-only update — never touches the vector bytes. `UACE-FR-1.3`.
    pub async fn update_scalars(
        &mut self,
        key: &[u8],
        scalars: std::collections::BTreeMap<String, Value>,
    ) -> UaceResult<()> {
        let seq = self.take_seq();
        let record = PaxRecord::encode_scalar_update(seq, key, &scalars)?;
        self.wal.append(&record).await?;
        self.memtable.put(record);
        self.maybe_flush().await
    }

    pub async fn delete(&mut self, key: &[u8]) -> UaceResult<()> {
        let seq = self.take_seq();
        let record = PaxRecord::encode_tombstone(seq, key);
        self.wal.append(&record).await?;
        self.memtable.put(record);
        self.maybe_flush().await
    }

    /// Point read. `UACE-FR-1.4` — resolves scalar + vector state across MemTable -> L0
    /// (newest first) -> L1.
    pub fn get(&self, key: &[u8]) -> UaceResult<Option<Row>> {
        let mut chain: Vec<&PaxRecord> = Vec::new();
        if let Some(r) = self.memtable.get(key) {
            chain.push(r);
        }
        for table in self.l0.iter().rev() {
            if let Some(r) = table.get(key) {
                chain.push(r);
            }
        }
        if let Some(table) = &self.l1
            && let Some(r) = table.get(key)
        {
            chain.push(r);
        }
        resolve_from_records(chain.into_iter())
    }

    /// `UACE-FR-1.8`: every live row this shard holds -- newest version per key,
    /// tombstoned keys excluded. Key-union across MemTable/L0/L1 (the same discipline
    /// `compact_l0_into_l1` uses), each key resolved through the existing newest-wins
    /// `get` path so scan and point-read can never disagree.
    pub fn scan(&self) -> UaceResult<Vec<Row>> {
        let mut keys: BTreeSet<Vec<u8>> = BTreeSet::new();
        for record in self.memtable.iter_sorted() {
            keys.insert(record.key.clone());
        }
        for table in &self.l0 {
            for record in table.iter() {
                keys.insert(record.key.clone());
            }
        }
        if let Some(table) = &self.l1 {
            for record in table.iter() {
                keys.insert(record.key.clone());
            }
        }
        let mut rows = Vec::new();
        for key in keys {
            if let Some(row) = self.get(&key)? {
                rows.push(row);
            }
        }
        Ok(rows)
    }

    pub fn stats(&self) -> ShardStats {
        ShardStats {
            memtable_records: self.memtable.len(),
            l0_tables: self.l0.len(),
            l1_present: self.l1.is_some(),
        }
    }

    async fn maybe_flush(&mut self) -> UaceResult<()> {
        if self.memtable.approx_bytes() < self.cfg.memtable_flush_threshold_bytes {
            return Ok(());
        }
        let id = self.take_sst_id();
        let table = SsTable::flush_memtable(&self.cfg.dir, id, 0, &self.memtable).await?;
        self.l0.push(table);
        self.memtable.clear();
        self.wal = Wal::reset(&self.wal_path).await?;

        if self.l0.len() >= self.cfg.l0_compaction_trigger {
            self.compact_l0_into_l1().await?;
        }
        Ok(())
    }

    /// `UACE-FR-1.6`: merges all L0 tables + the existing L1 table into a new L1 table,
    /// keeping the newest version of every key and dropping tombstones (this is the bottom
    /// level in Phase 1's two-level scheme, so tombstones are safe to drop here). Vector
    /// bytes are only re-materialized here — once per compaction pass — even if many
    /// scalar-only updates queued up across the L0 tables being merged; see `CLAUDE.md`.
    async fn compact_l0_into_l1(&mut self) -> UaceResult<()> {
        let mut keys: BTreeSet<Vec<u8>> = BTreeSet::new();
        for table in &self.l0 {
            for record in table.iter() {
                keys.insert(record.key.clone());
            }
        }
        if let Some(table) = &self.l1 {
            for record in table.iter() {
                keys.insert(record.key.clone());
            }
        }

        let mut merged: Vec<PaxRecord> = Vec::with_capacity(keys.len());
        for key in &keys {
            let mut chain: Vec<&PaxRecord> = Vec::new();
            for table in self.l0.iter().rev() {
                if let Some(r) = table.get(key) {
                    chain.push(r);
                }
            }
            if let Some(table) = &self.l1
                && let Some(r) = table.get(key)
            {
                chain.push(r);
            }
            if let Some(row) = resolve_from_records(chain.into_iter())? {
                merged.push(PaxRecord::encode_full(0, &row)?); // seq assigned below
            }
        }
        // Assign fresh monotonic seqs now that we're done borrowing self.l0/self.l1.
        for record in &mut merged {
            record.seq = self.take_seq();
        }

        let new_id = self.take_sst_id();
        let old_l0: Vec<SsTable> = std::mem::take(&mut self.l0);
        let old_l1 = self.l1.take();

        let new_l1 = SsTable::write_records(&self.cfg.dir, new_id, 1, merged.into_iter()).await?;

        for table in old_l0 {
            let _ = std::fs::remove_file(&table.path);
        }
        if let Some(table) = old_l1 {
            let _ = std::fs::remove_file(&table.path);
        }
        self.l1 = Some(new_l1);
        Ok(())
    }
}

/// Resolve a key's current state from an iterator of candidate records in
/// newest-to-oldest priority order. Returns `Ok(None)` if the key is absent or its newest
/// record is a tombstone. `UACE-FR-1.4`, `UACE-FR-1.5`.
pub(crate) fn resolve_from_records<'a>(
    records: impl Iterator<Item = &'a PaxRecord>,
) -> UaceResult<Option<Row>> {
    let mut scalar_answer: Option<&PaxRecord> = None;
    let mut vector_bytes: Option<Vec<u8>> = None;
    let mut vector_resolved = false;

    for record in records {
        if scalar_answer.is_none() {
            scalar_answer = Some(record);
            if record.tombstone {
                return Ok(None);
            }
        }
        if !vector_resolved {
            match &record.vector {
                VectorField::Set(bytes) => {
                    vector_bytes = Some(bytes.clone());
                    vector_resolved = true;
                }
                VectorField::Unset => {
                    vector_resolved = true;
                }
                VectorField::Unchanged => {}
            }
        }
        if scalar_answer.is_some() && vector_resolved {
            break;
        }
    }

    let scalar_answer = match scalar_answer {
        Some(r) => r,
        None => return Ok(None),
    };
    if scalar_answer.tombstone {
        return Ok(None);
    }

    let resolved_vector = match vector_bytes {
        Some(bytes) => Some(
            bincode::deserialize::<Vec<f32>>(&bytes)
                .map_err(|e| UaceError::Codec(format!("vector decode: {e}")))?,
        ),
        None => None,
    };
    scalar_answer.decode_with_vector(resolved_vector).map(Some)
}

fn max_seq_in(table: &SsTable) -> u64 {
    table.iter().map(|r| r.seq).max().unwrap_or(0)
}

/// Parses `l{level}-{id:016x}.sst` back into `(level, id)`.
fn parse_sst_name(name: &str) -> Option<(u32, u64)> {
    let name = name.strip_suffix(".sst")?;
    let rest = name.strip_prefix('l')?;
    let (level_str, id_str) = rest.split_once('-')?;
    let level = level_str.parse::<u32>().ok()?;
    let id = u64::from_str_radix(id_str, 16).ok()?;
    Some((level, id))
}