Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9328183ce4 | ||
|
|
93d7382a46 | ||
|
|
c0792d54e9 | ||
|
|
2967b9cb5f |
No files matched your search
+1
-1
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "tg-archive-bot"
|
||||
version = "0.3.0"
|
||||
version = "0.4.0"
|
||||
edition = "2024"
|
||||
license = "GPL-3.0-or-later"
|
||||
description = "Bot for archiving messages in TeenGovernment and related servers"
|
||||
|
||||
+460
-26
@@ -1,4 +1,6 @@
|
||||
use super::{WebData, WebUser, accessible_channels};
|
||||
use super::{
|
||||
CHANNEL_ACCESS_TTL_SECONDS, ChannelAccess, WebData, WebUser, accessible_channels, safe_filename,
|
||||
};
|
||||
use axum::{
|
||||
Extension, Json, Router,
|
||||
extract::{ConnectInfo, Path, Query, Request, State},
|
||||
@@ -7,7 +9,7 @@ use axum::{
|
||||
response::{IntoResponse, Response},
|
||||
routing::{get, post},
|
||||
};
|
||||
use rand::{Rng, distr::Alphanumeric, RngExt};
|
||||
use rand::{RngExt, distr::Alphanumeric};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
@@ -21,6 +23,52 @@ use tower_sessions::Session;
|
||||
const RESULTS_PER_PAGE: i64 = 100;
|
||||
const TOKEN_RATE_LIMIT: Duration = Duration::from_secs(5 * 60);
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub(super) struct ApiPermissionCache {
|
||||
users: Arc<Mutex<HashMap<i64, (Instant, Vec<ChannelAccess>)>>>,
|
||||
}
|
||||
|
||||
impl ApiPermissionCache {
|
||||
fn get_at(&self, discord_id: i64, now: Instant) -> Option<Vec<ChannelAccess>> {
|
||||
let mut users = self.users.lock().unwrap_or_else(|error| error.into_inner());
|
||||
users.retain(|_, (refreshed_at, _)| {
|
||||
now.saturating_duration_since(*refreshed_at).as_secs() < CHANNEL_ACCESS_TTL_SECONDS
|
||||
});
|
||||
users
|
||||
.get(&discord_id)
|
||||
.map(|(_, channel_access)| channel_access.clone())
|
||||
}
|
||||
|
||||
fn get(&self, discord_id: i64) -> Option<Vec<ChannelAccess>> {
|
||||
self.get_at(discord_id, Instant::now())
|
||||
}
|
||||
|
||||
fn insert_at(&self, discord_id: i64, channel_access: Vec<ChannelAccess>, now: Instant) {
|
||||
self.users
|
||||
.lock()
|
||||
.unwrap_or_else(|error| error.into_inner())
|
||||
.insert(discord_id, (now, channel_access));
|
||||
}
|
||||
|
||||
fn insert(&self, discord_id: i64, channel_access: Vec<ChannelAccess>) {
|
||||
self.insert_at(discord_id, channel_access, Instant::now());
|
||||
}
|
||||
}
|
||||
|
||||
async fn cached_accessible_channels(
|
||||
data: &WebData,
|
||||
discord_id: i64,
|
||||
) -> Result<Vec<ChannelAccess>, Box<dyn std::error::Error + Send + Sync>> {
|
||||
if let Some(channel_access) = data.api_permission_cache.get(discord_id) {
|
||||
return Ok(channel_access);
|
||||
}
|
||||
|
||||
let channel_access = accessible_channels(data, discord_id).await?;
|
||||
data.api_permission_cache
|
||||
.insert(discord_id, channel_access.clone());
|
||||
Ok(channel_access)
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(super) struct TokenRateLimiter {
|
||||
attempts: Arc<Mutex<HashMap<IpAddr, Instant>>>,
|
||||
@@ -115,6 +163,28 @@ struct ChannelResponse {
|
||||
struct UserResponse {
|
||||
discord_id: i64,
|
||||
username: String,
|
||||
avatar_url: Option<String>,
|
||||
first_seen_at: Option<String>,
|
||||
last_seen_at: Option<String>,
|
||||
message_count: i64,
|
||||
server_count: i64,
|
||||
channel_count: i64,
|
||||
attachment_count: i64,
|
||||
embed_count: i64,
|
||||
usernames: Vec<UserNameResponse>,
|
||||
avatars: Vec<UserAvatarResponse>,
|
||||
}
|
||||
#[derive(Serialize, sqlx::FromRow)]
|
||||
struct UserNameResponse {
|
||||
username: String,
|
||||
first_seen_at: String,
|
||||
last_seen_at: String,
|
||||
}
|
||||
#[derive(Serialize, sqlx::FromRow)]
|
||||
struct UserAvatarResponse {
|
||||
url: String,
|
||||
first_seen_at: String,
|
||||
last_seen_at: String,
|
||||
}
|
||||
#[derive(Serialize, sqlx::FromRow)]
|
||||
struct AttachmentResponse {
|
||||
@@ -160,12 +230,31 @@ struct MessageSummary {
|
||||
content: Option<String>,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct FilteredMessageSummary {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
discord_id: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
author_id: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
author_username: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
server_id: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
server_name: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
channel_id: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
channel_name: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
timestamp: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
content: Option<Option<String>>,
|
||||
}
|
||||
#[derive(Serialize)]
|
||||
struct SearchResponse {
|
||||
query: String,
|
||||
page: i64,
|
||||
per_page: i64,
|
||||
total: i64,
|
||||
results: Vec<MessageSummary>,
|
||||
limit: i64,
|
||||
results: Vec<FilteredMessageSummary>,
|
||||
}
|
||||
#[derive(Default, Deserialize)]
|
||||
struct VersionQuery {
|
||||
@@ -176,7 +265,67 @@ struct SearchQuery {
|
||||
#[serde(default)]
|
||||
q: String,
|
||||
page: Option<i64>,
|
||||
limit: Option<i64>,
|
||||
filter: Option<String>,
|
||||
}
|
||||
|
||||
const SEARCH_FIELDS: [&str; 9] = [
|
||||
"discord_id",
|
||||
"author_id",
|
||||
"author_username",
|
||||
"server_id",
|
||||
"server_name",
|
||||
"channel_id",
|
||||
"channel_name",
|
||||
"timestamp",
|
||||
"content",
|
||||
];
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
enum MetadataType {
|
||||
#[default]
|
||||
Message,
|
||||
Server,
|
||||
Channel,
|
||||
User,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
enum MetadataLookup {
|
||||
#[default]
|
||||
MessageCount,
|
||||
FirstMessage,
|
||||
LastMessage,
|
||||
AttachmentCount,
|
||||
EmbedCount,
|
||||
// For channels or servers
|
||||
UserCount,
|
||||
// For servers
|
||||
ChannelCount,
|
||||
}
|
||||
|
||||
impl MetadataLookup {
|
||||
fn supports(self, metadata_type: MetadataType) -> bool {
|
||||
match self {
|
||||
Self::UserCount => {
|
||||
matches!(metadata_type, MetadataType::Server | MetadataType::Channel)
|
||||
}
|
||||
Self::ChannelCount => matches!(metadata_type, MetadataType::Server),
|
||||
_ => true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default, Deserialize)]
|
||||
struct MetadataQuery {
|
||||
#[serde(default)]
|
||||
mtype: MetadataType,
|
||||
id: i64,
|
||||
ltype: MetadataLookup,
|
||||
}
|
||||
|
||||
pub(super) fn router(data: WebData) -> Router {
|
||||
let protected = Router::new()
|
||||
.route("/me", get(me))
|
||||
@@ -184,8 +333,14 @@ pub(super) fn router(data: WebData) -> Router {
|
||||
.route("/view/server/{id}", get(view_server))
|
||||
.route("/view/channel/{id}", get(view_channel))
|
||||
.route("/view/user/{id}", get(view_user))
|
||||
// Added attachement and attachment because I make the mistake often
|
||||
.route("/view/attachment/{id}", get(view_attachment))
|
||||
.route("/view/attachement/{id}", get(view_attachment))
|
||||
.route("/download/attachment/{id}", get(download_attachment))
|
||||
.route("/download/attachement/{id}", get(download_attachment))
|
||||
.route("/search/timestamp", get(search_timestamp))
|
||||
.route("/search/content", get(search_content))
|
||||
.route("/metadata", get(metadata_lookup))
|
||||
.route_layer(middleware::from_fn_with_state(
|
||||
data.clone(),
|
||||
require_api_user,
|
||||
@@ -219,7 +374,7 @@ async fn require_api_user(
|
||||
return api_error(StatusCode::INTERNAL_SERVER_ERROR, "database error");
|
||||
}
|
||||
};
|
||||
let channel_ids = match accessible_channels(&data, discord_id).await {
|
||||
let channel_ids = match cached_accessible_channels(&data, discord_id).await {
|
||||
Ok(access) => access
|
||||
.into_iter()
|
||||
.map(|channel| channel.channel_id)
|
||||
@@ -319,6 +474,71 @@ async fn view_server(
|
||||
}))
|
||||
}
|
||||
|
||||
async fn view_attachment(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Path(id): Path<i64>,
|
||||
Query(query): Query<VersionQuery>,
|
||||
) -> ApiResult<AttachmentResponse> {
|
||||
let row = sqlx::query_as::<_, AttachmentResponse>(
|
||||
"SELECT attachment_id AS discord_id, filename, description, content_type, size FROM attachments a
|
||||
JOIN messages m ON m.message_id = a.message_id
|
||||
WHERE a.attachment_id = $1 AND (m.channel_id = ANY($2) OR m.author_id = $3)
|
||||
AND ($4::bigint IS NULL OR a.message_version = $4)
|
||||
ORDER BY a.message_version DESC
|
||||
LIMIT 1",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(&user.channel_ids)
|
||||
.bind(user.discord_id)
|
||||
.bind(query.version)
|
||||
.fetch_optional(&data.pool)
|
||||
.await
|
||||
.map_err(database)?
|
||||
.ok_or_else(not_found)?;
|
||||
Ok(Json(row))
|
||||
}
|
||||
|
||||
async fn download_attachment(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Path(id): Path<i64>,
|
||||
Query(query): Query<VersionQuery>,
|
||||
) -> Result<Response, (StatusCode, Json<ErrorResponse>)> {
|
||||
let row = sqlx::query_as::<_, (String, Option<String>, Vec<u8>)>(
|
||||
"SELECT a.filename, a.content_type, a.data FROM attachments a
|
||||
JOIN messages m ON m.message_id = a.message_id
|
||||
WHERE a.attachment_id = $1 AND (m.channel_id = ANY($2) OR m.author_id = $3)
|
||||
AND ($4::bigint IS NULL OR a.message_version = $4)
|
||||
ORDER BY a.message_version DESC
|
||||
LIMIT 1",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(&user.channel_ids)
|
||||
.bind(user.discord_id)
|
||||
.bind(query.version)
|
||||
.fetch_optional(&data.pool)
|
||||
.await
|
||||
.map_err(database)?
|
||||
.ok_or_else(not_found)?;
|
||||
|
||||
let content_type = row
|
||||
.1
|
||||
.unwrap_or_else(|| "application/octet-stream".to_owned());
|
||||
Ok((
|
||||
[
|
||||
(header::CONTENT_TYPE, content_type),
|
||||
(
|
||||
header::CONTENT_DISPOSITION,
|
||||
format!("attachment; filename=\"{}\"", safe_filename(&row.0)),
|
||||
),
|
||||
(header::X_CONTENT_TYPE_OPTIONS, "nosniff".to_owned()),
|
||||
],
|
||||
row.2,
|
||||
)
|
||||
.into_response())
|
||||
}
|
||||
|
||||
async fn view_channel(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
@@ -353,9 +573,66 @@ async fn view_user(
|
||||
.await
|
||||
.map_err(database)?
|
||||
.ok_or_else(not_found)?;
|
||||
let history = sqlx::query_as::<_, (Option<String>, Option<String>, Option<String>)>(
|
||||
"SELECT
|
||||
(ARRAY_AGG(discord_avatar_url ORDER BY last_seen_at DESC)
|
||||
FILTER (WHERE discord_avatar_url IS NOT NULL))[1],
|
||||
MIN(first_seen_at)::text,
|
||||
MAX(last_seen_at)::text
|
||||
FROM discord_user_history WHERE discord_id = $1",
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_one(&data.pool)
|
||||
.await
|
||||
.map_err(database)?;
|
||||
let counts = sqlx::query_as::<_, (i64, i64, i64, i64, i64)>(
|
||||
"SELECT COUNT(DISTINCT m.message_id), COUNT(DISTINCT m.guild_id),
|
||||
COUNT(DISTINCT m.channel_id), COUNT(DISTINCT a.uuid), COUNT(DISTINCT e.uuid)
|
||||
FROM messages m
|
||||
LEFT JOIN attachments a ON a.message_id = m.message_id
|
||||
LEFT JOIN embeds e ON e.message_id = m.message_id
|
||||
WHERE m.author_id = $1 AND (m.channel_id = ANY($2) OR m.author_id = $3)",
|
||||
)
|
||||
.bind(id)
|
||||
.bind(&user.channel_ids)
|
||||
.bind(user.discord_id)
|
||||
.fetch_one(&data.pool)
|
||||
.await
|
||||
.map_err(database)?;
|
||||
let usernames = sqlx::query_as::<_, UserNameResponse>(
|
||||
"SELECT discord_username AS username, MIN(first_seen_at)::text AS first_seen_at,
|
||||
MAX(last_seen_at)::text AS last_seen_at
|
||||
FROM discord_user_history WHERE discord_id = $1
|
||||
GROUP BY discord_username ORDER BY MIN(first_seen_at), discord_username",
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_all(&data.pool)
|
||||
.await
|
||||
.map_err(database)?;
|
||||
let avatars = sqlx::query_as::<_, UserAvatarResponse>(
|
||||
"SELECT discord_avatar_url AS url, MIN(first_seen_at)::text AS first_seen_at,
|
||||
MAX(last_seen_at)::text AS last_seen_at
|
||||
FROM discord_user_history
|
||||
WHERE discord_id = $1 AND discord_avatar_url IS NOT NULL
|
||||
GROUP BY discord_avatar_url ORDER BY MIN(first_seen_at), discord_avatar_url",
|
||||
)
|
||||
.bind(id)
|
||||
.fetch_all(&data.pool)
|
||||
.await
|
||||
.map_err(database)?;
|
||||
Ok(Json(UserResponse {
|
||||
discord_id: row.0,
|
||||
username: row.1,
|
||||
avatar_url: history.0,
|
||||
first_seen_at: history.1,
|
||||
last_seen_at: history.2,
|
||||
message_count: counts.0,
|
||||
server_count: counts.1,
|
||||
channel_count: counts.2,
|
||||
attachment_count: counts.3,
|
||||
embed_count: counts.4,
|
||||
usernames,
|
||||
avatars,
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -417,47 +694,163 @@ async fn search_timestamp(
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Query(query): Query<SearchQuery>,
|
||||
) -> ApiResult<SearchResponse> {
|
||||
search(&data, &user, query, true).await
|
||||
search(&data, &user, &query, &query.limit.unwrap_or(100), 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
|
||||
search(&data, &user, &query, &query.limit.unwrap_or(100), false).await
|
||||
}
|
||||
async fn search(
|
||||
data: &WebData,
|
||||
user: &ApiUser,
|
||||
query: SearchQuery,
|
||||
query: &SearchQuery,
|
||||
limit: &i64,
|
||||
timestamp: bool,
|
||||
) -> ApiResult<SearchResponse> {
|
||||
let page = query.page.unwrap_or(1).max(1);
|
||||
let limit = limit.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 fields = parse_search_filter(query.filter.as_deref())?;
|
||||
|
||||
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
|
||||
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)?;
|
||||
ORDER BY m.timestamp DESC, m.message_id DESC
|
||||
LIMIT $5"
|
||||
)
|
||||
.bind(&pattern)
|
||||
.bind(timestamp)
|
||||
.bind(&user.channel_ids)
|
||||
.bind(user.discord_id)
|
||||
.bind(limit)
|
||||
.fetch_all(&data.pool)
|
||||
.await
|
||||
.map_err(database)?;
|
||||
|
||||
let results = results
|
||||
.into_iter()
|
||||
.map(|result| FilteredMessageSummary {
|
||||
discord_id: fields[0].then_some(result.discord_id),
|
||||
author_id: fields[1].then_some(result.author_id),
|
||||
author_username: fields[2].then_some(result.author_username),
|
||||
server_id: fields[3].then_some(result.server_id),
|
||||
server_name: fields[4].then_some(result.server_name),
|
||||
channel_id: fields[5].then_some(result.channel_id),
|
||||
channel_name: fields[6].then_some(result.channel_name),
|
||||
timestamp: fields[7].then_some(result.timestamp),
|
||||
content: fields[8].then_some(result.content),
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(Json(SearchResponse {
|
||||
query: query.q,
|
||||
page,
|
||||
per_page: RESULTS_PER_PAGE,
|
||||
total,
|
||||
query: query.q.clone(),
|
||||
limit: *limit,
|
||||
results,
|
||||
}))
|
||||
}
|
||||
|
||||
fn parse_search_filter(
|
||||
filter: Option<&str>,
|
||||
) -> Result<[bool; 9], (StatusCode, Json<ErrorResponse>)> {
|
||||
let Some(filter) = filter else {
|
||||
return Ok([true; 9]);
|
||||
};
|
||||
let mut selected = [false; 9];
|
||||
for field in filter.split(',').map(str::trim) {
|
||||
let Some(index) = SEARCH_FIELDS
|
||||
.iter()
|
||||
.position(|candidate| *candidate == field)
|
||||
else {
|
||||
return Err(error(
|
||||
StatusCode::BAD_REQUEST,
|
||||
"invalid search filter field",
|
||||
));
|
||||
};
|
||||
selected[index] = true;
|
||||
}
|
||||
Ok(selected)
|
||||
}
|
||||
|
||||
async fn metadata_lookup(
|
||||
State(data): State<WebData>,
|
||||
Extension(user): Extension<ApiUser>,
|
||||
Query(query): Query<MetadataQuery>,
|
||||
) -> ApiResult<i64> {
|
||||
if !query.ltype.supports(query.mtype) {
|
||||
return Err(error(
|
||||
StatusCode::BAD_REQUEST,
|
||||
"metadata lookup is not supported for this type",
|
||||
));
|
||||
}
|
||||
let entity_filter = match query.mtype {
|
||||
MetadataType::Message => "m.message_id = $1",
|
||||
MetadataType::Server => "m.guild_id = $1",
|
||||
MetadataType::Channel => "m.channel_id = $1",
|
||||
MetadataType::User => "m.author_id = $1",
|
||||
};
|
||||
let aggregate = match query.ltype {
|
||||
MetadataLookup::MessageCount => "COUNT(*)",
|
||||
MetadataLookup::FirstMessage => {
|
||||
"(ARRAY_AGG(message_id ORDER BY timestamp ASC, message_id ASC))[1]"
|
||||
}
|
||||
MetadataLookup::LastMessage => {
|
||||
"(ARRAY_AGG(message_id ORDER BY timestamp DESC, message_id DESC))[1]"
|
||||
}
|
||||
MetadataLookup::AttachmentCount => {
|
||||
"(SELECT COUNT(*) FROM attachments a WHERE a.message_id = visible.message_id)"
|
||||
}
|
||||
MetadataLookup::EmbedCount => {
|
||||
"(SELECT COUNT(*) FROM embeds e WHERE e.message_id = visible.message_id)"
|
||||
}
|
||||
MetadataLookup::UserCount => "COUNT(DISTINCT author_id)",
|
||||
MetadataLookup::ChannelCount => "COUNT(DISTINCT channel_id)",
|
||||
};
|
||||
|
||||
let statement = if matches!(
|
||||
query.ltype,
|
||||
MetadataLookup::AttachmentCount | MetadataLookup::EmbedCount
|
||||
) {
|
||||
format!(
|
||||
"SELECT COALESCE(SUM({aggregate}), 0)::bigint FROM (
|
||||
SELECT m.message_id, m.author_id, m.channel_id, m.timestamp
|
||||
FROM messages m
|
||||
WHERE {entity_filter}
|
||||
AND (m.channel_id = ANY($2) OR m.author_id = $3)
|
||||
) visible"
|
||||
)
|
||||
} else {
|
||||
format!(
|
||||
"SELECT {aggregate} FROM (
|
||||
SELECT m.message_id, m.author_id, m.channel_id, m.timestamp
|
||||
FROM messages m
|
||||
WHERE {entity_filter}
|
||||
AND (m.channel_id = ANY($2) OR m.author_id = $3)
|
||||
) visible"
|
||||
)
|
||||
};
|
||||
|
||||
let value = sqlx::query_scalar::<_, Option<i64>>(sqlx::AssertSqlSafe(statement))
|
||||
.bind(query.id)
|
||||
.bind(&user.channel_ids)
|
||||
.bind(user.discord_id)
|
||||
.fetch_one(&data.pool)
|
||||
.await
|
||||
.map_err(database)?
|
||||
.ok_or_else(not_found)?;
|
||||
|
||||
Ok(Json(value))
|
||||
}
|
||||
|
||||
fn bearer_token(value: Option<&axum::http::HeaderValue>) -> Option<&str> {
|
||||
let value = value?.to_str().ok()?;
|
||||
let (scheme, token) = value.split_once(' ')?;
|
||||
@@ -525,6 +918,22 @@ mod tests {
|
||||
assert_eq!(bearer_token(Some(&spaced)), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn search_filter_selects_requested_fields() {
|
||||
let selected = parse_search_filter(Some("discord_id, content"));
|
||||
assert_eq!(
|
||||
selected.ok(),
|
||||
Some([true, false, false, false, false, false, false, false, true])
|
||||
);
|
||||
assert_eq!(parse_search_filter(None).ok(), Some([true; 9]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn search_filter_rejects_unknown_and_empty_fields() {
|
||||
assert!(parse_search_filter(Some("unknown")).is_err());
|
||||
assert!(parse_search_filter(Some("")).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn only_trusts_forwarded_ip_from_configured_proxies() {
|
||||
let mut headers = HeaderMap::new();
|
||||
@@ -549,4 +958,29 @@ mod tests {
|
||||
assert!(limiter.check(ip).is_err());
|
||||
assert_eq!(limiter.check("192.0.2.6".parse().unwrap()), Ok(()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn caches_api_permissions_by_user_id_until_the_ttl_expires() {
|
||||
let cache = ApiPermissionCache::default();
|
||||
let refreshed_at = Instant::now();
|
||||
let access = vec![ChannelAccess {
|
||||
channel_id: 10,
|
||||
reason: "test".into(),
|
||||
}];
|
||||
|
||||
cache.insert_at(1, access, refreshed_at);
|
||||
|
||||
assert_eq!(
|
||||
cache
|
||||
.get_at(1, refreshed_at + Duration::from_secs(299))
|
||||
.unwrap()[0]
|
||||
.channel_id,
|
||||
10
|
||||
);
|
||||
assert!(
|
||||
cache
|
||||
.get_at(1, refreshed_at + Duration::from_secs(300))
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
}
|
||||
+48
-3
@@ -156,12 +156,14 @@ pub(super) async fn request_is_allowed(pool: &PgPool, user: &WebUser, path: &str
|
||||
channel_id.is_some_and(|channel_id| user.channel_ids.contains(&channel_id))
|
||||
}
|
||||
|
||||
pub(super) async fn login(session: Session) -> Response {
|
||||
pub(super) async fn login(session: Session, Query(query): Query<LoginQuery>) -> Response {
|
||||
match session.get::<WebUser>("user").await {
|
||||
Ok(Some(_)) => Redirect::to("/").into_response(),
|
||||
Ok(None) => Html(render_page(
|
||||
"Log in",
|
||||
&render_template(&LoginTemplate),
|
||||
&render_template(&LoginTemplate {
|
||||
error: query.error.as_deref(),
|
||||
}),
|
||||
false,
|
||||
))
|
||||
.into_response(),
|
||||
@@ -354,6 +356,9 @@ pub(super) async fn discord_callback(
|
||||
session: Session,
|
||||
Query(query): Query<OAuthQuery>,
|
||||
) -> Response {
|
||||
if !query.code.is_some() {
|
||||
return oauth2_error_handling(query.error.as_deref().unwrap_or("none"));
|
||||
}
|
||||
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();
|
||||
@@ -365,7 +370,7 @@ pub(super) async fn discord_callback(
|
||||
("client_id", data.client_id.as_str()),
|
||||
("client_secret", data.client_secret.as_str()),
|
||||
("grant_type", "authorization_code"),
|
||||
("code", query.code.as_str()),
|
||||
("code", query.code.expect("Missing code in query").as_str()),
|
||||
("redirect_uri", data.redirect_uri.as_str()),
|
||||
])
|
||||
.send()
|
||||
@@ -424,6 +429,46 @@ pub(super) async fn discord_callback(
|
||||
Redirect::to("/").into_response()
|
||||
}
|
||||
|
||||
fn oauth2_error_handling(error_code: &str) -> Response {
|
||||
return match error_code {
|
||||
"access_denied" => Redirect::to(
|
||||
"/login?error=You denied the request. Please log in again and hit Authorize to continue."
|
||||
).into_response(),
|
||||
|
||||
"invalid_request" => Redirect::to(
|
||||
"/login?error=Error code: invalid_request. Please try again. If this keeps happening, report this issue."
|
||||
).into_response(),
|
||||
|
||||
"unauthorized_client" => Redirect::to(
|
||||
"/login?error=Error code: unauthorized_client. Please try again. If this keeps happening, report this issue."
|
||||
).into_response(),
|
||||
|
||||
"unsupported_response_type" => Redirect::to(
|
||||
"/login?error=Error code: unsupported_response_type. Please try again. If this keeps happening, report this issue."
|
||||
).into_response(),
|
||||
|
||||
"invalid_scope" => Redirect::to(
|
||||
"/login?error=Error code: invalid_scope. Please try again. If this keeps happening, report this issue. Also if this is just you messing with the oauth2 scope, stop."
|
||||
).into_response(),
|
||||
|
||||
"server_error" => Redirect::to(
|
||||
"/login?error=Discord encountered an internal error. Please try again."
|
||||
).into_response(),
|
||||
|
||||
"temporarily_unavailable" => Redirect::to(
|
||||
"/login?error=Discord authentication is temporarily unavailable. Please try again later and check discordstatus.com for updates."
|
||||
).into_response(),
|
||||
|
||||
"none" => Redirect::to(
|
||||
"/login?error=Error code: none. Please try again. If this keeps happening, report this issue. Also if this is just you messing with the oauth2 scope, stop."
|
||||
).into_response(),
|
||||
|
||||
_ => Redirect::to(
|
||||
"/login?error=Error code: unknown. Please try again. If this keeps happening, report this issue."
|
||||
).into_response(),
|
||||
};
|
||||
}
|
||||
|
||||
pub(super) async fn logout(session: Session) -> Redirect {
|
||||
let _ = session.delete().await;
|
||||
Redirect::to("/login")
|
||||
|
||||
+23
-3
@@ -12,7 +12,7 @@ use axum::{
|
||||
routing::get,
|
||||
};
|
||||
use axum_extra::extract::cookie::{Cookie, CookieJar, SameSite};
|
||||
use rand::{Rng, distr::Alphanumeric};
|
||||
use rand::{distr::Alphanumeric};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sqlx::PgPool;
|
||||
use std::collections::{HashMap, HashSet};
|
||||
@@ -43,6 +43,7 @@ struct WebData {
|
||||
client_secret: String,
|
||||
redirect_uri: String,
|
||||
token_rate_limiter: api::TokenRateLimiter,
|
||||
api_permission_cache: api::ApiPermissionCache,
|
||||
}
|
||||
|
||||
#[derive(Clone, Deserialize, Serialize)]
|
||||
@@ -89,7 +90,9 @@ struct TimezoneContext {
|
||||
|
||||
#[derive(Template)]
|
||||
#[template(path = "login.html")]
|
||||
struct LoginTemplate;
|
||||
struct LoginTemplate<'a> {
|
||||
error: Option<&'a str>,
|
||||
}
|
||||
|
||||
#[derive(Template)]
|
||||
#[template(path = "privacy.html")]
|
||||
@@ -431,9 +434,16 @@ struct ErrorTemplate<'a> {
|
||||
message: &'a str,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct LoginQuery {
|
||||
error: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct OAuthQuery {
|
||||
code: String,
|
||||
code: Option<String>,
|
||||
error: Option<String>,
|
||||
error_description: Option<String>,
|
||||
state: String,
|
||||
}
|
||||
|
||||
@@ -574,6 +584,7 @@ pub async fn run(
|
||||
client_secret,
|
||||
redirect_uri,
|
||||
token_rate_limiter,
|
||||
api_permission_cache: api::ApiPermissionCache::default(),
|
||||
};
|
||||
let archive = Router::new()
|
||||
.route("/", get(index))
|
||||
@@ -989,6 +1000,15 @@ mod tests {
|
||||
assert!(html.contains("<new>"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn login_page_renders_and_escapes_query_error() {
|
||||
let html = render_template(&LoginTemplate {
|
||||
error: Some("Login failed: <try again>"),
|
||||
});
|
||||
|
||||
assert!(html.contains("Login failed: <try again>"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn renders_short_pagination_without_arbitrary_page_form() {
|
||||
let pagination = Pagination::new(2, ITEMS_PER_PAGE * 3);
|
||||
|
||||
@@ -1,4 +1,15 @@
|
||||
<h2>Log in</h2>
|
||||
{% if let Some(error) = error %}
|
||||
<div style="
|
||||
background-color: #fee2e2;
|
||||
color: #991b1b;
|
||||
border: 1px solid #ef4444;
|
||||
border-radius: 6px;
|
||||
padding: 12px 16px;
|
||||
margin: 12px 0;">
|
||||
{{ error }}
|
||||
</div>
|
||||
{% endif %}
|
||||
<p>It is required to log in with Discord to view the archive.</p>
|
||||
<p>The login happens on the official Discord website. This site does not have access to your email or password.</p>
|
||||
<form method="get" action="/login/discord">
|
||||
|
||||
Reference in new issue
Block a user