use crate::error::InvalidPolicyError;
use std::time::Duration;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RateLimitPolicy {
MinInterval { interval: Duration },
PerMinute { rpm: u32 },
}
impl RateLimitPolicy {
pub fn min_interval(interval: Duration) -> Result<Self, InvalidPolicyError> {
Self::MinInterval { interval }.validated()
}
pub fn per_minute(rpm: u32) -> Result<Self, InvalidPolicyError> {
Self::PerMinute { rpm }.validated()
}
pub fn interval(&self) -> Result<Duration, InvalidPolicyError> {
match self {
Self::MinInterval { interval } => {
if interval.is_zero() {
return Err(InvalidPolicyError::zero_interval());
}
Ok(*interval)
}
Self::PerMinute { rpm } => {
if *rpm == 0 {
return Err(InvalidPolicyError::zero_rpm());
}
let nanos_per_minute: u128 = Duration::from_secs(60).as_nanos();
let interval_nanos = nanos_per_minute / u128::from(*rpm);
if interval_nanos == 0 {
return Err(InvalidPolicyError::interval_out_of_range());
}
let interval_nanos = u64::try_from(interval_nanos)
.map_err(|_| InvalidPolicyError::interval_out_of_range())?;
Ok(Duration::from_nanos(interval_nanos))
}
}
}
pub fn validated(self) -> Result<Self, InvalidPolicyError> {
self.interval()?;
Ok(self)
}
}
#[cfg(test)]
mod tests {
use super::RateLimitPolicy;
use crate::InvalidPolicyKind;
use std::time::Duration;
#[test]
fn per_minute_converts_to_even_interval() {
let interval = RateLimitPolicy::per_minute(120)
.unwrap()
.interval()
.unwrap();
assert_eq!(interval, Duration::from_millis(500));
}
#[test]
fn zero_rpm_is_invalid() {
let error = RateLimitPolicy::per_minute(0).unwrap_err();
assert_eq!(error.kind(), InvalidPolicyKind::ZeroRatePerMinute);
}
#[test]
fn zero_interval_is_invalid() {
let error = RateLimitPolicy::min_interval(Duration::ZERO).unwrap_err();
assert_eq!(error.kind(), InvalidPolicyKind::ZeroInterval);
}
}