cairn_mod/server/
limits.rs1use std::collections::HashMap;
20use std::net::IpAddr;
21use std::sync::{Arc, Mutex};
22
23#[derive(Debug)]
25pub struct Limiter {
26 global_cap: usize,
27 per_ip_cap: usize,
28 inner: Mutex<LimiterInner>,
29}
30
31#[derive(Debug, Default)]
32struct LimiterInner {
33 global: usize,
34 per_ip: HashMap<IpAddr, usize>,
35}
36
37pub struct Permit {
39 limiter: Arc<Limiter>,
40 ip: IpAddr,
41}
42
43impl std::fmt::Debug for Permit {
44 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
45 f.debug_struct("Permit").field("ip", &self.ip).finish()
46 }
47}
48
49impl Limiter {
50 pub fn new(global_cap: usize, per_ip_cap: usize) -> Arc<Self> {
54 Arc::new(Self {
55 global_cap,
56 per_ip_cap,
57 inner: Mutex::new(LimiterInner::default()),
58 })
59 }
60
61 pub fn try_acquire(self: &Arc<Self>, ip: IpAddr) -> Option<Permit> {
64 let mut inner = self.inner.lock().expect("limiter mutex poisoned");
65 if inner.global >= self.global_cap {
66 return None;
67 }
68 let per_ip = inner.per_ip.get(&ip).copied().unwrap_or(0);
69 if per_ip >= self.per_ip_cap {
70 return None;
71 }
72 inner.global += 1;
73 *inner.per_ip.entry(ip).or_insert(0) += 1;
74 drop(inner);
75 Some(Permit {
76 limiter: Arc::clone(self),
77 ip,
78 })
79 }
80
81 pub fn snapshot(&self) -> (usize, usize) {
84 let inner = self.inner.lock().expect("limiter mutex poisoned");
85 (inner.global, inner.per_ip.len())
86 }
87}
88
89impl Drop for Permit {
90 fn drop(&mut self) {
91 let mut inner = self.limiter.inner.lock().expect("limiter mutex poisoned");
92 inner.global = inner.global.saturating_sub(1);
93 if let Some(c) = inner.per_ip.get_mut(&self.ip) {
94 *c = c.saturating_sub(1);
95 if *c == 0 {
96 inner.per_ip.remove(&self.ip);
97 }
98 }
99 }
100}
101
102#[cfg(test)]
103mod tests {
104 use super::*;
105
106 fn ip(n: u8) -> IpAddr {
107 IpAddr::from([127, 0, 0, n])
108 }
109
110 #[test]
111 fn releases_on_drop_decrements_counters() {
112 let l = Limiter::new(10, 10);
113 {
114 let _p = l.try_acquire(ip(1)).expect("acquire");
115 assert_eq!(l.snapshot().0, 1);
116 }
117 assert_eq!(l.snapshot().0, 0);
118 }
119
120 #[test]
121 fn global_cap_rejects_excess() {
122 let l = Limiter::new(2, 10);
123 let _a = l.try_acquire(ip(1)).expect("first");
124 let _b = l.try_acquire(ip(2)).expect("second");
125 assert!(l.try_acquire(ip(3)).is_none(), "third must be rejected");
126 }
127
128 #[test]
129 fn per_ip_cap_rejects_excess_from_same_ip() {
130 let l = Limiter::new(100, 2);
131 let _a = l.try_acquire(ip(1)).expect("first");
132 let _b = l.try_acquire(ip(1)).expect("second");
133 assert!(
134 l.try_acquire(ip(1)).is_none(),
135 "third from same ip rejected"
136 );
137 let _c = l.try_acquire(ip(2)).expect("different ip ok");
139 }
140
141 #[test]
142 fn rejection_does_not_increment_counters() {
143 let l = Limiter::new(1, 10);
144 let _a = l.try_acquire(ip(1)).expect("first");
145 assert_eq!(l.snapshot().0, 1);
146 let rejected = l.try_acquire(ip(2));
147 assert!(rejected.is_none());
148 assert_eq!(l.snapshot().0, 1, "rejection must not increment");
149 }
150}