use crate::split::bilock::{bilock, BiLock};
use futures::future::join;
use futures::task;
use std::ops::{Deref, DerefMut};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::Context;
#[derive(Default, Debug)]
pub struct TestWaker(AtomicBool);
impl TestWaker {
pub fn woken(&self) -> bool {
self.0.load(Ordering::SeqCst)
}
}
impl task::ArcWake for TestWaker {
fn wake_by_ref(arc_self: &Arc<Self>) {
arc_self.0.store(true, Ordering::SeqCst);
}
}
#[test]
fn bounds() {
fn f<T: Send + Sync>() {}
f::<BiLock<()>>();
}
#[tokio::test]
async fn simple_lock() {
let value = 13;
let (left, right) = bilock(value);
let guard = left.lock().await;
assert_eq!(*guard.deref(), 13);
drop(guard);
let guard = right.lock().await;
assert_eq!(*guard.deref(), 13);
}
#[tokio::test]
async fn guards() {
let value = 13;
let (left, right) = bilock(value);
let test_waker = Arc::new(TestWaker::default());
let waker = task::waker(test_waker.clone());
let mut ctx = Context::from_waker(&waker);
let poll = right.poll_lock(&mut ctx);
assert!(poll.is_ready());
drop(poll);
let mut guard = left.lock().await;
*guard.deref_mut() = 15;
let poll = right.poll_lock(&mut ctx);
assert!(poll.is_pending());
let poll = right.poll_lock(&mut ctx);
assert!(poll.is_pending());
drop(guard);
let poll = right.poll_lock(&mut ctx);
assert!(poll.is_ready());
assert!(test_waker.woken());
}
#[tokio::test]
async fn two_tasks() {
let value = 13;
let (left, right) = bilock(value);
let left_task = tokio::spawn(async move {
let mut guard = left.lock().await;
*guard.deref_mut() += 1;
drop(guard);
left
});
let right_task = tokio::spawn(async move {
let mut guard = right.lock().await;
*guard.deref_mut() += 100;
});
let (left_result, right_result) = join(left_task, right_task).await;
assert!(left_result.is_ok());
let left = left_result.unwrap();
assert!(right_result.is_ok());
let guard = left.lock().await;
assert_eq!(*guard.deref(), 114);
}
#[test]
fn reunite_ok() {
let (left, right) = bilock(13);
let reunite_result = left.reunite(right);
assert!(reunite_result.is_ok());
assert_eq!(reunite_result.unwrap(), 13);
}
#[test]
fn reunite_err() {
let (left, _right) = bilock(13);
let (_left, right) = bilock(13);
let reunite_result = left.reunite(right);
assert!(reunite_result.is_err());
assert_eq!(
reunite_result.unwrap_err().to_string(),
"Attempted to reunite two BiLocks that don't form a pair".to_string()
);
}