1use std::{convert::Infallible, future::ready, net::IpAddr, sync::Arc, time::Duration};
9
10use axum::extract::FromRequestParts;
11use governor::{RateLimiter, clock::QuantaClock, state::keyed::DashMapStateStore};
12use mas_config::RateLimitingConfig;
13use mas_data_model::{User, UserEmailAuthentication};
14use ulid::Ulid;
15
16use crate::ClientIp;
17
18#[derive(Debug, Clone, thiserror::Error)]
19pub enum AccountRecoveryLimitedError {
20 #[error("Too many account recovery requests for requester {0}")]
21 Requester(RequesterFingerprint),
22
23 #[error("Too many account recovery requests for e-mail {0}")]
24 Email(String),
25}
26
27#[derive(Debug, Clone, Copy, thiserror::Error)]
28pub enum PasswordCheckLimitedError {
29 #[error("Too many password checks for requester {0}")]
30 Requester(RequesterFingerprint),
31
32 #[error("Too many password checks for user {0}")]
33 User(Ulid),
34}
35
36#[derive(Debug, Clone, thiserror::Error)]
37pub enum RegistrationLimitedError {
38 #[error("Too many account registration requests for requester {0}")]
39 Requester(RequesterFingerprint),
40}
41
42#[derive(Debug, Clone, thiserror::Error)]
43pub enum EmailAuthenticationLimitedError {
44 #[error("Too many email authentication requests for requester {0}")]
45 Requester(RequesterFingerprint),
46
47 #[error("Too many email authentication requests for authentication session {0}")]
48 Authentication(Ulid),
49
50 #[error("Too many email authentication requests for email {0}")]
51 Email(String),
52}
53
54#[derive(Debug, Clone, thiserror::Error)]
55pub enum DeviceCodeLinkLimitedError {
56 #[error("Too many device code link attempts for requester {0}")]
57 Requester(RequesterFingerprint),
58}
59
60#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
62pub struct RequesterFingerprint {
63 ip: Option<IpAddr>,
64}
65
66impl std::fmt::Display for RequesterFingerprint {
67 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
68 if let Some(ip) = self.ip {
69 write!(f, "{ip}")
70 } else {
71 write!(f, "(NO CLIENT IP)")
72 }
73 }
74}
75
76impl RequesterFingerprint {
77 pub const EMPTY: Self = Self { ip: None };
80
81 #[must_use]
83 pub const fn new(ip: IpAddr) -> Self {
84 Self { ip: Some(ip) }
85 }
86}
87
88impl<S: Send + Sync> FromRequestParts<S> for RequesterFingerprint {
89 type Rejection = Infallible;
90
91 fn from_request_parts(
92 parts: &mut axum::http::request::Parts,
93 _state: &S,
94 ) -> impl Future<Output = Result<Self, Self::Rejection>> + Send {
95 let ip = parts.extensions.get::<ClientIp>().and_then(|ip| ip.0);
96
97 let fingerprint = if let Some(ip) = ip {
98 Self::new(ip)
99 } else {
100 tracing::warn!(
103 "Could not infer client IP address for an operation which rate-limits based on IP addresses"
104 );
105 Self::EMPTY
106 };
107
108 ready(Ok(fingerprint))
109 }
110}
111
112#[derive(Debug, Clone)]
114pub struct Limiter {
115 inner: Arc<LimiterInner>,
116}
117
118type KeyedRateLimiter<K> = RateLimiter<K, DashMapStateStore<K>, QuantaClock>;
119
120#[derive(Debug)]
121struct LimiterInner {
122 account_recovery_per_requester: KeyedRateLimiter<RequesterFingerprint>,
123 account_recovery_per_email: KeyedRateLimiter<String>,
124 password_check_for_requester: KeyedRateLimiter<RequesterFingerprint>,
125 password_check_for_user: KeyedRateLimiter<Ulid>,
126 registration_per_requester: KeyedRateLimiter<RequesterFingerprint>,
127 email_authentication_per_requester: KeyedRateLimiter<RequesterFingerprint>,
128 email_authentication_per_email: KeyedRateLimiter<String>,
129 email_authentication_emails_per_session: KeyedRateLimiter<Ulid>,
130 email_authentication_attempt_per_session: KeyedRateLimiter<Ulid>,
131 device_code_link_per_requester: KeyedRateLimiter<RequesterFingerprint>,
132}
133
134impl LimiterInner {
135 fn new(config: &RateLimitingConfig) -> Option<Self> {
136 Some(Self {
137 account_recovery_per_requester: RateLimiter::keyed(
138 config.account_recovery.per_ip.to_quota()?,
139 ),
140 account_recovery_per_email: RateLimiter::keyed(
141 config.account_recovery.per_address.to_quota()?,
142 ),
143 password_check_for_requester: RateLimiter::keyed(config.login.per_ip.to_quota()?),
144 password_check_for_user: RateLimiter::keyed(config.login.per_account.to_quota()?),
145 registration_per_requester: RateLimiter::keyed(config.registration.to_quota()?),
146 email_authentication_per_email: RateLimiter::keyed(
147 config.email_authentication.per_address.to_quota()?,
148 ),
149 email_authentication_per_requester: RateLimiter::keyed(
150 config.email_authentication.per_ip.to_quota()?,
151 ),
152 email_authentication_emails_per_session: RateLimiter::keyed(
153 config.email_authentication.emails_per_session.to_quota()?,
154 ),
155 email_authentication_attempt_per_session: RateLimiter::keyed(
156 config.email_authentication.attempt_per_session.to_quota()?,
157 ),
158 device_code_link_per_requester: RateLimiter::keyed(config.device_code_link.to_quota()?),
159 })
160 }
161}
162
163impl Limiter {
164 #[must_use]
169 pub fn new(config: &RateLimitingConfig) -> Option<Self> {
170 Some(Self {
171 inner: Arc::new(LimiterInner::new(config)?),
172 })
173 }
174
175 pub fn start(&self) {
180 let this = self.clone();
182 tokio::spawn(async move {
183 let mut interval = tokio::time::interval(Duration::from_mins(1));
185 interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
186
187 loop {
188 this.inner.account_recovery_per_email.retain_recent();
190 this.inner.account_recovery_per_requester.retain_recent();
191 this.inner.password_check_for_requester.retain_recent();
192 this.inner.password_check_for_user.retain_recent();
193 this.inner.registration_per_requester.retain_recent();
194 this.inner.email_authentication_per_email.retain_recent();
195 this.inner
196 .email_authentication_per_requester
197 .retain_recent();
198 this.inner
199 .email_authentication_emails_per_session
200 .retain_recent();
201 this.inner
202 .email_authentication_attempt_per_session
203 .retain_recent();
204 this.inner.device_code_link_per_requester.retain_recent();
205
206 interval.tick().await;
207 }
208 });
209 }
210
211 pub fn check_account_recovery(
217 &self,
218 requester: RequesterFingerprint,
219 email_address: &str,
220 ) -> Result<(), AccountRecoveryLimitedError> {
221 self.inner
222 .account_recovery_per_requester
223 .check_key(&requester)
224 .map_err(|_| AccountRecoveryLimitedError::Requester(requester))?;
225
226 let canonical_email = email_address.to_lowercase();
230 self.inner
231 .account_recovery_per_email
232 .check_key(&canonical_email)
233 .map_err(|_| AccountRecoveryLimitedError::Email(canonical_email))?;
234
235 Ok(())
236 }
237
238 pub fn check_password(
244 &self,
245 key: RequesterFingerprint,
246 user: &User,
247 ) -> Result<(), PasswordCheckLimitedError> {
248 self.inner
249 .password_check_for_requester
250 .check_key(&key)
251 .map_err(|_| PasswordCheckLimitedError::Requester(key))?;
252
253 self.inner
254 .password_check_for_user
255 .check_key(&user.id)
256 .map_err(|_| PasswordCheckLimitedError::User(user.id))?;
257
258 Ok(())
259 }
260
261 pub fn check_registration(
267 &self,
268 requester: RequesterFingerprint,
269 ) -> Result<(), RegistrationLimitedError> {
270 self.inner
271 .registration_per_requester
272 .check_key(&requester)
273 .map_err(|_| RegistrationLimitedError::Requester(requester))?;
274
275 Ok(())
276 }
277
278 pub fn check_email_authentication_email(
285 &self,
286 requester: RequesterFingerprint,
287 email: &str,
288 ) -> Result<(), EmailAuthenticationLimitedError> {
289 self.inner
290 .email_authentication_per_requester
291 .check_key(&requester)
292 .map_err(|_| EmailAuthenticationLimitedError::Requester(requester))?;
293
294 let canonical_email = email.to_lowercase();
298 self.inner
299 .email_authentication_per_email
300 .check_key(&canonical_email)
301 .map_err(|_| EmailAuthenticationLimitedError::Email(email.to_owned()))?;
302 Ok(())
303 }
304
305 pub fn check_email_authentication_attempt(
311 &self,
312 authentication: &UserEmailAuthentication,
313 ) -> Result<(), EmailAuthenticationLimitedError> {
314 self.inner
315 .email_authentication_attempt_per_session
316 .check_key(&authentication.id)
317 .map_err(|_| EmailAuthenticationLimitedError::Authentication(authentication.id))
318 }
319
320 pub fn check_email_authentication_send_code(
327 &self,
328 requester: RequesterFingerprint,
329 authentication: &UserEmailAuthentication,
330 ) -> Result<(), EmailAuthenticationLimitedError> {
331 self.check_email_authentication_email(requester, &authentication.email)?;
332 self.inner
333 .email_authentication_emails_per_session
334 .check_key(&authentication.id)
335 .map_err(|_| EmailAuthenticationLimitedError::Authentication(authentication.id))
336 }
337
338 pub fn check_device_code_link(
347 &self,
348 requester: RequesterFingerprint,
349 ) -> Result<(), DeviceCodeLinkLimitedError> {
350 self.inner
351 .device_code_link_per_requester
352 .check_key(&requester)
353 .map_err(|_| DeviceCodeLinkLimitedError::Requester(requester))
354 }
355}
356
357#[cfg(test)]
358mod tests {
359 use mas_data_model::{Clock, UlidExt as _, User, clock::MockClock};
360 use rand::SeedableRng;
361
362 use super::*;
363
364 #[test]
365 fn test_password_check_limiter() {
366 let now = MockClock::default().now();
367 let mut rng = rand_chacha::ChaChaRng::seed_from_u64(42);
368
369 let limiter = Limiter::new(&RateLimitingConfig::default()).unwrap();
370
371 let requesters: [_; 768] = (0..=255)
373 .flat_map(|a| (0..3).map(move |b| RequesterFingerprint::new([a, a, b, b].into())))
374 .collect::<Vec<_>>()
375 .try_into()
376 .unwrap();
377
378 let alice = User {
379 id: Ulid::from_datetime_with_rng(now, &mut rng),
380 username: "alice".to_owned(),
381 sub: "123-456".to_owned(),
382 created_at: now,
383 locked_at: None,
384 deactivated_at: None,
385 can_request_admin: false,
386 is_guest: true,
387 };
388
389 let bob = User {
390 id: Ulid::from_datetime_with_rng(now, &mut rng),
391 username: "bob".to_owned(),
392 sub: "123-456".to_owned(),
393 created_at: now,
394 locked_at: None,
395 deactivated_at: None,
396 can_request_admin: false,
397 is_guest: true,
398 };
399
400 assert!(limiter.check_password(requesters[0], &alice).is_ok());
402 assert!(limiter.check_password(requesters[0], &alice).is_ok());
403 assert!(limiter.check_password(requesters[0], &alice).is_ok());
404
405 assert!(limiter.check_password(requesters[0], &alice).is_err());
407 assert!(limiter.check_password(requesters[0], &bob).is_err());
409
410 assert!(limiter.check_password(requesters[1], &alice).is_ok());
412
413 for requester in requesters.iter().skip(2).take(598) {
416 assert!(limiter.check_password(*requester, &alice).is_ok());
417 assert!(limiter.check_password(*requester, &alice).is_ok());
418 assert!(limiter.check_password(*requester, &alice).is_ok());
419 assert!(limiter.check_password(*requester, &alice).is_err());
420 }
421
422 assert!(limiter.check_password(requesters[600], &alice).is_ok());
425 assert!(limiter.check_password(requesters[601], &alice).is_ok());
426 assert!(limiter.check_password(requesters[602], &alice).is_err());
427
428 assert!(limiter.check_password(requesters[603], &bob).is_ok());
430 }
431
432 #[test]
433 fn test_device_code_link_limiter() {
434 let limiter = Limiter::new(&RateLimitingConfig::default()).unwrap();
435
436 let alice = RequesterFingerprint::new([1, 2, 3, 4].into());
437 let bob = RequesterFingerprint::new([4, 3, 2, 1].into());
438
439 for _ in 0..10 {
441 assert!(limiter.check_device_code_link(alice).is_ok());
442 }
443 assert!(limiter.check_device_code_link(alice).is_err());
444
445 assert!(limiter.check_device_code_link(bob).is_ok());
447 }
448}