use std::{
ops::{Deref, DerefMut},
sync::{Arc, atomic::Ordering},
task::Poll,
};
use crate::{Closed, Counts, State, consumer::Consumer, lock::*, waiter::*, weak::ProducerWeak};
#[derive(Debug)]
pub struct Producer<T> {
pub(crate) state: Lock<State<T>>,
pub(crate) counts: Arc<Counts>,
}
impl<T: Default> Default for Producer<T> {
fn default() -> Self {
Self {
state: Lock::new(State::default()),
counts: Arc::new(Counts::default()),
}
}
}
impl<T> Producer<T> {
pub fn new(value: T) -> Self {
Self {
state: Lock::new(State::new(value)),
counts: Arc::new(Counts::default()),
}
}
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 close(&self) -> Result<(), Ref<'_, T>> {
self.write()?.close();
Ok(())
}
pub fn write(&self) -> Result<Mut<'_, T>, Ref<'_, T>> {
let state = self.state.lock();
if state.closed {
Err(Ref { state })
} else {
Ok(Mut::new(state))
}
}
pub fn poll<F>(&self, waiter: &Waiter, mut f: F) -> Poll<Result<Mut<'_, T>, Ref<'_, T>>>
where
F: FnMut(&Ref<'_, T>) -> Poll<()>,
{
let state = self.state.lock();
if state.closed {
return Poll::Ready(Err(Ref { state }));
}
let mut guard = Ref { state };
match f(&guard) {
Poll::Ready(()) => Poll::Ready(Ok(Mut::new(guard.state))),
Poll::Pending => {
waiter.register(&mut guard.state.waiters_value);
Poll::Pending
}
}
}
pub fn poll_ref<F, R>(&self, waiter: &Waiter, mut f: F) -> Poll<Result<R, Ref<'_, T>>>
where
F: FnMut(&Ref<'_, T>) -> Poll<R>,
{
let state = self.state.lock();
let mut guard = Ref { state };
if let Poll::Ready(res) = f(&guard) {
return Poll::Ready(Ok(res));
}
if guard.state.closed {
return Poll::Ready(Err(guard));
}
waiter.register(&mut guard.state.waiters_value);
Poll::Pending
}
pub async fn wait<F>(&self, mut f: F) -> Result<Mut<'_, T>, Closed>
where
F: FnMut(&Ref<'_, T>) -> Poll<()> + Unpin,
{
crate::wait(move |waiter| self.poll(waiter, &mut f).map(|res| res.map_err(|_| Closed))).await
}
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 read(&self) -> Ref<'_, T> {
Ref {
state: self.state.lock(),
}
}
pub fn same_channel(&self, other: &Self) -> bool {
self.state.is_clone(&other.state)
}
pub fn is_last(&self) -> bool {
self.counts.producers.load(Ordering::Acquire) == 1
}
pub fn weak(&self) -> ProducerWeak<T> {
ProducerWeak {
state: self.state.clone(),
counts: self.counts.clone(),
}
}
}
impl<T> Clone for Producer<T> {
fn clone(&self) -> Self {
self.counts.producers.fetch_add(1, Ordering::Relaxed);
Self {
state: self.state.clone(),
counts: self.counts.clone(),
}
}
}
impl<T> Drop for Producer<T> {
fn drop(&mut self) {
let prev = self.counts.producers.fetch_sub(1, Ordering::AcqRel);
if prev > 1 {
return;
}
let mut waiters = {
let mut state = self.state.lock();
if state.closed {
return;
}
state.closed = true;
state.take_close_waiters()
};
for list in &mut waiters {
list.wake();
}
}
}
#[derive(Debug)]
pub struct Mut<'a, T> {
pub(crate) state: Option<LockGuard<'a, State<T>>>,
pub(crate) modified: bool,
}
impl<'a, T> Mut<'a, T> {
pub(crate) fn new(state: LockGuard<'a, State<T>>) -> Self {
Self {
state: Some(state),
modified: false,
}
}
pub fn close(mut self) {
let state = self.state.as_mut().unwrap();
state.closed = true;
self.modified = true;
}
}
impl<T> Deref for Mut<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.state.as_ref().unwrap().value
}
}
impl<T> DerefMut for Mut<'_, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
self.modified = true;
&mut self.state.as_mut().unwrap().value
}
}
impl<T> Drop for Mut<'_, T> {
fn drop(&mut self) {
let mut state = self.state.take().unwrap();
if !self.modified {
return;
}
let mut waiters_value = state.waiters_value.take();
let extra = state
.closed
.then(|| [state.waiters_closed.take(), state.waiters_consumer.take()]);
drop(state);
waiters_value.wake();
if let Some(mut extra) = extra {
for list in &mut extra {
list.wake();
}
}
}
}
pub struct Ref<'a, T> {
pub(crate) state: LockGuard<'a, State<T>>,
}
impl<T> Ref<'_, T> {
pub fn is_closed(&self) -> bool {
self.state.closed
}
}
impl<T> Deref for Ref<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.state.value
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn is_last_tracks_producer_count() {
let producer = Producer::new(0u8);
assert!(producer.is_last());
let clone = producer.clone();
assert!(!producer.is_last());
assert!(!clone.is_last());
drop(clone);
assert!(producer.is_last());
let _consumer = producer.consume();
let _weak = producer.weak();
assert!(producer.is_last());
}
#[test]
fn poll_gates_on_predicate_then_writes() {
let producer = Producer::<Vec<u32>>::default();
let waiter = Waiter::noop();
let predicate = |state: &Ref<'_, Vec<u32>>| {
if state.is_empty() {
Poll::Pending
} else {
Poll::Ready(())
}
};
assert!(matches!(producer.poll(&waiter, predicate), Poll::Pending));
let Ok(mut write) = producer.write() else {
panic!("channel should be open");
};
write.push(1);
drop(write);
let Poll::Ready(Ok(mut state)) = producer.poll(&waiter, predicate) else {
panic!("expected a writable guard");
};
assert_eq!(state.pop(), Some(1));
drop(state);
assert!(producer.close().is_ok());
assert!(matches!(producer.poll(&waiter, predicate), Poll::Ready(Err(_))));
}
}