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, 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.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
}
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
}
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);
}
self.maybe_flush().await
}
pub fn get_batch(&self, keys: &[Vec<u8>]) -> UaceResult<Vec<Option<Row>>> {
keys.iter().map(|key| self.get(key)).collect()
}
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
}
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())
}
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(())
}
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)?); }
}
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(())
}
}
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)
}
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))
}