use std::{
future::Future,
ops::Deref,
panic::Location,
pin::Pin,
sync::Arc,
task::{Context, Poll, Waker},
};
use parking_lot::{Mutex, MutexGuard};
use crate::flash::{
diag::PrimKind,
flash_ambient,
ids::{Backend, trace_native_from_ambient},
system,
};
pub mod error {
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecvError;
impl std::fmt::Display for RecvError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("watch channel closed")
}
}
impl std::error::Error for RecvError {}
}
pub use error::RecvError;
struct State<T> {
value: T,
wakers: Vec<Waker>,
closed: bool,
version: u64,
}
struct Shared<T> {
backend: Backend,
senders: Mutex<usize>,
state: Mutex<State<T>>,
}
impl<T> Shared<T> {
fn signal(&self, drained: Vec<Waker>) {
match self.backend {
Backend::Engine(cvid) => system::signal_channel(cvid, true),
Backend::Native => {
trace_native_from_ambient("watch", "signal");
for waker in drained {
waker.wake();
}
}
}
}
}
#[must_use]
#[track_caller]
pub fn channel<T>(init: T) -> (Sender<T>, Receiver<T>) {
let shared = Arc::new(Shared {
state: Mutex::new(State {
value: init,
version: 0,
closed: false,
wakers: Vec::new(),
}),
senders: Mutex::new(1),
backend: if flash_ambient() {
let cvid = system::next_condvar_id();
system::describe_cvid(cvid, PrimKind::Watch, Location::caller());
Backend::Engine(cvid)
} else {
Backend::Native
},
});
(
Sender {
shared: Arc::clone(&shared),
},
Receiver {
shared,
seen: 0,
pending: None,
},
)
}
pub struct Sender<T> {
shared: Arc<Shared<T>>,
}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Self {
*self.shared.senders.lock() += 1;
Self {
shared: Arc::clone(&self.shared),
}
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
let mut senders = self.shared.senders.lock();
*senders -= 1;
let last = *senders == 0;
drop(senders);
if last {
let mut state = self.shared.state.lock();
state.closed = true;
let drained = std::mem::take(&mut state.wakers);
drop(state);
self.shared.signal(drained);
}
}
}
impl<T> Sender<T> {
pub fn send(&self, value: T) -> Result<(), SendError<T>> {
let mut state = self.shared.state.lock();
state.value = value;
state.version += 1;
let drained = std::mem::take(&mut state.wakers);
drop(state);
self.shared.signal(drained);
Ok(())
}
}
pub struct SendError<T>(pub T);
impl<T> std::fmt::Debug for SendError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("SendError(..)")
}
}
impl<T> std::fmt::Display for SendError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("sending on a watch channel with no receivers")
}
}
impl<T> std::error::Error for SendError<T> {}
pub struct Receiver<T> {
shared: Arc<Shared<T>>,
pending: Option<Parked>,
seen: u64,
}
impl<T> Clone for Receiver<T> {
fn clone(&self) -> Self {
Self {
shared: Arc::clone(&self.shared),
seen: self.seen,
pending: None,
}
}
}
enum Parked {
Engine(system::AsyncHandle),
Real(Waker),
}
pub struct Ref<'a, T> {
guard: MutexGuard<'a, State<T>>,
}
impl<T> Deref for Ref<'_, T> {
type Target = T;
fn deref(&self) -> &T {
&self.guard.value
}
}
impl<T> Receiver<T> {
#[must_use]
pub fn borrow(&self) -> Ref<'_, T> {
Ref {
guard: self.shared.state.lock(),
}
}
#[must_use]
pub fn borrow_and_update(&mut self) -> Ref<'_, T> {
let guard = self.shared.state.lock();
self.seen = guard.version;
Ref { guard }
}
pub fn changed(&mut self) -> Changed<'_, T> {
Changed { rx: self }
}
}
pub struct Changed<'a, T> {
rx: &'a mut Receiver<T>,
}
impl<T> Future for Changed<'_, T> {
type Output = Result<(), RecvError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let rx = &mut *self.get_mut().rx;
match rx.pending.as_ref() {
Some(Parked::Engine(handle)) => {
if handle.granted() {
rx.pending = None;
} else {
return Poll::Pending;
}
}
Some(Parked::Real(_)) => rx.pending = None,
None => {}
}
let mut state = rx.shared.state.lock();
if state.version > rx.seen {
rx.seen = state.version;
drop(state);
return Poll::Ready(Ok(()));
}
if state.closed {
drop(state);
return Poll::Ready(Err(RecvError));
}
match rx.shared.backend {
Backend::Engine(cvid) => {
let (handle, adv) = system::register_channel_async(cvid, cx.waker().clone());
rx.pending = Some(Parked::Engine(handle));
drop(state);
adv.fire();
}
Backend::Native => {
trace_native_from_ambient("watch", "changed_park");
let waker = cx.waker().clone();
state.wakers.push(waker.clone());
rx.pending = Some(Parked::Real(waker));
drop(state);
}
}
Poll::Pending
}
}
impl<T> Drop for Changed<'_, T> {
fn drop(&mut self) {
match self.rx.pending.take() {
Some(Parked::Real(waker)) => {
self.rx
.shared
.state
.lock()
.wakers
.retain(|w| !w.will_wake(&waker));
}
Some(Parked::Engine(handle)) => system::cancel_async_wait(&handle),
None => {}
}
}
}
#[cfg(test)]
mod tests {
use kithara_test_utils::kithara;
use super::channel;
use crate::{flash, tokio::task::spawn};
#[kithara::test(tokio, multi_thread)]
async fn changed_no_lost_wakeup() {
flash::reset();
let (tx, mut rx) = channel::<u32>(0);
let waiter = spawn(async move {
rx.changed().await.expect("sender delivered");
*rx.borrow()
});
drop(spawn(async move {
tx.send(7).expect("receiver present");
}));
assert_eq!(waiter.await.expect("task joined"), 7);
}
#[kithara::test(tokio, multi_thread)]
async fn drop_sender_resolves_recv_error() {
flash::reset();
let (tx, mut rx) = channel::<u32>(0);
let waiter = spawn(async move { rx.changed().await });
drop(spawn(async move {
drop(tx);
}));
assert_eq!(waiter.await.expect("task joined"), Err(super::RecvError));
}
#[kithara::test(tokio, multi_thread)]
async fn borrow_and_update_marks_seen() {
flash::reset();
let (tx, mut rx) = channel::<u32>(0);
tx.send(1).expect("receiver present");
assert_eq!(*rx.borrow_and_update(), 1);
let waiter = spawn(async move {
rx.changed().await.expect("second change delivered");
*rx.borrow()
});
drop(spawn(async move {
tx.send(2).expect("receiver present");
}));
assert_eq!(waiter.await.expect("task joined"), 2);
}
}