1use std::collections::HashMap;
10use std::net::IpAddr;
11use std::num::NonZeroU32;
12use std::sync::Arc;
13use std::sync::atomic::{AtomicU64, Ordering};
14use std::time::{Duration, Instant};
15
16use governor::clock::{Clock, DefaultClock};
17use governor::state::keyed::DashMapStateStore;
18use governor::{Quota, RateLimiter};
19use lru::LruCache;
20use parking_lot::RwLock;
21use std::num::NonZeroUsize;
22use tracing::{debug, warn};
23
24use super::api::{
25 BanEntry, PenaltyReason, PowPenaltyReason, RateLimitApi, RateLimitStatus, RateLimiterStats,
26};
27use super::config::{EndpointCategoryConfig, PowConfig, RateLimitConfig, RateLimitTierConfig};
28use super::error::{PowError, RateLimitError};
29use super::extractors::AddressKey;
30use super::pow::PowCounterStore;
31use crate::prelude::*;
32
33type KeyedLimiter = RateLimiter<AddressKey, DashMapStateStore<AddressKey>, DefaultClock>;
35
36struct TierLimiters {
38 short_term: Arc<KeyedLimiter>,
39 long_term: Arc<KeyedLimiter>,
40}
41
42impl TierLimiters {
43 const ONE: NonZeroU32 = match NonZeroU32::new(1) {
45 Some(v) => v,
46 None => unreachable!(),
47 };
48
49 fn new(config: &RateLimitTierConfig) -> Self {
50 let short_quota =
52 Quota::per_second(config.short_term_rps).allow_burst(config.short_term_burst);
53 let short_term = Arc::new(RateLimiter::keyed(short_quota));
54
55 let period_nanos = 3_600_000_000_000_u64 / u64::from(config.long_term_rph.get());
59 let long_quota = Quota::with_period(Duration::from_nanos(period_nanos))
60 .unwrap_or_else(|| Quota::per_second(Self::ONE))
61 .allow_burst(config.long_term_burst);
62 let long_term = Arc::new(RateLimiter::keyed(long_quota));
63
64 Self { short_term, long_term }
65 }
66
67 fn check(&self, key: &AddressKey) -> Result<(), Duration> {
69 if let Err(not_until) = self.short_term.check_key(key) {
71 return Err(not_until.wait_time_from(DefaultClock::default().now()));
72 }
73
74 if let Err(not_until) = self.long_term.check_key(key) {
76 return Err(not_until.wait_time_from(DefaultClock::default().now()));
77 }
78
79 Ok(())
80 }
81}
82
83struct CategoryLimiters {
85 ipv4_individual: TierLimiters,
86 ipv4_network: TierLimiters,
87 ipv6_subnet: TierLimiters,
88 ipv6_provider: TierLimiters,
89}
90
91impl CategoryLimiters {
92 fn new(config: &EndpointCategoryConfig) -> Self {
93 Self {
94 ipv4_individual: TierLimiters::new(&config.ipv4_individual),
95 ipv4_network: TierLimiters::new(&config.ipv4_network),
96 ipv6_subnet: TierLimiters::new(&config.ipv6_subnet),
97 ipv6_provider: TierLimiters::new(&config.ipv6_provider),
98 }
99 }
100
101 fn check(&self, addr: &IpAddr) -> Result<(), RateLimitError> {
103 let keys = AddressKey::extract_all(addr);
104
105 for key in keys {
106 let limiter = self.get_limiter_for_key(&key);
107 if let Err(wait_time) = limiter.check(&key) {
108 return Err(RateLimitError::RateLimited {
109 level: key.level_name(),
110 retry_after: wait_time,
111 });
112 }
113 }
114
115 Ok(())
116 }
117
118 fn get_limiter_for_key(&self, key: &AddressKey) -> &TierLimiters {
119 match key {
120 AddressKey::Ipv4Individual(_) => &self.ipv4_individual,
121 AddressKey::Ipv4Network(_) => &self.ipv4_network,
122 AddressKey::Ipv6Subnet(_) => &self.ipv6_subnet,
123 AddressKey::Ipv6Provider(_) => &self.ipv6_provider,
124 }
125 }
126}
127
128#[derive(Debug, Clone, Default)]
130struct PenaltyEntry {
131 count: u32,
132 last_penalty: Option<Instant>,
133 reason: Option<PenaltyReason>,
134}
135
136pub struct RateLimitManager {
138 categories: HashMap<String, CategoryLimiters>,
140 bans: RwLock<LruCache<AddressKey, BanEntry>>,
142 penalties: RwLock<LruCache<AddressKey, PenaltyEntry>>,
144 pow_store: PowCounterStore,
146 total_limited: AtomicU64,
148 total_bans: AtomicU64,
149}
150
151impl RateLimitManager {
152 const TEN_THOUSAND: NonZeroUsize = match NonZeroUsize::new(10_000) {
154 Some(v) => v,
155 None => unreachable!(),
156 };
157 const TWENTY_THOUSAND: NonZeroUsize = match NonZeroUsize::new(20_000) {
158 Some(v) => v,
159 None => unreachable!(),
160 };
161
162 pub fn new(config: &RateLimitConfig) -> Self {
164 let mut categories = HashMap::new();
165
166 categories.insert("auth".to_string(), CategoryLimiters::new(&config.auth));
168 categories.insert("dav".to_string(), CategoryLimiters::new(&config.dav));
169 categories.insert("federation".to_string(), CategoryLimiters::new(&config.federation));
170 categories.insert("general".to_string(), CategoryLimiters::new(&config.general));
171 categories.insert("search".to_string(), CategoryLimiters::new(&config.search));
172 categories.insert("websocket".to_string(), CategoryLimiters::new(&config.websocket));
173
174 let ban_cap = NonZeroUsize::new(config.max_tracked_ips / 10).unwrap_or(Self::TEN_THOUSAND);
175 let penalty_cap =
176 NonZeroUsize::new(config.max_tracked_ips / 5).unwrap_or(Self::TWENTY_THOUSAND);
177
178 Self {
179 categories,
180 bans: RwLock::new(LruCache::new(ban_cap)),
181 penalties: RwLock::new(LruCache::new(penalty_cap)),
182 pow_store: PowCounterStore::new(PowConfig::default()),
183 total_limited: AtomicU64::new(0),
184 total_bans: AtomicU64::new(0),
185 }
186 }
187
188 pub fn with_pow_config(config: &RateLimitConfig, pow_config: PowConfig) -> Self {
190 let mut manager = Self::new(config);
191 manager.pow_store = PowCounterStore::new(pow_config);
192 manager
193 }
194
195 pub fn check(&self, addr: &IpAddr, category: &str) -> Result<(), RateLimitError> {
197 if let Some(ban) = self.check_ban(addr) {
199 return Err(RateLimitError::Banned { remaining: ban.remaining_duration() });
200 }
201
202 let cat_limiters = self
204 .categories
205 .get(category)
206 .ok_or_else(|| RateLimitError::UnknownCategory(category.to_string()))?;
207
208 if let Err(e) = cat_limiters.check(addr) {
209 self.total_limited.fetch_add(1, Ordering::Relaxed);
210 return Err(e);
211 }
212
213 Ok(())
214 }
215
216 pub fn check_skip_ban(&self, addr: &IpAddr, category: &str) -> Result<(), RateLimitError> {
222 let cat_limiters = self
223 .categories
224 .get(category)
225 .ok_or_else(|| RateLimitError::UnknownCategory(category.to_string()))?;
226
227 if let Err(e) = cat_limiters.check(addr) {
228 self.total_limited.fetch_add(1, Ordering::Relaxed);
229 return Err(e);
230 }
231
232 Ok(())
233 }
234
235 fn check_ban(&self, addr: &IpAddr) -> Option<BanEntry> {
237 let keys = AddressKey::extract_all(addr);
238 let mut bans = self.bans.write();
239
240 for key in keys {
241 if let Some(ban) = bans.get(&key) {
242 if ban.is_expired() {
243 bans.pop(&key);
244 } else {
245 return Some(ban.clone());
246 }
247 }
248 }
249
250 None
251 }
252
253 fn record_penalty(&self, addr: &IpAddr, reason: PenaltyReason, amount: u32) {
255 let key = AddressKey::from_ip_individual(addr);
256 let mut penalties = self.penalties.write();
257
258 let entry = penalties.get_or_insert_mut(key.clone(), PenaltyEntry::default);
259 entry.count = entry.count.saturating_add(amount);
260 entry.last_penalty = Some(Instant::now());
261 entry.reason = Some(reason);
262
263 if entry.count >= reason.failures_to_ban() {
265 drop(penalties);
266 if let Err(e) = self.ban(addr, reason.ban_duration(), reason) {
267 warn!("Failed to auto-ban address: {}", e);
268 }
269 }
270 }
271}
272
273impl Default for RateLimitManager {
274 fn default() -> Self {
275 Self::new(&RateLimitConfig::default())
276 }
277}
278
279impl RateLimitApi for RateLimitManager {
280 fn get_status(
281 &self,
282 addr: &IpAddr,
283 category: &str,
284 ) -> ClResult<Vec<(AddressKey, RateLimitStatus)>> {
285 let _cat_limiters = self.categories.get(category).ok_or(Error::NotFound)?;
286
287 let keys = AddressKey::extract_all(addr);
288 let bans = self.bans.read();
289
290 let statuses = keys
291 .into_iter()
292 .map(|key| {
293 let is_banned = bans.peek(&key).is_some_and(|b| !b.is_expired());
294 let ban_expires = bans.peek(&key).and_then(|b| {
295 if b.is_expired() {
296 None
297 } else {
298 Some(
299 b.expires_at
300 .unwrap_or_else(|| Instant::now() + Duration::from_hours(24 * 365)),
301 )
302 }
303 });
304
305 let status = RateLimitStatus {
306 is_limited: false, remaining: None,
308 reset_at: None,
309 quota: 0,
310 is_banned,
311 ban_expires_at: ban_expires,
312 };
313
314 (key, status)
315 })
316 .collect();
317
318 Ok(statuses)
319 }
320
321 fn penalize(&self, addr: &IpAddr, reason: PenaltyReason, amount: u32) -> ClResult<()> {
322 debug!("Penalizing {:?} for {:?} (amount: {})", addr, reason, amount);
323 self.record_penalty(addr, reason, amount);
324 Ok(())
325 }
326
327 fn grant(&self, addr: &IpAddr, amount: u32) -> ClResult<()> {
328 let key = AddressKey::from_ip_individual(addr);
329 let mut penalties = self.penalties.write();
330
331 if let Some(entry) = penalties.get_mut(&key) {
332 entry.count = entry.count.saturating_sub(amount);
333 if entry.count == 0 {
334 penalties.pop(&key);
335 }
336 }
337
338 Ok(())
339 }
340
341 fn reset(&self, addr: &IpAddr) -> ClResult<()> {
342 let keys = AddressKey::extract_all(addr);
343
344 let mut penalties = self.penalties.write();
346 for key in &keys {
347 penalties.pop(key);
348 }
349 drop(penalties);
350
351 let mut bans = self.bans.write();
353 for key in &keys {
354 bans.pop(key);
355 }
356
357 self.pow_store.decrement(addr, u32::MAX);
359
360 Ok(())
361 }
362
363 fn ban(&self, addr: &IpAddr, duration: Duration, reason: PenaltyReason) -> ClResult<()> {
364 let keys = AddressKey::extract_all(addr);
365 let now = Instant::now();
366 let expires_at = Some(now + duration);
367
368 let mut bans = self.bans.write();
369 for key in keys {
370 let entry = BanEntry { key: key.clone(), reason, created_at: now, expires_at };
371 bans.put(key, entry);
372 }
373
374 self.total_bans.fetch_add(1, Ordering::Relaxed);
375 debug!("Banned {:?} for {:?} due to {:?}", addr, duration, reason);
376
377 Ok(())
378 }
379
380 fn unban(&self, addr: &IpAddr) -> ClResult<()> {
381 let keys = AddressKey::extract_all(addr);
382 let mut bans = self.bans.write();
383
384 for key in keys {
385 bans.pop(&key);
386 }
387
388 Ok(())
389 }
390
391 fn is_banned(&self, addr: &IpAddr) -> bool {
392 self.check_ban(addr).is_some()
393 }
394
395 fn list_bans(&self) -> Vec<BanEntry> {
396 self.bans
397 .read()
398 .iter()
399 .filter(|(_, b)| !b.is_expired())
400 .map(|(_, b)| b.clone())
401 .collect()
402 }
403
404 fn stats(&self) -> RateLimiterStats {
405 let tracked = self
407 .categories
408 .values()
409 .map(|c| {
410 c.ipv4_individual.short_term.len()
411 + c.ipv4_network.short_term.len()
412 + c.ipv6_subnet.short_term.len()
413 + c.ipv6_provider.short_term.len()
414 })
415 .sum();
416
417 RateLimiterStats {
418 tracked_addresses: tracked,
419 active_bans: self.bans.read().len(),
420 total_requests_limited: self.total_limited.load(Ordering::Relaxed),
421 total_bans_issued: self.total_bans.load(Ordering::Relaxed),
422 pow_individual_entries: self.pow_store.individual_count(),
423 pow_network_entries: self.pow_store.network_count(),
424 }
425 }
426
427 fn get_pow_requirement(&self, addr: &IpAddr) -> u32 {
428 self.pow_store.get_requirement(addr)
429 }
430
431 fn increment_pow_counter(&self, addr: &IpAddr, reason: PowPenaltyReason) -> ClResult<()> {
432 self.pow_store.increment(addr, reason);
433 Ok(())
434 }
435
436 fn decrement_pow_counter(&self, addr: &IpAddr, amount: u32) -> ClResult<()> {
437 self.pow_store.decrement(addr, amount);
438 Ok(())
439 }
440
441 fn verify_pow(&self, addr: &IpAddr, token: &str) -> Result<(), PowError> {
442 self.pow_store.verify(addr, token)
443 }
444}
445
446#[cfg(test)]
447#[allow(clippy::unwrap_used, clippy::expect_used)]
448mod tests {
449 use super::*;
450 use std::net::Ipv4Addr;
451
452 #[test]
453 fn test_rate_limit_manager_creation() {
454 let manager = RateLimitManager::default();
455 assert!(manager.categories.contains_key("auth"));
456 assert!(manager.categories.contains_key("federation"));
457 assert!(manager.categories.contains_key("general"));
458 assert!(manager.categories.contains_key("search"));
461 assert!(manager.categories.contains_key("websocket"));
462 }
463
464 #[test]
465 fn test_rate_limit_check() {
466 let manager = RateLimitManager::default();
467 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
468
469 for _ in 0..5 {
471 assert!(manager.check(&ip, "general").is_ok());
472 }
473 }
474
475 #[test]
476 fn test_unknown_category() {
477 let manager = RateLimitManager::default();
478 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
479
480 let result = manager.check(&ip, "nonexistent");
481 assert!(matches!(result, Err(RateLimitError::UnknownCategory(_))));
482 }
483
484 #[test]
485 fn test_ban_functionality() {
486 let manager = RateLimitManager::default();
487 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
488
489 assert!(!manager.is_banned(&ip));
490
491 manager.ban(&ip, Duration::from_mins(1), PenaltyReason::AuthFailure).unwrap();
492 assert!(manager.is_banned(&ip));
493
494 let result = manager.check(&ip, "general");
495 assert!(matches!(result, Err(RateLimitError::Banned { .. })));
496
497 manager.unban(&ip).unwrap();
498 assert!(!manager.is_banned(&ip));
499 }
500
501 #[test]
502 fn test_penalty_auto_ban() {
503 let manager = RateLimitManager::default();
504 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
505
506 for _ in 0..19 {
508 manager.penalize(&ip, PenaltyReason::AuthFailure, 1).unwrap();
509 assert!(!manager.is_banned(&ip));
510 }
511
512 manager.penalize(&ip, PenaltyReason::AuthFailure, 1).unwrap();
514 assert!(manager.is_banned(&ip));
515 }
516
517 #[test]
518 fn test_pow_integration() {
519 let manager = RateLimitManager::default();
520 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
521
522 assert_eq!(manager.get_pow_requirement(&ip), 0);
524 assert!(manager.verify_pow(&ip, "any_token").is_ok());
525
526 manager
528 .increment_pow_counter(&ip, PowPenaltyReason::ConnSignatureFailure)
529 .unwrap();
530 assert_eq!(manager.get_pow_requirement(&ip), 1);
531
532 assert!(manager.verify_pow(&ip, "any_token").is_err());
534 assert!(manager.verify_pow(&ip, "any_tokenA").is_ok());
535 }
536
537 #[test]
538 fn test_stats() {
539 let manager = RateLimitManager::default();
540 let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
541
542 let stats = manager.stats();
543 assert_eq!(stats.active_bans, 0);
544 assert_eq!(stats.total_bans_issued, 0);
545
546 manager.ban(&ip, Duration::from_mins(1), PenaltyReason::AuthFailure).unwrap();
547
548 let stats = manager.stats();
549 assert!(stats.active_bans > 0);
550 assert_eq!(stats.total_bans_issued, 1);
551 }
552}
553
554