use batchloader::{BatchController, BatchRules, KeySet};
use cooked_waker::{IntoWaker, Wake, WakeRef};
use futures::{channel::oneshot, FutureExt};
use std::{
cell::Cell,
collections::HashMap,
future::Future,
pin::Pin,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
task::{Context, Poll, Waker},
};
#[derive(Debug, Clone, Default, IntoWaker)]
struct BoolWaker {
cell: Arc<AtomicBool>,
}
impl BoolWaker {
fn reset(&self) {
self.cell.store(false, Ordering::SeqCst)
}
fn is_signaled(&self) -> bool {
self.cell.load(Ordering::SeqCst)
}
}
impl WakeRef for BoolWaker {
fn wake_by_ref(&self) {
self.cell.store(true, Ordering::SeqCst)
}
}
impl Wake for BoolWaker {}
#[derive(Debug, Clone)]
struct Skipper {
remaining_skips: usize,
}
impl Skipper {
fn new(count: usize) -> Self {
Skipper {
remaining_skips: count,
}
}
}
impl Future for Skipper {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
match &mut self.get_mut().remaining_skips {
0 => Poll::Ready(()),
skips => {
*skips -= 1;
cx.waker().wake_by_ref();
Poll::Pending
}
}
}
}
struct Task<F: Future + Unpin> {
fut: F,
signal: BoolWaker,
waker: Waker,
}
impl<F: Future + Unpin> Task<F> {
fn new(fut: F) -> Self {
let signal = BoolWaker::default();
Task {
fut,
waker: signal.clone().into_waker(),
signal,
}
}
fn poll(&mut self) -> Poll<F::Output> {
self.signal.reset();
self.fut.poll_unpin(&mut Context::from_waker(&self.waker))
}
fn is_signaled(&self) -> bool {
self.signal.is_signaled()
}
}
#[test]
fn test_notify_lifecycle() {
let delay_trigger = Cell::new(None);
let rules = BatchRules {
max_keys: None,
window: || {
let (trigger, fut) = oneshot::channel();
let old = delay_trigger.replace(Some(trigger));
assert!(old.is_none());
fut.map(|_| ())
},
batcher: |keys: KeySet<i32>| async {
Skipper::new(1).await;
if false {
Err(())
} else {
Ok(keys.into_values(|key| *key))
}
},
};
let controller = BatchController::new(&rules);
let mut task1 = Task::new(controller.load(1));
let mut task2 = Task::new(controller.load(2));
let mut task3 = Task::new(controller.load(3));
assert_eq!(task3.poll(), Poll::Pending);
assert_eq!(task2.poll(), Poll::Pending);
assert_eq!(task1.poll(), Poll::Pending);
assert!(!task1.is_signaled());
assert!(!task2.is_signaled());
assert!(!task3.is_signaled());
delay_trigger.take().unwrap().send(()).unwrap();
assert!(task1.is_signaled());
assert!(!task2.is_signaled());
assert!(!task3.is_signaled());
assert_eq!(task1.poll(), Poll::Pending);
assert!(task1.is_signaled());
assert!(!task2.is_signaled());
assert!(!task3.is_signaled());
assert_eq!(task1.poll(), Poll::Ready(Ok(1)));
assert!(task2.is_signaled());
assert!(task3.is_signaled());
assert_eq!(task2.poll(), Poll::Ready(Ok(2)));
assert_eq!(task3.poll(), Poll::Ready(Ok(3)));
}
#[test]
fn test_notify_lifecycle_drops() {
let delay_trigger = Cell::new(None);
let rules = BatchRules {
window: || {
let (trigger, fut) = oneshot::channel();
let old = delay_trigger.replace(Some(trigger));
assert!(old.is_none());
fut.map(|_| ())
},
max_keys: None,
batcher: |keys: KeySet<i32>| async {
Skipper::new(1).await;
if false {
Err(())
} else {
Ok(keys.into_values(|key| *key))
}
},
};
let controller = BatchController::new(&rules);
let mut tasks: HashMap<i32, _> = (1..=5)
.map(|key| (key, Task::new(controller.load(key))))
.collect();
for i in 1..=5 {
assert_eq!(tasks.get_mut(&i).unwrap().poll(), Poll::Pending);
}
assert!(tasks.values().all(|task| !task.is_signaled()));
tasks.remove(&5);
let mut driving_task = None;
for (&i, task) in tasks.iter() {
if task.is_signaled() {
match driving_task {
None => driving_task = Some(i),
Some(..) => panic!("Test failure: multiple tasks awoken after drop"),
}
}
}
let driving_task = driving_task.expect("Test failure: no task was awakened after a drop");
delay_trigger.take().unwrap().send(()).unwrap();
for (&i, task) in tasks.iter() {
if i == driving_task {
assert!(task.is_signaled());
} else {
assert!(!task.is_signaled());
}
}
assert_eq!(tasks.get_mut(&driving_task).unwrap().poll(), Poll::Pending);
tasks.remove(&driving_task);
let mut driving_task = None;
for (&i, task) in tasks.iter() {
if task.is_signaled() {
match driving_task {
None => driving_task = Some(i),
Some(..) => panic!("Test failure: multiple tasks awoken after drop"),
}
}
}
let driving_task = driving_task.expect("Test failure: no task was awakened after a drop");
assert_eq!(
tasks.get_mut(&driving_task).unwrap().poll(),
Poll::Ready(Ok(driving_task))
);
for (&i, task) in tasks.iter() {
if i == driving_task {
assert!(!task.is_signaled())
} else {
assert!(task.is_signaled())
}
}
tasks.remove(&driving_task);
for (&i, task) in tasks.iter_mut() {
assert_eq!(task.poll(), Poll::Ready(Ok(i)));
}
}