use parking_lot::{Condvar, Mutex};
use std::time::Duration;
pub struct GroupCommitter {
state: Mutex<State>,
cv: Condvar,
window: Duration,
}
struct State {
accumulating: u64,
completed: u64,
leader_present: bool,
flush_in_progress: bool,
failed: std::collections::BTreeMap<u64, String>,
}
const ERROR_RETENTION_GENS: u64 = 1024;
impl GroupCommitter {
pub fn new(window: Duration) -> Self {
Self {
state: Mutex::new(State {
accumulating: 1,
completed: 0,
leader_present: false,
flush_in_progress: false,
failed: std::collections::BTreeMap::new(),
}),
cv: Condvar::new(),
window,
}
}
pub fn window(&self) -> Duration {
self.window
}
pub fn wait_durable<F>(&self, flush: F) -> std::result::Result<(), String>
where
F: FnOnce() -> std::result::Result<(), String>,
{
let mut st = self.state.lock();
let my_gen = st.accumulating;
if st.leader_present {
while st.completed < my_gen {
self.cv.wait(&mut st);
}
return match st.failed.get(&my_gen) {
Some(msg) => Err(msg.clone()),
None => Ok(()),
};
}
st.leader_present = true;
if !self.window.is_zero() {
drop(st);
std::thread::sleep(self.window);
st = self.state.lock();
}
while st.flush_in_progress {
self.cv.wait(&mut st);
}
debug_assert_eq!(st.accumulating, my_gen);
st.accumulating = my_gen + 1;
st.leader_present = false;
st.flush_in_progress = true;
drop(st);
let result = flush();
let mut st = self.state.lock();
st.flush_in_progress = false;
st.completed = st.completed.max(my_gen);
if let Err(msg) = &result {
st.failed.insert(my_gen, msg.clone());
let cutoff = st.completed.saturating_sub(ERROR_RETENTION_GENS);
st.failed.retain(|gen, _| *gen > cutoff);
}
self.cv.notify_all();
result
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
#[test]
fn single_committer_flushes_once() {
let gc = GroupCommitter::new(Duration::ZERO);
let flushes = AtomicUsize::new(0);
gc.wait_durable(|| {
flushes.fetch_add(1, Ordering::SeqCst);
Ok(())
})
.unwrap();
assert_eq!(flushes.load(Ordering::SeqCst), 1);
}
#[test]
fn cohort_shares_one_flush() {
let gc = Arc::new(GroupCommitter::new(Duration::from_millis(50)));
let flushes = Arc::new(AtomicUsize::new(0));
let threads = 8;
let mut handles = Vec::new();
for _ in 0..threads {
let gc = Arc::clone(&gc);
let flushes = Arc::clone(&flushes);
handles.push(std::thread::spawn(move || {
gc.wait_durable(|| {
flushes.fetch_add(1, Ordering::SeqCst);
std::thread::sleep(Duration::from_millis(20));
Ok(())
})
.unwrap();
}));
}
for h in handles {
h.join().unwrap();
}
let n = flushes.load(Ordering::SeqCst);
assert!(
n >= 1 && n < threads,
"expected grouped flushes, got {n} for {threads} committers"
);
}
#[test]
fn flush_error_propagates_to_whole_generation() {
let gc = Arc::new(GroupCommitter::new(Duration::from_millis(30)));
let attempts = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..4 {
let gc = Arc::clone(&gc);
let attempts = Arc::clone(&attempts);
handles.push(std::thread::spawn(move || {
gc.wait_durable(|| {
attempts.fetch_add(1, Ordering::SeqCst);
std::thread::sleep(Duration::from_millis(10));
Err("disk on fire".to_string())
})
}));
}
let mut errs = 0;
for h in handles {
if h.join().unwrap().is_err() {
errs += 1;
}
}
assert_eq!(errs, 4);
}
#[test]
fn generations_advance_after_completion() {
let gc = GroupCommitter::new(Duration::ZERO);
for _ in 0..3 {
gc.wait_durable(|| Ok(())).unwrap();
}
let st = gc.state.lock();
assert_eq!(st.accumulating, 4);
assert_eq!(st.completed, 3);
assert!(!st.leader_present);
assert!(!st.flush_in_progress);
}
#[test]
fn error_does_not_poison_next_generation() {
let gc = GroupCommitter::new(Duration::ZERO);
assert!(gc.wait_durable(|| Err("boom".into())).is_err());
assert!(gc.wait_durable(|| Ok(())).is_ok());
}
#[test]
fn failures_are_recorded_per_generation() {
let gc = GroupCommitter::new(Duration::ZERO);
assert!(gc.wait_durable(|| Err("gen1 fsync lost".into())).is_err());
assert!(gc.wait_durable(|| Err("gen2 fsync lost".into())).is_err());
assert!(gc.wait_durable(|| Ok(())).is_ok());
let st = gc.state.lock();
assert_eq!(st.failed.get(&1).map(String::as_str), Some("gen1 fsync lost"));
assert_eq!(st.failed.get(&2).map(String::as_str), Some("gen2 fsync lost"));
assert!(!st.failed.contains_key(&3));
}
}