use std::{
sync::{
Arc,
atomic::{AtomicBool, AtomicU64, Ordering},
},
time::Duration,
};
use redis::aio::ConnectionManager;
use crate::{
BucketSize, HardLimitFactor, RateLimiterBuilder, SuppressionFactorCachePeriod, TrypemaError,
WindowSize,
builder::{CleanupConfig, ProviderConfig},
runtime::{new_interval, spawn_task, tick},
};
use super::{AbsoluteRedisRateLimiter, RedisKey, SuppressedRedisRateLimiter};
#[derive(Clone, Debug)]
pub(crate) struct RedisRateLimiterConfig {
pub connection_manager: ConnectionManager,
pub prefix: Option<RedisKey>,
pub provider: ProviderConfig,
}
impl RedisRateLimiterConfig {
pub(crate) fn new(connection_manager: ConnectionManager, provider: ProviderConfig) -> Self {
Self {
connection_manager,
prefix: None,
provider,
}
}
}
#[derive(Clone, Debug)]
pub struct RedisRateLimiterBuilder {
connection_manager: ConnectionManager,
prefix: Option<RedisKey>,
provider: ProviderConfig,
cleanup: CleanupConfig,
}
impl RedisRateLimiterBuilder {
fn new(connection_manager: ConnectionManager) -> Self {
Self {
connection_manager,
prefix: None,
provider: ProviderConfig::default(),
cleanup: CleanupConfig::default(),
}
}
pub fn prefix(mut self, value: RedisKey) -> Self {
self.prefix = Some(value);
self
}
}
impl RateLimiterBuilder for RedisRateLimiterBuilder {
type Provider = RedisRateLimiterProvider;
fn window_size(mut self, value: WindowSize) -> Self {
self.provider.window_size = value;
self
}
fn bucket_size(mut self, value: BucketSize) -> Self {
self.provider.bucket_size = value;
self
}
fn hard_limit_factor(mut self, value: HardLimitFactor) -> Self {
self.provider.hard_limit_factor = value;
self
}
fn suppression_factor_cache_period(mut self, value: SuppressionFactorCachePeriod) -> Self {
self.provider.suppression_factor_cache_period = value;
self
}
fn stale_after(mut self, value: Duration) -> Self {
self.cleanup.stale_after = value;
self
}
fn cleanup_interval(mut self, value: Duration) -> Self {
self.cleanup.interval = value;
self
}
fn cleanup_enabled(mut self, enabled: bool) -> Self {
self.cleanup.enabled = enabled;
self
}
fn build(self) -> Result<Arc<Self::Provider>, TrypemaError> {
let provider_config = self.provider.validate()?;
let cleanup = self.cleanup.validate()?;
let mut config = RedisRateLimiterConfig::new(self.connection_manager, provider_config);
config.prefix = self.prefix;
let provider = Arc::new(RedisRateLimiterProvider::new(config, cleanup));
if cleanup.enabled {
provider.start_cleanup_loop();
}
Ok(provider)
}
}
#[derive(Debug)]
pub struct RedisRateLimiterProvider {
absolute: AbsoluteRedisRateLimiter,
suppressed: SuppressedRedisRateLimiter,
cleanup: CleanupConfig,
is_cleanup_loop_running: AtomicBool,
cleanup_generation: AtomicU64,
}
impl RedisRateLimiterProvider {
pub fn builder(connection_manager: ConnectionManager) -> RedisRateLimiterBuilder {
RedisRateLimiterBuilder::new(connection_manager)
}
pub(crate) fn new(options: RedisRateLimiterConfig, cleanup: CleanupConfig) -> Self {
Self {
absolute: AbsoluteRedisRateLimiter::new(options.clone()),
suppressed: SuppressedRedisRateLimiter::new(options),
cleanup,
is_cleanup_loop_running: AtomicBool::new(false),
cleanup_generation: AtomicU64::new(0),
}
}
pub fn start_cleanup_loop(self: &Arc<Self>) {
if self.is_cleanup_loop_running.swap(true, Ordering::AcqRel) {
return;
}
let generation = self
.cleanup_generation
.fetch_add(1, Ordering::AcqRel)
.wrapping_add(1);
let provider = Arc::downgrade(self);
let cleanup_interval = self.cleanup.interval;
spawn_task(async move {
let mut interval = new_interval(cleanup_interval);
#[cfg(feature = "redis-tokio")]
tick(&mut interval).await;
loop {
tick(&mut interval).await;
let Some(provider) = provider.upgrade() else {
break;
};
if !provider.is_cleanup_loop_running.load(Ordering::Acquire)
|| provider.cleanup_generation.load(Ordering::Acquire) != generation
{
break;
}
if let Err(error) = provider.cleanup(provider.cleanup.stale_after_ms()).await {
tracing::warn!(?error, "Redis cleanup failed, will retry");
}
}
});
}
pub fn stop_cleanup_loop(&self) {
self.cleanup_generation.fetch_add(1, Ordering::AcqRel);
self.is_cleanup_loop_running.store(false, Ordering::Release);
}
pub fn absolute(&self) -> &AbsoluteRedisRateLimiter {
&self.absolute
}
pub fn suppressed(&self) -> &SuppressedRedisRateLimiter {
&self.suppressed
}
pub(crate) async fn cleanup(&self, stale_after_ms: u64) -> Result<(), TrypemaError> {
self.absolute.cleanup(stale_after_ms).await?;
self.suppressed.cleanup(stale_after_ms).await
}
}