#![allow(dead_code)]
use std::sync::{Arc, Mutex};
use std::task::Waker;
enum SlotState<T> {
Empty,
Waiting(Waker),
Ready(T),
Taken,
}
pub struct Oneshot<T> {
state: Mutex<SlotState<T>>,
}
impl<T> Oneshot<T> {
pub(crate) fn new() -> Arc<Self> {
Arc::new(Self {
state: Mutex::new(SlotState::Empty),
})
}
pub(crate) fn resolve(self: &Arc<Self>, value: T) {
let mut guard = self.state.lock().expect("oneshot mutex poisoned");
let old = std::mem::replace(&mut *guard, SlotState::Ready(value));
if let SlotState::Waiting(waker) = old {
drop(guard);
waker.wake();
}
}
pub(crate) fn poll(self: &Arc<Self>, cx: &mut std::task::Context<'_>) -> std::task::Poll<T> {
let mut guard = self.state.lock().expect("oneshot mutex poisoned");
match &*guard {
SlotState::Ready(_) => {
let SlotState::Ready(value) = std::mem::replace(&mut *guard, SlotState::Taken)
else {
unreachable!()
};
std::task::Poll::Ready(value)
}
SlotState::Taken => {
if cfg!(debug_assertions) {
panic!("oneshot polled after completion");
}
std::task::Poll::Pending
}
_ => {
*guard = SlotState::Waiting(cx.waker().clone());
std::task::Poll::Pending
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::task::{Context, Poll, Wake};
fn noop_cx() -> std::task::Waker {
std::task::Waker::noop().clone()
}
struct UnparkWaker(std::thread::Thread);
impl Wake for UnparkWaker {
fn wake(self: Arc<Self>) {
self.0.unpark();
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.unpark();
}
}
#[test]
fn resolve_before_poll() {
let os = Oneshot::new();
os.resolve(42u32);
let waker = noop_cx();
let mut cx = Context::from_waker(&waker);
assert_eq!(os.poll(&mut cx), Poll::Ready(42u32));
}
#[test]
fn poll_before_resolve() {
let os = Oneshot::new();
let waker = noop_cx();
let mut cx = Context::from_waker(&waker);
assert_eq!(os.poll(&mut cx), Poll::Pending);
os.resolve(99u32);
assert_eq!(os.poll(&mut cx), Poll::Ready(99u32));
}
#[test]
fn multiple_polls_before_resolve() {
let os = Oneshot::new();
let w1 = noop_cx();
let mut cx1 = Context::from_waker(&w1);
assert_eq!(os.poll(&mut cx1), Poll::Pending);
let w2 = noop_cx();
let mut cx2 = Context::from_waker(&w2);
assert_eq!(os.poll(&mut cx2), Poll::Pending);
os.resolve(7u32);
let w3 = noop_cx();
let mut cx3 = Context::from_waker(&w3);
assert_eq!(os.poll(&mut cx3), Poll::Ready(7u32));
}
#[cfg(debug_assertions)]
#[test]
#[should_panic(expected = "oneshot polled after completion")]
fn poll_after_taken_panics_in_debug() {
let os = Oneshot::new();
os.resolve(1u32);
let waker = noop_cx();
let mut cx = Context::from_waker(&waker);
assert_eq!(os.poll(&mut cx), Poll::Ready(1u32));
let _ = os.poll(&mut cx);
}
#[cfg(not(debug_assertions))]
#[test]
fn poll_after_taken_returns_pending_in_release() {
let os = Oneshot::new();
os.resolve(1u32);
let waker = noop_cx();
let mut cx = Context::from_waker(&waker);
assert_eq!(os.poll(&mut cx), Poll::Ready(1u32));
assert_eq!(os.poll(&mut cx), Poll::Pending);
}
#[test]
fn concurrent_resolve_and_poll() {
let os = Oneshot::new();
let os2 = Arc::clone(&os);
let main_thread = std::thread::current();
let unpark_waker = Arc::new(UnparkWaker(main_thread));
let waker = std::task::Waker::from(unpark_waker);
let mut cx = Context::from_waker(&waker);
assert_eq!(os.poll(&mut cx), Poll::Pending);
let handle = std::thread::spawn(move || {
std::thread::sleep(std::time::Duration::from_millis(5));
os2.resolve(123u32);
});
loop {
match os.poll(&mut cx) {
Poll::Ready(v) => {
assert_eq!(v, 123u32);
break;
}
Poll::Pending => std::thread::park(),
}
}
handle.join().unwrap();
}
#[test]
fn drop_without_resolve() {
let os = Oneshot::<u32>::new();
drop(os);
}
}