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("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	/// Create with custom PoW config
189	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	/// Check if a request should be rate limited
196	pub fn check(&self, addr: &IpAddr, category: &str) -> Result<(), RateLimitError> {
197		// Check ban list first
198		if let Some(ban) = self.check_ban(addr) {
199			return Err(RateLimitError::Banned { remaining: ban.remaining_duration() });
200		}
201
202		// Check rate limits
203		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	/// Check rate limits for a request WITHOUT consulting the global ban list.
217	///
218	/// Behaves exactly like [`Self::check`] minus the ban gate. Used by routes
219	/// that must remain reachable from a banned IP (e.g. the password-recovery
220	/// flow) while still being subject to the normal per-category rate limit.
221	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	/// Check if address is banned
236	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	/// Record a penalty for an address
254	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		// Check for auto-ban
264		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, // Would need to check governor state
307					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		// Clear penalties
345		let mut penalties = self.penalties.write();
346		for key in &keys {
347			penalties.pop(key);
348		}
349		drop(penalties);
350
351		// Clear bans
352		let mut bans = self.bans.write();
353		for key in &keys {
354			bans.pop(key);
355		}
356
357		// Clear PoW counters
358		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		// Count tracked addresses across all categories
406		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		// A `RateLimitLayer` naming a missing bucket 500s every request; `/api/search`
459		// names this one.
460		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		// First few requests should pass
470		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		// AuthFailure requires 20 failures for auto-ban
507		for _ in 0..19 {
508			manager.penalize(&ip, PenaltyReason::AuthFailure, 1).unwrap();
509			assert!(!manager.is_banned(&ip));
510		}
511
512		// 20th failure should trigger auto-ban
513		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		// Initially no PoW required
523		assert_eq!(manager.get_pow_requirement(&ip), 0);
524		assert!(manager.verify_pow(&ip, "any_token").is_ok());
525
526		// Increment counter
527		manager
528			.increment_pow_counter(&ip, PowPenaltyReason::ConnSignatureFailure)
529			.unwrap();
530		assert_eq!(manager.get_pow_requirement(&ip), 1);
531
532		// Now need PoW
533		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// vim: ts=4