1use std::sync::Arc;
15use std::time::{Duration, Instant};
16
17use tokio::sync::{Mutex, OwnedSemaphorePermit, Semaphore};
18
19#[derive(Debug, thiserror::Error)]
21#[non_exhaustive]
22pub enum RateLimitError {
23 #[error("request rate limit exceeded")]
25 TooManyRequests,
26 #[error("too many concurrent requests")]
28 ConcurrencyLimitExceeded,
29}
30
31struct WindowState {
32 window_start: Instant,
33 count: usize,
34}
35
36pub struct RateLimiter {
38 semaphore: Arc<Semaphore>,
40 window: Duration,
42 max_requests: usize,
44 state: Mutex<WindowState>,
46}
47
48impl RateLimiter {
49 pub fn new(max_concurrent: usize, max_requests_per_minute: usize) -> Self {
55 let permits = if max_concurrent == 0 {
56 Semaphore::MAX_PERMITS
57 } else {
58 max_concurrent
59 };
60 Self {
61 semaphore: Arc::new(Semaphore::new(permits)),
62 window: Duration::from_secs(60),
63 max_requests: max_requests_per_minute,
64 state: Mutex::new(WindowState {
65 window_start: Instant::now(),
66 count: 0,
67 }),
68 }
69 }
70
71 pub async fn try_acquire(&self) -> Result<RateLimitPermit, RateLimitError> {
76 let mut state = self.state.lock().await;
81 let now = Instant::now();
82 if now.duration_since(state.window_start) >= self.window {
83 state.window_start = now;
84 state.count = 0;
85 }
86 if self.max_requests > 0 && state.count >= self.max_requests {
87 return Err(RateLimitError::TooManyRequests);
88 }
89
90 let permit = self
92 .semaphore
93 .clone()
94 .try_acquire_owned()
95 .map_err(|_| RateLimitError::ConcurrencyLimitExceeded)?;
96
97 if self.max_requests > 0 {
98 state.count += 1;
99 }
100 drop(state);
101
102 Ok(RateLimitPermit { _permit: permit })
103 }
104}
105
106pub struct RateLimitPermit {
108 _permit: OwnedSemaphorePermit,
109}
110
111#[cfg(test)]
112mod tests {
113 use super::*;
114
115 #[test]
116 fn unlimited_acquires() {
117 let limiter = RateLimiter::new(0, 0);
118 let rt = tokio::runtime::Runtime::new().unwrap();
119 rt.block_on(async {
120 let permit = limiter.try_acquire().await;
121 assert!(permit.is_ok());
122 drop(permit);
123 assert!(limiter.try_acquire().await.is_ok());
124 });
125 }
126
127 #[test]
128 fn enforces_per_minute_limit() {
129 let limiter = RateLimiter::new(0, 2);
130 let rt = tokio::runtime::Runtime::new().unwrap();
131 rt.block_on(async {
132 assert!(limiter.try_acquire().await.is_ok());
133 assert!(limiter.try_acquire().await.is_ok());
134 let err = limiter.try_acquire().await;
135 assert!(matches!(err, Err(RateLimitError::TooManyRequests)));
136 });
137 }
138
139 #[test]
140 fn enforces_concurrency_limit() {
141 let limiter = RateLimiter::new(1, 0);
142 let rt = tokio::runtime::Runtime::new().unwrap();
143 rt.block_on(async {
144 let p1 = limiter.try_acquire().await.unwrap();
145 let err = limiter.try_acquire().await;
146 assert!(matches!(err, Err(RateLimitError::ConcurrencyLimitExceeded)));
147 drop(p1);
148 assert!(limiter.try_acquire().await.is_ok());
150 });
151 }
152
153 #[tokio::test]
156 async fn permit_is_held_across_await() {
157 let limiter = Arc::new(RateLimiter::new(1, 0));
158 let permit = limiter.try_acquire().await.unwrap();
159 let other = limiter.clone();
160 let rejected = tokio::spawn(async move { other.try_acquire().await.is_err() });
161 tokio::time::sleep(Duration::from_millis(20)).await;
162 assert!(
163 rejected.await.unwrap(),
164 "concurrency cap must apply while the permit is held across an await"
165 );
166 drop(permit);
167 assert!(limiter.try_acquire().await.is_ok());
168 }
169
170 #[tokio::test]
171 async fn concurrency_rejection_does_not_consume_window_budget() {
172 let limiter = RateLimiter::new(1, 2);
173 let permit = limiter.try_acquire().await.unwrap();
174 assert!(matches!(
176 limiter.try_acquire().await,
177 Err(RateLimitError::ConcurrencyLimitExceeded)
178 ));
179 drop(permit);
180 assert!(limiter.try_acquire().await.is_ok());
181 assert!(matches!(
184 limiter.try_acquire().await,
185 Err(RateLimitError::TooManyRequests)
186 ));
187 }
188}