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, LocalActor> = HashMap::new();
221
222 'outer: loop {
223 let selected = {
224 let mut sel = Select::new();
225 let _ctrl_idx = sel.recv(&ctrl_rx);
226
227 let ids: Vec<u64> = actors.keys().cloned().collect();
228 for id in &ids {
229 let actor = actors.get(id).unwrap();
230 sel.recv(&actor.rx);
231 }
232
233 let oper = sel.select();
234 let idx = oper.index();
235
236 if idx == 0 {
237 None
238 } else {
239 let id = ids[idx - 1];
240 let actor = actors.get(&id).unwrap();
241 let msg = oper.recv(&actor.rx).unwrap();
242 Some((id, msg))
243 }
244 };
245
246 match selected {
247 None => {
248 match ctrl_rx.recv() {
249 Ok(WorkerMsg::Spawn(params)) => {
250 let actor = LocalActor::new(params);
251 actors.insert(actor.id, actor);
252 }
253 Ok(WorkerMsg::Shutdown) => {
254 for (_, actor) in actors.iter_mut() {
255 if let Some(ref ctx) = actor.ctx {
256 let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
257 if let Some(ref h) = ctx.unmount_handler {
258 h();
259 }
260 }));
261 }
262 actor.registry.unregister(&actor.name);
263 }
264 break 'outer;
265 }
266 Err(_) => break 'outer,
267 }
268 }
269 Some((id, msg)) => {
270 let actor = actors.get_mut(&id).unwrap();
271 if run_actor_batch(actor, Some(msg)).is_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