use crate::weak_waker::WeakWakerFuture;
use crate::worker::Worker;
use core::pin::Pin;
use libdd_capabilities::spawn::SpawnError;
use libdd_capabilities::MaybeSend;
use std::fmt::Display;
use std::future::Future;
use tokio::select;
use tokio_util::sync::CancellationToken;
use tracing::debug;
#[cfg(not(target_arch = "wasm32"))]
type WorkerFuture<T> = Pin<Box<dyn Future<Output = T> + Send + 'static>>;
#[cfg(target_arch = "wasm32")]
type WorkerFuture<T> = Pin<Box<dyn Future<Output = T> + 'static>>;
#[cfg(not(target_arch = "wasm32"))]
type WorkerJoinHandle<T> = Pin<Box<dyn Future<Output = Result<T, SpawnError>> + Send>>;
#[cfg(target_arch = "wasm32")]
type WorkerJoinHandle<T> = Pin<Box<dyn Future<Output = Result<T, SpawnError>>>>;
#[cfg(not(target_arch = "wasm32"))]
pub(super) fn tokio_spawn_fn<T: Send + 'static>(
handle: &tokio::runtime::Handle,
) -> impl FnOnce(WorkerFuture<T>) -> WorkerJoinHandle<T> {
let h = handle.clone();
move |future| {
let jh = h.spawn(future);
Box::pin(async { jh.await.map_err(|e| SpawnError::new(e.to_string())) })
}
}
pub enum PausableWorker<T: Worker + MaybeSend + Sync + 'static> {
Running {
handle: WorkerJoinHandle<T>,
stop_token: CancellationToken,
},
Paused {
worker: T,
},
InvalidState,
}
impl<T: Worker + MaybeSend + Sync + 'static> std::fmt::Debug for PausableWorker<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Running { .. } => f.debug_struct("PausableWorker::Running").finish(),
Self::Paused { worker } => f
.debug_struct("PausableWorker::Paused")
.field("worker", worker)
.finish(),
Self::InvalidState => write!(f, "PausableWorker::InvalidState"),
}
}
}
#[derive(Debug)]
pub enum PausableWorkerError {
InvalidState,
TaskAborted,
}
impl Display for PausableWorkerError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PausableWorkerError::InvalidState => {
write!(f, "Worker is in an invalid state and must be recreated.")
}
PausableWorkerError::TaskAborted => {
write!(f, "Worker task has been aborted and state has been lost.")
}
}
}
}
impl core::error::Error for PausableWorkerError {}
impl<T: Worker + MaybeSend + Sync + 'static> PausableWorker<T> {
pub fn new(worker: T) -> Self {
Self::Paused { worker }
}
pub fn start(
&mut self,
spawn_fn: impl FnOnce(WorkerFuture<T>) -> WorkerJoinHandle<T>,
) -> Result<(), PausableWorkerError> {
match self {
PausableWorker::Running { .. } => Ok(()),
PausableWorker::Paused { worker: _ } => {
debug!(?self, "Starting pausable worker");
let PausableWorker::Paused { mut worker } =
std::mem::replace(self, PausableWorker::InvalidState)
else {
return Ok(());
};
let stop_token = CancellationToken::new();
let cloned_token = stop_token.clone();
let future = Box::pin(async move {
select! {
biased;
_ = cloned_token.cancelled() => {
return worker;
}
_ = WeakWakerFuture::new(worker.initial_trigger()) => {
worker.run().await;
}
}
loop {
select! {
biased;
_ = cloned_token.cancelled() => {
break;
}
_ = WeakWakerFuture::new(worker.trigger()) => {
worker.run().await;
}
}
}
worker
});
let handle = spawn_fn(future);
*self = PausableWorker::Running { handle, stop_token };
Ok(())
}
PausableWorker::InvalidState => Err(PausableWorkerError::InvalidState),
}
}
pub async fn pause(&mut self) -> Result<(), PausableWorkerError> {
match self {
PausableWorker::Running { .. } => {
debug!("Waiting for worker to pause");
let PausableWorker::Running { handle, stop_token } =
std::mem::replace(self, PausableWorker::InvalidState)
else {
return Ok(());
};
if !stop_token.is_cancelled() {
stop_token.cancel();
}
if let Ok(worker) = handle.await {
debug!(?worker, "Worker paused successfully");
*self = PausableWorker::Paused { worker };
Ok(())
} else {
*self = PausableWorker::InvalidState;
Err(PausableWorkerError::TaskAborted)
}
}
PausableWorker::Paused { .. } => Ok(()),
PausableWorker::InvalidState => Err(PausableWorkerError::InvalidState),
}
}
pub fn reset(&mut self) {
if let PausableWorker::Paused { worker } = self {
worker.reset();
}
}
pub async fn shutdown(&mut self) {
if let PausableWorker::Paused { worker } = self {
worker.shutdown().await;
}
}
}
#[cfg(test)]
mod tests {
use async_trait::async_trait;
use tokio::{runtime::Builder, time::sleep};
use super::*;
use std::{
sync::mpsc::{channel, Sender},
time::Duration,
};
#[derive(Debug)]
struct TestWorker {
state: u32,
sender: Sender<u32>,
}
#[async_trait]
impl Worker for TestWorker {
async fn run(&mut self) {
let _ = self.sender.send(self.state);
self.state += 1;
}
async fn trigger(&mut self) {
sleep(Duration::from_millis(100)).await;
}
}
#[test]
fn test_restart() {
let (sender, receiver) = channel::<u32>();
let worker = TestWorker { state: 0, sender };
let runtime = Builder::new_multi_thread().enable_time().build().unwrap();
let handle = runtime.handle().clone();
let mut pausable_worker: PausableWorker<Box<dyn Worker + Sync>> =
PausableWorker::new(Box::new(worker));
pausable_worker.start(tokio_spawn_fn(&handle)).unwrap();
assert_eq!(receiver.recv().unwrap(), 0);
runtime.block_on(async { pausable_worker.pause().await.unwrap() });
let mut next_message = 1;
for message in receiver.try_iter() {
next_message = message + 1;
}
pausable_worker.start(tokio_spawn_fn(&handle)).unwrap();
assert_eq!(receiver.recv().unwrap(), next_message);
}
}