use std::{
hint::spin_loop,
sync::{
Arc,
atomic::{AtomicBool, AtomicI32, AtomicI64, Ordering},
},
time::{Duration, Instant},
};
use parking_lot::{Condvar, Mutex};
struct Round {
target: AtomicI64,
remaining: AtomicI32,
released: AtomicBool,
release_pair: Mutex<bool>,
release_signal: Condvar,
}
impl Round {
fn new(target: i64, participant_count: i32) -> Self {
Self {
target: AtomicI64::new(target),
remaining: AtomicI32::new(participant_count),
released: AtomicBool::new(false),
release_pair: Mutex::new(false),
release_signal: Condvar::new(),
}
}
fn release(&self) {
self.released.store(true, Ordering::Release);
*self.release_pair.lock() = true;
self.release_signal.notify_all();
}
}
pub struct ReplayAlignBarrier {
participant_count: i32,
last_arrived_round: Vec<Mutex<Option<Arc<Round>>>>,
current_round: Mutex<Option<Arc<Round>>>,
replica_sync_timeout: Option<Duration>,
}
impl ReplayAlignBarrier {
pub fn new(participant_count: usize, replica_sync_timeout: Option<Duration>) -> Self {
Self {
participant_count: participant_count as i32,
last_arrived_round: (0..participant_count).map(|_| Mutex::new(None)).collect(),
current_round: Mutex::new(None),
replica_sync_timeout,
}
}
pub fn in_progress(&self) -> bool {
self.current_round.lock().is_some()
}
pub fn try_open_round(&self, target: i64) {
let mut current = self.current_round.lock();
if current.is_some() {
return;
}
*current = Some(Arc::new(Round::new(target, self.participant_count)));
}
pub fn signal_arrival_and_wait(&self, virtual_sublog_idx: usize, frontier: i64) {
let Some(round) = self.arrive(virtual_sublog_idx, frontier) else {
return;
};
self.wait_for_all_arrivals(&round, virtual_sublog_idx);
}
pub fn signal_arrival(&self, virtual_sublog_idx: usize, frontier: i64) {
if let Some(round) = self.arrive(virtual_sublog_idx, frontier) {
if round.remaining.load(Ordering::Acquire) > 0 {
return;
}
self.release_round(&round);
}
}
fn arrive(&self, virtual_sublog_idx: usize, frontier: i64) -> Option<Arc<Round>> {
let round = Arc::clone(self.current_round.lock().as_ref()?);
if frontier < round.target.load(Ordering::Acquire) {
return None;
}
let mut last = self.last_arrived_round[virtual_sublog_idx].lock();
if last.as_ref().is_some_and(|r| Arc::ptr_eq(r, &round)) {
return None;
}
*last = Some(Arc::clone(&round));
round.remaining.fetch_sub(1, Ordering::AcqRel);
Some(round)
}
fn wait_for_all_arrivals(&self, round: &Arc<Round>, _virtual_sublog_idx: usize) {
if round.remaining.load(Ordering::Acquire) <= 0 {
self.release_round(round);
return;
}
let spin_deadline = Instant::now() + Duration::from_micros(64);
while Instant::now() < spin_deadline {
if round.released.load(Ordering::Acquire) {
return;
}
spin_loop();
}
let deadline = self.replica_sync_timeout.map(|t| Instant::now() + t);
let mut released = round.release_pair.lock();
while !round.released.load(Ordering::Acquire) {
let remain = match deadline.and_then(|d| d.checked_duration_since(Instant::now())) {
Some(remain) => remain.min(Duration::from_millis(100)),
None => Duration::from_millis(100),
};
round.release_signal.wait_for(&mut released, remain);
if deadline.is_some_and(|d| Instant::now() >= d) {
return;
}
}
}
fn release_round(&self, round: &Arc<Round>) {
round.release();
let mut current = self.current_round.lock();
if current.as_ref().is_some_and(|r| Arc::ptr_eq(r, round)) {
*current = None;
}
}
fn signal_all(&self, round: &Arc<Round>) {
round.release();
}
pub fn disable(&self) {
let mut current = self.current_round.lock();
if let Some(round) = current.take() {
self.signal_all(&round);
}
let inert = Arc::new(Round::new(i64::MAX, i32::MAX));
inert.release();
*current = Some(inert);
}
pub fn enable(&self) {
let mut current = self.current_round.lock();
if let Some(round) = current.take() {
self.signal_all(&round);
}
}
}
#[cfg(test)]
mod tests {
use std::thread;
use super::*;
#[test]
fn open_round_and_release_on_all_arrivals() {
let barrier = ReplayAlignBarrier::new(2, Some(Duration::from_secs(1)));
assert!(!barrier.in_progress());
barrier.try_open_round(100);
assert!(barrier.in_progress());
barrier.try_open_round(200);
assert!(barrier.in_progress());
barrier.signal_arrival(0, 50);
barrier.signal_arrival(0, 150);
assert!(barrier.in_progress());
barrier.signal_arrival(1, 150);
assert!(!barrier.in_progress());
}
#[test]
fn blocking_arrival_released_by_peer() {
let barrier = Arc::new(ReplayAlignBarrier::new(2, Some(Duration::from_secs(5))));
barrier.try_open_round(10);
let b2 = Arc::clone(&barrier);
let handle = thread::spawn(move || b2.signal_arrival_and_wait(1, 12));
thread::sleep(Duration::from_millis(20));
barrier.signal_arrival_and_wait(0, 12);
handle.join().unwrap();
assert!(!barrier.in_progress());
}
#[test]
fn timeout_proceeds_unaligned() {
let barrier = ReplayAlignBarrier::new(2, Some(Duration::from_millis(30)));
barrier.try_open_round(10);
let started = Instant::now();
barrier.signal_arrival_and_wait(0, 12);
assert!(started.elapsed() >= Duration::from_millis(30));
assert!(barrier.in_progress(), "轮次未被放行者摘除(仅超时退出)");
}
#[test]
fn disable_rejects_and_enable_restores() {
let barrier = ReplayAlignBarrier::new(2, None);
barrier.try_open_round(10);
barrier.disable();
assert!(barrier.in_progress());
barrier.try_open_round(20);
barrier.enable();
assert!(!barrier.in_progress());
barrier.try_open_round(30);
barrier.signal_arrival(0, 40);
barrier.signal_arrival(1, 40);
assert!(!barrier.in_progress());
}
}