Compare commits

...
11 Commits
43 changed files with 424 additions and 79 deletions
+3
View File
@@ -59,6 +59,9 @@ REVOLT_UNSAFE_NO_EMAIL=1
## Application Settings ## Application Settings
## ##
# Whether to enable staging only features
REVOLT_IS_STAGING=1
# Whether to only allow users to sign up if they have an invite code # Whether to only allow users to sign up if they have an invite code
REVOLT_INVITE_ONLY=0 REVOLT_INVITE_ONLY=0
Generated
+41 -8
View File
@@ -806,6 +806,12 @@ dependencies = [
"uuid", "uuid",
] ]
[[package]]
name = "decancer"
version = "1.6.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "808127a7de612079ec37bfc1abc48ed77a6015a971a8bd7d4178d79147cbc839"
[[package]] [[package]]
name = "derivative" name = "derivative"
version = "2.2.0" version = "2.2.0"
@@ -2490,6 +2496,20 @@ dependencies = [
"yansi", "yansi",
] ]
[[package]]
name = "prometheus"
version = "0.13.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "449811d15fbdf5ceb5c1144416066429cf82316e2ec8ce0c1f6f8a02e7bbcf8c"
dependencies = [
"cfg-if 1.0.0",
"fnv",
"lazy_static",
"memchr",
"parking_lot",
"thiserror",
]
[[package]] [[package]]
name = "querystring" name = "querystring"
version = "1.1.0" version = "1.1.0"
@@ -2837,7 +2857,7 @@ dependencies = [
[[package]] [[package]]
name = "revolt-bonfire" name = "revolt-bonfire"
version = "0.6.0" version = "0.6.5"
dependencies = [ dependencies = [
"async-std", "async-std",
"async-tungstenite", "async-tungstenite",
@@ -2854,7 +2874,7 @@ dependencies = [
[[package]] [[package]]
name = "revolt-database" name = "revolt-database"
version = "0.6.0" version = "0.6.5"
dependencies = [ dependencies = [
"async-recursion", "async-recursion",
"async-std", "async-std",
@@ -2885,7 +2905,7 @@ dependencies = [
[[package]] [[package]]
name = "revolt-delta" name = "revolt-delta"
version = "0.6.0" version = "0.6.5"
dependencies = [ dependencies = [
"async-channel", "async-channel",
"async-std", "async-std",
@@ -2914,6 +2934,7 @@ dependencies = [
"rocket", "rocket",
"rocket_authifier", "rocket_authifier",
"rocket_empty", "rocket_empty",
"rocket_prometheus",
"schemars", "schemars",
"serde", "serde",
"serde_json", "serde_json",
@@ -2925,7 +2946,7 @@ dependencies = [
[[package]] [[package]]
name = "revolt-models" name = "revolt-models"
version = "0.6.0" version = "0.6.5"
dependencies = [ dependencies = [
"revolt-permissions", "revolt-permissions",
"revolt_optional_struct", "revolt_optional_struct",
@@ -2936,7 +2957,7 @@ dependencies = [
[[package]] [[package]]
name = "revolt-permissions" name = "revolt-permissions"
version = "0.6.0" version = "0.6.5"
dependencies = [ dependencies = [
"async-std", "async-std",
"async-trait", "async-trait",
@@ -2950,7 +2971,7 @@ dependencies = [
[[package]] [[package]]
name = "revolt-presence" name = "revolt-presence"
version = "0.6.0" version = "0.6.5"
dependencies = [ dependencies = [
"async-std", "async-std",
"log", "log",
@@ -2961,7 +2982,7 @@ dependencies = [
[[package]] [[package]]
name = "revolt-quark" name = "revolt-quark"
version = "0.6.0" version = "0.6.5"
dependencies = [ dependencies = [
"async-lock", "async-lock",
"async-recursion", "async-recursion",
@@ -2974,6 +2995,7 @@ dependencies = [
"bson", "bson",
"dashmap", "dashmap",
"deadqueue", "deadqueue",
"decancer",
"dotenv", "dotenv",
"futures", "futures",
"impl_ops", "impl_ops",
@@ -2991,6 +3013,7 @@ dependencies = [
"redis-kiss", "redis-kiss",
"regex", "regex",
"reqwest", "reqwest",
"revolt-database",
"revolt-models", "revolt-models",
"revolt-presence", "revolt-presence",
"revolt-result", "revolt-result",
@@ -3012,7 +3035,7 @@ dependencies = [
[[package]] [[package]]
name = "revolt-result" name = "revolt-result"
version = "0.6.0" version = "0.6.5"
dependencies = [ dependencies = [
"revolt_okapi", "revolt_okapi",
"revolt_rocket_okapi", "revolt_rocket_okapi",
@@ -3235,6 +3258,16 @@ dependencies = [
"uncased", "uncased",
] ]
[[package]]
name = "rocket_prometheus"
version = "0.10.0-rc.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e3e70efcf6b0234723d0d295b95d64283ed02bb98a8a3f0c8c7ebb5b8da69165"
dependencies = [
"prometheus",
"rocket",
]
[[package]] [[package]]
name = "rust-argon2" name = "rust-argon2"
version = "1.0.0" version = "1.0.0"
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "revolt-bonfire" name = "revolt-bonfire"
version = "0.6.0" version = "0.6.5"
license = "AGPL-3.0-or-later" license = "AGPL-3.0-or-later"
edition = "2021" edition = "2021"
+5 -5
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "revolt-database" name = "revolt-database"
version = "0.6.0" version = "0.6.5"
edition = "2021" edition = "2021"
license = "AGPL-3.0-or-later" license = "AGPL-3.0-or-later"
authors = [ "Paul Makles <me@insrt.uk>" ] authors = [ "Paul Makles <me@insrt.uk>" ]
@@ -22,10 +22,10 @@ default = [ "mongodb", "async-std-runtime" ]
[dependencies] [dependencies]
# Core # Core
revolt-result = { version = "0.6.0", path = "../result" } revolt-result = { version = "0.6.5", path = "../result" }
revolt-models = { version = "0.6.0", path = "../models" } revolt-models = { version = "0.6.5", path = "../models" }
revolt-presence = { version = "0.6.0", path = "../presence" } revolt-presence = { version = "0.6.5", path = "../presence" }
revolt-permissions = { version = "0.6.0", path = "../permissions", features = [ "serde", "bson" ] } revolt-permissions = { version = "0.6.5", path = "../permissions", features = [ "serde", "bson" ] }
# Utility # Utility
log = "0.4" log = "0.4"
@@ -76,6 +76,10 @@ pub async fn create_database(db: &MongoDb) {
.await .await
.expect("Failed to create bots collection."); .expect("Failed to create bots collection.");
db.create_collection("ratelimit_events", None)
.await
.expect("Failed to create ratelimit_events collection.");
db.create_collection( db.create_collection(
"pubsub", "pubsub",
CreateCollectionOptions::builder() CreateCollectionOptions::builder()
@@ -209,5 +213,24 @@ pub async fn create_database(db: &MongoDb) {
.await .await
.expect("Failed to save migration info."); .expect("Failed to save migration info.");
db.run_command(
doc! {
"createIndexes": "ratelimit_events",
"indexes": [
{
"key": {
"_id": 1_i32,
"target_id": 1_i32,
"event_type": 1_i32,
},
"name": "compound_key"
}
]
},
None,
)
.await
.expect("Failed to create ratelimit_events index.");
info!("Created database."); info!("Created database.");
} }
@@ -9,6 +9,7 @@ use crate::{
}; };
use futures::StreamExt; use futures::StreamExt;
use rand::seq::SliceRandom; use rand::seq::SliceRandom;
use revolt_permissions::DEFAULT_WEBHOOK_PERMISSIONS;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use unicode_segmentation::UnicodeSegmentation; use unicode_segmentation::UnicodeSegmentation;
@@ -18,7 +19,7 @@ struct MigrationInfo {
revision: i32, revision: i32,
} }
pub const LATEST_REVISION: i32 = 25; pub const LATEST_REVISION: i32 = 26;
pub async fn migrate_database(db: &MongoDb) { pub async fn migrate_database(db: &MongoDb) {
let migrations = db.col::<Document>("migrations"); let migrations = db.col::<Document>("migrations");
@@ -945,10 +946,56 @@ pub async fn run_migrations(db: &MongoDb, revision: i32) -> i32 {
) )
.await .await
.expect("Failed to create username index."); .expect("Failed to create username index.");
};
if revision <= 25 {
info!("Running migration [revision 25 / 11-06-2023]: Add permissions to webhooks.");
db.col::<Document>("webhooks")
.update_many(
doc! {},
doc! {
"$set": {
"permissions": *DEFAULT_WEBHOOK_PERMISSIONS as i64
}
},
None,
)
.await
.expect("Failed to update webhooks.");
}
if revision <= 25 {
info!("Running migration [revision 25 / 15-06-2023]: Add collection `ratelimit_events` with index.");
db.db()
.create_collection("ratelimit_events", None)
.await
.ok();
db.db()
.run_command(
doc! {
"createIndexes": "ratelimit_events",
"indexes": [
{
"key": {
"_id": 1_i32,
"target_id": 1_i32,
"event_type": 1_i32,
},
"name": "compound_key"
}
]
},
None,
)
.await
.expect("Failed to create ratelimit_events index.");
} }
// Need to migrate fields on attachments, change `user_id`, `object_id`, etc to `parent`. // Need to migrate fields on attachments, change `user_id`, `object_id`, etc to `parent`.
// Reminder to update LATEST_REVISION when adding new migrations. // Reminder to update LATEST_REVISION when adding new migrations.
LATEST_REVISION LATEST_REVISION.max(revision)
} }
@@ -20,6 +20,9 @@ auto_derived_partial!(
/// The channel this webhook belongs to /// The channel this webhook belongs to
pub channel_id: String, pub channel_id: String,
/// The permissions of the webhook
pub permissions: u64,
/// The private token for the webhook /// The private token for the webhook
pub token: Option<String>, pub token: Option<String>,
}, },
+3
View File
@@ -3,6 +3,7 @@ mod bots;
mod channel_webhooks; mod channel_webhooks;
mod channels; mod channels;
mod files; mod files;
mod ratelimit_events;
mod safety_strikes; mod safety_strikes;
mod server_members; mod server_members;
mod servers; mod servers;
@@ -14,6 +15,7 @@ pub use bots::*;
pub use channel_webhooks::*; pub use channel_webhooks::*;
pub use channels::*; pub use channels::*;
pub use files::*; pub use files::*;
pub use ratelimit_events::*;
pub use safety_strikes::*; pub use safety_strikes::*;
pub use server_members::*; pub use server_members::*;
pub use servers::*; pub use servers::*;
@@ -30,6 +32,7 @@ pub trait AbstractDatabase:
+ channels::AbstractChannels + channels::AbstractChannels
+ channel_webhooks::AbstractWebhooks + channel_webhooks::AbstractWebhooks
+ files::AbstractAttachments + files::AbstractAttachments
+ ratelimit_events::AbstractRatelimitEvents
+ safety_strikes::AbstractAccountStrikes + safety_strikes::AbstractAccountStrikes
+ server_members::AbstractServerMembers + server_members::AbstractServerMembers
+ servers::AbstractServers + servers::AbstractServers
@@ -0,0 +1,5 @@
mod model;
mod ops;
pub use model::*;
pub use ops::*;
@@ -0,0 +1,25 @@
use std::fmt;
auto_derived!(
/// Ratelimit Event
pub struct RatelimitEvent {
/// Id
#[serde(rename = "_id")]
pub id: String,
/// Relevant Object Id
pub target_id: String,
/// Type of event
pub event_type: RatelimitEventType,
}
/// Event type
pub enum RatelimitEventType {
DiscriminatorChange,
}
);
impl fmt::Display for RatelimitEventType {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
fmt::Debug::fmt(self, f)
}
}
@@ -0,0 +1,20 @@
use std::time::Duration;
use crate::{revolt_result::Result, RatelimitEvent, RatelimitEventType};
mod mongodb;
mod reference;
#[async_trait]
pub trait AbstractRatelimitEvents: Sync + Send {
/// Insert a new ratelimit event
async fn insert_ratelimit_event(&self, event: &RatelimitEvent) -> Result<()>;
/// Count number of events in given duration and check if we've hit the limit
async fn has_ratelimited(
&self,
target_id: &str,
event_type: RatelimitEventType,
period: Duration,
count: usize,
) -> Result<bool>;
}
@@ -0,0 +1,40 @@
use std::time::{Duration, SystemTime};
use super::AbstractRatelimitEvents;
use crate::{MongoDb, RatelimitEvent, RatelimitEventType};
use revolt_result::Result;
use ulid::Ulid;
static COL: &str = "ratelimit_events";
#[async_trait]
impl AbstractRatelimitEvents for MongoDb {
/// Insert a new ratelimit event
async fn insert_ratelimit_event(&self, event: &RatelimitEvent) -> Result<()> {
query!(self, insert_one, COL, &event).map(|_| ())
}
/// Count number of events in given duration and check if we've hit the limit
async fn has_ratelimited(
&self,
target_id: &str,
event_type: RatelimitEventType,
period: Duration,
count: usize,
) -> Result<bool> {
self.col::<RatelimitEvent>(COL)
.count_documents(
doc! {
"_id": {
"$gte": Ulid::from_datetime(SystemTime::now() - period).to_string()
},
"target_id": target_id,
"event_type": event_type.to_string()
},
None,
)
.await
.map(|c| c as usize >= count)
.map_err(|_| create_database_error!("count_documents", COL))
}
}
@@ -0,0 +1,28 @@
use std::time::Duration;
use super::AbstractRatelimitEvents;
use crate::RatelimitEvent;
use crate::RatelimitEventType;
use crate::ReferenceDb;
use revolt_result::Result;
#[async_trait]
impl AbstractRatelimitEvents for ReferenceDb {
/// Insert a new ratelimit event
async fn insert_ratelimit_event(&self, _event: &RatelimitEvent) -> Result<()> {
// TODO: implement
unimplemented!()
}
/// Count number of events in given duration and check if we've hit the limit
async fn has_ratelimited(
&self,
_target_id: &str,
_event_type: RatelimitEventType,
_period: Duration,
_count: usize,
) -> Result<bool> {
// TODO: implement
unimplemented!()
}
}
@@ -52,6 +52,7 @@ impl From<crate::Webhook> for Webhook {
avatar: value.avatar.map(|file| file.into()), avatar: value.avatar.map(|file| file.into()),
channel_id: value.channel_id, channel_id: value.channel_id,
token: value.token, token: value.token,
permissions: value.permissions
} }
} }
} }
@@ -64,6 +65,7 @@ impl From<crate::PartialWebhook> for PartialWebhook {
avatar: value.avatar.map(|file| file.into()), avatar: value.avatar.map(|file| file.into()),
channel_id: value.channel_id, channel_id: value.channel_id,
token: value.token, token: value.token,
permissions: value.permissions
} }
} }
} }
+2 -2
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "revolt-models" name = "revolt-models"
version = "0.6.0" version = "0.6.5"
edition = "2021" edition = "2021"
license = "AGPL-3.0-or-later" license = "AGPL-3.0-or-later"
authors = [ "Paul Makles <me@insrt.uk>" ] authors = [ "Paul Makles <me@insrt.uk>" ]
@@ -18,7 +18,7 @@ default = [ "serde", "partials" ]
[dependencies] [dependencies]
# Core # Core
revolt-permissions = { version = "0.6.0", path = "../permissions" } revolt-permissions = { version = "0.6.5", path = "../permissions" }
# Serialisation # Serialisation
revolt_optional_struct = { version = "0.2.0", optional = true } revolt_optional_struct = { version = "0.2.0", optional = true }
@@ -16,6 +16,9 @@ auto_derived_partial!(
/// The channel this webhook belongs to /// The channel this webhook belongs to
pub channel_id: String, pub channel_id: String,
/// The permissions for the webhook
pub permissions: u64,
/// The private token for the webhook /// The private token for the webhook
pub token: Option<String>, pub token: Option<String>,
}, },
@@ -43,6 +46,9 @@ auto_derived!(
#[cfg_attr(feature = "validator", validate(length(min = 1, max = 128)))] #[cfg_attr(feature = "validator", validate(length(min = 1, max = 128)))]
pub avatar: Option<String>, pub avatar: Option<String>,
/// Webhook permissions
pub permissions: Option<u64>,
/// Fields to remove from webhook /// Fields to remove from webhook
#[cfg_attr(feature = "serde", serde(default))] #[cfg_attr(feature = "serde", serde(default))]
pub remove: Vec<FieldsWebhook>, pub remove: Vec<FieldsWebhook>,
@@ -61,6 +67,9 @@ auto_derived!(
/// The channel this webhook belongs to /// The channel this webhook belongs to
pub channel_id: String, pub channel_id: String,
/// The permissions for the webhook
pub permissions: u64
} }
/// Optional fields on webhook object /// Optional fields on webhook object
@@ -85,6 +94,7 @@ impl From<Webhook> for ResponseWebhook {
name: value.name, name: value.name,
avatar: value.avatar.map(|file| file.id), avatar: value.avatar.map(|file| file.id),
channel_id: value.channel_id, channel_id: value.channel_id,
permissions: value.permissions
} }
} }
} }
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "revolt-permissions" name = "revolt-permissions"
version = "0.6.0" version = "0.6.5"
edition = "2021" edition = "2021"
license = "AGPL-3.0-or-later" license = "AGPL-3.0-or-later"
authors = [ "Paul Makles <me@insrt.uk>" ] authors = [ "Paul Makles <me@insrt.uk>" ]
@@ -135,3 +135,5 @@ pub static DEFAULT_PERMISSION_SERVER: Lazy<u64> = Lazy::new(|| {
+ ChannelPermission::ChangeAvatar, + ChannelPermission::ChangeAvatar,
) )
}); });
pub static DEFAULT_WEBHOOK_PERMISSIONS: Lazy<u64> = Lazy::new(|| ChannelPermission::SendMessage + ChannelPermission::SendEmbeds + ChannelPermission::Masquerade + ChannelPermission::React);
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "revolt-presence" name = "revolt-presence"
version = "0.6.0" version = "0.6.5"
edition = "2021" edition = "2021"
license = "AGPL-3.0-or-later" license = "AGPL-3.0-or-later"
authors = [ "Paul Makles <me@insrt.uk>" ] authors = [ "Paul Makles <me@insrt.uk>" ]
+1 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "revolt-result" name = "revolt-result"
version = "0.6.0" version = "0.6.5"
edition = "2021" edition = "2021"
license = "AGPL-3.0-or-later" license = "AGPL-3.0-or-later"
authors = [ "Paul Makles <me@insrt.uk>" ] authors = [ "Paul Makles <me@insrt.uk>" ]
+2 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "revolt-delta" name = "revolt-delta"
version = "0.6.0" version = "0.6.5"
license = "AGPL-3.0-or-later" license = "AGPL-3.0-or-later"
authors = ["Paul Makles <paulmakles@gmail.com>"] authors = ["Paul Makles <paulmakles@gmail.com>"]
edition = "2018" edition = "2018"
@@ -47,6 +47,7 @@ lettre = "0.10.0-alpha.4"
rocket = { version = "0.5.0-rc.2", default-features = false, features = ["json"] } rocket = { version = "0.5.0-rc.2", default-features = false, features = ["json"] }
rocket_empty = { version = "0.1.1", features = ["schema"] } rocket_empty = { version = "0.1.1", features = ["schema"] }
rocket_authifier = { version = "1.0.7" } rocket_authifier = { version = "1.0.7" }
rocket_prometheus = "0.10.0-rc.3"
# spec generation # spec generation
schemars = "0.8.8" schemars = "0.8.8"
+5
View File
@@ -8,6 +8,7 @@ extern crate serde_json;
pub mod routes; pub mod routes;
pub mod util; pub mod util;
use rocket_prometheus::PrometheusMetrics;
use std::net::Ipv4Addr; use std::net::Ipv4Addr;
use async_std::channel::unbounded; use async_std::channel::unbounded;
@@ -65,7 +66,11 @@ async fn rocket() -> _ {
// Configure Rocket // Configure Rocket
let rocket = rocket::build(); let rocket = rocket::build();
let prometheus = PrometheusMetrics::new();
routes::mount(rocket) routes::mount(rocket)
.attach(prometheus.clone())
.mount("/metrics", prometheus)
.mount("/", revolt_quark::web::cors::catch_all_options_routes()) .mount("/", revolt_quark::web::cors::catch_all_options_routes())
.mount("/", revolt_quark::web::ratelimiter::routes()) .mount("/", revolt_quark::web::ratelimiter::routes())
.mount("/swagger/", revolt_quark::web::swagger::routes()) .mount("/swagger/", revolt_quark::web::swagger::routes())
@@ -2,6 +2,7 @@ use revolt_database::{Database, Webhook};
use revolt_quark::{ use revolt_quark::{
models::{Channel, User}, models::{Channel, User},
perms, Db, Error, Permission, Ref, Result, perms, Db, Error, Permission, Ref, Result,
DEFAULT_WEBHOOK_PERMISSIONS,
}; };
use rocket::{serde::json::Json, State}; use rocket::{serde::json::Json, State};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
@@ -60,6 +61,7 @@ pub async fn req(
name: data.name, name: data.name,
avatar, avatar,
channel_id: channel.id().to_string(), channel_id: channel.id().to_string(),
permissions: *DEFAULT_WEBHOOK_PERMISSIONS,
token: Some(nanoid::nanoid!(64)), token: Some(nanoid::nanoid!(64)),
}; };
@@ -14,6 +14,8 @@ struct BannedUser {
pub id: String, pub id: String,
/// Username of the banned user /// Username of the banned user
pub username: String, pub username: String,
/// Discriminator of the banned user
pub discriminator: String,
/// Avatar of the banned user /// Avatar of the banned user
pub avatar: Option<File>, pub avatar: Option<File>,
} }
@@ -32,6 +34,7 @@ impl From<User> for BannedUser {
BannedUser { BannedUser {
id: user.id, id: user.id,
username: user.username, username: user.username,
discriminator: user.discriminator,
avatar: user.avatar, avatar: user.avatar,
} }
} }
@@ -13,7 +13,7 @@ pub async fn req(db: &Db, user: User, target: Ref, role_id: String) -> Result<Em
.throw_permission(db, Permission::ManageRole) .throw_permission(db, Permission::ManageRole)
.await?; .await?;
let member_rank = permissions.get_member_rank().unwrap_or(0); let member_rank = permissions.get_member_rank().unwrap_or(i64::MIN);
if let Some(role) = server.roles.remove(&role_id) { if let Some(role) = server.roles.remove(&role_id) {
if role.rank <= member_rank { if role.rank <= member_rank {
+1 -1
View File
@@ -27,7 +27,6 @@ pub struct UserProfileData {
#[derive(Validate, Serialize, Deserialize, JsonSchema)] #[derive(Validate, Serialize, Deserialize, JsonSchema)]
pub struct DataEditUser { pub struct DataEditUser {
/// New display name /// New display name
#[serde(rename = "displayName")]
#[validate(length(min = 2, max = 32), regex = "RE_DISPLAY_NAME")] #[validate(length(min = 2, max = 32), regex = "RE_DISPLAY_NAME")]
display_name: Option<String>, display_name: Option<String>,
/// Attachment Id for avatar /// Attachment Id for avatar
@@ -123,6 +122,7 @@ pub async fn req(
} }
let mut partial: PartialUser = PartialUser { let mut partial: PartialUser = PartialUser {
display_name: data.display_name,
badges: data.badges, badges: data.badges,
flags: data.flags, flags: data.flags,
..Default::default() ..Default::default()
@@ -35,11 +35,13 @@ pub async fn webhook_edit(
let DataEditWebhook { let DataEditWebhook {
name, name,
avatar, avatar,
permissions,
remove, remove,
} = data; } = data;
let mut partial = PartialWebhook { let mut partial = PartialWebhook {
name, name,
permissions,
..Default::default() ..Default::default()
}; };
@@ -33,11 +33,13 @@ pub async fn webhook_edit_token(
let DataEditWebhook { let DataEditWebhook {
name, name,
avatar, avatar,
remove, permissions,
remove
} = data; } = data;
let mut partial = PartialWebhook { let mut partial = PartialWebhook {
name, name,
permissions,
..Default::default() ..Default::default()
}; };
@@ -29,8 +29,7 @@ pub async fn webhook_execute(
let webhook = webhook_id.as_webhook(db).await.map_err(Error::from_core)?; let webhook = webhook_id.as_webhook(db).await.map_err(Error::from_core)?;
webhook.assert_token(&token).map_err(Error::from_core)?; webhook.assert_token(&token).map_err(Error::from_core)?;
// TODO: webhooks can currently always send masquerades, files, embeds, reactions (interactions) data.validate_webhook_permissions(webhook.permissions)?;
// TODO: they can also mention anyone
let channel = legacy_db.fetch_channel(&webhook.channel_id).await?; let channel = legacy_db.fetch_channel(&webhook.channel_id).await?;
let message = channel let message = channel
+3 -1
View File
@@ -1,6 +1,6 @@
[package] [package]
name = "revolt-quark" name = "revolt-quark"
version = "0.6.0" version = "0.6.5"
edition = "2021" edition = "2021"
license = "AGPL-3.0-or-later" license = "AGPL-3.0-or-later"
@@ -64,6 +64,7 @@ nanoid = "0.4.0"
linkify = "0.8.1" linkify = "0.8.1"
dotenv = "0.15.0" dotenv = "0.15.0"
indexmap = "1.9.1" indexmap = "1.9.1"
decancer = "1.6.2"
impl_ops = "0.1.1" impl_ops = "0.1.1"
num_enum = "0.5.6" num_enum = "0.5.6"
reqwest = "0.11.10" reqwest = "0.11.10"
@@ -93,4 +94,5 @@ sentry = "0.25.0"
# Core # Core
revolt-result = { path = "../core/result", features = [ "serde", "schemas" ] } revolt-result = { path = "../core/result", features = [ "serde", "schemas" ] }
revolt-presence = { path = "../core/presence", features = [ "redis-is-patched" ] } revolt-presence = { path = "../core/presence", features = [ "redis-is-patched" ] }
revolt-database = { path = "../core/database" }
revolt-models = { path = "../core/models" } revolt-models = { path = "../core/models" }
+11
View File
@@ -71,3 +71,14 @@ impl From<Database> for authifier::Database {
} }
} }
} }
impl From<Database> for revolt_database::Database {
fn from(val: Database) -> Self {
match val {
Database::Dummy(_) => revolt_database::Database::Reference(Default::default()),
Database::MongoDb(MongoDb(client)) => revolt_database::Database::MongoDb(
revolt_database::MongoDb(client, "revolt".to_string()),
),
}
}
}
+19 -10
View File
@@ -23,8 +23,7 @@ impl Cache {
pub async fn can_view_channel(&self, db: &Database, channel: &Channel) -> bool { pub async fn can_view_channel(&self, db: &Database, channel: &Channel) -> bool {
match &channel { match &channel {
Channel::TextChannel { server, .. } | Channel::VoiceChannel { server, .. } => { Channel::TextChannel { server, .. } | Channel::VoiceChannel { server, .. } => {
let member = self.members.values().find(|x| &x.id.server == server); let member = self.members.get(server);
let server = self.servers.get(server); let server = self.servers.get(server);
let mut perms = perms(self.users.get(&self.user_id).unwrap()).channel(channel); let mut perms = perms(self.users.get(&self.user_id).unwrap()).channel(channel);
@@ -107,9 +106,15 @@ impl State {
// Fetch all memberships with their corresponding servers. // Fetch all memberships with their corresponding servers.
let members: Vec<Member> = db.fetch_all_memberships(&user.id).await?; let members: Vec<Member> = db.fetch_all_memberships(&user.id).await?;
self.cache.members = members
.iter()
.cloned()
.map(|x| (x.id.server.clone(), x))
.collect();
let server_ids: Vec<String> = members.iter().map(|x| x.id.server.clone()).collect(); let server_ids: Vec<String> = members.iter().map(|x| x.id.server.clone()).collect();
let servers = db.fetch_servers(&server_ids).await?; let servers = db.fetch_servers(&server_ids).await?;
self.cache.servers = servers.iter().cloned().map(|x| (x.id.clone(), x)).collect();
// Collect channel ids from servers. // Collect channel ids from servers.
let mut channel_ids = vec![]; let mut channel_ids = vec![];
@@ -164,17 +169,11 @@ impl State {
self.cache self.cache
.users .users
.insert(self.cache.user_id.clone(), user.clone()); .insert(self.cache.user_id.clone(), user.clone());
self.cache.servers = servers.iter().cloned().map(|x| (x.id.clone(), x)).collect();
self.cache.channels = channels self.cache.channels = channels
.iter() .iter()
.cloned() .cloned()
.map(|x| (x.id().to_string(), x)) .map(|x| (x.id().to_string(), x))
.collect(); .collect();
self.cache.members = members
.iter()
.cloned()
.map(|x| (x.id.server.clone(), x))
.collect();
// Make all users appear from our perspective. // Make all users appear from our perspective.
let mut users: Vec<User> = users let mut users: Vec<User> = users
@@ -353,7 +352,7 @@ impl State {
let could_view: bool = if let Some(channel) = self.cache.channels.get(id) { let could_view: bool = if let Some(channel) = self.cache.channels.get(id) {
self.cache.can_view_channel(db, channel).await self.cache.can_view_channel(db, channel).await
} else { } else {
true false
}; };
if let Some(channel) = self.cache.channels.get_mut(id) { if let Some(channel) = self.cache.channels.get_mut(id) {
@@ -364,6 +363,12 @@ impl State {
channel.apply_options(data.clone()); channel.apply_options(data.clone());
} }
if !self.cache.channels.contains_key(id) {
if let Ok(channel) = db.fetch_channel(id).await {
self.cache.channels.insert(id.clone(), channel);
}
}
if let Some(channel) = self.cache.channels.get(id) { if let Some(channel) = self.cache.channels.get(id) {
let can_view = self.cache.can_view_channel(db, channel).await; let can_view = self.cache.can_view_channel(db, channel).await;
if could_view != can_view { if could_view != can_view {
@@ -398,7 +403,9 @@ impl State {
channels, channels,
} => { } => {
self.insert_subscription(id.clone()); self.insert_subscription(id.clone());
self.cache.servers.insert(id.to_string(), server.clone()); self.cache.servers.insert(id.clone(), server.clone());
let member = Member::new(id.clone(), self.cache.user_id.clone());
self.cache.members.insert(id.clone(), member);
for channel in channels { for channel in channels {
self.cache self.cache
@@ -436,6 +443,7 @@ impl State {
self.cache.channels.remove(channel); self.cache.channels.remove(channel);
} }
} }
self.cache.members.remove(id);
} }
} }
EventV1::ServerDelete { id } => { EventV1::ServerDelete { id } => {
@@ -447,6 +455,7 @@ impl State {
self.cache.channels.remove(channel); self.cache.channels.remove(channel);
} }
} }
self.cache.members.remove(id);
} }
EventV1::ServerMemberUpdate { id, data, clear } => { EventV1::ServerMemberUpdate { id, data, clear } => {
if id.user == self.cache.user_id { if id.user == self.cache.user_id {
@@ -3,22 +3,10 @@ use crate::{AbstractServerMember, Result};
use super::super::DummyDb; use super::super::DummyDb;
use iso8601_timestamp::Timestamp;
#[async_trait] #[async_trait]
impl AbstractServerMember for DummyDb { impl AbstractServerMember for DummyDb {
async fn fetch_member(&self, server: &str, user: &str) -> Result<Member> { async fn fetch_member(&self, server: &str, user: &str) -> Result<Member> {
Ok(Member { Ok(Member::new(server.into(), user.into()))
id: MemberCompositeKey {
server: server.into(),
user: user.into(),
},
joined_at: Timestamp::now_utc(),
nickname: None,
avatar: None,
roles: vec![],
timeout: None,
})
} }
async fn insert_member(&self, member: &Member) -> Result<()> { async fn insert_member(&self, member: &Member) -> Result<()> {
@@ -10,7 +10,7 @@ use crate::{
models::{ models::{
message::{ message::{
AppendMessage, BulkMessageResponse, Interactions, PartialMessage, SendableEmbed, AppendMessage, BulkMessageResponse, Interactions, PartialMessage, SendableEmbed,
SystemMessage, SystemMessage, DataMessageSend,
}, },
Channel, Emoji, Message, User, Channel, Emoji, Message, User,
}, },
@@ -451,3 +451,39 @@ impl Interactions {
!self.restrict_reactions && self.reactions.is_none() !self.restrict_reactions && self.reactions.is_none()
} }
} }
fn throw_permission(permissions: u64, permission: Permission) -> Result<()> {
if (permission as u64) & permissions == (permission as u64) {
Ok(())
} else {
Err(Error::MissingPermission { permission })
}
}
impl DataMessageSend {
pub fn validate_webhook_permissions(
&self,
permissions: u64,
) -> Result<()> {
throw_permission(permissions, Permission::SendMessage)?;
if self.attachments.as_ref().map_or(false, |v| !v.is_empty()) {
throw_permission(permissions, Permission::UploadFiles)?;
};
if self.embeds.as_ref().map_or(false, |v| !v.is_empty()) {
throw_permission(permissions, Permission::SendEmbeds)?;
};
if self.masquerade.is_some() {
throw_permission(permissions, Permission::Masquerade)?;
};
if self.interactions.is_some() {
throw_permission(permissions, Permission::React)?;
};
Ok(())
}
}
@@ -1,6 +1,5 @@
use std::collections::HashSet; use std::collections::HashSet;
use iso8601_timestamp::Timestamp;
use ulid::Ulid; use ulid::Ulid;
use crate::{ use crate::{
@@ -186,18 +185,7 @@ impl Server {
return Err(Error::Banned); return Err(Error::Banned);
} }
let member = Member { let member = Member::new(self.id.clone(), user.id.clone());
id: MemberCompositeKey {
server: self.id.clone(),
user: user.id.clone(),
},
joined_at: Timestamp::now_utc(),
nickname: None,
avatar: None,
roles: vec![],
timeout: None,
};
db.insert_member(&member).await?; db.insert_member(&member).await?;
let should_fetch = channels.is_none(); let should_fetch = channels.is_none();
@@ -3,13 +3,27 @@ use iso8601_timestamp::Timestamp;
use crate::{ use crate::{
events::client::EventV1, events::client::EventV1,
models::{ models::{
server_member::{FieldsMember, PartialMember}, server_member::{FieldsMember, MemberCompositeKey, PartialMember},
Member, Server, Member, Server,
}, },
Database, Result, Database, Result,
}; };
impl Member { impl Member {
pub fn new(server_id: String, user_id: String) -> Self {
Self {
id: MemberCompositeKey {
server: server_id,
user: user_id,
},
joined_at: Timestamp::now_utc(),
nickname: None,
avatar: None,
roles: vec![],
timeout: None,
}
}
/// Update member data /// Update member data
pub async fn update<'a>( pub async fn update<'a>(
&mut self, &mut self,
+33 -3
View File
@@ -10,9 +10,11 @@ use futures::try_join;
use impl_ops::impl_op_ex_commutative; use impl_ops::impl_op_ex_commutative;
use once_cell::sync::Lazy; use once_cell::sync::Lazy;
use rand::seq::SliceRandom; use rand::seq::SliceRandom;
use revolt_database::RatelimitEventType;
use revolt_presence::filter_online; use revolt_presence::filter_online;
use std::collections::HashSet; use std::collections::HashSet;
use std::ops; use std::ops;
use std::time::Duration;
impl_op_ex_commutative!(+ |a: &i32, b: &Badges| -> i32 { *a | *b as i32 }); impl_op_ex_commutative!(+ |a: &i32, b: &Badges| -> i32 { *a | *b as i32 });
@@ -68,6 +70,7 @@ impl User {
x.background = None; x.background = None;
} }
} }
FieldsUser::DisplayName => self.display_name = None,
} }
} }
@@ -176,6 +179,11 @@ impl User {
// Copy the username for validation // Copy the username for validation
let username_lowercase = username.to_lowercase(); let username_lowercase = username.to_lowercase();
// Block homoglyphs
if decancer::cure(&username_lowercase).into_str() != username_lowercase {
return Err(Error::InvalidUsername);
}
// Ensure the username itself isn't blocked // Ensure the username itself isn't blocked
const BLOCKED_USERNAMES: &[&str] = &["admin", "revolt"]; const BLOCKED_USERNAMES: &[&str] = &["admin", "revolt"];
@@ -201,7 +209,7 @@ impl User {
pub async fn find_discriminator( pub async fn find_discriminator(
db: &Database, db: &Database,
username: &str, username: &str,
preferred: Option<String>, preferred: Option<(String, String)>,
) -> Result<String> { ) -> Result<String> {
let search_space: &HashSet<String> = &DISCRIMINATOR_SEARCH_SPACE_QUARK; let search_space: &HashSet<String> = &DISCRIMINATOR_SEARCH_SPACE_QUARK;
let used_discriminators: HashSet<String> = db let used_discriminators: HashSet<String> = db
@@ -217,9 +225,31 @@ impl User {
return Err(Error::UsernameTaken); return Err(Error::UsernameTaken);
} }
if let Some(preferred) = preferred { if let Some((preferred, target_id)) = preferred {
if available_discriminators.contains(&&preferred) { if available_discriminators.contains(&&preferred) {
return Ok(preferred); return Ok(preferred);
} else {
let rvdb: revolt_database::Database = db.clone().into();
if rvdb
.has_ratelimited(
&target_id,
RatelimitEventType::DiscriminatorChange,
Duration::from_secs(60 * 60 * 24),
1,
)
.await
.map_err(Error::from_core)?
{
return Err(Error::DiscriminatorChangeRatelimited);
}
rvdb.insert_ratelimit_event(&revolt_database::RatelimitEvent {
id: ulid::Ulid::new().to_string(),
target_id,
event_type: RatelimitEventType::DiscriminatorChange,
})
.await
.map_err(Error::from_core)?;
} }
} }
@@ -251,7 +281,7 @@ impl User {
User::find_discriminator( User::find_discriminator(
db, db,
&username, &username,
Some(self.discriminator.to_string()), Some((self.discriminator.to_string(), self.id.clone())),
) )
.await?, .await?,
), ),
@@ -337,6 +337,7 @@ impl IntoDocumentPath for FieldsUser {
FieldsUser::ProfileContent => "profile.content", FieldsUser::ProfileContent => "profile.content",
FieldsUser::StatusPresence => "status.presence", FieldsUser::StatusPresence => "status.presence",
FieldsUser::StatusText => "status.text", FieldsUser::StatusText => "status.text",
FieldsUser::DisplayName => "display_name",
}) })
} }
} }
+1
View File
@@ -177,6 +177,7 @@ pub enum FieldsUser {
StatusPresence, StatusPresence,
ProfileContent, ProfileContent,
ProfileBackground, ProfileBackground,
DisplayName,
} }
/// Enumeration providing a hint to the type of user we are handling /// Enumeration providing a hint to the type of user we are handling
@@ -110,6 +110,7 @@ pub static DEFAULT_PERMISSION: Lazy<u64> = Lazy::new(|| DEFAULT_PERMISSION_VIEW_
pub static DEFAULT_PERMISSION_SAVED_MESSAGES: u64 = Permission::GrantAllSafe as u64; pub static DEFAULT_PERMISSION_SAVED_MESSAGES: u64 = Permission::GrantAllSafe as u64;
pub static DEFAULT_PERMISSION_DIRECT_MESSAGE: Lazy<u64> = Lazy::new(|| DEFAULT_PERMISSION.add(Permission::ManageChannel + Permission::React)); pub static DEFAULT_PERMISSION_DIRECT_MESSAGE: Lazy<u64> = Lazy::new(|| DEFAULT_PERMISSION.add(Permission::ManageChannel + Permission::React));
pub static DEFAULT_PERMISSION_SERVER: Lazy<u64> = Lazy::new(|| DEFAULT_PERMISSION.add(Permission::React + Permission::ChangeNickname + Permission::ChangeAvatar)); pub static DEFAULT_PERMISSION_SERVER: Lazy<u64> = Lazy::new(|| DEFAULT_PERMISSION.add(Permission::React + Permission::ChangeNickname + Permission::ChangeAvatar));
pub static DEFAULT_WEBHOOK_PERMISSIONS: Lazy<u64> = Lazy::new(|| Permission::SendMessage + Permission::SendEmbeds + Permission::Masquerade + Permission::React);
bitfield! { bitfield! {
#[derive(Default)] #[derive(Default)]
+1 -1
View File
@@ -8,6 +8,6 @@ pub fn prefix_keys<T: Serialize>(t: &T, prefix: &str) -> HashMap<String, serde_j
let v: HashMap<String, serde_json::Value> = serde_json::from_str(&v).unwrap(); let v: HashMap<String, serde_json::Value> = serde_json::from_str(&v).unwrap();
v.into_iter() v.into_iter()
.filter(|(_k, v)| !v.is_null()) .filter(|(_k, v)| !v.is_null())
.map(|(k, v)| (prefix.to_owned() + &k, v)) .map(|(k, v)| (format!("{}{}", prefix.to_owned(), k), v))
.collect() .collect()
} }
+2
View File
@@ -31,6 +31,7 @@ pub enum Error {
// ? User related errors // ? User related errors
UsernameTaken, UsernameTaken,
InvalidUsername, InvalidUsername,
DiscriminatorChangeRatelimited,
UnknownUser, UnknownUser,
AlreadyFriends, AlreadyFriends,
AlreadySentRequest, AlreadySentRequest,
@@ -165,6 +166,7 @@ impl<'r> Responder<'r, 'static> for Error {
Error::UnknownUser => Status::NotFound, Error::UnknownUser => Status::NotFound,
Error::InvalidUsername => Status::BadRequest, Error::InvalidUsername => Status::BadRequest,
Error::DiscriminatorChangeRatelimited => Status::TooManyRequests,
Error::UsernameTaken => Status::Conflict, Error::UsernameTaken => Status::Conflict,
Error::AlreadyFriends => Status::Conflict, Error::AlreadyFriends => Status::Conflict,
Error::AlreadySentRequest => Status::Conflict, Error::AlreadySentRequest => Status::Conflict,
+13 -9
View File
@@ -101,16 +101,19 @@ pub struct Ratelimiter {
fn resolve_bucket<'r>(request: &'r rocket::Request<'_>) -> (&'r str, Option<&'r str>) { fn resolve_bucket<'r>(request: &'r rocket::Request<'_>) -> (&'r str, Option<&'r str>) {
if let Some(segment) = request.routed_segment(0) { if let Some(segment) = request.routed_segment(0) {
let resource = request.routed_segment(1); let resource = request.routed_segment(1);
match (segment, resource) {
("users", _) => { let method = request.method();
match (segment, resource, method) {
("users", target, Method::Patch) => ("user_edit", target),
("users", _, _) => {
if let Some("default_avatar") = request.routed_segment(2) { if let Some("default_avatar") = request.routed_segment(2) {
return ("default_avatar", None); return ("default_avatar", None);
} }
("users", None) ("users", None)
} }
("bots", _) => ("bots", None), ("bots", _, _) => ("bots", None),
("channels", Some(id)) => { ("channels", Some(id), _) => {
if request.method() == Method::Post { if request.method() == Method::Post {
if let Some("messages") = request.routed_segment(2) { if let Some("messages") = request.routed_segment(2) {
return ("messaging", Some(id)); return ("messaging", Some(id));
@@ -119,17 +122,17 @@ fn resolve_bucket<'r>(request: &'r rocket::Request<'_>) -> (&'r str, Option<&'r
("channels", Some(id)) ("channels", Some(id))
} }
("servers", Some(id)) => ("servers", Some(id)), ("servers", Some(id), _) => ("servers", Some(id)),
("auth", _) => { ("auth", _, _) => {
if request.method() == Method::Delete { if request.method() == Method::Delete {
("auth_delete", None) ("auth_delete", None)
} else { } else {
("auth", None) ("auth", None)
} }
} }
("swagger", _) => ("swagger", None), ("swagger", _, _) => ("swagger", None),
("safety", Some("report")) => ("safety_report", Some("report")), ("safety", Some("report"), _) => ("safety_report", Some("report")),
("safety", _) => ("safety", None), ("safety", _, _) => ("safety", None),
_ => ("any", None), _ => ("any", None),
} }
} else { } else {
@@ -140,6 +143,7 @@ fn resolve_bucket<'r>(request: &'r rocket::Request<'_>) -> (&'r str, Option<&'r
/// Resolve per-bucket limits /// Resolve per-bucket limits
fn resolve_bucket_limit(bucket: &str) -> u8 { fn resolve_bucket_limit(bucket: &str) -> u8 {
match bucket { match bucket {
"user_edit" => 2,
"users" => 20, "users" => 20,
"bots" => 10, "bots" => 10,
"messaging" => 10, "messaging" => 10,