Skip to main content

sprite_core/
engine.rs

1use std::collections::HashMap;
2use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
3use std::sync::{Arc, Weak};
4use std::time::{Duration, Instant};
5use parking_lot::RwLock;
6use crossbeam_channel::{unbounded, Sender, Receiver, Select, TryRecvError};
7
8use crate::actor::{Context, StateStore};
9use crate::arena::Arena;
10use crate::message::Message;
11use crate::registry::Registry;
12use crate::error::SpriteError;
13
14#[derive(Clone)]
15pub struct Handle {
16    pub id: u64,
17    pub name: String,
18    pub(crate) tx: Sender<Message>,
19}
20
21impl Handle {
22    pub fn send(&self, msg: Message) {
23        let _ = self.tx.send(msg);
24    }
25    pub fn send_msg<T: crate::util::IntoMessage>(&self, msg: T) {
26        self.send(msg.into_message());
27    }
28    pub fn request(&self, msg: Message, timeout: Duration) -> Result<crate::request::Response, SpriteError> {
29        let (req, rx) = crate::request::Request::new(msg);
30        self.send(req.payload);
31        rx.recv_timeout(timeout)
32            .map_err(|_| SpriteError::RequestTimeout)
33    }
34}
35
36pub struct Engine {
37    pub(crate) inner: Arc<EngineInner>,
38}
39
40pub(crate) struct EngineInner {
41    pub(crate) next_id: AtomicU64,
42    pub(crate) registry: Arc<Registry>,
43    pub(crate) channels: RwLock<HashMap<u64, Sender<Message>>>,
44    pub(crate) running: AtomicBool,
45    workers: Vec<Sender<WorkerMsg>>,
46    next_worker: AtomicU64,
47}
48
49enum WorkerMsg {
50    Spawn(SpawnParams),
51    Shutdown,
52}
53
54struct SpawnParams {
55    id: u64,
56    name: String,
57    rx: Receiver<Message>,
58    tx: Sender<Message>,
59    setup: Arc<dyn Fn(&mut Context) + Send + Sync>,
60    state_store: StateStore,
61    arena_size: usize,
62    max_recoveries: u32,
63    recovery_window: Duration,
64    engine: Weak<EngineInner>,
65    registry: Arc<Registry>,
66}
67
68struct LocalActor {
69    id: u64,
70    name: String,
71    rx: Receiver<Message>,
72    tx: Sender<Message>,
73    state_store: StateStore,
74    setup: Arc<dyn Fn(&mut Context) + Send + Sync>,
75    arena_size: usize,
76    max_recoveries: u32,
77    recovery_window: Duration,
78    engine: Weak<EngineInner>,
79    registry: Arc<Registry>,
80    arena: Arena,
81    ctx: Option<Context>,
82    is_first_mount: bool,
83    recovery_count: u32,
84    last_recovery: Instant,
85}
86
87impl LocalActor {
88    fn new(params: SpawnParams) -> Self {
89        Self {
90            id: params.id,
91            name: params.name,
92            rx: params.rx,
93            tx: params.tx,
94            state_store: params.state_store,
95            setup: params.setup,
96            arena_size: params.arena_size,
97            max_recoveries: params.max_recoveries,
98            recovery_window: params.recovery_window,
99            engine: params.engine,
100            registry: params.registry,
101            arena: Arena::with_capacity(params.arena_size),
102            ctx: None,
103            is_first_mount: true,
104            recovery_count: 0,
105            last_recovery: Instant::now(),
106        }
107    }
108}
109
110impl EngineInner {
111    pub(crate) fn new() -> Self {
112        let num_workers = std::thread::available_parallelism()
113            .map(|n| n.get())
114            .unwrap_or(4);
115
116        let mut workers = Vec::with_capacity(num_workers);
117        for _ in 0..num_workers {
118            let (tx, rx) = unbounded::<WorkerMsg>();
119            std::thread::spawn(move || {
120                worker_loop(rx);
121            });
122            workers.push(tx);
123        }
124
125        Self {
126            next_id: AtomicU64::new(1),
127            registry: Arc::new(Registry::new()),
128            channels: RwLock::new(HashMap::new()),
129            running: AtomicBool::new(true),
130            workers,
131            next_worker: AtomicU64::new(0),
132        }
133    }
134
135    pub(crate) fn send_to(&self, id: u64, msg: Message) {
136        let channels = self.channels.read();
137        if let Some(tx) = channels.get(&id) {
138            let _ = tx.send(msg);
139        }
140    }
141
142    pub(crate) fn request(&self, id: u64, msg: Message, timeout: Duration) -> Option<Message> {
143        let channels = self.channels.read();
144        if let Some(tx) = channels.get(&id) {
145            let (req, rx) = crate::request::Request::new(msg);
146            let _ = tx.send(req.payload);
147            rx.recv_timeout(timeout).ok().map(|r| r.into_message())
148        } else {
149            None
150        }
151    }
152
153    pub(crate) fn spawn_simple<F>(&self, name: &str, setup: F) -> Handle
154    where F: Fn(&mut Context) + Send + Sync + 'static,
155    {
156        self.spawn(name, setup, 1024 * 64, 10, Duration::from_secs(5), Arc::new(self.clone_shallow()))
157    }
158
159    pub(crate) fn spawn<F>(
160        &self,
161        name: &str,
162        setup: F,
163        arena_size: usize,
164        max_recoveries: u32,
165        recovery_window: Duration,
166        engine_arc: Arc<EngineInner>,
167    ) -> Handle
168    where
169        F: Fn(&mut Context) + Send + Sync + 'static,
170    {
171        let id = self.next_id.fetch_add(1, Ordering::SeqCst);
172        let (tx, rx) = unbounded();
173
174        {
175            let mut channels = self.channels.write();
176            channels.insert(id, tx.clone());
177        }
178        self.registry.register(name, id);
179
180        let state_store: StateStore = Arc::new(RwLock::new(HashMap::new()));
181        let setup = Arc::new(setup);
182        let name_owned = name.to_string();
183        let tx_for_handle = tx.clone();
184        let registry = self.registry.clone();
185        let engine_weak = Arc::downgrade(&engine_arc);
186
187        let params = SpawnParams {
188            id,
189            name: name_owned,
190            rx,
191            tx: tx.clone(),
192            setup,
193            state_store,
194            arena_size,
195            max_recoveries,
196            recovery_window,
197            engine: engine_weak,
198            registry,
199        };
200
201        let worker_idx = (self.next_worker.fetch_add(1, Ordering::Relaxed) as usize) % self.workers.len();
202        let _ = self.workers[worker_idx].send(WorkerMsg::Spawn(params));
203
204        Handle { id, name: name.to_string(), tx: tx_for_handle }
205    }
206
207    fn clone_shallow(&self) -> Self {
208        Self {
209            next_id: AtomicU64::new(self.next_id.load(Ordering::SeqCst)),
210            registry: self.registry.clone(),
211            channels: RwLock::new(self.channels.read().clone()),
212            running: AtomicBool::new(self.running.load(Ordering::SeqCst)),
213            workers: self.workers.clone(),
214            next_worker: AtomicU64::new(self.next_worker.load(Ordering::SeqCst)),
215        }
216    }
217}
218
219fn worker_loop(ctrl_rx: Receiver<WorkerMsg>) {
220    let mut actors: HashMap<u64, Box<LocalActor>> = HashMap::new();
221
222    'outer: loop {
223        let ids: Vec<u64> = actors.keys().cloned().collect();
224
225        // Use raw pointers to rx so Select can hold references without
226        // borrowing actors mutably. Box guarantees stable addresses.
227        let rx_ptrs: Vec<*const Receiver<Message>> = ids.iter()
228            .map(|id| &actors.get(id).unwrap().rx as *const _)
229            .collect();
230
231        let mut sel = Select::new();
232        let ctrl_idx = sel.recv(&ctrl_rx);
233        for &rx in &rx_ptrs {
234            sel.recv(unsafe { &*rx });
235        }
236
237        let oper = sel.select();
238        let idx = oper.index();
239
240        if idx == ctrl_idx {
241            match oper.recv(&ctrl_rx) {
242                Ok(WorkerMsg::Spawn(params)) => {
243                    let actor = Box::new(LocalActor::new(params));
244                    let id = actor.id;
245                    actors.insert(id, actor);
246                }
247                Ok(WorkerMsg::Shutdown) => {
248                    for (_, actor) in actors.iter_mut() {
249                        if let Some(ref ctx) = actor.ctx {
250                            let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
251                                if let Some(ref h) = ctx.unmount_handler {
252                                    h();
253                                }
254                            }));
255                        }
256                        actor.registry.unregister(&actor.name);
257                    }
258                    break 'outer;
259                }
260                Err(_) => break 'outer,
261            }
262        } else {
263            let id = ids[idx - 1];
264            let actor = actors.get_mut(&id).unwrap();
265            match oper.recv(&actor.rx) {
266                Ok(msg) => {
267                    if run_actor_batch(&mut **actor, Some(msg)).is_err() {
268                        actors.remove(&id);
269                    }
270                }
271                Err(_) => {
272                    actors.remove(&id);
273                }
274            }
275        }
276    }
277}
278
279fn run_actor_batch(actor: &mut LocalActor, first_msg: Option<Message>) -> Result<(), ()> {
280    if actor.ctx.is_none() {
281        let engine_ref = match actor.engine.upgrade() {
282            Some(arc) => arc,
283            None => return Err(()),
284        };
285
286        let mut ctx = Context::new(
287            actor.id,
288            actor.name.clone(),
289            actor.state_store.clone(),
290            actor.rx.clone(),
291            actor.tx.clone(),
292            engine_ref,
293        );
294
295        let setup = actor.setup.clone();
296        let is_first = actor.is_first_mount;
297        let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
298            setup(&mut ctx);
299            if is_first {
300                if let Some(ref h) = ctx.mount_handler {
301                    h();
302                }
303            }
304        }));
305
306        match result {
307            Ok(()) => {
308                actor.is_first_mount = false;
309                actor.ctx = Some(ctx);
310            }
311            Err(_) => {
312                return handle_actor_panic(actor);
313            }
314        }
315    }
316
317    let ctx = match actor.ctx.as_mut() {
318        Some(c) => c,
319        None => return Err(()),
320    };
321
322    if !ctx.engine.running.load(Ordering::SeqCst) {
323        let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
324            if let Some(ref h) = ctx.unmount_handler {
325                h();
326            }
327        }));
328        actor.registry.unregister(&actor.name);
329        return Err(());
330    }
331
332    if let Some(msg) = first_msg {
333        ctx.metrics.inc_received();
334        if let Some(ref handler) = ctx.message_handler {
335            let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
336                handler(msg);
337            }));
338            if result.is_err() {
339                return handle_actor_panic(actor);
340            }
341        }
342    }
343
344    let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
345        loop {
346            match ctx.rx.try_recv() {
347                Ok(msg) => {
348                    ctx.metrics.inc_received();
349                    if let Some(ref handler) = ctx.message_handler {
350                        handler(msg);
351                    }
352                }
353                Err(TryRecvError::Empty) => break,
354                Err(TryRecvError::Disconnected) => break,
355            }
356        }
357    }));
358
359    match result {
360        Ok(()) => {
361            if !ctx.engine.running.load(Ordering::SeqCst) {
362                let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
363                    if let Some(ref h) = ctx.unmount_handler {
364                        h();
365                    }
366                }));
367                actor.registry.unregister(&actor.name);
368                return Err(());
369            }
370            Ok(())
371        }
372        Err(_) => handle_actor_panic(actor),
373    }
374}
375
376fn handle_actor_panic(actor: &mut LocalActor) -> Result<(), ()> {
377    actor.recovery_count += 1;
378    if actor.recovery_count > actor.max_recoveries && actor.last_recovery.elapsed() < actor.recovery_window {
379        if let Some(ref ctx) = actor.ctx {
380            let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
381                if let Some(ref h) = ctx.unmount_handler {
382                    h();
383                }
384            }));
385        }
386        actor.registry.unregister(&actor.name);
387        tracing::error!(
388            "[Actor {}] CIRCUIT BREAKER TRIPPED after {} recoveries — halting.",
389            actor.id, actor.recovery_count
390        );
391        return Err(());
392    }
393
394    actor.last_recovery = Instant::now();
395    let start = Instant::now();
396    actor.arena.reset();
397    let elapsed = start.elapsed();
398
399    if let Some(ref ctx) = actor.ctx {
400        ctx.metrics.inc_panic();
401        ctx.metrics.inc_recovery();
402        tracing::debug!("[Actor {}] recovered in {:?}", actor.id, elapsed);
403        if let Some(ref h) = ctx.panic_handler {
404            let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| h()));
405        }
406    }
407
408    actor.ctx = None;
409    Ok(())
410}
411
412impl Engine {
413    pub fn new() -> Self {
414        Self { inner: Arc::new(EngineInner::new()) }
415    }
416
417    pub fn spawn<F>(&self, name: &str, setup: F) -> Handle
418    where F: Fn(&mut Context) + Send + Sync + 'static,
419    {
420        self.inner.spawn(name, setup, 1024 * 64, 10, Duration::from_secs(5), self.inner.clone())
421    }
422
423    pub fn spawn_with_config<F>(
424        &self, name: &str, setup: F,
425        arena_size: usize, max_recoveries: u32, recovery_window: Duration,
426    ) -> Handle
427    where F: Fn(&mut Context) + Send + Sync + 'static,
428    {
429        self.inner.spawn(name, setup, arena_size, max_recoveries, recovery_window, self.inner.clone())
430    }
431
432    pub fn send_to(&self, id: u64, msg: Message) {
433        self.inner.send_to(id, msg);
434    }
435
436    pub fn send_named(&self, name: &str, msg: Message) {
437        if let Some(id) = self.inner.registry.lookup(name) {
438            self.inner.send_to(id, msg);
439        }
440    }
441
442    pub fn lookup(&self, name: &str) -> Option<u64> {
443        self.inner.registry.lookup(name)
444    }
445
446    pub fn broadcast(&self, msg: Message) -> usize {
447        let channels = self.inner.channels.read();
448        let mut sent = 0;
449        for (_, tx) in channels.iter() {
450            if tx.send(msg.clone()).is_ok() { sent += 1; }
451        }
452        sent
453    }
454
455    pub fn shutdown(&self) {
456        self.inner.running.store(false, Ordering::SeqCst);
457        for worker in &self.inner.workers {
458            let _ = worker.send(WorkerMsg::Shutdown);
459        }
460    }
461
462    pub fn is_running(&self) -> bool {
463        self.inner.running.load(Ordering::SeqCst)
464    }
465
466    pub fn actor_count(&self) -> usize {
467        self.inner.channels.read().len()
468    }
469}
470
471#[cfg(test)]
472mod tests {
473    use super::*;
474    use std::time::Duration;
475
476    #[test]
477    fn spawn_and_send() {
478        let engine = Engine::new();
479        let handle = engine.spawn("test", |ctx| {
480            ctx.on_message(|msg| { println!("got: {:?}", msg); });
481        });
482        std::thread::sleep(Duration::from_millis(20));
483        handle.send(Message::text("hello"));
484        std::thread::sleep(Duration::from_millis(50));
485    }
486
487    #[test]
488    fn state_persists_across_panics() {
489        let engine = Engine::new();
490        let handle = engine.spawn("fragile", |ctx| {
491            let count = ctx.use_state("count", 0i64);
492            ctx.on_message(move |msg| {
493                if msg == "set" { count.set(42); }
494                if msg == "panic" { panic!("boom"); }
495                if msg == "check" { assert_eq!(count.get(), 42); }
496            });
497        });
498        std::thread::sleep(Duration::from_millis(20));
499        handle.send(Message::text("set"));
500        std::thread::sleep(Duration::from_millis(20));
501        handle.send(Message::text("panic"));
502        std::thread::sleep(Duration::from_millis(50));
503        handle.send(Message::text("check"));
504        std::thread::sleep(Duration::from_millis(50));
505    }
506
507    #[test]
508    fn named_lookup() {
509        let engine = Engine::new();
510        let h = engine.spawn("logger", |ctx| {
511            ctx.on_message(|msg| println!("{:?}", msg));
512        });
513        std::thread::sleep(Duration::from_millis(10));
514        assert_eq!(engine.lookup("logger"), Some(h.id));
515        engine.send_named("logger", Message::text("hi"));
516        std::thread::sleep(Duration::from_millis(50));
517    }
518
519    #[test]
520    fn broadcast_works() {
521        let engine = Engine::new();
522        let _ = engine.spawn("a", |ctx| {
523            ctx.on_message(|msg| println!("a: {:?}", msg));
524        });
525        let _ = engine.spawn("b", |ctx| {
526            ctx.on_message(|msg| println!("b: {:?}", msg));
527        });
528        std::thread::sleep(Duration::from_millis(20));
529        let sent = engine.broadcast(Message::text("all"));
530        assert_eq!(sent, 2);
531        std::thread::sleep(Duration::from_millis(50));
532    }
533
534    #[test]
535    fn mount_and_unmount() {
536        let engine = Engine::new();
537        let handle = engine.spawn("lifecycle", |ctx| {
538            ctx.on_mount(|| println!("mounted"));
539            ctx.on_unmount(|| println!("unmounted"));
540            ctx.on_message(|_| {});
541        });
542        std::thread::sleep(Duration::from_millis(20));
543        handle.send(Message::text("hi"));
544        std::thread::sleep(Duration::from_millis(50));
545    }
546}
547