use std::path::{
Path,
PathBuf,
};
use futures_util::stream;
use serde::Serialize;
use serde::de::DeserializeOwned;
use crate::engine::cleanup::AsyncCleanupGuard;
use crate::engine::merge::{
ValueStream,
cascade_and_stream,
merge_group_to_file,
};
use crate::engine::row::RunRow;
use crate::engine::run::RunWriter;
use crate::error::SorterError;
use crate::plan::SortPlan;
#[cfg(test)]
mod tests;
pub struct SortSession<K, V> {
plan: SortPlan,
dir: PathBuf,
dedup: bool,
buffer: Vec<RunRow<K, V>>,
buffer_bytes: usize,
spills: [Option<PathBuf>; u64::BITS as usize],
next_run: u64,
failed: bool,
hold: Option<Box<dyn Send>>,
cleanup: AsyncCleanupGuard,
}
impl<K, V> SortSession<K, V>
where
K: Ord + Clone + Serialize + DeserializeOwned + Send + 'static,
V: Serialize + DeserializeOwned + Send + 'static,
{
#[must_use]
pub fn new(plan: SortPlan, dir: PathBuf, dedup: bool) -> Self {
Self {
plan,
dir,
dedup,
buffer: Vec::new(),
buffer_bytes: 0,
spills: std::array::from_fn(|_| None),
next_run: 0,
failed: false,
hold: None,
cleanup: AsyncCleanupGuard::disarmed(),
}
}
#[must_use]
pub fn with_temp_dir(plan: SortPlan, temp_root: &Path, dedup: bool) -> Self {
let dir = temp_root.join(format!("sort-{}", uuid::Uuid::new_v4()));
Self::new(plan, dir, dedup)
}
#[must_use]
pub fn hold_resource(mut self, resource: Box<dyn Send>) -> Self {
self.hold = Some(resource);
self
}
#[must_use]
pub fn scratch_dir(&self) -> &Path {
&self.dir
}
pub async fn push_with_size(
&mut self,
key: K,
value: V,
estimated_bytes: usize,
) -> Result<(), SorterError> {
self.ensure_usable()?;
let bytes = estimated_bytes.max(1);
if let SortPlan::External {
run_buffer_bytes, ..
} = self.plan
&& !self.buffer.is_empty()
&& self.buffer_bytes.saturating_add(bytes) > run_buffer_bytes
{
self.spill().await?;
}
self.buffer.push(RunRow::new(key, value));
self.buffer_bytes = self.buffer_bytes.saturating_add(bytes);
Ok(())
}
pub async fn push(&mut self, key: K, value: V) -> Result<(), SorterError> {
let estimate = std::mem::size_of::<(K, V)>().max(1);
self.push_with_size(key, value, estimate).await
}
pub async fn finish(mut self) -> Result<ValueStream<V>, SorterError> {
self.ensure_usable()?;
self.buffer.sort_by(|left, right| left.key.cmp(&right.key));
let fan_in = match self.plan {
SortPlan::InMemory => return Ok(self.in_memory_stream()),
SortPlan::External { .. } if self.next_run == 0 => {
return Ok(self.in_memory_stream());
}
SortPlan::External { max_fan_in, .. } => max_fan_in as usize,
};
if !self.buffer.is_empty() {
self.spill().await?;
}
let hold = self.hold.take();
let guard = std::mem::take(&mut self.cleanup);
let spills = self
.spills
.iter_mut()
.rev()
.filter_map(Option::take)
.collect();
let dir = self.dir.clone();
cascade_and_stream::<K, V>(spills, fan_in.max(2), dir, guard, self.dedup, hold).await
}
fn in_memory_stream(&mut self) -> ValueStream<V> {
let rows = std::mem::take(&mut self.buffer);
let mut values = Vec::with_capacity(rows.len());
let mut last: Option<K> = None;
for row in rows {
if self.dedup {
if last.as_ref() == Some(&row.key) {
continue;
}
last = Some(row.key.clone());
}
values.push(row.value);
}
let hold = self.hold.take();
Box::pin(stream::unfold(
(values.into_iter(), hold),
|(mut values, hold)| async move { values.next().map(|value| (Ok(value), (values, hold))) },
))
}
async fn spill(&mut self) -> Result<(), SorterError> {
if self.buffer.is_empty() {
return Ok(());
}
self.failed = true;
let ordinal = self.next_run;
self.next_run = ordinal.checked_add(1).ok_or(SorterError::RunLimit)?;
self.buffer.sort_by(|left, right| left.key.cmp(&right.key));
async_fs_io::ensure_dir(&self.dir).await?;
self.cleanup.arm(self.dir.clone());
let path = self.dir.join(format!("run-{ordinal:020}.cbor"));
let mut writer = RunWriter::create(path).await?;
for row in &self.buffer {
writer.write_row(row).await?;
}
let mut run = writer.finish().await?;
self.buffer.clear();
self.buffer_bytes = 0;
for (level, slot) in self.spills.iter_mut().enumerate() {
let Some(older) = slot.take() else {
*slot = Some(run);
self.failed = false;
return Ok(());
};
let output = self
.dir
.join(format!("carry-{ordinal:020}-{level:02}.cbor"));
run = merge_group_to_file::<K, V>(&[older, run], output, false).await?;
}
unreachable!("a checked u64 spill count always fits in 64 frontier levels")
}
fn ensure_usable(&self) -> Result<(), SorterError> {
if self.failed {
return Err(SorterError::SessionFailed);
}
Ok(())
}
}