1use core::cell::UnsafeCell;
2use std::{
3 any::{Any, TypeId},
4 collections::HashMap,
5 rc::Rc,
6};
7
8use cubecl::prelude::*;
9use cubecl_core::{self as cubecl, intrinsic};
10
11#[derive(CubeType, Clone)]
12pub struct ComptimeEventBus {
19 #[allow(unused)]
20 #[cube(comptime)]
21 listener_family: Rc<UnsafeCell<HashMap<TypeId, Vec<EventItem>>>>,
22}
23
24type EventItem = Box<dyn Any>;
25type Call<E> = Box<
26 dyn FnMut(&Scope, &mut Box<dyn Any>, <E as CubeType>::ExpandType, &mut ComptimeEventBusExpand),
27>;
28
29struct Payload<E: CubeType> {
30 listener: Box<dyn Any>,
31 call: Call<E>,
32}
33
34impl Default for ComptimeEventBus {
35 fn default() -> Self {
36 Self::new()
37 }
38}
39
40#[cube]
41impl ComptimeEventBus {
42 pub fn new() -> Self {
44 intrinsic!(|_| {
45 ComptimeEventBusExpand {
46 listener_family: Rc::new(UnsafeCell::new(HashMap::new())),
47 }
48 })
49 }
50
51 pub fn listener<L: EventListener>(&mut self, listener: L) {
58 intrinsic!(|_| {
59 let type_id = TypeId::of::<L::Event>();
60 let mut listeners = unsafe { self.listener_family.get().as_mut().unwrap() };
61
62 let call =
67 |scope: &Scope,
68 listener: &mut Box<dyn Any>,
69 event: <L::Event as cubecl::prelude::CubeType>::ExpandType,
70 bus: &mut <ComptimeEventBus as cubecl::prelude::CubeType>::ExpandType| {
71 let listener: &mut L::ExpandType = listener.downcast_mut().unwrap();
72 listener.__expand_on_event_method(scope, event, bus)
73 };
74 let call: Call<L::Event> = Box::new(call);
75
76 let listener: Box<dyn Any> = Box::new(listener);
77 let payload = Payload::<L::Event> { listener, call };
78
79 let listener_dyn: Box<dyn Any> = Box::new(payload);
82
83 match listeners.get_mut(&type_id) {
84 Some(list) => list.push(listener_dyn),
85 None => {
86 listeners.insert(type_id, vec![listener_dyn]);
87 }
88 }
89 })
90 }
91
92 pub fn event<E: CubeType + 'static>(&mut self, event: E) {
94 intrinsic!(|scope| {
95 let type_id = TypeId::of::<E>();
96 let family = self.listener_family.clone();
97 let family = unsafe { family.get().as_mut().unwrap() };
98 let listeners = match family.get_mut(&type_id) {
99 Some(val) => val,
100 None => return,
101 };
102
103 for listener in listeners.iter_mut() {
104 let mut payload = listener.downcast_mut::<Payload<E>>().unwrap();
105 let call = &mut payload.call;
106 call(scope, &mut payload.listener, event.clone_unchecked(), self);
107 }
108 })
109 }
110}
111
112#[cube]
113pub trait EventListener: 'static {
116 type Event: CubeType + 'static;
118
119 fn on_event(&mut self, event: Self::Event, bus: &mut ComptimeEventBus);
122}