use crossbeam_channel::{Receiver as CbReceiver, Sender as CbSender, unbounded};
use std::sync::mpsc::{Receiver, TryRecvError, channel};
use crate::ecs::plugin::Plugin;
type Job = Box<dyn FnOnce() + Send + 'static>;
fn panic_message(payload: Box<dyn std::any::Any + Send>) -> String {
if let Some(s) = payload.downcast_ref::<&str>() {
(*s).to_string()
} else if let Some(s) = payload.downcast_ref::<String>() {
s.clone()
} else {
"<panic payload was not a string>".to_string()
}
}
#[cfg(not(target_arch = "wasm32"))]
pub trait SpawnableFuture<T>: std::future::Future<Output = T> + Send + 'static {}
#[cfg(not(target_arch = "wasm32"))]
impl<T, F: std::future::Future<Output = T> + Send + 'static> SpawnableFuture<T> for F {}
#[cfg(target_arch = "wasm32")]
pub trait SpawnableFuture<T>: std::future::Future<Output = T> + 'static {}
#[cfg(target_arch = "wasm32")]
impl<T, F: std::future::Future<Output = T> + 'static> SpawnableFuture<T> for F {}
#[derive(Clone)]
pub struct BackgroundTasks {
job_tx: CbSender<Job>,
}
impl BackgroundTasks {
pub fn new(worker_count: usize) -> Self {
let (job_tx, job_rx): (CbSender<Job>, CbReceiver<Job>) = unbounded();
#[cfg(not(target_arch = "wasm32"))]
for _ in 0..worker_count.max(1) {
let job_rx = job_rx.clone(); std::thread::spawn(move || {
while let Ok(job) = job_rx.recv() {
if let Err(payload) = std::panic::catch_unwind(std::panic::AssertUnwindSafe(job)) {
tracing::error!(
"BackgroundTasks: a job panicked without going through its own \
error reporting — the worker thread survived regardless: {}",
panic_message(payload)
);
}
}
});
}
#[cfg(target_arch = "wasm32")]
let _ = (worker_count, &job_rx);
Self { job_tx }
}
pub fn spawn_blocking<T: Send + 'static>(
&self,
work: impl FnOnce() -> T + Send + 'static,
) -> TaskHandle<T> {
let (result_tx, result_rx) = channel::<Result<T, String>>();
let job: Job = Box::new(move || {
let outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(work)).map_err(|payload| {
let message = panic_message(payload);
tracing::error!("BackgroundTasks: a spawned task panicked: {message}");
message
});
let _ = result_tx.send(outcome);
});
let _ = self.job_tx.send(job);
TaskHandle { rx: result_rx }
}
#[cfg(not(target_arch = "wasm32"))]
pub fn spawn_async<T: Send + 'static>(
&self,
future: impl SpawnableFuture<T>,
) -> TaskHandle<T> {
self.spawn_blocking(move || pollster::block_on(future))
}
#[cfg(target_arch = "wasm32")]
pub fn spawn_async<T: 'static>(&self, future: impl SpawnableFuture<T>) -> TaskHandle<T> {
let (result_tx, result_rx) = channel::<Result<T, String>>();
wasm_bindgen_futures::spawn_local(async move {
let result = future.await;
let _ = result_tx.send(Ok(result)); });
TaskHandle { rx: result_rx }
}
}
pub enum TaskStatus<T> {
Pending,
Ready(T),
Panicked(String),
}
pub struct TaskHandle<T> {
rx: Receiver<Result<T, String>>,
}
impl<T> TaskHandle<T> {
pub fn poll(&mut self) -> TaskStatus<T> {
match self.rx.try_recv() {
Ok(Ok(value)) => TaskStatus::Ready(value),
Ok(Err(message)) => TaskStatus::Panicked(message),
Err(TryRecvError::Empty) => TaskStatus::Pending,
Err(TryRecvError::Disconnected) => TaskStatus::Panicked(
"the task's sender was dropped without ever sending a result — on native, this \
also means the panic that caused it was already logged via tracing::error! at \
the time it happened"
.to_string(),
),
}
}
pub fn try_recv(&mut self) -> Option<T> {
match self.poll() {
TaskStatus::Ready(value) => Some(value),
TaskStatus::Pending | TaskStatus::Panicked(_) => None,
}
}
}
pub struct BackgroundTasksPlugin {
worker_count: usize,
}
impl BackgroundTasksPlugin {
pub fn new(worker_count: usize) -> Self {
Self { worker_count }
}
}
impl Plugin for BackgroundTasksPlugin {
fn build(&self, app: &mut crate::prelude::App) {
app.add_resource(BackgroundTasks::new(self.worker_count));
}
}
#[cfg(all(test, not(target_arch = "wasm32")))]
mod tests {
use super::*;
use std::time::{Duration, Instant};
fn poll_until<T>(handle: &mut TaskHandle<T>, timeout: Duration) -> TaskStatus<T> {
let deadline = Instant::now() + timeout;
loop {
match handle.poll() {
TaskStatus::Pending => {
assert!(Instant::now() < deadline, "task did not resolve within {timeout:?}");
std::thread::sleep(Duration::from_millis(5));
}
status => return status,
}
}
}
#[test]
fn a_panicking_task_reports_panicked_instead_of_hanging_forever() {
let previous_hook = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let pool = BackgroundTasks::new(1);
let mut handle = pool.spawn_blocking(|| -> u32 { panic!("deliberate test panic") });
let status = poll_until(&mut handle, Duration::from_secs(2));
std::panic::set_hook(previous_hook);
match status {
TaskStatus::Panicked(message) => assert!(message.contains("deliberate test panic")),
_ => panic!("expected TaskStatus::Panicked"),
}
}
#[test]
fn the_worker_pool_survives_a_panic_and_keeps_processing_later_tasks() {
let previous_hook = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let pool = BackgroundTasks::new(1);
let mut doomed = pool.spawn_blocking(|| -> u32 { panic!("first task panics") });
poll_until(&mut doomed, Duration::from_secs(2));
let mut healthy = pool.spawn_blocking(|| 42u32);
let status = poll_until(&mut healthy, Duration::from_secs(2));
std::panic::set_hook(previous_hook);
match status {
TaskStatus::Ready(value) => assert_eq!(value, 42),
_ => panic!("expected the pool's sole worker thread to still be alive and processing"),
}
}
}