use super::wal::{Lsn, TxnId, WriteAheadLog};
use crate::error::Result;
use std::sync::{Arc, Condvar, Mutex};
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct GroupCommitConfig {
pub max_batch_size: usize,
pub max_wait_time: Duration,
pub enabled: bool,
}
impl Default for GroupCommitConfig {
fn default() -> Self {
Self {
max_batch_size: 100,
max_wait_time: Duration::from_millis(10),
enabled: true,
}
}
}
#[derive(Debug)]
struct PendingCommit {
txn_id: TxnId,
commit_lsn: Lsn,
requested_at: Instant,
}
pub struct GroupCommitCoordinator {
config: GroupCommitConfig,
wal: Arc<WriteAheadLog>,
pending: Arc<Mutex<PendingCommitQueue>>,
commit_cv: Arc<Condvar>,
stats: Arc<Mutex<GroupCommitStats>>,
}
struct PendingCommitQueue {
commits: Vec<PendingCommit>,
last_flush: Instant,
last_flushed_lsn: Lsn,
}
impl PendingCommitQueue {
fn new() -> Self {
Self {
commits: Vec::new(),
last_flush: Instant::now(),
last_flushed_lsn: Lsn::ZERO,
}
}
fn is_empty(&self) -> bool {
self.commits.is_empty()
}
fn len(&self) -> usize {
self.commits.len()
}
fn should_flush(&self, config: &GroupCommitConfig) -> bool {
if self.commits.is_empty() {
return false;
}
if self.commits.len() >= config.max_batch_size {
return true;
}
if let Some(oldest) = self.commits.first() {
if oldest.requested_at.elapsed() >= config.max_wait_time {
return true;
}
}
false
}
fn drain_batch(&mut self) -> Vec<PendingCommit> {
std::mem::take(&mut self.commits)
}
}
#[derive(Debug, Default)]
pub struct GroupCommitStats {
pub total_commits: u64,
pub total_flushes: u64,
pub avg_batch_size: f64,
pub total_wait_time_us: u64,
pub max_wait_time_us: u64,
}
impl GroupCommitCoordinator {
pub fn new(wal: Arc<WriteAheadLog>, config: GroupCommitConfig) -> Self {
Self {
config,
wal,
pending: Arc::new(Mutex::new(PendingCommitQueue::new())),
commit_cv: Arc::new(Condvar::new()),
stats: Arc::new(Mutex::new(GroupCommitStats::default())),
}
}
pub fn commit(&self, txn_id: TxnId, commit_lsn: Lsn) -> Result<()> {
let requested_at = Instant::now();
if !self.config.enabled {
self.wal.flush()?;
let wait_time_us = requested_at.elapsed().as_micros() as u64;
let mut stats = self.stats.lock().expect("lock poisoned");
stats.total_commits += 1;
stats.total_flushes += 1; stats.total_wait_time_us += wait_time_us;
stats.max_wait_time_us = stats.max_wait_time_us.max(wait_time_us);
return Ok(());
}
{
let mut pending = self.pending.lock().expect("lock poisoned");
pending.commits.push(PendingCommit {
txn_id,
commit_lsn,
requested_at,
});
}
self.try_flush()?;
let timeout = self.config.max_wait_time * 2; self.wait_for_flush(commit_lsn, timeout)?;
let wait_time_us = requested_at.elapsed().as_micros() as u64;
let mut stats = self.stats.lock().expect("lock poisoned");
stats.total_commits += 1;
stats.total_wait_time_us += wait_time_us;
stats.max_wait_time_us = stats.max_wait_time_us.max(wait_time_us);
Ok(())
}
fn try_flush(&self) -> Result<()> {
let mut pending = self.pending.lock().expect("lock poisoned");
if !pending.should_flush(&self.config) {
return Ok(());
}
let batch = pending.drain_batch();
let batch_size = batch.len();
pending.last_flush = Instant::now();
let max_lsn = batch
.iter()
.map(|c| c.commit_lsn)
.max()
.unwrap_or(Lsn::ZERO);
drop(pending);
self.wal.flush()?;
{
let mut pending = self.pending.lock().expect("lock poisoned");
pending.last_flushed_lsn = max_lsn;
}
self.commit_cv.notify_all();
{
let mut stats = self.stats.lock().expect("lock poisoned");
stats.total_flushes += 1;
let total_commits = stats.total_commits as f64;
stats.avg_batch_size = (stats.avg_batch_size * (total_commits - batch_size as f64)
+ batch_size as f64)
/ total_commits.max(1.0);
}
Ok(())
}
fn wait_for_flush(&self, target_lsn: Lsn, timeout: Duration) -> Result<()> {
let deadline = Instant::now() + timeout;
let mut pending = self.pending.lock().expect("lock poisoned");
loop {
if pending.last_flushed_lsn >= target_lsn {
return Ok(());
}
let now = Instant::now();
if now >= deadline {
drop(pending);
self.force_flush()?;
return Ok(());
}
let remaining = deadline.duration_since(now);
let (guard, timeout_result) = self
.commit_cv
.wait_timeout(pending, remaining)
.expect("lock poisoned");
pending = guard;
if timeout_result.timed_out() {
drop(pending);
self.force_flush()?;
return Ok(());
}
}
}
pub fn force_flush(&self) -> Result<()> {
let mut pending = self.pending.lock().expect("lock poisoned");
if pending.is_empty() {
return Ok(());
}
let batch = pending.drain_batch();
let batch_size = batch.len();
let max_lsn = batch
.iter()
.map(|c| c.commit_lsn)
.max()
.unwrap_or(Lsn::ZERO);
pending.last_flush = Instant::now();
drop(pending);
self.wal.flush()?;
{
let mut pending = self.pending.lock().expect("lock poisoned");
pending.last_flushed_lsn = max_lsn;
}
self.commit_cv.notify_all();
{
let mut stats = self.stats.lock().expect("lock poisoned");
stats.total_flushes += 1;
let total_commits = stats.total_commits as f64;
stats.avg_batch_size = (stats.avg_batch_size * (total_commits - batch_size as f64)
+ batch_size as f64)
/ total_commits.max(1.0);
}
Ok(())
}
pub fn stats(&self) -> GroupCommitStats {
let stats = self.stats.lock().expect("lock poisoned");
GroupCommitStats {
total_commits: stats.total_commits,
total_flushes: stats.total_flushes,
avg_batch_size: stats.avg_batch_size,
total_wait_time_us: stats.total_wait_time_us,
max_wait_time_us: stats.max_wait_time_us,
}
}
pub fn avg_commits_per_flush(&self) -> f64 {
let stats = self.stats.lock().expect("lock poisoned");
if stats.total_flushes == 0 {
0.0
} else {
stats.total_commits as f64 / stats.total_flushes as f64
}
}
pub fn pending_count(&self) -> usize {
self.pending.lock().expect("lock poisoned").len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::env;
use std::thread;
#[test]
fn test_group_commit_config() {
let config = GroupCommitConfig::default();
assert_eq!(config.max_batch_size, 100);
assert_eq!(config.max_wait_time, Duration::from_millis(10));
assert!(config.enabled);
}
#[test]
fn test_single_commit() {
let temp_dir = env::temp_dir().join("oxirs_group_commit_single");
std::fs::create_dir_all(&temp_dir).unwrap();
let wal = Arc::new(WriteAheadLog::new(&temp_dir).unwrap());
let coordinator = GroupCommitCoordinator::new(wal.clone(), GroupCommitConfig::default());
let lsn = wal
.append(super::super::wal::LogRecord::Commit {
txn_id: TxnId::new(1),
})
.unwrap();
coordinator.commit(TxnId::new(1), lsn).unwrap();
coordinator.force_flush().unwrap();
let stats = coordinator.stats();
assert_eq!(stats.total_commits, 1);
assert!(stats.total_flushes >= 1);
std::fs::remove_dir_all(&temp_dir).ok();
}
#[test]
fn test_batch_commit() {
let temp_dir = env::temp_dir().join("oxirs_group_commit_batch");
std::fs::create_dir_all(&temp_dir).unwrap();
let wal = Arc::new(WriteAheadLog::new(&temp_dir).unwrap());
let config = GroupCommitConfig {
max_batch_size: 5,
max_wait_time: Duration::from_millis(50), enabled: true,
};
let coordinator = Arc::new(GroupCommitCoordinator::new(wal.clone(), config));
let mut handles = vec![];
for i in 0..5 {
let coordinator = Arc::clone(&coordinator);
let wal = Arc::clone(&wal);
let handle = thread::spawn(move || {
let lsn = wal
.append(super::super::wal::LogRecord::Commit {
txn_id: TxnId::new(i),
})
.unwrap();
coordinator.commit(TxnId::new(i), lsn).unwrap();
});
handles.push(handle);
}
for handle in handles {
handle.join().unwrap();
}
let stats = coordinator.stats();
assert_eq!(stats.total_commits, 5);
assert!(stats.total_flushes <= 5);
assert!(stats.avg_batch_size >= 1.0);
std::fs::remove_dir_all(&temp_dir).ok();
}
#[test]
fn test_force_flush() {
let temp_dir = env::temp_dir().join("oxirs_group_commit_force");
std::fs::create_dir_all(&temp_dir).unwrap();
let wal = Arc::new(WriteAheadLog::new(&temp_dir).unwrap());
let config = GroupCommitConfig {
max_batch_size: 100,
max_wait_time: Duration::from_millis(50), enabled: false, };
let coordinator = Arc::new(GroupCommitCoordinator::new(wal.clone(), config));
for i in 0..3 {
let lsn = wal
.append(super::super::wal::LogRecord::Commit {
txn_id: TxnId::new(i),
})
.unwrap();
coordinator.commit(TxnId::new(i), lsn).unwrap();
}
let stats = coordinator.stats();
assert_eq!(stats.total_commits, 3);
assert!(stats.total_flushes >= 1);
std::fs::remove_dir_all(&temp_dir).ok();
}
#[test]
fn test_disabled_group_commit() {
let temp_dir = env::temp_dir().join("oxirs_group_commit_disabled");
std::fs::create_dir_all(&temp_dir).unwrap();
let wal = Arc::new(WriteAheadLog::new(&temp_dir).unwrap());
let config = GroupCommitConfig {
enabled: false,
..Default::default()
};
let coordinator = GroupCommitCoordinator::new(wal.clone(), config);
for i in 0..5 {
let lsn = wal
.append(super::super::wal::LogRecord::Commit {
txn_id: TxnId::new(i),
})
.unwrap();
coordinator.commit(TxnId::new(i), lsn).unwrap();
}
let stats = coordinator.stats();
assert_eq!(stats.total_commits, 5);
assert_eq!(stats.avg_batch_size, 0.0);
std::fs::remove_dir_all(&temp_dir).ok();
}
#[test]
#[ignore] fn test_timeout_flush() {
let temp_dir = env::temp_dir().join("oxirs_group_commit_timeout");
std::fs::create_dir_all(&temp_dir).unwrap();
let wal = Arc::new(WriteAheadLog::new(&temp_dir).unwrap());
let config = GroupCommitConfig {
max_batch_size: 100,
max_wait_time: Duration::from_millis(50), enabled: true,
};
let coordinator = Arc::new(GroupCommitCoordinator::new(wal.clone(), config));
let lsn = wal
.append(super::super::wal::LogRecord::Commit {
txn_id: TxnId::new(1),
})
.unwrap();
coordinator.commit(TxnId::new(1), lsn).unwrap();
let stats = coordinator.stats();
assert_eq!(stats.total_commits, 1);
assert!(stats.max_wait_time_us > 0);
std::fs::remove_dir_all(&temp_dir).ok();
}
#[test]
fn test_concurrent_commits() {
let temp_dir = env::temp_dir().join("oxirs_group_commit_concurrent");
std::fs::create_dir_all(&temp_dir).unwrap();
let wal = Arc::new(WriteAheadLog::new(&temp_dir).unwrap());
let config = GroupCommitConfig {
max_batch_size: 10,
max_wait_time: Duration::from_millis(50), enabled: true,
};
let coordinator = Arc::new(GroupCommitCoordinator::new(wal.clone(), config));
let mut handles = vec![];
for i in 0..5 {
let coordinator = Arc::clone(&coordinator);
let wal = Arc::clone(&wal);
let handle = thread::spawn(move || {
let lsn = wal
.append(super::super::wal::LogRecord::Commit {
txn_id: TxnId::new(i),
})
.unwrap();
coordinator.commit(TxnId::new(i), lsn).unwrap();
});
handles.push(handle);
}
for handle in handles {
handle.join().unwrap();
}
let stats = coordinator.stats();
assert_eq!(stats.total_commits, 5);
assert!(stats.total_flushes <= 5);
std::fs::remove_dir_all(&temp_dir).ok();
}
}