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_DELETE_LUA, ABSOLUTE_GET_TOTAL_LUA, ABSOLUTE_INC_LUA,
ABSOLUTE_IS_ALLOWED_LUA, ABSOLUTE_SET_IF_LUA, ABSOLUTE_SET_RATE_LIMIT_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,
set_rate_limit_script: Script,
delete_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),
set_rate_limit_script: absolute_lua_script(ABSOLUTE_SET_RATE_LIMIT_LUA),
delete_script: absolute_lua_script(ABSOLUTE_DELETE_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)
}
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 window_limit = self.window_size.as_seconds() as f64 * rate_limit.as_per_second();
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())
.arg(key.to_string())
.arg(self.window_size.as_seconds())
.arg(window_limit)
.invoke_async(&mut connection_manager)
.await?;
match status.as_str() {
"missing" => Ok(None),
"found" => {
let previous_window_limit = previous.parse::<f64>().map_err(|_| {
TrypemaError::CustomError("invalid stored absolute window limit".to_string())
})?;
RateLimit::from_stored_window_limit(previous_window_limit, self.window_size, 1.0)
.map(Some)
}
"invalid" => Err(TrypemaError::CustomError(
"invalid stored absolute window limit".to_string(),
)),
_ => Err(TrypemaError::UnexpectedRedisScriptResult {
operation: "absolute.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, total_count): (u8, u64) = self
.delete_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())
.invoke_async(&mut connection_manager)
.await?;
Ok((existed == 1).then_some(total_count))
}
pub async fn clear(&self) -> Result<(), TrypemaError> {
self.cleanup(0).await
}
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(())
}
}