use std::cmp::Reverse;
use std::collections::BinaryHeap;
use std::fs::File;
use std::io::{BufReader, BufWriter, Read as _, Seek as _, SeekFrom, Write as _};
use std::path::{Path, PathBuf};
use crate::error::{Error, Result};
pub const MAXIMUM_FAN_IN: usize = 64;
const SPILL_BUFFER_BYTES: usize = 256 * 1024;
pub trait SpillRecord: Ord + Sized {
fn encode(&self, buffer: &mut Vec<u8>);
fn decode(bytes: &[u8]) -> Result<Self>;
fn resident_size(&self) -> usize;
}
#[derive(Clone, PartialEq, Eq, Debug)]
pub struct RunLocation {
directory: PathBuf,
name_prefix: String,
}
impl RunLocation {
#[must_use]
pub fn new(directory: impl Into<PathBuf>, name_prefix: impl Into<String>) -> Self {
Self {
directory: directory.into(),
name_prefix: name_prefix.into(),
}
}
#[must_use]
pub fn directory(&self) -> &Path {
&self.directory
}
#[must_use]
pub fn name_prefix(&self) -> &str {
&self.name_prefix
}
fn run_path(&self, number: usize) -> PathBuf {
self.directory
.join(format!("{}-{number:06}.run", self.name_prefix))
}
}
#[derive(Clone, Debug)]
pub struct SortBudget {
inner: std::rc::Rc<std::cell::Cell<BudgetState>>,
}
#[derive(Clone, Copy, Debug)]
struct BudgetState {
limit: usize,
charged: usize,
}
impl SortBudget {
#[must_use]
pub fn of_bytes(limit: usize) -> Self {
Self {
inner: std::rc::Rc::new(std::cell::Cell::new(BudgetState { limit, charged: 0 })),
}
}
#[must_use]
pub fn limit(&self) -> usize {
self.inner.get().limit
}
#[must_use]
pub fn charged(&self) -> usize {
self.inner.get().charged
}
fn charge(&self, bytes: usize) -> bool {
let mut state = self.inner.get();
state.charged = state.charged.saturating_add(bytes);
self.inner.set(state);
state.charged > state.limit
}
fn release(&self, bytes: usize) {
let mut state = self.inner.get();
state.charged = state.charged.saturating_sub(bytes);
self.inner.set(state);
}
}
pub(crate) struct SortedRuns<Record: SpillRecord> {
location: RunLocation,
budget: SortBudget,
resident: Vec<Record>,
resident_bytes: usize,
spilled: Vec<PathBuf>,
next_run_number: usize,
}
impl<Record: SpillRecord> SortedRuns<Record> {
#[must_use]
pub(crate) fn new(location: RunLocation, budget: SortBudget) -> Self {
Self {
location,
budget,
resident: Vec::new(),
resident_bytes: 0,
spilled: Vec::new(),
next_run_number: 0,
}
}
pub(crate) fn push(&mut self, record: Record) -> Result<()> {
let size = record.resident_size();
self.resident.push(record);
self.resident_bytes = self.resident_bytes.saturating_add(size);
if self.budget.charge(size) {
self.spill()?;
}
Ok(())
}
#[must_use]
#[cfg_attr(
not(test),
expect(dead_code, reason = "task 0707's plan is the first production caller")
)]
pub(crate) fn spilled_run_count(&self) -> usize {
self.spilled.len()
}
fn spill(&mut self) -> Result<()> {
if self.resident.is_empty() {
return Ok(());
}
self.resident.sort_unstable();
let path = self.location.run_path(self.next_run_number);
self.next_run_number += 1;
write_run(&path, self.resident.drain(..).map(Ok))?;
self.spilled.push(path);
self.budget.release(self.resident_bytes);
self.resident_bytes = 0;
Ok(())
}
pub(crate) fn into_sorted(mut self) -> Result<SortedPass<'static, Record>> {
self.spill()?;
let spilled = std::mem::take(&mut self.spilled);
let runs = self.reduce_to_fan_in(spilled)?;
SortedPass::over_owned(&runs)
}
pub(crate) fn into_sorted_passes(mut self) -> Result<SortedPasses<Record>> {
self.spill()?;
let spilled = std::mem::take(&mut self.spilled);
let runs = self.reduce_to_fan_in(spilled)?;
Ok(SortedPasses {
runs,
resident: Vec::new(),
})
}
fn reduce_to_fan_in(&mut self, mut runs: Vec<PathBuf>) -> Result<Vec<PathBuf>> {
while runs.len() > MAXIMUM_FAN_IN {
record_merge_pass();
let mut reduced = Vec::with_capacity(runs.len().div_ceil(MAXIMUM_FAN_IN));
for group in runs.chunks(MAXIMUM_FAN_IN) {
let path = self.location.run_path(self.next_run_number);
self.next_run_number += 1;
let merged = SortedPass::<Record>::over_borrowed(group)?;
write_run(&path, merged)?;
for input in group {
let _ = std::fs::remove_file(input);
}
reduced.push(path);
}
runs = reduced;
}
Ok(runs)
}
}
impl<Record: SpillRecord> Drop for SortedRuns<Record> {
fn drop(&mut self) {
self.budget.release(self.resident_bytes);
for path in &self.spilled {
let _ = std::fs::remove_file(path);
}
}
}
pub struct SortedPasses<Record: SpillRecord> {
runs: Vec<PathBuf>,
resident: Vec<Record>,
}
impl<Record: SpillRecord> SortedPasses<Record> {
#[must_use]
pub fn from_sorted_records(records: Vec<Record>) -> Self {
Self {
runs: Vec::new(),
resident: records,
}
}
pub fn pass(&mut self) -> Result<SortedPass<'_, Record>> {
if self.runs.is_empty() {
return Ok(SortedPass::over_slice(&self.resident));
}
SortedPass::over_borrowed(&self.runs)
}
}
impl<Record: SpillRecord> Drop for SortedPasses<Record> {
fn drop(&mut self) {
for path in &self.runs {
let _ = std::fs::remove_file(path);
}
}
}
pub struct SortedPass<'records, Record: SpillRecord> {
heap: BinaryHeap<Reverse<HeapEntry<Record>>>,
cursors: Vec<RunCursor>,
unlink_exhausted: bool,
resident: std::slice::Iter<'records, Record>,
resident_only: bool,
}
impl<Record: SpillRecord> SortedPass<'static, Record> {
fn over_owned(runs: &[PathBuf]) -> Result<Self> {
let mut pass = Self::open(runs, true)?;
pass.unlink_exhausted = true;
Ok(pass)
}
}
impl<'records, Record: SpillRecord> SortedPass<'records, Record> {
fn over_borrowed(runs: &[PathBuf]) -> Result<Self> {
Self::open(runs, false)
}
fn over_slice(records: &'records [Record]) -> Self {
Self {
heap: BinaryHeap::new(),
cursors: Vec::new(),
unlink_exhausted: false,
resident: records.iter(),
resident_only: true,
}
}
fn open(runs: &[PathBuf], owned: bool) -> Result<Self> {
let mut cursors = Vec::with_capacity(runs.len());
let mut heap = BinaryHeap::with_capacity(runs.len());
for path in runs {
let mut cursor = RunCursor::open(path)?;
if let Some(record) = cursor.next_record::<Record>()? {
heap.push(Reverse(HeapEntry {
record,
run: cursors.len(),
}));
}
cursors.push(cursor);
}
record_open_files(cursors.len());
Ok(Self {
heap,
cursors,
unlink_exhausted: owned,
resident: [].iter(),
resident_only: false,
})
}
}
impl<Record: SpillRecord> Iterator for SortedPass<'_, Record> {
type Item = Result<Record>;
fn next(&mut self) -> Option<Self::Item> {
if self.resident_only {
let record = self.resident.next()?;
let mut buffer = Vec::new();
record.encode(&mut buffer);
return Some(Record::decode(&buffer));
}
let Reverse(HeapEntry { record, run }) = self.heap.pop()?;
match self.cursors[run].next_record::<Record>() {
Ok(Some(next)) => self.heap.push(Reverse(HeapEntry { record: next, run })),
Ok(None) => {
if self.unlink_exhausted {
self.cursors[run].unlink();
}
}
Err(error) => return Some(Err(error)),
}
Some(Ok(record))
}
}
impl<Record: SpillRecord> Drop for SortedPass<'_, Record> {
fn drop(&mut self) {
if self.unlink_exhausted {
for cursor in &mut self.cursors {
cursor.unlink();
}
}
}
}
struct HeapEntry<Record: SpillRecord> {
record: Record,
run: usize,
}
impl<Record: SpillRecord> Ord for HeapEntry<Record> {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.record
.cmp(&other.record)
.then_with(|| self.run.cmp(&other.run))
}
}
impl<Record: SpillRecord> PartialOrd for HeapEntry<Record> {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl<Record: SpillRecord> PartialEq for HeapEntry<Record> {
fn eq(&self, other: &Self) -> bool {
self.cmp(other) == std::cmp::Ordering::Equal
}
}
impl<Record: SpillRecord> Eq for HeapEntry<Record> {}
struct RunCursor {
path: PathBuf,
reader: Option<BufReader<File>>,
}
impl RunCursor {
fn open(path: &Path) -> Result<Self> {
let file = File::open(path)?;
Ok(Self {
path: path.to_path_buf(),
reader: Some(BufReader::with_capacity(SPILL_BUFFER_BYTES, file)),
})
}
fn next_record<Record: SpillRecord>(&mut self) -> Result<Option<Record>> {
let Some(reader) = self.reader.as_mut() else {
return Ok(None);
};
let mut length = [0u8; 4];
match reader.read_exact(&mut length) {
Ok(()) => {}
Err(error) if error.kind() == std::io::ErrorKind::UnexpectedEof => {
self.reader = None;
record_closed_file();
return Ok(None);
}
Err(error) => return Err(Error::InputOutput(error)),
}
let length = u32::from_le_bytes(length) as usize;
let mut bytes = vec![0u8; length];
reader.read_exact(&mut bytes)?;
Record::decode(&bytes).map(Some)
}
fn unlink(&mut self) {
if self.reader.take().is_some() {
record_closed_file();
}
let _ = std::fs::remove_file(&self.path);
}
}
impl Drop for RunCursor {
fn drop(&mut self) {
if self.reader.is_some() {
record_closed_file();
}
}
}
fn write_run<Record: SpillRecord>(
path: &Path,
records: impl Iterator<Item = Result<Record>>,
) -> Result<()> {
let file = std::fs::OpenOptions::new()
.write(true)
.create_new(true)
.open(path)
.map_err(|error| {
if error.kind() == std::io::ErrorKind::AlreadyExists {
Error::InvalidFormat {
details: format!(
"the spill file {} already exists; a run directory must not be \
shared with an earlier run",
path.display()
),
}
} else {
Error::InputOutput(error)
}
})?;
let mut writer = BufWriter::with_capacity(SPILL_BUFFER_BYTES, file);
let mut buffer = Vec::new();
for record in records {
buffer.clear();
record?.encode(&mut buffer);
let length = u32::try_from(buffer.len()).map_err(|_| Error::InvalidFormat {
details: format!(
"a spill record of {} bytes exceeds the format",
buffer.len()
),
})?;
writer.write_all(&length.to_le_bytes())?;
writer.write_all(&buffer)?;
}
let mut file = writer
.into_inner()
.map_err(|error| Error::InputOutput(std::io::Error::other(error.to_string())))?;
file.flush()?;
file.sync_all()?;
let _ = file.seek(SeekFrom::Start(0));
Ok(())
}
#[cfg(test)]
thread_local! {
static OPEN_RUN_FILES: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
static PEAK_OPEN_RUN_FILES: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
static MERGE_PASSES: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
}
#[cfg(test)]
fn record_open_files(count: usize) {
OPEN_RUN_FILES.with(|open| {
let total = open.get() + count;
open.set(total);
PEAK_OPEN_RUN_FILES.with(|peak| peak.set(peak.get().max(total)));
});
}
#[cfg(not(test))]
fn record_open_files(_count: usize) {}
#[cfg(test)]
fn record_closed_file() {
OPEN_RUN_FILES.with(|open| open.set(open.get().saturating_sub(1)));
}
#[cfg(not(test))]
fn record_closed_file() {}
#[cfg(test)]
fn record_merge_pass() {
MERGE_PASSES.with(|passes| passes.set(passes.get() + 1));
}
#[cfg(not(test))]
fn record_merge_pass() {}
#[cfg(test)]
pub(crate) fn reset_sort_accounting() {
OPEN_RUN_FILES.with(|open| open.set(0));
PEAK_OPEN_RUN_FILES.with(|peak| peak.set(0));
MERGE_PASSES.with(|passes| passes.set(0));
}
#[cfg(test)]
pub(crate) fn peak_open_run_files() -> usize {
PEAK_OPEN_RUN_FILES.with(std::cell::Cell::get)
}
#[cfg(test)]
pub(crate) fn merge_passes() -> usize {
MERGE_PASSES.with(std::cell::Cell::get)
}
#[cfg(test)]
mod tests;