Skip to main content

revolt_database/models/mfa_tickets/ops/
reference.rs

1use std::time::{Duration, SystemTime};
2
3use crate::{AbstractMFATickets, MFATicket, ReferenceDb};
4use iso8601_timestamp::Timestamp;
5use revolt_result::Result;
6use ulid::Ulid;
7
8#[async_trait]
9impl AbstractMFATickets for ReferenceDb {
10    /// Find ticket by token
11    async fn fetch_ticket_by_token(&self, token: &str) -> Result<MFATicket> {
12        let tickets = self.tickets.lock().await;
13        let ticket = tickets
14            .values()
15            .find(|ticket| ticket.token == token)
16            .ok_or_else(|| create_error!(InvalidToken))?;
17
18        if let Ok(ulid) = Ulid::from_string(&ticket.id) {
19            if Timestamp::from(ulid.datetime() + Duration::from_mins(5)) > Timestamp::now_utc() {
20                Ok(ticket.clone())
21            } else {
22                Err(create_error!(InvalidToken))
23            }
24        } else {
25            Err(create_error!(InvalidToken))
26        }
27    }
28
29    /// Save ticket
30    async fn save_ticket(&self, ticket: &MFATicket) -> Result<()> {
31        let mut tickets = self.tickets.lock().await;
32        tickets.insert(ticket.id.to_string(), ticket.clone());
33        Ok(())
34    }
35
36    /// Delete ticket
37    async fn delete_ticket(&self, id: &str) -> Result<()> {
38        let mut tickets = self.tickets.lock().await;
39        if tickets.remove(id).is_some() {
40            Ok(())
41        } else {
42            Err(create_error!(InvalidToken))
43        }
44    }
45
46    /// Delete all expired tickets
47    async fn delete_expired_tickets(&self) -> Result<usize> {
48        let threshhold =
49            Ulid::from_datetime(SystemTime::now() - Duration::from_mins(5)).to_string();
50        let mut tickets = self.tickets.lock().await;
51
52        let before = tickets.len();
53        tickets.retain(|_, ticket| ticket.id >= threshhold);
54
55        Ok(before - tickets.len())
56    }
57}