use std::time::Duration;
use chrono::{NaiveDate, TimeZone, Utc};
use reqwest::StatusCode;
use serde_json::Value;
use crate::{
AppError, Result,
discord::{MessageInfo, MessageSearchPage, MessageSearchQuery, gateway::parse_message_info},
};
use super::{DiscordRest, clone_array, extra_fields};
const MESSAGE_SEARCH_PAGE_LIMIT: u16 = 25;
const MESSAGE_SEARCH_MAX_OFFSET: usize = 9_975;
const MESSAGE_SEARCH_MAX_ATTEMPTS: usize = 3;
const MESSAGE_SEARCH_DEFAULT_RETRY: Duration = Duration::from_secs(1);
const DISCORD_EPOCH_MILLIS: i64 = 1_420_070_400_000;
impl DiscordRest {
pub async fn search_messages(&self, query: MessageSearchQuery) -> Result<MessageSearchPage> {
if query.is_empty() {
return Ok(MessageSearchPage {
query,
messages: Vec::new(),
total_results: Some(0),
has_more: false,
});
}
let endpoint = match (query.guild_id, query.channel_id) {
(Some(guild_id), _) => format!(
"https://discord.com/api/v9/guilds/{}/messages/search",
guild_id.get()
),
(None, Some(channel_id)) => format!(
"https://discord.com/api/v9/channels/{}/messages/search",
channel_id.get()
),
(None, None) => {
return Err(AppError::DiscordRequest(
"message search requires a server or channel".to_owned(),
));
}
};
let params = message_search_query_params(&query);
let mut raw = None;
for attempt in 0..MESSAGE_SEARCH_MAX_ATTEMPTS {
let response = self
.execute_authenticated(
self.raw_http.get(&endpoint).query(¶ms),
"message search",
)
.await?;
let status = response.status();
if status != StatusCode::ACCEPTED
&& let Err(error) = response.error_for_status_ref()
{
return Err(super::request_error(error, response, "message search").await);
}
let response_body: Value = response.json().await.map_err(|error| {
AppError::DiscordRequest(format!("message search decode failed: {error}"))
})?;
if status != StatusCode::ACCEPTED {
raw = Some(response_body);
break;
}
if attempt + 1 == MESSAGE_SEARCH_MAX_ATTEMPTS {
return Err(AppError::DiscordRequest(message_search_indexing_message(
&response_body,
)));
}
tokio::time::sleep(message_search_retry_delay(&response_body)).await;
}
let raw = raw.expect("message search attempt loop returns a response");
let response = parse_message_search_response(&raw)?;
let next_offset = query
.offset
.saturating_add(MESSAGE_SEARCH_PAGE_LIMIT as usize);
let has_more = message_search_has_more(&response, next_offset);
Ok(MessageSearchPage {
query,
messages: response.messages,
total_results: response.total_results,
has_more,
})
}
}
#[derive(Clone, Debug, PartialEq)]
pub(super) struct MessageSearchResponse {
pub(super) total_results: Option<usize>,
pub(super) messages: Vec<MessageInfo>,
pub(super) message_groups: Vec<Value>,
pub(super) extra_fields: std::collections::BTreeMap<String, Value>,
}
pub(super) fn parse_message_search_response(raw: &Value) -> Result<MessageSearchResponse> {
Ok(MessageSearchResponse {
total_results: raw
.get("total_results")
.and_then(Value::as_u64)
.map(|value| usize::try_from(value).unwrap_or(usize::MAX)),
messages: parse_message_search_messages(raw)?,
message_groups: clone_array(raw.get("messages")),
extra_fields: extra_fields(raw, &["total_results", "messages"]),
})
}
pub(super) fn message_search_query_params(
query: &MessageSearchQuery,
) -> Vec<(&'static str, String)> {
let mut params = vec![
("limit", MESSAGE_SEARCH_PAGE_LIMIT.clamp(1, 25).to_string()),
(
"offset",
query.offset.min(MESSAGE_SEARCH_MAX_OFFSET).to_string(),
),
("sort_by", "timestamp".to_owned()),
("sort_order", "desc".to_owned()),
];
if let Some(content) = query.content.as_deref().filter(|value| !value.is_empty()) {
params.push(("content", content.chars().take(1024).collect()));
}
if let Some(channel_id) = query.channel_id
&& query.guild_id.is_some()
{
params.push(("channel_id", channel_id.to_string()));
}
if let Some(author_id) = query.author_id {
params.push(("author_id", author_id.to_string()));
}
if let Some(user_id) = query.mentions_user_id {
params.push(("mentions", user_id.to_string()));
}
for has in &query.has {
params.push(("has", has.as_query_value().to_owned()));
}
for author_type in &query.author_type {
params.push(("author_type", author_type.as_query_value().to_owned()));
}
if let Some(pinned) = query.pinned {
params.push(("pinned", pinned.to_string()));
}
if let Some(bounds) = query
.date
.as_deref()
.and_then(message_search_date_snowflake_bounds)
{
if let Some(min_id) = bounds.min_id {
params.push(("min_id", min_id.to_string()));
}
if let Some(max_id) = bounds.max_id {
params.push(("max_id", max_id.to_string()));
}
}
params
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(super) struct MessageSearchDateBounds {
pub(super) min_id: Option<u64>,
pub(super) max_id: Option<u64>,
}
pub(super) fn message_search_date_snowflake_bounds(value: &str) -> Option<MessageSearchDateBounds> {
let mut min_id = None;
let mut max_id = None;
for token in value.split(',') {
let token = token.trim();
if token.is_empty() {
return None;
}
let (operator, date) = token
.split_once(':')
.map(|(operator, date)| (operator.trim(), date.trim()))
.unwrap_or(("equal", token));
let (lower, upper) = match operator {
"gte" => (
Some(message_search_date_start_snowflake(date)?.saturating_sub(1)),
None,
),
"lte" => (None, Some(message_search_date_next_snowflake(date)?)),
"equal" => (
Some(message_search_date_start_snowflake(date)?.saturating_sub(1)),
Some(message_search_date_next_snowflake(date)?),
),
_ => return None,
};
if let Some(lower) = lower {
min_id = Some(min_id.map_or(lower, |current: u64| current.max(lower)));
}
if let Some(upper) = upper {
max_id = Some(max_id.map_or(upper, |current: u64| current.min(upper)));
}
}
Some(MessageSearchDateBounds { min_id, max_id })
}
fn message_search_date_start_snowflake(value: &str) -> Option<u64> {
let date = NaiveDate::parse_from_str(value.trim(), "%Y-%m-%d").ok()?;
let start = date.and_hms_opt(0, 0, 0)?;
let start_millis = Utc.from_utc_datetime(&start).timestamp_millis();
Some(timestamp_millis_to_snowflake(start_millis))
}
fn message_search_date_next_snowflake(value: &str) -> Option<u64> {
let date = NaiveDate::parse_from_str(value.trim(), "%Y-%m-%d").ok()?;
let end = date.succ_opt()?.and_hms_opt(0, 0, 0)?;
let end_millis = Utc.from_utc_datetime(&end).timestamp_millis();
Some(timestamp_millis_to_snowflake(end_millis))
}
fn timestamp_millis_to_snowflake(timestamp_millis: i64) -> u64 {
let discord_millis = timestamp_millis.saturating_sub(DISCORD_EPOCH_MILLIS).max(0);
u64::try_from(discord_millis).unwrap_or_default() << 22
}
fn parse_message_search_messages(raw: &Value) -> Result<Vec<MessageInfo>> {
let Some(groups) = raw.get("messages").and_then(Value::as_array) else {
return Ok(Vec::new());
};
let mut messages = Vec::new();
for group in groups {
let Some(group_messages) = group.as_array() else {
continue;
};
for raw_message in group_messages {
let message = parse_message_info(raw_message).ok_or_else(|| {
AppError::DiscordRequest(
"search message response was missing required fields".to_owned(),
)
})?;
messages.push(message);
}
}
Ok(messages)
}
fn message_search_indexing_message(raw: &Value) -> String {
let retry_after = raw
.get("retry_after")
.and_then(Value::as_f64)
.unwrap_or(0.0);
if retry_after > 0.0 {
format!("message search index is not ready, retry after {retry_after:.1}s")
} else {
"message search index is not ready, try again shortly".to_owned()
}
}
pub(super) fn message_search_retry_delay(raw: &Value) -> Duration {
raw.get("retry_after")
.and_then(Value::as_f64)
.and_then(super::rate_limit_delay)
.filter(|delay| !delay.is_zero())
.unwrap_or(MESSAGE_SEARCH_DEFAULT_RETRY)
}
pub(super) fn message_search_has_more(
response: &MessageSearchResponse,
next_offset: usize,
) -> bool {
let reported_results_remain = response
.total_results
.is_some_and(|total_results| next_offset < total_results);
let full_page = response.message_groups.len() >= usize::from(MESSAGE_SEARCH_PAGE_LIMIT);
next_offset <= MESSAGE_SEARCH_MAX_OFFSET && (reported_results_remain || full_page)
}