use std::task::Poll;
use crate::{
State,
lock::Lock,
producer::{Mut, Ref},
waiter::*,
};
#[derive(Debug)]
pub struct Shared<T> {
state: Lock<State<T>>,
}
impl<T: Default> Default for Shared<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T> Shared<T> {
pub fn new(value: T) -> Self {
Self {
state: Lock::new(State::new(value)),
}
}
pub fn lock(&self) -> Mut<'_, T> {
Mut::new(self.state.lock())
}
pub fn read(&self) -> Ref<'_, T> {
Ref {
state: self.state.lock(),
}
}
pub fn poll<F>(&self, waiter: &Waiter, mut f: F) -> Poll<Mut<'_, T>>
where
F: FnMut(&Ref<'_, T>) -> Poll<()>,
{
let mut guard = Ref {
state: self.state.lock(),
};
match f(&guard) {
Poll::Ready(()) => Poll::Ready(Mut::new(guard.state)),
Poll::Pending => {
waiter.register(&mut guard.state.waiters_value);
Poll::Pending
}
}
}
pub async fn wait<F>(&self, mut f: F) -> Mut<'_, T>
where
F: FnMut(&Ref<'_, T>) -> Poll<()> + Unpin,
{
crate::wait(move |waiter| self.poll(waiter, &mut f)).await
}
pub fn same_channel(&self, other: &Self) -> bool {
self.state.is_clone(&other.state)
}
}
impl<T> Clone for Shared<T> {
fn clone(&self) -> Self {
Self {
state: self.state.clone(),
}
}
}
#[cfg(all(test, not(loom)))]
mod test {
use std::{
future::Future,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
task::{Context, Wake, Waker},
};
use super::*;
struct CountWaker(AtomicUsize);
impl CountWaker {
fn count(&self) -> usize {
self.0.load(Ordering::SeqCst)
}
}
impl Wake for CountWaker {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
fn wake_by_ref(self: &Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
fn counting() -> (Arc<CountWaker>, Waker) {
let waker = Arc::new(CountWaker(AtomicUsize::new(0)));
let w = Waker::from(waker.clone());
(waker, w)
}
fn nonempty(queue: &Ref<'_, Vec<u32>>) -> Poll<()> {
if queue.is_empty() {
Poll::Pending
} else {
Poll::Ready(())
}
}
#[test]
fn enqueue_then_drain() {
let shared = Shared::<Vec<u32>>::default();
let drain = shared.clone();
let waiter = Waiter::noop();
assert!(drain.poll(&waiter, nonempty).is_pending());
shared.lock().push(1);
let Poll::Ready(mut guard) = drain.poll(&waiter, nonempty) else {
panic!("expected a drainable guard");
};
assert_eq!(guard.pop(), Some(1));
}
#[test]
fn mutation_wakes_parked_poll() {
let shared = Shared::<Vec<u32>>::default();
let drain = shared.clone();
let (waker, w) = counting();
let mut cx = Context::from_waker(&w);
let mut fut = Box::pin(crate::wait(|waiter| drain.poll(waiter, nonempty)));
assert!(fut.as_mut().poll(&mut cx).is_pending(), "pending until enqueue");
shared.lock().push(7);
assert!(waker.count() >= 1, "enqueue should wake the parked poll");
let Poll::Ready(mut guard) = fut.as_mut().poll(&mut cx) else {
panic!("expected a drainable guard after enqueue");
};
assert_eq!(guard.pop(), Some(7));
}
#[test]
fn read_does_not_wake() {
let shared = Shared::<Vec<u32>>::default();
let drain = shared.clone();
let (waker, w) = counting();
let mut cx = Context::from_waker(&w);
let mut fut = Box::pin(crate::wait(|waiter| drain.poll(waiter, nonempty)));
assert!(fut.as_mut().poll(&mut cx).is_pending());
assert!(shared.read().is_empty());
let guard = shared.lock();
assert!(guard.is_empty());
drop(guard);
assert_eq!(waker.count(), 0, "reads spuriously woke a parked poll");
}
#[tokio::test]
async fn wait_parks_until_enqueued() {
let shared = Shared::<Vec<u32>>::default();
let drain = shared.clone();
let task = tokio::spawn(async move {
let mut guard = drain.wait(nonempty).await;
guard.pop()
});
tokio::task::yield_now().await;
shared.lock().push(3);
assert_eq!(task.await.unwrap(), Some(3));
}
#[test]
fn same_channel_tracks_identity() {
let shared = Shared::<Vec<u32>>::default();
let clone = shared.clone();
let other = Shared::<Vec<u32>>::default();
assert!(shared.same_channel(&clone));
assert!(!shared.same_channel(&other));
}
}