use std::time::Duration;
use reqwest::StatusCode;
use serde_json::Value;
use serde_json::json;
use crate::discord::ids::{
Id,
marker::{ChannelMarker, ForumTagMarker, GuildMarker},
};
use crate::{
AppError, Result,
discord::{
ChannelInfo, ForumPostArchiveState, ForumPostCreate, MessageAttachmentUpload, MessageInfo,
gateway::{parse_channel_info, parse_message_info},
},
};
use super::messages::{message_multipart_form, validate_message_payload};
use super::{DiscordRest, clone_array, extra_fields};
const FORUM_POST_SEARCH_PAGE_LIMIT: u16 = 25;
const FORUM_POST_SEARCH_MAX_ATTEMPTS: usize = 3;
const FORUM_POST_SEARCH_DEFAULT_RETRY: Duration = Duration::from_secs(5);
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct ForumPostPage {
pub threads: Vec<ChannelInfo>,
pub first_messages: Vec<MessageInfo>,
pub has_more: bool,
pub next_offset: usize,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct CreatedForumPost {
pub thread: ChannelInfo,
pub first_message: Option<MessageInfo>,
}
impl DiscordRest {
pub async fn follow_thread(&self, thread_id: Id<ChannelMarker>) -> Result<()> {
self.send_unit(
self.raw_http.put(format!(
"https://discord.com/api/v9/channels/{}/thread-members/@me",
thread_id.get()
)),
"follow post",
)
.await
}
pub async fn unfollow_thread(&self, thread_id: Id<ChannelMarker>) -> Result<()> {
self.send_unit(
self.raw_http.delete(format!(
"https://discord.com/api/v9/channels/{}/thread-members/@me",
thread_id.get()
)),
"unfollow post",
)
.await
}
pub async fn set_thread_archived(
&self,
thread_id: Id<ChannelMarker>,
archived: bool,
) -> Result<()> {
self.edit_thread(thread_id, &json!({ "archived": archived }))
.await
}
pub async fn set_thread_locked(
&self,
thread_id: Id<ChannelMarker>,
locked: bool,
) -> Result<()> {
self.edit_thread(thread_id, &json!({ "locked": locked }))
.await
}
pub async fn set_thread_pinned(
&self,
thread_id: Id<ChannelMarker>,
pinned: bool,
current_flags: u64,
) -> Result<()> {
const THREAD_FLAG_PINNED: u64 = 1 << 1;
let flags = if pinned {
current_flags | THREAD_FLAG_PINNED
} else {
current_flags & !THREAD_FLAG_PINNED
};
self.edit_thread(thread_id, &json!({ "flags": flags }))
.await
}
pub async fn edit_thread_settings(
&self,
thread_id: Id<ChannelMarker>,
name: &str,
applied_tags: &[Id<ForumTagMarker>],
rate_limit_per_user: Option<u64>,
auto_archive_duration: u64,
) -> Result<()> {
let mut body = json!({
"name": name,
"applied_tags": applied_tags
.iter()
.map(|tag_id| Value::String(tag_id.to_string()))
.collect::<Vec<_>>(),
"auto_archive_duration": auto_archive_duration,
});
if let Some(rate_limit_per_user) = rate_limit_per_user {
body.as_object_mut()
.expect("thread edit body is an object")
.insert(
"rate_limit_per_user".to_owned(),
Value::from(rate_limit_per_user),
);
}
self.edit_thread(thread_id, &body).await
}
pub async fn delete_thread(&self, thread_id: Id<ChannelMarker>) -> Result<()> {
self.send_unit(
self.raw_http.delete(format!(
"https://discord.com/api/v9/channels/{}",
thread_id.get()
)),
"delete thread",
)
.await
}
async fn edit_thread(&self, thread_id: Id<ChannelMarker>, body: &Value) -> Result<()> {
self.send_unit(
self.raw_http
.patch(format!(
"https://discord.com/api/v9/channels/{}",
thread_id.get()
))
.json(body),
"edit thread",
)
.await
}
pub async fn create_forum_post(
&self,
post: &ForumPostCreate,
upload_limit: u64,
slow_mode: Option<Duration>,
) -> Result<CreatedForumPost> {
let _channel_guard = self.message_sends.acquire(post.channel_id).await;
self.message_sends
.ensure_cooldown_elapsed(post.channel_id)?;
let body = create_forum_post_request_body(
&post.title,
&post.content,
&post.applied_tags,
&post.attachments,
upload_limit,
)?;
let request = self.raw_http.post(format!(
"https://discord.com/api/v9/channels/{}/threads",
post.channel_id.get()
));
let request = if post.attachments.is_empty() {
request.json(&body)
} else {
request.multipart(message_multipart_form(body, &post.attachments, upload_limit).await?)
};
let result = self
.send_json(request, "create forum post")
.await
.and_then(|raw: Value| parse_create_forum_post_response(&raw, Some(post.channel_id)));
match &result {
Ok(_) => {
if let Some(slow_mode) = slow_mode {
self.message_sends
.record_cooldown(post.channel_id, slow_mode);
}
}
Err(AppError::DiscordRateLimited {
retry_after_millis, ..
}) => self
.message_sends
.record_cooldown(post.channel_id, Duration::from_millis(*retry_after_millis)),
Err(_) => {}
}
result
}
pub async fn load_forum_posts(
&self,
guild_id: Id<GuildMarker>,
channel_id: Id<ChannelMarker>,
archive_state: ForumPostArchiveState,
offset: usize,
) -> Result<ForumPostPage> {
if offset == 0 {
let harvest_pins = archive_state == ForumPostArchiveState::Active;
let (activity, recent, pins) = tokio::join!(
self.load_forum_post_search_page(
guild_id,
channel_id,
archive_state,
offset,
ForumSearchSort::LastMessageTime,
),
self.load_forum_post_search_page(
guild_id,
channel_id,
archive_state,
offset,
ForumSearchSort::CreationTime,
),
async {
if harvest_pins {
self.load_forum_post_search_page(
guild_id,
channel_id,
archive_state,
offset,
ForumSearchSort::Relevance,
)
.await
.ok()
} else {
None
}
},
);
let page = merge_forum_pages(activity?, recent?);
return Ok(match pins {
Some(pins) => merge_pinned_forum_posts(page, pins),
None => page,
});
}
self.load_forum_post_search_page(
guild_id,
channel_id,
archive_state,
offset,
ForumSearchSort::LastMessageTime,
)
.await
}
async fn load_forum_post_search_page(
&self,
guild_id: Id<GuildMarker>,
channel_id: Id<ChannelMarker>,
archive_state: ForumPostArchiveState,
offset: usize,
sort_by: ForumSearchSort,
) -> Result<ForumPostPage> {
for attempt in 0..FORUM_POST_SEARCH_MAX_ATTEMPTS {
match self
.request_forum_post_search_page(
guild_id,
channel_id,
archive_state,
offset,
sort_by,
)
.await
{
Ok(page) => return Ok(page),
Err(error) => {
let Some(delay) = forum_search_retry_after(&error) else {
return Err(error);
};
if attempt + 1 == FORUM_POST_SEARCH_MAX_ATTEMPTS {
return Err(error);
}
tokio::time::sleep(delay).await;
}
}
}
unreachable!("forum search attempt loop always returns")
}
async fn request_forum_post_search_page(
&self,
guild_id: Id<GuildMarker>,
channel_id: Id<ChannelMarker>,
archive_state: ForumPostArchiveState,
offset: usize,
sort_by: ForumSearchSort,
) -> Result<ForumPostPage> {
let request = self
.raw_http
.get(format!(
"https://discord.com/api/v9/channels/{}/threads/search",
channel_id.get()
))
.query(&[
("archived", archive_state.as_query_value().to_owned()),
("sort_by", sort_by.as_str().to_owned()),
("sort_order", "desc".to_owned()),
("limit", FORUM_POST_SEARCH_PAGE_LIMIT.to_string()),
("tag_setting", "match_some".to_owned()),
("offset", offset.to_string()),
]);
let response = self
.execute_authenticated(request, "forum post search")
.await?;
if response.status() == StatusCode::ACCEPTED {
let raw = response.json::<Value>().await.unwrap_or(Value::Null);
let delay = raw
.get("retry_after")
.and_then(Value::as_f64)
.and_then(super::rate_limit_delay)
.filter(|delay| !delay.is_zero())
.unwrap_or(FORUM_POST_SEARCH_DEFAULT_RETRY);
return Err(AppError::ForumSearchIndexWarming {
retry_after_millis: super::duration_millis_ceil(delay),
});
}
if let Err(error) = response.error_for_status_ref() {
return Err(super::request_error(error, response, "forum post search").await);
}
let raw: Value = response.json().await.map_err(|error| {
AppError::DiscordRequest(format!("forum post search decode failed: {error}"))
})?;
let response = parse_forum_thread_search_response(&raw, Some(guild_id), channel_id, true);
let has_more = forum_search_has_more(&response);
Ok(ForumPostPage {
next_offset: forum_search_next_offset(offset, &response),
threads: response.threads,
first_messages: response.first_messages,
has_more,
})
}
}
pub(super) fn forum_search_next_offset(
offset: usize,
response: &ForumThreadSearchResponse,
) -> usize {
offset.saturating_add(response.raw_threads.len())
}
pub(super) fn forum_search_has_more(response: &ForumThreadSearchResponse) -> bool {
response.has_more && !response.raw_threads.is_empty()
}
pub(super) fn create_forum_post_request_body(
title: &str,
content: &str,
applied_tags: &[Id<ForumTagMarker>],
attachments: &[MessageAttachmentUpload],
upload_limit: u64,
) -> Result<Value> {
let title = validate_forum_post_title(title)?;
validate_message_payload(content, attachments, upload_limit)?;
let mut body = json!({
"name": title,
"message": {
"content": content,
},
});
if !applied_tags.is_empty() {
body["applied_tags"] = Value::Array(
applied_tags
.iter()
.map(|tag_id| Value::String(tag_id.to_string()))
.collect(),
);
}
if !attachments.is_empty() {
body["message"]["attachments"] = Value::Array(
attachments
.iter()
.enumerate()
.map(|(index, attachment)| {
json!({
"id": index,
"filename": attachment.filename,
})
})
.collect(),
);
}
Ok(body)
}
fn validate_forum_post_title(title: &str) -> Result<&str> {
let title = title.trim();
let len = title.chars().count();
if len == 0 {
return Err(AppError::DiscordRequest(
"forum post title cannot be empty".to_owned(),
));
}
if len > 100 {
return Err(AppError::DiscordRequest(format!(
"forum post title is too long: {len}/100"
)));
}
Ok(title)
}
pub(super) fn parse_create_forum_post_response(
raw: &Value,
parent_channel_id: Option<Id<ChannelMarker>>,
) -> Result<CreatedForumPost> {
let mut thread = parse_channel_info(raw, None).ok_or_else(|| {
AppError::DiscordRequest("create forum post response was missing thread".to_owned())
})?;
if thread.parent_id.is_none() {
thread.parent_id = parent_channel_id;
}
let first_message = raw.get("message").and_then(parse_message_info);
Ok(CreatedForumPost {
thread,
first_message,
})
}
#[derive(Clone, Debug, PartialEq)]
pub(super) struct ForumThreadSearchResponse {
pub(super) threads: Vec<ChannelInfo>,
pub(super) first_messages: Vec<MessageInfo>,
pub(super) has_more: bool,
pub(super) raw_threads: Vec<Value>,
pub(super) raw_first_messages: Vec<Value>,
pub(super) extra_fields: std::collections::BTreeMap<String, Value>,
}
pub(super) fn parse_forum_thread_search_response(
raw: &Value,
guild_id: Option<Id<GuildMarker>>,
parent_channel_id: Id<ChannelMarker>,
fill_missing_parent: bool,
) -> ForumThreadSearchResponse {
let threads = parse_forum_threads(raw, guild_id, parent_channel_id, fill_missing_parent);
let first_messages = parse_forum_first_messages(raw, &threads);
ForumThreadSearchResponse {
threads,
first_messages,
has_more: raw
.get("has_more")
.and_then(Value::as_bool)
.unwrap_or(false),
raw_threads: clone_array(raw.get("threads")),
raw_first_messages: clone_array(raw.get("first_messages")),
extra_fields: extra_fields(raw, &["threads", "first_messages", "has_more"]),
}
}
pub(super) fn parse_forum_threads(
raw: &Value,
guild_id: Option<Id<GuildMarker>>,
parent_channel_id: Id<ChannelMarker>,
fill_missing_parent: bool,
) -> Vec<ChannelInfo> {
raw.get("threads")
.and_then(Value::as_array)
.map(|threads| {
threads
.iter()
.filter_map(|thread| {
let mut channel = parse_channel_info(thread, guild_id)?;
if fill_missing_parent && channel.parent_id.is_none() {
channel.parent_id = Some(parent_channel_id);
}
if channel.parent_id != Some(parent_channel_id) {
return None;
}
Some(channel)
})
.collect()
})
.unwrap_or_default()
}
pub(super) fn parse_forum_first_messages(raw: &Value, threads: &[ChannelInfo]) -> Vec<MessageInfo> {
let mut seen = std::collections::HashSet::new();
parse_forum_messages_from_field(raw, threads, "first_messages")
.into_iter()
.filter(|message| seen.insert(message.message_id))
.collect()
}
fn parse_forum_messages_from_field(
raw: &Value,
threads: &[ChannelInfo],
field: &str,
) -> Vec<MessageInfo> {
raw.get(field)
.and_then(Value::as_array)
.map(|messages| {
messages
.iter()
.filter_map(parse_message_info)
.filter(|message| {
threads
.iter()
.any(|thread| thread.channel_id == message.channel_id)
})
.collect()
})
.unwrap_or_default()
}
pub(super) fn forum_search_retry_after(error: &AppError) -> Option<Duration> {
match error {
AppError::ForumSearchIndexWarming { retry_after_millis } => {
Some(Duration::from_millis(*retry_after_millis))
}
_ => None,
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(super) enum ForumSearchSort {
LastMessageTime,
CreationTime,
Relevance,
}
impl ForumSearchSort {
pub(super) fn as_str(self) -> &'static str {
match self {
Self::LastMessageTime => "last_message_time",
Self::CreationTime => "creation_time",
Self::Relevance => "relevance",
}
}
}
pub(super) fn merge_forum_pages(active: ForumPostPage, recent: ForumPostPage) -> ForumPostPage {
let mut seen_threads = std::collections::HashSet::new();
let mut threads = Vec::with_capacity(active.threads.len() + recent.threads.len());
for thread in active.threads.into_iter().chain(recent.threads) {
if seen_threads.insert(thread.channel_id) {
threads.push(thread);
}
}
let mut seen_first_messages = std::collections::HashSet::new();
let mut first_messages =
Vec::with_capacity(active.first_messages.len() + recent.first_messages.len());
for message in active
.first_messages
.into_iter()
.chain(recent.first_messages)
{
if seen_first_messages.insert(message.message_id) {
first_messages.push(message);
}
}
ForumPostPage {
next_offset: active.next_offset,
threads,
first_messages,
has_more: active.has_more,
}
}
pub(super) fn merge_pinned_forum_posts(
mut page: ForumPostPage,
pins: ForumPostPage,
) -> ForumPostPage {
let mut seen_threads = page
.threads
.iter()
.map(|thread| thread.channel_id)
.collect::<std::collections::HashSet<_>>();
let mut new_pin_ids = std::collections::HashSet::new();
let mut pinned_threads = Vec::new();
for thread in pins.threads {
if !thread.thread_pinned().unwrap_or(false) {
continue;
}
if seen_threads.insert(thread.channel_id) {
new_pin_ids.insert(thread.channel_id);
pinned_threads.push(thread);
}
}
if pinned_threads.is_empty() {
return page;
}
let mut seen_messages = page
.first_messages
.iter()
.map(|message| message.message_id)
.collect::<std::collections::HashSet<_>>();
for message in pins.first_messages {
if new_pin_ids.contains(&message.channel_id) && seen_messages.insert(message.message_id) {
page.first_messages.push(message);
}
}
pinned_threads.extend(page.threads);
page.threads = pinned_threads;
page
}