Large 0.3 changes
- Split the gigantic mess that web.rs was into seperate files - Split the app into two clear functions: web and bot - Add privacy policy - Add privacy - Make the site respect GDPR and Discord ToS - Add basic API - Add timezone detection and automatic time changes
This commit is contained in:
1 parent
96a8c9ea8c
commit
126b9e9e9f
23 files changed
+3876
-2767
No files matched your search
@@ -0,0 +1,57 @@
|
||||
use sqlx::PgPool;
|
||||
|
||||
pub struct ArchiveStats {
|
||||
pub messages: i64,
|
||||
pub users: i64,
|
||||
pub servers: i64,
|
||||
pub channels: i64,
|
||||
pub total_storage: i64,
|
||||
pub message_storage: i64,
|
||||
pub attachment_storage: i64,
|
||||
}
|
||||
|
||||
pub async fn load(pool: &PgPool) -> Result<ArchiveStats, sqlx::Error> {
|
||||
let stats = sqlx::query_as::<_, (i64, i64, i64, i64, i64, i64, i64)>(
|
||||
"SELECT
|
||||
(SELECT COUNT(*) FROM messages),
|
||||
(SELECT COUNT(DISTINCT author_id) FROM messages),
|
||||
(SELECT COUNT(*) FROM guilds),
|
||||
(SELECT COUNT(*) FROM channels),
|
||||
(SELECT COALESCE(SUM(OCTET_LENGTH(content)), 0)
|
||||
FROM message_versions)
|
||||
+ (SELECT COALESCE(SUM(OCTET_LENGTH(data)), 0)
|
||||
FROM attachments),
|
||||
(SELECT COALESCE(SUM(OCTET_LENGTH(content)), 0)
|
||||
FROM message_versions),
|
||||
(SELECT COALESCE(SUM(OCTET_LENGTH(data)), 0)
|
||||
FROM attachments);",
|
||||
)
|
||||
.fetch_one(pool)
|
||||
.await?;
|
||||
Ok(ArchiveStats {
|
||||
messages: stats.0,
|
||||
users: stats.1,
|
||||
servers: stats.2,
|
||||
channels: stats.3,
|
||||
total_storage: stats.4,
|
||||
message_storage: stats.5,
|
||||
attachment_storage: stats.6,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn format_bytes(bytes: i64) -> String {
|
||||
const UNITS: [&str; 5] = ["bytes", "KB", "MB", "GB", "TB"];
|
||||
|
||||
let mut size = bytes.max(0) as f64;
|
||||
let mut unit = 0;
|
||||
while size >= 1024.0 && unit < UNITS.len() - 1 {
|
||||
size /= 1024.0;
|
||||
unit += 1;
|
||||
}
|
||||
|
||||
if unit == 0 {
|
||||
format!("{} {}", bytes.max(0), UNITS[unit])
|
||||
} else {
|
||||
format!("{:.1} {}", size, UNITS[unit])
|
||||
}
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
use crate::messaging::{edit_response_message, trace_message};
|
||||
use crate::bot::messaging::{edit_response_message, trace_message};
|
||||
use crate::{Context, Error};
|
||||
use tracing::debug;
|
||||
|
||||
File renamed without changes.
@@ -1,4 +1,4 @@
|
||||
use crate::messaging::trace_message;
|
||||
use crate::bot::messaging::trace_message;
|
||||
use crate::{Context, Error};
|
||||
use poise::serenity_prelude as sere;
|
||||
use tracing::{debug, trace};
|
||||
@@ -1,5 +1,5 @@
|
||||
use crate::messaging::trace_message;
|
||||
use crate::{Context, Error, web};
|
||||
use crate::bot::messaging::{trace_message, edit_response_message};
|
||||
use crate::{Context, Error, archive_stats};
|
||||
use tracing::{debug, trace};
|
||||
|
||||
#[doc = "Show archive statistics"]
|
||||
@@ -10,24 +10,27 @@ pub async fn stats(ctx: Context<'_>) -> Result<(), Error> {
|
||||
ctx.author().id.get(),
|
||||
ctx.guild_id().unwrap().get()
|
||||
);
|
||||
let stats = web::archive_stats(&ctx.data().pool).await?;
|
||||
let response = ctx.say("Processing...").await?;
|
||||
|
||||
let stats = archive_stats::load(&ctx.data().pool).await?;
|
||||
let msg = format!(
|
||||
"# Archive Statistics\n## Messages\nArchived messages: {}\nArchived users: {}\nChannels: {}\nServers: {}\n## Storage usage\nArchived data: {}\nMessage content: {}\nAttachments: {}",
|
||||
stats.messages,
|
||||
stats.users,
|
||||
stats.channels,
|
||||
stats.servers,
|
||||
web::format_bytes(stats.total_storage),
|
||||
web::format_bytes(stats.message_storage),
|
||||
web::format_bytes(stats.attachment_storage),
|
||||
archive_stats::format_bytes(stats.total_storage),
|
||||
archive_stats::format_bytes(stats.message_storage),
|
||||
archive_stats::format_bytes(stats.attachment_storage),
|
||||
);
|
||||
|
||||
edit_response_message(&response, ctx, &msg, true).await?;
|
||||
trace_message(
|
||||
&msg,
|
||||
ctx.channel_id().to_string(),
|
||||
ctx.guild_id().unwrap().to_string(),
|
||||
)
|
||||
.await;
|
||||
ctx.say(msg).await?;
|
||||
debug!("Stats command performed for user {}", ctx.author().name);
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
use crate::messaging::trace_message;
|
||||
use crate::bot::messaging::{trace_message, edit_response_message};
|
||||
use crate::{Context, Error};
|
||||
use tracing::{debug, trace};
|
||||
|
||||
@@ -13,6 +13,7 @@ pub async fn version(ctx: Context<'_>) -> Result<(), Error> {
|
||||
ctx.guild_id().unwrap().get()
|
||||
);
|
||||
trace!("Loading embedded version information");
|
||||
let response = ctx.say("Processing...").await?;
|
||||
let msg = format!("Current version:\n{VERSION}");
|
||||
trace_message(
|
||||
&msg,
|
||||
@@ -20,7 +21,7 @@ pub async fn version(ctx: Context<'_>) -> Result<(), Error> {
|
||||
ctx.guild_id().unwrap().to_string(),
|
||||
)
|
||||
.await;
|
||||
ctx.say(msg).await?;
|
||||
edit_response_message(&response, ctx, &msg, true).await?;
|
||||
debug!(
|
||||
"Version command performed for user {} with version information",
|
||||
ctx.author().name
|
||||
@@ -10,7 +10,7 @@ pub async fn event_handler(
|
||||
) -> Result<(), Error> {
|
||||
match event {
|
||||
sere::FullEvent::Message { new_message } => {
|
||||
crate::message_archive::record_message(ctx, new_message, &user_data.pool).await?;
|
||||
super::message_archive::record_message(ctx, new_message, &user_data.pool).await?;
|
||||
}
|
||||
sere::FullEvent::MessageUpdate {
|
||||
old_if_available,
|
||||
@@ -43,7 +43,7 @@ pub async fn event_handler(
|
||||
message = event.channel_id.message(ctx, event.id).await?;
|
||||
}
|
||||
|
||||
crate::message_archive::record_message_edit(ctx, &message, &user_data.pool).await?;
|
||||
super::message_archive::record_message_edit(ctx, &message, &user_data.pool).await?;
|
||||
}
|
||||
sere::FullEvent::MessageDelete {
|
||||
deleted_message_id, ..
|
||||
File renamed without changes.
File renamed without changes.
@@ -0,0 +1,4 @@
|
||||
pub mod commands;
|
||||
pub mod event_handler;
|
||||
pub mod message_archive;
|
||||
pub mod messaging;
|
||||
+8
-10
@@ -2,11 +2,9 @@ use poise::serenity_prelude as sere;
|
||||
use std::env;
|
||||
use tracing::{debug, error, info, trace};
|
||||
use tracing_subscriber::EnvFilter;
|
||||
pub mod commands;
|
||||
pub mod event_handler;
|
||||
pub mod message_archive;
|
||||
pub mod messaging;
|
||||
pub mod web;
|
||||
mod archive_stats;
|
||||
mod bot;
|
||||
mod web;
|
||||
|
||||
pub struct Data {
|
||||
pub pool: sqlx::PgPool,
|
||||
@@ -95,7 +93,7 @@ async fn main() {
|
||||
let framework = poise::Framework::builder()
|
||||
.options(poise::FrameworkOptions {
|
||||
event_handler: |ctx, event, framework, user_data| {
|
||||
Box::pin(event_handler::event_handler(
|
||||
Box::pin(bot::event_handler::event_handler(
|
||||
ctx, event, framework, user_data,
|
||||
))
|
||||
},
|
||||
@@ -114,10 +112,10 @@ async fn main() {
|
||||
})
|
||||
},
|
||||
commands: vec![
|
||||
commands::ping(),
|
||||
commands::version(),
|
||||
commands::git(),
|
||||
commands::stats(),
|
||||
bot::commands::ping(),
|
||||
bot::commands::version(),
|
||||
bot::commands::git(),
|
||||
bot::commands::stats(),
|
||||
],
|
||||
..Default::default()
|
||||
})
|
||||
|
||||
-2742
File diff suppressed because it is too large.
Load diff
+552
@@ -0,0 +1,552 @@
|
||||
use super::{WebData, WebUser, accessible_channels};
|
||||
use axum::{
|
||||
Extension, Json, Router,
|
||||
extract::{ConnectInfo, Path, Query, Request, State},
|
||||
http::{HeaderMap, HeaderValue, StatusCode, header},
|
||||
middleware::{self, Next},
|
||||
response::{IntoResponse, Response},
|
||||
routing::{get, post},
|
||||
};
|
||||
use rand::{Rng, distr::Alphanumeric};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
env,
|
||||
net::{IpAddr, SocketAddr},
|
||||
sync::{Arc, Mutex},
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
use tower_sessions::Session;
|
||||
|
||||
const RESULTS_PER_PAGE: i64 = 100;
|
||||
const TOKEN_RATE_LIMIT: Duration = Duration::from_secs(5 * 60);
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(super) struct TokenRateLimiter {
|
||||
attempts: Arc<Mutex<HashMap<IpAddr, Instant>>>,
|
||||
trusted_proxies: Arc<Vec<ipnet::IpNet>>,
|
||||
}
|
||||
|
||||
impl TokenRateLimiter {
|
||||
pub(super) fn from_env() -> Result<Self, String> {
|
||||
let trusted_proxies = env::var("TG_BOT_TRUSTED_PROXY_RANGES")
|
||||
.unwrap_or_default()
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.filter(|range| !range.is_empty())
|
||||
.map(|range| {
|
||||
range
|
||||
.parse::<ipnet::IpNet>()
|
||||
.or_else(|_| range.parse::<IpAddr>().map(ipnet::IpNet::from))
|
||||
.map_err(|_| format!("invalid trusted proxy IP range: {range}"))
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
Ok(Self {
|
||||
attempts: Arc::new(Mutex::new(HashMap::new())),
|
||||
trusted_proxies: Arc::new(trusted_proxies),
|
||||
})
|
||||
}
|
||||
|
||||
fn client_ip(&self, peer: IpAddr, headers: &HeaderMap) -> IpAddr {
|
||||
if self
|
||||
.trusted_proxies
|
||||
.iter()
|
||||
.any(|range| range.contains(&peer))
|
||||
{
|
||||
headers
|
||||
.get("x-real-ip")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| value.trim().parse().ok())
|
||||
.unwrap_or(peer)
|
||||
} else {
|
||||
peer
|
||||
}
|
||||
}
|
||||
|
||||
fn check(&self, ip: IpAddr) -> Result<(), u64> {
|
||||
let now = Instant::now();
|
||||
let mut attempts = self
|
||||
.attempts
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner());
|
||||
attempts.retain(|_, attempt| now.duration_since(*attempt) < TOKEN_RATE_LIMIT);
|
||||
if let Some(attempt) = attempts.get(&ip) {
|
||||
let remaining = TOKEN_RATE_LIMIT.saturating_sub(now.duration_since(*attempt));
|
||||
return Err(remaining.as_secs().max(1));
|
||||
}
|
||||
attempts.insert(ip, now);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ApiUser {
|
||||
discord_id: i64,
|
||||
channel_ids: Vec<i64>,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct ErrorResponse {
|
||||
error: &'static str,
|
||||
}
|
||||
type ApiResult<T> = Result<Json<T>, (StatusCode, Json<ErrorResponse>)>;
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct MeResponse {
|
||||
discord_id: i64,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct TokenResponse {
|
||||
token: String,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct ServerResponse {
|
||||
discord_id: i64,
|
||||
name: String,
|
||||
icon_url: Option<String>,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct ChannelResponse {
|
||||
discord_id: i64,
|
||||
name: String,
|
||||
server_id: i64,
|
||||
server_name: String,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct UserResponse {
|
||||
discord_id: i64,
|
||||
username: String,
|
||||
}
|
||||
#[derive(Serialize, sqlx::FromRow)]
|
||||
struct AttachmentResponse {
|
||||
discord_id: i64,
|
||||
filename: String,
|
||||
description: Option<String>,
|
||||
content_type: Option<String>,
|
||||
size: i64,
|
||||
}
|
||||
#[derive(Serialize, sqlx::FromRow)]
|
||||
struct EmbedResponse {
|
||||
index: i32,
|
||||
title: Option<String>,
|
||||
description: Option<String>,
|
||||
url: Option<String>,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct MessageResponse {
|
||||
discord_id: i64,
|
||||
author_id: i64,
|
||||
author_username: String,
|
||||
server_id: i64,
|
||||
server_name: String,
|
||||
channel_id: i64,
|
||||
channel_name: String,
|
||||
timestamp: String,
|
||||
version: i64,
|
||||
archived_at: String,
|
||||
content: Option<String>,
|
||||
attachments: Vec<AttachmentResponse>,
|
||||
embeds: Vec<EmbedResponse>,
|
||||
}
|
||||
#[derive(Serialize, sqlx::FromRow)]
|
||||
struct MessageSummary {
|
||||
discord_id: i64,
|
||||
author_id: i64,
|
||||
author_username: String,
|
||||
server_id: i64,
|
||||
server_name: String,
|
||||
channel_id: i64,
|
||||
channel_name: String,
|
||||
timestamp: String,
|
||||
content: Option<String>,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct SearchResponse {
|
||||
query: String,
|
||||
page: i64,
|
||||
per_page: i64,
|
||||
total: i64,
|
||||
results: Vec<MessageSummary>,
|
||||
}
|
||||
#[derive(Default, Deserialize)]
|
||||
struct VersionQuery {
|
||||
version: Option<i64>,
|
||||
}
|
||||
#[derive(Default, Deserialize)]
|
||||
struct SearchQuery {
|
||||
#[serde(default)]
|
||||
q: String,
|
||||
page: Option<i64>,
|
||||
}
|
||||
pub(super) fn router(data: WebData) -> Router {
|
||||
let protected = Router::new()
|
||||
.route("/me", get(me))
|
||||
.route("/view/message/{id}", get(view_message))
|
||||
.route("/view/server/{id}", get(view_server))
|
||||
.route("/view/channel/{id}", get(view_channel))
|
||||
.route("/view/user/{id}", get(view_user))
|
||||
.route("/search/timestamp", get(search_timestamp))
|
||||
.route("/search/content", get(search_content))
|
||||
.route_layer(middleware::from_fn_with_state(
|
||||
data.clone(),
|
||||
require_api_user,
|
||||
));
|
||||
Router::new()
|
||||
.route("/token", post(create_token))
|
||||
.merge(protected)
|
||||
.with_state(data)
|
||||
}
|
||||
|
||||
async fn require_api_user(
|
||||
State(data): State<WebData>,
|
||||
mut request: Request,
|
||||
next: Next,
|
||||
) -> Response {
|
||||
let Some(token) = bearer_token(request.headers().get(header::AUTHORIZATION)) else {
|
||||
return api_error(StatusCode::UNAUTHORIZED, "invalid or missing bearer token");
|
||||
};
|
||||
let discord_id =
|
||||
match sqlx::query_scalar::<_, i64>("SELECT discord_id FROM api_tokens WHERE token = $1")
|
||||
.bind(token)
|
||||
.fetch_optional(&data.pool)
|
||||
.await
|
||||
{
|
||||
Ok(Some(id)) => id,
|
||||
Ok(None) => {
|
||||
return api_error(StatusCode::UNAUTHORIZED, "invalid or missing bearer token");
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::error!("API token lookup failed: {}", error);
|
||||
return api_error(StatusCode::INTERNAL_SERVER_ERROR, "database error");
|
||||
}
|
||||
};
|
||||
let channel_ids = match accessible_channels(&data, discord_id).await {
|
||||
Ok(access) => access
|
||||
.into_iter()
|
||||
.map(|channel| channel.channel_id)
|
||||
.collect(),
|
||||
Err(error) => {
|
||||
tracing::error!("API permission refresh failed: {}", error);
|
||||
return api_error(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"could not refresh Discord permissions",
|
||||
);
|
||||
}
|
||||
};
|
||||
request.extensions_mut().insert(ApiUser {
|
||||
discord_id,
|
||||
channel_ids,
|
||||
});
|
||||
next.run(request).await
|
||||
}
|
||||
|
||||
async fn create_token(
|
||||
State(data): State<WebData>,
|
||||
session: Session,
|
||||
ConnectInfo(peer): ConnectInfo<SocketAddr>,
|
||||
headers: HeaderMap,
|
||||
) -> Response {
|
||||
let user = match session.get::<WebUser>("user").await {
|
||||
Ok(Some(user)) => user,
|
||||
Ok(None) => return api_error(StatusCode::UNAUTHORIZED, "browser login required"),
|
||||
Err(_) => return api_error(StatusCode::INTERNAL_SERVER_ERROR, "session error"),
|
||||
};
|
||||
let client_ip = data.token_rate_limiter.client_ip(peer.ip(), &headers);
|
||||
if let Err(retry_after) = data.token_rate_limiter.check(client_ip) {
|
||||
return rate_limited(retry_after);
|
||||
}
|
||||
let token: String = rand::rng()
|
||||
.sample_iter(&Alphanumeric)
|
||||
.take(64)
|
||||
.map(char::from)
|
||||
.collect();
|
||||
let mut transaction = match data.pool.begin().await {
|
||||
Ok(transaction) => transaction,
|
||||
Err(error_value) => return database(error_value).into_response(),
|
||||
};
|
||||
let result = sqlx::query(
|
||||
"INSERT INTO discord_users (discord_id, discord_username)
|
||||
VALUES ($1, $2)
|
||||
ON CONFLICT (discord_id) DO UPDATE
|
||||
SET discord_username = EXCLUDED.discord_username",
|
||||
)
|
||||
.bind(user.id)
|
||||
.bind(&user.username)
|
||||
.execute(&mut *transaction)
|
||||
.await;
|
||||
if let Err(error_value) = result {
|
||||
return database(error_value).into_response();
|
||||
}
|
||||
let result = sqlx::query("INSERT INTO api_tokens (token, discord_id) VALUES ($1, $2)")
|
||||
.bind(&token)
|
||||
.bind(user.id)
|
||||
.execute(&mut *transaction)
|
||||
.await;
|
||||
if let Err(error_value) = result {
|
||||
return database(error_value).into_response();
|
||||
}
|
||||
if let Err(error_value) = transaction.commit().await {
|
||||
return database(error_value).into_response();
|
||||
}
|
||||
Json(TokenResponse { token }).into_response()
|
||||
}
|
||||
|
||||
async fn me(Extension(user): Extension<ApiUser>) -> Json<MeResponse> {
|
||||
Json(MeResponse {
|
||||
discord_id: user.discord_id,
|
||||
})
|
||||
}
|
||||
|
||||
async fn view_server(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Path(id): Path<i64>,
|
||||
) -> ApiResult<ServerResponse> {
|
||||
let row = sqlx::query_as::<_, (i64, String, Option<String>)>(
|
||||
"SELECT DISTINCT g.guild_id, g.guild_name, g.guild_icon_url FROM guilds g
|
||||
JOIN channels c ON c.guild_id = g.guild_id
|
||||
WHERE g.guild_id = $1 AND c.channel_id = ANY($2)",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(&user.channel_ids)
|
||||
.fetch_optional(&data.pool)
|
||||
.await
|
||||
.map_err(database)?
|
||||
.ok_or_else(not_found)?;
|
||||
Ok(Json(ServerResponse {
|
||||
discord_id: row.0,
|
||||
name: row.1,
|
||||
icon_url: row.2,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn view_channel(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Path(id): Path<i64>,
|
||||
) -> ApiResult<ChannelResponse> {
|
||||
let row = sqlx::query_as::<_, (i64, String, i64, String)>(
|
||||
"SELECT c.channel_id, c.channel_name, g.guild_id, g.guild_name FROM channels c
|
||||
JOIN guilds g ON g.guild_id = c.guild_id WHERE c.channel_id = $1 AND c.channel_id = ANY($2)")
|
||||
.bind(id).bind(&user.channel_ids).fetch_optional(&data.pool).await.map_err(database)?.ok_or_else(not_found)?;
|
||||
Ok(Json(ChannelResponse {
|
||||
discord_id: row.0,
|
||||
name: row.1,
|
||||
server_id: row.2,
|
||||
server_name: row.3,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn view_user(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Path(id): Path<i64>,
|
||||
) -> ApiResult<UserResponse> {
|
||||
let row = sqlx::query_as::<_, (i64, String)>(
|
||||
"SELECT u.discord_id, u.discord_username FROM discord_users u WHERE u.discord_id = $1
|
||||
AND (u.discord_id = $2 OR EXISTS (SELECT 1 FROM messages m
|
||||
WHERE m.author_id = u.discord_id AND m.channel_id = ANY($3)))",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(user.discord_id)
|
||||
.bind(&user.channel_ids)
|
||||
.fetch_optional(&data.pool)
|
||||
.await
|
||||
.map_err(database)?
|
||||
.ok_or_else(not_found)?;
|
||||
Ok(Json(UserResponse {
|
||||
discord_id: row.0,
|
||||
username: row.1,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn view_message(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Path(id): Path<i64>,
|
||||
Query(query): Query<VersionQuery>,
|
||||
) -> ApiResult<MessageResponse> {
|
||||
let message = sqlx::query_as::<_, (i64, i64, String, i64, String, i64, String, String)>(
|
||||
"SELECT m.message_id, m.author_id, m.author_username, m.guild_id, g.guild_name,
|
||||
m.channel_id, c.channel_name, m.timestamp::text FROM messages m
|
||||
JOIN guilds g ON g.guild_id = m.guild_id JOIN channels c ON c.channel_id = m.channel_id
|
||||
WHERE m.message_id = $1 AND (m.channel_id = ANY($2) OR m.author_id = $3)",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(&user.channel_ids)
|
||||
.bind(user.discord_id)
|
||||
.fetch_optional(&data.pool)
|
||||
.await
|
||||
.map_err(database)?
|
||||
.ok_or_else(not_found)?;
|
||||
let version = sqlx::query_as::<_, (i64, Option<String>, String)>(
|
||||
"SELECT version, content, archived_at::text FROM message_versions
|
||||
WHERE message_id = $1 AND ($2::bigint IS NULL OR version = $2) ORDER BY version DESC LIMIT 1")
|
||||
.bind(id).bind(query.version).fetch_optional(&data.pool).await.map_err(database)?.ok_or_else(not_found)?;
|
||||
let attachments = sqlx::query_as::<_, AttachmentResponse>(
|
||||
"SELECT attachment_id AS discord_id, filename, description, content_type, size FROM attachments
|
||||
WHERE message_id = $1 AND message_version = $2 ORDER BY attachment_id")
|
||||
.bind(id).bind(version.0).fetch_all(&data.pool).await.map_err(database)?;
|
||||
let embeds = sqlx::query_as::<_, EmbedResponse>(
|
||||
"SELECT embed_index AS index, title, description, url FROM embeds
|
||||
WHERE message_id = $1 AND message_version = $2 ORDER BY embed_index",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(version.0)
|
||||
.fetch_all(&data.pool)
|
||||
.await
|
||||
.map_err(database)?;
|
||||
Ok(Json(MessageResponse {
|
||||
discord_id: message.0,
|
||||
author_id: message.1,
|
||||
author_username: message.2,
|
||||
server_id: message.3,
|
||||
server_name: message.4,
|
||||
channel_id: message.5,
|
||||
channel_name: message.6,
|
||||
timestamp: message.7,
|
||||
version: version.0,
|
||||
content: version.1,
|
||||
archived_at: version.2,
|
||||
attachments,
|
||||
embeds,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn search_timestamp(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Query(query): Query<SearchQuery>,
|
||||
) -> ApiResult<SearchResponse> {
|
||||
search(&data, &user, query, true).await
|
||||
}
|
||||
async fn search_content(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Query(query): Query<SearchQuery>,
|
||||
) -> ApiResult<SearchResponse> {
|
||||
search(&data, &user, query, false).await
|
||||
}
|
||||
async fn search(
|
||||
data: &WebData,
|
||||
user: &ApiUser,
|
||||
query: SearchQuery,
|
||||
timestamp: bool,
|
||||
) -> ApiResult<SearchResponse> {
|
||||
let page = query.page.unwrap_or(1).max(1);
|
||||
let pattern = format!("%{}%", query.q);
|
||||
let total = sqlx::query_scalar::<_, i64>(
|
||||
"SELECT COUNT(*) FROM messages m WHERE CASE WHEN $2 THEN m.timestamp::text ILIKE $1
|
||||
ELSE COALESCE(m.content, '') ILIKE $1 END AND (m.channel_id = ANY($3) OR m.author_id = $4)")
|
||||
.bind(&pattern).bind(timestamp).bind(&user.channel_ids).bind(user.discord_id)
|
||||
.fetch_one(&data.pool).await.map_err(database)?;
|
||||
let results = sqlx::query_as::<_, MessageSummary>(
|
||||
"SELECT m.message_id AS discord_id, m.author_id, m.author_username, m.guild_id AS server_id,
|
||||
g.guild_name AS server_name, m.channel_id, c.channel_name, m.timestamp::text AS timestamp, m.content
|
||||
FROM messages m JOIN guilds g ON g.guild_id = m.guild_id JOIN channels c ON c.channel_id = m.channel_id
|
||||
WHERE CASE WHEN $2 THEN m.timestamp::text ILIKE $1 ELSE COALESCE(m.content, '') ILIKE $1 END
|
||||
AND (m.channel_id = ANY($3) OR m.author_id = $4)
|
||||
ORDER BY m.timestamp DESC, m.message_id DESC LIMIT $5 OFFSET $6")
|
||||
.bind(&pattern).bind(timestamp).bind(&user.channel_ids).bind(user.discord_id)
|
||||
.bind(RESULTS_PER_PAGE).bind((page - 1) * RESULTS_PER_PAGE)
|
||||
.fetch_all(&data.pool).await.map_err(database)?;
|
||||
Ok(Json(SearchResponse {
|
||||
query: query.q,
|
||||
page,
|
||||
per_page: RESULTS_PER_PAGE,
|
||||
total,
|
||||
results,
|
||||
}))
|
||||
}
|
||||
|
||||
fn bearer_token(value: Option<&axum::http::HeaderValue>) -> Option<&str> {
|
||||
let value = value?.to_str().ok()?;
|
||||
let (scheme, token) = value.split_once(' ')?;
|
||||
(scheme.eq_ignore_ascii_case("Bearer")
|
||||
&& !token.is_empty()
|
||||
&& !token.contains(char::is_whitespace))
|
||||
.then_some(token)
|
||||
}
|
||||
fn database(error_value: sqlx::Error) -> (StatusCode, Json<ErrorResponse>) {
|
||||
tracing::error!("API database query failed: {}", error_value);
|
||||
error(StatusCode::INTERNAL_SERVER_ERROR, "database error")
|
||||
}
|
||||
fn not_found() -> (StatusCode, Json<ErrorResponse>) {
|
||||
error(StatusCode::NOT_FOUND, "not found")
|
||||
}
|
||||
fn error(status: StatusCode, message: &'static str) -> (StatusCode, Json<ErrorResponse>) {
|
||||
(status, Json(ErrorResponse { error: message }))
|
||||
}
|
||||
fn api_error(status: StatusCode, message: &'static str) -> Response {
|
||||
error(status, message).into_response()
|
||||
}
|
||||
|
||||
fn rate_limited(retry_after: u64) -> Response {
|
||||
let mut response = api_error(
|
||||
StatusCode::TOO_MANY_REQUESTS,
|
||||
"token creation rate limit exceeded",
|
||||
);
|
||||
response.headers_mut().insert(
|
||||
header::RETRY_AFTER,
|
||||
HeaderValue::from_str(&retry_after.to_string()).unwrap(),
|
||||
);
|
||||
response
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn limiter(trusted_proxies: &[&str]) -> TokenRateLimiter {
|
||||
TokenRateLimiter {
|
||||
attempts: Arc::new(Mutex::new(HashMap::new())),
|
||||
trusted_proxies: Arc::new(
|
||||
trusted_proxies
|
||||
.iter()
|
||||
.map(|range| range.parse().unwrap())
|
||||
.collect(),
|
||||
),
|
||||
}
|
||||
}
|
||||
#[test]
|
||||
fn parses_bearer_tokens() {
|
||||
let header = axum::http::HeaderValue::from_static("Bearer secret-token");
|
||||
assert_eq!(bearer_token(Some(&header)), Some("secret-token"));
|
||||
let lowercase = axum::http::HeaderValue::from_static("bearer another-token");
|
||||
assert_eq!(bearer_token(Some(&lowercase)), Some("another-token"));
|
||||
}
|
||||
#[test]
|
||||
fn rejects_invalid_authorization_headers() {
|
||||
let basic = axum::http::HeaderValue::from_static("Basic credentials");
|
||||
let empty = axum::http::HeaderValue::from_static("Bearer ");
|
||||
let spaced = axum::http::HeaderValue::from_static("Bearer two tokens");
|
||||
assert_eq!(bearer_token(None), None);
|
||||
assert_eq!(bearer_token(Some(&basic)), None);
|
||||
assert_eq!(bearer_token(Some(&empty)), None);
|
||||
assert_eq!(bearer_token(Some(&spaced)), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn only_trusts_forwarded_ip_from_configured_proxies() {
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-real-ip", HeaderValue::from_static("203.0.113.10"));
|
||||
let limiter = limiter(&["10.0.0.0/8"]);
|
||||
|
||||
assert_eq!(
|
||||
limiter.client_ip("10.1.2.3".parse().unwrap(), &headers),
|
||||
"203.0.113.10".parse::<IpAddr>().unwrap()
|
||||
);
|
||||
assert_eq!(
|
||||
limiter.client_ip("192.0.2.5".parse().unwrap(), &headers),
|
||||
"192.0.2.5".parse::<IpAddr>().unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn limits_repeated_token_creation_by_ip() {
|
||||
let limiter = limiter(&[]);
|
||||
let ip = "192.0.2.5".parse().unwrap();
|
||||
assert_eq!(limiter.check(ip), Ok(()));
|
||||
assert!(limiter.check(ip).is_err());
|
||||
assert_eq!(limiter.check("192.0.2.6".parse().unwrap()), Ok(()));
|
||||
}
|
||||
}
|
||||
+669
@@ -0,0 +1,669 @@
|
||||
use super::*;
|
||||
|
||||
pub(super) async fn timezone_request(jar: CookieJar, request: Request, next: Next) -> Response {
|
||||
let timezone = jar
|
||||
.get("timezone")
|
||||
.and_then(|cookie| cookie.value().parse::<chrono_tz::Tz>().ok());
|
||||
let detect = timezone.is_none() && request.uri().path() == "/";
|
||||
ACTIVE_TIMEZONE
|
||||
.scope(
|
||||
TimezoneContext {
|
||||
timezone: timezone.unwrap_or(chrono_tz::UTC),
|
||||
detect,
|
||||
},
|
||||
next.run(request),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(super) async fn theme_request(jar: CookieJar, request: Request, next: Next) -> Response {
|
||||
let theme = Theme::from_cookie(jar.get("theme").map(Cookie::value));
|
||||
ACTIVE_THEME.scope(theme, next.run(request)).await
|
||||
}
|
||||
|
||||
pub(super) async fn require_user(
|
||||
State(data): State<WebData>,
|
||||
session: Session,
|
||||
mut request: Request,
|
||||
next: Next,
|
||||
) -> Response {
|
||||
match session.get::<WebUser>("user").await {
|
||||
Ok(Some(mut user)) => {
|
||||
if channel_access_needs_refresh(&user, request.uri().path()) {
|
||||
let channel_access = match accessible_channels(&data, user.id).await {
|
||||
Ok(channel_access) => channel_access,
|
||||
Err(error) => {
|
||||
error!("Discord permission refresh failed: {}", error);
|
||||
return StatusCode::BAD_GATEWAY.into_response();
|
||||
}
|
||||
};
|
||||
user.channel_ids = channel_access
|
||||
.iter()
|
||||
.map(|access| access.channel_id)
|
||||
.collect();
|
||||
user.channel_access = channel_access;
|
||||
user.channel_access_refreshed_at = unix_timestamp();
|
||||
if let Err(error) = session.insert("user", &user).await {
|
||||
error!("Could not update refreshed session permissions: {}", error);
|
||||
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
|
||||
}
|
||||
}
|
||||
if !request_is_allowed(&data.pool, &user, request.uri().path()).await {
|
||||
return StatusCode::NOT_FOUND.into_response();
|
||||
}
|
||||
request.extensions_mut().insert(user);
|
||||
next.run(request).await
|
||||
}
|
||||
Ok(None) => Redirect::to("/login").into_response(),
|
||||
Err(error) => {
|
||||
error!("Session error: {}", error);
|
||||
StatusCode::INTERNAL_SERVER_ERROR.into_response()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn unix_timestamp() -> u64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs()
|
||||
}
|
||||
|
||||
pub(super) fn channel_access_needs_refresh_at(user: &WebUser, path: &str, now: u64) -> bool {
|
||||
let expired =
|
||||
now.saturating_sub(user.channel_access_refreshed_at) >= CHANNEL_ACCESS_TTL_SECONDS;
|
||||
let requested_channel_is_missing = path
|
||||
.trim_matches('/')
|
||||
.split('/')
|
||||
.collect::<Vec<_>>()
|
||||
.as_slice()
|
||||
.get(0..2)
|
||||
.and_then(|parts| match parts {
|
||||
["channels", id] => id.parse::<i64>().ok(),
|
||||
_ => None,
|
||||
})
|
||||
.is_some_and(|channel_id| !user.channel_ids.contains(&channel_id));
|
||||
expired || requested_channel_is_missing
|
||||
}
|
||||
|
||||
pub(super) fn channel_access_needs_refresh(user: &WebUser, path: &str) -> bool {
|
||||
channel_access_needs_refresh_at(user, path, unix_timestamp())
|
||||
}
|
||||
|
||||
pub(super) async fn request_is_allowed(pool: &PgPool, user: &WebUser, path: &str) -> bool {
|
||||
let parts = path.trim_matches('/').split('/').collect::<Vec<_>>();
|
||||
let channel_id = match parts.as_slice() {
|
||||
["channels", id, ..] => id.parse::<i64>().ok(),
|
||||
["messages", id] => return match id.parse::<i64>() {
|
||||
Ok(id) => sqlx::query_scalar::<_, bool>(
|
||||
"SELECT EXISTS(
|
||||
SELECT 1 FROM messages
|
||||
WHERE message_id = $1
|
||||
AND (author_id = $2 OR channel_id = ANY($3))
|
||||
)",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(user.id)
|
||||
.bind(&user.channel_ids)
|
||||
.fetch_one(pool)
|
||||
.await
|
||||
.unwrap_or(false),
|
||||
Err(_) => false,
|
||||
},
|
||||
["attachments", id] => return match id.parse::<i64>() {
|
||||
Ok(id) => sqlx::query_scalar::<_, bool>(
|
||||
"SELECT EXISTS(
|
||||
SELECT 1 FROM attachments a
|
||||
JOIN messages m ON m.message_id = a.message_id
|
||||
WHERE a.attachment_id = $1
|
||||
AND (m.author_id = $2 OR m.channel_id = ANY($3))
|
||||
)",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(user.id)
|
||||
.bind(&user.channel_ids)
|
||||
.fetch_one(pool)
|
||||
.await
|
||||
.unwrap_or(false),
|
||||
Err(_) => false,
|
||||
},
|
||||
["servers", id, ..] => return match id.parse::<i64>() {
|
||||
Ok(id) => sqlx::query_scalar::<_, bool>(
|
||||
"SELECT EXISTS(SELECT 1 FROM channels WHERE guild_id = $1 AND channel_id = ANY($2))",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(&user.channel_ids)
|
||||
.fetch_one(pool)
|
||||
.await
|
||||
.unwrap_or(false),
|
||||
Err(_) => false,
|
||||
},
|
||||
["users", id, ..] => return match id.parse::<i64>() {
|
||||
Ok(id) if id == user.id => true,
|
||||
Ok(id) => sqlx::query_scalar::<_, bool>(
|
||||
"SELECT EXISTS(SELECT 1 FROM messages WHERE author_id = $1 AND channel_id = ANY($2))",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(&user.channel_ids)
|
||||
.fetch_one(pool)
|
||||
.await
|
||||
.unwrap_or(false),
|
||||
Err(_) => false,
|
||||
},
|
||||
_ => return true,
|
||||
};
|
||||
channel_id.is_some_and(|channel_id| user.channel_ids.contains(&channel_id))
|
||||
}
|
||||
|
||||
pub(super) async fn login(session: Session) -> Response {
|
||||
match session.get::<WebUser>("user").await {
|
||||
Ok(Some(_)) => Redirect::to("/").into_response(),
|
||||
Ok(None) => Html(render_page(
|
||||
"Log in",
|
||||
&render_template(&LoginTemplate),
|
||||
false,
|
||||
))
|
||||
.into_response(),
|
||||
Err(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn privacy(session: Session) -> Response {
|
||||
match session.get::<WebUser>("user").await {
|
||||
Ok(user) => {
|
||||
let logged_in = user.is_some();
|
||||
let body = render_template(&PrivacyTemplate { logged_in });
|
||||
Html(render_page("Privacy", &body, logged_in)).into_response()
|
||||
}
|
||||
Err(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn privacy_policy(session: Session) -> Response {
|
||||
match session.get::<WebUser>("user").await {
|
||||
Ok(user) => Html(render_page(
|
||||
"Privacy policy",
|
||||
&privacy_policy_html(),
|
||||
user.is_some(),
|
||||
))
|
||||
.into_response(),
|
||||
Err(_) => StatusCode::INTERNAL_SERVER_ERROR.into_response(),
|
||||
}
|
||||
}
|
||||
|
||||
fn privacy_policy_html() -> String {
|
||||
const MARKDOWN: &str = include_str!("../../static/privacy-policy.md");
|
||||
|
||||
let options = pulldown_cmark::Options::all();
|
||||
let parser = pulldown_cmark::Parser::new_ext(MARKDOWN, options);
|
||||
|
||||
let mut output = String::new();
|
||||
pulldown_cmark::html::push_html(&mut output, parser);
|
||||
|
||||
output
|
||||
}
|
||||
|
||||
pub(super) async fn anonymize_confirmation() -> Html<String> {
|
||||
let body = render_template(&AnonymizeConfirmationTemplate);
|
||||
Html(page("Confirm anonymization", &body))
|
||||
}
|
||||
|
||||
pub(super) async fn anonymize_all(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<WebUser>,
|
||||
session: Session,
|
||||
) -> Response {
|
||||
let result = async {
|
||||
let mut transaction = data.pool.begin().await?;
|
||||
sqlx::query(
|
||||
"INSERT INTO discord_users (discord_id, discord_username)
|
||||
VALUES (0, 'Deleted user')
|
||||
ON CONFLICT (discord_id) DO UPDATE SET discord_username = EXCLUDED.discord_username",
|
||||
)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
sqlx::query(
|
||||
"INSERT INTO discord_user_history
|
||||
(discord_id, discord_username, discord_avatar_url)
|
||||
VALUES (0, 'Deleted user', NULL)
|
||||
ON CONFLICT (discord_id, discord_username, discord_avatar_url)
|
||||
DO UPDATE SET last_seen_at = EXCLUDED.last_seen_at",
|
||||
)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
sqlx::query(
|
||||
"INSERT INTO guild_users (guild_id, discord_id, first_seen_at, last_seen_at)
|
||||
SELECT guild_id, 0, first_seen_at, last_seen_at
|
||||
FROM guild_users
|
||||
WHERE discord_id = $1
|
||||
ON CONFLICT (guild_id, discord_id) DO UPDATE SET
|
||||
first_seen_at = LEAST(guild_users.first_seen_at, EXCLUDED.first_seen_at),
|
||||
last_seen_at = GREATEST(guild_users.last_seen_at, EXCLUDED.last_seen_at)",
|
||||
)
|
||||
.bind(user.id)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
sqlx::query(
|
||||
"UPDATE messages
|
||||
SET author_id = 0, author_username = 'Deleted user'
|
||||
WHERE author_id = $1",
|
||||
)
|
||||
.bind(user.id)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
sqlx::query("DELETE FROM discord_users WHERE discord_id = $1")
|
||||
.bind(user.id)
|
||||
.execute(&mut *transaction)
|
||||
.await?;
|
||||
transaction.commit().await
|
||||
}
|
||||
.await;
|
||||
|
||||
if let Err(error) = result {
|
||||
error!("Could not anonymize Discord user {}: {}", user.id, error);
|
||||
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
|
||||
}
|
||||
|
||||
let _ = session.delete().await;
|
||||
Redirect::to("/privacy").into_response()
|
||||
}
|
||||
|
||||
pub(super) async fn set_theme(
|
||||
State(data): State<WebData>,
|
||||
jar: CookieJar,
|
||||
headers: HeaderMap,
|
||||
Form(form): Form<ThemeForm>,
|
||||
) -> Response {
|
||||
let theme = Theme::from_cookie(Some(&form.theme));
|
||||
let cookie = Cookie::build(("theme", theme.as_str()))
|
||||
.path("/")
|
||||
.http_only(true)
|
||||
.same_site(SameSite::Lax)
|
||||
.secure(!data.redirect_uri.starts_with("http://"))
|
||||
.max_age(Duration::days(365))
|
||||
.build();
|
||||
let return_to = headers
|
||||
.get(header::REFERER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| reqwest::Url::parse(value).ok())
|
||||
.map(|url| {
|
||||
let mut path = url.path().to_string();
|
||||
if let Some(query) = url.query() {
|
||||
path.push('?');
|
||||
path.push_str(query);
|
||||
}
|
||||
path
|
||||
})
|
||||
.unwrap_or_else(|| "/".into());
|
||||
(jar.add(cookie), Redirect::to(&return_to)).into_response()
|
||||
}
|
||||
|
||||
pub(super) async fn theme_css() -> impl IntoResponse {
|
||||
(
|
||||
[(header::CONTENT_TYPE, "text/css; charset=utf-8")],
|
||||
include_str!("../../static/theme.css"),
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) async fn timezone_js() -> impl IntoResponse {
|
||||
(
|
||||
[(header::CONTENT_TYPE, "text/javascript; charset=utf-8")],
|
||||
include_str!("../../static/timezone.js"),
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) async fn set_timezone(
|
||||
State(data): State<WebData>,
|
||||
jar: CookieJar,
|
||||
Json(form): Json<TimezoneForm>,
|
||||
) -> Response {
|
||||
if form.timezone.parse::<chrono_tz::Tz>().is_err() {
|
||||
return StatusCode::BAD_REQUEST.into_response();
|
||||
}
|
||||
let cookie = Cookie::build(("timezone", form.timezone))
|
||||
.path("/")
|
||||
.http_only(true)
|
||||
.same_site(SameSite::Lax)
|
||||
.secure(!data.redirect_uri.starts_with("http://"))
|
||||
.max_age(Duration::days(365))
|
||||
.build();
|
||||
(jar.add(cookie), StatusCode::NO_CONTENT).into_response()
|
||||
}
|
||||
|
||||
pub(super) async fn discord_login(State(data): State<WebData>, session: Session) -> Response {
|
||||
let state: String = rand::rng()
|
||||
.sample_iter(&Alphanumeric)
|
||||
.take(48)
|
||||
.map(char::from)
|
||||
.collect();
|
||||
if session.insert("oauth_state", &state).await.is_err() {
|
||||
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
|
||||
}
|
||||
let url = format!(
|
||||
"https://discord.com/oauth2/authorize?response_type=code&client_id={}&scope=identify&state={}&redirect_uri={}",
|
||||
data.client_id,
|
||||
state,
|
||||
urlencoding::encode(&data.redirect_uri),
|
||||
);
|
||||
Redirect::to(&url).into_response()
|
||||
}
|
||||
|
||||
pub(super) async fn discord_callback(
|
||||
State(data): State<WebData>,
|
||||
session: Session,
|
||||
Query(query): Query<OAuthQuery>,
|
||||
) -> Response {
|
||||
let state = session.remove::<String>("oauth_state").await;
|
||||
if !matches!(state, Ok(Some(state)) if state == query.state) {
|
||||
return (StatusCode::BAD_REQUEST, "Invalid OAuth state").into_response();
|
||||
}
|
||||
let token = data
|
||||
.http
|
||||
.post("https://discord.com/api/v10/oauth2/token")
|
||||
.form(&[
|
||||
("client_id", data.client_id.as_str()),
|
||||
("client_secret", data.client_secret.as_str()),
|
||||
("grant_type", "authorization_code"),
|
||||
("code", query.code.as_str()),
|
||||
("redirect_uri", data.redirect_uri.as_str()),
|
||||
])
|
||||
.send()
|
||||
.await;
|
||||
let Ok(token) = token else {
|
||||
return oauth_error("Discord token request failed");
|
||||
};
|
||||
let Ok(token) = token.error_for_status() else {
|
||||
return oauth_error("Discord rejected the login request");
|
||||
};
|
||||
let Ok(token) = token.json::<OAuthToken>().await else {
|
||||
return oauth_error("Discord returned an invalid token response");
|
||||
};
|
||||
let user = data
|
||||
.http
|
||||
.get("https://discord.com/api/v10/users/@me")
|
||||
.bearer_auth(&token.access_token)
|
||||
.send()
|
||||
.await;
|
||||
let Ok(user) = user else {
|
||||
return oauth_error("Discord user request failed");
|
||||
};
|
||||
let Ok(user) = user.error_for_status() else {
|
||||
return oauth_error("Discord rejected the user request");
|
||||
};
|
||||
let Ok(user) = user.json::<DiscordUser>().await else {
|
||||
return oauth_error("Discord returned an invalid user response");
|
||||
};
|
||||
let Ok(user_id) = user.id.parse::<i64>() else {
|
||||
return oauth_error("Discord returned an invalid user ID");
|
||||
};
|
||||
let channel_access = match accessible_channels(&data, user_id).await {
|
||||
Ok(channel_access) => channel_access,
|
||||
Err(error) => {
|
||||
error!("Discord permission check failed: {}", error);
|
||||
return oauth_error("Could not check your current Discord permissions");
|
||||
}
|
||||
};
|
||||
let channel_ids = channel_access
|
||||
.iter()
|
||||
.map(|access| access.channel_id)
|
||||
.collect();
|
||||
let user = WebUser {
|
||||
id: user_id,
|
||||
username: user.username,
|
||||
channel_ids,
|
||||
channel_access,
|
||||
channel_access_refreshed_at: unix_timestamp(),
|
||||
};
|
||||
if session.cycle_id().await.is_err() {
|
||||
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
|
||||
}
|
||||
if session.insert("user", user).await.is_err() {
|
||||
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
|
||||
}
|
||||
Redirect::to("/").into_response()
|
||||
}
|
||||
|
||||
pub(super) async fn logout(session: Session) -> Redirect {
|
||||
let _ = session.delete().await;
|
||||
Redirect::to("/login")
|
||||
}
|
||||
|
||||
pub(super) fn oauth_error(message: &str) -> Response {
|
||||
(StatusCode::BAD_GATEWAY, message.to_string()).into_response()
|
||||
}
|
||||
|
||||
pub(super) async fn accessible_channels(
|
||||
data: &WebData,
|
||||
user_id: i64,
|
||||
) -> Result<Vec<ChannelAccess>, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let guilds = sqlx::query_as::<_, (i64, String)>(
|
||||
"SELECT guild_id, guild_name FROM guilds ORDER BY guild_id;",
|
||||
)
|
||||
.fetch_all(&data.pool)
|
||||
.await?;
|
||||
let bot = discord_get::<DiscordUser>(data, "/users/@me").await?;
|
||||
let bot_id = bot.id.parse::<i64>()?;
|
||||
let archived_channels = sqlx::query_scalar::<_, i64>("SELECT channel_id FROM channels;")
|
||||
.fetch_all(&data.pool)
|
||||
.await?
|
||||
.into_iter()
|
||||
.collect::<HashSet<_>>();
|
||||
let mut visible = Vec::new();
|
||||
|
||||
for (guild_id, guild_name) in guilds {
|
||||
let user = discord_get_optional::<DiscordMember>(
|
||||
data,
|
||||
&format!("/guilds/{}/members/{}", guild_id, user_id),
|
||||
)
|
||||
.await?;
|
||||
let Some(user) = user else {
|
||||
continue;
|
||||
};
|
||||
let bot = discord_get_optional::<DiscordMember>(
|
||||
data,
|
||||
&format!("/guilds/{}/members/{}", guild_id, bot_id),
|
||||
)
|
||||
.await?;
|
||||
let Some(bot) = bot else {
|
||||
continue;
|
||||
};
|
||||
let roles =
|
||||
discord_get::<Vec<DiscordRole>>(data, &format!("/guilds/{}/roles", guild_id)).await?;
|
||||
let channels =
|
||||
discord_get::<Vec<DiscordChannel>>(data, &format!("/guilds/{}/channels", guild_id))
|
||||
.await?;
|
||||
let role_names = roles
|
||||
.iter()
|
||||
.map(|role| Ok((role.id.parse::<i64>()?, role.name.clone())))
|
||||
.collect::<Result<HashMap<_, _>, Box<dyn std::error::Error + Send + Sync>>>()?;
|
||||
let role_permissions = roles
|
||||
.iter()
|
||||
.map(|role| Ok((role.id.parse::<i64>()?, role.permissions.parse::<u64>()?)))
|
||||
.collect::<Result<HashMap<_, _>, Box<dyn std::error::Error + Send + Sync>>>()?;
|
||||
let user_roles = user
|
||||
.roles
|
||||
.iter()
|
||||
.map(|role| role.parse::<i64>())
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
let bot_roles = bot
|
||||
.roles
|
||||
.iter()
|
||||
.map(|role| role.parse::<i64>())
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
|
||||
for channel in channels {
|
||||
let channel_id = channel.id.parse::<i64>()?;
|
||||
if !archived_channels.contains(&channel_id) {
|
||||
continue;
|
||||
}
|
||||
let overwrites = channel.permission_overwrites.as_deref().unwrap_or_default();
|
||||
let reason = view_channel_reason(
|
||||
guild_id,
|
||||
user_id,
|
||||
&user_roles,
|
||||
&role_permissions,
|
||||
&role_names,
|
||||
overwrites,
|
||||
&guild_name,
|
||||
);
|
||||
if let Some(reason) = reason
|
||||
&& can_view_channel(guild_id, bot_id, &bot_roles, &role_permissions, overwrites)
|
||||
{
|
||||
visible.push(ChannelAccess { channel_id, reason });
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
visible.sort_unstable_by_key(|access| access.channel_id);
|
||||
Ok(visible)
|
||||
}
|
||||
|
||||
pub(super) async fn discord_get<T: serde::de::DeserializeOwned>(
|
||||
data: &WebData,
|
||||
path: &str,
|
||||
) -> Result<T, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let response = data
|
||||
.http
|
||||
.get(format!("https://discord.com/api/v10{}", path))
|
||||
.header("Authorization", format!("Bot {}", data.bot_token))
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?;
|
||||
Ok(response.json().await?)
|
||||
}
|
||||
|
||||
pub(super) async fn discord_get_optional<T: serde::de::DeserializeOwned>(
|
||||
data: &WebData,
|
||||
path: &str,
|
||||
) -> Result<Option<T>, Box<dyn std::error::Error + Send + Sync>> {
|
||||
let response = data
|
||||
.http
|
||||
.get(format!("https://discord.com/api/v10{}", path))
|
||||
.header("Authorization", format!("Bot {}", data.bot_token))
|
||||
.send()
|
||||
.await?;
|
||||
if response.status() == StatusCode::NOT_FOUND {
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(Some(response.error_for_status()?.json().await?))
|
||||
}
|
||||
|
||||
pub(super) fn can_view_channel(
|
||||
guild_id: i64,
|
||||
user_id: i64,
|
||||
member_roles: &[i64],
|
||||
roles: &HashMap<i64, u64>,
|
||||
overwrites: &[DiscordOverwrite],
|
||||
) -> bool {
|
||||
const ADMINISTRATOR: u64 = 1 << 3;
|
||||
const VIEW_CHANNEL: u64 = 1 << 10;
|
||||
let mut permissions = roles.get(&guild_id).copied().unwrap_or_default();
|
||||
for role_id in member_roles {
|
||||
permissions |= roles.get(role_id).copied().unwrap_or_default();
|
||||
}
|
||||
if permissions & ADMINISTRATOR != 0 {
|
||||
return true;
|
||||
}
|
||||
apply_overwrite(&mut permissions, overwrites, guild_id, 0);
|
||||
let mut allow = 0;
|
||||
let mut deny = 0;
|
||||
for overwrite in overwrites {
|
||||
let Ok(role_id) = overwrite.id.parse::<i64>() else {
|
||||
continue;
|
||||
};
|
||||
if overwrite.kind == 0 && member_roles.contains(&role_id) {
|
||||
allow |= overwrite.allow.parse::<u64>().unwrap_or_default();
|
||||
deny |= overwrite.deny.parse::<u64>().unwrap_or_default();
|
||||
}
|
||||
}
|
||||
permissions &= !deny;
|
||||
permissions |= allow;
|
||||
apply_overwrite(&mut permissions, overwrites, user_id, 1);
|
||||
permissions & VIEW_CHANNEL != 0
|
||||
}
|
||||
|
||||
pub(super) fn view_channel_reason(
|
||||
guild_id: i64,
|
||||
user_id: i64,
|
||||
member_roles: &[i64],
|
||||
roles: &HashMap<i64, u64>,
|
||||
role_names: &HashMap<i64, String>,
|
||||
overwrites: &[DiscordOverwrite],
|
||||
guild_name: &str,
|
||||
) -> Option<String> {
|
||||
const ADMINISTRATOR: u64 = 1 << 3;
|
||||
const VIEW_CHANNEL: u64 = 1 << 10;
|
||||
if !can_view_channel(guild_id, user_id, member_roles, roles, overwrites) {
|
||||
return None;
|
||||
}
|
||||
for role_id in member_roles {
|
||||
if roles.get(role_id).copied().unwrap_or_default() & ADMINISTRATOR != 0 {
|
||||
let role_name = role_names
|
||||
.get(role_id)
|
||||
.map(String::as_str)
|
||||
.unwrap_or("unknown");
|
||||
return Some(format!(
|
||||
"you have the {} administrator role in {}",
|
||||
role_name, guild_name
|
||||
));
|
||||
}
|
||||
}
|
||||
let member_allows_view = overwrites.iter().any(|overwrite| {
|
||||
overwrite.kind == 1
|
||||
&& overwrite.id.parse::<i64>() == Ok(user_id)
|
||||
&& overwrite.allow.parse::<u64>().unwrap_or_default() & VIEW_CHANNEL != 0
|
||||
});
|
||||
if member_allows_view {
|
||||
return Some(format!(
|
||||
"you have a channel-specific permission in {}",
|
||||
guild_name
|
||||
));
|
||||
}
|
||||
for role_id in member_roles {
|
||||
let role_allows_view = overwrites.iter().any(|overwrite| {
|
||||
overwrite.kind == 0
|
||||
&& overwrite.id.parse::<i64>() == Ok(*role_id)
|
||||
&& overwrite.allow.parse::<u64>().unwrap_or_default() & VIEW_CHANNEL != 0
|
||||
});
|
||||
if role_allows_view {
|
||||
let role_name = role_names
|
||||
.get(role_id)
|
||||
.map(String::as_str)
|
||||
.unwrap_or("unknown");
|
||||
return Some(format!("you have the {} role in {}", role_name, guild_name));
|
||||
}
|
||||
}
|
||||
let everyone_allows_view = overwrites.iter().any(|overwrite| {
|
||||
overwrite.kind == 0
|
||||
&& overwrite.id.parse::<i64>() == Ok(guild_id)
|
||||
&& overwrite.allow.parse::<u64>().unwrap_or_default() & VIEW_CHANNEL != 0
|
||||
});
|
||||
if everyone_allows_view {
|
||||
return Some(format!("you are a member of {}", guild_name));
|
||||
}
|
||||
for role_id in member_roles {
|
||||
if roles.get(role_id).copied().unwrap_or_default() & VIEW_CHANNEL != 0 {
|
||||
let role_name = role_names
|
||||
.get(role_id)
|
||||
.map(String::as_str)
|
||||
.unwrap_or("unknown");
|
||||
return Some(format!("you have the {} role in {}", role_name, guild_name));
|
||||
}
|
||||
}
|
||||
Some(format!("you are a member of {}", guild_name))
|
||||
}
|
||||
|
||||
pub(super) fn apply_overwrite(
|
||||
permissions: &mut u64,
|
||||
overwrites: &[DiscordOverwrite],
|
||||
id: i64,
|
||||
kind: u8,
|
||||
) {
|
||||
let overwrite = overwrites
|
||||
.iter()
|
||||
.find(|overwrite| overwrite.kind == kind && overwrite.id.parse::<i64>() == Ok(id));
|
||||
if let Some(overwrite) = overwrite {
|
||||
*permissions &= !overwrite.deny.parse::<u64>().unwrap_or_default();
|
||||
*permissions |= overwrite.allow.parse::<u64>().unwrap_or_default();
|
||||
}
|
||||
}
|
||||
+1081
File diff suppressed because it is too large.
Load diff
+1262
File diff suppressed because it is too large.
Load diff
Reference in new issue
Block a user