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:
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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(())
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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(())
|
||||
}
|
||||
@@ -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()
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user