use std::sync::{
atomic::{AtomicU64, Ordering},
Arc,
};
use tokio::sync::Notify;
use super::Frame;
#[derive(Clone, Debug)]
pub(crate) struct ByteBudget {
inner: Arc<ByteBudgetInner>,
}
#[derive(Debug)]
struct ByteBudgetInner {
max_bytes: u64,
reserved_bytes: AtomicU64,
capacity: Notify,
}
impl ByteBudget {
pub(crate) fn new(max_bytes: u64) -> Self {
Self {
inner: Arc::new(ByteBudgetInner {
max_bytes,
reserved_bytes: AtomicU64::new(0),
capacity: Notify::new(),
}),
}
}
pub(crate) fn available(&self) -> u64 {
self.inner
.max_bytes
.saturating_sub(self.inner.reserved_bytes.load(Ordering::Acquire))
}
#[cfg(test)]
pub(crate) fn max_bytes_for_test(&self) -> u64 {
self.inner.max_bytes
}
pub(crate) fn reserved(&self) -> u64 {
self.inner.reserved_bytes.load(Ordering::Acquire)
}
pub(crate) fn try_reserve(&mut self, bytes: u64) -> bool {
if bytes == 0 {
return false;
}
let mut reserved = self.inner.reserved_bytes.load(Ordering::Acquire);
loop {
let available = self.inner.max_bytes.saturating_sub(reserved);
if bytes > available {
return false;
}
let next = reserved.saturating_add(bytes);
match self.inner.reserved_bytes.compare_exchange_weak(
reserved,
next,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return true,
Err(observed) => reserved = observed,
}
}
}
pub(crate) fn release(&mut self, bytes: u64) {
if bytes == 0 {
return;
}
let mut reserved = self.inner.reserved_bytes.load(Ordering::Acquire);
loop {
let next = reserved.saturating_sub(bytes);
match self.inner.reserved_bytes.compare_exchange_weak(
reserved,
next,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => break,
Err(observed) => reserved = observed,
}
}
self.inner.capacity.notify_waiters();
}
pub(crate) fn charge(&mut self, bytes: u64) {
if bytes == 0 {
return;
}
let mut reserved = self.inner.reserved_bytes.load(Ordering::Acquire);
loop {
let next = reserved.saturating_add(bytes);
match self.inner.reserved_bytes.compare_exchange_weak(
reserved,
next,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => break,
Err(observed) => reserved = observed,
}
}
}
pub(crate) fn audit(&self, expected: u64, _context: &'static str) -> bool {
let actual = self.reserved();
let ok = actual >= expected;
if !ok {
metrics::counter!("sync.block.budget.audit_drift").increment(1);
}
ok
}
pub(crate) fn subscribe_capacity(&self) -> &Notify {
&self.inner.capacity
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub(crate) enum Admit {
Pass,
Throttle,
Reject(&'static str),
}
#[derive(Debug)]
pub(crate) struct PeerMeters;
impl PeerMeters {
pub(crate) fn new() -> Self {
Self
}
pub(crate) fn try_take(&mut self, _message_type: u8) -> bool {
true
}
}
#[derive(Debug)]
pub(crate) struct SessionGuard {
allowed: &'static [u8],
max_bytes: u32,
byte_budget: Option<ByteBudget>,
meters: PeerMeters,
}
impl SessionGuard {
#[allow(dead_code)] pub(crate) fn new(
allowed: &'static [u8],
max_bytes: u32,
byte_budget: Option<ByteBudget>,
) -> Self {
Self {
allowed,
max_bytes,
byte_budget,
meters: PeerMeters::new(),
}
}
pub(crate) fn oversize_only(max_bytes: u32) -> Self {
const ALL_TYPES: &[u8] = &{
let mut all = [0u8; 256];
let mut ty = 0usize;
while ty < all.len() {
all[ty] = ty as u8;
ty += 1;
}
all
};
Self {
allowed: ALL_TYPES,
max_bytes,
byte_budget: None,
meters: PeerMeters::new(),
}
}
pub(crate) fn admit(&mut self, frame: &Frame) -> Admit {
let Ok(ty) = u8::try_from(frame.message_type) else {
return Admit::Reject("bad type");
};
if !self.allowed.contains(&ty) {
return Admit::Reject("disallowed type");
}
if frame.payload.len() > self.max_bytes as usize {
return Admit::Reject("oversize");
}
if let Some(budget) = &mut self.byte_budget {
if !budget.try_reserve(frame.payload.len() as u64) {
return Admit::Throttle;
}
}
if !self.meters.try_take(ty) {
return Admit::Throttle;
}
Admit::Pass
}
#[allow(dead_code)] pub(crate) fn release(&mut self, bytes: u64) {
if let Some(budget) = &mut self.byte_budget {
budget.release(bytes);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const ALLOWED: &[u8] = &[1, 2];
fn frame(message_type: u16, payload_len: usize) -> Frame {
Frame {
message_type,
flags: 0,
payload: vec![0u8; payload_len],
}
}
#[test]
fn byte_budget_reserves_and_releases() {
let mut budget = ByteBudget::new(1_000);
assert_eq!(budget.available(), 1_000);
assert!(budget.try_reserve(400));
assert_eq!(budget.reserved(), 400);
assert_eq!(budget.available(), 600);
assert!(!budget.try_reserve(0));
assert!(!budget.try_reserve(601));
assert_eq!(budget.reserved(), 400);
budget.release(400);
assert_eq!(budget.reserved(), 0);
}
#[test]
fn byte_budget_charge_overdrafts_past_the_max() {
let mut budget = ByteBudget::new(1_000);
assert!(budget.try_reserve(900));
budget.charge(300);
assert_eq!(budget.reserved(), 1_200);
assert_eq!(budget.available(), 0);
assert!(!budget.try_reserve(1));
budget.release(1_200);
assert_eq!(budget.reserved(), 0);
}
#[test]
fn byte_budget_concurrent_reservations_never_over_commit() {
use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
use std::sync::{Arc, Barrier};
use std::thread;
const CHUNK: u64 = 4_096;
const CAP: u64 = 8;
const THREADS: usize = 16; let max = CHUNK * CAP;
let budget = ByteBudget::new(max);
let start = Arc::new(Barrier::new(THREADS));
let all_reserved = Arc::new(Barrier::new(THREADS));
let successes = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..THREADS {
let mut budget = budget.clone();
let start = start.clone();
let all_reserved = all_reserved.clone();
let successes = successes.clone();
handles.push(thread::spawn(move || {
start.wait();
let ok = budget.try_reserve(CHUNK);
if ok {
successes.fetch_add(1, AtomicOrdering::AcqRel);
}
all_reserved.wait();
assert!(
budget.reserved() <= max,
"reserved {} exceeded the budget max {max} (over-commit)",
budget.reserved(),
);
if ok {
budget.release(CHUNK);
}
}));
}
for handle in handles {
handle.join().expect("reserver thread panicked");
}
assert_eq!(
u64::try_from(successes.load(AtomicOrdering::Acquire)).unwrap(),
CAP,
"exactly CAP reservations may be admitted; the rest must be rejected",
);
assert_eq!(
budget.reserved(),
0,
"every admitted reservation was released"
);
let mut handles = Vec::new();
for _ in 0..THREADS {
let mut budget = budget.clone();
handles.push(thread::spawn(move || {
for _ in 0..5_000 {
if budget.try_reserve(CHUNK) {
assert!(budget.reserved() <= max, "over-commit during churn");
budget.release(CHUNK);
}
}
}));
}
for handle in handles {
handle.join().expect("churn thread panicked");
}
assert_eq!(budget.reserved(), 0);
}
#[test]
fn admit_rejects_disallowed_type() {
let mut guard = SessionGuard::new(ALLOWED, 1_024, None);
assert_eq!(guard.admit(&frame(99, 0)), Admit::Reject("disallowed type"));
}
#[test]
fn admit_rejects_non_u8_type() {
let mut guard = SessionGuard::new(ALLOWED, 1_024, None);
assert_eq!(guard.admit(&frame(0x0100, 0)), Admit::Reject("bad type"));
}
#[test]
fn admit_rejects_oversize() {
let mut guard = SessionGuard::new(ALLOWED, 4, None);
assert_eq!(guard.admit(&frame(1, 5)), Admit::Reject("oversize"));
}
#[test]
fn admit_throttles_when_budget_exhausted() {
let mut guard = SessionGuard::new(ALLOWED, 1_024, Some(ByteBudget::new(8)));
assert_eq!(guard.admit(&frame(1, 8)), Admit::Pass);
assert_eq!(guard.admit(&frame(1, 8)), Admit::Throttle);
guard.release(8);
assert_eq!(guard.admit(&frame(1, 8)), Admit::Pass);
}
#[test]
fn admit_passes_allowed_under_caps() {
let mut guard = SessionGuard::new(ALLOWED, 1_024, None);
assert_eq!(guard.admit(&frame(1, 16)), Admit::Pass);
assert_eq!(guard.admit(&frame(2, 16)), Admit::Pass);
}
#[test]
fn oversize_only_admits_all_u8_types_under_cap() {
let mut guard = SessionGuard::oversize_only(1_024);
assert_eq!(guard.admit(&frame(0, 0)), Admit::Pass);
assert_eq!(guard.admit(&frame(99, 16)), Admit::Pass);
assert_eq!(guard.admit(&frame(255, 16)), Admit::Pass);
}
#[test]
fn oversize_only_rejects_non_u8_type() {
let mut guard = SessionGuard::oversize_only(1_024);
assert_eq!(guard.admit(&frame(0x0100, 0)), Admit::Reject("bad type"));
}
#[test]
fn oversize_only_rejects_oversize() {
let mut guard = SessionGuard::oversize_only(4);
assert_eq!(guard.admit(&frame(1, 5)), Admit::Reject("oversize"));
}
}