dig_rpc/middleware/
rate_limit.rs1use std::collections::HashMap;
14use std::time::Instant;
15
16use dig_rpc_protocol::Tier;
17use parking_lot::Mutex;
18
19pub type PeerKey = Vec<u8>;
23
24#[derive(Debug, Clone, Copy)]
26pub struct BucketSpec {
27 pub fill_per_sec: f64,
29 pub capacity: f64,
31}
32
33#[derive(Debug, Clone)]
35pub struct RateLimitConfig {
36 pub buckets: HashMap<Tier, BucketSpec>,
38}
39
40impl RateLimitConfig {
41 pub fn defaults() -> Self {
43 let mut buckets = HashMap::new();
44 buckets.insert(
45 Tier::PublicRead,
46 BucketSpec {
47 fill_per_sec: 50.0,
48 capacity: 100.0,
49 },
50 );
51 buckets.insert(
52 Tier::Peer,
53 BucketSpec {
54 fill_per_sec: 20.0,
55 capacity: 40.0,
56 },
57 );
58 buckets.insert(
59 Tier::Control,
60 BucketSpec {
61 fill_per_sec: 5.0,
62 capacity: 10.0,
63 },
64 );
65 Self { buckets }
66 }
67}
68
69impl Default for RateLimitConfig {
70 fn default() -> Self {
71 Self::defaults()
72 }
73}
74
75#[derive(Debug, Clone, Copy, PartialEq, Eq)]
77pub enum RateLimitOutcome {
78 Allow,
80 Deny {
82 retry_after_secs: u64,
84 },
85}
86
87#[derive(Debug, Clone)]
89pub struct RateLimitState {
90 inner: std::sync::Arc<Mutex<HashMap<(PeerKey, Tier), Bucket>>>,
91 config: std::sync::Arc<RateLimitConfig>,
92}
93
94#[derive(Debug)]
95struct Bucket {
96 tokens: f64,
97 last_refill: Instant,
98}
99
100impl RateLimitState {
101 pub fn new(config: RateLimitConfig) -> Self {
103 Self {
104 inner: std::sync::Arc::new(Mutex::new(HashMap::new())),
105 config: std::sync::Arc::new(config),
106 }
107 }
108
109 pub fn check(&self, peer: &PeerKey, tier: Tier) -> RateLimitOutcome {
111 let Some(spec) = self.config.buckets.get(&tier).copied() else {
112 tracing::warn!(?tier, "rate tier not configured; allowing");
114 return RateLimitOutcome::Allow;
115 };
116
117 let mut g = self.inner.lock();
118 let now = Instant::now();
119 let b = g.entry((peer.clone(), tier)).or_insert(Bucket {
120 tokens: spec.capacity,
121 last_refill: now,
122 });
123 let elapsed = now.duration_since(b.last_refill).as_secs_f64();
124 b.tokens = (b.tokens + spec.fill_per_sec * elapsed).min(spec.capacity);
125 b.last_refill = now;
126
127 if b.tokens >= 1.0 {
128 b.tokens -= 1.0;
129 RateLimitOutcome::Allow
130 } else {
131 let deficit = 1.0 - b.tokens;
132 let wait_s = (deficit / spec.fill_per_sec).ceil() as u64;
133 RateLimitOutcome::Deny {
134 retry_after_secs: wait_s.max(1),
135 }
136 }
137 }
138}
139
140#[cfg(test)]
141mod tests {
142 use super::*;
143
144 #[test]
147 fn first_request_allowed() {
148 let s = RateLimitState::new(RateLimitConfig::defaults());
149 assert_eq!(
150 s.check(&vec![0; 32], Tier::PublicRead),
151 RateLimitOutcome::Allow
152 );
153 }
154
155 #[test]
159 fn exhaust_bucket_denies() {
160 let mut buckets = HashMap::new();
161 buckets.insert(
162 Tier::Control,
163 BucketSpec {
164 fill_per_sec: 1.0,
165 capacity: 3.0,
166 },
167 );
168 let s = RateLimitState::new(RateLimitConfig { buckets });
169 for _ in 0..3 {
170 assert_eq!(
171 s.check(&vec![0; 32], Tier::Control),
172 RateLimitOutcome::Allow
173 );
174 }
175 match s.check(&vec![0; 32], Tier::Control) {
176 RateLimitOutcome::Deny { retry_after_secs } => assert!(retry_after_secs >= 1),
177 _ => panic!("expected Deny"),
178 }
179 }
180
181 #[test]
185 fn buckets_are_per_peer() {
186 let mut buckets = HashMap::new();
187 buckets.insert(
188 Tier::Peer,
189 BucketSpec {
190 fill_per_sec: 1.0,
191 capacity: 2.0,
192 },
193 );
194 let s = RateLimitState::new(RateLimitConfig { buckets });
195 let a = vec![0xAA; 32];
196 let b = vec![0xBB; 32];
197 for _ in 0..2 {
198 assert_eq!(s.check(&a, Tier::Peer), RateLimitOutcome::Allow);
199 }
200 assert!(matches!(
201 s.check(&a, Tier::Peer),
202 RateLimitOutcome::Deny { .. }
203 ));
204 assert_eq!(s.check(&b, Tier::Peer), RateLimitOutcome::Allow);
205 }
206
207 #[test]
210 fn buckets_are_per_tier() {
211 let mut buckets = HashMap::new();
212 buckets.insert(
213 Tier::Control,
214 BucketSpec {
215 fill_per_sec: 1.0,
216 capacity: 1.0,
217 },
218 );
219 buckets.insert(
220 Tier::PublicRead,
221 BucketSpec {
222 fill_per_sec: 1.0,
223 capacity: 1.0,
224 },
225 );
226 let s = RateLimitState::new(RateLimitConfig { buckets });
227 let p = vec![1; 32];
228 assert_eq!(s.check(&p, Tier::Control), RateLimitOutcome::Allow);
229 assert!(matches!(
230 s.check(&p, Tier::Control),
231 RateLimitOutcome::Deny { .. }
232 ));
233 assert_eq!(s.check(&p, Tier::PublicRead), RateLimitOutcome::Allow);
235 }
236
237 #[test]
240 fn unconfigured_tier_allows() {
241 let s = RateLimitState::new(RateLimitConfig {
242 buckets: HashMap::new(),
243 });
244 assert_eq!(
245 s.check(&vec![0; 32], Tier::PublicRead),
246 RateLimitOutcome::Allow
247 );
248 }
249}