use std::future::Future;
use std::pin::Pin;
use std::sync::OnceLock;
use std::task::{Context, Poll};
use tokio::runtime::{Builder, Runtime};
use tokio::sync::oneshot;
fn runtime() -> &'static Runtime {
static RUNTIME: OnceLock<Runtime> = OnceLock::new();
RUNTIME.get_or_init(|| {
Builder::new_multi_thread()
.worker_threads(2)
.enable_all()
.thread_name("zippa-db")
.build()
.expect("failed to start the database runtime")
})
}
pub struct Task<T> {
receiver: oneshot::Receiver<T>,
running: Option<tokio::task::JoinHandle<()>>,
abort_on_drop: bool,
}
impl<T> Drop for Task<T> {
fn drop(&mut self) {
if self.abort_on_drop
&& let Some(running) = &self.running
{
running.abort();
}
}
}
impl<T> Task<T> {
pub fn abort_handle(&self) -> Option<tokio::task::AbortHandle> {
self.running.as_ref().map(|running| running.abort_handle())
}
pub fn try_recv(&mut self) -> Option<T> {
self.receiver.try_recv().ok()
}
pub fn abort_on_drop(mut self) -> Self {
self.abort_on_drop = true;
self
}
}
impl<T> Future for Task<T> {
type Output = Result<T, oneshot::error::RecvError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
Pin::new(&mut self.receiver).poll(cx)
}
}
#[cfg(not(test))]
pub fn spawn<F>(future: F) -> Task<F::Output>
where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
let (tx, rx) = oneshot::channel();
let running = runtime().spawn(async move {
let output = future.await;
let _ = tx.send(output);
});
Task {
receiver: rx,
running: Some(running),
abort_on_drop: false,
}
}
#[cfg(test)]
pub fn spawn<F>(future: F) -> Task<F::Output>
where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
let (tx, rx) = oneshot::channel();
let _ = tx.send(runtime().block_on(future));
Task {
receiver: rx,
running: None,
abort_on_drop: false,
}
}
#[cfg(test)]
pub fn block_on<F: Future>(future: F) -> F::Output {
runtime().block_on(future)
}