use std::sync::{Arc, Condvar, Mutex};
use std::time::{Duration, Instant};
use super::error::{report_fault, RuntimeError};
use super::sync_lock;
struct Inner<T> {
slot: Mutex<Slot<T>>,
cvar: Condvar,
}
enum Slot<T> {
Pending,
Ready(T),
Abandoned,
Taken,
}
fn collect<T>(slot: &mut Slot<T>) -> Option<Result<T, RuntimeError>> {
match std::mem::replace(slot, Slot::Taken) {
Slot::Ready(value) => Some(Ok(value)),
Slot::Pending => {
*slot = Slot::Pending;
None
}
Slot::Abandoned => {
*slot = Slot::Abandoned;
Some(Err(RuntimeError::Abandoned("oneshot")))
}
Slot::Taken => Some(Err(RuntimeError::AlreadyCollected("oneshot"))),
}
}
#[derive(Debug)]
pub enum JoinState<T> {
Ready(T),
Pending,
}
pub struct Sender<T> {
inner: Arc<Inner<T>>,
sent: bool,
}
pub struct Receiver<T> {
inner: Arc<Inner<T>>,
}
pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
let inner = Arc::new(Inner {
slot: Mutex::new(Slot::Pending),
cvar: Condvar::new(),
});
(
Sender {
inner: inner.clone(),
sent: false,
},
Receiver { inner },
)
}
impl<T> Sender<T> {
pub fn send(mut self, value: T) {
match sync_lock::lock(&self.inner.slot, "oneshot::send") {
Ok(mut slot) => {
debug_assert!(
matches!(*slot, Slot::Pending),
"oneshot slot must be Pending before send: one sender, sends once"
);
*slot = Slot::Ready(value);
self.sent = true;
drop(slot);
self.inner.cvar.notify_one();
}
Err(e) => report_fault(e),
}
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
if self.sent {
return;
}
match sync_lock::lock(&self.inner.slot, "oneshot::disconnect") {
Ok(mut slot) => {
if matches!(*slot, Slot::Pending) {
*slot = Slot::Abandoned;
drop(slot);
self.inner.cvar.notify_one();
}
}
Err(e) => report_fault(e),
}
}
}
impl<T> Receiver<T> {
pub fn try_join(&self) -> Result<JoinState<T>, RuntimeError> {
let mut slot = sync_lock::lock(&self.inner.slot, "oneshot::try_join")?;
match collect(&mut slot) {
Some(result) => result.map(JoinState::Ready),
None => Ok(JoinState::Pending),
}
}
pub fn join_deadline(&self, deadline: Instant) -> Result<JoinState<T>, RuntimeError> {
let mut slot = sync_lock::lock(&self.inner.slot, "oneshot::join_deadline")?;
loop {
if let Some(result) = collect(&mut slot) {
return result.map(JoinState::Ready);
}
let Some(remaining) = deadline.checked_duration_since(Instant::now()) else {
return Ok(JoinState::Pending);
};
if remaining.is_zero() {
return Ok(JoinState::Pending);
}
let (guard, _) = sync_lock::wait_timeout(
&self.inner.cvar,
slot,
remaining,
"oneshot::join_deadline",
)?;
slot = guard;
}
}
pub fn join_timeout(&self, timeout: Duration) -> Result<JoinState<T>, RuntimeError> {
match Instant::now().checked_add(timeout) {
Some(deadline) => self.join_deadline(deadline),
None => self.wait_forever().map(JoinState::Ready),
}
}
pub fn join(self) -> Result<T, RuntimeError> {
self.wait_forever()
}
fn wait_forever(&self) -> Result<T, RuntimeError> {
let mut slot = sync_lock::lock(&self.inner.slot, "oneshot::join")?;
loop {
if let Some(result) = collect(&mut slot) {
return result;
}
slot = sync_lock::wait(&self.inner.cvar, slot, "oneshot::wait")?;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
use std::thread;
use std::time::{Duration, Instant};
type TestResult = Result<(), Box<dyn std::error::Error>>;
const WAKE_BUDGET: Duration = Duration::from_secs(5);
#[test]
fn join_returns_the_sent_value() -> TestResult {
let (tx, rx) = channel::<u32>();
tx.send(42);
match rx.join() {
Ok(v) => {
assert_eq!(v, 42);
Ok(())
}
Err(e) => Err(format!("expected the value, got {e}").into()),
}
}
#[test]
fn a_value_sent_before_the_drop_still_wins() -> TestResult {
let (tx, rx) = channel::<u32>();
tx.send(7);
match rx.join() {
Ok(v) => {
assert_eq!(v, 7);
Ok(())
}
Err(e) => Err(format!("a sent value must outrank the disconnect: {e}").into()),
}
}
#[test]
fn dropping_the_sender_before_the_join_reports_abandoned() -> TestResult {
let (tx, rx) = channel::<u32>();
drop(tx);
match rx.join() {
Err(RuntimeError::Abandoned(_)) => Ok(()),
Err(e) => Err(format!("expected Abandoned, got {e}").into()),
Ok(_) => Err("a dropped sender cannot yield a value".into()),
}
}
#[test]
fn try_join_is_pending_while_the_flow_is_still_running() -> TestResult {
let (tx, rx) = channel::<u32>();
match rx.try_join()? {
JoinState::Pending => {}
JoinState::Ready(v) => return Err(format!("nothing was sent, got {v}").into()),
}
drop(tx);
Ok(())
}
#[test]
fn the_value_is_handed_out_exactly_once() -> TestResult {
let (tx, rx) = channel::<u32>();
tx.send(9);
match rx.try_join()? {
JoinState::Ready(v) => assert_eq!(v, 9),
JoinState::Pending => return Err("the value was already sent".into()),
}
match rx.try_join() {
Err(RuntimeError::AlreadyCollected(_)) => Ok(()),
Err(e) => Err(format!("expected AlreadyCollected, got {e}").into()),
Ok(_) => Err("the value must not be handed out twice".into()),
}
}
#[test]
fn a_timeout_leaves_the_receiver_usable() -> TestResult {
let (tx, rx) = channel::<u32>();
match rx.join_timeout(Duration::from_millis(30))? {
JoinState::Pending => {}
JoinState::Ready(v) => return Err(format!("nothing was sent, got {v}").into()),
}
tx.send(4);
match rx.join_timeout(Duration::from_millis(30))? {
JoinState::Ready(v) => {
assert_eq!(v, 4);
Ok(())
}
JoinState::Pending => Err("the value was sent before this wait".into()),
}
}
#[test]
fn abandonment_outranks_a_pending_timeout() -> TestResult {
let (tx, rx) = channel::<u32>();
drop(tx);
match rx.join_timeout(Duration::from_secs(30)) {
Err(RuntimeError::Abandoned(_)) => Ok(()),
Err(e) => Err(format!("expected Abandoned, got {e}").into()),
Ok(_) => Err("a dropped sender cannot yield a value".into()),
}
}
#[test]
fn the_deadline_holds_under_a_storm_of_spurious_wakeups() -> TestResult {
let (tx, rx) = channel::<u32>();
let inner = rx.inner.clone();
let stop = Arc::new(AtomicBool::new(false));
let halt = stop.clone();
let noise = thread::spawn(move || {
while !halt.load(Ordering::SeqCst) {
inner.cvar.notify_all();
thread::sleep(Duration::from_millis(1));
}
});
let bound = Duration::from_millis(150);
let started = Instant::now();
let state = rx.join_timeout(bound);
let elapsed = started.elapsed();
stop.store(true, Ordering::SeqCst);
if noise.join().is_err() {
return Err("the notifier thread panicked".into());
}
drop(tx);
match state? {
JoinState::Pending => {}
JoinState::Ready(v) => return Err(format!("nothing was sent, got {v}").into()),
}
if elapsed > Duration::from_secs(5) {
return Err(format!("the bound was {bound:?} but the wait took {elapsed:?}").into());
}
if elapsed < Duration::from_millis(100) {
return Err(format!("returned after {elapsed:?}, before the bound").into());
}
Ok(())
}
#[test]
fn dropping_the_sender_wakes_a_joiner_that_is_already_blocked() -> TestResult {
let (tx, rx) = channel::<u32>();
let finished = Arc::new(AtomicBool::new(false));
let flag = finished.clone();
let joiner = thread::spawn(move || {
let outcome = rx.join();
flag.store(true, Ordering::SeqCst);
outcome
});
thread::sleep(Duration::from_millis(50));
drop(tx);
let deadline = Instant::now() + WAKE_BUDGET;
while !finished.load(Ordering::SeqCst) {
if Instant::now() > deadline {
return Err("join() never returned after the sender was dropped".into());
}
thread::sleep(Duration::from_millis(10));
}
match joiner.join() {
Ok(Err(RuntimeError::Abandoned(_))) => Ok(()),
Ok(Err(e)) => Err(format!("expected Abandoned, got {e}").into()),
Ok(Ok(_)) => Err("a dropped sender cannot yield a value".into()),
Err(_) => Err("the joining thread panicked".into()),
}
}
}