Skip to main content

cloudillo_core/rate_limit/
limiter.rs

1// SPDX-FileCopyrightText: Szilárd Hajba
2// SPDX-License-Identifier: LGPL-3.0-or-later
3
4//! Rate Limit Manager
5//!
6//! Core rate limiting implementation using the governor crate's GCRA algorithm.
7//! Supports hierarchical address levels with dual-tier (short + long term) limits.
8
9use 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
33/// Type alias for a keyed rate limiter
34type KeyedLimiter = RateLimiter<AddressKey, DashMapStateStore<AddressKey>, DefaultClock>;
35
36/// Holds both short-term and long-term limiters for an address level
37struct TierLimiters {
38	short_term: Arc<KeyedLimiter>,
39	long_term: Arc<KeyedLimiter>,
40}
41
42impl TierLimiters {
43	// SAFETY: 1 is non-zero
44	const ONE: NonZeroU32 = match NonZeroU32::new(1) {
45		Some(v) => v,
46		None => unreachable!(),
47	};
48
49	fn new(config: &RateLimitTierConfig) -> Self {
50		// Short-term: per-second with burst
51		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		// Long-term: per-hour with burst
56		// Convert RPH to nanosecond period using integer math:
57		// period_nanos = 3_600_000_000_000 / rph
58		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	/// Check if both short and long term limits allow the request
68	fn check(&self, key: &AddressKey) -> Result<(), Duration> {
69		// Check short-term first
70		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		// Check long-term
75		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
83/// Per-category rate limiters (one for each hierarchical level)
84struct 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	/// Check all applicable limits for an address
102	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/// Penalty tracking for an address
129#[derive(Debug, Clone, Default)]
130struct PenaltyEntry {
131	count: u32,
132	last_penalty: Option<Instant>,
133	reason: Option<PenaltyReason>,
134}
135
136/// Main rate limit manager
137pub struct RateLimitManager {
138	/// Per-category limiters
139	categories: HashMap<String, CategoryLimiters>,
140	/// Global ban list
141	bans: RwLock<LruCache<AddressKey, BanEntry>>,
142	/// Penalty tracking per address
143	penalties: RwLock<LruCache<AddressKey, PenaltyEntry>>,
144	/// Proof-of-work counter store
145	pow_store: PowCounterStore,
146	/// Statistics
147	total_limited: AtomicU64,
148	total_bans: AtomicU64,
149}
150
151impl RateLimitManager {
152	// SAFETY: These are non-zero constants
153	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	/// Create a new rate limit manager
163	pub fn new(config: &RateLimitConfig) -> Self {
164		let mut categories = HashMap::new();
165
166		// Initialize category limiters
167		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("websocket".to_string(), CategoryLimiters::new(&config.websocket));
172
173		let ban_cap = NonZeroUsize::new(config.max_tracked_ips / 10).unwrap_or(Self::TEN_THOUSAND);
174		let penalty_cap =
175			NonZeroUsize::new(config.max_tracked_ips / 5).unwrap_or(Self::TWENTY_THOUSAND);
176
177		Self {
178			categories,
179			bans: RwLock::new(LruCache::new(ban_cap)),
180			penalties: RwLock::new(LruCache::new(penalty_cap)),
181			pow_store: PowCounterStore::new(PowConfig::default()),
182			total_limited: AtomicU64::new(0),
183			total_bans: AtomicU64::new(0),
184		}
185	}
186
187	/// Create with custom PoW config
188	pub fn with_pow_config(config: &RateLimitConfig, pow_config: PowConfig) -> Self {
189		let mut manager = Self::new(config);
190		manager.pow_store = PowCounterStore::new(pow_config);
191		manager
192	}
193
194	/// Check if a request should be rate limited
195	pub fn check(&self, addr: &IpAddr, category: &str) -> Result<(), RateLimitError> {
196		// Check ban list first
197		if let Some(ban) = self.check_ban(addr) {
198			return Err(RateLimitError::Banned { remaining: ban.remaining_duration() });
199		}
200
201		// Check rate limits
202		let cat_limiters = self
203			.categories
204			.get(category)
205			.ok_or_else(|| RateLimitError::UnknownCategory(category.to_string()))?;
206
207		if let Err(e) = cat_limiters.check(addr) {
208			self.total_limited.fetch_add(1, Ordering::Relaxed);
209			return Err(e);
210		}
211
212		Ok(())
213	}
214
215	/// Check rate limits for a request WITHOUT consulting the global ban list.
216	///
217	/// Behaves exactly like [`Self::check`] minus the ban gate. Used by routes
218	/// that must remain reachable from a banned IP (e.g. the password-recovery
219	/// flow) while still being subject to the normal per-category rate limit.
220	pub fn check_skip_ban(&self, addr: &IpAddr, category: &str) -> Result<(), RateLimitError> {
221		let cat_limiters = self
222			.categories
223			.get(category)
224			.ok_or_else(|| RateLimitError::UnknownCategory(category.to_string()))?;
225
226		if let Err(e) = cat_limiters.check(addr) {
227			self.total_limited.fetch_add(1, Ordering::Relaxed);
228			return Err(e);
229		}
230
231		Ok(())
232	}
233
234	/// Check if address is banned
235	fn check_ban(&self, addr: &IpAddr) -> Option<BanEntry> {
236		let keys = AddressKey::extract_all(addr);
237		let mut bans = self.bans.write();
238
239		for key in keys {
240			if let Some(ban) = bans.get(&key) {
241				if ban.is_expired() {
242					bans.pop(&key);
243				} else {
244					return Some(ban.clone());
245				}
246			}
247		}
248
249		None
250	}
251
252	/// Record a penalty for an address
253	fn record_penalty(&self, addr: &IpAddr, reason: PenaltyReason, amount: u32) {
254		let key = AddressKey::from_ip_individual(addr);
255		let mut penalties = self.penalties.write();
256
257		let entry = penalties.get_or_insert_mut(key.clone(), PenaltyEntry::default);
258		entry.count = entry.count.saturating_add(amount);
259		entry.last_penalty = Some(Instant::now());
260		entry.reason = Some(reason);
261
262		// Check for auto-ban
263		if entry.count >= reason.failures_to_ban() {
264			drop(penalties);
265			if let Err(e) = self.ban(addr, reason.ban_duration(), reason) {
266				warn!("Failed to auto-ban address: {}", e);
267			}
268		}
269	}
270}
271
272impl Default for RateLimitManager {
273	fn default() -> Self {
274		Self::new(&RateLimitConfig::default())
275	}
276}
277
278impl RateLimitApi for RateLimitManager {
279	fn get_status(
280		&self,
281		addr: &IpAddr,
282		category: &str,
283	) -> ClResult<Vec<(AddressKey, RateLimitStatus)>> {
284		let _cat_limiters = self.categories.get(category).ok_or(Error::NotFound)?;
285
286		let keys = AddressKey::extract_all(addr);
287		let bans = self.bans.read();
288
289		let statuses = keys
290			.into_iter()
291			.map(|key| {
292				let is_banned = bans.peek(&key).is_some_and(|b| !b.is_expired());
293				let ban_expires = bans.peek(&key).and_then(|b| {
294					if b.is_expired() {
295						None
296					} else {
297						Some(
298							b.expires_at
299								.unwrap_or_else(|| Instant::now() + Duration::from_hours(24 * 365)),
300						)
301					}
302				});
303
304				let status = RateLimitStatus {
305					is_limited: false, // Would need to check governor state
306					remaining: None,
307					reset_at: None,
308					quota: 0,
309					is_banned,
310					ban_expires_at: ban_expires,
311				};
312
313				(key, status)
314			})
315			.collect();
316
317		Ok(statuses)
318	}
319
320	fn penalize(&self, addr: &IpAddr, reason: PenaltyReason, amount: u32) -> ClResult<()> {
321		debug!("Penalizing {:?} for {:?} (amount: {})", addr, reason, amount);
322		self.record_penalty(addr, reason, amount);
323		Ok(())
324	}
325
326	fn grant(&self, addr: &IpAddr, amount: u32) -> ClResult<()> {
327		let key = AddressKey::from_ip_individual(addr);
328		let mut penalties = self.penalties.write();
329
330		if let Some(entry) = penalties.get_mut(&key) {
331			entry.count = entry.count.saturating_sub(amount);
332			if entry.count == 0 {
333				penalties.pop(&key);
334			}
335		}
336
337		Ok(())
338	}
339
340	fn reset(&self, addr: &IpAddr) -> ClResult<()> {
341		let keys = AddressKey::extract_all(addr);
342
343		// Clear penalties
344		let mut penalties = self.penalties.write();
345		for key in &keys {
346			penalties.pop(key);
347		}
348		drop(penalties);
349
350		// Clear bans
351		let mut bans = self.bans.write();
352		for key in &keys {
353			bans.pop(key);
354		}
355
356		// Clear PoW counters
357		self.pow_store.decrement(addr, u32::MAX);
358
359		Ok(())
360	}
361
362	fn ban(&self, addr: &IpAddr, duration: Duration, reason: PenaltyReason) -> ClResult<()> {
363		let keys = AddressKey::extract_all(addr);
364		let now = Instant::now();
365		let expires_at = Some(now + duration);
366
367		let mut bans = self.bans.write();
368		for key in keys {
369			let entry = BanEntry { key: key.clone(), reason, created_at: now, expires_at };
370			bans.put(key, entry);
371		}
372
373		self.total_bans.fetch_add(1, Ordering::Relaxed);
374		debug!("Banned {:?} for {:?} due to {:?}", addr, duration, reason);
375
376		Ok(())
377	}
378
379	fn unban(&self, addr: &IpAddr) -> ClResult<()> {
380		let keys = AddressKey::extract_all(addr);
381		let mut bans = self.bans.write();
382
383		for key in keys {
384			bans.pop(&key);
385		}
386
387		Ok(())
388	}
389
390	fn is_banned(&self, addr: &IpAddr) -> bool {
391		self.check_ban(addr).is_some()
392	}
393
394	fn list_bans(&self) -> Vec<BanEntry> {
395		self.bans
396			.read()
397			.iter()
398			.filter(|(_, b)| !b.is_expired())
399			.map(|(_, b)| b.clone())
400			.collect()
401	}
402
403	fn stats(&self) -> RateLimiterStats {
404		// Count tracked addresses across all categories
405		let tracked = self
406			.categories
407			.values()
408			.map(|c| {
409				c.ipv4_individual.short_term.len()
410					+ c.ipv4_network.short_term.len()
411					+ c.ipv6_subnet.short_term.len()
412					+ c.ipv6_provider.short_term.len()
413			})
414			.sum();
415
416		RateLimiterStats {
417			tracked_addresses: tracked,
418			active_bans: self.bans.read().len(),
419			total_requests_limited: self.total_limited.load(Ordering::Relaxed),
420			total_bans_issued: self.total_bans.load(Ordering::Relaxed),
421			pow_individual_entries: self.pow_store.individual_count(),
422			pow_network_entries: self.pow_store.network_count(),
423		}
424	}
425
426	fn get_pow_requirement(&self, addr: &IpAddr) -> u32 {
427		self.pow_store.get_requirement(addr)
428	}
429
430	fn increment_pow_counter(&self, addr: &IpAddr, reason: PowPenaltyReason) -> ClResult<()> {
431		self.pow_store.increment(addr, reason);
432		Ok(())
433	}
434
435	fn decrement_pow_counter(&self, addr: &IpAddr, amount: u32) -> ClResult<()> {
436		self.pow_store.decrement(addr, amount);
437		Ok(())
438	}
439
440	fn verify_pow(&self, addr: &IpAddr, token: &str) -> Result<(), PowError> {
441		self.pow_store.verify(addr, token)
442	}
443}
444
445#[cfg(test)]
446#[allow(clippy::unwrap_used, clippy::expect_used)]
447mod tests {
448	use super::*;
449	use std::net::Ipv4Addr;
450
451	#[test]
452	fn test_rate_limit_manager_creation() {
453		let manager = RateLimitManager::default();
454		assert!(manager.categories.contains_key("auth"));
455		assert!(manager.categories.contains_key("federation"));
456		assert!(manager.categories.contains_key("general"));
457		assert!(manager.categories.contains_key("websocket"));
458	}
459
460	#[test]
461	fn test_rate_limit_check() {
462		let manager = RateLimitManager::default();
463		let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
464
465		// First few requests should pass
466		for _ in 0..5 {
467			assert!(manager.check(&ip, "general").is_ok());
468		}
469	}
470
471	#[test]
472	fn test_unknown_category() {
473		let manager = RateLimitManager::default();
474		let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
475
476		let result = manager.check(&ip, "nonexistent");
477		assert!(matches!(result, Err(RateLimitError::UnknownCategory(_))));
478	}
479
480	#[test]
481	fn test_ban_functionality() {
482		let manager = RateLimitManager::default();
483		let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
484
485		assert!(!manager.is_banned(&ip));
486
487		manager.ban(&ip, Duration::from_mins(1), PenaltyReason::AuthFailure).unwrap();
488		assert!(manager.is_banned(&ip));
489
490		let result = manager.check(&ip, "general");
491		assert!(matches!(result, Err(RateLimitError::Banned { .. })));
492
493		manager.unban(&ip).unwrap();
494		assert!(!manager.is_banned(&ip));
495	}
496
497	#[test]
498	fn test_penalty_auto_ban() {
499		let manager = RateLimitManager::default();
500		let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
501
502		// AuthFailure requires 20 failures for auto-ban
503		for _ in 0..19 {
504			manager.penalize(&ip, PenaltyReason::AuthFailure, 1).unwrap();
505			assert!(!manager.is_banned(&ip));
506		}
507
508		// 20th failure should trigger auto-ban
509		manager.penalize(&ip, PenaltyReason::AuthFailure, 1).unwrap();
510		assert!(manager.is_banned(&ip));
511	}
512
513	#[test]
514	fn test_pow_integration() {
515		let manager = RateLimitManager::default();
516		let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
517
518		// Initially no PoW required
519		assert_eq!(manager.get_pow_requirement(&ip), 0);
520		assert!(manager.verify_pow(&ip, "any_token").is_ok());
521
522		// Increment counter
523		manager
524			.increment_pow_counter(&ip, PowPenaltyReason::ConnSignatureFailure)
525			.unwrap();
526		assert_eq!(manager.get_pow_requirement(&ip), 1);
527
528		// Now need PoW
529		assert!(manager.verify_pow(&ip, "any_token").is_err());
530		assert!(manager.verify_pow(&ip, "any_tokenA").is_ok());
531	}
532
533	#[test]
534	fn test_stats() {
535		let manager = RateLimitManager::default();
536		let ip = IpAddr::V4(Ipv4Addr::new(192, 168, 1, 100));
537
538		let stats = manager.stats();
539		assert_eq!(stats.active_bans, 0);
540		assert_eq!(stats.total_bans_issued, 0);
541
542		manager.ban(&ip, Duration::from_mins(1), PenaltyReason::AuthFailure).unwrap();
543
544		let stats = manager.stats();
545		assert!(stats.active_bans > 0);
546		assert_eq!(stats.total_bans_issued, 1);
547	}
548}
549
550// vim: ts=4