chore: migrate authifier into codebase (#658)

Co-authored-by: izzy <me@insrt.uk>
Signed-off-by: Zomatree <me@zomatree.live>
Signed-off-by: izzy <me@insrt.uk>
This commit is contained in:
Zomatree
2026-06-21 00:50:06 +01:00
committed by GitHub
co-authored by izzy
parent a7af24b38d
commit d27917b824
145 changed files with 108392 additions and 1189 deletions
@@ -1422,3 +1422,92 @@ impl From<crate::VoiceInformation> for VoiceInformation {
}
}
}
impl From<crate::Account> for AccountInfo {
fn from(item: crate::Account) -> Self {
AccountInfo {
id: item.id,
email: item.email,
}
}
}
impl From<crate::MFATicket> for MFATicket {
fn from(value: crate::MFATicket) -> Self {
MFATicket {
id: value.id,
account_id: value.account_id,
token: value.token,
validated: value.validated,
authorised: value.authorised,
last_totp_code: value.last_totp_code,
}
}
}
impl From<crate::MultiFactorAuthentication> for MultiFactorStatus {
fn from(item: crate::MultiFactorAuthentication) -> Self {
MultiFactorStatus {
// email_otp: item.enable_email_otp,
// trusted_handover: item.enable_trusted_handover,
// email_mfa: item.enable_email_mfa,
totp_mfa: !item.totp_token.is_disabled(),
// security_key_mfa: item.security_key_token.is_some(),
recovery_active: !item.recovery_codes.is_empty(),
..Default::default()
}
}
}
impl From<crate::MFAMethod> for MFAMethod {
fn from(value: crate::MFAMethod) -> Self {
match value {
crate::MFAMethod::Password => MFAMethod::Password,
crate::MFAMethod::Recovery => MFAMethod::Recovery,
crate::MFAMethod::Totp => MFAMethod::Totp,
}
}
}
impl From<crate::Session> for SessionInfo {
fn from(item: crate::Session) -> Self {
SessionInfo {
id: item.id,
name: item.name,
}
}
}
impl From<crate::Session> for Session {
fn from(value: crate::Session) -> Self {
Session {
id: value.id,
user_id: value.user_id,
token: value.token,
name: value.name,
last_seen: value.last_seen,
origin: value.origin,
subscription: value.subscription.map(Into::into),
}
}
}
impl From<crate::WebPushSubscription> for WebPushSubscription {
fn from(value: crate::WebPushSubscription) -> Self {
WebPushSubscription {
endpoint: value.endpoint,
p256dh: value.p256dh,
auth: value.auth,
}
}
}
impl From<WebPushSubscription> for crate::WebPushSubscription {
fn from(value: WebPushSubscription) -> Self {
crate::WebPushSubscription {
endpoint: value.endpoint,
p256dh: value.p256dh,
auth: value.auth,
}
}
}
+42
View File
@@ -0,0 +1,42 @@
use reqwest::Client;
use revolt_config::config;
use revolt_result::Result;
use std::sync::LazyLock;
static CLIENT: LazyLock<Client> = LazyLock::new(Client::new);
#[derive(Serialize, Deserialize)]
struct CaptchaResponse {
success: bool,
}
pub async fn check_captcha(token: Option<&str>) -> Result<()> {
let config = config().await;
if !config.api.security.captcha.hcaptcha_key.is_empty() {
let Some(token) = token else {
return Err(create_error!(CaptchaFailed));
};
let response = CLIENT
.post("https://hcaptcha.com/siteverify")
.form(&[
("secret", config.api.security.captcha.hcaptcha_key.as_str()),
("response", token),
])
.send()
.await
.map_err(|_| create_error!(CaptchaFailed))?
.json::<CaptchaResponse>()
.await
.map_err(|_| create_error!(CaptchaFailed))?;
if response.success {
Ok(())
} else {
Err(create_error!(CaptchaFailed))
}
} else {
Ok(())
}
}
+63
View File
@@ -0,0 +1,63 @@
#[cfg(feature = "mongodb")]
use ::mongodb::{ClientSession, SessionCursor};
use revolt_result::{Result, ToRevoltError};
use serde::Deserialize;
#[derive(Debug)]
#[allow(clippy::large_enum_variant)]
pub enum ChunkedDatabaseGenerator<T> {
#[cfg(feature = "mongodb")]
MongoDb {
session: ClientSession,
cursor: SessionCursor<T>,
},
Reference {
offset: usize,
data: Vec<T>,
},
}
impl<T: for<'d> Deserialize<'d> + Clone> ChunkedDatabaseGenerator<T> {
#[cfg(feature = "mongodb")]
pub fn new_mongo(session: ClientSession, cursor: SessionCursor<T>) -> Self {
Self::MongoDb { session, cursor }
}
pub fn new_reference(data: Vec<T>) -> Self {
Self::Reference { offset: 0, data }
}
pub async fn next(&mut self) -> Result<Option<T>> {
match self {
#[cfg(feature = "mongodb")]
Self::MongoDb { session, cursor } => {
cursor.next(session).await.transpose().to_internal_error()
}
Self::Reference { offset, data } => {
if let Some(value) = data.get(*offset) {
*offset += 1;
Ok(Some(value.clone()))
} else {
Ok(None)
}
}
}
}
pub async fn next_n(&mut self, n: usize) -> Result<Option<Vec<T>>> {
let mut docs = Vec::new();
while docs.len() < n {
if let Some(doc) = self.next().await? {
docs.push(doc);
} else if docs.is_empty() {
return Ok(None);
} else {
break;
}
}
Ok(Some(docs))
}
}
+271
View File
@@ -0,0 +1,271 @@
use std::{collections::HashSet, sync::LazyLock};
use lettre::{
transport::smtp::{authentication::Credentials, client::Tls},
SmtpTransport,
};
use regex::Regex;
use revolt_config::{config, ApiSmtp};
use revolt_result::Result;
static SPLIT: LazyLock<Regex> = LazyLock::new(|| Regex::new("([^@]+)(@.+)").unwrap());
static SYMBOL_RE: LazyLock<Regex> = LazyLock::new(|| Regex::new("\\+.+|\\.").unwrap());
static HANDLEBARS: LazyLock<handlebars::Handlebars<'static>> =
LazyLock::new(handlebars::Handlebars::new);
static REVOLT_SOURCE_LIST: LazyLock<HashSet<String>> = LazyLock::new(|| {
include_str!("../../assets/revolt_source_list.txt")
.split('\n')
.map(|x| x.into())
.collect()
});
/// Strip special characters and aliases from emails
pub fn normalise_email(original: String) -> String {
let split = SPLIT.captures(&original).unwrap();
let mut clean = SYMBOL_RE
.replace_all(split.get(1).unwrap().as_str(), "")
.to_string();
clean.push_str(split.get(2).unwrap().as_str());
clean.to_lowercase()
}
/// Email template
#[derive(Clone)]
pub struct Template {
/// Title of the email
pub title: String,
/// Plain text version of this email
pub text: String,
/// HTML version of this email
pub html: Option<String>,
/// URL to redirect people to from the email
///
/// Use `{{url}}` to fill this field.
///
/// Any given URL will be suffixed with a unique token if applicable.
///
/// e.g. `https://example.com?t=` becomes `https://example.com?t=UNIQUE_CODE`
pub url: String,
}
/// Email templates
#[derive(Clone)]
pub struct Templates {
/// Template for email verification
pub verify: Template,
/// Template for password reset
pub reset: Template,
/// Template for password reset when the account already exists on creation
pub reset_existing: Template,
/// Template for account deletion
pub deletion: Template,
/// Template for suspention
pub suspension: Template,
}
pub async fn email_templates() -> Templates {
let config = config().await;
if std::env::var("TEST_DB").is_ok() {
Templates {
verify: Template {
title: "verify".into(),
text: "[[{{url}}]]".into(),
url: "".into(),
html: None,
},
reset: Template {
title: "reset".into(),
text: "[[{{url}}]]".into(),
url: "".into(),
html: None,
},
reset_existing: Template {
title: "reset_existing".into(),
text: "[[{{url}}]]".into(),
url: "".into(),
html: None,
},
deletion: Template {
title: "deletion".into(),
text: "[[{{url}}]]".into(),
url: "".into(),
html: None,
},
suspension: Template {
title: "suspension".into(),
text: "[[dummy]]".into(),
url: "".into(),
html: None,
},
}
} else if config.production {
Templates {
verify: Template {
title: "Verify your Stoat account.".into(),
text: include_str!("../../templates/verify.txt").into(),
url: format!("{}/login/verify/", config.hosts.app),
html: Some(include_str!("../../templates/verify.html").into()),
},
reset: Template {
title: "Reset your Stoat password.".into(),
text: include_str!("../../templates/reset.txt").into(),
url: format!("{}/login/reset/", config.hosts.app),
html: Some(include_str!("../../templates/reset.html").into()),
},
reset_existing: Template {
title: "You already have a Stoat account, reset your password.".into(),
text: include_str!("../../templates/reset-existing.txt").into(),
url: format!("{}/login/reset/", config.hosts.app),
html: Some(include_str!("../../templates/reset-existing.html").into()),
},
deletion: Template {
title: "Confirm account deletion.".into(),
text: include_str!("../../templates/deletion.txt").into(),
url: format!("{}/delete/", config.hosts.app),
html: Some(include_str!("../../templates/deletion.html").into()),
},
suspension: Template {
title: "Account Suspension".to_string(),
html: Some(include_str!("../../templates/suspension.html").to_owned()),
text: include_str!("../../templates/suspension.txt").to_owned(),
url: Default::default(),
},
}
} else {
Templates {
verify: Template {
title: "Verify your account.".into(),
text: include_str!("../../templates/verify.whitelabel.txt").into(),
url: format!("{}/login/verify/", config.hosts.app),
html: None,
},
reset: Template {
title: "Reset your password.".into(),
text: include_str!("../../templates/reset.whitelabel.txt").into(),
url: format!("{}/login/reset/", config.hosts.app),
html: None,
},
reset_existing: Template {
title: "Reset your password.".into(),
text: include_str!("../../templates/reset.whitelabel.txt").into(),
url: format!("{}/login/reset/", config.hosts.app),
html: None,
},
deletion: Template {
title: "Confirm account deletion.".into(),
text: include_str!("../../templates/deletion.whitelabel.txt").into(),
url: format!("{}/delete/", config.hosts.app),
html: None,
},
suspension: Template {
title: "Account Suspension".to_string(),
text: include_str!("../../templates/suspension.whitelabel.txt").to_owned(),
url: Default::default(),
html: None,
},
}
}
}
/// Create SMTP transport
pub fn create_transport(smtp: &ApiSmtp) -> SmtpTransport {
let relay = if smtp.use_starttls == Some(true) {
SmtpTransport::starttls_relay(&smtp.host).unwrap()
} else {
SmtpTransport::relay(&smtp.host).unwrap()
};
let relay = if let Some(port) = smtp.port {
relay.port(port.try_into().unwrap())
} else {
relay
};
let relay = if smtp.use_tls == Some(false) {
relay.tls(Tls::None)
} else {
relay
};
relay
.credentials(Credentials::new(
smtp.username.clone(),
smtp.password.clone(),
))
.build()
}
/// Render an email template
fn render_template(text: &str, variables: &handlebars::JsonValue) -> Result<String> {
HANDLEBARS
.render_template(text, variables)
.map_err(|_| create_error!(RenderFail))
}
/// Send an email
pub fn send_email(
smtp: &ApiSmtp,
address: String,
template: &Template,
variables: handlebars::JsonValue,
) -> Result<()> {
let m = lettre::Message::builder()
.from(smtp.from_address.parse().expect("valid `smtp_from`"))
.to(address.parse().expect("valid `smtp_to`"))
.subject(template.title.clone());
let m = if let Some(reply_to) = &smtp.reply_to {
m.reply_to(reply_to.parse().expect("valid `smtp_reply_to`"))
} else {
m
};
let text = render_template(&template.text, &variables).expect("valid `template`");
let m = if let Some(html) = &template.html {
m.multipart(lettre::message::MultiPart::alternative_plain_html(
text,
render_template(html, &variables).expect("valid `template`"),
))
} else {
m.body(text)
}
.expect("valid `message`");
use lettre::Transport;
let sender = create_transport(smtp);
match sender.send(&m) {
Ok(_) => Ok(()),
Err(error) => {
error!(
"Failed to send email to {}!\nlettre error: {}",
address, error
);
revolt_config::capture_error(&error);
Err(create_error!(EmailFailed))
}
}
}
pub fn validate_email(email: &str) -> Result<()> {
// Make sure this is an actual email
if !validator::validate_email(email) {
return Err(create_error!(IncorrectData {
with: "email".to_string()
}));
}
// Check if the email is blacklisted
if let Some(domain) = email.split('@').next_back() {
if REVOLT_SOURCE_LIST.contains(&domain.to_string()) {
return Err(create_error!(Blacklisted));
}
}
Ok(())
}
+56
View File
@@ -0,0 +1,56 @@
#[cfg(feature = "rocket-impl")]
pub mod rocket {
use revolt_config::config;
use rocket::Request;
pub fn to_ip(request: &'_ Request<'_>) -> String {
request
.client_ip()
.map(|x| x.to_string())
.unwrap_or_default()
}
/// Find the actual IP of the client
pub async fn to_real_ip(request: &'_ Request<'_>) -> String {
if config().await.api.security.trust_cloudflare {
request
.headers()
.get_one("CF-Connecting-IP")
.map(|x| x.to_string())
.unwrap_or_else(|| to_ip(request))
} else {
to_ip(request)
}
}
}
#[cfg(feature = "axum-impl")]
pub mod axum {
use axum::{
extract::ConnectInfo,
http::request::Parts,
};
use revolt_config::config;
use std::net::SocketAddr;
pub fn to_ip(parts: &Parts) -> String {
parts
.extensions
.get::<ConnectInfo<SocketAddr>>()
.map(|info| info.ip().to_string())
.unwrap_or_default()
}
/// Find the actual IP of the client
pub async fn to_real_ip(parts: &Parts) -> String {
if config().await.api.security.trust_cloudflare {
parts
.headers
.get("CF-Connecting-IP")
.map(|x| x.to_str().unwrap().to_string())
.unwrap_or_else(|| to_ip(parts))
} else {
to_ip(parts)
}
}
}
+7
View File
@@ -1,10 +1,17 @@
pub mod acker;
pub mod bridge;
pub mod bulk_permissions;
pub mod captcha;
pub mod chunked;
pub mod email;
mod funcs;
pub mod idempotency;
pub mod ip;
pub mod password;
pub mod permissions;
pub mod reference;
pub mod shield;
pub mod test_fixtures;
pub use funcs::*;
pub use chunked::ChunkedDatabaseGenerator;
+72
View File
@@ -0,0 +1,72 @@
use reqwest::Client;
use sha1::Digest;
use std::{collections::HashSet, sync::LazyLock};
use revolt_config::config;
use revolt_result::{Result, ToRevoltError};
static CLIENT: LazyLock<Client> = LazyLock::new(Client::new);
static ARGON_CONFIG: LazyLock<argon2::Config<'static>> = LazyLock::new(argon2::Config::default);
static TOP_100K_COMPROMISED: LazyLock<HashSet<String>> = LazyLock::new(|| {
include_str!("../../assets/pwned100k.txt")
.split('\n')
.map(|x| x.into())
.collect()
});
#[derive(Deserialize)]
struct EasyPwnedResult {
secure: bool,
}
/// Hash a password using argon2
pub fn hash_password(plaintext_password: String) -> Result<String> {
argon2::hash_encoded(
plaintext_password.as_bytes(),
nanoid::nanoid!(24).as_bytes(),
&ARGON_CONFIG,
)
.to_internal_error()
}
pub async fn assert_safe(password: &str) -> Result<()> {
// Make sure the password is long enough.
if password.len() < 8 {
return Err(create_error!(ShortPassword));
}
let config = config().await;
if !config.api.security.easypwned.is_empty() {
let mut hasher = sha1::Sha1::new();
hasher.update(password);
let pwd_hash = hasher.finalize();
let result = match CLIENT
.get(format!(
"{}/hash/{pwd_hash:#02x}",
&config.api.security.easypwned
))
.send()
.await
{
Ok(response) => match response.json::<EasyPwnedResult>().await {
Ok(result) => Ok(result.secure),
Err(e) => Err(e),
},
Err(e) => Err(e),
};
if let Err(e) = &result {
revolt_config::capture_error(e);
} else if result.is_ok_and(|b| b) {
return Ok(());
}
};
if TOP_100K_COMPROMISED.contains(password) {
return Err(create_error!(CompromisedPassword));
};
Ok(())
}
+118
View File
@@ -0,0 +1,118 @@
use reqwest::Client;
use revolt_config::config;
use revolt_result::Result;
use serde::{Deserialize, Serialize};
use std::{collections::HashMap, sync::LazyLock};
use crate::util::ip;
static CLIENT: LazyLock<Client> = LazyLock::new(Client::new);
#[derive(Serialize, Deserialize, Default, Debug)]
pub struct ShieldValidationInput {
/// Remote user IP
pub ip: Option<String>,
/// User provided email
pub email: Option<String>,
/// Request headers
pub headers: Option<HashMap<String, String>>,
/// Skip alerts and monitoring for this request
pub dry_run: bool,
}
#[derive(Serialize, Deserialize)]
pub struct ValidationResult {
/// Whether this request was blocked
blocked: bool,
/// Reasons for the request being blocked
reasons: Vec<String>,
}
pub async fn validate_shield(input: ShieldValidationInput) -> Result<()> {
let shield = config().await.api.security.shield;
if !shield.host.is_empty() {
if let Ok(response) = CLIENT
.post(format!("{}/validate", &shield.host))
.json(&input)
.header("Authorization", &shield.key)
.send()
.await
{
let result = response
.json::<ValidationResult>()
.await
.map_err(|_| create_error!(InternalError))?;
if result.blocked {
return Err(create_error!(BlockedByShield));
}
}
}
Ok(())
}
#[cfg(feature = "rocket-impl")]
#[async_trait]
impl<'r> rocket::request::FromRequest<'r> for ShieldValidationInput {
type Error = revolt_result::Error;
#[allow(clippy::collapsible_match)]
async fn from_request(
request: &'r rocket::Request<'_>,
) -> rocket::request::Outcome<Self, Self::Error> {
rocket::request::Outcome::Success(ShieldValidationInput {
ip: Some(ip::rocket::to_real_ip(request).await),
headers: Some(
request
.headers()
.iter()
.map(|entry| (entry.name.to_string(), entry.value.to_string()))
.collect(),
),
..Default::default()
})
}
}
#[cfg(feature = "rocket-impl")]
impl<'r> revolt_rocket_okapi::request::OpenApiFromRequest<'r> for ShieldValidationInput {
fn from_request_input(
_gen: &mut revolt_rocket_okapi::r#gen::OpenApiGenerator,
_name: String,
_required: bool,
) -> revolt_rocket_okapi::Result<revolt_rocket_okapi::request::RequestHeaderInput> {
Ok(revolt_rocket_okapi::request::RequestHeaderInput::None)
}
}
#[cfg(feature = "axum-impl")]
#[async_trait]
impl<S> axum::extract::FromRequestParts<S> for ShieldValidationInput {
type Rejection = axum::Json<revolt_result::Error> ;
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
Ok(ShieldValidationInput {
ip: Some(ip::axum::to_real_ip(parts).await),
headers: Some(
parts
.headers
.iter()
.map(|(name, value)| {
(
name.to_string(),
value.to_str().map(|s| s.to_string()).unwrap_or_default(),
)
})
.collect(),
),
..Default::default()
})
}
}