Skip to main content

rate_limiter/
policy.rs

1use crate::error::InvalidPolicyError;
2use std::time::Duration;
3
4#[derive(Debug, Clone, Copy, PartialEq, Eq)]
5pub enum RateLimitPolicy {
6    MinInterval { interval: Duration },
7    PerMinute { rpm: u32 },
8}
9
10impl RateLimitPolicy {
11    pub fn min_interval(interval: Duration) -> Result<Self, InvalidPolicyError> {
12        Self::MinInterval { interval }.validated()
13    }
14
15    pub fn per_minute(rpm: u32) -> Result<Self, InvalidPolicyError> {
16        Self::PerMinute { rpm }.validated()
17    }
18
19    pub fn interval(&self) -> Result<Duration, InvalidPolicyError> {
20        match self {
21            Self::MinInterval { interval } => {
22                if interval.is_zero() {
23                    return Err(InvalidPolicyError::zero_interval());
24                }
25                Ok(*interval)
26            }
27            Self::PerMinute { rpm } => {
28                if *rpm == 0 {
29                    return Err(InvalidPolicyError::zero_rpm());
30                }
31                let nanos_per_minute: u128 = Duration::from_secs(60).as_nanos();
32                let interval_nanos = nanos_per_minute / u128::from(*rpm);
33                if interval_nanos == 0 {
34                    return Err(InvalidPolicyError::interval_out_of_range());
35                }
36                let interval_nanos = u64::try_from(interval_nanos)
37                    .map_err(|_| InvalidPolicyError::interval_out_of_range())?;
38                Ok(Duration::from_nanos(interval_nanos))
39            }
40        }
41    }
42
43    pub fn validated(self) -> Result<Self, InvalidPolicyError> {
44        self.interval()?;
45        Ok(self)
46    }
47}
48
49#[cfg(test)]
50mod tests {
51    use super::RateLimitPolicy;
52    use crate::InvalidPolicyKind;
53    use std::time::Duration;
54
55    #[test]
56    fn per_minute_converts_to_even_interval() {
57        let interval = RateLimitPolicy::per_minute(120)
58            .unwrap()
59            .interval()
60            .unwrap();
61        assert_eq!(interval, Duration::from_millis(500));
62    }
63
64    #[test]
65    fn zero_rpm_is_invalid() {
66        let error = RateLimitPolicy::per_minute(0).unwrap_err();
67        assert_eq!(error.kind(), InvalidPolicyKind::ZeroRatePerMinute);
68    }
69
70    #[test]
71    fn zero_interval_is_invalid() {
72        let error = RateLimitPolicy::min_interval(Duration::ZERO).unwrap_err();
73        assert_eq!(error.kind(), InvalidPolicyKind::ZeroInterval);
74    }
75}