use futures::Future;
use std::{
collections::HashMap,
fmt,
ops::DerefMut,
sync::{Arc, Mutex, MutexGuard},
task::{Poll, Waker},
};
#[derive(Clone)]
pub struct Observable<T>
where
T: Clone,
{
inner: Arc<Mutex<Inner<T>>>,
version: u128,
}
impl<T> Observable<T>
where
T: Clone,
{
pub fn new(value: T) -> Self {
Observable {
inner: Arc::new(Mutex::new(Inner::new(value))),
version: 0,
}
}
pub fn publish(&mut self, value: T) {
self.modify(|v| *v = value);
}
pub fn modify<M>(&mut self, modify: M)
where
M: FnOnce(&mut T),
{
self.modify_conditional(|_| true, modify);
}
pub fn modify_conditional<C, M>(&mut self, condition: C, modify: M) -> bool
where
C: FnOnce(&T) -> bool,
M: FnOnce(&mut T),
{
self.apply(|value| {
if condition(value) {
modify(value);
true
} else {
false
}
})
}
#[doc(hidden)]
pub(crate) fn apply<F>(&mut self, change: F) -> bool
where
F: FnOnce(&mut T) -> bool,
{
let mut inner = self.lock();
if !change(&mut inner.value) {
return false;
}
inner.version += 1;
for waker in inner.waker.values() {
waker.wake_by_ref();
}
inner.waker.clear();
true
}
pub fn clone_and_reset(&self) -> Observable<T> {
Self {
inner: self.inner.clone(),
version: 0,
}
}
pub fn latest(&self) -> T {
let inner = self.lock();
inner.value.clone()
}
pub async fn next(&mut self) -> T {
AwaitObservableUpdate::from(self).await
}
pub fn synchronize(&mut self) -> T {
let (value, version) = {
let inner = self.lock();
(inner.value.clone(), inner.version)
};
self.version = version;
value
}
pub fn split(self) -> (Self, Self) {
(self.clone(), self)
}
pub(crate) fn lock(&self) -> MutexGuard<Inner<T>> {
match self.inner.lock() {
Ok(guard) => guard,
Err(e) => e.into_inner(),
}
}
#[cfg(test)]
pub(crate) fn waker_count(&self) -> usize {
self.inner.lock().unwrap().waker.len()
}
}
impl<T> Observable<T>
where
T: Clone + Eq,
{
pub fn publish_if_changed(&mut self, value: T) -> bool {
self.apply(|v| {
if *v != value {
*v = value;
true
} else {
false
}
})
}
}
impl<T> From<T> for Observable<T>
where
T: Clone,
{
fn from(value: T) -> Self {
Observable::new(value)
}
}
impl<T> fmt::Debug for Observable<T>
where
T: Clone + fmt::Debug,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let inner = self.lock();
f.debug_struct("Observable")
.field("inner", &inner)
.field("version", &self.version)
.finish()
}
}
struct Inner<T>
where
T: Clone,
{
version: u128,
future_count: u128,
value: T,
waker: HashMap<u128, Waker>,
}
impl<T> Inner<T>
where
T: Clone,
{
pub fn new(value: T) -> Self {
Self {
version: 0,
future_count: 0,
value,
waker: HashMap::new(),
}
}
pub fn add_waker(&mut self, id: u128, waker: Waker) {
self.waker.insert(id, waker);
}
pub fn remove_waker(&mut self, id: u128) {
self.waker.remove(&id);
}
}
impl<T> fmt::Debug for Inner<T>
where
T: Clone + fmt::Debug,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Inner")
.field("value", &self.value)
.field("version", &self.version)
.finish()
}
}
#[doc(hidden)]
struct AwaitObservableUpdate<'a, T>
where
T: Clone,
{
id: u128,
observable: &'a mut Observable<T>,
}
impl<'a, T: Clone> From<&'a mut Observable<T>> for AwaitObservableUpdate<'a, T> {
fn from(obs: &'a mut Observable<T>) -> Self {
let id = {
let mut guard = obs.lock();
let mut inner = guard.deref_mut();
inner.future_count += 1;
inner.future_count
};
Self {
id,
observable: obs,
}
}
}
impl<'a, T> Future for AwaitObservableUpdate<'a, T>
where
T: Clone,
{
type Output = T;
fn poll(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> Poll<Self::Output> {
let mut guard = self.observable.lock();
let inner = guard.deref_mut();
if self.observable.version == inner.version {
inner.add_waker(self.id, cx.waker().clone());
Poll::Pending
} else {
inner.remove_waker(self.id);
let (version, value) = (inner.version, inner.value.clone());
drop(guard);
self.observable.version = version;
Poll::Ready(value)
}
}
}
impl<'a, T> Drop for AwaitObservableUpdate<'a, T>
where
T: Clone,
{
fn drop(&mut self) {
let mut guard = self.observable.lock();
let inner = guard.deref_mut();
inner.remove_waker(self.id);
}
}
#[cfg(test)]
mod test {
use super::Observable;
use async_std::future::timeout;
use async_std::task::{sleep, spawn};
use std::time::Duration;
const SLEEP_DURATION: Duration = Duration::from_millis(25);
const TIMEOUT_DURATION: Duration = Duration::from_millis(500);
mod publishing {
use super::*;
use async_std::test;
#[test]
async fn should_get_notified_sync() {
let mut int = Observable::new(1);
let mut other = int.clone();
int.publish(2);
assert_eq!(other.next().await, 2);
int.publish(3);
assert_eq!(other.next().await, 3);
int.publish(0);
assert_eq!(other.next().await, 0);
}
#[test]
async fn should_get_notified_sync_multiple() {
let mut int = Observable::new(1);
let mut fork_one = int.clone();
let mut fork_two = int.clone();
int.publish(2);
assert_eq!(fork_one.next().await, 2);
assert_eq!(fork_two.next().await, 2);
int.publish(3);
assert_eq!(fork_one.next().await, 3);
assert_eq!(fork_two.next().await, 3);
int.publish(0);
assert_eq!(fork_one.next().await, 0);
assert_eq!(fork_two.next().await, 0);
}
#[test]
async fn should_publish_after_modify() {
let mut int = Observable::new(1);
let mut fork = int.clone();
int.modify(|i| *i += 1);
assert_eq!(fork.next().await, 2);
int.modify(|i| *i += 1);
assert_eq!(fork.next().await, 3);
int.modify(|i| *i -= 2);
assert_eq!(fork.next().await, 1);
int.modify(|i| *i -= 2);
assert_eq!(fork.next().await, -1);
}
#[test]
async fn should_conditionally_modify() {
let mut int = Observable::new(1);
let modified = int.modify_conditional(|i| i % 2 == 0, |i| *i *= 2);
assert!(!modified);
assert_eq!(int.latest(), 1);
let modified = int.modify_conditional(|i| i % 2 == 1, |i| *i *= 2);
assert!(modified);
assert_eq!(int.latest(), 2);
let modified = int.modify_conditional(|i| i % 2 == 0, |i| *i = 1000);
assert!(modified);
assert_eq!(int.latest(), 1000);
}
#[test]
async fn shouldnt_publish_same_change() {
let mut int = Observable::new(1);
let published = int.publish_if_changed(1);
assert!(!published);
assert!(timeout(TIMEOUT_DURATION, int.next()).await.is_err());
}
#[test]
async fn should_publish_changed() {
let mut int = Observable::new(1);
let published = int.publish_if_changed(2);
assert!(published);
assert_eq!(int.synchronize(), 2);
let published = int.publish_if_changed(2);
assert!(!published);
assert!(timeout(TIMEOUT_DURATION, int.next()).await.is_err());
}
}
mod versions {
use super::*;
use async_std::test;
#[test]
async fn should_skip_versions() {
let mut int = Observable::new(1);
let mut fork = int.clone();
int.publish(2);
int.publish(3);
int.publish(0);
assert_eq!(fork.next().await, 0);
}
#[test]
async fn should_wait_after_skiped_versions() {
let mut int = Observable::new(1);
let mut fork = int.clone();
int.publish(2);
int.publish(3);
int.publish(0);
assert_eq!(fork.next().await, 0);
assert!(timeout(TIMEOUT_DURATION, fork.next()).await.is_err());
}
#[test]
async fn should_skip_unchecked_updates() {
let mut int = Observable::new(1);
let mut fork = int.clone();
int.publish(2);
assert_eq!(fork.next().await, 2);
int.publish(3);
int.publish(0);
assert_eq!(fork.next().await, 0);
}
}
mod asynchronous {
use super::*;
use async_std::test;
#[test]
async fn should_wait_for_publisher_task() {
let mut int = Observable::new(1);
let mut fork = int.clone();
spawn(async move {
sleep(SLEEP_DURATION).await;
int.publish(2);
sleep(SLEEP_DURATION).await;
int.publish(3);
sleep(SLEEP_DURATION).await;
int.publish(0);
});
assert_eq!(fork.next().await, 2);
assert_eq!(fork.next().await, 3);
assert_eq!(fork.next().await, 0);
}
}
mod synchronization {
use super::*;
use async_std::test;
#[test]
async fn should_get_latest_without_loosing_updates() {
let mut int = Observable::new(1);
let mut fork = int.clone();
int.publish(2);
assert_eq!(fork.latest(), 2);
assert_eq!(fork.latest(), 2);
assert_eq!(fork.next().await, 2);
}
#[test]
async fn should_skip_updates_while_synchronizing() {
let mut int = Observable::new(1);
let mut fork = int.clone();
int.publish(2);
int.publish(3);
assert_eq!(fork.synchronize(), 3);
assert!(timeout(TIMEOUT_DURATION, fork.next()).await.is_err());
}
#[test]
async fn should_synchronize_multiple_times() {
let mut int = Observable::new(1);
let mut fork = int.clone();
int.publish(2);
int.publish(3);
assert_eq!(fork.synchronize(), 3);
assert_eq!(fork.synchronize(), 3);
int.publish(4);
assert_eq!(fork.synchronize(), 4);
assert!(timeout(TIMEOUT_DURATION, fork.next()).await.is_err());
}
}
mod future {
use super::*;
use async_std::test;
#[test]
async fn should_remove_waker_on_future_drop() {
let int = Observable::new(1);
let mut fork = int.clone();
for _ in 0..100 {
timeout(Duration::from_millis(10), fork.next()).await.ok();
assert_eq!(int.waker_count(), 0);
}
}
#[test]
async fn should_wait_forever() {
let int = Observable::new(1);
let mut fork = int.clone();
assert!(timeout(TIMEOUT_DURATION, fork.next()).await.is_err());
}
}
}