use redis::{Script, aio::ConnectionManager};
use crate::{
BucketSize, ConditionalSetOutcome, HardLimitFactor, HistoryPreservation, RateLimit,
RateLimitComparator, RateLimitDecision, SuppressedRateLimitSnapshot, TrypemaError, WindowSize,
common::{HistoryUpdateMode, RateType, SuppressionFactorCachePeriod},
redis::{
RedisKey, RedisKeyGenerator,
redis_rate_limiter_provider::RedisRateLimiterConfig,
scripts::{
SUPPRESSED_CLEANUP_LUA, SUPPRESSED_DELETE_LUA, SUPPRESSED_GET_FACTOR_LUA,
SUPPRESSED_GET_STATE_LUA, SUPPRESSED_INC_LUA, SUPPRESSED_SET_IF_LUA,
SUPPRESSED_SET_RATE_LIMIT_LUA, lua_script, suppressed_lua_script,
},
},
};
#[derive(Clone, Debug)]
pub struct SuppressedRedisRateLimiter {
connection_manager: ConnectionManager,
key_generator: RedisKeyGenerator,
hard_limit_factor: HardLimitFactor,
bucket_size: BucketSize,
window_size: WindowSize,
suppression_factor_cache_period: SuppressionFactorCachePeriod,
inc_script: Script,
cleanup_script: Script,
suppression_factor_script: Script,
get_state_script: Script,
set_if_script: Script,
set_rate_limit_script: Script,
delete_script: Script,
}
impl SuppressedRedisRateLimiter {
pub(crate) fn new(options: RedisRateLimiterConfig) -> Self {
let prefix = options.prefix.unwrap_or_else(RedisKey::default_prefix);
let key_generator = RedisKeyGenerator::new(prefix, RateType::Suppressed);
Self {
connection_manager: options.connection_manager,
window_size: options.provider.window_size,
bucket_size: options.provider.bucket_size,
hard_limit_factor: options.provider.hard_limit_factor,
suppression_factor_cache_period: options.provider.suppression_factor_cache_period,
key_generator,
inc_script: suppressed_lua_script(SUPPRESSED_INC_LUA),
cleanup_script: lua_script(SUPPRESSED_CLEANUP_LUA),
suppression_factor_script: suppressed_lua_script(SUPPRESSED_GET_FACTOR_LUA),
get_state_script: suppressed_lua_script(SUPPRESSED_GET_STATE_LUA),
set_if_script: lua_script(SUPPRESSED_SET_IF_LUA),
set_rate_limit_script: lua_script(SUPPRESSED_SET_RATE_LIMIT_LUA),
delete_script: suppressed_lua_script(SUPPRESSED_DELETE_LUA),
}
}
pub async fn inc(
&self,
key: &RedisKey,
rate_limit: &RateLimit,
count: u64,
) -> Result<RateLimitDecision, TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let hard_window_limit = self.window_size.as_seconds() as f64
* rate_limit.as_per_second()
* self.hard_limit_factor.as_multiplier();
let (result, suppression_factor, should_allow): (String, f64, u8) = self
.inc_script
.key(self.key_generator.get_hash_key(key))
.key(self.key_generator.get_active_keys(key))
.key(self.key_generator.get_window_limit_key(key))
.key(self.key_generator.get_total_count_key(key))
.key(self.key_generator.get_active_entities_key())
.key(self.key_generator.get_suppression_factor_key(key))
.key(self.key_generator.get_total_declined_key(key))
.key(self.key_generator.get_hash_declined_key(key))
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.arg(hard_window_limit)
.arg(self.bucket_size.as_milliseconds())
.arg(self.suppression_factor_cache_period.as_milliseconds())
.arg(self.hard_limit_factor.as_multiplier())
.arg(count)
.invoke_async(&mut connection_manager)
.await?;
match result.as_str() {
"allowed" => Ok(RateLimitDecision::Allowed),
"suppressed" => Ok(RateLimitDecision::Suppressed {
suppression_factor,
is_allowed: should_allow == 1,
}),
_ => Err(TrypemaError::UnexpectedRedisScriptResult {
operation: "suppressed.inc",
key: key.to_string(),
result,
}),
}
}
async fn set_if_with_history_mode(
&self,
key: &RedisKey,
rate_limit: &RateLimit,
comparator: RateLimitComparator,
count: u64,
mode: HistoryUpdateMode,
) -> Result<(u64, u64), TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let (comparator_op, comparator_operand) = comparator.redis_args();
let (new_total, old_total, _changed): (u64, u64, u64) = self
.set_if_script
.key(self.key_generator.get_hash_key(key))
.key(self.key_generator.get_active_keys(key))
.key(self.key_generator.get_window_limit_key(key))
.key(self.key_generator.get_total_count_key(key))
.key(self.key_generator.get_active_entities_key())
.key(self.key_generator.get_suppression_factor_key(key))
.key(self.key_generator.get_total_declined_key(key))
.key(self.key_generator.get_hash_declined_key(key))
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.arg(
self.window_size.as_seconds() as f64
* rate_limit.as_per_second()
* self.hard_limit_factor.as_multiplier(),
)
.arg(comparator_op)
.arg(comparator_operand)
.arg(count)
.arg(mode.redis_arg())
.arg(0_u64)
.arg(0_u64)
.invoke_async(&mut connection_manager)
.await?;
Ok((new_total, old_total))
}
pub async fn get_suppression_factor(&self, key: &RedisKey) -> Result<f64, TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let suppression_factor: f64 = self
.suppression_factor_script
.key(self.key_generator.get_hash_key(key))
.key(self.key_generator.get_active_keys(key))
.key(self.key_generator.get_window_limit_key(key))
.key(self.key_generator.get_total_count_key(key))
.key(self.key_generator.get_active_entities_key())
.key(self.key_generator.get_suppression_factor_key(key))
.key(self.key_generator.get_total_declined_key(key))
.key(self.key_generator.get_hash_declined_key(key))
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.arg(self.bucket_size.as_milliseconds())
.arg(self.suppression_factor_cache_period.as_milliseconds())
.arg(self.hard_limit_factor.as_multiplier())
.invoke_async(&mut connection_manager)
.await?;
Ok(suppression_factor)
}
pub async fn get(&self, key: &RedisKey) -> Result<SuppressedRateLimitSnapshot, TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let (total, total_declined, suppression_factor): (u64, u64, f64) = self
.get_state_script
.key(self.key_generator.get_hash_key(key))
.key(self.key_generator.get_active_keys(key))
.key(self.key_generator.get_total_count_key(key))
.key(self.key_generator.get_active_entities_key())
.key(self.key_generator.get_total_declined_key(key))
.key(self.key_generator.get_hash_declined_key(key))
.key(self.key_generator.get_window_limit_key(key))
.key(self.key_generator.get_suppression_factor_key(key))
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.arg(self.suppression_factor_cache_period.as_milliseconds())
.arg(self.hard_limit_factor.as_multiplier())
.invoke_async(&mut connection_manager)
.await?;
Ok(SuppressedRateLimitSnapshot {
total,
total_declined,
suppression_factor,
})
}
pub async fn set_rate_limit(
&self,
key: &RedisKey,
rate_limit: &RateLimit,
) -> Result<Option<RateLimit>, TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let hard_window_limit = self.window_size.as_seconds() as f64
* rate_limit.as_per_second()
* self.hard_limit_factor.as_multiplier();
let (status, previous, _changed): (String, String, u8) = self
.set_rate_limit_script
.key(self.key_generator.get_window_limit_key(key))
.key(self.key_generator.get_active_entities_key())
.key(self.key_generator.get_suppression_factor_key(key))
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.arg(hard_window_limit)
.arg(self.hard_limit_factor.as_multiplier())
.invoke_async(&mut connection_manager)
.await?;
match status.as_str() {
"missing" => Ok(None),
"found" => {
let previous_hard_window_limit = previous.parse::<f64>().map_err(|_| {
TrypemaError::CustomError(
"invalid stored suppressed hard window limit".to_string(),
)
})?;
RateLimit::from_stored_window_limit(
previous_hard_window_limit,
self.window_size,
self.hard_limit_factor.as_multiplier(),
)
.map(Some)
}
"invalid" => Err(TrypemaError::CustomError(
"invalid stored suppressed hard window limit".to_string(),
)),
_ => Err(TrypemaError::UnexpectedRedisScriptResult {
operation: "suppressed.set_rate_limit",
key: key.to_string(),
result: status,
}),
}
}
pub async fn delete(&self, key: &RedisKey) -> Result<Option<u64>, TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let (existed, accepted_count): (u8, u64) = self
.delete_script
.key(self.key_generator.get_hash_key(key))
.key(self.key_generator.get_hash_declined_key(key))
.key(self.key_generator.get_active_keys(key))
.key(self.key_generator.get_window_limit_key(key))
.key(self.key_generator.get_total_count_key(key))
.key(self.key_generator.get_total_declined_key(key))
.key(self.key_generator.get_suppression_factor_key(key))
.key(self.key_generator.get_active_entities_key())
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.invoke_async(&mut connection_manager)
.await?;
Ok((existed == 1).then_some(accepted_count))
}
pub async fn clear(&self) -> Result<(), TrypemaError> {
self.cleanup(0).await
}
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,
})
}
pub(crate) async fn cleanup(&self, stale_after_ms: u64) -> Result<(), TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let _: () = self
.cleanup_script
.key(self.key_generator.prefix.to_string())
.key(self.key_generator.rate_type.to_string())
.key(self.key_generator.get_active_entities_key())
.arg(stale_after_ms)
.arg(self.key_generator.hash_key_suffix.to_string())
.arg(self.key_generator.window_limit_key_suffix.to_string())
.arg(self.key_generator.total_count_key_suffix.to_string())
.arg(self.key_generator.active_keys_key_suffix.to_string())
.arg(self.key_generator.suppression_factor_key_suffix.to_string())
.arg(self.key_generator.total_declined_key_suffix.to_string())
.arg(self.key_generator.hash_declined_key_suffix.to_string())
.invoke_async(&mut connection_manager)
.await?;
Ok(())
}
}