use std::collections::VecDeque;
use std::io::{self};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Condvar, Mutex, MutexGuard};
use std::time::{Duration, Instant};
pub const DEFAULT_DRAIN_TIMEOUT: Duration = Duration::from_micros(200);
pub const DEFAULT_LEADER_TIMEOUT: Duration = Duration::from_millis(50);
pub trait WalLike: Send + Sync {
fn flush_to_disk(&self) -> io::Result<()>;
}
impl WalLike for std::sync::Mutex<crate::wal::WAL> {
fn flush_to_disk(&self) -> io::Result<()> {
let mut wal = self
.lock()
.map_err(|e| io::Error::other(format!("WAL lock poisoned: {e}")))?;
wal.flush_to_disk()
}
}
#[derive(Debug, Clone, Copy)]
pub struct GroupCommitConfig {
pub enabled: bool,
pub drain_timeout: Duration,
pub leader_timeout: Duration,
}
impl Default for GroupCommitConfig {
fn default() -> Self {
Self {
enabled: true,
drain_timeout: DEFAULT_DRAIN_TIMEOUT,
leader_timeout: DEFAULT_LEADER_TIMEOUT,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct GroupCommitResult {
pub lsn: u64,
pub durable_batch_lsn: u64,
pub group_size: usize,
}
#[derive(Debug, Default, Clone, Copy)]
pub struct GroupCommitStats {
pub batches: u64,
pub commits: u64,
pub total_group_entries: u64,
}
impl GroupCommitStats {
pub fn avg_group_size(&self) -> f64 {
if self.batches == 0 {
0.0
} else {
self.total_group_entries as f64 / self.batches as f64
}
}
}
struct PendingCommit {
lsn: u64,
slot: Arc<(Mutex<Option<io::Result<GroupCommitResult>>>, Condvar)>,
}
fn new_slot() -> Arc<(Mutex<Option<io::Result<GroupCommitResult>>>, Condvar)> {
Arc::new((Mutex::new(None), Condvar::new()))
}
pub struct GroupCommit<W: WalLike> {
wal: Arc<W>,
config: GroupCommitConfig,
next_lsn: AtomicU64,
queue: Mutex<VecDeque<PendingCommit>>,
leader_lock: Mutex<()>,
stats: Mutex<GroupCommitStats>,
}
impl<W: WalLike> GroupCommit<W> {
pub fn new(wal: Arc<W>, config: GroupCommitConfig) -> Self {
Self {
wal,
config,
next_lsn: AtomicU64::new(0),
queue: Mutex::new(VecDeque::new()),
leader_lock: Mutex::new(()),
stats: Mutex::new(GroupCommitStats::default()),
}
}
pub fn config(&self) -> &GroupCommitConfig {
&self.config
}
pub fn stats(&self) -> GroupCommitStats {
*self.stats.lock().unwrap()
}
pub fn flush(&self) -> io::Result<GroupCommitResult> {
let lsn = self.next_lsn.fetch_add(1, Ordering::Relaxed) + 1;
let slot = new_slot();
{
let mut queue = self.queue.lock().unwrap();
queue.push_back(PendingCommit {
lsn,
slot: slot.clone(),
});
}
if let Ok(guard) = self.leader_lock.try_lock() {
self.run_leader(guard);
}
let deadline = Instant::now() + self.config.leader_timeout;
let (lock, cvar) = &*slot;
let mut result = lock.lock().unwrap();
loop {
if let Some(outcome) = result.take() {
return match outcome {
Ok(ok) => Ok(ok),
Err(e) => Err(io::Error::new(e.kind(), e.to_string())),
};
}
if Instant::now() >= deadline {
drop(result);
if let Ok(guard) = self.leader_lock.try_lock() {
self.run_leader(guard);
}
result = lock.lock().unwrap();
continue;
}
let (guard, _to) = cvar.wait_timeout(result, Duration::from_millis(1)).unwrap();
result = guard;
}
}
fn run_leader(&self, _guard: MutexGuard<'_, ()>) {
let deadline = Instant::now() + self.config.drain_timeout;
let mut batch: Vec<PendingCommit> = Vec::new();
loop {
let mut drained = {
let mut queue = self.queue.lock().unwrap();
let mut v = Vec::with_capacity(queue.len());
while let Some(entry) = queue.pop_front() {
v.push(entry);
}
drop(queue);
v
};
if !drained.is_empty() {
batch.append(&mut drained);
}
if batch.is_empty() {
return;
}
if Instant::now() >= deadline {
break;
}
std::thread::sleep(Duration::from_micros(100));
}
let group_size = batch.len();
let durable_batch_lsn = batch.iter().map(|e| e.lsn).max().unwrap_or(0);
let io_result = self.wal.flush_to_disk();
{
let mut stats = self.stats.lock().unwrap();
stats.batches += 1;
stats.commits += group_size as u64;
stats.total_group_entries += group_size as u64;
}
for entry in batch {
let outcome = match &io_result {
Ok(()) => Ok(GroupCommitResult {
lsn: entry.lsn,
durable_batch_lsn,
group_size,
}),
Err(e) => Err(io::Error::new(e.kind(), format!("group fsync failed: {e}"))),
};
let (lock, cvar) = &*entry.slot;
let mut slot = lock.lock().unwrap();
*slot = Some(outcome);
drop(slot);
cvar.notify_all();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicU64;
use std::sync::atomic::Ordering as AtomicOrdering;
use std::thread;
struct MockWal {
fsyncs: AtomicU64,
delay: Duration,
}
impl WalLike for MockWal {
fn flush_to_disk(&self) -> io::Result<()> {
if !self.delay.is_zero() {
thread::sleep(self.delay);
}
self.fsyncs.fetch_add(1, AtomicOrdering::SeqCst);
Ok(())
}
}
fn make_gc(drain: Duration) -> (Arc<GroupCommit<MockWal>>, Arc<MockWal>) {
let wal = Arc::new(MockWal {
fsyncs: AtomicU64::new(0),
delay: Duration::from_millis(2),
});
let gc = Arc::new(GroupCommit::new(
wal.clone(),
GroupCommitConfig {
enabled: true,
drain_timeout: drain,
leader_timeout: Duration::from_millis(100),
},
));
(gc, wal)
}
#[test]
fn test_single_commit_flushes_once() {
let (gc, wal) = make_gc(Duration::from_millis(2));
let result = gc.flush().unwrap();
assert_eq!(result.lsn, 1);
assert_eq!(result.durable_batch_lsn, 1);
assert_eq!(result.group_size, 1);
assert_eq!(wal.fsyncs.load(AtomicOrdering::SeqCst), 1);
let stats = gc.stats();
assert_eq!(stats.batches, 1);
assert_eq!(stats.commits, 1);
}
#[test]
fn test_concurrent_commits_coalesce_into_one_fsync() {
let (gc, wal) = make_gc(Duration::from_millis(10));
let threads: Vec<_> = (0..12)
.map(|_| {
let gc = gc.clone();
thread::spawn(move || gc.flush().unwrap())
})
.collect();
let results: Vec<GroupCommitResult> = threads.into_iter().map(|t| t.join().unwrap()).collect();
assert_eq!(results.len(), 12);
let mut sorted: Vec<u64> = results.iter().map(|r| r.lsn).collect();
sorted.sort_unstable();
assert_eq!(sorted, (1..=12).collect::<Vec<u64>>());
for r in &results {
assert!(r.durable_batch_lsn >= r.lsn);
}
assert!(results.iter().any(|r| r.group_size > 1), "expected coalescing");
let fsyncs = wal.fsyncs.load(AtomicOrdering::SeqCst);
assert!(fsyncs < 12, "group commit should batch fsyncs, got {fsyncs}");
assert_eq!(fsyncs, gc.stats().batches);
assert_eq!(gc.stats().commits, 12);
assert!(gc.stats().avg_group_size() > 1.0);
}
#[test]
fn test_concurrent_commits_no_coalescing_guarantee_needed() {
let (gc, wal) = make_gc(Duration::from_micros(0));
let threads: Vec<_> = (0..8)
.map(|_| {
let gc = gc.clone();
thread::spawn(move || gc.flush().unwrap())
})
.collect();
let results: Vec<GroupCommitResult> = threads.into_iter().map(|t| t.join().unwrap()).collect();
assert_eq!(results.len(), 8);
let mut sorted: Vec<u64> = results.iter().map(|r| r.lsn).collect();
sorted.sort_unstable();
assert_eq!(sorted, (1..=8).collect::<Vec<u64>>());
assert!(wal.fsyncs.load(AtomicOrdering::SeqCst) <= 8);
}
#[test]
fn test_stale_follower_self_heals() {
let (gc, wal) = make_gc(Duration::from_millis(1));
let garbage: Arc<(Mutex<Option<io::Result<GroupCommitResult>>>, Condvar)> = new_slot();
{
let mut queue = gc.queue.lock().unwrap();
queue.push_back(PendingCommit {
lsn: 999,
slot: garbage,
});
}
let result = gc.flush().unwrap();
assert!(result.lsn >= 1);
assert!(wal.fsyncs.load(AtomicOrdering::SeqCst) >= 1);
}
#[test]
fn test_fsync_error_propagates() {
struct FailingWal;
impl WalLike for FailingWal {
fn flush_to_disk(&self) -> io::Result<()> {
Err(io::Error::other("disk on fire"))
}
}
let gc = Arc::new(GroupCommit::new(Arc::new(FailingWal), GroupCommitConfig::default()));
let err = gc.flush().unwrap_err();
assert!(err.to_string().contains("disk on fire"));
}
}