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}