use crate::task::{InFlightTaskCancellation, TaskID};
use atomic_waker::AtomicWaker;
use std::any::Any;
use std::convert::Infallible;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use std::task::{Context, Poll};
use wasm_lite_std::Mutex;
#[derive(Debug, PartialEq, Eq, Hash, Clone)]
pub enum Observation<T> {
Pending,
Ready(T),
Done,
Cancelled,
}
#[derive(Debug, PartialEq, Eq, Hash, Clone)]
pub enum FinishedObservation<T> {
Ready(T),
Cancelled,
}
#[derive(Debug)]
struct Shared<T> {
lock: Mutex<Observation<T>>,
waker: AtomicWaker,
observer_cancelled: AtomicBool,
in_flight_task_cancellation: InFlightTaskCancellation,
}
#[must_use]
pub trait Observer: 'static + Future<Output = FinishedObservation<Self::Value>> {
type Value;
fn observe(&self) -> Observation<Self::Value>;
fn task_id(&self) -> &TaskID;
fn detach_in_place(&mut self);
fn detach(mut self)
where
Self: Sized,
{
self.detach_in_place();
}
}
#[derive(Debug)]
pub struct TypedObserver<T, ENotifier: ExecutorNotified> {
shared: Arc<Shared<T>>,
task_id: TaskID,
notifier: Option<ENotifier>,
detached: bool,
}
impl<T, ENotifier: ExecutorNotified> Drop for TypedObserver<T, ENotifier> {
fn drop(&mut self) {
if !self.detached {
let should_cancel = {
let state = self.shared.lock.lock_sync();
if matches!(*state, Observation::Pending) {
self.shared
.observer_cancelled
.store(true, std::sync::atomic::Ordering::Relaxed);
self.shared.in_flight_task_cancellation.cancel();
true
} else {
false
}
};
if should_cancel && let Some(mut n) = self.notifier.take() {
n.request_cancel()
}
}
}
}
impl<T: 'static, ENotifier: ExecutorNotified> Observer for TypedObserver<T, ENotifier> {
type Value = T;
fn observe(&self) -> Observation<Self::Value> {
TypedObserver::observe(self)
}
fn task_id(&self) -> &TaskID {
TypedObserver::task_id(self)
}
fn detach_in_place(&mut self) {
self.notifier.take();
self.detached = true;
}
}
#[derive(Debug)]
pub(crate) struct ObserverSender<T, Notifier> {
shared: Arc<Shared<T>>,
pub(crate) notifier: Option<Notifier>,
}
impl<T, Notifier> ObserverSender<T, Notifier> {
pub(crate) fn send(&mut self, value: T)
where
Notifier: ObserverNotified<T>,
{
if let Some(n) = self.notifier.as_mut() {
n.notify(&value)
}
{
let mut lock = self.shared.lock.lock_sync();
match *lock {
Observation::Pending => {
*lock = Observation::Ready(value);
}
Observation::Ready(_) => {
panic!("Observer already has a value");
}
Observation::Done => {
panic!("Observer already completed");
}
Observation::Cancelled => {
panic!("Observer cancelled");
}
}
}
self.shared.waker.wake();
}
pub(crate) fn observer_cancelled(&self) -> bool {
self.shared
.observer_cancelled
.load(std::sync::atomic::Ordering::Relaxed)
}
}
impl<T, Notifier> Drop for ObserverSender<T, Notifier> {
fn drop(&mut self) {
let should_wake = {
let mut lock = self.shared.lock.lock_sync();
match *lock {
Observation::Pending => {
self.shared.in_flight_task_cancellation.cancel();
*lock = Observation::Cancelled;
true
}
Observation::Ready(_) => {
false
}
Observation::Done => {
false
}
Observation::Cancelled => {
panic!("Observer cancelled");
}
}
};
if should_wake {
self.shared.waker.wake();
}
}
}
impl<T, E: ExecutorNotified> TypedObserver<T, E> {
pub fn observe(&self) -> Observation<T> {
let mut lock = self.shared.lock.lock_sync();
match *lock {
Observation::Pending => Observation::Pending,
Observation::Ready(..) => std::mem::replace(&mut *lock, Observation::Done),
Observation::Done => Observation::Done,
Observation::Cancelled => Observation::Cancelled,
}
}
pub fn task_id(&self) -> &TaskID {
&self.task_id
}
pub fn detach(mut self) {
self.notifier.take();
self.detached = true;
}
}
impl<T, E> Future for TypedObserver<T, E>
where
E: ExecutorNotified,
{
type Output = FinishedObservation<T>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.shared.waker.register(cx.waker());
let o = self.observe();
match o {
Observation::Pending => Poll::Pending,
Observation::Ready(v) => Poll::Ready(FinishedObservation::Ready(v)),
Observation::Done => {
panic!(
"Observer polled after the value was already taken (e.g. by a prior observe() or await)"
);
}
Observation::Cancelled => Poll::Ready(FinishedObservation::Cancelled),
}
}
}
pub trait ObserverNotified<T: ?Sized>: Unpin + 'static {
fn notify(&mut self, value: &T);
}
pub trait ExecutorNotified: 'static {
fn request_cancel(&mut self);
}
impl<T> ObserverNotified<T> for Infallible {
fn notify(&mut self, _value: &T) {
panic!("NoNotified should not be used");
}
}
impl ObserverNotified<Box<dyn std::any::Any + 'static>>
for Box<dyn ObserverNotified<dyn std::any::Any + 'static>>
{
fn notify(&mut self, value: &Box<dyn Any + 'static>) {
let r = Box::as_mut(self);
r.notify(&**value);
}
}
impl ObserverNotified<Box<dyn Any + Send + 'static>>
for Box<dyn ObserverNotified<dyn Any + Send + 'static> + Send>
{
fn notify(&mut self, value: &Box<dyn Any + Send + 'static>) {
let r = Box::as_mut(self);
r.notify(&**value);
}
}
impl ExecutorNotified for Infallible {
fn request_cancel(&mut self) {
panic!("NoNotified should not be used");
}
}
pub(crate) fn observer_channel<R, ONotifier, ENotifier: ExecutorNotified>(
observer_notify: Option<ONotifier>,
executor_notify: Option<ENotifier>,
task_cancellation: InFlightTaskCancellation,
task_id: TaskID,
) -> (ObserverSender<R, ONotifier>, TypedObserver<R, ENotifier>) {
let shared = Arc::new(Shared {
lock: Mutex::new(Observation::Pending),
waker: AtomicWaker::new(),
observer_cancelled: AtomicBool::new(false),
in_flight_task_cancellation: task_cancellation,
});
(
ObserverSender {
shared: shared.clone(),
notifier: observer_notify,
},
TypedObserver {
shared,
task_id,
notifier: executor_notify,
detached: false,
},
)
}
impl ExecutorNotified for Box<dyn ExecutorNotified + '_> {
fn request_cancel(&mut self) {
(**self).request_cancel();
}
}
impl ExecutorNotified for Box<dyn ExecutorNotified + Send> {
fn request_cancel(&mut self) {
(**self).request_cancel();
}
}
impl<T> From<T> for Observation<T> {
fn from(value: T) -> Self {
Observation::Ready(value)
}
}
impl<T> From<Observation<T>> for Option<T> {
fn from(value: Observation<T>) -> Self {
match value {
Observation::Ready(v) => Some(v),
_ => None,
}
}
}
#[cfg(test)]
mod tests {
use super::{ExecutorNotified, TypedObserver, observer_channel};
use crate::task::{InFlightTaskCancellation, TaskID};
use std::convert::Infallible;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
struct CancelCounter(Arc<AtomicUsize>);
impl ExecutorNotified for CancelCounter {
fn request_cancel(&mut self) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
#[cfg(not(target_arch = "wasm32"))]
struct ReentrantWakeProbe {
shared: Arc<super::Shared<u8>>,
blocked_on_observer_lock: std::sync::atomic::AtomicBool,
}
#[cfg(not(target_arch = "wasm32"))]
impl ReentrantWakeProbe {
fn probe(&self) {
use std::sync::mpsc;
use std::time::Duration;
let shared = self.shared.clone();
let (acquired, wait_for_acquire) = mpsc::channel();
std::thread::spawn(move || {
let _guard = shared.lock.lock_sync();
let _ = acquired.send(());
});
if wait_for_acquire
.recv_timeout(Duration::from_millis(250))
.is_err()
{
self.blocked_on_observer_lock.store(true, Ordering::Relaxed);
}
}
}
#[cfg(not(target_arch = "wasm32"))]
impl std::task::Wake for ReentrantWakeProbe {
fn wake(self: Arc<Self>) {
self.probe();
}
fn wake_by_ref(self: &Arc<Self>) {
self.probe();
}
}
#[cfg_attr(not(target_arch = "wasm32"), test)]
#[cfg_attr(target_arch = "wasm32", wasm_lite::wasm_lite_test)]
fn dropping_completed_observer_does_not_request_cancellation() {
let cancellation_requests = Arc::new(AtomicUsize::new(0));
let (mut sender, observer) = observer_channel::<u8, Infallible, CancelCounter>(
None,
Some(CancelCounter(cancellation_requests.clone())),
InFlightTaskCancellation::default(),
TaskID::from_u64(1),
);
sender.send(42);
drop(sender);
assert_eq!(observer.observe(), super::Observation::Ready(42));
drop(observer);
assert_eq!(
cancellation_requests.load(Ordering::Relaxed),
0,
"dropping an observer after successful completion must not request cancellation"
);
}
#[cfg(not(target_arch = "wasm32"))]
#[test]
fn observer_state_lock_is_released_before_waking() {
use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll, Waker};
let (mut sender, mut observer) = observer_channel::<u8, Infallible, Infallible>(
None,
None,
InFlightTaskCancellation::default(),
TaskID::from_u64(2),
);
let probe = Arc::new(ReentrantWakeProbe {
shared: observer.shared.clone(),
blocked_on_observer_lock: std::sync::atomic::AtomicBool::new(false),
});
let waker = Waker::from(probe.clone());
let mut context = Context::from_waker(&waker);
assert_eq!(Pin::new(&mut observer).poll(&mut context), Poll::Pending);
sender.send(7);
assert!(
!probe.blocked_on_observer_lock.load(Ordering::Relaxed),
"the observer's waker ran while the shared state mutex was still locked"
);
observer.detach();
}
#[cfg_attr(not(target_arch = "wasm32"), test)]
#[cfg_attr(target_arch = "wasm32", wasm_lite::wasm_lite_test)]
fn completing_task_leaves_cancellation_token_unset() {
let cancellation = InFlightTaskCancellation::default();
let (mut sender, observer) = observer_channel::<u8, Infallible, Infallible>(
None,
None,
cancellation.clone(),
TaskID::from_u64(3),
);
sender.send(42);
drop(sender);
assert!(
!cancellation.is_cancelled(),
"a task that ran to completion must not mark its cancellation token"
);
drop(observer);
}
#[cfg_attr(not(target_arch = "wasm32"), test)]
#[cfg_attr(target_arch = "wasm32", wasm_lite::wasm_lite_test)]
fn dropping_sender_before_completion_sets_cancellation_token() {
let cancellation = InFlightTaskCancellation::default();
let (sender, observer) = observer_channel::<u8, Infallible, Infallible>(
None,
None,
cancellation.clone(),
TaskID::from_u64(4),
);
drop(sender);
assert!(
cancellation.is_cancelled(),
"a task dropped before completion must mark its cancellation token"
);
drop(observer);
}
#[cfg_attr(not(target_arch = "wasm32"), test)]
#[cfg_attr(target_arch = "wasm32", wasm_lite::wasm_lite_test)]
fn boxed_notifier_unboxes_the_value() {
use std::any::Any;
struct Inner(Arc<AtomicUsize>);
impl super::ObserverNotified<dyn Any + Send> for Inner {
fn notify(&mut self, value: &(dyn Any + Send)) {
assert_eq!(
value.downcast_ref::<u32>().copied(),
Some(7),
"the erased notifier must receive the value, not the box around it"
);
self.0.fetch_add(1, Ordering::Relaxed);
}
}
let notifications = Arc::new(AtomicUsize::new(0));
let mut boxed: Box<dyn super::ObserverNotified<dyn Any + Send> + Send> =
Box::new(Inner(notifications.clone()));
let value: Box<dyn Any + Send> = Box::new(7u32);
super::ObserverNotified::notify(&mut boxed, &value);
assert_eq!(notifications.load(Ordering::Relaxed), 1);
}
#[cfg_attr(not(target_arch = "wasm32"), test)]
#[cfg_attr(target_arch = "wasm32", wasm_lite::wasm_lite_test)]
fn test_send() {
#[allow(unused)]
fn ex<T: Send, E: ExecutorNotified + Send>(_observer: TypedObserver<T, E>) {
fn assert_send<T: Send>() {}
assert_send::<TypedObserver<T, E>>();
}
}
#[cfg_attr(not(target_arch = "wasm32"), test)]
#[cfg_attr(target_arch = "wasm32", wasm_lite::wasm_lite_test)]
fn test_unpin() {
#[allow(unused)]
fn ex<T, E: ExecutorNotified + Unpin>(_observer: TypedObserver<T, E>) {
fn assert_unpin<T: Unpin>() {}
assert_unpin::<TypedObserver<T, E>>();
}
}
}