use parking_lot::RwLock;
use smol::prelude::*;
use std::any::{Any, TypeId};
use std::error::Error;
use std::fmt::{Debug, Display, Formatter};
use std::future::Future;
use std::ops::{Deref, DerefMut};
use std::pin::Pin;
use std::sync::{Arc, Weak};
use std::task::{Context, Poll};
use ahash::{HashMap, HashMapExt, HashSet, HashSetExt};
pub trait EventCancellable {}
pub struct Event<T> {
inner: T,
cancelled: bool,
}
impl<T> Event<T> {
fn new(inner: T) -> Self {
Self {
inner,
cancelled: false,
}
}
#[inline]
pub fn data(&mut self) -> &mut T { &mut self.inner }
#[inline]
pub fn into_data(self) -> T { self.inner }
}
impl<T> Event<T>
where
T: EventCancellable
{
pub fn cancel(&mut self) {
self.cancelled = true;
}
#[inline]
pub fn is_cancelled(&self) -> bool {
self.cancelled
}
}
impl<T> Deref for Event<T> {
type Target = T;
#[inline]
fn deref(&self) -> &Self::Target { &self.inner }
}
impl<T> DerefMut for Event<T> {
#[inline]
fn deref_mut(&mut self) -> &mut Self::Target { &mut self.inner }
}
pub type SubscriberFnOnce<T> = Box<dyn FnOnce(&mut Event<T>) + Send>;
pub trait EventHandler<T>: Send + Sync + 'static {
fn handle_event(&self, event: &mut Event<T>);
#[inline]
fn is_valid(&self) -> bool { true }
}
impl<T, F> EventHandler<T> for F
where
F: Fn(&mut Event<T>) + Send + Sync + 'static,
{
#[inline]
fn handle_event(&self, event: &mut Event<T>) {
(*self)(event)
}
}
impl<T, H> EventHandler<T> for Arc<H>
where
H: EventHandler<T>,
{
#[inline]
fn handle_event(&self, event: &mut Event<T>) {
(**self).handle_event(event)
}
}
impl<T, H> EventHandler<T> for Weak<H>
where
H: EventHandler<T>,
{
#[inline]
fn handle_event(&self, event: &mut Event<T>) {
if let Some(handler) = self.upgrade() {
handler.handle_event(event)
}
}
#[inline]
fn is_valid(&self) -> bool {
if let Some(handler) = self.upgrade() {
handler.is_valid()
} else {
false
}
}
}
struct SubscriberEntry(Box<dyn Any + Send + Sync + 'static>);
impl SubscriberEntry {
fn new<T>(handler: Box<dyn EventHandler<T>>) -> Self
where
for<'a> T: 'a,
{
Self(Box::new(handler))
}
fn handle_event<T>(&self, event: &mut Event<T>) -> bool
where
for<'a> T: 'a,
{
let handler = self.0.downcast_ref::<Box<dyn EventHandler<T>>>().unwrap();
if handler.is_valid() {
handler.handle_event(event);
true
} else {
false
}
}
}
#[derive(Debug)]
pub enum SubscriberOnceError {
Dropped,
}
impl Display for SubscriberOnceError {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
match self {
SubscriberOnceError::Dropped => f.write_str("QueuedAction was dropped before being executed"),
}
}
}
impl Error for SubscriberOnceError {}
impl From<ump::Error<()>> for SubscriberOnceError {
#[inline]
fn from(value: ump::Error<()>) -> Self {
match value {
ump::Error::ServerDisappeared | ump::Error::ClientsDisappeared | ump::Error::NoReply =>
Self::Dropped,
ump::Error::App(e) => panic!("Unexpected App error: {e:?}"),
}
}
}
enum SubscriberOnceFutureState {
Init(Result<ump::WaitReply<(), ()>, SubscriberOnceError>),
Awaiting(Pin<Box<dyn Future<Output = Result<(), ump::Error<()>>> + Send>>),
Uninit,
}
impl SubscriberOnceFutureState {
fn poll(self, cx: &mut Context<'_>) -> (SubscriberOnceFutureState, Poll<Result<(), SubscriberOnceError>>) {
match self {
Self::Init(result) => {
match result {
Ok(wait_reply) =>
Self::Awaiting(Box::pin(wait_reply.wait_async())).poll(cx),
Err(err) =>
(Self::Uninit, Poll::Ready(Err(err))),
}
},
Self::Awaiting(mut fut) => {
match fut.poll(cx) {
Poll::Ready(result) =>
(Self::Uninit, Poll::Ready(result.map_err(Into::into))),
Poll::Pending => (Self::Awaiting(fut), Poll::Pending),
}
},
Self::Uninit => unreachable!("Found uninit state"),
}
}
}
pub struct SubscriberOnceFuture(SubscriberOnceFutureState);
impl SubscriberOnceFuture {
fn map_ump_error(error: ump::Error<()>) -> SubscriberOnceError {
match error {
ump::Error::ServerDisappeared | ump::Error::ClientsDisappeared | ump::Error::NoReply
=> SubscriberOnceError::Dropped,
ump::Error::App(e) => panic!("Unexpected App error: {e:?}"),
}
}
fn new(wait_reply: Result<ump::WaitReply<(), ()>, ump::Error<()>>) -> Self {
Self(SubscriberOnceFutureState::Init(wait_reply.map_err(Self::map_ump_error)))
}
fn with_state<R>(
&mut self,
f: impl FnOnce(SubscriberOnceFutureState) -> (SubscriberOnceFutureState, R),
) -> R {
let mut state = SubscriberOnceFutureState::Uninit;
std::mem::swap(&mut self.0, &mut state);
let (mut state, result) = f(state);
std::mem::swap(&mut self.0, &mut state);
result
}
#[inline]
pub fn wait(self) -> Result<(), SubscriberOnceError> {
match self.0 {
SubscriberOnceFutureState::Init(wait_reply) =>
wait_reply?.wait().map_err(Self::map_ump_error),
SubscriberOnceFutureState::Awaiting(_) =>
panic!("Can't wait synchronously once async polling has started"),
SubscriberOnceFutureState::Uninit =>
unreachable!("Found uninit state"),
}
}
}
impl Future for SubscriberOnceFuture {
type Output = Result<(), SubscriberOnceError>;
#[inline]
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.with_state(|state| state.poll(cx))
}
}
struct SubscriberOncePipe {
client: Box<dyn Any + Send + Sync>, server: Box<dyn Any + Send + Sync>, }
impl SubscriberOncePipe {
fn new<T>() -> Self
where
for<'a> T: 'a,
{
let (server, client) = ump::channel::<SubscriberFnOnce<T>, (), ()>();
Self {
client: Box::new(client),
server: Box::new(server),
}
}
fn send<T>(&self, subscriber: SubscriberFnOnce<T>) -> SubscriberOnceFuture
where
for<'a> T: 'a,
{
let client = self.client.downcast_ref::<ump::Client<SubscriberFnOnce<T>, (), ()>>()
.unwrap();
SubscriberOnceFuture::new(client.req_async(subscriber))
}
fn handle_event<T>(&self, event: &mut Event<T>)
where
for<'a> T: 'a,
{
let server = self.server.downcast_ref::<ump::Server<SubscriberFnOnce<T>, (), ()>>()
.unwrap();
while let Some((subscriber, ctx)) = server.try_pop().unwrap() {
subscriber(event);
ctx.reply(()).unwrap();
}
}
}
#[derive(Debug)]
pub struct InvalidTypeIdError {
type_id: TypeId,
}
impl Display for InvalidTypeIdError {
#[inline]
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "TypeId '{:?}' is not valid for the EventBus", self.type_id)
}
}
impl Error for InvalidTypeIdError {}
pub struct EventBusRegistry {
valid_type_ids: HashSet<TypeId>,
subscribers: RwLock<HashMap<TypeId, Vec<SubscriberEntry>>>,
once_pipes: RwLock<HashMap<TypeId, SubscriberOncePipe>>,
}
impl EventBusRegistry {
pub fn register<T>(
&self,
subscriber: impl EventHandler<T>,
) -> Result<(), InvalidTypeIdError>
where
for<'a> T: 'a,
{
let type_id = TypeId::of::<T>();
if self.valid_type_ids.contains(&type_id) {
self.subscribers.write()
.entry(type_id)
.or_default()
.push(SubscriberEntry::new(Box::new(subscriber)));
Ok(())
} else {
Err(InvalidTypeIdError { type_id })
}
}
pub fn register_once<T>(
&self,
subscriber_once: impl FnOnce(&mut Event<T>) + Send + 'static,
) -> Result<SubscriberOnceFuture, InvalidTypeIdError>
where
for<'a> T: 'a,
{
let type_id = TypeId::of::<T>();
if self.valid_type_ids.contains(&type_id) {
let subscriber_once: SubscriberFnOnce<T> = Box::new(subscriber_once);
let mut once_pipes_map = self.once_pipes.upgradable_read();
let once_pipe = if once_pipes_map.contains_key(&type_id) {
&once_pipes_map[&type_id]
} else {
once_pipes_map.with_upgraded(|map| {
map.entry(type_id.clone())
.or_insert_with(SubscriberOncePipe::new::<T>);
});
&once_pipes_map[&type_id]
};
Ok(once_pipe.send(subscriber_once))
} else {
Err(InvalidTypeIdError { type_id })
}
}
}
pub struct EventBus {
registry: EventBusRegistry,
}
impl EventBus {
#[inline]
pub fn builder() -> EventBusBuilder {
EventBusBuilder::new()
}
#[inline]
pub fn registry(&self) -> &EventBusRegistry { &self.registry }
pub fn fire<T>(&self, event_data: T) -> Event<T>
where
for<'a> T: 'a,
{
let type_id = TypeId::of::<T>();
let mut event = Event::new(event_data);
let mut subscribers_map = self.registry.subscribers.write();
let once_pipes_map = self.registry.once_pipes.read();
if let Some(subscribers) = subscribers_map.get_mut(&type_id) {
subscribers.retain(|subscriber| subscriber.handle_event(&mut event));
}
if let Some(once_pipe) = once_pipes_map.get(&type_id) {
once_pipe.handle_event(&mut event);
}
event
}
}
pub struct EventBusBuilder {
valid_type_ids: HashSet<TypeId>,
}
impl EventBusBuilder {
#[inline]
pub fn new() -> Self {
Self {
valid_type_ids: HashSet::new(),
}
}
#[inline]
pub fn event_type<T>(&mut self) -> &mut Self
where
for<'a> T: 'a,
{
self.valid_type_ids.insert(TypeId::of::<T>());
self
}
#[inline]
pub fn build(&self) -> EventBus {
EventBus {
registry: EventBusRegistry {
valid_type_ids: self.valid_type_ids.clone(),
subscribers: RwLock::new(HashMap::new()),
once_pipes: RwLock::new(HashMap::new()),
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use smol::future::FutureExt;
use smol::{future, Timer};
use std::future::Future;
use std::time::Duration;
async fn timeout<T>(fut: impl Future<Output = T>, time: Duration) -> T {
let timeout_fn = async move {
Timer::after(time).await;
panic!("Timeout reached: {time:?}");
};
fut.or(timeout_fn).await
}
struct IncrementEvent(u32);
struct CancellableIncrementEvent(u32);
impl EventCancellable for CancellableIncrementEvent {}
#[derive(Default)]
struct IncrementHandler {}
impl EventHandler<IncrementEvent> for IncrementHandler {
fn handle_event(&self, event: &mut Event<IncrementEvent>) {
event.data().0 += 1;
}
}
fn create_event_bus() -> EventBus {
EventBus::builder()
.event_type::<IncrementEvent>()
.event_type::<CancellableIncrementEvent>()
.build()
}
#[test]
fn test_fire_event_no_subscriber() {
let event_bus = create_event_bus();
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 0);
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 0);
}
#[test]
fn test_fire_event_single_subscriber() {
let event_bus = create_event_bus();
event_bus.registry().register(IncrementHandler::default()).unwrap();
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 1);
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 1);
}
#[test]
fn test_fire_event_multiple_subscribers() {
let event_bus = create_event_bus();
for _ in 0..3 {
event_bus.registry().register(IncrementHandler::default()).unwrap();
}
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 3);
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 3);
}
#[test]
fn test_fire_event_multiple_types() {
struct Inc1Event(u32);
struct Inc2Event(u32);
struct Inc4Event(u32);
let event_bus = EventBus::builder()
.event_type::<Inc1Event>()
.event_type::<Inc2Event>()
.event_type::<Inc4Event>()
.build();
event_bus.registry().register(|event: &mut Event<Inc1Event>| { event.data().0 += 1; }).unwrap();
event_bus.registry().register(|event: &mut Event<Inc2Event>| { event.data().0 += 2; }).unwrap();
event_bus.registry().register(|event: &mut Event<Inc4Event>| { event.data().0 += 4; }).unwrap();
let result = event_bus.fire(Inc1Event(0));
assert_eq!(result.into_data().0, 1);
let result = event_bus.fire(Inc2Event(0));
assert_eq!(result.into_data().0, 2);
let result = event_bus.fire(Inc4Event(0));
assert_eq!(result.into_data().0, 4);
}
#[test]
fn test_fire_event_once_subscriber() {
let event_bus = create_event_bus();
let once_fut = event_bus.registry().register_once(|event: &mut Event<IncrementEvent>| {
event.data().0 += 1;
}).unwrap();
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 1);
future::block_on(timeout(once_fut, Duration::from_secs(1)))
.expect("SubscriberOnce should not error");
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 0);
}
#[test]
fn test_fire_event_subscriber_and_once() {
let event_bus = create_event_bus();
event_bus.registry().register(IncrementHandler::default()).unwrap();
let once_fut = event_bus.registry().register_once(|event: &mut Event<IncrementEvent>| {
event.data().0 += 1;
}).unwrap();
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 2);
future::block_on(timeout(once_fut, Duration::from_secs(1)))
.expect("SubscriberOnce should not error");
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 1);
}
#[test]
fn test_fire_event_cancelled() {
let event_bus = create_event_bus();
event_bus.registry().register(|event: &mut Event<CancellableIncrementEvent>| {
event.data().0 += 1;
}).unwrap();
event_bus.registry().register(|event: &mut Event<CancellableIncrementEvent>| {
event.data().0 += 1;
event.cancel();
}).unwrap();
event_bus.registry().register(|event: &mut Event<CancellableIncrementEvent>| {
event.data().0 += 1;
}).unwrap();
let result = event_bus.fire(CancellableIncrementEvent(0));
assert!(result.is_cancelled());
assert_eq!(result.into_data().0, 3);
let result = event_bus.fire(CancellableIncrementEvent(0));
assert!(result.is_cancelled());
assert_eq!(result.into_data().0, 3);
}
#[test]
fn test_fire_event_weak() {
let event_bus = create_event_bus();
event_bus.registry().register(IncrementHandler::default()).unwrap();
let handler2 = Arc::new(IncrementHandler::default());
event_bus.registry().register(handler2.clone()).unwrap();
let handler3 = Arc::new(IncrementHandler::default());
event_bus.registry().register(Arc::downgrade(&handler3)).unwrap();
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 3);
drop(handler3);
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 2);
drop(handler2);
let result = event_bus.fire(IncrementEvent(0));
assert_eq!(result.into_data().0, 2);
}
}