use std::{
ops::Deref,
sync::{
Arc, Mutex,
atomic::{AtomicU64, Ordering},
},
time::{Duration, Instant},
};
use dashmap::{DashMap, mapref::entry::Entry, mapref::one::Ref};
use tokio::sync::{mpsc, watch};
use crate::{
ConditionalSetOutcome, HistoryPreservation, RateLimit, RateLimitComparator, RateLimitDecision,
TrypemaError, WindowSize,
common::{HistoryUpdateMode, RandomState, duration_from_milliseconds},
hybrid::{
AbsoluteHybridCommitterSignal, RedisCommitter, RedisCommitterOptions, RedisProxyCommitter,
absolute_hybrid_redis_proxy::{
AbsoluteHybridCommit, AbsoluteHybridRedisProxy, AbsoluteHybridRedisProxyOptions,
AbsoluteHybridRedisProxyReadStateResult,
},
common::{EPOCH_CHANGE_INTERVAL, RedisRateLimiterSignal},
hybrid_rate_limiter_provider::HybridRateLimiterConfig,
},
redis::{RedisKey, mutex_lock, spawn_task},
runtime,
};
mod helpers;
#[derive(Debug)]
enum AbsoluteRedisLimitingState {
Accepting {
window_limit: Mutex<u64>,
accept_limit: Mutex<u64>,
starting_count: Mutex<u64>,
count: AtomicU64,
time_instant: Mutex<Instant>,
oldest_bucket_ttl: Mutex<Option<u64>>,
oldest_bucket_count: Mutex<Option<u64>>,
last_modified: Mutex<Instant>,
},
Undefined,
Rejecting {
time_instant: Mutex<Instant>,
ttl_ms: Mutex<u64>,
count_after_release: Mutex<u64>,
committed_count: u64,
committed_at: Instant,
},
}
#[derive(Debug)]
pub struct AbsoluteHybridRateLimiter {
window_size: WindowSize,
commiter_sender: mpsc::Sender<AbsoluteHybridCommitterSignal<AbsoluteHybridCommit>>,
redis_proxy: AbsoluteHybridRedisProxy,
limiting_state: DashMap<RedisKey, AbsoluteRedisLimitingState, RandomState>,
reset_locks: DashMap<RedisKey, Arc<tokio::sync::Mutex<()>>, RandomState>,
epoch: AtomicU64,
last_commited_epoch: AtomicU64,
is_active_watch: watch::Sender<u64>,
}
impl AbsoluteHybridRateLimiter {
pub(crate) fn new(options: HybridRateLimiterConfig) -> Arc<Self> {
let prefix = options.prefix.unwrap_or_else(RedisKey::default_prefix);
let (tx, rx) = mpsc::channel::<RedisRateLimiterSignal>(1);
let redis_proxy = AbsoluteHybridRedisProxy::new(AbsoluteHybridRedisProxyOptions {
prefix: prefix.clone(),
connection_manager: options.connection_manager,
window_size: options.provider.window_size,
bucket_size: options.provider.bucket_size,
});
let is_active_watch = watch::Sender::new(0u64);
let commiter_sender = RedisCommitter::run(RedisCommitterOptions {
sync_interval: Duration::from_millis(options.sync_interval.as_milliseconds()),
channel_capacity: 8192,
max_batch_size: 4,
limiter_sender: tx,
redis_proxy: Box::new(redis_proxy.clone()),
is_active_watch: is_active_watch.subscribe(),
});
let limiter = Self {
window_size: options.provider.window_size,
commiter_sender,
redis_proxy,
limiting_state: DashMap::default(),
reset_locks: DashMap::default(),
epoch: AtomicU64::new(0),
last_commited_epoch: AtomicU64::new(0),
is_active_watch,
};
let limiter = Arc::new(limiter);
limiter.listen_for_committer_signals(rx);
limiter.epoch_change_task();
limiter
}
fn epoch_change_task(self: &Arc<Self>) {
self.epoch.fetch_add(1, Ordering::AcqRel);
let limiter = Arc::downgrade(self);
spawn_task(async move {
loop {
runtime::sleep(EPOCH_CHANGE_INTERVAL).await;
let Some(limiter) = limiter.upgrade() else {
break;
};
limiter.epoch.fetch_add(1, Ordering::AcqRel);
}
});
}
fn listen_for_committer_signals(
self: &Arc<Self>,
mut rx: mpsc::Receiver<RedisRateLimiterSignal>,
) {
let limitter = Arc::downgrade(self);
spawn_task(async move {
while let Some(signal) = rx.recv().await {
let Some(limiter) = limitter.upgrade() else {
break;
};
match signal {
RedisRateLimiterSignal::Flush => {
if let Err(err) = limiter.flush().await {
tracing::error!(error = ?err, "Failed to flush redis rate limiter");
}
}
}
}
});
}
fn send_epoch_change_if_needed(&self) {
let epoch = self.epoch.load(Ordering::Acquire);
if self.last_commited_epoch.load(Ordering::Acquire) < epoch {
let _ = self.is_active_watch.send(epoch);
self.last_commited_epoch.store(epoch, Ordering::Release);
}
}
fn get_or_create_reset_lock(&self, key: &RedisKey) -> Arc<tokio::sync::Mutex<()>> {
if let Some(lock) = self.reset_locks.get(key) {
return Arc::clone(&lock);
}
self.reset_locks
.entry(key.clone())
.or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
.downgrade()
.clone()
}
fn get_or_create_limiting_state(
&self,
key: &RedisKey,
) -> Ref<'_, RedisKey, AbsoluteRedisLimitingState> {
match self.limiting_state.get(key) {
Some(state) => state,
None => self
.limiting_state
.entry(key.clone())
.or_insert_with(|| AbsoluteRedisLimitingState::Undefined)
.downgrade(),
}
}
pub async fn inc(
&self,
key: &RedisKey,
rate_limit: &RateLimit,
count: u64,
) -> Result<RateLimitDecision, TrypemaError> {
self.send_epoch_change_if_needed();
let decision = self
.is_allowed_with_count_increment(key, count, count, Some(rate_limit))
.await?;
Ok(decision)
}
pub async fn is_allowed(&self, key: &RedisKey) -> Result<RateLimitDecision, TrypemaError> {
self.send_epoch_change_if_needed();
self.is_allowed_with_count_increment(key, 1, 0, None).await
}
pub async fn get_estimate(&self, key: &RedisKey) -> Result<u64, TrypemaError> {
self.send_epoch_change_if_needed();
self.is_allowed_with_count_increment(key, 0, 0, None)
.await?;
let Some(state) = self.limiting_state.get(key) else {
return Ok(0);
};
match state.deref() {
AbsoluteRedisLimitingState::Accepting {
starting_count,
count,
..
} => {
let starting_count = *mutex_lock(starting_count, "accepting.starting_count")?;
Ok(starting_count.saturating_add(count.load(Ordering::Acquire)))
}
AbsoluteRedisLimitingState::Rejecting {
committed_count,
committed_at,
..
} if committed_at.elapsed() < Duration::from_secs(self.window_size.as_seconds()) => {
Ok(*committed_count)
}
AbsoluteRedisLimitingState::Undefined
| AbsoluteRedisLimitingState::Rejecting { .. } => Ok(0),
}
}
pub async fn get(&self, key: &RedisKey) -> Result<u64, TrypemaError> {
self.send_epoch_change_if_needed();
let lock = self.get_or_create_reset_lock(key);
let _guard = lock.lock().await;
let read_state_result = self.redis_proxy.read_state(key).await?;
let mut total_count = read_state_result.current_total_count;
if let Some(state) = self.limiting_state.get(key) {
match state.deref() {
AbsoluteRedisLimitingState::Accepting { count, .. } => {
total_count = total_count.saturating_add(count.load(Ordering::Acquire));
}
AbsoluteRedisLimitingState::Rejecting {
committed_count,
committed_at,
..
} if committed_at.elapsed()
< Duration::from_secs(self.window_size.as_seconds()) =>
{
total_count = total_count.max(*committed_count);
}
AbsoluteRedisLimitingState::Undefined
| AbsoluteRedisLimitingState::Rejecting { .. } => {}
}
}
Ok(total_count)
}
pub async fn set_if(
&self,
key: &RedisKey,
rate_limit: &RateLimit,
comparator: RateLimitComparator,
count: u64,
) -> Result<ConditionalSetOutcome, TrypemaError> {
let (current_total, previous_total) = self
.set_if_with_history_mode(
key,
rate_limit,
comparator,
count,
HistoryUpdateMode::Replace,
)
.await?;
Ok(ConditionalSetOutcome {
matched: comparator.matches(previous_total),
previous_total,
current_total,
})
}
pub async fn set_if_preserve_history(
&self,
key: &RedisKey,
rate_limit: &RateLimit,
comparator: RateLimitComparator,
count: u64,
preservation: HistoryPreservation,
) -> Result<ConditionalSetOutcome, TrypemaError> {
let (current_total, previous_total) = self
.set_if_with_history_mode(
key,
rate_limit,
comparator,
count,
HistoryUpdateMode::Preserve(preservation),
)
.await?;
Ok(ConditionalSetOutcome {
matched: comparator.matches(previous_total),
previous_total,
current_total,
})
}
#[cfg(test)]
pub(crate) fn local_state_count(&self) -> usize {
self.limiting_state.len()
}
fn should_cleanup_local_state(state: &AbsoluteRedisLimitingState, stale_after_ms: u64) -> bool {
match state {
AbsoluteRedisLimitingState::Undefined => true,
AbsoluteRedisLimitingState::Accepting { last_modified, .. } => {
let last_modified = match last_modified.lock() {
Ok(last_modified) => last_modified,
Err(err) => {
tracing::warn!("last_modified is poisoned: {err:?}");
return true;
}
};
last_modified.elapsed().as_millis() >= stale_after_ms as u128
}
AbsoluteRedisLimitingState::Rejecting {
time_instant,
ttl_ms,
..
} => {
let elapsed_ms = match time_instant.lock() {
Ok(time_instant) => time_instant.elapsed().as_millis(),
Err(err) => {
tracing::warn!("time_instant is poisoned: {err:?}");
return true;
}
};
let retry_ttl_ms = match ttl_ms.lock() {
Ok(ttl_ms) => *ttl_ms as u128,
Err(err) => {
tracing::warn!("ttl_ms is poisoned: {err:?}");
return true;
}
};
elapsed_ms.saturating_sub(retry_ttl_ms) >= stale_after_ms as u128
}
}
}
async fn set_if_with_history_mode(
&self,
key: &RedisKey,
rate_limit: &RateLimit,
comparator: RateLimitComparator,
count: u64,
mode: HistoryUpdateMode,
) -> Result<(u64, u64), TrypemaError> {
self.send_epoch_change_if_needed();
let lock = self.get_or_create_reset_lock(key);
let _guard = lock.lock().await;
let pending_count = self
.limiting_state
.get(key)
.and_then(|state| match state.deref() {
AbsoluteRedisLimitingState::Accepting { count, .. } => {
Some(count.load(Ordering::Acquire))
}
_ => None,
})
.unwrap_or(0);
let window_limit =
((self.window_size.as_seconds() as f64) * rate_limit.as_per_second()) as u64;
let (new_total, old_total, changed) = self
.redis_proxy
.set_if(key, window_limit, comparator, count, mode, pending_count)
.await?;
if !changed {
return Ok((new_total, old_total));
}
if let Some((_, state)) = self.limiting_state.remove(key)
&& let AbsoluteRedisLimitingState::Accepting {
count: local_count, ..
} = state
{
let extra_count = local_count.into_inner().saturating_sub(pending_count);
if extra_count > 0 {
let commit = AbsoluteHybridCommit {
key: key.clone(),
window_limit,
count: extra_count,
};
self.send_commit(commit).await?;
}
}
Ok((new_total, old_total))
}
pub(crate) async fn cleanup(&self, stale_after_ms: u64) -> Result<(), TrypemaError> {
self.redis_proxy.cleanup(stale_after_ms).await?;
let candidates = self
.limiting_state
.iter()
.filter(|state| Self::should_cleanup_local_state(state.value(), stale_after_ms))
.map(|state| state.key().clone())
.collect::<Vec<_>>();
for key in candidates {
let lock = self.get_or_create_reset_lock(&key);
let _guard = lock.lock().await;
let should_remove = self.limiting_state.get(&key).is_some_and(|state| {
Self::should_cleanup_local_state(state.value(), stale_after_ms)
});
if !should_remove {
continue;
}
match self.limiting_state.entry(key) {
Entry::Occupied(entry)
if Self::should_cleanup_local_state(entry.get(), stale_after_ms) =>
{
entry.remove();
}
Entry::Occupied(_) | Entry::Vacant(_) => {}
}
}
Ok(())
}
async fn flush(&self) -> Result<(), TrypemaError> {
let mut resets: Vec<RedisKey> = Vec::new();
let mut stale: Vec<RedisKey> = Vec::new();
for state in self.limiting_state.iter() {
let key = state.key();
if let AbsoluteRedisLimitingState::Accepting {
last_modified,
count,
..
} = state.deref()
{
let elapsed = {
let last_modified = match last_modified.lock() {
Ok(last_modified) => last_modified,
Err(err) => {
tracing::warn!("last_modified is poisoned: {err:?}");
continue;
}
};
last_modified.elapsed()
};
if elapsed.as_millis() > self.window_size.as_milliseconds() {
stale.push(key.clone());
continue;
}
if count.load(Ordering::Acquire) == 0 {
continue;
}
resets.push(key.clone());
}
}
let mut guards = Vec::with_capacity(resets.len() + stale.len());
for key in resets.iter().chain(&stale) {
guards.push(self.get_or_create_reset_lock(key).lock_owned().await);
}
for key in stale {
let should_reset = self.limiting_state.get(&key).is_some_and(|state| {
matches!(
state.deref(),
AbsoluteRedisLimitingState::Accepting { last_modified, .. }
if mutex_lock(last_modified, "accepting.last_modified")
.is_ok_and(|last_modified| {
last_modified.elapsed().as_millis()
> self.window_size.as_milliseconds()
})
)
});
if !should_reset {
continue;
}
if let Some(mut state) = self.limiting_state.get_mut(&key)
&& let AbsoluteRedisLimitingState::Accepting { last_modified, .. } = state.deref()
{
let is_stale = mutex_lock(last_modified, "accepting.last_modified")?
.elapsed()
.as_millis()
> self.window_size.as_milliseconds();
if is_stale {
*state = AbsoluteRedisLimitingState::Undefined;
}
}
}
resets.retain(|key| {
self.limiting_state.get(key).is_some_and(|state| {
matches!(
state.deref(),
AbsoluteRedisLimitingState::Accepting { count, .. }
if count.load(Ordering::Acquire) > 0
)
})
});
if resets.is_empty() {
return Ok(());
}
let read_state_results = self.redis_proxy.batch_read_state(&resets).await?;
for result in read_state_results {
if let Err(err) = self
.reset_single_state_from_read_result(result, 0, 0, None)
.await
{
tracing::error!(error = ?err, "Failed to reset state from redis read result");
continue;
}
}
Ok(())
} }