Skip to main content

mas_handlers/
rate_limit.rs

1// Copyright 2025, 2026 Element Creations Ltd.
2// Copyright 2024, 2025 New Vector Ltd.
3// Copyright 2024 The Matrix.org Foundation C.I.C.
4//
5// SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
6// Please see LICENSE files in the repository root for full details.
7
8use 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/// Key used to rate limit requests per requester
61#[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    /// An anonymous key with no IP address set. This should not be used in
78    /// production, and we should warn users if we can't find their client IPs.
79    pub const EMPTY: Self = Self { ip: None };
80
81    /// Create a new anonymous key with the given IP address
82    #[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            // If we can't infer the IP address, we'll just use an empty fingerprint and
101            // warn about it
102            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/// Rate limiters for the different operations
113#[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    /// Creates a new `Limiter` based on a `RateLimitingConfig`.
165    ///
166    /// If the config is not valid, returns `None`.
167    /// (This should not happen if the config was validated, though.)
168    #[must_use]
169    pub fn new(config: &RateLimitingConfig) -> Option<Self> {
170        Some(Self {
171            inner: Arc::new(LimiterInner::new(config)?),
172        })
173    }
174
175    /// Start the rate limiter housekeeping task
176    ///
177    /// This task will periodically remove old entries from the rate limiters,
178    /// to make sure we don't build up a huge number of entries in memory.
179    pub fn start(&self) {
180        // Spawn a task that will periodically clean the rate limiters
181        let this = self.clone();
182        tokio::spawn(async move {
183            // Run the task every minute
184            let mut interval = tokio::time::interval(Duration::from_mins(1));
185            interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
186
187            loop {
188                // Call the retain_recent method on each rate limiter
189                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    /// Check if an account recovery can be performed
212    ///
213    /// # Errors
214    ///
215    /// Returns an error if the operation is rate limited.
216    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        // Convert to lowercase to prevent bypassing the limit by enumerating different
227        // case variations.
228        // A case-folding transformation may be more proper.
229        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    /// Check if a password check can be performed
239    ///
240    /// # Errors
241    ///
242    /// Returns an error if the operation is rate limited
243    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    /// Check if an account registration can be performed
262    ///
263    /// # Errors
264    ///
265    /// Returns an error if the operation is rate limited.
266    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    /// Check if an email can be sent to the address for an email
279    /// authentication session
280    ///
281    /// # Errors
282    ///
283    /// Returns an error if the operation is rate limited.
284    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        // Convert to lowercase to prevent bypassing the limit by enumerating different
295        // case variations.
296        // A case-folding transformation may be more proper.
297        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    /// Check if an attempt can be done on an email authentication session
306    ///
307    /// # Errors
308    ///
309    /// Returns an error if the operation is rate limited.
310    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    /// Check if a new authentication code can be sent for an email
321    /// authentication session
322    ///
323    /// # Errors
324    ///
325    /// Returns an error if the operation is rate limited.
326    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    /// Check if a user code can be submitted on the device link page
339    ///
340    /// This protects against brute-forcing the user code of a device code
341    /// grant, as described in RFC 8628 section 5.1.
342    ///
343    /// # Errors
344    ///
345    /// Returns an error if the operation is rate limited.
346    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's create a lot of requesters to test account-level rate limiting
372        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        // Three times the same IP address should be allowed
401        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        // But the fourth time should be rejected
406        assert!(limiter.check_password(requesters[0], &alice).is_err());
407        // Using another user should also be rejected
408        assert!(limiter.check_password(requesters[0], &bob).is_err());
409
410        // Using a different IP address should be allowed, the account isn't locked yet
411        assert!(limiter.check_password(requesters[1], &alice).is_ok());
412
413        // At this point, we consumed 4 cells out of 1800 on alice, let's distribute the
414        // requests with other IPs so that we get rate-limited on the account-level
415        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        // We now have consumed 4+598*3 = 1798 cells on the account, so we should be
423        // rejected soon
424        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        // The other account isn't rate-limited
429        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        // The default burst allowance is 10 attempts per requester
440        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        // Another requester is unaffected
446        assert!(limiter.check_device_code_link(bob).is_ok());
447    }
448}