use crate::container::is_running_in_container;
use crate::shutdown;
use crate::sysinfo::{cpu::CPU, memory::Memory};
use parking_lot::Mutex;
use ringbuf::traits::*;
use ringbuf::HeapRb;
use serde::{Deserialize, Serialize};
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicU8, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tracing::{debug, info, warn};
#[inline]
pub fn default_bucket_count() -> u32 {
50
}
#[inline]
pub fn default_bucket_interval() -> Duration {
Duration::from_millis(200)
}
#[inline]
pub fn default_cpu_threshold() -> u8 {
100
}
#[inline]
pub fn default_memory_threshold() -> u8 {
90
}
#[inline]
pub fn default_shed_cooldown() -> Duration {
Duration::from_secs(5)
}
#[inline]
pub fn default_collect_interval() -> Duration {
Duration::from_secs(3)
}
pub struct RequestGuard<'a> {
limiter: &'a BBR,
created_at: Instant,
}
impl Drop for RequestGuard<'_> {
fn drop(&mut self) {
let rt = self.created_at.elapsed().as_millis().max(1) as u64;
self.limiter.rolling_window.add(rt);
self.limiter.rolling_window.sub_in_flight();
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default, rename_all = "camelCase")]
pub struct BBRConfig {
#[serde(default = "default_bucket_count")]
pub bucket_count: u32,
#[serde(default = "default_bucket_interval", with = "humantime_serde")]
pub bucket_interval: Duration,
#[serde(default = "default_cpu_threshold")]
pub cpu_threshold: u8,
#[serde(default = "default_memory_threshold")]
pub memory_threshold: u8,
#[serde(default = "default_shed_cooldown", with = "humantime_serde")]
pub shed_cooldown: Duration,
#[serde(default = "default_collect_interval", with = "humantime_serde")]
pub collect_interval: Duration,
}
impl Default for BBRConfig {
fn default() -> Self {
Self {
bucket_count: default_bucket_count(),
bucket_interval: default_bucket_interval(),
cpu_threshold: default_cpu_threshold(),
memory_threshold: default_memory_threshold(),
shed_cooldown: default_shed_cooldown(),
collect_interval: default_collect_interval(),
}
}
}
pub struct BBR {
rolling_window: RollingWindow,
overload_collector: Arc<OverloadCollector>,
shed_cooldown: Duration,
shed_at: Mutex<Option<Instant>>,
}
impl BBR {
pub async fn new(config: BBRConfig) -> Self {
let overload_collector = Arc::new(OverloadCollector::new(config.clone()));
overload_collector.collect_overloaded().await;
let overload_collector_clone = overload_collector.clone();
tokio::spawn(async move {
overload_collector_clone.run().await;
});
Self {
rolling_window: RollingWindow::new(config.bucket_count, config.bucket_interval),
overload_collector,
shed_cooldown: config.shed_cooldown,
shed_at: Mutex::new(None),
}
}
pub async fn acquire(&self) -> Option<RequestGuard<'_>> {
self.rolling_window.add_in_flight();
if self.should_shed().await {
self.rolling_window.sub_in_flight();
return None;
}
Some(RequestGuard {
limiter: self,
created_at: Instant::now(),
})
}
async fn should_shed(&self) -> bool {
if self.is_in_cooldown() {
debug!("in cooldown period after shedding, continuing to shed requests");
return true;
}
if !self.overload_collector.is_overloaded() {
return false;
}
let (max_pass, min_rt, in_flight) = self.rolling_window.get_stats();
if max_pass == 0 || in_flight == 0 {
return false;
}
let estimated_limit =
(max_pass as f64 * min_rt as f64 * self.rolling_window.bucket_count() as f64 / 1000.0)
.round() as u64;
if estimated_limit >= in_flight {
return false;
}
warn!(
"overloaded: cpu={}%, memory={}%, estimated_limit={}, in_flight={}",
self.overload_collector.cpu_used_percent(),
self.overload_collector.memory_used_percent(),
estimated_limit,
in_flight
);
self.shed_at.lock().replace(Instant::now());
true
}
#[inline]
fn is_in_cooldown(&self) -> bool {
self.shed_at
.lock()
.is_some_and(|shed_at| shed_at.elapsed() < self.shed_cooldown)
}
}
struct OverloadCollector {
is_overloaded: AtomicBool,
cpu: CPU,
cpu_threshold: u8,
cpu_used_percent: AtomicU8,
memory: Memory,
memory_threshold: u8,
memory_used_percent: AtomicU8,
collect_interval: Duration,
pid: u32,
is_running_in_container: bool,
}
impl OverloadCollector {
pub fn new(config: BBRConfig) -> Self {
Self {
is_overloaded: AtomicBool::new(false),
cpu: CPU::new(),
cpu_threshold: config.cpu_threshold,
cpu_used_percent: AtomicU8::new(0),
memory: Memory::default(),
memory_threshold: config.memory_threshold,
memory_used_percent: AtomicU8::new(0),
collect_interval: config.collect_interval,
pid: std::process::id(),
is_running_in_container: is_running_in_container(),
}
}
pub async fn run(&self) {
let mut interval = tokio::time::interval(self.collect_interval);
loop {
tokio::select! {
_ = interval.tick() => {
self.collect_overloaded().await;
}
_ = shutdown::shutdown_signal() => {
info!("ratelimiter's collecting server shutting down");
return
}
}
}
}
pub async fn collect_overloaded(&self) {
self.is_overloaded.store(
self.is_cpu_overloaded().await || self.is_memory_overloaded(),
Ordering::Relaxed,
);
}
pub fn is_overloaded(&self) -> bool {
self.is_overloaded.load(Ordering::Relaxed)
}
pub fn cpu_used_percent(&self) -> u8 {
self.cpu_used_percent.load(Ordering::Relaxed)
}
pub fn memory_used_percent(&self) -> u8 {
self.memory_used_percent.load(Ordering::Relaxed)
}
#[inline]
fn is_memory_overloaded(&self) -> bool {
if self.memory_threshold == 100 {
return false;
}
let used_percent = if self.is_running_in_container {
match self.memory.get_cgroup_stats(self.pid) {
Some(stats) => stats.used_percent.round() as u8,
None => {
warn!("container detected but cgroup memory stats unavailable, falling back to process stats");
self.memory.get_process_stats(self.pid).used_percent.round() as u8
}
}
} else {
self.memory.get_process_stats(self.pid).used_percent.round() as u8
};
self.memory_used_percent
.store(used_percent, Ordering::Relaxed);
used_percent >= self.memory_threshold
}
#[inline]
async fn is_cpu_overloaded(&self) -> bool {
if self.cpu_threshold == 100 {
return false;
}
let used_percent = if self.is_running_in_container {
match self.cpu.get_cgroup_stats(self.pid).await {
Some(stats) => stats.used_percent.round() as u8,
None => {
warn!("container detected but cgroup CPU stats unavailable, falling back to process stats");
self.cpu
.get_process_stats(self.pid)
.await
.used_percent
.round() as u8
}
}
} else {
self.cpu
.get_process_stats(self.pid)
.await
.used_percent
.round() as u8
};
self.cpu_used_percent.store(used_percent, Ordering::Relaxed);
used_percent >= self.cpu_threshold
}
}
#[derive(Clone, Copy)]
struct Sample {
pass: u64,
min_rt: u64,
}
pub struct RollingWindow {
ring: Mutex<HeapRb<Sample>>,
bucket_count: u32,
bucket_interval: Duration,
current_bucket: Mutex<(Instant, u64, u64)>,
in_flight: AtomicU64,
}
impl RollingWindow {
pub fn new(bucket_count: u32, bucket_interval: Duration) -> Self {
Self {
ring: Mutex::new(HeapRb::new(bucket_count as usize)),
bucket_count,
bucket_interval,
current_bucket: Mutex::new((Instant::now(), 0, u64::MAX)),
in_flight: AtomicU64::new(0),
}
}
pub fn add(&self, rt: u64) {
let now = Instant::now();
let mut current_bucket = self.current_bucket.lock();
if now.duration_since(current_bucket.0) >= self.bucket_interval {
if current_bucket.1 > 0 {
let mut ring = self.ring.lock();
ring.push_overwrite(Sample {
pass: current_bucket.1,
min_rt: current_bucket.2,
});
}
*current_bucket = (now, 0, u64::MAX);
}
current_bucket.1 += 1;
current_bucket.2 = current_bucket.2.min(rt);
}
pub fn get_stats(&self) -> (u64, u64, u64) {
let ring = self.ring.lock();
let (max_pass, min_rt) = ring.iter().fold((0, u64::MAX), |(max_pass, min_rt), s| {
(max_pass.max(s.pass), min_rt.min(s.min_rt))
});
let min_rt = if min_rt == u64::MAX { 1 } else { min_rt };
(max_pass, min_rt, self.in_flight())
}
#[inline]
pub fn in_flight(&self) -> u64 {
self.in_flight.load(Ordering::Relaxed)
}
#[inline]
pub fn add_in_flight(&self) -> u64 {
self.in_flight.fetch_add(1, Ordering::Relaxed) + 1
}
#[inline]
pub fn sub_in_flight(&self) {
self.in_flight.fetch_sub(1, Ordering::Relaxed);
}
#[inline]
pub fn bucket_count(&self) -> u32 {
self.bucket_count
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::type_complexity)]
use super::*;
use std::thread;
#[test]
fn default_config_uses_the_default_fns() {
let config = BBRConfig::default();
assert_eq!(config.bucket_count, default_bucket_count());
assert_eq!(config.bucket_interval, default_bucket_interval());
assert_eq!(config.cpu_threshold, default_cpu_threshold());
assert_eq!(config.memory_threshold, default_memory_threshold());
assert_eq!(config.shed_cooldown, default_shed_cooldown());
assert_eq!(config.collect_interval, default_collect_interval());
}
#[test]
fn new_window_starts_empty() {
let window = RollingWindow::new(5, Duration::from_millis(50));
assert_eq!(window.bucket_count(), 5);
assert_eq!(window.in_flight(), 0);
assert_eq!(window.get_stats(), (0, 1, 0));
assert_eq!(window.ring.lock().occupied_len(), 0);
}
#[test]
fn add_counts_passes_and_tracks_the_min_rt_within_a_bucket() {
let test_cases: Vec<(Vec<u64>, fn((u64, u64)))> = vec![
(vec![100], |bucket| assert_eq!(bucket, (1, 100))),
(vec![100, 50, 200], |bucket| assert_eq!(bucket, (3, 50))),
(vec![0], |bucket| assert_eq!(bucket, (1, 0))),
(vec![u64::MAX - 1], |bucket| {
assert_eq!(bucket, (1, u64::MAX - 1));
}),
];
for (rts, expect) in test_cases {
let window = RollingWindow::new(10, Duration::from_millis(100));
for rt in rts {
window.add(rt);
}
assert_eq!(window.ring.lock().occupied_len(), 0);
let current_bucket = window.current_bucket.lock();
expect((current_bucket.1, current_bucket.2));
}
}
#[test]
fn add_flushes_the_expired_bucket_into_the_ring() {
let window = RollingWindow::new(10, Duration::from_millis(100));
window.add(100);
window.add(80);
thread::sleep(Duration::from_millis(150));
window.add(150);
let ring = window.ring.lock();
assert_eq!(ring.occupied_len(), 1);
let sample = ring.iter().next().unwrap();
assert_eq!(sample.pass, 2);
assert_eq!(sample.min_rt, 80);
drop(ring);
let current_bucket = window.current_bucket.lock();
assert_eq!((current_bucket.1, current_bucket.2), (1, 150));
}
#[test]
fn add_skips_flushing_an_empty_bucket() {
let window = RollingWindow::new(10, Duration::from_millis(100));
thread::sleep(Duration::from_millis(120));
window.add(100);
assert_eq!(window.ring.lock().occupied_len(), 0);
let current_bucket = window.current_bucket.lock();
assert_eq!((current_bucket.1, current_bucket.2), (1, 100));
}
#[test]
fn add_overwrites_the_oldest_sample_when_the_ring_is_full() {
let window = RollingWindow::new(1, Duration::from_millis(100));
window.add(100);
thread::sleep(Duration::from_millis(150));
window.add(50);
thread::sleep(Duration::from_millis(150));
window.add(25);
let ring = window.ring.lock();
assert_eq!(ring.occupied_len(), 1);
assert_eq!(ring.iter().next().unwrap().min_rt, 50);
}
#[test]
fn in_flight_follows_add_and_sub() {
let test_cases: Vec<(Vec<fn(&RollingWindow)>, u64)> = vec![
(vec![], 0),
(
vec![
|window| assert_eq!(window.add_in_flight(), 1),
|window| assert_eq!(window.add_in_flight(), 2),
|window| assert_eq!(window.add_in_flight(), 3),
],
3,
),
(
vec![
|window| assert_eq!(window.add_in_flight(), 1),
|window| assert_eq!(window.add_in_flight(), 2),
RollingWindow::sub_in_flight,
],
1,
),
(
vec![
|window| assert_eq!(window.add_in_flight(), 1),
RollingWindow::sub_in_flight,
|window| assert_eq!(window.add_in_flight(), 1),
RollingWindow::sub_in_flight,
],
0,
),
];
for (ops, expected) in test_cases {
let window = RollingWindow::new(10, Duration::from_millis(100));
for op in ops {
op(&window);
}
assert_eq!(window.in_flight(), expected);
assert_eq!(window.get_stats().2, expected);
}
}
#[test]
fn get_stats_aggregates_the_flushed_buckets() {
let window = RollingWindow::new(10, Duration::from_millis(100));
window.add(100);
window.add(50);
thread::sleep(Duration::from_millis(120));
window.add(200);
window.add(150);
window.add(120);
thread::sleep(Duration::from_millis(110));
window.add(80);
window.add_in_flight();
window.add_in_flight();
assert_eq!(window.get_stats(), (3, 50, 2));
}
#[test]
fn get_stats_drops_samples_beyond_the_bucket_count() {
let window = RollingWindow::new(3, Duration::from_millis(100));
for rt in [100, 80, 60, 40] {
window.add(rt);
thread::sleep(Duration::from_millis(120));
}
window.add(30);
let (max_pass, min_rt, _) = window.get_stats();
assert_eq!(max_pass, 1);
assert_eq!(min_rt, 40);
}
#[test]
fn concurrent_adds_land_in_the_current_bucket() {
let window = Arc::new(RollingWindow::new(10, Duration::from_millis(100)));
let mut handles = vec![];
for _ in 0..4 {
let window = window.clone();
handles.push(thread::spawn(move || {
for rt in 0..100 {
window.add(rt);
}
}));
}
for handle in handles {
handle.join().unwrap();
}
let current_bucket = window.current_bucket.lock();
assert!(current_bucket.1 > 0);
assert!(current_bucket.2 <= 99);
}
#[tokio::test]
async fn concurrent_adds_and_stats_stay_consistent() {
let window = Arc::new(RollingWindow::new(10, Duration::from_millis(100)));
let mut handles = vec![];
for _ in 0..2 {
let window = window.clone();
handles.push(tokio::spawn(async move {
for rt in 0..50 {
window.add(rt + 10);
tokio::time::sleep(Duration::from_micros(100)).await;
}
}));
}
for _ in 0..2 {
let window = window.clone();
handles.push(tokio::spawn(async move {
for _ in 0..50 {
let (max_pass, min_rt, _) = window.get_stats();
assert!(min_rt > 0);
assert!(max_pass <= 100);
tokio::time::sleep(Duration::from_micros(100)).await;
}
}));
}
for handle in handles {
handle.await.unwrap();
}
}
#[tokio::test]
async fn acquire_admits_requests_when_not_overloaded() {
let bbr = BBR::new(BBRConfig {
cpu_threshold: 100,
memory_threshold: 100,
..Default::default()
})
.await;
assert!(!bbr.overload_collector.is_overloaded());
let guard = bbr.acquire().await.unwrap();
assert_eq!(bbr.rolling_window.in_flight(), 1);
drop(guard);
assert_eq!(bbr.rolling_window.in_flight(), 0);
assert_eq!(bbr.rolling_window.current_bucket.lock().1, 1);
assert!(bbr.shed_at.lock().is_none());
}
#[tokio::test]
async fn acquire_sheds_only_when_in_flight_exceeds_the_estimated_limit() {
let bbr = BBR::new(BBRConfig {
cpu_threshold: 100,
memory_threshold: 100,
..Default::default()
})
.await;
bbr.overload_collector
.is_overloaded
.store(true, Ordering::Relaxed);
let _baseline = bbr.acquire().await.unwrap();
assert_eq!(bbr.rolling_window.in_flight(), 1);
bbr.rolling_window.add(40);
thread::sleep(bbr.rolling_window.bucket_interval + Duration::from_millis(10));
bbr.rolling_window.add(40);
assert_eq!(bbr.rolling_window.get_stats(), (1, 40, 1));
let guard = bbr.acquire().await;
assert!(guard.is_some());
assert_eq!(bbr.rolling_window.in_flight(), 2);
assert!(!bbr.is_in_cooldown());
drop(guard);
bbr.rolling_window.add_in_flight();
assert!(bbr.acquire().await.is_none());
assert_eq!(bbr.rolling_window.in_flight(), 2);
assert!(bbr.is_in_cooldown());
}
#[tokio::test]
async fn acquire_keeps_shedding_during_the_cooldown() {
let bbr = BBR::new(BBRConfig {
cpu_threshold: 100,
memory_threshold: 100,
shed_cooldown: Duration::from_millis(50),
..Default::default()
})
.await;
bbr.shed_at.lock().replace(Instant::now());
assert!(bbr.is_in_cooldown());
assert!(bbr.acquire().await.is_none());
assert_eq!(bbr.rolling_window.in_flight(), 0);
tokio::time::sleep(Duration::from_millis(60)).await;
assert!(!bbr.is_in_cooldown());
assert!(bbr.acquire().await.is_some());
}
#[tokio::test]
async fn collector_ignores_resources_at_the_full_threshold() {
let collector = OverloadCollector::new(BBRConfig {
cpu_threshold: 100,
memory_threshold: 100,
..Default::default()
});
collector.collect_overloaded().await;
assert!(!collector.is_overloaded());
assert!(!collector.is_cpu_overloaded().await);
assert!(!collector.is_memory_overloaded());
assert_eq!(collector.cpu_used_percent(), 0);
assert_eq!(collector.memory_used_percent(), 0);
}
#[tokio::test]
async fn collector_flags_overload_at_a_zero_threshold() {
let collector = OverloadCollector::new(BBRConfig {
cpu_threshold: 0,
memory_threshold: 0,
..Default::default()
});
collector.collect_overloaded().await;
assert!(collector.is_overloaded());
assert!(collector.is_cpu_overloaded().await);
assert!(collector.is_memory_overloaded());
}
}