tuwunel_service/threepid/
ratelimit.rs1use std::{borrow::Borrow, hash::Hash, net::IpAddr, time::Instant};
7
8use http::StatusCode;
9use ruma::api::error::{ErrorKind, LimitExceededErrorData};
10use tuwunel_core::{Error, Result, implement};
11
12use super::{EmailKey, Ratelimiter};
13
14const RC_PER_SECOND: f64 = 0.2;
17const RC_BURST: f64 = 5.0;
18
19const RATELIMIT_MAP_CAP: usize = 1 << 16;
22
23#[cfg(test)]
24mod tests;
25
26#[implement(super::Service)]
31pub fn check_ip_rate_limit(&self, client: IpAddr) -> Result {
32 check_bucket(&self.ip_ratelimiter, &client, || client, RC_PER_SECOND, RC_BURST)
33}
34
35#[implement(super::Service)]
40pub fn check_address_rate_limit(&self, address: &str) -> Result {
41 check_bucket(
42 &self.address_ratelimiter,
43 address,
44 || EmailKey::from(address),
45 RC_PER_SECOND,
46 RC_BURST,
47 )
48}
49
50fn check_bucket<K, Q>(
51 table: &Ratelimiter<K>,
52 key: &Q,
53 make_key: impl FnOnce() -> K,
54 rate: f64,
55 burst: f64,
56) -> Result
57where
58 K: Borrow<Q> + Clone + Eq + Hash,
59 Q: Eq + Hash + ?Sized,
60{
61 check_bucket_at(table, key, make_key, rate, burst, Instant::now(), RATELIMIT_MAP_CAP)
62}
63
64fn check_bucket_at<K, Q>(
65 table: &Ratelimiter<K>,
66 key: &Q,
67 make_key: impl FnOnce() -> K,
68 rate: f64,
69 burst: f64,
70 now: Instant,
71 cap: usize,
72) -> Result
73where
74 K: Borrow<Q> + Clone + Eq + Hash,
75 Q: Eq + Hash + ?Sized,
76{
77 let mut buckets = table.lock()?;
78 debug_assert!(cap > 0, "rate-limit table cap must be positive");
79 debug_assert!(buckets.len() <= cap, "rate-limit table exceeded its cap");
80
81 if let Some(bucket) = buckets.get_mut(key) {
82 return debit_bucket(bucket, rate, burst, now);
83 }
84
85 if buckets.len() >= cap {
86 let mut oldest = None;
87
88 buckets.retain(|key, bucket| {
89 let (last, toks) = *bucket;
90 let refilled = now
91 .duration_since(last)
92 .as_secs_f64()
93 .mul_add(rate, toks);
94
95 let retain = refilled < burst;
96
97 if retain
98 && oldest
99 .as_ref()
100 .is_none_or(|(_, oldest_at)| last < *oldest_at)
101 {
102 oldest = Some((key.clone(), last));
103 }
104
105 retain
106 });
107
108 if buckets.len() >= cap
109 && let Some((oldest, _)) = oldest
110 {
111 buckets.remove::<K>(&oldest);
112 }
113 }
114
115 let bucket = buckets
116 .entry(make_key())
117 .or_insert_with(|| (now, burst));
118
119 debit_bucket(bucket, rate, burst, now)
120}
121
122fn debit_bucket(bucket: &mut (Instant, f64), rate: f64, burst: f64, now: Instant) -> Result {
123 let (last_time, tokens) = bucket;
124 let new_tokens = now
125 .duration_since(*last_time)
126 .as_secs_f64()
127 .mul_add(rate, *tokens)
128 .min(burst);
129
130 if new_tokens < 1.0 {
131 return Err(Error::Request(
132 ErrorKind::LimitExceeded(LimitExceededErrorData { retry_after: None }),
133 "Too many verification requests.".into(),
134 StatusCode::TOO_MANY_REQUESTS,
135 ));
136 }
137
138 *last_time = now;
139 *tokens = new_tokens - 1.0;
140
141 Ok(())
142}