use std::future::Future;
use std::sync::atomic::{AtomicPtr, AtomicU32, Ordering};
use std::sync::{mpsc, Mutex};
use tokio::runtime::{Builder, Handle, Runtime};
static ACTOR_RUNTIME: AtomicPtr<Runtime> = AtomicPtr::new(std::ptr::null_mut());
static ACTOR_RUNTIME_PID: AtomicU32 = AtomicU32::new(0);
static ACTOR_RUNTIME_INIT: Mutex<()> = Mutex::new(());
pub(crate) fn runtime() -> &'static Runtime {
let pid = std::process::id();
if let Some(runtime) = current_for(pid) {
return runtime;
}
let _init = ACTOR_RUNTIME_INIT
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(runtime) = current_for(pid) {
return runtime;
}
let runtime: &'static Runtime = Box::leak(Box::new(build_runtime()));
ACTOR_RUNTIME.store(std::ptr::from_ref(runtime).cast_mut(), Ordering::Release);
ACTOR_RUNTIME_PID.store(pid, Ordering::Release);
runtime
}
fn current_for(pid: u32) -> Option<&'static Runtime> {
if ACTOR_RUNTIME_PID.load(Ordering::Acquire) != pid {
return None;
}
let pointer = ACTOR_RUNTIME.load(Ordering::Acquire);
unsafe { pointer.as_ref() }
}
fn build_runtime() -> Runtime {
let mut builder = Builder::new_multi_thread();
builder
.worker_threads(runtime_worker_threads())
.enable_time()
.thread_name("running-process-actor");
#[cfg(feature = "async-process")]
builder.enable_io();
builder.build().expect("process runtime must initialize")
}
pub(crate) fn runtime_worker_threads() -> usize {
std::thread::available_parallelism()
.map(usize::from)
.unwrap_or(2)
.clamp(2, 4)
}
pub(crate) fn block_on_anywhere<F>(future: F) -> F::Output
where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
let actor = runtime();
let Ok(current) = Handle::try_current() else {
return actor.block_on(future);
};
let (reply_tx, reply_rx) = mpsc::sync_channel(1);
actor.spawn(async move {
let _ = reply_tx.send(future.await);
});
let receive = move || {
reply_rx
.recv()
.expect("actor runtime dropped a sync adapter task")
};
if current.id() == actor.handle().id() {
tokio::task::block_in_place(receive)
} else {
receive()
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::{block_on_anywhere, runtime, runtime_worker_threads};
#[test]
fn worker_count_is_bounded() {
assert!((2..=4).contains(&runtime_worker_threads()));
}
#[test]
fn adapter_runs_outside_any_runtime() {
assert_eq!(block_on_anywhere(async { 7 }), 7);
}
#[tokio::test(flavor = "current_thread")]
async fn adapter_is_safe_on_a_current_thread_runtime() {
let value = block_on_anywhere(async {
tokio::time::sleep(Duration::from_millis(5)).await;
11
});
assert_eq!(value, 11);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn adapter_is_safe_on_a_multi_thread_runtime() {
assert_eq!(block_on_anywhere(async { 13 }), 13);
}
#[test]
fn adapter_is_safe_on_the_actor_runtime_itself() {
let value = runtime()
.block_on(async { tokio::spawn(async { block_on_anywhere(async { 17 }) }).await })
.expect("task joins");
assert_eq!(value, 17);
}
}