sz_orm_pool/
rate_limiter.rs1#[derive(Debug, Clone)]
22pub struct RateLimitResult {
23 pub allowed: bool,
25 pub remaining: u64,
27 pub reset_at: i64,
29}
30
31impl RateLimitResult {
32 pub fn allowed(remaining: u64, reset_at: i64) -> Self {
37 Self {
38 allowed: true,
39 remaining,
40 reset_at,
41 }
42 }
43
44 pub fn rejected(remaining: u64, reset_at: i64) -> Self {
49 Self {
50 allowed: false,
51 remaining,
52 reset_at,
53 }
54 }
55}
56
57#[derive(Debug)]
61pub enum RateLimitError {
62 KeyNotFound(String),
64 Internal(String),
66}
67
68impl std::fmt::Display for RateLimitError {
69 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
70 match self {
71 RateLimitError::KeyNotFound(key) => write!(f, "Key not found: {}", key),
72 RateLimitError::Internal(msg) => write!(f, "Internal error: {}", msg),
73 }
74 }
75}
76
77impl std::error::Error for RateLimitError {}
78
79impl serde::Serialize for RateLimitError {
80 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
81 where
82 S: serde::Serializer,
83 {
84 serializer.serialize_str(&self.to_string())
85 }
86}
87
88pub trait RateLimiter: Send + Sync {
101 fn acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError>;
105
106 fn try_acquire(&self, key: &str) -> Result<RateLimitResult, RateLimitError> {
110 self.acquire(key)
111 }
112
113 fn reset(&self, key: &str) -> Result<(), RateLimitError>;
115}
116
117#[cfg(test)]
118mod tests {
119 use super::*;
120 use std::sync::atomic::{AtomicU64, Ordering};
121 use std::sync::Arc;
122
123 #[test]
124 fn test_rate_limit_result_allowed() {
125 let result = RateLimitResult::allowed(5, 1000);
126 assert!(result.allowed);
127 assert_eq!(result.remaining, 5);
128 assert_eq!(result.reset_at, 1000);
129 }
130
131 #[test]
132 fn test_rate_limit_result_rejected() {
133 let result = RateLimitResult::rejected(0, 2000);
134 assert!(!result.allowed);
135 assert_eq!(result.remaining, 0);
136 }
137
138 #[test]
139 fn test_rate_limit_error_display() {
140 let e = RateLimitError::KeyNotFound("user-1".to_string());
141 assert!(format!("{}", e).contains("user-1"));
142 let e = RateLimitError::Internal("lock poisoned".to_string());
143 assert!(format!("{}", e).contains("lock poisoned"));
144 }
145
146 struct CounterLimiter {
148 max: u64,
149 count: AtomicU64,
150 }
151
152 impl CounterLimiter {
153 fn new(max: u64) -> Self {
154 Self {
155 max,
156 count: AtomicU64::new(0),
157 }
158 }
159 }
160
161 impl RateLimiter for CounterLimiter {
162 fn acquire(&self, _key: &str) -> Result<RateLimitResult, RateLimitError> {
163 let prev = self.count.fetch_add(1, Ordering::SeqCst);
164 if prev < self.max {
165 Ok(RateLimitResult::allowed(self.max - prev - 1, 0))
166 } else {
167 Ok(RateLimitResult::rejected(0, 0))
168 }
169 }
170
171 fn reset(&self, _key: &str) -> Result<(), RateLimitError> {
172 self.count.store(0, Ordering::SeqCst);
173 Ok(())
174 }
175 }
176
177 #[test]
178 fn test_counter_limiter_via_trait() {
179 let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(2));
180 let r1 = limiter.acquire("k").unwrap();
181 assert!(r1.allowed);
182 let r2 = limiter.acquire("k").unwrap();
183 assert!(r2.allowed);
184 let r3 = limiter.acquire("k").unwrap();
185 assert!(!r3.allowed);
186 }
187
188 #[test]
189 fn test_counter_limiter_reset() {
190 let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(1));
191 assert!(limiter.acquire("k").unwrap().allowed);
192 assert!(!limiter.acquire("k").unwrap().allowed);
193 limiter.reset("k").unwrap();
194 assert!(limiter.acquire("k").unwrap().allowed);
195 }
196
197 #[test]
198 fn test_default_try_acquire_delegates_to_acquire() {
199 let limiter: Arc<dyn RateLimiter> = Arc::new(CounterLimiter::new(1));
200 let r1 = limiter.try_acquire("k").unwrap();
202 assert!(r1.allowed);
203 let r2 = limiter.try_acquire("k").unwrap();
204 assert!(!r2.allowed);
205 }
206
207 #[test]
208 fn test_rate_limiter_is_send_sync() {
209 fn assert_send_sync<T: Send + Sync>() {}
210 assert_send_sync::<Arc<dyn RateLimiter>>();
211 }
212}