use std::sync::Arc;
use crate::error::Error;
use crate::events::{Event, EventKind};
use crate::extract::ExtractBag;
use crate::filters::{FilterOutcome, MessageFilter};
use crate::middleware::Middleware;
use crate::models::{Interaction, Message, Ready};
pub type HandlerResult = Result<(), Error>;
pub type MessageHandler = Box<dyn Fn(&Message, &ExtractBag) -> HandlerResult + Send + Sync>;
pub(crate) type SharedMessageFilter =
Arc<dyn Fn(&Message, &ExtractBag) -> Result<FilterOutcome, Error> + Send + Sync>;
pub(crate) type SharedEventHandler =
Arc<dyn Fn(&Event, &ExtractBag) -> HandlerResult + Send + Sync>;
type SharedMessageHandler = Arc<dyn Fn(&Message, &ExtractBag) -> HandlerResult + Send + Sync>;
pub struct MessageHandlerDef {
name: &'static str,
filters: Vec<MessageFilter>,
handler: MessageHandler,
}
impl MessageHandlerDef {
pub fn new<F>(name: &'static str, handler: F, filters: Vec<MessageFilter>) -> Self
where
F: Fn(&Message) -> HandlerResult + Send + Sync + 'static,
{
Self {
name,
filters,
handler: Box::new(move |message, _bag| handler(message)),
}
}
pub fn name(&self) -> &'static str {
self.name
}
}
pub(crate) struct Route {
pub(crate) kind: EventKind,
pub(crate) filters: Vec<SharedMessageFilter>,
pub(crate) handler: SharedEventHandler,
}
pub(crate) struct CollectedRoute {
pub(crate) kind: EventKind,
pub(crate) filters: Vec<SharedMessageFilter>,
pub(crate) middlewares: Vec<Middleware>,
pub(crate) handler: SharedEventHandler,
}
#[derive(Default)]
pub struct Router {
name: &'static str,
routes: Vec<Route>,
children: Vec<Router>,
middlewares: Vec<Middleware>,
}
impl Router {
pub fn new() -> Self {
Self::default()
}
pub fn named(name: &'static str) -> Self {
Self {
name,
..Self::default()
}
}
pub fn name(&self) -> &'static str {
self.name
}
pub fn use_middleware<F>(&mut self, middleware: F)
where
F: for<'a> Fn(&Event, &ExtractBag, crate::middleware::Next<'a>) -> HandlerResult
+ Send
+ Sync
+ 'static,
{
self.middlewares.push(Arc::new(middleware));
}
pub fn include(&mut self, child: Router) {
self.children.push(child);
}
pub fn on_message<F>(&mut self, handler: F)
where
F: Fn(&Message) -> HandlerResult + Send + Sync + 'static,
{
self.on_message_filtered(handler, Vec::new());
}
pub fn on_message_filtered<F>(&mut self, handler: F, filters: Vec<MessageFilter>)
where
F: Fn(&Message) -> HandlerResult + Send + Sync + 'static,
{
self.push_message_route(
filters,
Box::new(move |message, _bag| handler(message)),
);
}
pub fn add_message_handler(&mut self, definition: MessageHandlerDef) {
self.push_message_route(definition.filters, definition.handler);
}
fn push_message_route(&mut self, filters: Vec<MessageFilter>, handler: MessageHandler) {
let handler: SharedMessageHandler = Arc::from(handler);
self.routes.push(Route {
kind: EventKind::MessageCreate,
filters: filters.into_iter().map(SharedMessageFilter::from).collect(),
handler: Arc::new(move |event: &Event, bag: &ExtractBag| match event.message() {
Some(message) => handler(message, bag),
None => Ok(()),
}),
});
}
pub fn on_event<F>(&mut self, kind: EventKind, handler: F)
where
F: Fn(&Event) -> HandlerResult + Send + Sync + 'static,
{
self.routes.push(Route {
kind,
filters: Vec::new(),
handler: Arc::new(move |event: &Event, _bag: &ExtractBag| handler(event)),
});
}
pub fn on_ready<F>(&mut self, handler: F)
where
F: Fn(&Ready) -> HandlerResult + Send + Sync + 'static,
{
self.on_event(EventKind::Ready, move |event| match event {
Event::Ready(ready) => handler(ready),
_ => Ok(()),
});
}
pub fn on_interaction<F>(&mut self, handler: F)
where
F: Fn(&Interaction) -> HandlerResult + Send + Sync + 'static,
{
self.on_event(EventKind::InteractionCreate, move |event| match event {
Event::InteractionCreate(interaction) => handler(interaction.as_ref()),
_ => Ok(()),
});
}
pub(crate) fn collect_routes(&self, parent: &[Middleware]) -> Vec<CollectedRoute> {
let chain: Vec<Middleware> = parent
.iter()
.cloned()
.chain(self.middlewares.iter().cloned())
.collect();
let mut out = Vec::with_capacity(self.routes.len());
for route in &self.routes {
out.push(CollectedRoute {
kind: route.kind,
filters: route.filters.clone(),
middlewares: chain.clone(),
handler: Arc::clone(&route.handler),
});
}
for child in &self.children {
out.extend(child.collect_routes(&chain));
}
out
}
pub fn dispatch_message(&self, message: &Message) -> HandlerResult {
let mut dispatcher = crate::dispatcher::Dispatcher::new();
dispatcher.include(self);
dispatcher.dispatch(&Event::MessageCreate(message.clone()))
}
}
pub use crate::filters::{author_id, content_starts_with};
#[macro_export]
macro_rules! register_on_message {
($router:expr, $handler:expr) => {
$router.on_message($handler)
};
($router:expr, $handler:expr, filters = [$($filter:expr),* $(,)?]) => {
$router.on_message_filtered($handler, vec![$($filter),*])
};
($router:expr, $handler:expr, any = [$($filter:expr),+ $(,)?]) => {{
let mut __any: ::std::vec::Vec<$crate::MessageFilter> = vec![$($filter),+];
let mut __iter = __any.drain(..);
let mut __acc = __iter.next().expect("`any` requires at least one filter");
for __f in __iter {
__acc = $crate::or(__acc, __f);
}
$router.on_message_filtered($handler, vec![__acc])
}};
(
$router:expr,
$handler:expr,
filters = [$($filter:expr),* $(,)?],
any = [$($any_filter:expr),+ $(,)?]
) => {{
let mut __any: ::std::vec::Vec<$crate::MessageFilter> = vec![$($any_filter),+];
let mut __iter = __any.drain(..);
let mut __acc = __iter.next().expect("`any` requires at least one filter");
for __f in __iter {
__acc = $crate::or(__acc, __f);
}
let mut __all: ::std::vec::Vec<$crate::MessageFilter> = vec![$($filter),*];
__all.push(__acc);
$router.on_message_filtered($handler, __all)
}};
}
#[cfg(test)]
mod tests {
use super::*;
use crate::models::User;
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
fn message(content: &str) -> Message {
Message {
id: "1".to_string(),
channel_id: "2".to_string(),
guild_id: None,
author: User {
id: "42".to_string(),
username: "tester".to_string(),
discriminator: None,
global_name: None,
bot: false,
system: false,
avatar: None,
banner: None,
public_flags: None,
},
content: content.to_string(),
timestamp: None,
edited_timestamp: None,
tts: false,
mention_everyone: false,
mentions: Vec::new(),
embeds: Vec::new(),
attachments: Vec::new(),
components: Vec::new(),
flags: None,
}
}
#[test]
fn router_dispatches_filtered_messages() {
let mut router = Router::new();
let calls = Arc::new(AtomicUsize::new(0));
let handler_calls = Arc::clone(&calls);
register_on_message!(
router,
move |message: &Message| {
handler_calls.fetch_add(1, Ordering::SeqCst);
assert_eq!(message.content, "!ping");
Ok(())
},
filters = [content_starts_with("!")]
);
router.dispatch_message(&message("plain")).unwrap();
router.dispatch_message(&message("!ping")).unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
fn ping(message: &Message) -> HandlerResult {
assert_eq!(message.content, "!ping");
Ok(())
}
#[test]
fn router_accepts_handler_definitions() {
let mut router = Router::new();
let definition = MessageHandlerDef::new("ping", ping, vec![content_starts_with("!")]);
assert_eq!(definition.name(), "ping");
router.add_message_handler(definition);
router.dispatch_message(&message("plain")).unwrap();
router.dispatch_message(&message("!ping")).unwrap();
}
#[test]
fn router_runs_handler_only_when_all_filters_pass() {
let mut router = Router::new();
let calls = Arc::new(AtomicUsize::new(0));
let handler_calls = Arc::clone(&calls);
register_on_message!(
router,
move |_message: &Message| {
handler_calls.fetch_add(1, Ordering::SeqCst);
Ok(())
},
filters = [
content_starts_with("!"),
crate::command("ping")
]
);
router.dispatch_message(&message("hello")).unwrap(); router.dispatch_message(&message("!hello")).unwrap(); router.dispatch_message(&message("/ping")).unwrap(); router.dispatch_message(&message("!ping")).unwrap(); assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[test]
fn router_any_macro_matches_when_either_filter_passes() {
let mut router = Router::new();
let calls = Arc::new(AtomicUsize::new(0));
let handler_calls = Arc::clone(&calls);
register_on_message!(
router,
move |_message: &Message| {
handler_calls.fetch_add(1, Ordering::SeqCst);
Ok(())
},
any = [crate::command("ping"), crate::command("pong")]
);
router.dispatch_message(&message("/ping")).unwrap();
router.dispatch_message(&message("/pong")).unwrap();
router.dispatch_message(&message("/foo")).unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[test]
fn router_combines_filters_and_any_branches() {
let mut router = Router::new();
let calls = Arc::new(AtomicUsize::new(0));
let handler_calls = Arc::clone(&calls);
register_on_message!(
router,
move |_message: &Message| {
handler_calls.fetch_add(1, Ordering::SeqCst);
Ok(())
},
filters = [content_starts_with("/")],
any = [crate::command("ping"), crate::command("pong")]
);
router.dispatch_message(&message("/ping")).unwrap(); router.dispatch_message(&message("!ping")).unwrap(); router.dispatch_message(&message("/foo")).unwrap(); assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[test]
fn command_filter_extracts_args_for_handler() {
let mut router = Router::new();
let calls = Arc::new(AtomicUsize::new(0));
let handler_calls = Arc::clone(&calls);
router.on_message_filtered(
move |_message: &Message| {
handler_calls.fetch_add(1, Ordering::SeqCst);
Ok(())
},
vec![crate::command("echo")],
);
router.dispatch_message(&message("plain")).unwrap();
router.dispatch_message(&message("!echo hello")).unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[test]
fn nested_router_middleware_wraps_child_routes() {
let log: Arc<std::sync::Mutex<Vec<String>>> = Arc::new(std::sync::Mutex::new(Vec::new()));
let mut child = Router::named("child");
let child_log = Arc::clone(&log);
child.use_middleware(move |event, bag, next| {
child_log.lock().unwrap().push("child:in".into());
let result = next.run(event, bag);
child_log.lock().unwrap().push("child:out".into());
result
});
let handler_log = Arc::clone(&log);
child.on_message(move |_message| {
handler_log.lock().unwrap().push("handler".into());
Ok(())
});
let mut parent = Router::named("parent");
let parent_log = Arc::clone(&log);
parent.use_middleware(move |event, bag, next| {
parent_log.lock().unwrap().push("parent:in".into());
let result = next.run(event, bag);
parent_log.lock().unwrap().push("parent:out".into());
result
});
parent.include(child);
parent.dispatch_message(&message("hi")).unwrap();
assert_eq!(
*log.lock().unwrap(),
vec!["parent:in", "child:in", "handler", "child:out", "parent:out"]
);
}
}