use crate::error::Result;
use std::any::{Any, TypeId};
use std::collections::HashMap;
pub trait Event: Send + Sync + 'static {
fn event_type_id(&self) -> TypeId;
fn as_any(&self) -> &dyn Any;
fn event_name(&self) -> &str {
"UnnamedEvent"
}
fn validate(&self) -> Result<()> {
Ok(())
}
}
pub trait EventSubscriber: Send + Sync {
fn on_event(&mut self, event: &dyn Event) -> Result<()>;
fn name(&self) -> &str {
"UnnamedSubscriber"
}
fn can_handle(&self, _event_type: TypeId) -> bool {
true }
}
trait EventStorage: Any + Send + Sync {
#[allow(dead_code)]
fn as_any(&self) -> &dyn Any;
fn as_any_mut(&mut self) -> &mut dyn Any;
fn process(&mut self, subscribers: &mut [Box<dyn EventSubscriber>]) -> Result<()>;
fn clear(&mut self);
fn len(&self) -> usize;
}
struct TypedEventQueue<T: Event> {
events: Vec<T>,
}
impl<T: Event> EventStorage for TypedEventQueue<T> {
fn as_any(&self) -> &dyn Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn Any {
self
}
fn process(&mut self, subscribers: &mut [Box<dyn EventSubscriber>]) -> Result<()> {
for event in &self.events {
for subscriber in subscribers.iter_mut() {
subscriber.on_event(event)?;
}
}
Ok(())
}
fn clear(&mut self) {
self.events.clear();
}
fn len(&self) -> usize {
self.events.len()
}
}
pub struct EventBus {
queues: HashMap<TypeId, Box<dyn EventStorage>>,
subscribers: HashMap<TypeId, Vec<Box<dyn EventSubscriber>>>,
processed_count: u64,
}
impl EventBus {
pub fn new() -> Self {
Self {
queues: HashMap::new(),
subscribers: HashMap::new(),
processed_count: 0,
}
}
pub fn subscribe<E: Event + 'static>(&mut self, subscriber: Box<dyn EventSubscriber>) {
let type_id = TypeId::of::<E>();
self.subscribers
.entry(type_id)
.or_default()
.push(subscriber);
}
pub fn subscribe_all(&mut self, subscriber: Box<dyn EventSubscriber>) {
let type_id = TypeId::of::<()>(); self.subscribers
.entry(type_id)
.or_default()
.push(subscriber);
}
pub fn publish_event<E: Event + 'static>(&mut self, event: E) -> Result<()> {
let type_id = TypeId::of::<E>();
let queue = self
.queues
.entry(type_id)
.or_insert_with(|| Box::new(TypedEventQueue::<E> { events: Vec::new() }));
let typed_queue = queue
.as_any_mut()
.downcast_mut::<TypedEventQueue<E>>()
.unwrap();
typed_queue.events.push(event);
Ok(())
}
pub fn process_events(&mut self) -> Result<()> {
let wildcard_type = TypeId::of::<()>();
for (type_id, queue) in self.queues.iter_mut() {
if let Some(subs) = self.subscribers.get_mut(type_id) {
queue.process(subs)?;
}
if let Some(wildcard_subs) = self.subscribers.get_mut(&wildcard_type) {
queue.process(wildcard_subs)?;
}
self.processed_count += queue.len() as u64;
queue.clear();
}
Ok(())
}
pub fn queue_size(&self) -> usize {
self.queues.values().map(|q| q.len()).sum()
}
pub fn processed_count(&self) -> u64 {
self.processed_count
}
pub fn clear_queue(&mut self) {
for queue in self.queues.values_mut() {
queue.clear();
}
}
pub fn subscriber_count(&self, event_type: TypeId) -> usize {
self.subscribers
.get(&event_type)
.map(|subs| subs.len())
.unwrap_or(0)
}
pub fn total_subscribers(&self) -> usize {
self.subscribers.values().map(|subs| subs.len()).sum()
}
pub fn clear_subscribers(&mut self) {
self.subscribers.clear();
}
}
impl Default for EventBus {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{Arc, Mutex};
struct TestEvent;
impl Event for TestEvent {
fn event_type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn as_any(&self) -> &dyn Any {
self
}
}
struct TestSubscriber {
call_count: Arc<Mutex<usize>>,
}
impl EventSubscriber for TestSubscriber {
fn on_event(&mut self, _event: &dyn Event) -> Result<()> {
*self.call_count.lock().unwrap() += 1;
Ok(())
}
}
#[test]
fn test_publish_and_process() {
let mut bus = EventBus::new();
let count = Arc::new(Mutex::new(0));
let subscriber = TestSubscriber {
call_count: count.clone(),
};
bus.subscribe::<TestEvent>(Box::new(subscriber));
bus.publish_event(TestEvent).unwrap();
assert_eq!(bus.queue_size(), 1);
bus.process_events().unwrap();
assert_eq!(*count.lock().unwrap(), 1);
}
#[test]
fn test_multiple_subscribers() {
let mut bus = EventBus::new();
let count1 = Arc::new(Mutex::new(0));
let count2 = Arc::new(Mutex::new(0));
bus.subscribe::<TestEvent>(Box::new(TestSubscriber {
call_count: count1.clone(),
}));
bus.subscribe::<TestEvent>(Box::new(TestSubscriber {
call_count: count2.clone(),
}));
bus.publish_event(TestEvent).unwrap();
bus.process_events().unwrap();
assert_eq!(*count1.lock().unwrap(), 1);
assert_eq!(*count2.lock().unwrap(), 1);
}
}