use std::collections::BinaryHeap;
use std::fs::File;
use std::io::{BufReader, BufWriter, Read, Seek, SeekFrom, Write};
use std::path::PathBuf;
use crate::types::StoreError;
type KvResult = Result<(Vec<u8>, Vec<u8>), StoreError>;
type KvPair = (Vec<u8>, Vec<u8>);
const MAX_OPEN_CHUNKS: usize = 128;
#[derive(Eq, PartialEq)]
struct HeapEntry {
key: Vec<u8>,
idx: usize,
val: Vec<u8>,
}
impl Ord for HeapEntry {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other.key.cmp(&self.key).then(other.idx.cmp(&self.idx))
}
}
impl PartialOrd for HeapEntry {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
pub(crate) struct ExternalSorter {
work_dir: PathBuf,
max_memory_bytes: usize,
raw_data: Vec<u8>,
offsets: Vec<(usize, usize, usize, usize)>,
buffer_bytes: usize,
chunk_files: Vec<PathBuf>,
chunk_counter: usize,
}
impl ExternalSorter {
pub(crate) fn new(work_dir: PathBuf, max_memory_bytes: usize) -> Self {
let _ = std::fs::create_dir_all(&work_dir);
Self {
work_dir,
max_memory_bytes,
raw_data: Vec::new(),
offsets: Vec::new(),
buffer_bytes: 0,
chunk_files: Vec::new(),
chunk_counter: 0,
}
}
pub(crate) fn push(&mut self, key: Vec<u8>, value: Vec<u8>) -> Result<(), StoreError> {
self.buffer_bytes += key.len() + value.len();
let key_start = self.raw_data.len();
let key_len = key.len();
self.raw_data.extend_from_slice(&key);
let val_start = self.raw_data.len();
let val_len = value.len();
self.raw_data.extend_from_slice(&value);
self.offsets.push((key_start, key_len, val_start, val_len));
if self.buffer_bytes >= self.max_memory_bytes {
self.spill()?;
}
Ok(())
}
pub(crate) fn finish(mut self) -> Result<impl Iterator<Item = KvResult>, StoreError> {
if self.chunk_files.is_empty() {
let raw = &self.raw_data;
self.offsets
.sort_unstable_by(|&(ks1, kl1, _, _), &(ks2, kl2, _, _)| raw[ks1..ks1 + kl1].cmp(&raw[ks2..ks2 + kl2]));
let raw_data = std::mem::take(&mut self.raw_data);
let pairs: Vec<KvPair> = std::mem::take(&mut self.offsets)
.into_iter()
.map(|(ks, kl, vs, vl)| (raw_data[ks..ks + kl].to_vec(), raw_data[vs..vs + vl].to_vec()))
.collect();
return Ok(SortedIter::Memory(pairs.into_iter()));
}
if !self.offsets.is_empty() {
if let Err(e) = self.spill() {
self.cleanup_chunk_files();
return Err(e);
}
}
if let Err(e) = self.cascade_merge() {
self.cleanup_chunk_files();
return Err(e);
}
let chunk_files = std::mem::take(&mut self.chunk_files);
let readers: Vec<_> = chunk_files.into_iter().map(ChunkReader::open).collect::<Result<_, _>>()?;
Ok(SortedIter::Merge(Merger::new(readers)))
}
fn spill(&mut self) -> Result<(), StoreError> {
let raw = &self.raw_data;
self.offsets
.sort_unstable_by(|&(ks1, kl1, _, _), &(ks2, kl2, _, _)| raw[ks1..ks1 + kl1].cmp(&raw[ks2..ks2 + kl2]));
let path = self.work_dir.join(format!("chunk_{:06}.bin", self.chunk_counter));
self.chunk_counter += 1;
let file = File::create(&path).map_err(StoreError::Io)?;
let mut w = BufWriter::new(file);
let count = self.offsets.len() as u64;
w.write_all(&count.to_le_bytes()).map_err(StoreError::Io)?;
for &(ks, kl, vs, vl) in &self.offsets {
w.write_all(&(kl as u32).to_le_bytes()).map_err(StoreError::Io)?;
w.write_all(&self.raw_data[ks..ks + kl]).map_err(StoreError::Io)?;
w.write_all(&(vl as u32).to_le_bytes()).map_err(StoreError::Io)?;
w.write_all(&self.raw_data[vs..vs + vl]).map_err(StoreError::Io)?;
}
w.flush().map_err(StoreError::Io)?;
self.chunk_files.push(path);
self.raw_data.clear();
self.offsets.clear();
self.buffer_bytes = 0;
Ok(())
}
fn cascade_merge(&mut self) -> Result<(), StoreError> {
while self.chunk_files.len() > MAX_OPEN_CHUNKS {
let input_chunks: Vec<PathBuf> = self.chunk_files.drain(..).collect();
if let Err(e) = self.cascade_merge_pass(&input_chunks) {
for path in &input_chunks {
let _ = std::fs::remove_file(path);
}
self.cleanup_chunk_files();
return Err(e);
}
}
Ok(())
}
fn cascade_merge_pass(&mut self, input_chunks: &[PathBuf]) -> Result<(), StoreError> {
for group in input_chunks.chunks(MAX_OPEN_CHUNKS) {
if group.len() == 1 {
self.chunk_files.push(group[0].clone());
continue;
}
let merged_path = self.work_dir.join(format!("chunk_{:06}.bin", self.chunk_counter));
self.chunk_counter += 1;
self.chunk_files.push(merged_path.clone());
let readers: Vec<_> = group.iter().map(|p| ChunkReader::open(p.clone())).collect::<Result<_, _>>()?;
let mut merger = Merger::new(readers);
let mut file = File::create(&merged_path).map_err(StoreError::Io)?;
file.write_all(&0u64.to_le_bytes()).map_err(StoreError::Io)?;
let mut w = BufWriter::new(file);
let mut count = 0u64;
for item in &mut merger {
let (key, val) = item?;
w.write_all(&(key.len() as u32).to_le_bytes()).map_err(StoreError::Io)?;
w.write_all(&key).map_err(StoreError::Io)?;
w.write_all(&(val.len() as u32).to_le_bytes()).map_err(StoreError::Io)?;
w.write_all(&val).map_err(StoreError::Io)?;
count += 1;
}
w.flush().map_err(StoreError::Io)?;
drop(merger);
let mut file = w.into_inner().map_err(|e| StoreError::Io(e.into_error()))?;
file.seek(SeekFrom::Start(0)).map_err(StoreError::Io)?;
file.write_all(&count.to_le_bytes()).map_err(StoreError::Io)?;
}
Ok(())
}
fn cleanup_chunk_files(&mut self) {
for path in self.chunk_files.drain(..) {
let _ = std::fs::remove_file(path);
}
}
}
impl Drop for ExternalSorter {
fn drop(&mut self) {
self.cleanup_chunk_files();
}
}
enum SortedIter {
Memory(std::vec::IntoIter<KvPair>),
Merge(Merger),
}
impl Iterator for SortedIter {
type Item = KvResult;
fn next(&mut self) -> Option<Self::Item> {
match self {
SortedIter::Memory(iter) => iter.next().map(Ok),
SortedIter::Merge(merger) => merger.next(),
}
}
}
struct ChunkReader {
reader: BufReader<File>,
remaining: u64,
path: Option<PathBuf>,
}
impl ChunkReader {
fn open(path: PathBuf) -> Result<Self, StoreError> {
let file = File::open(&path).map_err(StoreError::Io)?;
let mut reader = BufReader::new(file);
let mut count_buf = [0u8; 8];
reader.read_exact(&mut count_buf).map_err(StoreError::Io)?;
let remaining = u64::from_le_bytes(count_buf);
Ok(Self { reader, remaining, path: Some(path) })
}
fn take_front(&mut self) -> Option<KvResult> {
if self.remaining == 0 {
return None;
}
self.remaining -= 1;
let mut len_buf = [0u8; 4];
if let Err(e) = self.reader.read_exact(&mut len_buf) {
return Some(Err(StoreError::Io(e)));
}
let key_len = u32::from_le_bytes(len_buf) as usize;
let mut key = vec![0u8; key_len];
if let Err(e) = self.reader.read_exact(&mut key) {
return Some(Err(StoreError::Io(e)));
}
if let Err(e) = self.reader.read_exact(&mut len_buf) {
return Some(Err(StoreError::Io(e)));
}
let val_len = u32::from_le_bytes(len_buf) as usize;
let mut val = vec![0u8; val_len];
if let Err(e) = self.reader.read_exact(&mut val) {
return Some(Err(StoreError::Io(e)));
}
Some(Ok((key, val)))
}
}
impl Drop for ChunkReader {
fn drop(&mut self) {
if let Some(ref p) = self.path {
let _ = std::fs::remove_file(p);
}
}
}
struct Merger {
readers: Vec<ChunkReader>,
heap: BinaryHeap<HeapEntry>,
}
impl Merger {
fn new(mut readers: Vec<ChunkReader>) -> Self {
let mut heap = BinaryHeap::with_capacity(readers.len());
for (i, r) in readers.iter_mut().enumerate() {
if let Some(Ok((key, val))) = r.take_front() {
heap.push(HeapEntry { key, idx: i, val });
}
}
Self { readers, heap }
}
}
impl Iterator for Merger {
type Item = KvResult;
fn next(&mut self) -> Option<Self::Item> {
let HeapEntry { key, idx, val } = self.heap.pop()?;
if let Some(next) = self.readers[idx].take_front() {
match next {
Ok((nk, nv)) => self.heap.push(HeapEntry { key: nk, idx, val: nv }),
Err(e) => return Some(Err(e)),
}
}
Some(Ok((key, val)))
}
}
#[cfg(test)]
mod tests {
use tempfile::tempdir;
use super::*;
fn check_sorted(results: &[(Vec<u8>, Vec<u8>)]) {
for w in results.windows(2) {
assert!(w[0].0 <= w[1].0, "out of order: {:?} > {:?}", w[0].0, w[1].0);
}
}
#[test]
fn test_in_memory_sort() {
let dir = tempdir().unwrap();
let mut sorter = ExternalSorter::new(dir.path().join("s"), 1024 * 1024);
for i in (0i32..200).rev() {
sorter.push(i.to_be_bytes().to_vec(), Vec::new()).unwrap();
}
let results: Vec<_> = sorter.finish().unwrap().map(|r| r.unwrap()).collect();
assert_eq!(results.len(), 200);
check_sorted(&results);
}
#[test]
fn test_spill_and_merge() {
let dir = tempdir().unwrap();
let mut sorter = ExternalSorter::new(dir.path().join("s"), 50);
for i in (0i32..500).rev() {
sorter.push(i.to_be_bytes().to_vec(), vec![0u8; 30]).unwrap();
}
let results: Vec<_> = sorter.finish().unwrap().map(|r| r.unwrap()).collect();
assert_eq!(results.len(), 500);
check_sorted(&results);
}
#[test]
fn test_empty() {
let dir = tempdir().unwrap();
let sorter = ExternalSorter::new(dir.path().join("s"), 1024);
assert_eq!(sorter.finish().unwrap().count(), 0);
}
#[test]
fn test_duplicate_keys_stable() {
let dir = tempdir().unwrap();
let mut sorter = ExternalSorter::new(dir.path().join("s"), 1024 * 1024);
for _ in 0..5 {
sorter.push(vec![0], vec![]).unwrap();
sorter.push(vec![1], vec![]).unwrap();
}
let results: Vec<_> = sorter.finish().unwrap().map(|r| r.unwrap()).collect();
assert_eq!(results.len(), 10);
check_sorted(&results);
}
#[test]
fn test_cascaded_merge() {
let dir = tempdir().unwrap();
let mut sorter = ExternalSorter::new(dir.path().join("s"), 1);
for i in (0u32..300).rev() {
sorter.push(i.to_be_bytes().to_vec(), Vec::new()).unwrap();
}
let results: Vec<_> = sorter.finish().unwrap().map(|r| r.unwrap()).collect();
assert_eq!(results.len(), 300);
check_sorted(&results);
let leftover = std::fs::read_dir(dir.path().join("s")).map(|d| d.count()).unwrap_or(0);
assert_eq!(leftover, 0, "chunk files leaked");
}
}