mirror of
https://github.com/TheFunny/TelegramTwitterMediaBot.git
synced 2026-09-23 23:32:05 +00:00
finish rust rewrite, add docker, drop python
This commit is contained in:
@@ -0,0 +1,21 @@
|
||||
[package]
|
||||
name = "xmedia-bot"
|
||||
version = "0.1.0"
|
||||
edition = "2024"
|
||||
|
||||
[dependencies]
|
||||
teloxide = { version = "0.17", features = ["webhooks-axum", "macros"] }
|
||||
tokio = { version = "1.40", features = ["rt-multi-thread", "macros"] }
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
log = "0.4"
|
||||
pretty_env_logger = "0.5"
|
||||
dotenv = "0.15"
|
||||
url = "2.5.2"
|
||||
regex = "1.12"
|
||||
html-escape = "0.2"
|
||||
rusqlite = { version = "0.32", features = ["bundled"] }
|
||||
rand = "0.8"
|
||||
tempfile = "3"
|
||||
parking_lot = "0.12"
|
||||
x-media = { path = "../x-media" }
|
||||
@@ -0,0 +1,58 @@
|
||||
//! Central env handling. The only other places that read env are
|
||||
//! `Bot::from_env` (TELOXIDE_TOKEN) and x-media (PIXIV_REFRESH_TOKEN).
|
||||
|
||||
use std::env;
|
||||
use std::net::IpAddr;
|
||||
use std::time::Duration;
|
||||
|
||||
pub struct Config {
|
||||
/// BOT_ADMIN: comma-separated ints; empty when unset.
|
||||
pub admin_ids: Vec<i64>,
|
||||
/// EDIT_MESSAGE_TTL_SECONDS, default 86400 (24h).
|
||||
pub edit_message_ttl: Duration,
|
||||
// Webhook settings (moved out of main; names/defaults unchanged).
|
||||
pub webhook_enabled: bool,
|
||||
pub webhook_url: Option<url::Url>,
|
||||
pub webhook_listen: Option<IpAddr>,
|
||||
pub webhook_port: Option<u16>,
|
||||
pub webhook_cert: Option<String>,
|
||||
pub webhook_secret_token: Option<String>,
|
||||
}
|
||||
|
||||
impl Config {
|
||||
pub fn load() -> Config {
|
||||
let admin_ids = env::var("BOT_ADMIN")
|
||||
.ok()
|
||||
.map(|s| {
|
||||
s.split(',')
|
||||
.filter_map(|part| part.trim().parse::<i64>().ok())
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
|
||||
let edit_message_ttl = env::var("EDIT_MESSAGE_TTL_SECONDS")
|
||||
.ok()
|
||||
.and_then(|s| s.parse::<u64>().ok())
|
||||
.map(Duration::from_secs)
|
||||
.unwrap_or(Duration::from_secs(86400));
|
||||
|
||||
let webhook_enabled = env::var("WEBHOOK")
|
||||
.is_ok_and(|v| matches!(v.to_lowercase().as_str(), "true" | "yes" | "1"));
|
||||
let webhook_url = env::var("WEBHOOK_URL").ok().and_then(|s| s.parse().ok());
|
||||
let webhook_listen = env::var("WEBHOOK_LISTEN").ok().and_then(|s| s.parse().ok());
|
||||
let webhook_port = env::var("WEBHOOK_PORT").ok().and_then(|s| s.parse().ok());
|
||||
let webhook_cert = env::var("WEBHOOK_CERT").ok();
|
||||
let webhook_secret_token = env::var("WEBHOOK_SECRET_TOKEN").ok();
|
||||
|
||||
Config {
|
||||
admin_ids,
|
||||
edit_message_ttl,
|
||||
webhook_enabled,
|
||||
webhook_url,
|
||||
webhook_listen,
|
||||
webhook_port,
|
||||
webhook_cert,
|
||||
webhook_secret_token,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,639 @@
|
||||
use crate::config::Config;
|
||||
use crate::queue::PersistentTaskQueue;
|
||||
use crate::send::{self, MediaItemPayload, Task};
|
||||
use crate::state::{ChatStore, unix_now};
|
||||
use std::collections::HashSet;
|
||||
use std::sync::LazyLock;
|
||||
use teloxide::prelude::*;
|
||||
use teloxide::types::{
|
||||
CallbackQuery, ChatAction, ChatId, ChatKind, InlineQuery, InlineQueryResult,
|
||||
InlineQueryResultMpeg4Gif, InlineQueryResultPhoto, InlineQueryResultVideo, Message,
|
||||
MessageEntityKind, MessageId, ParseMode, Recipient, ReplyParameters,
|
||||
};
|
||||
use teloxide::utils::command::BotCommands;
|
||||
use teloxide::RequestError;
|
||||
use x_media::media::Media;
|
||||
|
||||
pub static CHAT_STORE: LazyLock<ChatStore> = LazyLock::new(|| {
|
||||
ChatStore::open("data/task_queue.db").expect("failed to open chat store")
|
||||
});
|
||||
pub static TASK_QUEUE: LazyLock<PersistentTaskQueue> =
|
||||
LazyLock::new(|| PersistentTaskQueue::new("data/task_queue.db"));
|
||||
pub static CONFIG: LazyLock<Config> = LazyLock::new(Config::load);
|
||||
|
||||
#[derive(BotCommands, Clone)]
|
||||
#[command(rename_rule = "snake_case", description = "")]
|
||||
enum Command {
|
||||
#[command(description = "")]
|
||||
Start,
|
||||
#[command(description = "")]
|
||||
Help,
|
||||
#[command(description = "", parse_with = "split")]
|
||||
SetForwardChannel(String),
|
||||
#[command(description = "")]
|
||||
RemoveForwardChannel,
|
||||
#[command(description = "")]
|
||||
EditBeforeForward,
|
||||
#[command(description = "", parse_with = "split")]
|
||||
SetTemplate(String),
|
||||
#[command(description = "")]
|
||||
BotDict,
|
||||
#[command(description = "", parse_with = "split")]
|
||||
SetFormat(String),
|
||||
}
|
||||
|
||||
async fn reply<T>(bot: Bot, message: Message, text: T) -> Result<Message, RequestError>
|
||||
where
|
||||
T: Into<String>,
|
||||
{
|
||||
bot.send_message(message.chat.id, text)
|
||||
.reply_parameters(ReplyParameters::new(message.id).allow_sending_without_reply())
|
||||
.await
|
||||
}
|
||||
|
||||
fn now_f64() -> f64 {
|
||||
std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.map(|d| d.as_secs_f64())
|
||||
.unwrap_or(0.0)
|
||||
}
|
||||
|
||||
/// Extracts URL and text-link entities (text + caption), deduped in order.
|
||||
pub fn extract_urls(message: &Message) -> Vec<String> {
|
||||
let mut urls = Vec::new();
|
||||
for entity in message.parse_entities().into_iter().flatten() {
|
||||
match entity.kind() {
|
||||
MessageEntityKind::Url => urls.push(entity.text().to_string()),
|
||||
MessageEntityKind::TextLink { url } => urls.push(url.to_string()),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
for entity in message.parse_caption_entities().into_iter().flatten() {
|
||||
match entity.kind() {
|
||||
MessageEntityKind::Url => urls.push(entity.text().to_string()),
|
||||
MessageEntityKind::TextLink { url } => urls.push(url.to_string()),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
let mut seen = HashSet::new();
|
||||
urls.retain(|url| seen.insert(url.clone()));
|
||||
urls
|
||||
}
|
||||
|
||||
/// Edit-before-forward: a reply to the prompt swaps the caption of the first
|
||||
/// forwarded message. Returns true when the message was consumed as an edit.
|
||||
async fn edit_message_handler(bot: &Bot, message: &Message) -> bool {
|
||||
let Some(reply) = message.reply_to_message() else {
|
||||
return false;
|
||||
};
|
||||
let chat_id = message.chat.id.0;
|
||||
let Some(text) = message.text() else {
|
||||
return false;
|
||||
};
|
||||
let chat_data = CHAT_STORE.get(chat_id).await;
|
||||
let Some(edit) = chat_data.edit_message.get(&(reply.id.0 as i64)) else {
|
||||
return false;
|
||||
};
|
||||
let Some(first_forward_id) = edit.forward_message_ids.first() else {
|
||||
return false;
|
||||
};
|
||||
let link = format!(
|
||||
"<a href=\"{0}\">{1}</a>",
|
||||
edit.url,
|
||||
html_escape::encode_text(text)
|
||||
);
|
||||
let new_text = if edit.template.is_empty() {
|
||||
link
|
||||
} else {
|
||||
chat_data
|
||||
.template
|
||||
.get(&edit.template)
|
||||
.map(|template| template.replace("[]", &link))
|
||||
.unwrap_or(link)
|
||||
};
|
||||
let result = bot
|
||||
.edit_message_caption(ChatId(chat_id), MessageId(*first_forward_id as i32))
|
||||
.caption(new_text)
|
||||
.parse_mode(ParseMode::Html)
|
||||
.await;
|
||||
match result {
|
||||
Ok(_) => log::info!(
|
||||
"edit-before-forward: caption swapped on message {first_forward_id} for prompt {}",
|
||||
reply.id.0
|
||||
),
|
||||
Err(e) => log::error!("edit_message_caption failed: {e}"),
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
enum SetForwardChannelError {
|
||||
EmptyParameter,
|
||||
NotChannel,
|
||||
NotAdmin,
|
||||
NotBotAdmin(RequestError),
|
||||
NotBotCanPost,
|
||||
}
|
||||
|
||||
async fn set_forward_channel_handler(
|
||||
bot: &Bot,
|
||||
message: &Message,
|
||||
channel: String,
|
||||
) -> Result<i64, SetForwardChannelError> {
|
||||
if channel.is_empty() {
|
||||
return Err(SetForwardChannelError::EmptyParameter);
|
||||
}
|
||||
let channel = match channel.parse::<i64>() {
|
||||
Ok(id) => Recipient::Id(ChatId(id)),
|
||||
Err(_) => Recipient::ChannelUsername(channel),
|
||||
};
|
||||
if let Some(from) = &message.from {
|
||||
log::info!(
|
||||
"Set forward channel for {} ({}) to {}",
|
||||
from.full_name(),
|
||||
message.chat.id,
|
||||
channel
|
||||
);
|
||||
}
|
||||
let chat = match bot.get_chat(channel.clone()).await {
|
||||
Err(e) => {
|
||||
log::error!("Failed to get channel {}: {}", channel, e);
|
||||
return Err(SetForwardChannelError::NotBotAdmin(e));
|
||||
}
|
||||
Ok(chat) => chat,
|
||||
};
|
||||
if !chat.is_channel() {
|
||||
return Err(SetForwardChannelError::NotChannel);
|
||||
}
|
||||
let channel_id = chat.id.0;
|
||||
match bot.get_chat_administrators(channel.clone()).await {
|
||||
Err(e) => {
|
||||
log::error!("Failed to get channel administrators {}: {}", channel, e);
|
||||
return Err(SetForwardChannelError::NotBotAdmin(e));
|
||||
}
|
||||
Ok(admins) => {
|
||||
if !admins.iter().any(|admin| admin.user.id == message.chat.id) {
|
||||
return Err(SetForwardChannelError::NotAdmin);
|
||||
}
|
||||
let bot_id = bot.get_me().await.expect("Failed get bot id").user.id;
|
||||
if let Some(bot_admin) = admins.iter().find(|admin| admin.user.id == bot_id)
|
||||
&& !bot_admin.can_post_messages()
|
||||
{
|
||||
return Err(SetForwardChannelError::NotBotCanPost);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(channel_id)
|
||||
}
|
||||
|
||||
async fn execute_command(bot: &Bot, message: &Message, command: Command) -> Result<(), RequestError> {
|
||||
match command {
|
||||
Command::Start => {
|
||||
bot.send_message(message.chat.id, "Hello!").await?;
|
||||
}
|
||||
Command::Help => {
|
||||
bot.send_message(message.chat.id, Command::descriptions().to_string())
|
||||
.await?;
|
||||
}
|
||||
Command::SetForwardChannel(channel) => {
|
||||
let result = match set_forward_channel_handler(bot, message, channel).await {
|
||||
Ok(channel_id) => {
|
||||
let mut chat_data = CHAT_STORE.get(message.chat.id.0).await;
|
||||
chat_data.forward_channel_id = Some(channel_id);
|
||||
CHAT_STORE.set(message.chat.id.0, &chat_data).await;
|
||||
"Add successfully.".to_string()
|
||||
}
|
||||
Err(SetForwardChannelError::EmptyParameter) => {
|
||||
"Receive empty parameter.\nYou should enter a channel id or username".to_string()
|
||||
}
|
||||
Err(SetForwardChannelError::NotChannel) => {
|
||||
"Given id / username is not a channel".to_string()
|
||||
}
|
||||
Err(SetForwardChannelError::NotAdmin) => {
|
||||
"You are not an administrator of the channel".to_string()
|
||||
}
|
||||
Err(SetForwardChannelError::NotBotAdmin(e)) => {
|
||||
e.to_string() + "\nPlease add the bot as an admin to the channel"
|
||||
}
|
||||
Err(SetForwardChannelError::NotBotCanPost) => {
|
||||
"Bot can't post messages to the channel".to_string()
|
||||
}
|
||||
};
|
||||
reply(bot.clone(), message.clone(), result).await?;
|
||||
}
|
||||
Command::RemoveForwardChannel => {
|
||||
let chat_id = message.chat.id.0;
|
||||
let mut chat_data = CHAT_STORE.get(chat_id).await;
|
||||
let text = if chat_data.forward_channel_id.is_some() {
|
||||
chat_data.forward_channel_id = None;
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
"Remove successfully.".to_string()
|
||||
} else {
|
||||
"No channel to remove.".to_string()
|
||||
};
|
||||
reply(bot.clone(), message.clone(), text).await?;
|
||||
}
|
||||
Command::EditBeforeForward => {
|
||||
let chat_id = message.chat.id.0;
|
||||
let mut chat_data = CHAT_STORE.get(chat_id).await;
|
||||
let text = if chat_data.forward_channel_id.is_none() {
|
||||
"Please enable forward channel first.".to_string()
|
||||
} else if chat_data.edit_before_forward {
|
||||
chat_data.edit_before_forward = false;
|
||||
chat_data.edit_message.clear();
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
"Disable edit before forward.".to_string()
|
||||
} else {
|
||||
chat_data.edit_before_forward = true;
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
"Enable edit before forward.".to_string()
|
||||
};
|
||||
reply(bot.clone(), message.clone(), text).await?;
|
||||
}
|
||||
Command::SetTemplate(name) => {
|
||||
let chat_id = message.chat.id.0;
|
||||
let text = match message.reply_to_message() {
|
||||
None => "Please reply to a message to set as template.".to_string(),
|
||||
Some(reply) => {
|
||||
let reply_text = reply.text().unwrap_or_default();
|
||||
if !reply_text.contains("[]") {
|
||||
"Please reply to a message with [] to set as template.".to_string()
|
||||
} else if name.is_empty() {
|
||||
"Please provide a name for the template.".to_string()
|
||||
} else {
|
||||
let mut chat_data = CHAT_STORE.get(chat_id).await;
|
||||
chat_data
|
||||
.template
|
||||
.insert(name, html_escape::encode_text(reply_text).into_owned());
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
"Template set.".to_string()
|
||||
}
|
||||
}
|
||||
};
|
||||
reply(bot.clone(), message.clone(), text).await?;
|
||||
}
|
||||
Command::BotDict => {
|
||||
let chat_data = CHAT_STORE.get(message.chat.id.0).await;
|
||||
let debug = format!("{chat_data:?}");
|
||||
let text = html_escape::encode_text(&debug).into_owned();
|
||||
reply(bot.clone(), message.clone(), text).await?;
|
||||
}
|
||||
Command::SetFormat(arg) => {
|
||||
let chat_id = message.chat.id.0;
|
||||
let (site, format) = match arg.split_once(char::is_whitespace) {
|
||||
Some((site, format)) if !format.trim().is_empty() => (site.trim(), format.trim().to_string()),
|
||||
_ => {
|
||||
reply(
|
||||
bot.clone(),
|
||||
message.clone(),
|
||||
"Usage: /set_format <site> <format>",
|
||||
)
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
if !["twitter", "bsky", "pixiv"].contains(&site) {
|
||||
reply(
|
||||
bot.clone(),
|
||||
message.clone(),
|
||||
"Unknown site. Use twitter, bsky or pixiv.",
|
||||
)
|
||||
.await?;
|
||||
return Ok(());
|
||||
}
|
||||
let mut chat_data = CHAT_STORE.get(chat_id).await;
|
||||
chat_data.message_format.insert(site.to_string(), format);
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
reply(bot.clone(), message.clone(), "Format set.").await?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// For locally produced media (encoded ugoira MP4) the thumbnail URL is a
|
||||
/// hotlink-protected remote URL Telegram may not fetch; let Telegram generate
|
||||
/// its own thumbnail instead.
|
||||
fn thumbnail_for(media: &Media) -> Option<String> {
|
||||
let url = media.url();
|
||||
if url.starts_with("http://") || url.starts_with("https://") {
|
||||
media.thumbnail_url().map(str::to_string)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn media_to_payload(media: &Media, sensitive: bool) -> MediaItemPayload {
|
||||
match media {
|
||||
// A gif inside a group becomes a video item; a lone gif takes the
|
||||
// animation path (see url_media).
|
||||
Media::Illustration { .. } => MediaItemPayload::Photo {
|
||||
media: media.url().to_string(),
|
||||
has_spoiler: sensitive,
|
||||
},
|
||||
Media::Video { .. } => MediaItemPayload::Video {
|
||||
media: media.url().to_string(),
|
||||
has_spoiler: sensitive,
|
||||
thumbnail: thumbnail_for(media),
|
||||
},
|
||||
Media::Animated { .. } => MediaItemPayload::Video {
|
||||
media: media.url().to_string(),
|
||||
has_spoiler: sensitive,
|
||||
thumbnail: thumbnail_for(media),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
async fn enqueue_retry(task: Task, delay_seconds: f64) {
|
||||
let payload = serde_json::to_value(task).expect("task serializes");
|
||||
let run_after = now_f64() + delay_seconds;
|
||||
if let Err(e) = TASK_QUEUE.enqueue(payload, run_after).await {
|
||||
log::error!("failed to enqueue retry: {e}");
|
||||
}
|
||||
}
|
||||
|
||||
async fn url_media(bot: Bot, message: &Message, url: &str) {
|
||||
let chat_id = message.chat.id.0;
|
||||
if let Err(e) = bot.send_chat_action(ChatId(chat_id), ChatAction::Typing).await {
|
||||
log::error!("send_chat_action failed: {e}");
|
||||
}
|
||||
log::info!("fetching {url}");
|
||||
match x_media::site::fetch(url).await {
|
||||
// Unsupported links are ignored silently (Python parity).
|
||||
Ok(None) => {
|
||||
log::info!("no site pattern matches {url}; ignoring");
|
||||
}
|
||||
// Retries exhausted: notify the user (Rust-only requirement 3).
|
||||
Err(e) => {
|
||||
log::error!("fetch {url}: {e}");
|
||||
let _ = reply(bot, message.clone(), "Failed to fetch media from this link.").await;
|
||||
}
|
||||
Ok(Some(fetched)) => {
|
||||
if fetched.media.is_empty() {
|
||||
let _ = reply(
|
||||
bot,
|
||||
message.clone(),
|
||||
"No media found or media type is not supported.",
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
let chat_data = CHAT_STORE.get(chat_id).await;
|
||||
// Per-site caption format override (empty -> built-in caption).
|
||||
let format = chat_data
|
||||
.message_format
|
||||
.get(fetched.site_name())
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
let caption = fetched.caption_with(&format);
|
||||
let task = if fetched.media.len() == 1
|
||||
&& matches!(fetched.media[0], Media::Animated { .. })
|
||||
{
|
||||
Task::SendAnimation {
|
||||
chat_id,
|
||||
reply_to_message_id: message.id.0 as i64,
|
||||
caption: caption.clone(),
|
||||
animation: MediaItemPayload::Animation {
|
||||
media: fetched.media[0].url().to_string(),
|
||||
has_spoiler: fetched.sensitive,
|
||||
},
|
||||
source_url: fetched.source_url.clone(),
|
||||
edit_before_forward: chat_data.edit_before_forward,
|
||||
forward_channel_id: chat_data.forward_channel_id,
|
||||
notify_chat_id: Some(chat_id),
|
||||
notify_message_id: Some(message.id.0 as i64),
|
||||
}
|
||||
} else {
|
||||
let items: Vec<MediaItemPayload> = fetched
|
||||
.media
|
||||
.iter()
|
||||
.map(|media| media_to_payload(media, fetched.sensitive))
|
||||
.collect();
|
||||
Task::SendMediaSequence {
|
||||
chat_id,
|
||||
reply_to_message_id: message.id.0 as i64,
|
||||
caption: caption.clone(),
|
||||
media_batches: send::chunk_media_items(items),
|
||||
batch_index: 0,
|
||||
sent_message_ids: vec![],
|
||||
source_url: fetched.source_url.clone(),
|
||||
edit_before_forward: chat_data.edit_before_forward,
|
||||
forward_channel_id: chat_data.forward_channel_id,
|
||||
notify_chat_id: Some(chat_id),
|
||||
notify_message_id: Some(message.id.0 as i64),
|
||||
}
|
||||
};
|
||||
let result = match &task {
|
||||
Task::SendAnimation { .. } => send::send_animation(&bot, &task).await,
|
||||
Task::SendMediaSequence { .. } => send::send_media_sequence(&bot, &task).await,
|
||||
Task::ForwardMessages { .. } => unreachable!(),
|
||||
};
|
||||
match result {
|
||||
Ok(message_ids) => {
|
||||
log::info!("sent {} message(s) for {url}", message_ids.len());
|
||||
send::post_send_actions(&bot, &task, message_ids).await;
|
||||
}
|
||||
Err(send::SendError::Retryable { delay_seconds, task }) => {
|
||||
log::info!("send for {url} failed, queued for retry in {delay_seconds:.1}s");
|
||||
enqueue_retry(task, delay_seconds).await;
|
||||
let _ = reply(bot, message.clone(), "Send failed. Task queued for retry.").await;
|
||||
}
|
||||
Err(send::SendError::Permanent {
|
||||
message: err_message,
|
||||
..
|
||||
}) => {
|
||||
log::error!("send for {url} failed permanently: {err_message}");
|
||||
let _ = reply(bot, message.clone(), format!("Send failed: {err_message}")).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn message_handler(bot: Bot, message: Message) -> Result<(), RequestError> {
|
||||
let is_private = matches!(message.chat.kind, ChatKind::Private(_));
|
||||
let sender = message
|
||||
.from
|
||||
.as_ref()
|
||||
.map(|from| from.full_name())
|
||||
.unwrap_or_else(|| "unknown".to_string());
|
||||
let text_preview = message
|
||||
.text()
|
||||
.map(|t| if t.len() > 120 { &t[..120] } else { t })
|
||||
.unwrap_or("<no text>");
|
||||
log::info!("message from {sender} in {} (private={is_private}): {text_preview}", message.chat.id);
|
||||
// URL/edit flows only run in private chats; commands run in any chat.
|
||||
if is_private && edit_message_handler(&bot, &message).await {
|
||||
return respond(());
|
||||
}
|
||||
if let Some(text) = message.text()
|
||||
&& let Ok(command) = Command::parse(text, "")
|
||||
{
|
||||
log::info!("command from {}: {text_preview}", message.chat.id);
|
||||
execute_command(&bot, &message, command).await?;
|
||||
return respond(());
|
||||
}
|
||||
if is_private {
|
||||
let urls = extract_urls(&message);
|
||||
if !urls.is_empty() {
|
||||
log::info!("extracted {} URL(s): {urls:?}", urls.len());
|
||||
}
|
||||
for url in urls {
|
||||
url_media(bot.clone(), &message, &url).await;
|
||||
}
|
||||
}
|
||||
respond(())
|
||||
}
|
||||
|
||||
pub async fn inline_query_handler(bot: Bot, query: InlineQuery) -> Result<(), RequestError> {
|
||||
if query.query.is_empty() {
|
||||
return respond(());
|
||||
}
|
||||
log::info!("inline query: {}", query.query);
|
||||
match x_media::site::fetch(&query.query).await {
|
||||
Ok(Some(fetched)) => {
|
||||
let mut results: Vec<InlineQueryResult> = Vec::new();
|
||||
for (i, media) in fetched.media.iter().enumerate() {
|
||||
let id = format!("{i}");
|
||||
let Some(url) = url::Url::parse(media.url()).ok() else {
|
||||
continue;
|
||||
};
|
||||
let thumbnail = media
|
||||
.thumbnail_url()
|
||||
.and_then(|t| url::Url::parse(t).ok())
|
||||
.unwrap_or_else(|| url.clone());
|
||||
let caption = fetched.caption.clone();
|
||||
let result = match media {
|
||||
Media::Illustration { .. } => InlineQueryResult::Photo(
|
||||
InlineQueryResultPhoto::new(id, url, thumbnail)
|
||||
.caption(caption)
|
||||
.parse_mode(ParseMode::Html),
|
||||
),
|
||||
Media::Video { .. } => InlineQueryResult::Video(
|
||||
InlineQueryResultVideo::new(
|
||||
id,
|
||||
url,
|
||||
"video/mp4".parse().expect("valid mime"),
|
||||
thumbnail,
|
||||
fetched.title.clone(),
|
||||
)
|
||||
.caption(caption)
|
||||
.parse_mode(ParseMode::Html),
|
||||
),
|
||||
Media::Animated { .. } => InlineQueryResult::Mpeg4Gif(
|
||||
InlineQueryResultMpeg4Gif::new(id, url, thumbnail)
|
||||
.caption(caption)
|
||||
.parse_mode(ParseMode::Html),
|
||||
),
|
||||
};
|
||||
results.push(result);
|
||||
}
|
||||
if !results.is_empty() {
|
||||
bot.answer_inline_query(query.id, results).await?;
|
||||
}
|
||||
}
|
||||
Ok(None) => {}
|
||||
Err(e) => log::error!("inline fetch {}: {e}", query.query),
|
||||
}
|
||||
respond(())
|
||||
}
|
||||
|
||||
pub async fn callback_query_handler(bot: Bot, query: CallbackQuery) -> Result<(), RequestError> {
|
||||
let callback_query_id = query.id;
|
||||
let data = query.data.clone();
|
||||
let Some(message) = &query.message else {
|
||||
return respond(());
|
||||
};
|
||||
let chat_id = message.chat().id.0;
|
||||
let prompt_message_id = message.id().0 as i64;
|
||||
let ttl_secs = CONFIG.edit_message_ttl.as_secs() as i64;
|
||||
let mut chat_data = CHAT_STORE.get(chat_id).await;
|
||||
let edit = chat_data.edit_message.get(&prompt_message_id).cloned();
|
||||
let Some(edit) = edit else {
|
||||
log::info!("callback from {}: no edit record for prompt {prompt_message_id}", chat_id);
|
||||
bot.answer_callback_query(callback_query_id)
|
||||
.text("Expired")
|
||||
.await?;
|
||||
return respond(());
|
||||
};
|
||||
// Lazy expiry: a stale record (past the TTL, not yet swept) is dropped.
|
||||
if edit.created_at + ttl_secs <= unix_now() {
|
||||
chat_data.edit_message.remove(&prompt_message_id);
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
bot.answer_callback_query(callback_query_id)
|
||||
.text("Expired")
|
||||
.await?;
|
||||
return respond(());
|
||||
}
|
||||
|
||||
let Some(data) = data else {
|
||||
return respond(());
|
||||
};
|
||||
log::info!("callback from {} on prompt {prompt_message_id}: {data}", chat_id);
|
||||
if data == "forward" {
|
||||
match chat_data.forward_channel_id {
|
||||
Some(channel_id) => {
|
||||
let forward_task = Task::ForwardMessages {
|
||||
from_chat_id: edit.chat_id,
|
||||
to_chat_id: channel_id,
|
||||
message_ids: edit.forward_message_ids.clone(),
|
||||
notify_chat_id: Some(chat_id),
|
||||
notify_message_id: Some(prompt_message_id),
|
||||
};
|
||||
match send::forward_messages(&bot, &forward_task).await {
|
||||
Ok(()) => {
|
||||
log::info!(
|
||||
"forwarded {} message(s) to channel {channel_id}",
|
||||
edit.forward_message_ids.len()
|
||||
);
|
||||
bot.answer_callback_query(callback_query_id)
|
||||
.text("✅ Forwarded")
|
||||
.await?;
|
||||
let _ = bot
|
||||
.delete_message(ChatId(chat_id), MessageId(prompt_message_id as i32))
|
||||
.await;
|
||||
chat_data.edit_message.remove(&prompt_message_id);
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
}
|
||||
Err(send::SendError::Retryable { delay_seconds, task }) => {
|
||||
log::info!("forward queued for retry in {delay_seconds:.1}s");
|
||||
enqueue_retry(task, delay_seconds).await;
|
||||
bot.answer_callback_query(callback_query_id)
|
||||
.text("Forward queued for retry.")
|
||||
.await?;
|
||||
}
|
||||
Err(send::SendError::Permanent { message, .. }) => {
|
||||
log::error!("forward failed permanently: {message}");
|
||||
bot.answer_callback_query(callback_query_id)
|
||||
.text(format!("Forward failed: {message}"))
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
}
|
||||
None => {
|
||||
log::info!("forward callback without a forward channel set");
|
||||
bot.answer_callback_query(callback_query_id)
|
||||
.text("No forward channel set.")
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
return respond(());
|
||||
}
|
||||
if let Some(name) = data.strip_prefix("template|") {
|
||||
if let Some(template_html) = chat_data.template.get(name).cloned()
|
||||
&& let Some(first_forward_id) = edit.forward_message_ids.first().copied()
|
||||
{
|
||||
// Raw template including the [] placeholder (Python parity).
|
||||
let _ = bot
|
||||
.edit_message_caption(ChatId(chat_id), MessageId(first_forward_id as i32))
|
||||
.caption(template_html)
|
||||
.parse_mode(ParseMode::Html)
|
||||
.await;
|
||||
if let Some(entry) = chat_data.edit_message.get_mut(&prompt_message_id) {
|
||||
entry.template = name.to_string();
|
||||
}
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
log::info!("template '{name}' applied to prompt {prompt_message_id}");
|
||||
}
|
||||
bot.answer_callback_query(callback_query_id).await?;
|
||||
}
|
||||
respond(())
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
use dotenv::dotenv;
|
||||
use teloxide::dptree::endpoint;
|
||||
use teloxide::types::{ChatId, InputFile, MessageId};
|
||||
use teloxide::update_listeners::webhooks;
|
||||
use teloxide::prelude::*;
|
||||
use tokio::sync::watch;
|
||||
use x_media::site;
|
||||
|
||||
mod config;
|
||||
mod handlers;
|
||||
mod queue;
|
||||
mod send;
|
||||
mod state;
|
||||
|
||||
use handlers::{CHAT_STORE, CONFIG, TASK_QUEUE};
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() {
|
||||
dotenv().ok();
|
||||
pretty_env_logger::init();
|
||||
log::info!("Starting bot");
|
||||
|
||||
let bot = Bot::from_env();
|
||||
|
||||
log::info!(
|
||||
"config: {} admin(s), edit-message TTL {}s",
|
||||
CONFIG.admin_ids.len(),
|
||||
CONFIG.edit_message_ttl.as_secs()
|
||||
);
|
||||
|
||||
// Queue worker: handles typed tasks, dead-letters failed sends to the
|
||||
// task's chat.
|
||||
TASK_QUEUE
|
||||
.start(send::handle_task, send::dead_letter_notify)
|
||||
.await;
|
||||
log::info!("task queue worker started");
|
||||
|
||||
// Pixiv login validation (user request): a failed login notifies the
|
||||
// admin and disables pixiv for this process.
|
||||
if site::pixiv::enabled() {
|
||||
match site::pixiv::validate().await {
|
||||
Ok(()) => log::info!("pixiv login validated"),
|
||||
Err(e) => {
|
||||
log::error!("pixiv login failed: {e}");
|
||||
if let Some(admin) = CONFIG.admin_ids.first() {
|
||||
let _ = bot
|
||||
.send_message(ChatId(*admin), format!("Pixiv login failed: {e}"))
|
||||
.await;
|
||||
}
|
||||
site::pixiv::disable();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Edit-expiry sweep: clears the prompt's buttons once the record expires.
|
||||
log::info!("edit-expiry sweep: every 300s, ttl {}", CONFIG.edit_message_ttl.as_secs());
|
||||
let (stop_tx, stop_rx) = watch::channel(false);
|
||||
{
|
||||
let bot = bot.clone();
|
||||
let mut stop_rx = stop_rx;
|
||||
tokio::spawn(async move {
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = stop_rx.changed() => break,
|
||||
_ = tokio::time::sleep(std::time::Duration::from_secs(300)) => {}
|
||||
}
|
||||
let ttl = CONFIG.edit_message_ttl;
|
||||
let removed = CHAT_STORE.prune_expired(ttl).await;
|
||||
for (chat_id, prompt_message_id) in removed {
|
||||
// If the prompt was already deleted, this fails with a
|
||||
// 400 "message to edit not found" — log and ignore.
|
||||
if let Err(e) = bot
|
||||
.edit_message_reply_markup(ChatId(chat_id), MessageId(prompt_message_id as i32))
|
||||
.await
|
||||
{
|
||||
log::info!("edit-expiry sweep: prompt message gone: {e}");
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
let handler = dptree::entry()
|
||||
.branch(Update::filter_message().branch(endpoint(handlers::message_handler)))
|
||||
.branch(Update::filter_inline_query().branch(endpoint(handlers::inline_query_handler)))
|
||||
.branch(Update::filter_callback_query().branch(endpoint(handlers::callback_query_handler)));
|
||||
|
||||
let mut dispatcher = Dispatcher::builder(bot.clone(), handler)
|
||||
.dependencies(dptree::deps![""])
|
||||
.enable_ctrlc_handler()
|
||||
.build();
|
||||
|
||||
if CONFIG.webhook_enabled {
|
||||
log::info!("running in webhook mode");
|
||||
let url = CONFIG
|
||||
.webhook_url
|
||||
.clone()
|
||||
.expect("WEBHOOK_URL is not set");
|
||||
bot.set_webhook(url.clone()).await.unwrap();
|
||||
let listen = CONFIG.webhook_listen.expect("WEBHOOK_LISTEN is not set");
|
||||
let port = CONFIG.webhook_port.expect("WEBHOOK_PORT is not set");
|
||||
let mut options = webhooks::Options::new((listen, port).into(), url);
|
||||
if let Some(cert) = &CONFIG.webhook_cert {
|
||||
options = options.certificate(InputFile::file(cert));
|
||||
}
|
||||
if let Some(secret) = &CONFIG.webhook_secret_token {
|
||||
options = options.secret_token(secret.clone());
|
||||
}
|
||||
|
||||
dispatcher
|
||||
.dispatch_with_listener(
|
||||
webhooks::axum(bot.clone(), options)
|
||||
.await
|
||||
.expect("Failed to create webhook listener"),
|
||||
LoggingErrorHandler::with_custom_text("Error from update listener"),
|
||||
)
|
||||
.await;
|
||||
} else {
|
||||
log::info!("running in polling mode");
|
||||
dispatcher.dispatch().await;
|
||||
}
|
||||
|
||||
// Graceful stop (Ctrl+C): stop the sweep, notify the admin, drain the queue.
|
||||
log::info!("Stopping bot");
|
||||
let _ = stop_tx.send(true);
|
||||
if let Some(admin) = CONFIG.admin_ids.first() {
|
||||
let _ = bot.send_message(ChatId(*admin), "Shutting down...").await;
|
||||
}
|
||||
TASK_QUEUE.stop().await;
|
||||
log::info!("Bot stopped");
|
||||
}
|
||||
@@ -0,0 +1,488 @@
|
||||
//! Generic persistent task queue backed by SQLite (table `tasks`).
|
||||
//!
|
||||
//! Concepts kept from the Python `utils/task_queue.py` (untrusted, redesigned):
|
||||
//! the table schema, the lease/lock/recovery model, and the retry→dead-letter
|
||||
//! flow. The Python dict-mutation hack (attempts inside the payload) is
|
||||
//! replaced by dedicated columns.
|
||||
|
||||
use parking_lot::Mutex;
|
||||
use rusqlite::{params, Connection};
|
||||
use serde_json::Value;
|
||||
use std::pin::Pin;
|
||||
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
use tokio::sync::Notify;
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
pub const MAX_RETRIES: u32 = 2;
|
||||
pub const LOCK_TTL_SECONDS: f64 = 120.0;
|
||||
|
||||
/// What a handler returns instead of throwing. The payload it carries is the
|
||||
/// (possibly updated) task state to persist for the next attempt.
|
||||
pub enum QueueError {
|
||||
/// Reschedule with the given delay; after `MAX_RETRIES` attempts the task
|
||||
/// is dead-lettered instead.
|
||||
Retryable {
|
||||
delay_seconds: f64,
|
||||
payload: Value,
|
||||
},
|
||||
/// Give up now.
|
||||
Permanent {
|
||||
message: String,
|
||||
payload: Value,
|
||||
},
|
||||
}
|
||||
|
||||
type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
|
||||
type Handler = dyn Fn(Value) -> BoxFuture<'static, Result<(), QueueError>> + Send + Sync;
|
||||
type DeadLetter = dyn Fn(Value, String) -> BoxFuture<'static, ()> + Send + Sync;
|
||||
|
||||
pub struct PersistentTaskQueue {
|
||||
db_path: String,
|
||||
notify: Arc<Notify>,
|
||||
stop: Arc<AtomicBool>,
|
||||
worker: Mutex<Option<JoinHandle<()>>>,
|
||||
counter: AtomicU64,
|
||||
}
|
||||
|
||||
struct LeasedRow {
|
||||
id: String,
|
||||
payload: String,
|
||||
attempts: i32,
|
||||
}
|
||||
|
||||
/// Owned worker state so the spawned loop does not borrow the queue handle.
|
||||
struct QueueWorker {
|
||||
db_path: String,
|
||||
notify: Arc<Notify>,
|
||||
stop: Arc<AtomicBool>,
|
||||
handler: Arc<Handler>,
|
||||
dead_letter: Arc<DeadLetter>,
|
||||
}
|
||||
|
||||
fn now_f64() -> f64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.as_secs_f64())
|
||||
.unwrap_or(0.0)
|
||||
}
|
||||
|
||||
fn ensure_schema(conn: &Connection) -> rusqlite::Result<()> {
|
||||
conn.execute_batch(
|
||||
"CREATE TABLE IF NOT EXISTS tasks (id TEXT PRIMARY KEY, payload TEXT NOT NULL, \
|
||||
run_after REAL NOT NULL, attempts INTEGER NOT NULL, status TEXT NOT NULL, \
|
||||
locked_until REAL NOT NULL, created_at REAL NOT NULL);",
|
||||
)
|
||||
}
|
||||
|
||||
impl PersistentTaskQueue {
|
||||
pub fn new(db_path: &str) -> Self {
|
||||
// Ensure the parent dir and table exist even if only the queue (not
|
||||
// ChatStore) is used — a fresh container without a mounted data dir
|
||||
// must still be able to open the DB.
|
||||
if let Some(parent) = std::path::Path::new(db_path).parent()
|
||||
&& !parent.as_os_str().is_empty()
|
||||
&& let Err(e) = std::fs::create_dir_all(parent)
|
||||
{
|
||||
log::error!("failed to create queue dir: {e}");
|
||||
}
|
||||
if let Ok(conn) = Connection::open(db_path)
|
||||
&& let Err(e) = ensure_schema(&conn)
|
||||
{
|
||||
log::error!("failed to initialize queue schema: {e}");
|
||||
}
|
||||
Self {
|
||||
db_path: db_path.to_string(),
|
||||
notify: Arc::new(Notify::new()),
|
||||
stop: Arc::new(AtomicBool::new(false)),
|
||||
worker: Mutex::new(None),
|
||||
counter: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// Starts the worker loop. Also recovers rows left `in_progress` by a
|
||||
/// previous process (lease expired).
|
||||
pub async fn start<H, F, D, G>(&self, handler: H, dead_letter: D)
|
||||
where
|
||||
H: Fn(Value) -> F + Send + Sync + 'static,
|
||||
F: Future<Output = Result<(), QueueError>> + Send + 'static,
|
||||
D: Fn(Value, String) -> G + Send + Sync + 'static,
|
||||
G: Future<Output = ()> + Send + 'static,
|
||||
{
|
||||
let handler: Arc<Handler> = Arc::new(move |payload| Box::pin(handler(payload)));
|
||||
let dead_letter: Arc<DeadLetter> =
|
||||
Arc::new(move |payload, message| Box::pin(dead_letter(payload, message)));
|
||||
self.recover_stale().await;
|
||||
let worker = QueueWorker {
|
||||
db_path: self.db_path.clone(),
|
||||
notify: Arc::clone(&self.notify),
|
||||
stop: Arc::clone(&self.stop),
|
||||
handler,
|
||||
dead_letter,
|
||||
};
|
||||
let worker = tokio::spawn(worker.run_loop());
|
||||
*self.worker.lock() = Some(worker);
|
||||
}
|
||||
|
||||
pub async fn stop(&self) {
|
||||
self.stop.store(true, Ordering::Relaxed);
|
||||
self.notify.notify_one();
|
||||
if let Some(handle) = self.worker.lock().take() {
|
||||
let _ = handle.await;
|
||||
}
|
||||
}
|
||||
|
||||
/// Persists a task. `run_after` is an absolute unix timestamp (seconds).
|
||||
/// Notifies the worker only after the insert has committed, so the worker
|
||||
/// never wakes to an invisible row.
|
||||
pub async fn enqueue(&self, payload: Value, run_after: f64) -> rusqlite::Result<()> {
|
||||
let id = format!(
|
||||
"task_{}_{}",
|
||||
(now_f64() * 1000.0) as u64,
|
||||
self.counter.fetch_add(1, Ordering::Relaxed)
|
||||
);
|
||||
let payload = payload.to_string();
|
||||
let db_path = self.db_path.clone();
|
||||
log::info!("enqueued {id} (run_after {run_after:.1})");
|
||||
let result = tokio::task::spawn_blocking(move || -> rusqlite::Result<()> {
|
||||
let conn = Connection::open(&db_path)?;
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO tasks (id, payload, run_after, attempts, status, locked_until, created_at) \
|
||||
VALUES (?1, ?2, ?3, 0, 'pending', 0, ?4)",
|
||||
params![id, payload, run_after, now_f64()],
|
||||
)?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
.expect("queue insert worker panicked")?;
|
||||
self.notify.notify_one();
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
async fn recover_stale(&self) {
|
||||
let db_path = self.db_path.clone();
|
||||
tokio::task::spawn_blocking(move || -> rusqlite::Result<()> {
|
||||
let conn = Connection::open(&db_path)?;
|
||||
conn.execute(
|
||||
"UPDATE tasks SET status='pending', locked_until=0 WHERE status='in_progress' AND locked_until < ?1",
|
||||
params![now_f64()],
|
||||
)?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
.expect("queue recovery worker panicked")
|
||||
.unwrap_or_else(|e| log::error!("queue recovery failed: {e}"));
|
||||
}
|
||||
}
|
||||
|
||||
impl QueueWorker {
|
||||
async fn run_loop(self) {
|
||||
while !self.stop.load(Ordering::Relaxed) {
|
||||
match self.lease_next().await {
|
||||
Some(row) => self.process(row).await,
|
||||
None => {
|
||||
let wait_until = self.earliest_run_after().await;
|
||||
let notified = self.notify.notified();
|
||||
tokio::pin!(notified);
|
||||
match wait_until {
|
||||
Some(until) => {
|
||||
let delay = (until - now_f64()).max(0.0);
|
||||
tokio::select! {
|
||||
_ = &mut notified => {}
|
||||
_ = tokio::time::sleep(Duration::from_secs_f64(delay)) => {}
|
||||
}
|
||||
}
|
||||
None => {
|
||||
notified.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Leases the oldest due row (sets it `in_progress` with a lock TTL).
|
||||
async fn lease_next(&self) -> Option<LeasedRow> {
|
||||
let db_path = self.db_path.clone();
|
||||
tokio::task::spawn_blocking(move || -> rusqlite::Result<Option<LeasedRow>> {
|
||||
let mut conn = Connection::open(&db_path)?;
|
||||
let tx = conn.transaction()?;
|
||||
let now = now_f64();
|
||||
let row = tx.query_row(
|
||||
"SELECT id, payload, attempts FROM tasks WHERE status='pending' AND run_after <= ?1 \
|
||||
ORDER BY run_after LIMIT 1",
|
||||
params![now],
|
||||
|r| {
|
||||
Ok((
|
||||
r.get::<_, String>(0)?,
|
||||
r.get::<_, String>(1)?,
|
||||
r.get::<_, i32>(2)?,
|
||||
))
|
||||
},
|
||||
);
|
||||
let (id, payload, attempts) = match row {
|
||||
Ok(row) => row,
|
||||
Err(rusqlite::Error::QueryReturnedNoRows) => {
|
||||
tx.commit()?;
|
||||
return Ok(None);
|
||||
}
|
||||
Err(e) => return Err(e),
|
||||
};
|
||||
tx.execute(
|
||||
"UPDATE tasks SET status='in_progress', locked_until=?1 WHERE id=?2",
|
||||
params![now + LOCK_TTL_SECONDS, id],
|
||||
)?;
|
||||
tx.commit()?;
|
||||
Ok(Some(LeasedRow {
|
||||
id,
|
||||
payload,
|
||||
attempts,
|
||||
}))
|
||||
})
|
||||
.await
|
||||
.expect("queue lease worker panicked")
|
||||
.unwrap_or_else(|e| {
|
||||
log::error!("queue lease failed: {e}");
|
||||
None
|
||||
})
|
||||
}
|
||||
|
||||
async fn earliest_run_after(&self) -> Option<f64> {
|
||||
let db_path = self.db_path.clone();
|
||||
tokio::task::spawn_blocking(move || -> rusqlite::Result<Option<f64>> {
|
||||
let conn = Connection::open(&db_path)?;
|
||||
let mut stmt = conn.prepare("SELECT MIN(run_after) FROM tasks WHERE status='pending'")?;
|
||||
let mut rows = stmt.query([])?;
|
||||
match rows.next()? {
|
||||
Some(row) => Ok(row.get::<_, Option<f64>>(0)?),
|
||||
None => Ok(None),
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("queue timing worker panicked")
|
||||
.unwrap_or_else(|e| {
|
||||
log::error!("queue timing query failed: {e}");
|
||||
None
|
||||
})
|
||||
}
|
||||
|
||||
async fn process(&self, row: LeasedRow) {
|
||||
let payload: Value = match serde_json::from_str(&row.payload) {
|
||||
Ok(value) => value,
|
||||
Err(e) => {
|
||||
log::error!("queue: unparseable payload for {}: {e}", row.id);
|
||||
self.delete_row(&row.id).await;
|
||||
(self.dead_letter)(Value::Null, format!("invalid stored payload: {e}")).await;
|
||||
return;
|
||||
}
|
||||
};
|
||||
log::info!("processing {} (attempt {})", row.id, row.attempts + 1);
|
||||
match (self.handler)(payload).await {
|
||||
Ok(()) => {
|
||||
log::info!("task {} completed", row.id);
|
||||
self.delete_row(&row.id).await;
|
||||
}
|
||||
Err(QueueError::Retryable {
|
||||
delay_seconds,
|
||||
payload,
|
||||
}) => {
|
||||
if row.attempts as u32 >= MAX_RETRIES {
|
||||
let message = format!("task failed after {MAX_RETRIES} retries");
|
||||
log::error!("dead-lettering {}: {message}", row.id);
|
||||
self.delete_row(&row.id).await;
|
||||
(self.dead_letter)(payload, message).await;
|
||||
} else {
|
||||
log::info!(
|
||||
"task {} rescheduled in {delay_seconds:.1}s (attempt {})",
|
||||
row.id,
|
||||
row.attempts + 1
|
||||
);
|
||||
self.reschedule(&row.id, payload, delay_seconds, row.attempts + 1)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
Err(QueueError::Permanent { message, payload }) => {
|
||||
log::error!("dead-lettering {}: {message}", row.id);
|
||||
self.delete_row(&row.id).await;
|
||||
(self.dead_letter)(payload, message).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn delete_row(&self, id: &str) {
|
||||
let db_path = self.db_path.clone();
|
||||
let id = id.to_string();
|
||||
tokio::task::spawn_blocking(move || -> rusqlite::Result<()> {
|
||||
let conn = Connection::open(&db_path)?;
|
||||
conn.execute("DELETE FROM tasks WHERE id = ?1", params![id])?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
.expect("queue delete worker panicked")
|
||||
.unwrap_or_else(|e| log::error!("queue delete failed: {e}"));
|
||||
}
|
||||
|
||||
async fn reschedule(&self, id: &str, payload: Value, delay_seconds: f64, attempts: i32) {
|
||||
let db_path = self.db_path.clone();
|
||||
let id = id.to_string();
|
||||
let payload = payload.to_string();
|
||||
tokio::task::spawn_blocking(move || -> rusqlite::Result<()> {
|
||||
let conn = Connection::open(&db_path)?;
|
||||
conn.execute(
|
||||
"UPDATE tasks SET payload=?1, run_after=?2, attempts=?3, status='pending', locked_until=0 WHERE id=?4",
|
||||
params![payload, now_f64() + delay_seconds, attempts, id],
|
||||
)?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
.expect("queue reschedule worker panicked")
|
||||
.unwrap_or_else(|e| log::error!("queue reschedule failed: {e}"));
|
||||
self.notify.notify_one();
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
|
||||
|
||||
async fn new_queue() -> (PersistentTaskQueue, tempfile::TempDir) {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("queue.db");
|
||||
let queue = PersistentTaskQueue::new(path.to_str().unwrap());
|
||||
(queue, dir)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enqueue_runs_handler_once() {
|
||||
let (queue, _dir) = new_queue().await;
|
||||
let calls = Arc::new(AtomicUsize::new(0));
|
||||
let calls_worker = calls.clone();
|
||||
queue
|
||||
.start(
|
||||
move |payload| {
|
||||
assert_eq!(payload["n"], 42);
|
||||
calls_worker.fetch_add(1, AtomicOrdering::SeqCst);
|
||||
async { Ok(()) }
|
||||
},
|
||||
|_payload, _message| async {},
|
||||
)
|
||||
.await;
|
||||
queue
|
||||
.enqueue(serde_json::json!({"n": 42}), now_f64())
|
||||
.await
|
||||
.unwrap();
|
||||
tokio::time::sleep(Duration::from_millis(300)).await;
|
||||
assert_eq!(calls.load(AtomicOrdering::SeqCst), 1);
|
||||
queue.stop().await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retryable_reschedules_then_dead_letters() {
|
||||
let (queue, _dir) = new_queue().await;
|
||||
let calls = Arc::new(AtomicUsize::new(0));
|
||||
let dead_calls = Arc::new(AtomicUsize::new(0));
|
||||
let c = calls.clone();
|
||||
let d = dead_calls.clone();
|
||||
queue
|
||||
.start(
|
||||
move |payload| {
|
||||
c.fetch_add(1, AtomicOrdering::SeqCst);
|
||||
let payload = payload.clone();
|
||||
async move {
|
||||
Err(QueueError::Retryable {
|
||||
delay_seconds: 0.001,
|
||||
payload,
|
||||
})
|
||||
}
|
||||
},
|
||||
move |_payload, _message| {
|
||||
d.fetch_add(1, AtomicOrdering::SeqCst);
|
||||
async {}
|
||||
},
|
||||
)
|
||||
.await;
|
||||
queue
|
||||
.enqueue(serde_json::json!({"a": 1}), now_f64())
|
||||
.await
|
||||
.unwrap();
|
||||
tokio::time::sleep(Duration::from_millis(600)).await;
|
||||
assert_eq!(
|
||||
calls.load(AtomicOrdering::SeqCst),
|
||||
MAX_RETRIES as usize + 1,
|
||||
"handler should run once per attempt"
|
||||
);
|
||||
assert_eq!(dead_calls.load(AtomicOrdering::SeqCst), 1);
|
||||
queue.stop().await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn permanent_error_dead_letters_immediately() {
|
||||
let (queue, _dir) = new_queue().await;
|
||||
let calls = Arc::new(AtomicUsize::new(0));
|
||||
let dead_calls = Arc::new(AtomicUsize::new(0));
|
||||
let c = calls.clone();
|
||||
let d = dead_calls.clone();
|
||||
queue
|
||||
.start(
|
||||
move |payload| {
|
||||
c.fetch_add(1, AtomicOrdering::SeqCst);
|
||||
let payload = payload.clone();
|
||||
async move {
|
||||
Err(QueueError::Permanent {
|
||||
message: "nope".into(),
|
||||
payload,
|
||||
})
|
||||
}
|
||||
},
|
||||
move |_payload, message| {
|
||||
assert_eq!(message, "nope");
|
||||
d.fetch_add(1, AtomicOrdering::SeqCst);
|
||||
async {}
|
||||
},
|
||||
)
|
||||
.await;
|
||||
queue
|
||||
.enqueue(serde_json::json!({"a": 1}), now_f64())
|
||||
.await
|
||||
.unwrap();
|
||||
tokio::time::sleep(Duration::from_millis(300)).await;
|
||||
assert_eq!(calls.load(AtomicOrdering::SeqCst), 1);
|
||||
assert_eq!(dead_calls.load(AtomicOrdering::SeqCst), 1);
|
||||
queue.stop().await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stale_in_progress_row_is_recovered_on_start() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let path = dir.path().join("queue.db");
|
||||
// Insert a stale leased row directly (lease expired).
|
||||
{
|
||||
let conn = Connection::open(&path).unwrap();
|
||||
ensure_schema(&conn).unwrap();
|
||||
conn.execute(
|
||||
"INSERT INTO tasks (id, payload, run_after, attempts, status, locked_until, created_at) \
|
||||
VALUES ('task_stale', '{\"s\":1}', 0, 0, 'in_progress', ?1, 0)",
|
||||
params![now_f64() - 10.0],
|
||||
)
|
||||
.unwrap();
|
||||
}
|
||||
let queue = PersistentTaskQueue::new(path.to_str().unwrap());
|
||||
let calls = Arc::new(AtomicUsize::new(0));
|
||||
let c = calls.clone();
|
||||
queue
|
||||
.start(
|
||||
move |payload| {
|
||||
assert_eq!(payload["s"], 1);
|
||||
c.fetch_add(1, AtomicOrdering::SeqCst);
|
||||
async { Ok(()) }
|
||||
},
|
||||
|_payload, _message| async {},
|
||||
)
|
||||
.await;
|
||||
tokio::time::sleep(Duration::from_millis(300)).await;
|
||||
assert_eq!(calls.load(AtomicOrdering::SeqCst), 1);
|
||||
queue.stop().await;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,917 @@
|
||||
//! Typed task payloads and send/forward executors with retry classification
|
||||
//! and the download-and-reupload fallback (Telegram's own fetch of a media
|
||||
//! URL is blocked by hotlink protection; the bot downloads the file itself
|
||||
//! and uploads it via multipart).
|
||||
|
||||
use crate::handlers::{CHAT_STORE, TASK_QUEUE};
|
||||
use crate::queue::QueueError;
|
||||
use crate::state::{EditMessage, unix_now};
|
||||
use rand::Rng;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use tempfile::NamedTempFile;
|
||||
use teloxide::prelude::*;
|
||||
use teloxide::types::{
|
||||
ChatId, InlineKeyboardButton, InlineKeyboardMarkup, InputFile, InputMedia,
|
||||
InputMediaAnimation, InputMediaPhoto, InputMediaVideo, Message, MessageId, ParseMode,
|
||||
ReplyParameters,
|
||||
};
|
||||
use teloxide::{ApiError, RequestError};
|
||||
use x_media::site::FetchError;
|
||||
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum MediaItemPayload {
|
||||
Photo {
|
||||
media: String,
|
||||
has_spoiler: bool,
|
||||
},
|
||||
Video {
|
||||
media: String,
|
||||
has_spoiler: bool,
|
||||
thumbnail: Option<String>,
|
||||
},
|
||||
Animation {
|
||||
media: String,
|
||||
has_spoiler: bool,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize, Clone, Debug)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum Task {
|
||||
SendMediaSequence {
|
||||
chat_id: i64,
|
||||
reply_to_message_id: i64,
|
||||
caption: String,
|
||||
media_batches: Vec<Vec<MediaItemPayload>>,
|
||||
batch_index: usize,
|
||||
sent_message_ids: Vec<i64>,
|
||||
source_url: String,
|
||||
edit_before_forward: bool,
|
||||
forward_channel_id: Option<i64>,
|
||||
notify_chat_id: Option<i64>,
|
||||
notify_message_id: Option<i64>,
|
||||
},
|
||||
SendAnimation {
|
||||
chat_id: i64,
|
||||
reply_to_message_id: i64,
|
||||
caption: String,
|
||||
animation: MediaItemPayload,
|
||||
source_url: String,
|
||||
edit_before_forward: bool,
|
||||
forward_channel_id: Option<i64>,
|
||||
notify_chat_id: Option<i64>,
|
||||
notify_message_id: Option<i64>,
|
||||
},
|
||||
ForwardMessages {
|
||||
from_chat_id: i64,
|
||||
to_chat_id: i64,
|
||||
message_ids: Vec<i64>,
|
||||
notify_chat_id: Option<i64>,
|
||||
notify_message_id: Option<i64>,
|
||||
},
|
||||
}
|
||||
|
||||
pub const MAX_MEDIA_GROUP: usize = 9;
|
||||
pub const MAX_UPLOAD_BYTES: u64 = 50 * 1024 * 1024; // Telegram Bot API upload cap
|
||||
|
||||
/// Splits media into batches of at most [`MAX_MEDIA_GROUP`] items.
|
||||
pub fn chunk_media_items<T: Clone>(items: Vec<T>) -> Vec<Vec<T>> {
|
||||
items.chunks(MAX_MEDIA_GROUP).map(|chunk| chunk.to_vec()).collect()
|
||||
}
|
||||
|
||||
/// Exponential backoff with jitter, capped at 30s.
|
||||
pub fn retry_delay_seconds(attempts: u32) -> f64 {
|
||||
let jitter: f64 = rand::thread_rng().gen_range(0.2..0.8);
|
||||
(2f64.powi(attempts as i32) + jitter).min(30.0)
|
||||
}
|
||||
|
||||
/// Telegram's servers failed to fetch a media URL (hotlink protection etc.):
|
||||
/// these errors are handled by the download-and-reupload fallback, NOT by a
|
||||
/// queue retry (resending the URL cannot succeed).
|
||||
pub fn is_media_fetch_failure(e: &ApiError) -> bool {
|
||||
const MARKERS: [&str; 5] = [
|
||||
"webpage_media_empty",
|
||||
"media_empty",
|
||||
"empty_web_media",
|
||||
"webpage_curl_failed",
|
||||
"timeout",
|
||||
];
|
||||
let description = e.to_string().to_lowercase();
|
||||
MARKERS.iter().any(|marker| description.contains(marker))
|
||||
}
|
||||
|
||||
/// Task-free classification of a Telegram request error. The callers attach
|
||||
/// the (updated) task when building a [`SendError`].
|
||||
pub enum Classification {
|
||||
Retryable { delay_seconds: f64 },
|
||||
Permanent { message: String },
|
||||
/// Handled by the download fallback, not a queue retry.
|
||||
MediaFetchFailure,
|
||||
}
|
||||
|
||||
pub fn classify_request_error(e: &RequestError) -> Classification {
|
||||
match e {
|
||||
RequestError::RetryAfter(seconds) => {
|
||||
Classification::Retryable { delay_seconds: seconds.seconds() as f64 }
|
||||
}
|
||||
RequestError::Network(_) => Classification::Retryable {
|
||||
delay_seconds: retry_delay_seconds(0),
|
||||
},
|
||||
RequestError::Api(api) if is_media_fetch_failure(api) => Classification::MediaFetchFailure,
|
||||
RequestError::Api(api) => Classification::Permanent { message: api.to_string() },
|
||||
RequestError::MigrateToChatId(_)
|
||||
| RequestError::InvalidJson { .. }
|
||||
| RequestError::Io(_) => Classification::Permanent { message: e.to_string() },
|
||||
}
|
||||
}
|
||||
|
||||
pub enum SendError {
|
||||
Retryable { delay_seconds: f64, task: Task },
|
||||
Permanent { message: String, task: Task },
|
||||
}
|
||||
|
||||
fn classify_to_send_error(e: &RequestError, task: Task) -> SendError {
|
||||
match classify_request_error(e) {
|
||||
Classification::Retryable { delay_seconds } => SendError::Retryable {
|
||||
delay_seconds,
|
||||
task,
|
||||
},
|
||||
Classification::Permanent { message } => SendError::Permanent { message, task },
|
||||
Classification::MediaFetchFailure => SendError::Permanent {
|
||||
message: "media fetch failed".into(),
|
||||
task,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_media_url(s: &str) -> Result<url::Url, String> {
|
||||
url::Url::parse(s).map_err(|e| format!("invalid media URL: {e}"))
|
||||
}
|
||||
|
||||
fn item_url(item: &MediaItemPayload) -> &str {
|
||||
match item {
|
||||
MediaItemPayload::Photo { media, .. }
|
||||
| MediaItemPayload::Video { media, .. }
|
||||
| MediaItemPayload::Animation { media, .. } => media,
|
||||
}
|
||||
}
|
||||
|
||||
/// Remote http(s) URLs are handed to Telegram to fetch; everything else
|
||||
/// (e.g. a locally encoded ugoira MP4) is uploaded directly.
|
||||
fn input_file_for(media: &str) -> Result<InputFile, String> {
|
||||
if media.starts_with("http://") || media.starts_with("https://") {
|
||||
Ok(InputFile::url(parse_media_url(media)?))
|
||||
} else {
|
||||
Ok(InputFile::file(media))
|
||||
}
|
||||
}
|
||||
|
||||
fn photo_media(file: InputFile, caption: Option<&str>, spoiler: bool) -> InputMedia {
|
||||
let mut photo = InputMediaPhoto::new(file).parse_mode(ParseMode::Html);
|
||||
if let Some(caption) = caption {
|
||||
photo = photo.caption(caption);
|
||||
}
|
||||
if spoiler {
|
||||
photo = photo.spoiler();
|
||||
}
|
||||
InputMedia::Photo(photo)
|
||||
}
|
||||
|
||||
fn video_media(file: InputFile, caption: Option<&str>, spoiler: bool) -> InputMedia {
|
||||
let mut video = InputMediaVideo::new(file).parse_mode(ParseMode::Html);
|
||||
if let Some(caption) = caption {
|
||||
video = video.caption(caption);
|
||||
}
|
||||
if spoiler {
|
||||
video = video.spoiler();
|
||||
}
|
||||
InputMedia::Video(video)
|
||||
}
|
||||
|
||||
fn animation_media(file: InputFile, caption: Option<&str>, spoiler: bool) -> InputMedia {
|
||||
let mut animation = InputMediaAnimation::new(file).parse_mode(ParseMode::Html);
|
||||
if let Some(caption) = caption {
|
||||
animation = animation.caption(caption);
|
||||
}
|
||||
if spoiler {
|
||||
animation = animation.spoiler();
|
||||
}
|
||||
InputMedia::Animation(animation)
|
||||
}
|
||||
|
||||
/// Builds a media group from payloads; only the first item of the batch gets
|
||||
/// the caption (Telegram rejects captions on later items).
|
||||
fn build_media_group(
|
||||
batch: &[MediaItemPayload],
|
||||
caption: Option<&str>,
|
||||
) -> Result<Vec<InputMedia>, String> {
|
||||
batch
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, item)| {
|
||||
let item_caption = if i == 0 { caption } else { None };
|
||||
Ok(match item {
|
||||
MediaItemPayload::Photo {
|
||||
media,
|
||||
has_spoiler,
|
||||
} => photo_media(input_file_for(media)?, item_caption, *has_spoiler),
|
||||
MediaItemPayload::Video {
|
||||
media,
|
||||
has_spoiler,
|
||||
thumbnail,
|
||||
} => {
|
||||
let mut video = video_media(input_file_for(media)?, item_caption, *has_spoiler);
|
||||
if let (Some(thumb), InputMedia::Video(v)) = (thumbnail, &mut video) {
|
||||
*v = v.clone().thumbnail(input_file_for(thumb)?);
|
||||
}
|
||||
video
|
||||
}
|
||||
MediaItemPayload::Animation {
|
||||
media,
|
||||
has_spoiler,
|
||||
} => animation_media(input_file_for(media)?, item_caption, *has_spoiler),
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Infers a file extension from magic bytes so Telegram detects the mime type
|
||||
/// on multipart uploads.
|
||||
fn sniff_ext(bytes: &[u8]) -> &'static str {
|
||||
if bytes.starts_with(&[0xFF, 0xD8]) {
|
||||
"jpg"
|
||||
} else if bytes.starts_with(b"\x89PNG") {
|
||||
"png"
|
||||
} else if bytes.starts_with(b"RIFF") && bytes.len() >= 12 && &bytes[8..12] == b"WEBP" {
|
||||
"webp"
|
||||
} else if bytes.starts_with(b"GIF8") {
|
||||
"gif"
|
||||
} else if bytes.len() >= 12 && &bytes[4..8] == b"ftyp" {
|
||||
"mp4"
|
||||
} else {
|
||||
"bin"
|
||||
}
|
||||
}
|
||||
|
||||
enum FallbackError {
|
||||
Retryable { delay_seconds: f64 },
|
||||
Permanent { message: String },
|
||||
}
|
||||
|
||||
/// Downloads one media item to a temp file (deleted on drop). Network errors
|
||||
/// are retryable; size over the upload cap and other download errors are not.
|
||||
async fn download_to_temp(item: &MediaItemPayload) -> Result<NamedTempFile, FallbackError> {
|
||||
let media_url = match item {
|
||||
MediaItemPayload::Photo { media, .. }
|
||||
| MediaItemPayload::Video { media, .. }
|
||||
| MediaItemPayload::Animation { media, .. } => media,
|
||||
};
|
||||
let bytes = match x_media::site::download_media(media_url).await {
|
||||
Ok(bytes) => bytes,
|
||||
Err(FetchError::Http(_)) => {
|
||||
return Err(FallbackError::Retryable {
|
||||
delay_seconds: retry_delay_seconds(0),
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
return Err(FallbackError::Permanent {
|
||||
message: format!("download failed: {e}"),
|
||||
});
|
||||
}
|
||||
};
|
||||
if bytes.len() as u64 > MAX_UPLOAD_BYTES {
|
||||
return Err(FallbackError::Permanent {
|
||||
message: "media too large".into(),
|
||||
});
|
||||
}
|
||||
let ext = sniff_ext(&bytes);
|
||||
let mut file = tempfile::Builder::new()
|
||||
.suffix(&format!(".{ext}"))
|
||||
.tempfile()
|
||||
.map_err(|e| FallbackError::Permanent {
|
||||
message: format!("temp file failed: {e}"),
|
||||
})?;
|
||||
use std::io::Write;
|
||||
file.as_file_mut()
|
||||
.write_all(&bytes)
|
||||
.map_err(|e| FallbackError::Permanent {
|
||||
message: format!("temp file write failed: {e}"),
|
||||
})?;
|
||||
Ok(file)
|
||||
}
|
||||
|
||||
/// Download-and-reupload fallback for one media batch.
|
||||
async fn send_batch_via_upload(
|
||||
bot: &Bot,
|
||||
chat_id: i64,
|
||||
reply_to: i64,
|
||||
batch: &[MediaItemPayload],
|
||||
caption: Option<&str>,
|
||||
) -> Result<Vec<Message>, FallbackError> {
|
||||
let mut files = Vec::new();
|
||||
let mut items = Vec::new();
|
||||
for (i, item) in batch.iter().enumerate() {
|
||||
let file = download_to_temp(item).await?;
|
||||
let path = file.path().to_path_buf();
|
||||
let item_caption = if i == 0 { caption } else { None };
|
||||
let media = match item {
|
||||
MediaItemPayload::Photo { has_spoiler, .. } => {
|
||||
photo_media(InputFile::file(path), item_caption, *has_spoiler)
|
||||
}
|
||||
MediaItemPayload::Video { has_spoiler, .. } => {
|
||||
video_media(InputFile::file(path), item_caption, *has_spoiler)
|
||||
}
|
||||
MediaItemPayload::Animation { has_spoiler, .. } => {
|
||||
animation_media(InputFile::file(path), item_caption, *has_spoiler)
|
||||
}
|
||||
};
|
||||
items.push(media);
|
||||
files.push(file);
|
||||
}
|
||||
let result = bot
|
||||
.send_media_group(ChatId(chat_id), items)
|
||||
.reply_parameters(ReplyParameters::new(MessageId(reply_to as i32)).allow_sending_without_reply())
|
||||
.await;
|
||||
match result {
|
||||
Ok(messages) => Ok(messages),
|
||||
Err(e) => Err(match classify_request_error(&e) {
|
||||
Classification::Retryable { delay_seconds } => FallbackError::Retryable {
|
||||
delay_seconds,
|
||||
},
|
||||
Classification::Permanent { message } => FallbackError::Permanent { message },
|
||||
Classification::MediaFetchFailure => FallbackError::Permanent {
|
||||
message: "upload failed".into(),
|
||||
},
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn updated_sequence_task(task: &Task, batch_index: usize, sent_message_ids: Vec<i64>) -> Task {
|
||||
match task {
|
||||
Task::SendMediaSequence {
|
||||
chat_id,
|
||||
reply_to_message_id,
|
||||
caption,
|
||||
media_batches,
|
||||
source_url,
|
||||
edit_before_forward,
|
||||
forward_channel_id,
|
||||
notify_chat_id,
|
||||
notify_message_id,
|
||||
..
|
||||
} => Task::SendMediaSequence {
|
||||
chat_id: *chat_id,
|
||||
reply_to_message_id: *reply_to_message_id,
|
||||
caption: caption.clone(),
|
||||
media_batches: media_batches.clone(),
|
||||
batch_index,
|
||||
sent_message_ids,
|
||||
source_url: source_url.clone(),
|
||||
edit_before_forward: *edit_before_forward,
|
||||
forward_channel_id: *forward_channel_id,
|
||||
notify_chat_id: *notify_chat_id,
|
||||
notify_message_id: *notify_message_id,
|
||||
},
|
||||
_ => unreachable!("updated_sequence_task requires a SendMediaSequence task"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Sends the media batches starting at `task.batch_index`, extending
|
||||
/// `sent_message_ids`. Returns all sent message ids on full success; on
|
||||
/// failure returns a [`SendError`] whose task carries the resumed state.
|
||||
pub async fn send_media_sequence(bot: &Bot, task: &Task) -> Result<Vec<i64>, SendError> {
|
||||
let Task::SendMediaSequence {
|
||||
chat_id,
|
||||
reply_to_message_id,
|
||||
caption,
|
||||
media_batches,
|
||||
batch_index,
|
||||
sent_message_ids,
|
||||
..
|
||||
} = task
|
||||
else {
|
||||
unreachable!("send_media_sequence requires a SendMediaSequence task")
|
||||
};
|
||||
let chat_id = *chat_id;
|
||||
let reply_to = *reply_to_message_id;
|
||||
let mut sent = sent_message_ids.clone();
|
||||
for idx in *batch_index..media_batches.len() {
|
||||
let batch = &media_batches[idx];
|
||||
let caption = if idx == 0 { Some(caption.as_str()) } else { None };
|
||||
let items = match build_media_group(batch, caption) {
|
||||
Ok(items) => items,
|
||||
Err(message) => {
|
||||
return Err(SendError::Permanent {
|
||||
message,
|
||||
task: updated_sequence_task(task, idx, sent),
|
||||
});
|
||||
}
|
||||
};
|
||||
match bot
|
||||
.send_media_group(ChatId(chat_id), items)
|
||||
.reply_parameters(ReplyParameters::new(MessageId(reply_to as i32)).allow_sending_without_reply())
|
||||
.await
|
||||
{
|
||||
Ok(messages) => {
|
||||
log::info!(
|
||||
"media group batch {idx}/{} sent ({} item(s))",
|
||||
media_batches.len(),
|
||||
batch.len()
|
||||
);
|
||||
sent.extend(messages.into_iter().map(|m| m.id.0 as i64));
|
||||
}
|
||||
Err(RequestError::Api(api)) if is_media_fetch_failure(&api) => {
|
||||
log::info!(
|
||||
"Telegram could not fetch media for batch {idx} ({}), downloading and reuploading",
|
||||
batch
|
||||
.first()
|
||||
.map(|item| item_url(item))
|
||||
.unwrap_or("?")
|
||||
);
|
||||
match send_batch_via_upload(bot, chat_id, reply_to, batch, caption).await {
|
||||
Ok(messages) => sent.extend(messages.into_iter().map(|m| m.id.0 as i64)),
|
||||
Err(FallbackError::Retryable { delay_seconds }) => {
|
||||
return Err(SendError::Retryable {
|
||||
delay_seconds,
|
||||
task: updated_sequence_task(task, idx, sent),
|
||||
});
|
||||
}
|
||||
Err(FallbackError::Permanent { message }) => {
|
||||
return Err(SendError::Permanent {
|
||||
message,
|
||||
task: updated_sequence_task(task, idx, sent),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
return Err(classify_to_send_error(
|
||||
&e,
|
||||
updated_sequence_task(task, idx, sent),
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(sent)
|
||||
}
|
||||
|
||||
async fn send_animation_inner(
|
||||
bot: &Bot,
|
||||
chat_id: i64,
|
||||
reply_to: i64,
|
||||
caption: &str,
|
||||
spoiler: bool,
|
||||
file: InputFile,
|
||||
) -> Result<Message, RequestError> {
|
||||
let mut request = bot
|
||||
.send_animation(ChatId(chat_id), file)
|
||||
.caption(caption)
|
||||
.parse_mode(ParseMode::Html)
|
||||
.reply_parameters(ReplyParameters::new(MessageId(reply_to as i32)).allow_sending_without_reply());
|
||||
if spoiler {
|
||||
request = request.has_spoiler(true);
|
||||
}
|
||||
request.await
|
||||
}
|
||||
|
||||
/// Sends a lone animation (gif), URL first with the download fallback.
|
||||
pub async fn send_animation(bot: &Bot, task: &Task) -> Result<Vec<i64>, SendError> {
|
||||
let Task::SendAnimation {
|
||||
chat_id,
|
||||
reply_to_message_id,
|
||||
caption,
|
||||
animation,
|
||||
..
|
||||
} = task
|
||||
else {
|
||||
unreachable!("send_animation requires a SendAnimation task")
|
||||
};
|
||||
let chat_id = *chat_id;
|
||||
let reply_to = *reply_to_message_id;
|
||||
let (media_url, has_spoiler) = match animation {
|
||||
MediaItemPayload::Animation {
|
||||
media,
|
||||
has_spoiler,
|
||||
} => (media, *has_spoiler),
|
||||
MediaItemPayload::Photo { .. } | MediaItemPayload::Video { .. } => {
|
||||
unreachable!("SendAnimation carries an Animation payload")
|
||||
}
|
||||
};
|
||||
let url_file = match input_file_for(media_url) {
|
||||
Ok(file) => file,
|
||||
Err(message) => return Err(SendError::Permanent { message, task: task.clone() }),
|
||||
};
|
||||
match send_animation_inner(bot, chat_id, reply_to, caption, has_spoiler, url_file)
|
||||
.await
|
||||
{
|
||||
Ok(message) => Ok(vec![message.id.0 as i64]),
|
||||
Err(RequestError::Api(api)) if is_media_fetch_failure(&api) => {
|
||||
log::info!(
|
||||
"Telegram could not fetch animation URL, downloading and reuploading: {}",
|
||||
media_url
|
||||
);
|
||||
let file = match download_to_temp(animation).await {
|
||||
Ok(file) => file,
|
||||
Err(FallbackError::Retryable { delay_seconds }) => {
|
||||
return Err(SendError::Retryable { delay_seconds, task: task.clone() });
|
||||
}
|
||||
Err(FallbackError::Permanent { message }) => {
|
||||
return Err(SendError::Permanent { message, task: task.clone() });
|
||||
}
|
||||
};
|
||||
let path = file.path().to_path_buf();
|
||||
match send_animation_inner(
|
||||
bot,
|
||||
chat_id,
|
||||
reply_to,
|
||||
caption,
|
||||
has_spoiler,
|
||||
InputFile::file(path),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(message) => Ok(vec![message.id.0 as i64]),
|
||||
Err(e) => Err(classify_to_send_error(&e, task.clone())),
|
||||
}
|
||||
}
|
||||
Err(e) => Err(classify_to_send_error(&e, task.clone())),
|
||||
}
|
||||
}
|
||||
|
||||
/// Copies already-sent messages to the forward channel. No download fallback:
|
||||
/// the files are already on Telegram's servers.
|
||||
pub async fn forward_messages(bot: &Bot, task: &Task) -> Result<(), SendError> {
|
||||
let Task::ForwardMessages {
|
||||
from_chat_id,
|
||||
to_chat_id,
|
||||
message_ids,
|
||||
..
|
||||
} = task
|
||||
else {
|
||||
unreachable!("forward_messages requires a ForwardMessages task")
|
||||
};
|
||||
let message_ids = message_ids
|
||||
.iter()
|
||||
.map(|id| MessageId(*id as i32))
|
||||
.collect::<Vec<_>>();
|
||||
match bot
|
||||
.copy_messages(ChatId(*to_chat_id), ChatId(*from_chat_id), message_ids.clone())
|
||||
.await
|
||||
{
|
||||
Ok(_) => {
|
||||
log::info!(
|
||||
"copied {} message(s) from {} to {}",
|
||||
message_ids.len(),
|
||||
from_chat_id,
|
||||
to_chat_id
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
Err(e) => Err(classify_to_send_error(&e, task.clone())),
|
||||
}
|
||||
}
|
||||
|
||||
/// One button per template name (column layout), then the confirm button.
|
||||
pub fn build_edit_markup(templates: &HashMap<String, String>) -> InlineKeyboardMarkup {
|
||||
let mut rows = Vec::new();
|
||||
for name in templates.keys() {
|
||||
rows.push(vec![InlineKeyboardButton::callback(
|
||||
name.clone(),
|
||||
format!("template|{name}"),
|
||||
)]);
|
||||
}
|
||||
rows.push(vec![InlineKeyboardButton::callback(
|
||||
"↩️ Confirm",
|
||||
"forward",
|
||||
)]);
|
||||
InlineKeyboardMarkup::new(rows)
|
||||
}
|
||||
|
||||
/// Notifies a chat about a dead-lettered task (skips when `notify_chat_id` is
|
||||
/// absent).
|
||||
pub async fn notify_failure(bot: &Bot, chat_id: Option<i64>, message_id: Option<i64>, message: &str) {
|
||||
let Some(chat_id) = chat_id else { return };
|
||||
let mut request = bot.send_message(ChatId(chat_id), message);
|
||||
if let Some(message_id) = message_id {
|
||||
request = request
|
||||
.reply_parameters(ReplyParameters::new(MessageId(message_id as i32)).allow_sending_without_reply());
|
||||
}
|
||||
if let Err(e) = request.await {
|
||||
log::error!("failed to notify about failed task: {e}");
|
||||
}
|
||||
}
|
||||
|
||||
/// After a successful send: either open the edit-before-forward prompt or
|
||||
/// forward to the configured channel (with retry/queue handling).
|
||||
pub async fn post_send_actions(bot: &Bot, task: &Task, message_ids: Vec<i64>) {
|
||||
let (chat_id, reply_to, source_url, edit_before_forward, forward_channel_id, notify_chat_id, notify_message_id) =
|
||||
match task {
|
||||
Task::SendMediaSequence {
|
||||
chat_id,
|
||||
reply_to_message_id,
|
||||
source_url,
|
||||
edit_before_forward,
|
||||
forward_channel_id,
|
||||
notify_chat_id,
|
||||
notify_message_id,
|
||||
..
|
||||
}
|
||||
| Task::SendAnimation {
|
||||
chat_id,
|
||||
reply_to_message_id,
|
||||
source_url,
|
||||
edit_before_forward,
|
||||
forward_channel_id,
|
||||
notify_chat_id,
|
||||
notify_message_id,
|
||||
..
|
||||
} => (
|
||||
*chat_id,
|
||||
*reply_to_message_id,
|
||||
source_url.clone(),
|
||||
*edit_before_forward,
|
||||
*forward_channel_id,
|
||||
*notify_chat_id,
|
||||
*notify_message_id,
|
||||
),
|
||||
Task::ForwardMessages { .. } => return,
|
||||
};
|
||||
|
||||
if edit_before_forward {
|
||||
let mut chat_data = CHAT_STORE.get(chat_id).await;
|
||||
let keyboard = build_edit_markup(&chat_data.template);
|
||||
match bot
|
||||
.send_message(ChatId(chat_id), "Reply to edit message.")
|
||||
.reply_markup(keyboard)
|
||||
.reply_parameters(
|
||||
ReplyParameters::new(MessageId(reply_to as i32)).allow_sending_without_reply(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(prompt) => {
|
||||
log::info!(
|
||||
"edit-before-forward prompt {} opened for {} message(s)",
|
||||
prompt.id.0,
|
||||
message_ids.len()
|
||||
);
|
||||
chat_data.edit_message.insert(
|
||||
prompt.id.0 as i64,
|
||||
EditMessage {
|
||||
url: source_url,
|
||||
chat_id,
|
||||
forward_message_ids: message_ids,
|
||||
template: String::new(),
|
||||
created_at: unix_now(),
|
||||
},
|
||||
);
|
||||
CHAT_STORE.set(chat_id, &chat_data).await;
|
||||
}
|
||||
Err(e) => log::error!("failed to send edit prompt: {e}"),
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
if let Some(channel_id) = forward_channel_id {
|
||||
log::info!("forwarding {} message(s) to channel {channel_id}", message_ids.len());
|
||||
let forward_task = Task::ForwardMessages {
|
||||
from_chat_id: chat_id,
|
||||
to_chat_id: channel_id,
|
||||
message_ids,
|
||||
notify_chat_id,
|
||||
notify_message_id,
|
||||
};
|
||||
match forward_messages(bot, &forward_task).await {
|
||||
Ok(()) => {}
|
||||
Err(SendError::Retryable { delay_seconds, task }) => {
|
||||
let payload = serde_json::to_value(task).expect("task serializes");
|
||||
let run_after = std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.map(|d| d.as_secs_f64())
|
||||
.unwrap_or(0.0)
|
||||
+ delay_seconds;
|
||||
if let Err(e) = TASK_QUEUE.enqueue(payload, run_after).await {
|
||||
log::error!("failed to enqueue forward retry: {e}");
|
||||
}
|
||||
}
|
||||
Err(SendError::Permanent { message, .. }) => {
|
||||
notify_failure(
|
||||
bot,
|
||||
notify_chat_id,
|
||||
notify_message_id,
|
||||
&format!("Task failed after retries: {message}"),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Queue entry point: parses the stored task and dispatches.
|
||||
pub async fn handle_task(payload: serde_json::Value) -> Result<(), QueueError> {
|
||||
let task: Task = match serde_json::from_value(payload.clone()) {
|
||||
Ok(task) => task,
|
||||
Err(e) => {
|
||||
return Err(QueueError::Permanent {
|
||||
message: format!("invalid task payload: {e}"),
|
||||
payload,
|
||||
});
|
||||
}
|
||||
};
|
||||
let bot = Bot::from_env();
|
||||
match task {
|
||||
Task::SendMediaSequence { .. } | Task::SendAnimation { .. } => {
|
||||
let message_ids = match send_media_or_animation(&bot, &task).await {
|
||||
Ok(ids) => ids,
|
||||
Err(SendError::Retryable { delay_seconds, task }) => {
|
||||
return Err(QueueError::Retryable {
|
||||
delay_seconds,
|
||||
payload: serde_json::to_value(task).expect("task serializes"),
|
||||
});
|
||||
}
|
||||
Err(SendError::Permanent { message, task }) => {
|
||||
return Err(QueueError::Permanent {
|
||||
message,
|
||||
payload: serde_json::to_value(task).expect("task serializes"),
|
||||
});
|
||||
}
|
||||
};
|
||||
post_send_actions(&bot, &task, message_ids).await;
|
||||
Ok(())
|
||||
}
|
||||
Task::ForwardMessages { .. } => match forward_messages(&bot, &task).await {
|
||||
Ok(()) => Ok(()),
|
||||
Err(SendError::Retryable { delay_seconds, task }) => Err(QueueError::Retryable {
|
||||
delay_seconds,
|
||||
payload: serde_json::to_value(task).expect("task serializes"),
|
||||
}),
|
||||
Err(SendError::Permanent { message, task }) => Err(QueueError::Permanent {
|
||||
message,
|
||||
payload: serde_json::to_value(task).expect("task serializes"),
|
||||
}),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
async fn send_media_or_animation(bot: &Bot, task: &Task) -> Result<Vec<i64>, SendError> {
|
||||
match task {
|
||||
Task::SendMediaSequence { .. } => send_media_sequence(bot, task).await,
|
||||
Task::SendAnimation { .. } => send_animation(bot, task).await,
|
||||
Task::ForwardMessages { .. } => unreachable!(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Dead-letter callback wired to the queue in main: notifies the task's chat.
|
||||
pub async fn dead_letter_notify(payload: serde_json::Value, message: String) {
|
||||
let notify_chat_id = payload.get("notify_chat_id").and_then(|v| v.as_i64());
|
||||
let notify_message_id = payload.get("notify_message_id").and_then(|v| v.as_i64());
|
||||
if notify_chat_id.is_some() {
|
||||
let bot = Bot::from_env();
|
||||
notify_failure(
|
||||
&bot,
|
||||
notify_chat_id,
|
||||
notify_message_id,
|
||||
&format!("Task failed after retries: {message}"),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn chunk_media_items_sizes() {
|
||||
assert_eq!(chunk_media_items::<i32>(vec![]), Vec::<Vec<i32>>::new());
|
||||
assert_eq!(chunk_media_items((0..9).collect()).len(), 1);
|
||||
assert_eq!(chunk_media_items((0..10).collect()).len(), 2);
|
||||
assert_eq!(chunk_media_items((0..25).collect()).len(), 3);
|
||||
assert_eq!(chunk_media_items((0..25).collect())[2].len(), 7);
|
||||
assert!(chunk_media_items((0..25).collect()).iter().all(|c| c.len() <= 9));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn retry_delay_seconds_bounds() {
|
||||
for attempts in 0..10 {
|
||||
let delay = retry_delay_seconds(attempts);
|
||||
assert!(delay >= 1.0, "attempts={attempts}: {delay}");
|
||||
assert!(delay <= 30.0, "attempts={attempts}: {delay}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_media_fetch_failure_matches_markers() {
|
||||
for description in [
|
||||
"Bad Request: WEBPAGE_MEDIA_EMPTY",
|
||||
"Bad Request: media_empty",
|
||||
"Bad Request: EMPTY_WEB_MEDIA",
|
||||
"Bad Request: webpage_curl_failed",
|
||||
"Bad Request: request timeout",
|
||||
] {
|
||||
let api = ApiError::Unknown(description.to_string());
|
||||
assert!(is_media_fetch_failure(&api), "{description}");
|
||||
}
|
||||
for description in ["Bad Request: message is not modified", "Forbidden: bot was blocked by the user"] {
|
||||
let api = ApiError::Unknown(description.to_string());
|
||||
assert!(!is_media_fetch_failure(&api), "{description}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classification_mapping() {
|
||||
use teloxide::types::Seconds;
|
||||
// RetryAfter -> Retryable with its delay
|
||||
let e = RequestError::RetryAfter(Seconds::from_seconds(7));
|
||||
assert!(matches!(
|
||||
classify_request_error(&e),
|
||||
Classification::Retryable { delay_seconds } if delay_seconds == 7.0
|
||||
));
|
||||
// Api error -> Permanent
|
||||
let e = RequestError::Api(ApiError::Unknown("Bad Request: something".into()));
|
||||
assert!(matches!(
|
||||
classify_request_error(&e),
|
||||
Classification::Permanent { .. }
|
||||
));
|
||||
// Api media-fetch marker -> MediaFetchFailure
|
||||
let e = RequestError::Api(ApiError::Unknown("Bad Request: WEBPAGE_MEDIA_EMPTY".into()));
|
||||
assert!(matches!(
|
||||
classify_request_error(&e),
|
||||
Classification::MediaFetchFailure
|
||||
));
|
||||
// MigrateToChatId -> Permanent
|
||||
let e = RequestError::MigrateToChatId(ChatId(123));
|
||||
assert!(matches!(
|
||||
classify_request_error(&e),
|
||||
Classification::Permanent { .. }
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn task_serde_round_trip_preserves_resume_state() {
|
||||
let task = Task::SendMediaSequence {
|
||||
chat_id: 111,
|
||||
reply_to_message_id: 222,
|
||||
caption: "cap".into(),
|
||||
media_batches: vec![
|
||||
vec![MediaItemPayload::Photo {
|
||||
media: "https://a/b.jpg".into(),
|
||||
has_spoiler: true,
|
||||
}],
|
||||
vec![MediaItemPayload::Video {
|
||||
media: "https://a/v.mp4".into(),
|
||||
has_spoiler: false,
|
||||
thumbnail: Some("https://a/t.jpg".into()),
|
||||
}],
|
||||
],
|
||||
batch_index: 1,
|
||||
sent_message_ids: vec![11, 12],
|
||||
source_url: "https://x.com/u/status/1".into(),
|
||||
edit_before_forward: true,
|
||||
forward_channel_id: Some(333),
|
||||
notify_chat_id: Some(111),
|
||||
notify_message_id: Some(222),
|
||||
};
|
||||
let json = serde_json::to_value(&task).unwrap();
|
||||
assert_eq!(json["type"], "send_media_sequence");
|
||||
assert_eq!(json["batch_index"], 1);
|
||||
let decoded: Task = serde_json::from_value(json).unwrap();
|
||||
match decoded {
|
||||
Task::SendMediaSequence {
|
||||
batch_index,
|
||||
sent_message_ids,
|
||||
forward_channel_id,
|
||||
media_batches,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(batch_index, 1);
|
||||
assert_eq!(sent_message_ids, vec![11, 12]);
|
||||
assert_eq!(forward_channel_id, Some(333));
|
||||
assert_eq!(media_batches.len(), 2);
|
||||
assert!(matches!(media_batches[0][0], MediaItemPayload::Photo { has_spoiler: true, .. }));
|
||||
}
|
||||
other => panic!("expected SendMediaSequence, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn media_item_payload_serde_tags() {
|
||||
let photo = MediaItemPayload::Photo {
|
||||
media: "https://a/b.jpg".into(),
|
||||
has_spoiler: false,
|
||||
};
|
||||
let json = serde_json::to_value(&photo).unwrap();
|
||||
assert_eq!(json["kind"], "photo");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sniff_ext_detects_formats() {
|
||||
assert_eq!(sniff_ext(&[0xFF, 0xD8, 0xFF, 0xE0]), "jpg");
|
||||
assert_eq!(sniff_ext(b"\x89PNG\r\n\x1a\n"), "png");
|
||||
assert_eq!(sniff_ext(b"RIFF\x00\x00\x00\x00WEBPVP8 "), "webp");
|
||||
assert_eq!(sniff_ext(b"GIF89a"), "gif");
|
||||
assert_eq!(sniff_ext(b"\x00\x00\x00\x18ftypisom"), "mp4");
|
||||
assert_eq!(sniff_ext(b"something else"), "bin");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
//! Per-chat state with SQLite persistence (table `chat_state` in
|
||||
//! `data/task_queue.db`, shared with the task queue).
|
||||
|
||||
use parking_lot::Mutex;
|
||||
use rusqlite::{params, Connection};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::path::Path;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
#[derive(Serialize, Deserialize, Default, Clone, Debug)]
|
||||
pub struct ChatData {
|
||||
pub forward_channel_id: Option<i64>,
|
||||
pub edit_before_forward: bool,
|
||||
/// Key: prompt message id.
|
||||
pub edit_message: HashMap<i64, EditMessage>,
|
||||
/// name -> HTML template containing "[]"
|
||||
pub template: HashMap<String, String>,
|
||||
/// site name (twitter/bsky/pixiv) -> user-supplied caption format with
|
||||
/// {url} {author} {author_url} {title} {tags} placeholders.
|
||||
pub message_format: HashMap<String, String>,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize, Clone, Debug, Default)]
|
||||
pub struct EditMessage {
|
||||
pub url: String,
|
||||
pub chat_id: i64,
|
||||
pub forward_message_ids: Vec<i64>,
|
||||
pub template: String,
|
||||
/// Unix seconds at registration; expiry = created_at + ttl.
|
||||
pub created_at: i64,
|
||||
}
|
||||
|
||||
pub struct ChatStore {
|
||||
/// In-memory cache; the DB is the source of truth on first access.
|
||||
cache: Mutex<HashMap<i64, ChatData>>,
|
||||
db_path: String,
|
||||
}
|
||||
|
||||
pub fn unix_now() -> i64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.as_secs() as i64)
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
impl ChatStore {
|
||||
/// Creates the parent directory and both tables (idempotent).
|
||||
pub fn open(path: &str) -> rusqlite::Result<Self> {
|
||||
if let Some(parent) = Path::new(path).parent()
|
||||
&& !parent.as_os_str().is_empty()
|
||||
{
|
||||
std::fs::create_dir_all(parent)
|
||||
.map_err(|e| rusqlite::Error::ToSqlConversionFailure(Box::new(e)))?;
|
||||
}
|
||||
let conn = Connection::open(path)?;
|
||||
conn.execute_batch(
|
||||
"CREATE TABLE IF NOT EXISTS tasks (id TEXT PRIMARY KEY, payload TEXT NOT NULL, \
|
||||
run_after REAL NOT NULL, attempts INTEGER NOT NULL, status TEXT NOT NULL, \
|
||||
locked_until REAL NOT NULL, created_at REAL NOT NULL); \
|
||||
CREATE TABLE IF NOT EXISTS chat_state (chat_id TEXT PRIMARY KEY, payload TEXT NOT NULL);",
|
||||
)?;
|
||||
drop(conn);
|
||||
Ok(ChatStore {
|
||||
cache: Mutex::new(HashMap::new()),
|
||||
db_path: path.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn get(&self, chat_id: i64) -> ChatData {
|
||||
if let Some(data) = self.cache.lock().get(&chat_id) {
|
||||
return data.clone();
|
||||
}
|
||||
let db_path = self.db_path.clone();
|
||||
let payload = tokio::task::spawn_blocking(move || -> rusqlite::Result<Option<String>> {
|
||||
let conn = Connection::open(&db_path)?;
|
||||
let mut stmt = conn.prepare("SELECT payload FROM chat_state WHERE chat_id = ?1")?;
|
||||
let mut rows = stmt.query(params![chat_id.to_string()])?;
|
||||
match rows.next()? {
|
||||
Some(row) => Ok(Some(row.get(0)?)),
|
||||
None => Ok(None),
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("chat_state worker panicked")
|
||||
.unwrap_or_else(|e| {
|
||||
log::error!("chat_state read failed: {e}");
|
||||
None
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let data: ChatData = serde_json::from_str(&payload).unwrap_or_default();
|
||||
self.cache.lock().insert(chat_id, data.clone());
|
||||
data
|
||||
}
|
||||
|
||||
/// Write-through: update the cache and the DB.
|
||||
pub async fn set(&self, chat_id: i64, data: &ChatData) {
|
||||
self.cache.lock().insert(chat_id, data.clone());
|
||||
let payload = serde_json::to_string(data).expect("chat state serializes");
|
||||
let db_path = self.db_path.clone();
|
||||
tokio::task::spawn_blocking(move || -> rusqlite::Result<()> {
|
||||
let conn = Connection::open(&db_path)?;
|
||||
conn.execute(
|
||||
"INSERT OR REPLACE INTO chat_state (chat_id, payload) VALUES (?1, ?2)",
|
||||
params![chat_id.to_string(), payload],
|
||||
)?;
|
||||
Ok(())
|
||||
})
|
||||
.await
|
||||
.expect("chat_state worker panicked")
|
||||
.unwrap_or_else(|e| log::error!("chat_state write failed: {e}"));
|
||||
}
|
||||
|
||||
/// Removes edit-before-forward records whose `created_at + ttl` is in the
|
||||
/// past. Returns the removed `(chat_id, prompt_message_id)` pairs so the
|
||||
/// caller can clear the prompt's buttons.
|
||||
pub async fn prune_expired(&self, ttl: Duration) -> Vec<(i64, i64)> {
|
||||
let now = unix_now();
|
||||
let ttl_secs = ttl.as_secs() as i64;
|
||||
let mut removed = Vec::new();
|
||||
let changed: Vec<(i64, ChatData)> = {
|
||||
let mut cache = self.cache.lock();
|
||||
let mut out = Vec::new();
|
||||
for (chat_id, data) in cache.iter_mut() {
|
||||
let keys: Vec<i64> = data.edit_message.keys().copied().collect();
|
||||
let mut kept = HashMap::new();
|
||||
for key in keys {
|
||||
if let Some(entry) = data.edit_message.get(&key) {
|
||||
if entry.created_at + ttl_secs > now {
|
||||
kept.insert(key, entry.clone());
|
||||
} else {
|
||||
removed.push((*chat_id, key));
|
||||
}
|
||||
}
|
||||
}
|
||||
if kept.len() != data.edit_message.len() {
|
||||
data.edit_message = kept;
|
||||
out.push((*chat_id, data.clone()));
|
||||
}
|
||||
}
|
||||
out
|
||||
};
|
||||
for (chat_id, data) in changed {
|
||||
self.set(chat_id, &data).await;
|
||||
}
|
||||
if !removed.is_empty() {
|
||||
log::info!("pruned {} expired edit-before-forward record(s)", removed.len());
|
||||
}
|
||||
removed
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user