use crate::clock::{Clock, SystemClock};
use crate::config::FluxLimiterConfig;
use crate::errors::FluxLimiterError;
use dashmap::DashMap;
use dashmap::mapref::entry::Entry;
use std::hash::Hash;
use std::sync::Arc;
impl<T, C> Clone for FluxLimiter<T, C>
where
T: Hash + Eq + Clone,
C: Clock + Clone,
{
fn clone(&self) -> Self {
Self {
rate_nanos: self.rate_nanos,
tolerance_nanos: self.tolerance_nanos,
client_state: Arc::clone(&self.client_state),
clock: self.clock.clone(),
}
}
}
#[derive(Debug)]
pub struct FluxLimiter<T, C = SystemClock>
where
T: Hash + Eq + Clone,
C: Clock,
{
rate_nanos: u64,
tolerance_nanos: u64,
pub(crate) client_state: Arc<DashMap<T, u64>>,
clock: C,
}
impl<T, C> FluxLimiter<T, C>
where
T: Hash + Eq + Clone,
C: Clock,
{
fn new(rate_per_second: f64, burst_capacity: f64, clock: C) -> Result<Self, FluxLimiterError> {
let rate_nanos = (1_000_000_000.0 / rate_per_second) as u64;
let tolerance_nanos = (burst_capacity * rate_nanos as f64) as u64;
Ok(Self {
rate_nanos,
tolerance_nanos,
client_state: Arc::new(DashMap::new()),
clock,
})
}
pub fn with_config(config: FluxLimiterConfig, clock: C) -> Result<Self, FluxLimiterError> {
config.validate()?;
Self::new(config.rate_per_second, config.burst_capacity, clock)
}
pub fn rate(&self) -> f64 {
1_000_000_000.0 / self.rate_nanos as f64
}
pub fn burst(&self) -> f64 {
self.tolerance_nanos as f64 / self.rate_nanos as f64
}
pub fn client_count(&self) -> usize {
self.client_state.len()
}
pub fn contains_client(&self, client_id: &T) -> bool {
self.client_state.contains_key(client_id)
}
pub fn check_request(&self, client_id: T) -> Result<FluxLimiterDecision, FluxLimiterError> {
let current_time_nanos = self.clock.now()?;
match self.client_state.entry(client_id) {
Entry::Occupied(mut occupied) => {
let previous_tat_nanos = *occupied.get();
let is_conforming =
current_time_nanos >= previous_tat_nanos.saturating_sub(self.tolerance_nanos);
if is_conforming {
let new_tat_nanos =
current_time_nanos.max(previous_tat_nanos) + self.rate_nanos;
occupied.insert(new_tat_nanos);
Ok(FluxLimiterDecision {
allowed: true,
retry_after_seconds: None,
remaining_capacity: Some(
self.calculate_remaining_capacity(current_time_nanos, new_tat_nanos),
),
reset_time_nanos: new_tat_nanos,
})
} else {
let retry_after_nanos = previous_tat_nanos
.saturating_sub(self.tolerance_nanos)
.saturating_sub(current_time_nanos);
Ok(FluxLimiterDecision {
allowed: false,
retry_after_seconds: Some(retry_after_nanos as f64 / 1_000_000_000.0),
remaining_capacity: Some(0.0),
reset_time_nanos: previous_tat_nanos,
})
}
}
Entry::Vacant(vacant) => {
let new_tat_nanos = current_time_nanos + self.rate_nanos;
vacant.insert(new_tat_nanos);
Ok(FluxLimiterDecision {
allowed: true,
retry_after_seconds: None,
remaining_capacity: Some(
self.calculate_remaining_capacity(current_time_nanos, new_tat_nanos),
),
reset_time_nanos: new_tat_nanos,
})
}
}
}
fn calculate_remaining_capacity(&self, current_time: u64, tat: u64) -> f64 {
if current_time >= tat.saturating_sub(self.tolerance_nanos) {
let time_until_tat = tat.saturating_sub(current_time) as f64 / 1_000_000_000.0;
let rate_per_second = self.rate();
(self.burst() - (time_until_tat * rate_per_second)).max(0.0)
} else {
0.0
}
}
pub fn cleanup_stale_clients(&self, max_stale_nanos: u64) -> Result<(), FluxLimiterError> {
let current_time_nanos = self.clock.now()?;
self.client_state.retain(|_, &mut tat| {
tat + self.tolerance_nanos > current_time_nanos.saturating_sub(max_stale_nanos)
});
Ok(())
}
}
#[non_exhaustive]
#[derive(Debug, Clone)]
pub struct FluxLimiterDecision {
pub allowed: bool,
pub retry_after_seconds: Option<f64>,
pub remaining_capacity: Option<f64>,
pub reset_time_nanos: u64,
}