use std::{sync::atomic::Ordering, task::Poll};
use crate::{
Closed, Counts, State,
consumer::Consumer,
lock::*,
producer::{Producer, Ref},
sync::Arc,
waiter::*,
};
pub struct Weak<T> {
pub(crate) state: WeakLock<State<T>>,
pub(crate) counts: Arc<Counts>,
}
impl<T> Weak<T> {
pub fn new() -> Self {
Self {
state: WeakLock::new(),
counts: Arc::new(Counts::default()),
}
}
pub fn upgrade(&self) -> Option<Producer<T>> {
ProducerWeak {
state: self.state.upgrade()?,
counts: self.counts.clone(),
}
.produce()
}
}
impl<T> Default for Weak<T> {
fn default() -> Self {
Self::new()
}
}
impl<T> Clone for Weak<T> {
fn clone(&self) -> Self {
Self {
state: self.state.clone(),
counts: self.counts.clone(),
}
}
}
impl<T> std::fmt::Debug for Weak<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Weak")
.field("alive", &self.state.upgrade().is_some())
.finish()
}
}
#[derive(Debug)]
pub struct ProducerWeak<T> {
pub(crate) state: Lock<State<T>>,
pub(crate) counts: Arc<Counts>,
}
impl<T> ProducerWeak<T> {
pub fn produce(&self) -> Option<Producer<T>> {
{
let state = self.state.lock();
if state.closed {
return None;
}
self.counts.producers.fetch_add(1, Ordering::Relaxed);
}
Some(Producer {
state: self.state.clone(),
counts: self.counts.clone(),
})
}
pub fn consume(&self) -> Consumer<T> {
let prev = self.counts.consumers.fetch_add(1, Ordering::AcqRel);
if prev == 0 {
let mut waiters = self.state.lock().waiters_consumer.take();
waiters.wake();
}
Consumer {
state: self.state.clone(),
counts: self.counts.clone(),
}
}
pub fn read(&self) -> Ref<'_, T> {
Ref {
state: self.state.lock(),
}
}
pub fn is_closed(&self) -> bool {
self.state.lock().closed
}
pub async fn closed(&self) {
crate::wait(move |waiter| self.poll_closed(waiter)).await
}
pub fn poll_closed(&self, waiter: &Waiter) -> Poll<()> {
let mut state = self.state.lock();
if state.closed {
return Poll::Ready(());
}
waiter.register(&mut state.waiters_closed);
Poll::Pending
}
pub async fn unused(&self) -> Result<(), Closed> {
match crate::wait(move |waiter| self.poll_unused(waiter)).await {
Some(()) => Ok(()),
None => Err(Closed),
}
}
pub fn poll_unused(&self, waiter: &Waiter) -> Poll<Option<()>> {
let mut state = self.state.lock();
if state.closed {
return Poll::Ready(None);
}
if self.counts.consumers.load(Ordering::Relaxed) == 0 {
return Poll::Ready(Some(()));
}
waiter.register(&mut state.waiters_consumer);
if self.counts.consumers.load(Ordering::Relaxed) == 0 {
return Poll::Ready(Some(()));
}
Poll::Pending
}
pub fn is_used(&self) -> bool {
self.counts.consumers.load(Ordering::Relaxed) > 0
}
pub async fn used(&self) -> Result<(), Closed> {
match crate::wait(move |waiter| self.poll_used(waiter)).await {
Some(()) => Ok(()),
None => Err(Closed),
}
}
pub fn poll_used(&self, waiter: &Waiter) -> Poll<Option<()>> {
let mut state = self.state.lock();
if state.closed {
return Poll::Ready(None);
}
if self.counts.consumers.load(Ordering::Relaxed) > 0 {
return Poll::Ready(Some(()));
}
waiter.register(&mut state.waiters_consumer);
if self.counts.consumers.load(Ordering::Relaxed) > 0 {
return Poll::Ready(Some(()));
}
Poll::Pending
}
pub fn same_channel(&self, other: &Self) -> bool {
self.state.is_clone(&other.state)
}
}
impl<T> Clone for ProducerWeak<T> {
fn clone(&self) -> Self {
Self {
state: self.state.clone(),
counts: self.counts.clone(),
}
}
}
#[derive(Debug)]
pub struct ConsumerWeak<T> {
pub(crate) state: Lock<State<T>>,
pub(crate) counts: Arc<Counts>,
}
impl<T> ConsumerWeak<T> {
pub fn consume(&self) -> Consumer<T> {
let prev = self.counts.consumers.fetch_add(1, Ordering::AcqRel);
if prev == 0 {
let mut waiters = self.state.lock().waiters_consumer.take();
waiters.wake();
}
Consumer {
state: self.state.clone(),
counts: self.counts.clone(),
}
}
pub fn read(&self) -> Ref<'_, T> {
Ref {
state: self.state.lock(),
}
}
pub fn is_closed(&self) -> bool {
self.state.lock().closed
}
pub async fn closed(&self) {
crate::wait(move |waiter| self.poll_closed(waiter)).await
}
pub fn poll_closed(&self, waiter: &Waiter) -> Poll<()> {
let mut state = self.state.lock();
if state.closed {
return Poll::Ready(());
}
waiter.register(&mut state.waiters_closed);
Poll::Pending
}
pub fn same_channel(&self, other: &Self) -> bool {
self.state.is_clone(&other.state)
}
}
impl<T> Clone for ConsumerWeak<T> {
fn clone(&self) -> Self {
Self {
state: self.state.clone(),
counts: self.counts.clone(),
}
}
}
#[cfg(all(test, not(loom)))]
mod test {
use super::*;
#[tokio::test]
async fn weak_and_producer_agree_once_closed() {
let producer = Producer::new(0u32);
let weak = producer.weak();
producer.close().ok().expect("open");
assert_eq!(producer.unused().await, Err(Closed));
assert_eq!(weak.unused().await, Err(Closed));
assert_eq!(producer.used().await, Err(Closed));
assert_eq!(weak.used().await, Err(Closed));
}
#[tokio::test]
async fn weak_and_producer_agree_while_open() {
let producer = Producer::new(0u32);
let weak = producer.weak();
assert_eq!(producer.unused().await, Ok(()));
assert_eq!(weak.unused().await, Ok(()));
let consumer = producer.consume();
assert_eq!(producer.used().await, Ok(()));
assert_eq!(weak.used().await, Ok(()));
drop(consumer);
assert_eq!(weak.unused().await, Ok(()));
}
#[test]
fn weak_breaks_a_self_reference() {
struct Node {
_alive: Arc<()>,
back: Option<Weak<Node>>,
}
let alive = Arc::new(());
let producer = Producer::new(Node {
_alive: alive.clone(),
back: None,
});
producer.write().ok().expect("open").back = Some(producer.downgrade());
assert_eq!(Arc::strong_count(&alive), 2, "the state is alive");
let weak = producer.downgrade();
assert!(weak.upgrade().is_some(), "upgradeable while the channel is live");
drop(producer);
assert_eq!(Arc::strong_count(&alive), 1, "the state is gone despite the cycle");
assert!(weak.upgrade().is_none());
}
#[test]
fn weak_does_not_upgrade_once_closed() {
let producer = Producer::new(0u32);
let weak = producer.downgrade();
let consumer = producer.consume();
drop(producer);
assert!(weak.upgrade().is_none(), "closed, despite the state being alive");
drop(consumer);
assert!(weak.upgrade().is_none());
}
#[test]
fn upgrade_never_wins_a_closing_channel() {
for _ in 0..2_000 {
let producer = Producer::new(0u32);
let weak = producer.downgrade();
let consumer = producer.consume();
let gate = std::sync::Arc::new(std::sync::Barrier::new(2));
let dropper = {
let gate = gate.clone();
std::thread::spawn(move || {
gate.wait();
drop(producer);
})
};
gate.wait();
if let Some(upgraded) = weak.upgrade() {
assert!(
upgraded.write().is_ok(),
"an upgrade must not resolve a closing channel"
);
}
dropper.join().expect("dropper panicked");
drop(consumer);
}
}
#[test]
fn upgrade_holds_the_channel_open() {
let producer = Producer::new(0u32);
let weak = producer.downgrade();
let consumer = producer.consume();
let upgraded = weak.upgrade().expect("open");
drop(producer);
assert!(!consumer.is_closed(), "the upgrade keeps it open");
drop(upgraded);
assert!(consumer.is_closed(), "dropping the last one closes it");
}
#[tokio::test]
async fn consumer_weak_reads_and_observes_close() {
let producer = Producer::new(7u32);
let consumer = producer.consume();
let weak = consumer.weak();
assert_eq!(*weak.read(), 7);
assert!(!weak.is_closed());
drop(producer);
weak.closed().await;
assert!(weak.is_closed());
assert!(weak.read().is_closed());
}
}