use std::{
future::Future,
panic::Location,
pin::Pin,
sync::OnceLock,
task::{Context, Poll, Waker},
};
use parking_lot::Mutex;
use crate::flash::{
diag::PrimKind,
flash_ambient,
ids::{Backend, trace_native_from_ambient},
system,
};
#[derive(Debug, Default)]
struct Init {
wakers: Vec<Waker>,
in_progress: bool,
}
pub struct OnceCell<T> {
backend: Backend,
init: Mutex<Init>,
value: OnceLock<T>,
}
impl<T> OnceCell<T> {
fn finish_init(&self) {
let mut init = self.init.lock();
init.in_progress = false;
let wakers = std::mem::take(&mut init.wakers);
drop(init);
match self.backend {
Backend::Engine(cvid) => system::signal_channel(cvid, true),
Backend::Native => {
trace_native_from_ambient("oncecell", "finish_init");
for waker in wakers {
waker.wake();
}
}
}
}
pub async fn get_or_try_init<E, F, Fut>(&self, f: F) -> Result<&T, E>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Result<T, E>>,
{
loop {
if let Some(value) = self.value.get() {
return Ok(value);
}
let claim = {
let mut init = self.init.lock();
if self.value.get().is_some() {
Claim::Ready
} else if init.in_progress {
Claim::Wait
} else {
init.in_progress = true;
Claim::Init
}
};
match claim {
Claim::Ready => continue,
Claim::Wait => AwaitChange::new(self).await,
Claim::Init => {
let mut guard = AbandonGuard {
cell: self,
armed: true,
};
let result = f().await;
guard.armed = false;
return match result {
Ok(value) => {
let stored = self.value.get_or_init(move || value);
self.finish_init();
Ok(stored)
}
Err(err) => {
self.finish_init();
Err(err)
}
};
}
}
}
}
}
impl<T> Default for OnceCell<T> {
#[track_caller]
fn default() -> Self {
Self {
value: OnceLock::new(),
init: Mutex::new(Init::default()),
backend: if flash_ambient() {
let cvid = system::next_condvar_id();
system::describe_cvid(cvid, PrimKind::OnceCell, Location::caller());
Backend::Engine(cvid)
} else {
Backend::Native
},
}
}
}
impl<T> std::fmt::Debug for OnceCell<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OnceCell")
.field("backend", &self.backend)
.field("init", &self.init)
.field("initialized", &self.value.get().is_some())
.finish()
}
}
enum Claim {
Ready,
Wait,
Init,
}
enum Parked {
Engine(system::AsyncHandle),
Real(Waker),
}
struct AwaitChange<'a, T> {
cell: &'a OnceCell<T>,
pending: Option<Parked>,
}
impl<'a, T> AwaitChange<'a, T> {
fn new(cell: &'a OnceCell<T>) -> Self {
Self {
cell,
pending: None,
}
}
}
impl<T> Unpin for AwaitChange<'_, T> {}
impl<T> Future for AwaitChange<'_, T> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
let this = self.get_mut();
match this.pending.as_ref() {
Some(Parked::Engine(handle)) => {
if handle.granted() {
this.pending = None;
return Poll::Ready(());
}
return Poll::Pending;
}
Some(Parked::Real(_)) => {
this.pending = None;
return Poll::Ready(());
}
None => {}
}
let mut init = this.cell.init.lock();
if this.cell.value.get().is_some() || !init.in_progress {
return Poll::Ready(());
}
match this.cell.backend {
Backend::Engine(cvid) => {
let (handle, adv) = system::register_channel_async(cvid, cx.waker().clone());
this.pending = Some(Parked::Engine(handle));
drop(init);
adv.fire();
}
Backend::Native => {
trace_native_from_ambient("oncecell", "await_init");
let waker = cx.waker().clone();
init.wakers.push(waker.clone());
this.pending = Some(Parked::Real(waker));
drop(init);
}
}
Poll::Pending
}
}
impl<T> Drop for AwaitChange<'_, T> {
fn drop(&mut self) {
match self.pending.take() {
Some(Parked::Real(waker)) => {
self.cell
.init
.lock()
.wakers
.retain(|w| !w.will_wake(&waker));
}
Some(Parked::Engine(handle)) => system::cancel_async_wait(&handle),
None => {}
}
}
}
struct AbandonGuard<'a, T> {
cell: &'a OnceCell<T>,
armed: bool,
}
impl<T> Drop for AbandonGuard<'_, T> {
fn drop(&mut self) {
if self.armed {
self.cell.finish_init();
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use kithara_test_utils::kithara;
use super::OnceCell;
use crate::{
flash,
tokio::task::{spawn, yield_now},
};
const WAITERS: usize = 8;
#[kithara::test(tokio, multi_thread)]
async fn concurrent_init_runs_once_no_lost_wakeup() {
flash::reset();
let cell: Arc<OnceCell<usize>> = Arc::new(OnceCell::default());
let calls = Arc::new(AtomicUsize::new(0));
let handles: Vec<_> = (0..WAITERS)
.map(|_| {
let cell = Arc::clone(&cell);
let calls = Arc::clone(&calls);
spawn(async move {
cell.get_or_try_init(|| async {
calls.fetch_add(1, Ordering::SeqCst);
yield_now().await;
Ok::<_, ()>(42usize)
})
.await
.copied()
})
})
.collect();
for handle in handles {
assert_eq!(handle.await.expect("task joined"), Ok(42));
}
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[kithara::test(tokio, multi_thread)]
async fn failed_init_retries() {
flash::reset();
let cell: OnceCell<usize> = OnceCell::default();
let first = cell.get_or_try_init(|| async { Err("boom") }).await;
assert_eq!(first, Err("boom"));
let second = cell.get_or_try_init(|| async { Ok::<_, &str>(7) }).await;
assert_eq!(second, Ok(&7));
}
}