use dashmap::DashMap;
use parking_lot::Mutex;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::time::{Duration, Instant};
const GLOBAL_BACKPRESSURE_ERROR_THRESHOLD: usize = 50;
const GLOBAL_BACKPRESSURE_PENALTY: Duration = Duration::from_secs(1);
const RATE_LIMIT_BACKOFF_MULTIPLIER: u32 = 2;
const RATE_LIMIT_MAX_INTERVAL: Duration = Duration::from_secs(8);
const RATE_LIMIT_RECOVERY_STEP_DIVISOR: u32 = 2;
struct ServiceLimit {
last_request: Instant,
interval: Duration,
base_interval: Duration,
}
pub struct RateLimiter {
services: DashMap<String, Mutex<ServiceLimit>>,
default_interval_nanos: AtomicU64,
global_error_count: AtomicUsize,
}
impl RateLimiter {
pub fn new(rps: f64) -> Self {
Self {
services: DashMap::new(),
default_interval_nanos: AtomicU64::new(rps_to_nanos(rps)),
global_error_count: AtomicUsize::new(0),
}
}
pub fn set_default_rps(&self, rps: f64) {
self.default_interval_nanos
.store(rps_to_nanos(rps), Ordering::Relaxed);
}
pub fn default_interval(&self) -> Duration {
Duration::from_nanos(self.default_interval_nanos.load(Ordering::Relaxed))
}
pub async fn wait(&self, service: &str) {
let bp = if self.global_error_count.load(Ordering::Relaxed)
> GLOBAL_BACKPRESSURE_ERROR_THRESHOLD
{
GLOBAL_BACKPRESSURE_PENALTY
} else {
Duration::ZERO
};
let wait_time = {
let default = self.default_interval();
if let Some(entry) = self.services.get(service) {
let mut limit = entry.value().lock();
reserve_service_slot(&mut limit, Instant::now())
} else {
let inserted = self.services.entry(service.to_string()).or_insert_with(|| {
Mutex::new(ServiceLimit {
last_request: initial_last_request(Instant::now(), default),
interval: default,
base_interval: default,
})
});
let mut limit = inserted.value().lock();
reserve_service_slot(&mut limit, Instant::now())
}
};
let delay = match wait_time {
Some(wait) => wait.max(bp),
None => bp,
};
if !delay.is_zero() {
tokio::time::sleep(delay).await;
}
}
pub fn record_error(&self) {
self.global_error_count.fetch_add(1, Ordering::Relaxed);
}
pub fn record_success(&self) {
let _ = self .global_error_count
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |n| {
Some(n.saturating_sub(1))
});
}
pub(crate) fn error_count_for_test(&self) -> usize {
self.global_error_count.load(Ordering::Relaxed)
}
pub async fn update_limit(&self, service: &str, rps: f64) {
let interval = Duration::from_nanos(rps_to_nanos(rps));
self.services.insert(
service.to_string(),
Mutex::new(ServiceLimit {
last_request: Instant::now(),
interval,
base_interval: interval,
}),
);
}
pub fn penalize_service(&self, service: &str) {
if let Some(entry) = self.services.get(service) {
let mut limit = entry.value().lock();
let ceiling = RATE_LIMIT_MAX_INTERVAL.max(limit.base_interval);
limit.interval = limit
.interval
.checked_mul(RATE_LIMIT_BACKOFF_MULTIPLIER)
.map_or(ceiling, |interval| interval)
.min(ceiling);
} else {
let default = self.default_interval();
let ceiling = RATE_LIMIT_MAX_INTERVAL.max(default);
let interval = default
.checked_mul(RATE_LIMIT_BACKOFF_MULTIPLIER)
.map_or(ceiling, |interval| interval)
.min(ceiling);
self.services.entry(service.to_string()).or_insert_with(|| {
Mutex::new(ServiceLimit {
last_request: Instant::now(),
interval,
base_interval: default,
})
});
}
}
pub fn reward_service(&self, service: &str) {
if let Some(entry) = self.services.get(service) {
let mut limit = entry.value().lock();
if limit.interval > limit.base_interval {
let step = limit.base_interval / RATE_LIMIT_RECOVERY_STEP_DIVISOR;
limit.interval = limit.interval.saturating_sub(step).max(limit.base_interval);
}
}
}
pub fn service_interval(&self, service: &str) -> Option<Duration> {
self.services
.get(service)
.map(|entry| entry.value().lock().interval)
}
}
pub(crate) fn initial_last_request(now: Instant, interval: Duration) -> Instant {
now.checked_sub(interval).map_or(now, |instant| instant)
}
fn reserve_service_slot(limit: &mut ServiceLimit, now: Instant) -> Option<Duration> {
let next_slot = limit.last_request + limit.interval;
if now >= next_slot {
limit.last_request = now;
None
} else {
let wait = next_slot.saturating_duration_since(now);
limit.last_request = next_slot;
Some(wait)
}
}
fn rps_to_nanos(rps: f64) -> u64 {
let rate = if rps.is_finite() && rps > 0.0 {
rps
} else {
1.0
};
let nanos = (1.0e9 / rate).round();
if nanos.is_finite() && nanos < 1.0 {
1
} else if nanos.is_finite() && nanos <= u64::MAX as f64 {
nanos as u64
} else {
1_000_000_000
}
}
use std::sync::OnceLock;
pub static GLOBAL_RATE_LIMITER: OnceLock<RateLimiter> = OnceLock::new();
pub fn get_rate_limiter() -> &'static RateLimiter {
GLOBAL_RATE_LIMITER.get_or_init(|| RateLimiter::new(5.0))
}
pub fn set_global_default_rps(rps: f64) {
get_rate_limiter().set_default_rps(rps);
}