Large 0.3 changes
CI and release / Detect release commit (push) Successful in 27s
CI and release / Run tests (push) Successful in 4m14s
CI and release / Build and publish container (push) Successful in 9m2s
CI and release / Create release (push) Skipped

- 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:
Elias Wendland committed 2026-08-31 17:10:02 +02:00
1 parent 96a8c9ea8c
commit 126b9e9e9f
23 files changed
+3876 -2767

No files matched your search

+57
View File
@@ -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.
+4
View File
@@ -0,0 +1,4 @@
pub mod commands;
pub mod event_handler;
pub mod message_archive;
pub mod messaging;
+8 -10
View File
@@ -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
View File
File diff suppressed because it is too large. Load diff
+552
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because it is too large. Load diff
+1262
View File
File diff suppressed because it is too large. Load diff