1use super::StoreError;
4use std::collections::HashMap;
5use std::sync::Mutex;
6
7pub trait ThrottleStore: Send + Sync {
12 fn record_failure(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError>;
15 fn failure_count(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError>;
20 fn is_banned(&self, key: &str, now: u64) -> Result<Option<u64>, StoreError>;
26 fn ban(&self, key: &str, until: u64) -> Result<(), StoreError>;
29 fn clear_failures(&self, key: &str) -> Result<(), StoreError>;
34 fn reset(&self, key: &str) -> Result<(), StoreError>;
36 fn purge_expired(&self, now: u64) -> Result<usize, StoreError>;
38}
39
40#[derive(Debug, Default)]
45struct Entry {
46 failures: Vec<u64>,
49 banned_until: Option<u64>,
52 window_secs: u64,
56}
57
58#[derive(Debug)]
70pub struct MemoryThrottleStore {
71 entries: Mutex<HashMap<String, Entry>>,
72}
73
74impl Default for MemoryThrottleStore {
75 fn default() -> Self {
76 Self::new()
77 }
78}
79
80impl MemoryThrottleStore {
81 pub fn new() -> Self {
82 Self {
83 entries: Mutex::new(HashMap::new()),
84 }
85 }
86
87 fn lock<T>(m: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
92 m.lock().unwrap_or_else(|e| e.into_inner())
93 }
94
95 fn cutoff(now: u64, window_secs: u64) -> u64 {
97 now.saturating_sub(window_secs)
98 }
99}
100
101impl ThrottleStore for MemoryThrottleStore {
102 fn record_failure(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError> {
103 let cutoff = Self::cutoff(now, window_secs);
104 let mut g = Self::lock(&self.entries);
105 let e = g.entry(key.to_string()).or_default();
106 e.window_secs = window_secs;
107 e.failures.retain(|t| *t > cutoff);
112 e.failures.push(now);
113 Ok(e.failures.len() as u32)
114 }
115
116 fn failure_count(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError> {
117 let cutoff = Self::cutoff(now, window_secs);
118 let g = Self::lock(&self.entries);
119 Ok(g.get(key)
120 .map(|e| e.failures.iter().filter(|t| **t > cutoff).count() as u32)
121 .unwrap_or(0))
122 }
123
124 fn is_banned(&self, key: &str, now: u64) -> Result<Option<u64>, StoreError> {
125 Ok(Self::lock(&self.entries)
126 .get(key)
127 .and_then(|e| e.banned_until)
128 .filter(|until| *until > now))
130 }
131
132 fn ban(&self, key: &str, until: u64) -> Result<(), StoreError> {
133 let mut g = Self::lock(&self.entries);
134 let e = g.entry(key.to_string()).or_default();
135 e.banned_until = Some(e.banned_until.map_or(until, |existing| existing.max(until)));
138 Ok(())
139 }
140
141 fn clear_failures(&self, key: &str) -> Result<(), StoreError> {
142 if let Some(e) = Self::lock(&self.entries).get_mut(key) {
143 e.failures.clear();
144 }
145 Ok(())
146 }
147
148 fn reset(&self, key: &str) -> Result<(), StoreError> {
149 Self::lock(&self.entries).remove(key);
150 Ok(())
151 }
152
153 fn purge_expired(&self, now: u64) -> Result<usize, StoreError> {
154 let mut g = Self::lock(&self.entries);
155 let before = g.len();
156 g.retain(|_, e| {
157 let banned = e.banned_until.is_some_and(|until| until > now);
161 let fresh = e
164 .failures
165 .iter()
166 .any(|t| *t > Self::cutoff(now, e.window_secs));
167 banned || fresh
168 });
169 Ok(before - g.len())
170 }
171}
172
173#[cfg(test)]
174mod tests {
175 use super::*;
176
177 const NOW: u64 = 1_000_000;
178
179 fn store() -> MemoryThrottleStore {
180 MemoryThrottleStore::new()
181 }
182
183 #[test]
184 fn record_failure_counts_within_window() {
185 let s = store();
186 assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 1);
187 assert_eq!(s.record_failure("k", NOW + 1, 60).unwrap(), 2);
188 }
189
190 #[test]
191 fn failures_at_window_edge_are_dropped() {
192 let s = store();
193 s.record_failure("k", NOW, 60).unwrap();
194 assert_eq!(s.record_failure("k", NOW + 60, 60).unwrap(), 1);
196 assert_eq!(s.record_failure("k", NOW + 61, 60).unwrap(), 2);
198 }
199
200 #[test]
201 fn failure_count_is_read_only() {
202 let s = store();
203 assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 0);
204 s.record_failure("k", NOW, 60).unwrap();
205 assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
206 for _ in 0..10 {
208 assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
209 }
210 assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 2);
211 }
212
213 #[test]
214 fn failure_count_of_unknown_key_is_zero() {
215 let s = store();
216 assert_eq!(s.failure_count("ghost", NOW, 60).unwrap(), 0);
217 }
218
219 #[test]
220 fn failures_vec_stays_bounded_to_window() {
221 let s = store();
223 for i in 0..100 {
224 s.record_failure("k", NOW + i, 60).unwrap();
225 }
226 assert_eq!(len(&s, "k"), 60, "窗口外的失败必须真被丢弃");
228 s.record_failure("k", NOW + 1_000, 60).unwrap();
230 assert_eq!(len(&s, "k"), 1);
231 }
232
233 #[test]
234 fn banned_expires_at_until_exclusive() {
235 let s = store();
236 s.ban("k", NOW + 100).unwrap();
237 assert_eq!(s.is_banned("k", NOW).unwrap(), Some(NOW + 100));
238 assert_eq!(s.is_banned("k", NOW + 99).unwrap(), Some(NOW + 100));
239 assert_eq!(s.is_banned("k", NOW + 100).unwrap(), None);
240 }
241
242 #[test]
243 fn is_banned_unknown_key_is_none() {
244 let s = store();
245 assert_eq!(s.is_banned("ghost", NOW).unwrap(), None);
246 }
247
248 #[test]
249 fn reset_clears_failures_and_ban() {
250 let s = store();
251 s.record_failure("k", NOW, 60).unwrap();
252 s.ban("k", NOW + 100).unwrap();
253 s.reset("k").unwrap();
254 assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 0);
255 assert_eq!(s.is_banned("k", NOW).unwrap(), None);
256 }
257
258 #[test]
259 fn reset_unknown_key_is_ok() {
260 let s = store();
261 assert!(s.reset("ghost").is_ok());
262 }
263
264 #[test]
265 fn purge_keeps_banned_and_fresh_entries() {
266 let s = store();
267 s.record_failure("stale", NOW, 60).unwrap();
268 s.record_failure("fresh", NOW + 990, 60).unwrap();
269 s.ban("banned", NOW + 5_000).unwrap();
270 assert_eq!(s.purge_expired(NOW + 1_000).unwrap(), 1);
271 assert_eq!(s.failure_count("stale", NOW + 1_000, 60).unwrap(), 0);
272 assert_eq!(s.failure_count("fresh", NOW + 1_000, 60).unwrap(), 1);
273 assert!(s.is_banned("banned", NOW + 1_000).unwrap().is_some());
274 }
275
276 #[test]
277 fn purge_with_expired_ban_drops_entry() {
278 let s = store();
279 s.record_failure("k", NOW, 60).unwrap();
280 s.ban("k", NOW + 10).unwrap();
281 assert_eq!(s.purge_expired(NOW + 100).unwrap(), 1);
283 assert_eq!(s.is_banned("k", NOW + 100).unwrap(), None);
284 }
285
286 #[test]
287 fn purge_empty_store_is_zero() {
288 let s = store();
289 assert_eq!(s.purge_expired(NOW).unwrap(), 0);
290 }
291
292 #[test]
293 fn purge_keeps_in_window_failures_after_clock_rollback() {
294 let s = store();
298 s.record_failure("k", NOW, 60).unwrap();
299 s.record_failure("k", NOW - 200, 60).unwrap();
300 assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
301 assert_eq!(
302 s.purge_expired(NOW).unwrap(),
303 0,
304 "窗口内仍有有效失败,不该清"
305 );
306 assert_eq!(
307 s.failure_count("k", NOW, 60).unwrap(),
308 1,
309 "计数不能被 purge 免费重置"
310 );
311 }
312
313 #[test]
314 fn ban_cannot_be_shortened_by_clock_rollback() {
315 let s = store();
316 s.ban("k", NOW + 900).unwrap();
317 s.ban("k", NOW - 5_000 + 900).unwrap();
319 assert_eq!(
320 s.is_banned("k", NOW).unwrap(),
321 Some(NOW + 900),
322 "封禁只能延长,回拨不能提前解封"
323 );
324 }
325
326 #[test]
327 fn empty_key_is_a_normal_key() {
328 let s = store();
330 assert_eq!(s.record_failure("", NOW, 60).unwrap(), 1);
331 assert_eq!(s.record_failure("", NOW, 60).unwrap(), 2);
332 assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 1);
333 }
334
335 #[test]
336 fn lock_recovers_from_poisoned_mutex() {
337 let m = Mutex::new(Entry {
339 failures: vec![NOW],
340 banned_until: Some(NOW + 1),
341 window_secs: 60,
342 });
343 std::panic::catch_unwind(|| {
344 let _guard = m.lock().unwrap();
345 panic!("poison");
346 })
347 .unwrap_err();
348 assert!(m.is_poisoned());
349
350 let g = MemoryThrottleStore::lock(&m);
351 assert_eq!(g.failures, vec![NOW]);
352 assert_eq!(g.banned_until, Some(NOW + 1));
353 }
354
355 fn len(s: &MemoryThrottleStore, key: &str) -> usize {
357 MemoryThrottleStore::lock(&s.entries)
358 .get(key)
359 .map(|e| e.failures.len())
360 .unwrap_or(0)
361 }
362}