tuwunel_service/login_ratelimit/
mod.rs1use std::{
12 collections::BTreeMap,
13 ops::Bound::{Excluded, Included, Unbounded},
14 sync::{Arc, Mutex},
15 time::{Duration, Instant},
16};
17
18use http::StatusCode;
19use ruma::{
20 UserId,
21 api::error::{ErrorKind, LimitExceededErrorData, RetryAfter},
22};
23use tuwunel_core::{Error, Result, Server, implement, warn};
24
25pub struct Service {
31 server: Arc<Server>,
32 account: Ratelimiter,
33 failed: Ratelimiter,
34}
35
36#[derive(Debug)]
45#[must_use]
46pub struct Reservation {
47 key: String,
48 hold: Hold,
49}
50
51type Ratelimiter = Mutex<Table>;
52
53#[derive(Default)]
60struct Table {
61 buckets: Buckets,
62
63 cursor: String,
65}
66
67type Buckets = BTreeMap<String, (Instant, f64)>;
68
69impl crate::Service for Service {
70 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
71 Ok(Arc::new(Self {
72 server: args.server.clone(),
73 account: Ratelimiter::default(),
74 failed: Ratelimiter::default(),
75 }))
76 }
77
78 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
79}
80
81const RATELIMIT_MAP_CAP: usize = 1 << 16;
86
87const PRUNE_SAMPLE: usize = 64;
89
90const MAX_RETRY_AFTER: Duration = Duration::from_hours(24);
92
93const FULL_WARNING_INTERVAL: Duration = Duration::from_mins(1);
95
96static FAILED_FULL_WARNING: Mutex<Option<Instant>> = Mutex::new(None);
97
98static ACCOUNT_FULL_WARNING: Mutex<Option<Instant>> = Mutex::new(None);
99
100#[cfg(test)]
101mod tests;
102
103#[derive(Clone, Copy)]
105struct Limit {
106 rate: f64,
107 burst: u32,
108}
109
110impl Limit {
111 fn enabled(self) -> bool { self.burst > 0 && self.rate > 0.0 && self.rate.is_finite() }
114
115 fn burst(self) -> f64 { f64::from(self.burst) }
116}
117
118#[derive(Clone, Copy, Debug, PartialEq, Eq)]
120enum Hold {
121 Held,
123
124 Untracked,
127}
128
129#[derive(Clone, Copy, Debug)]
134enum Axis {
135 Failed,
137
138 Account,
140}
141
142impl Axis {
143 fn when_full(self) -> Result<Hold> {
144 match self {
145 | Self::Failed => Ok(Hold::Untracked),
146 | Self::Account => Err(limit_exceeded(None)),
147 }
148 }
149
150 fn last_full_warning(self) -> &'static Mutex<Option<Instant>> {
151 match self {
152 | Self::Failed => &FAILED_FULL_WARNING,
153 | Self::Account => &ACCOUNT_FULL_WARNING,
154 }
155 }
156}
157
158#[implement(Service)]
168pub fn reserve_login_attempt(&self, user_id: &UserId) -> Result<Reservation> {
169 reserve_at(
170 &self.failed,
171 account_key(user_id),
172 self.failed_limit(),
173 Instant::now(),
174 RATELIMIT_MAP_CAP,
175 )
176}
177
178#[implement(Service)]
184pub fn record_login(&self, reservation: Reservation) -> Result {
185 let Reservation { key, hold } = reservation;
186 let now = Instant::now();
187
188 refund_at(&self.failed, &key, hold, self.failed_limit(), now)?;
189
190 debit_at(&self.account, &key, self.account_limit(), Axis::Account, now, RATELIMIT_MAP_CAP)?;
191
192 Ok(())
193}
194
195#[implement(Service)]
200pub fn refund_login_attempt(&self, reservation: Reservation) -> Result {
201 let Reservation { key, hold } = reservation;
202
203 refund_at(&self.failed, &key, hold, self.failed_limit(), Instant::now())
204}
205
206#[implement(Service)]
207fn failed_limit(&self) -> Limit {
208 let failed = &self.server.config.rate_limiting.login.failed;
209
210 Limit {
211 rate: failed.per_second,
212 burst: failed.burst_count,
213 }
214}
215
216#[implement(Service)]
217fn account_limit(&self) -> Limit {
218 let account = &self.server.config.rate_limiting.login.account;
219
220 Limit {
221 rate: account.per_second,
222 burst: account.burst_count,
223 }
224}
225
226fn account_key(user_id: &UserId) -> String { user_id.as_str().to_lowercase() }
229
230fn reserve_at(
231 table: &Ratelimiter,
232 key: String,
233 limit: Limit,
234 now: Instant,
235 cap: usize,
236) -> Result<Reservation> {
237 let hold = debit_at(table, &key, limit, Axis::Failed, now, cap)?;
238
239 Ok(Reservation { key, hold })
240}
241
242fn debit_at(
248 table: &Ratelimiter,
249 key: &str,
250 limit: Limit,
251 axis: Axis,
252 now: Instant,
253 cap: usize,
254) -> Result<Hold> {
255 if !limit.enabled() {
256 return Ok(Hold::Untracked);
257 }
258
259 let Limit { rate, .. } = limit;
260 let burst = limit.burst();
261 let mut table = table.lock()?;
262 let Table { buckets, cursor } = &mut *table;
263
264 debug_assert!(cap > 0, "rate-limit table cap must be positive");
265
266 let Some(bucket) = buckets.get_mut(key) else {
267 if buckets.len() >= cap {
268 prune_sample(buckets, cursor, rate, burst, now);
269 }
270
271 if buckets.len() >= cap {
272 drop(table);
273 warn_table_full(axis, now);
274 return axis.when_full();
275 }
276
277 buckets.insert(key.to_owned(), (now, burst - 1.0));
278 return Ok(Hold::Held);
279 };
280
281 let (last_time, tokens) = bucket;
282 let refilled = refill(*last_time, *tokens, rate, burst, now);
283
284 if refilled < 1.0 {
285 return Err(limit_exceeded(retry_after(rate, refilled)));
286 }
287
288 *last_time = now;
289 *tokens = refilled - 1.0;
290
291 Ok(Hold::Held)
292}
293
294fn refund_at(table: &Ratelimiter, key: &str, hold: Hold, limit: Limit, now: Instant) -> Result {
300 if hold == Hold::Untracked || !limit.enabled() {
301 return Ok(());
302 }
303
304 let burst = limit.burst();
305 let mut table = table.lock()?;
306 let Table { buckets, .. } = &mut *table;
307
308 let Some((last_time, tokens)) = buckets.get_mut(key) else {
309 return Ok(());
310 };
311
312 let level = burst.min(refill(*last_time, *tokens, limit.rate, burst, now) + 1.0);
313
314 if level >= burst {
315 buckets.remove(key);
316 } else {
317 *last_time = now;
318 *tokens = level;
319 }
320
321 Ok(())
322}
323
324fn refill(last: Instant, tokens: f64, rate: f64, burst: f64, now: Instant) -> f64 {
325 now.saturating_duration_since(last)
326 .as_secs_f64()
327 .mul_add(rate, tokens)
328 .min(burst)
329}
330
331fn prune_sample(buckets: &mut Buckets, cursor: &mut String, rate: f64, burst: f64, now: Instant) {
339 let after = (Excluded(cursor.as_str()), Unbounded);
340 let wrapped = (Unbounded, Included(cursor.as_str()));
341 let sample = buckets
342 .range::<str, _>(after)
343 .chain(buckets.range::<str, _>(wrapped))
344 .take(PRUNE_SAMPLE);
345
346 if let Some((key, _)) = sample.clone().last() {
347 cursor.clone_from(key);
348 }
349
350 let refilled: Vec<String> = sample
351 .filter(|(_, (last, tokens))| refill(*last, *tokens, rate, burst, now) >= burst)
352 .map(|(key, _)| key.clone())
353 .collect();
354
355 for key in refilled {
356 buckets.remove(&key);
357 }
358}
359
360fn warn_table_full(axis: Axis, now: Instant) {
365 let Ok(mut last) = axis.last_full_warning().lock() else {
366 return;
367 };
368
369 if last.is_some_and(|last| now.saturating_duration_since(last) < FULL_WARNING_INTERVAL) {
370 return;
371 }
372
373 *last = Some(now);
374
375 match axis {
376 | Axis::Failed => warn!(
377 table = ?axis,
378 cap = RATELIMIT_MAP_CAP,
379 "Login rate-limit table is full of accounts still being limited; wrong passwords for \
380 accounts not already in it go untracked. Likely a spray of distinct user names."
381 ),
382 | Axis::Account => warn!(
383 table = ?axis,
384 cap = RATELIMIT_MAP_CAP,
385 "Login rate-limit table is full of accounts still being limited; sign-ins for \
386 accounts not already in it are refused."
387 ),
388 }
389}
390
391fn retry_after(rate: f64, tokens: f64) -> Option<Duration> {
394 let secs = ((1.0 - tokens) / rate).ceil();
395
396 Duration::try_from_secs_f64(secs.min(MAX_RETRY_AFTER.as_secs_f64())).ok()
399}
400
401fn limit_exceeded(retry_after: Option<Duration>) -> Error {
402 Error::Request(
403 ErrorKind::LimitExceeded(LimitExceededErrorData {
404 retry_after: retry_after.map(RetryAfter::Delay),
405 }),
406 "Too many login attempts for this account.".into(),
407 StatusCode::TOO_MANY_REQUESTS,
408 )
409}