use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Weak};
use std::time::{Duration, Instant};
use parking_lot::RwLock;
use crossbeam_channel::{unbounded, Sender, Receiver, Select, TryRecvError};
use crate::actor::{Context, StateStore};
use crate::arena::Arena;
use crate::message::Message;
use crate::registry::Registry;
use crate::error::SpriteError;
#[derive(Clone)]
pub struct Handle {
pub id: u64,
pub name: String,
pub(crate) tx: Sender<Message>,
}
impl Handle {
pub fn send(&self, msg: Message) {
let _ = self.tx.send(msg);
}
pub fn send_msg<T: crate::util::IntoMessage>(&self, msg: T) {
self.send(msg.into_message());
}
pub fn request(&self, msg: Message, timeout: Duration) -> Result<crate::request::Response, SpriteError> {
let (req, rx) = crate::request::Request::new(msg);
self.send(req.payload);
rx.recv_timeout(timeout)
.map_err(|_| SpriteError::RequestTimeout)
}
}
pub struct Engine {
pub(crate) inner: Arc<EngineInner>,
}
pub(crate) struct EngineInner {
pub(crate) next_id: AtomicU64,
pub(crate) registry: Arc<Registry>,
pub(crate) channels: RwLock<HashMap<u64, Sender<Message>>>,
pub(crate) running: AtomicBool,
workers: Vec<Sender<WorkerMsg>>,
next_worker: AtomicU64,
}
enum WorkerMsg {
Spawn(SpawnParams),
Shutdown,
}
struct SpawnParams {
id: u64,
name: String,
rx: Receiver<Message>,
tx: Sender<Message>,
setup: Arc<dyn Fn(&mut Context) + Send + Sync>,
state_store: StateStore,
arena_size: usize,
max_recoveries: u32,
recovery_window: Duration,
engine: Weak<EngineInner>,
registry: Arc<Registry>,
}
struct LocalActor {
id: u64,
name: String,
rx: Receiver<Message>,
tx: Sender<Message>,
state_store: StateStore,
setup: Arc<dyn Fn(&mut Context) + Send + Sync>,
arena_size: usize,
max_recoveries: u32,
recovery_window: Duration,
engine: Weak<EngineInner>,
registry: Arc<Registry>,
arena: Arena,
ctx: Option<Context>,
is_first_mount: bool,
recovery_count: u32,
last_recovery: Instant,
}
impl LocalActor {
fn new(params: SpawnParams) -> Self {
Self {
id: params.id,
name: params.name,
rx: params.rx,
tx: params.tx,
state_store: params.state_store,
setup: params.setup,
arena_size: params.arena_size,
max_recoveries: params.max_recoveries,
recovery_window: params.recovery_window,
engine: params.engine,
registry: params.registry,
arena: Arena::with_capacity(params.arena_size),
ctx: None,
is_first_mount: true,
recovery_count: 0,
last_recovery: Instant::now(),
}
}
}
impl EngineInner {
pub(crate) fn new() -> Self {
let num_workers = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4);
let mut workers = Vec::with_capacity(num_workers);
for _ in 0..num_workers {
let (tx, rx) = unbounded::<WorkerMsg>();
std::thread::spawn(move || {
worker_loop(rx);
});
workers.push(tx);
}
Self {
next_id: AtomicU64::new(1),
registry: Arc::new(Registry::new()),
channels: RwLock::new(HashMap::new()),
running: AtomicBool::new(true),
workers,
next_worker: AtomicU64::new(0),
}
}
pub(crate) fn send_to(&self, id: u64, msg: Message) {
let channels = self.channels.read();
if let Some(tx) = channels.get(&id) {
let _ = tx.send(msg);
}
}
pub(crate) fn request(&self, id: u64, msg: Message, timeout: Duration) -> Option<Message> {
let channels = self.channels.read();
if let Some(tx) = channels.get(&id) {
let (req, rx) = crate::request::Request::new(msg);
let _ = tx.send(req.payload);
rx.recv_timeout(timeout).ok().map(|r| r.into_message())
} else {
None
}
}
pub(crate) fn spawn_simple<F>(&self, name: &str, setup: F) -> Handle
where F: Fn(&mut Context) + Send + Sync + 'static,
{
self.spawn(name, setup, 1024 * 64, 10, Duration::from_secs(5), Arc::new(self.clone_shallow()))
}
pub(crate) fn spawn<F>(
&self,
name: &str,
setup: F,
arena_size: usize,
max_recoveries: u32,
recovery_window: Duration,
engine_arc: Arc<EngineInner>,
) -> Handle
where
F: Fn(&mut Context) + Send + Sync + 'static,
{
let id = self.next_id.fetch_add(1, Ordering::SeqCst);
let (tx, rx) = unbounded();
{
let mut channels = self.channels.write();
channels.insert(id, tx.clone());
}
self.registry.register(name, id);
let state_store: StateStore = Arc::new(RwLock::new(HashMap::new()));
let setup = Arc::new(setup);
let name_owned = name.to_string();
let tx_for_handle = tx.clone();
let registry = self.registry.clone();
let engine_weak = Arc::downgrade(&engine_arc);
let params = SpawnParams {
id,
name: name_owned,
rx,
tx: tx.clone(),
setup,
state_store,
arena_size,
max_recoveries,
recovery_window,
engine: engine_weak,
registry,
};
let worker_idx = (self.next_worker.fetch_add(1, Ordering::Relaxed) as usize) % self.workers.len();
let _ = self.workers[worker_idx].send(WorkerMsg::Spawn(params));
Handle { id, name: name.to_string(), tx: tx_for_handle }
}
fn clone_shallow(&self) -> Self {
Self {
next_id: AtomicU64::new(self.next_id.load(Ordering::SeqCst)),
registry: self.registry.clone(),
channels: RwLock::new(self.channels.read().clone()),
running: AtomicBool::new(self.running.load(Ordering::SeqCst)),
workers: self.workers.clone(),
next_worker: AtomicU64::new(self.next_worker.load(Ordering::SeqCst)),
}
}
}
fn worker_loop(ctrl_rx: Receiver<WorkerMsg>) {
let mut actors: HashMap<u64, Box<LocalActor>> = HashMap::new();
'outer: loop {
let ids: Vec<u64> = actors.keys().cloned().collect();
let rx_ptrs: Vec<*const Receiver<Message>> = ids.iter()
.map(|id| &actors.get(id).unwrap().rx as *const _)
.collect();
let mut sel = Select::new();
let ctrl_idx = sel.recv(&ctrl_rx);
for &rx in &rx_ptrs {
sel.recv(unsafe { &*rx });
}
let oper = sel.select();
let idx = oper.index();
if idx == ctrl_idx {
match oper.recv(&ctrl_rx) {
Ok(WorkerMsg::Spawn(params)) => {
let actor = Box::new(LocalActor::new(params));
let id = actor.id;
actors.insert(id, actor);
}
Ok(WorkerMsg::Shutdown) => {
for (_, actor) in actors.iter_mut() {
if let Some(ref ctx) = actor.ctx {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
if let Some(ref h) = ctx.unmount_handler {
h();
}
}));
}
actor.registry.unregister(&actor.name);
}
break 'outer;
}
Err(_) => break 'outer,
}
} else {
let id = ids[idx - 1];
let actor = actors.get_mut(&id).unwrap();
match oper.recv(&actor.rx) {
Ok(msg) => {
if run_actor_batch(&mut **actor, Some(msg)).is_err() {
actors.remove(&id);
}
}
Err(_) => {
actors.remove(&id);
}
}
}
}
}
fn run_actor_batch(actor: &mut LocalActor, first_msg: Option<Message>) -> Result<(), ()> {
if actor.ctx.is_none() {
let engine_ref = match actor.engine.upgrade() {
Some(arc) => arc,
None => return Err(()),
};
let mut ctx = Context::new(
actor.id,
actor.name.clone(),
actor.state_store.clone(),
actor.rx.clone(),
actor.tx.clone(),
engine_ref,
);
let setup = actor.setup.clone();
let is_first = actor.is_first_mount;
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
setup(&mut ctx);
if is_first {
if let Some(ref h) = ctx.mount_handler {
h();
}
}
}));
match result {
Ok(()) => {
actor.is_first_mount = false;
actor.ctx = Some(ctx);
}
Err(_) => {
return handle_actor_panic(actor);
}
}
}
let ctx = match actor.ctx.as_mut() {
Some(c) => c,
None => return Err(()),
};
if !ctx.engine.running.load(Ordering::SeqCst) {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
if let Some(ref h) = ctx.unmount_handler {
h();
}
}));
actor.registry.unregister(&actor.name);
return Err(());
}
if let Some(msg) = first_msg {
ctx.metrics.inc_received();
if let Some(ref handler) = ctx.message_handler {
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
handler(msg);
}));
if result.is_err() {
return handle_actor_panic(actor);
}
}
}
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
loop {
match ctx.rx.try_recv() {
Ok(msg) => {
ctx.metrics.inc_received();
if let Some(ref handler) = ctx.message_handler {
handler(msg);
}
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Disconnected) => break,
}
}
}));
match result {
Ok(()) => {
if !ctx.engine.running.load(Ordering::SeqCst) {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
if let Some(ref h) = ctx.unmount_handler {
h();
}
}));
actor.registry.unregister(&actor.name);
return Err(());
}
Ok(())
}
Err(_) => handle_actor_panic(actor),
}
}
fn handle_actor_panic(actor: &mut LocalActor) -> Result<(), ()> {
actor.recovery_count += 1;
if actor.recovery_count > actor.max_recoveries && actor.last_recovery.elapsed() < actor.recovery_window {
if let Some(ref ctx) = actor.ctx {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
if let Some(ref h) = ctx.unmount_handler {
h();
}
}));
}
actor.registry.unregister(&actor.name);
tracing::error!(
"[Actor {}] CIRCUIT BREAKER TRIPPED after {} recoveries — halting.",
actor.id, actor.recovery_count
);
return Err(());
}
actor.last_recovery = Instant::now();
let start = Instant::now();
actor.arena.reset();
let elapsed = start.elapsed();
if let Some(ref ctx) = actor.ctx {
ctx.metrics.inc_panic();
ctx.metrics.inc_recovery();
tracing::debug!("[Actor {}] recovered in {:?}", actor.id, elapsed);
if let Some(ref h) = ctx.panic_handler {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| h()));
}
}
actor.ctx = None;
Ok(())
}
impl Engine {
pub fn new() -> Self {
Self { inner: Arc::new(EngineInner::new()) }
}
pub fn spawn<F>(&self, name: &str, setup: F) -> Handle
where F: Fn(&mut Context) + Send + Sync + 'static,
{
self.inner.spawn(name, setup, 1024 * 64, 10, Duration::from_secs(5), self.inner.clone())
}
pub fn spawn_with_config<F>(
&self, name: &str, setup: F,
arena_size: usize, max_recoveries: u32, recovery_window: Duration,
) -> Handle
where F: Fn(&mut Context) + Send + Sync + 'static,
{
self.inner.spawn(name, setup, arena_size, max_recoveries, recovery_window, self.inner.clone())
}
pub fn send_to(&self, id: u64, msg: Message) {
self.inner.send_to(id, msg);
}
pub fn send_named(&self, name: &str, msg: Message) {
if let Some(id) = self.inner.registry.lookup(name) {
self.inner.send_to(id, msg);
}
}
pub fn lookup(&self, name: &str) -> Option<u64> {
self.inner.registry.lookup(name)
}
pub fn broadcast(&self, msg: Message) -> usize {
let channels = self.inner.channels.read();
let mut sent = 0;
for (_, tx) in channels.iter() {
if tx.send(msg.clone()).is_ok() { sent += 1; }
}
sent
}
pub fn shutdown(&self) {
self.inner.running.store(false, Ordering::SeqCst);
for worker in &self.inner.workers {
let _ = worker.send(WorkerMsg::Shutdown);
}
}
pub fn is_running(&self) -> bool {
self.inner.running.load(Ordering::SeqCst)
}
pub fn actor_count(&self) -> usize {
self.inner.channels.read().len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn spawn_and_send() {
let engine = Engine::new();
let handle = engine.spawn("test", |ctx| {
ctx.on_message(|msg| { println!("got: {:?}", msg); });
});
std::thread::sleep(Duration::from_millis(20));
handle.send(Message::text("hello"));
std::thread::sleep(Duration::from_millis(50));
}
#[test]
fn state_persists_across_panics() {
let engine = Engine::new();
let handle = engine.spawn("fragile", |ctx| {
let count = ctx.use_state("count", 0i64);
ctx.on_message(move |msg| {
if msg == "set" { count.set(42); }
if msg == "panic" { panic!("boom"); }
if msg == "check" { assert_eq!(count.get(), 42); }
});
});
std::thread::sleep(Duration::from_millis(20));
handle.send(Message::text("set"));
std::thread::sleep(Duration::from_millis(20));
handle.send(Message::text("panic"));
std::thread::sleep(Duration::from_millis(50));
handle.send(Message::text("check"));
std::thread::sleep(Duration::from_millis(50));
}
#[test]
fn named_lookup() {
let engine = Engine::new();
let h = engine.spawn("logger", |ctx| {
ctx.on_message(|msg| println!("{:?}", msg));
});
std::thread::sleep(Duration::from_millis(10));
assert_eq!(engine.lookup("logger"), Some(h.id));
engine.send_named("logger", Message::text("hi"));
std::thread::sleep(Duration::from_millis(50));
}
#[test]
fn broadcast_works() {
let engine = Engine::new();
let _ = engine.spawn("a", |ctx| {
ctx.on_message(|msg| println!("a: {:?}", msg));
});
let _ = engine.spawn("b", |ctx| {
ctx.on_message(|msg| println!("b: {:?}", msg));
});
std::thread::sleep(Duration::from_millis(20));
let sent = engine.broadcast(Message::text("all"));
assert_eq!(sent, 2);
std::thread::sleep(Duration::from_millis(50));
}
#[test]
fn mount_and_unmount() {
let engine = Engine::new();
let handle = engine.spawn("lifecycle", |ctx| {
ctx.on_mount(|| println!("mounted"));
ctx.on_unmount(|| println!("unmounted"));
ctx.on_message(|_| {});
});
std::thread::sleep(Duration::from_millis(20));
handle.send(Message::text("hi"));
std::thread::sleep(Duration::from_millis(50));
}
}