use asupersync::runtime::{BlockingTaskHandle, Runtime, RuntimeBuilder};
use fsqlite_types::glossary::TxnId;
use parking_lot::{Condvar, Mutex};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::time::Duration;
#[derive(Debug)]
pub struct CommitWaiter {
pub txn_id: TxnId,
pub ready: Arc<AtomicBool>,
pub epoch_at_submit: u64,
notifier: Arc<WaiterNotifier>,
}
#[derive(Debug, Default)]
struct WaiterNotifier {
lock: Mutex<()>,
cv: Condvar,
}
#[derive(Debug)]
struct PendingEntry {
ready: Arc<AtomicBool>,
notifier: Arc<WaiterNotifier>,
epoch_at_submit: u64,
}
type FlushFn = Box<dyn Fn() + Send + Sync + 'static>;
pub struct EpochGroupCommit {
current_epoch: AtomicU64,
epoch_duration_us: u64,
pending: Mutex<Vec<PendingEntry>>,
shutdown: Arc<AtomicBool>,
advancer_wake: Arc<WaiterNotifier>,
flush: Arc<FlushFn>,
advancer_runtime: Runtime,
advancer: Mutex<Option<BlockingTaskHandle>>,
}
impl std::fmt::Debug for EpochGroupCommit {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EpochGroupCommit")
.field("current_epoch", &self.current_epoch.load(Ordering::Relaxed))
.field("epoch_duration_us", &self.epoch_duration_us)
.field(
"pending_count",
&self.pending.try_lock().map_or(usize::MAX, |p| p.len()),
)
.finish_non_exhaustive()
}
}
impl EpochGroupCommit {
#[must_use]
pub fn new(epoch_duration_us: u64) -> Arc<Self> {
Self::new_with_flush(epoch_duration_us, Box::new(|| {}))
}
#[must_use]
pub fn new_with_flush(epoch_duration_us: u64, flush: FlushFn) -> Arc<Self> {
let advancer_runtime = RuntimeBuilder::new()
.worker_threads(0)
.blocking_threads(1, 1)
.thread_name_prefix("silo-epoch")
.build()
.expect("silo epoch advancer runtime");
let this = Arc::new(Self {
current_epoch: AtomicU64::new(1),
epoch_duration_us,
pending: Mutex::new(Vec::new()),
shutdown: Arc::new(AtomicBool::new(false)),
advancer_wake: Arc::new(WaiterNotifier::default()),
flush: Arc::new(flush),
advancer_runtime,
advancer: Mutex::new(None),
});
let weak = Arc::downgrade(&this);
let wake = Arc::clone(&this.advancer_wake);
let shutdown = Arc::clone(&this.shutdown);
let advancer = this
.advancer_runtime
.spawn_blocking(move || {
let epoch_window = Duration::from_micros(epoch_duration_us);
loop {
if shutdown.load(Ordering::Acquire) {
return;
}
{
let mut guard = wake.lock.lock();
let _ = wake.cv.wait_for(&mut guard, epoch_window);
}
if shutdown.load(Ordering::Acquire) {
return;
}
let Some(state) = weak.upgrade() else {
return;
};
if state.shutdown.load(Ordering::Acquire) {
return;
}
state.advance_epoch();
}
})
.expect("silo epoch advancer runtime must configure a blocking pool");
*this.advancer.lock() = Some(advancer);
this
}
pub fn submit(&self, txn_id: TxnId) -> CommitWaiter {
let ready = Arc::new(AtomicBool::new(false));
let notifier = Arc::new(WaiterNotifier::default());
let mut pending = self.pending.lock();
let epoch_at_submit = self.current_epoch.load(Ordering::Acquire);
pending.push(PendingEntry {
ready: Arc::clone(&ready),
notifier: Arc::clone(¬ifier),
epoch_at_submit,
});
drop(pending);
CommitWaiter {
txn_id,
ready,
epoch_at_submit,
notifier,
}
}
pub fn advance_epoch(&self) {
let drained: Vec<PendingEntry> = {
let mut pending = self.pending.lock();
let drained = std::mem::take(&mut *pending);
self.current_epoch.fetch_add(1, Ordering::AcqRel);
drained
};
(self.flush)();
for entry in drained {
entry.ready.store(true, Ordering::Release);
let guard = entry.notifier.lock.lock();
entry.notifier.cv.notify_all();
drop(guard);
let _ = entry.epoch_at_submit;
}
}
pub fn wait_for_commit(&self, waiter: &CommitWaiter) {
if waiter.ready.load(Ordering::Acquire) {
return;
}
let mut guard = waiter.notifier.lock.lock();
while !waiter.ready.load(Ordering::Acquire) {
waiter.notifier.cv.wait(&mut guard);
}
}
pub fn wait_for_commit_timeout(&self, waiter: &CommitWaiter, timeout: Duration) -> bool {
if waiter.ready.load(Ordering::Acquire) {
return true;
}
let mut guard = waiter.notifier.lock.lock();
if waiter.ready.load(Ordering::Acquire) {
return true;
}
let result = waiter.notifier.cv.wait_for(&mut guard, timeout);
if result.timed_out() {
waiter.ready.load(Ordering::Acquire)
} else {
true
}
}
#[must_use]
pub fn current_epoch(&self) -> u64 {
self.current_epoch.load(Ordering::Acquire)
}
}
impl Drop for EpochGroupCommit {
fn drop(&mut self) {
self.shutdown.store(true, Ordering::Release);
{
let guard = self.advancer_wake.lock.lock();
self.advancer_wake.cv.notify_all();
drop(guard);
}
let drained: Vec<PendingEntry> = std::mem::take(&mut *self.pending.lock());
for entry in drained {
entry.ready.store(true, Ordering::Release);
let guard = entry.notifier.lock.lock();
entry.notifier.cv.notify_all();
drop(guard);
}
let handle = self.advancer.lock().take();
if let Some(h) = handle {
h.cancel();
h.wait();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicU64 as StdAtomicU64;
use std::thread;
use std::time::Instant;
fn txn(id: u64) -> TxnId {
TxnId::new(id).expect("test txn id")
}
#[test]
fn advance_epoch_resolves_all_pending() {
let gc = EpochGroupCommit::new(10_000_000);
let waiters: Vec<CommitWaiter> = (1..=100).map(|i| gc.submit(txn(i))).collect();
for w in &waiters {
assert!(!w.ready.load(Ordering::Acquire));
}
gc.advance_epoch();
for w in &waiters {
gc.wait_for_commit(w);
assert!(w.ready.load(Ordering::Acquire));
}
}
#[test]
fn waiter_blocks_without_advance() {
let gc = EpochGroupCommit::new(10_000_000);
let w = gc.submit(txn(1));
let start = Instant::now();
let resolved = gc.wait_for_commit_timeout(&w, Duration::from_micros(100));
let elapsed = start.elapsed();
assert!(
!resolved,
"waiter should not have been resolved without advance_epoch; elapsed={elapsed:?}"
);
assert!(!w.ready.load(Ordering::Acquire));
}
#[test]
fn multi_threaded_submit_and_advance() {
let gc = EpochGroupCommit::new(10_000_000); let per_thread = 200_u64;
let handles: Vec<_> = (0..4_u64)
.map(|tid| {
let gc = Arc::clone(&gc);
thread::spawn(move || {
let mut waiters = Vec::with_capacity(per_thread as usize);
for i in 0..per_thread {
let id = tid * 10_000 + i + 1;
waiters.push(gc.submit(txn(id)));
}
waiters
})
})
.collect();
let resolved_total = Arc::new(StdAtomicU64::new(0));
let expected = 4 * per_thread;
let advancer = {
let gc = Arc::clone(&gc);
let resolved_total = Arc::clone(&resolved_total);
thread::spawn(move || {
while resolved_total.load(Ordering::Acquire) < expected {
gc.advance_epoch();
thread::sleep(Duration::from_micros(50));
}
})
};
for h in handles {
let waiters = h.join().expect("submitter thread");
for w in &waiters {
gc.wait_for_commit(w);
assert!(w.ready.load(Ordering::Acquire));
resolved_total.fetch_add(1, Ordering::AcqRel);
}
}
advancer.join().expect("advancer thread");
assert_eq!(resolved_total.load(Ordering::Acquire), expected);
}
}