use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::OnceLock;
use std::sync::Weak;
use std::task::Context;
use std::task::Poll;
use crate::internal::mutex::Mutex;
use crate::internal::wake_all;
use crate::internal::wakerset::WakerSet;
use crate::internal::wakerset::WakerToken;
pub fn new<T>() -> (Completer<T>, Completion<T>) {
let shared = Arc::new(Shared {
value: OnceLock::new(),
state: Mutex::new(State {
status: Status::Pending,
waiters: WakerSet::new(),
}),
});
let completer = Completer {
shared: Arc::downgrade(&shared),
};
let completion = Completion { shared };
(completer, completion)
}
struct Shared<T> {
value: OnceLock<T>,
state: Mutex<State>,
}
struct State {
status: Status,
waiters: WakerSet,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum Status {
Pending,
Completed,
Abandoned,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Abandoned(());
impl fmt::Display for Abandoned {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("completion was abandoned before a value was provided")
}
}
impl std::error::Error for Abandoned {}
#[must_use = "dropping the completer abandons the completion"]
pub struct Completer<T> {
shared: Weak<Shared<T>>,
}
unsafe impl<T: Send> Send for Completer<T> {}
unsafe impl<T: Send> Sync for Completer<T> {}
impl<T> fmt::Debug for Completer<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Completer").finish_non_exhaustive()
}
}
impl<T> Completer<T> {
pub fn complete(mut self, value: T) -> Result<(), T> {
let Some(shared) = self.shared.upgrade() else {
return Err(value);
};
let wakers = {
let mut state = shared.state.lock();
assert_eq!(
state.status,
Status::Pending,
"a live completer must refer to a pending completion"
);
if let Err(value) = shared.value.set(value) {
drop(state);
drop(value);
panic!("pending completion value must be unset");
}
state.status = Status::Completed;
state.waiters.take_all()
};
self.shared = Weak::new();
wake_all(wakers);
Ok(())
}
}
impl<T> Drop for Completer<T> {
fn drop(&mut self) {
let Some(shared) = self.shared.upgrade() else {
return;
};
let wakers = {
let mut state = shared.state.lock();
if state.status != Status::Pending {
return;
}
state.status = Status::Abandoned;
state.waiters.take_all()
};
wake_all(wakers);
}
}
pub struct Completion<T> {
shared: Arc<Shared<T>>,
}
impl<T> Clone for Completion<T> {
fn clone(&self) -> Self {
Self {
shared: self.shared.clone(),
}
}
}
impl<T> fmt::Debug for Completion<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Completion").finish_non_exhaustive()
}
}
impl<T> Completion<T> {
pub async fn wait(&self) -> Result<&T, Abandoned> {
Wait {
completion: self,
token: None,
}
.await
}
}
struct Wait<'a, T> {
completion: &'a Completion<T>,
token: Option<WakerToken>,
}
impl<'a, T> Future for Wait<'a, T> {
type Output = Result<&'a T, Abandoned>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
let mut state = this.completion.shared.state.lock();
let (poll, retired_waker) = match state.status {
Status::Pending => {
let retired = state.waiters.register(&mut this.token, cx.waker());
(Poll::Pending, retired)
}
Status::Completed => {
this.token = None;
let completion: &'a Completion<T> = this.completion;
let value = completion
.shared
.value
.get()
.expect("completed value must be initialized");
(Poll::Ready(Ok(value)), None)
}
Status::Abandoned => {
this.token = None;
(Poll::Ready(Err(Abandoned(()))), None)
}
};
drop(state);
drop(retired_waker);
poll
}
}
impl<T> Drop for Wait<'_, T> {
fn drop(&mut self) {
if self.token.is_none() {
return;
}
let mut state = self.completion.shared.state.lock();
if state.status != Status::Pending {
self.token = None;
return;
}
let waker = state.waiters.unregister(&mut self.token);
drop(state);
drop(waker);
}
}