use glam::Vec3;
use std::{
collections::BTreeMap,
fmt,
sync::{Mutex, MutexGuard, PoisonError},
};
#[cfg(test)]
mod tests;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum MergeError {
WrongLength {
expected: usize,
got: usize,
},
Overlap {
first_sample: u32,
done: u32,
},
}
impl fmt::Display for MergeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::WrongLength { expected, got } => write!(
f,
"chunk buffer holds {got} pixels, the image has {expected}"
),
Self::Overlap { first_sample, done } => write!(
f,
"chunk [{first_sample}, +{done}) overlaps samples already merged"
),
}
}
}
impl std::error::Error for MergeError {}
#[derive(Debug)]
struct Parked {
done: u32,
sum: Vec<Vec3>,
}
#[derive(Debug)]
struct MergeState {
merged: Vec<Vec3>,
merged_count: u32,
frontier: u32,
parked: BTreeMap<u32, Parked>,
parked_count: u32,
}
#[derive(Debug)]
pub struct Merger {
pixels: usize,
first_sample: u32,
state: Mutex<MergeState>,
}
impl Merger {
#[must_use]
pub const fn new(pixels: usize, first_sample: u32) -> Self {
Self {
pixels,
first_sample,
state: Mutex::new(MergeState {
merged: Vec::new(),
merged_count: 0,
frontier: first_sample,
parked: BTreeMap::new(),
parked_count: 0,
}),
}
}
#[must_use]
pub const fn pixel_count(&self) -> usize {
self.pixels
}
#[must_use]
pub const fn first_sample(&self) -> u32 {
self.first_sample
}
fn lock(&self) -> MutexGuard<'_, MergeState> {
self.state.lock().unwrap_or_else(PoisonError::into_inner)
}
pub fn add(&self, first_sample: u32, done: u32, sum: Vec<Vec3>) -> Result<u32, MergeError> {
if done == 0 {
return Ok(self.total());
}
if sum.len() != self.pixels {
return Err(MergeError::WrongLength {
expected: self.pixels,
got: sum.len(),
});
}
let mut state = self.lock();
if overlaps(&state, first_sample, done) {
return Err(MergeError::Overlap { first_sample, done });
}
state.parked.insert(first_sample, Parked { done, sum });
state.parked_count += done;
advance_frontier(&mut state, self.pixels);
let total = state.merged_count + state.parked_count;
drop(state);
Ok(total)
}
#[must_use]
pub fn total(&self) -> u32 {
let state = self.lock();
state.merged_count + state.parked_count
}
pub fn snapshot_into(&self, dst: &mut [Vec3]) -> u32 {
if dst.len() != self.pixels {
return 0;
}
let state = self.lock();
add_assign(dst, &state.merged);
for parked in state.parked.values() {
add_assign(dst, &parked.sum);
}
let count = state.merged_count + state.parked_count;
drop(state);
count
}
#[must_use]
pub fn into_parts(self) -> (Vec<Vec3>, u32) {
let pixels = self.pixels;
let mut state = self
.state
.into_inner()
.unwrap_or_else(PoisonError::into_inner);
let parked = std::mem::take(&mut state.parked);
for chunk in parked.into_values() {
fold(&mut state.merged, pixels, &chunk.sum);
state.merged_count += chunk.done;
}
if state.merged.is_empty() {
state.merged = vec![Vec3::ZERO; pixels];
}
(state.merged, state.merged_count)
}
}
fn overlaps(state: &MergeState, first: u32, done: u32) -> bool {
let end = first + done;
if first < state.frontier {
return true;
}
let before = state.parked.range(..=first).next_back();
if before.is_some_and(|(&start, chunk)| start + chunk.done > first) {
return true;
}
state
.parked
.range(first..)
.next()
.is_some_and(|(&start, _)| start < end)
}
fn advance_frontier(state: &mut MergeState, pixels: usize) {
while let Some(chunk) = state.parked.remove(&state.frontier) {
fold(&mut state.merged, pixels, &chunk.sum);
state.frontier += chunk.done;
state.merged_count += chunk.done;
state.parked_count -= chunk.done;
}
}
fn fold(merged: &mut Vec<Vec3>, pixels: usize, sum: &[Vec3]) {
if merged.is_empty() {
*merged = vec![Vec3::ZERO; pixels];
}
add_assign(merged, sum);
}
fn add_assign(dst: &mut [Vec3], src: &[Vec3]) {
for (d, s) in dst.iter_mut().zip(src) {
*d += *s;
}
}