use redis::{Script, aio::ConnectionManager};
use crate::{
BucketSize, ConditionalSetOutcome, HistoryPreservation, RateLimit, RateLimitComparator,
RateLimitDecision, TrypemaError, WindowSize,
common::{HistoryUpdateMode, RateType, duration_from_milliseconds},
redis::{
RedisKey, RedisKeyGenerator,
redis_rate_limiter_provider::RedisRateLimiterConfig,
scripts::{
ABSOLUTE_CLEANUP_LUA, ABSOLUTE_GET_TOTAL_LUA, ABSOLUTE_INC_LUA,
ABSOLUTE_IS_ALLOWED_LUA, ABSOLUTE_SET_IF_LUA, absolute_lua_script,
},
},
};
#[derive(Clone, Debug)]
pub struct AbsoluteRedisRateLimiter {
connection_manager: ConnectionManager,
window_size: WindowSize,
bucket_size: BucketSize,
key_generator: RedisKeyGenerator,
inc_script: Script,
is_allowed_script: Script,
get_total_script: Script,
set_if_script: Script,
cleanup_script: Script,
}
impl AbsoluteRedisRateLimiter {
pub(crate) fn new(options: RedisRateLimiterConfig) -> Self {
let prefix = options.prefix.unwrap_or_else(RedisKey::default_prefix);
Self {
connection_manager: options.connection_manager,
window_size: options.provider.window_size,
bucket_size: options.provider.bucket_size,
key_generator: RedisKeyGenerator::new(prefix, RateType::Absolute),
inc_script: absolute_lua_script(ABSOLUTE_INC_LUA),
is_allowed_script: absolute_lua_script(ABSOLUTE_IS_ALLOWED_LUA),
get_total_script: absolute_lua_script(ABSOLUTE_GET_TOTAL_LUA),
set_if_script: absolute_lua_script(ABSOLUTE_SET_IF_LUA),
cleanup_script: absolute_lua_script(ABSOLUTE_CLEANUP_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 window_limit = self.window_size.as_seconds() as f64 * rate_limit.as_per_second();
let (result, retry_after_ms, remaining_after_waiting): (String, u128, u64) = 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())
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.arg(window_limit)
.arg(self.bucket_size.as_milliseconds())
.arg(count)
.invoke_async(&mut connection_manager)
.await?;
match result.as_str() {
"allowed" => Ok(RateLimitDecision::Allowed),
"rejected" => Ok(RateLimitDecision::Rejected {
window_size: self.window_size,
retry_after: duration_from_milliseconds(retry_after_ms),
remaining_after_waiting,
}),
_ => Err(TrypemaError::UnexpectedRedisScriptResult {
operation: "absolute.inc",
key: key.to_string(),
result,
}),
}
}
pub async fn is_allowed(&self, key: &RedisKey) -> Result<RateLimitDecision, TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let (result, retry_after_ms, remaining_after_waiting): (String, u128, u64) = self
.is_allowed_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))
.arg(self.window_size.as_seconds())
.invoke_async(&mut connection_manager)
.await?;
match result.as_str() {
"allowed" => Ok(RateLimitDecision::Allowed),
"rejected" => Ok(RateLimitDecision::Rejected {
window_size: self.window_size,
retry_after: duration_from_milliseconds(retry_after_ms),
remaining_after_waiting,
}),
_ => Err(TrypemaError::UnexpectedRedisScriptResult {
operation: "absolute.is_allowed",
key: key.to_string(),
result,
}),
}
}
pub async fn get(&self, key: &RedisKey) -> Result<u64, TrypemaError> {
let mut connection_manager = self.connection_manager.clone();
let total: u64 = self
.get_total_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_window_limit_key(key))
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.invoke_async(&mut connection_manager)
.await?;
Ok(total)
}
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())
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.arg(self.window_size.as_seconds() as f64 * rate_limit.as_per_second())
.arg(comparator_op)
.arg(comparator_operand)
.arg(count)
.arg(mode.redis_arg())
.arg(0_u64)
.invoke_async(&mut connection_manager)
.await?;
Ok((new_total, old_total))
}
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())
.invoke_async(&mut connection_manager)
.await?;
Ok(())
}
}