use std::{
hint::spin_loop,
sync::{
Arc,
atomic::{AtomicI64, Ordering},
},
time::{Duration, Instant},
};
use parking_lot::{Condvar, Mutex};
const SKETCH_SLOT_SIZE: usize = 1 << 15;
const SKETCH_SLOT_MASK: usize = SKETCH_SLOT_SIZE - 1;
const MAX_SPIN_COUNT: u32 = 64;
#[derive(Default)]
pub struct ReadSessionWaiter {
target: AtomicI64,
pair: Mutex<bool>,
signal: Condvar,
}
impl ReadSessionWaiter {
pub fn new() -> Self {
Self::default()
}
fn reset(&self, target: i64) {
self.target.store(target, Ordering::Release);
*self.pair.lock() = false;
}
fn wait(&self, timeout: Duration) -> bool {
let mut signaled = self.pair.lock();
let deadline = Instant::now() + timeout;
while !*signaled {
let Some(remain) = deadline.checked_duration_since(Instant::now()) else {
return false;
};
self.signal.wait_for(&mut signaled, remain);
}
true
}
}
pub struct VirtualSublogReplayState {
sketch: Box<[AtomicI64]>,
frontier: AtomicI64,
min_waiter_target: AtomicI64,
next_drift_check_window_lower_bound: AtomicI64,
waiters: Mutex<Vec<Arc<ReadSessionWaiter>>>,
}
impl VirtualSublogReplayState {
pub fn new(next_drift_check_seed: i64) -> Self {
Self {
sketch: (0..SKETCH_SLOT_SIZE).map(|_| AtomicI64::new(0)).collect(),
frontier: AtomicI64::new(0),
min_waiter_target: AtomicI64::new(i64::MAX),
next_drift_check_window_lower_bound: AtomicI64::new(next_drift_check_seed),
waiters: Mutex::new(Vec::new()),
}
}
#[inline]
pub fn max(&self) -> i64 {
self.frontier.load(Ordering::Acquire)
}
#[inline]
pub fn next_drift_check_window_lower_bound(&self) -> i64 {
self
.next_drift_check_window_lower_bound
.load(Ordering::Relaxed)
}
#[inline]
pub fn set_next_drift_check_window_lower_bound(&self, value: i64) {
self
.next_drift_check_window_lower_bound
.store(value, Ordering::Relaxed);
}
#[inline]
pub const fn get_sketch_slot(hash: i64) -> usize {
((hash as u64 >> 32) as usize) & SKETCH_SLOT_MASK
}
#[inline]
pub fn get_frontier_sequence_number(&self, hash: i64) -> i64 {
let slot = self.sketch[Self::get_sketch_slot(hash)].load(Ordering::Acquire);
slot.max(self.frontier.load(Ordering::Acquire))
}
#[inline]
pub fn get_key_sequence_number(&self, hash: i64) -> i64 {
self.sketch[Self::get_sketch_slot(hash)].load(Ordering::Acquire)
}
#[inline]
pub fn prefetch_key_sequence_number(&self, hash: i64) {
let _ = hash;
}
pub fn update_max_sequence_number(&self, sequence_number: i64) {
self.frontier.fetch_max(sequence_number, Ordering::AcqRel);
self.signal_waiters();
}
pub fn update_key_sequence_number(&self, hash: i64, sequence_number: i64) {
let slot = &self.sketch[Self::get_sketch_slot(hash)];
slot.fetch_max(sequence_number, Ordering::AcqRel);
self.signal_waiters();
}
fn signal_waiters(&self) {
if self.frontier.load(Ordering::Acquire) <= self.min_waiter_target.load(Ordering::Acquire) {
return;
}
let mut waiters = self.waiters.lock();
let current_max = self.frontier.load(Ordering::Acquire);
let mut satisfied = 0;
while satisfied < waiters.len()
&& waiters[satisfied].target.load(Ordering::Acquire) < current_max
{
let node = &waiters[satisfied];
*node.pair.lock() = true;
node.signal.notify_all();
satisfied += 1;
}
if satisfied > 0 {
waiters.drain(..satisfied);
self.update_min_waiter_target_locked(&waiters);
}
}
pub fn wait_for_sequence_number(
&self,
maximum_session_sequence_number: i64,
waiter: &Arc<ReadSessionWaiter>,
timeout: Duration,
) -> bool {
for _ in 0..MAX_SPIN_COUNT {
if maximum_session_sequence_number < self.frontier.load(Ordering::Acquire) {
return true;
}
spin_loop();
}
waiter.reset(maximum_session_sequence_number);
{
let mut waiters = self.waiters.lock();
if maximum_session_sequence_number < self.frontier.load(Ordering::Acquire) {
return true;
}
self.insert_waiter(&mut waiters, Arc::clone(waiter));
self.update_min_waiter_target_locked(&waiters);
if maximum_session_sequence_number < self.frontier.load(Ordering::Acquire) {
self.remove_waiter_locked(&mut waiters, waiter);
self.update_min_waiter_target_locked(&waiters);
return true;
}
}
waiter.wait(timeout)
}
fn insert_waiter(&self, waiters: &mut Vec<Arc<ReadSessionWaiter>>, node: Arc<ReadSessionWaiter>) {
let target = node.target.load(Ordering::Acquire);
let idx = waiters.partition_point(|w| w.target.load(Ordering::Acquire) <= target);
waiters.insert(idx, node);
}
fn remove_waiter_locked(
&self,
waiters: &mut Vec<Arc<ReadSessionWaiter>>,
node: &Arc<ReadSessionWaiter>,
) {
if let Some(idx) = waiters.iter().position(|w| Arc::ptr_eq(w, node)) {
waiters.remove(idx);
}
}
pub fn remove_waiter(&self, node: &Arc<ReadSessionWaiter>) {
let mut waiters = self.waiters.lock();
self.remove_waiter_locked(&mut waiters, node);
self.update_min_waiter_target_locked(&waiters);
}
fn update_min_waiter_target_locked(&self, waiters: &[Arc<ReadSessionWaiter>]) {
let min = waiters
.first()
.map_or(i64::MAX, |w| w.target.load(Ordering::Acquire));
self.min_waiter_target.store(min, Ordering::Release);
}
pub fn update_min_waiter_target(&self) {
let waiters = self.waiters.lock();
self.update_min_waiter_target_locked(&waiters);
}
}
#[cfg(test)]
mod tests {
use std::thread;
use super::*;
#[test]
fn sketch_slot_is_high_word_masked() {
let hash = 0x1234_5678_9abc_def0_i64;
assert_eq!(
VirtualSublogReplayState::get_sketch_slot(hash),
(hash as u64 >> 32) as usize & SKETCH_SLOT_MASK
);
assert!(VirtualSublogReplayState::get_sketch_slot(hash) < SKETCH_SLOT_SIZE);
}
#[test]
fn monotonic_updates_and_frontier() {
let state = VirtualSublogReplayState::new(i64::MAX);
state.update_max_sequence_number(10);
state.update_max_sequence_number(5);
assert_eq!(state.max(), 10);
state.update_key_sequence_number(0x1234, 7);
assert_eq!(state.get_key_sequence_number(0x1234), 7);
assert_eq!(state.get_frontier_sequence_number(0x1234), 10);
assert_eq!(state.get_frontier_sequence_number(0x9999), 10);
}
#[test]
fn waiter_signaled_when_frontier_passes_target() {
let state = VirtualSublogReplayState::new(i64::MAX);
let waiter = Arc::new(ReadSessionWaiter::new());
state.update_max_sequence_number(5);
assert!(state.wait_for_sequence_number(4, &waiter, Duration::from_millis(1)));
let state2 = Arc::new(VirtualSublogReplayState::new(i64::MAX));
let waiter2 = Arc::new(ReadSessionWaiter::new());
let state3 = Arc::clone(&state2);
let waiter3 = Arc::clone(&waiter2);
let handle =
thread::spawn(move || state3.wait_for_sequence_number(20, &waiter3, Duration::from_secs(5)));
thread::sleep(Duration::from_millis(10));
state2.update_max_sequence_number(21);
assert!(handle.join().unwrap());
}
#[test]
fn wait_times_out_when_replay_lags() {
let state = VirtualSublogReplayState::new(i64::MAX);
let waiter = Arc::new(ReadSessionWaiter::new());
assert!(!state.wait_for_sequence_number(100, &waiter, Duration::from_millis(20)));
state.remove_waiter(&waiter);
assert!(state.waiters.lock().is_empty());
}
#[test]
fn min_waiter_target_tracks_head() {
let state = VirtualSublogReplayState::new(i64::MAX);
assert_eq!(state.min_waiter_target.load(Ordering::Acquire), i64::MAX);
let w = Arc::new(ReadSessionWaiter::new());
w.reset(42);
{
let mut waiters = state.waiters.lock();
state.insert_waiter(&mut waiters, Arc::clone(&w));
state.update_min_waiter_target_locked(&waiters);
}
assert_eq!(state.min_waiter_target.load(Ordering::Acquire), 42);
state.update_max_sequence_number(43);
assert!(state.waiters.lock().is_empty());
assert_eq!(state.min_waiter_target.load(Ordering::Acquire), i64::MAX);
}
}