use std::sync::Arc;
use std::sync::atomic::AtomicU64;
use std::sync::atomic::Ordering;
use ursula_shard::RaftGroupId;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct RaftUncommittedAdmission {
pub max_uncommitted_bytes_per_group: Option<u64>,
}
impl RaftUncommittedAdmission {
pub fn is_enabled(self) -> bool {
self.max_uncommitted_bytes_per_group.is_some()
}
}
#[derive(Debug)]
pub(crate) struct RaftUncommittedBytesTracker {
per_group: Vec<AtomicU64>,
}
impl RaftUncommittedBytesTracker {
pub(crate) fn new(group_count: usize) -> Self {
Self {
per_group: (0..group_count).map(|_| AtomicU64::new(0)).collect(),
}
}
pub(crate) fn load(&self, group_id: RaftGroupId) -> u64 {
self.slot(group_id).load(Ordering::Relaxed)
}
pub(crate) fn add(&self, group_id: RaftGroupId, bytes: u64) {
self.slot(group_id).fetch_add(bytes, Ordering::Relaxed);
}
pub(crate) fn sub(&self, group_id: RaftGroupId, bytes: u64) {
let slot = self.slot(group_id);
let mut current = slot.load(Ordering::Relaxed);
loop {
let next = current.saturating_sub(bytes);
match slot.compare_exchange_weak(current, next, Ordering::Relaxed, Ordering::Relaxed) {
Ok(_) => return,
Err(observed) => current = observed,
}
}
}
fn slot(&self, group_id: RaftGroupId) -> &AtomicU64 {
let index = usize::try_from(group_id.0).expect("u32 fits usize");
&self.per_group[index]
}
}
pub(crate) type SharedRaftUncommittedBytes = Arc<RaftUncommittedBytesTracker>;
pub(crate) struct UncommittedBytesGuard {
tracker: SharedRaftUncommittedBytes,
group_id: RaftGroupId,
bytes: u64,
armed: bool,
}
impl UncommittedBytesGuard {
pub(crate) fn new(
tracker: SharedRaftUncommittedBytes,
group_id: RaftGroupId,
bytes: u64,
) -> Self {
tracker.add(group_id, bytes);
Self {
tracker,
group_id,
bytes,
armed: true,
}
}
#[allow(dead_code)]
pub(crate) fn disarm(mut self) {
self.armed = false;
}
}
impl Drop for UncommittedBytesGuard {
fn drop(&mut self) {
if self.armed {
self.tracker.sub(self.group_id, self.bytes);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn disabled_admission_reports_disabled() {
let admission = RaftUncommittedAdmission::default();
assert!(!admission.is_enabled());
}
#[test]
fn enabled_admission_reports_enabled() {
let admission = RaftUncommittedAdmission {
max_uncommitted_bytes_per_group: Some(1024),
};
assert!(admission.is_enabled());
}
#[test]
fn tracker_add_load_sub_round_trips() {
let tracker = RaftUncommittedBytesTracker::new(2);
tracker.add(RaftGroupId(0), 32);
tracker.add(RaftGroupId(0), 8);
tracker.add(RaftGroupId(1), 4);
assert_eq!(tracker.load(RaftGroupId(0)), 40);
assert_eq!(tracker.load(RaftGroupId(1)), 4);
tracker.sub(RaftGroupId(0), 16);
assert_eq!(tracker.load(RaftGroupId(0)), 24);
}
#[test]
fn tracker_sub_saturates_at_zero() {
let tracker = RaftUncommittedBytesTracker::new(1);
tracker.add(RaftGroupId(0), 4);
tracker.sub(RaftGroupId(0), 10);
assert_eq!(tracker.load(RaftGroupId(0)), 0);
}
#[test]
fn guard_releases_on_drop() {
let tracker = Arc::new(RaftUncommittedBytesTracker::new(1));
{
let _guard = UncommittedBytesGuard::new(tracker.clone(), RaftGroupId(0), 32);
assert_eq!(tracker.load(RaftGroupId(0)), 32);
}
assert_eq!(tracker.load(RaftGroupId(0)), 0);
}
}