use std::cell::{Cell, RefCell};
use std::collections::VecDeque;
use std::future::Future;
use std::pin::Pin;
use std::rc::Rc;
use std::task::{Context, Poll, Waker};
struct GetWaiter {
amount: f64,
waker: Waker,
done: Rc<Cell<bool>>,
canceled: Rc<Cell<bool>>,
}
struct PutWaiter {
amount: f64,
waker: Waker,
done: Rc<Cell<bool>>,
canceled: Rc<Cell<bool>>,
}
struct ContainerState {
capacity: f64,
level: f64,
get_waiters: VecDeque<GetWaiter>,
put_waiters: VecDeque<PutWaiter>,
}
fn has_live_get_waiter(state: &ContainerState) -> bool {
state.get_waiters.iter().any(|w| !w.canceled.get())
}
fn has_live_put_waiter(state: &ContainerState) -> bool {
state.put_waiters.iter().any(|w| !w.canceled.get())
}
fn wake_get_waiters(state: &mut ContainerState) -> bool {
let mut serviced = false;
while let Some(front) = state.get_waiters.front() {
if front.canceled.get() {
state.get_waiters.pop_front();
continue;
}
if state.level >= front.amount {
let w = state.get_waiters.pop_front().unwrap();
state.level -= w.amount;
w.done.set(true);
w.waker.wake();
serviced = true;
} else {
break; }
}
serviced
}
fn wake_put_waiters(state: &mut ContainerState) -> bool {
let mut serviced = false;
while let Some(front) = state.put_waiters.front() {
if front.canceled.get() {
state.put_waiters.pop_front();
continue;
}
if state.level + front.amount <= state.capacity {
let w = state.put_waiters.pop_front().unwrap();
state.level += w.amount;
w.done.set(true);
w.waker.wake();
serviced = true;
} else {
break;
}
}
serviced
}
fn trigger_cascade(state: &mut ContainerState) {
loop {
let serviced_get = wake_get_waiters(state);
let serviced_put = wake_put_waiters(state);
if !serviced_get && !serviced_put {
break;
}
}
}
#[derive(Clone)]
pub struct Container {
state: Rc<RefCell<ContainerState>>,
}
impl std::fmt::Debug for Container {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut d = f.debug_struct("Container");
if let Ok(s) = self.state.try_borrow() {
d.field("level", &s.level)
.field("capacity", &s.capacity)
.field("get_waiters", &s.get_waiters.len())
.field("put_waiters", &s.put_waiters.len());
}
d.finish_non_exhaustive()
}
}
impl Container {
#[must_use]
pub fn empty(capacity: f64) -> Self {
Self::new(capacity, 0.0)
}
#[must_use]
pub fn new(capacity: f64, initial_level: f64) -> Self {
assert!(capacity > 0.0, "Container capacity must be positive");
assert!(
initial_level >= 0.0,
"Container initial_level must be non-negative"
);
assert!(
initial_level <= capacity,
"Container initial_level must not exceed capacity"
);
Container {
state: Rc::new(RefCell::new(ContainerState {
capacity,
level: initial_level,
get_waiters: VecDeque::new(),
put_waiters: VecDeque::new(),
})),
}
}
#[must_use]
pub fn level(&self) -> f64 {
self.state.borrow().level
}
#[must_use]
pub fn capacity(&self) -> f64 {
self.state.borrow().capacity
}
#[must_use]
pub fn get_queue_len(&self) -> usize {
self.state
.borrow()
.get_waiters
.iter()
.filter(|w| !w.canceled.get())
.count()
}
#[must_use]
pub fn put_queue_len(&self) -> usize {
self.state
.borrow()
.put_waiters
.iter()
.filter(|w| !w.canceled.get())
.count()
}
#[must_use = "futures do nothing unless awaited"]
pub fn put(&self, amount: f64) -> ContainerPutRequest {
assert!(amount > 0.0, "Container::put amount must be positive");
let capacity = self.state.borrow().capacity;
assert!(
amount <= capacity,
"Container::put amount ({amount}) exceeds capacity ({capacity}); it could never complete"
);
ContainerPutRequest {
state: Rc::clone(&self.state),
amount,
registered: false,
done: Rc::new(Cell::new(false)),
canceled: Rc::new(Cell::new(false)),
}
}
#[must_use = "futures do nothing unless awaited"]
pub fn get(&self, amount: f64) -> ContainerGetRequest {
assert!(amount > 0.0, "Container::get amount must be positive");
let capacity = self.state.borrow().capacity;
assert!(
amount <= capacity,
"Container::get amount ({amount}) exceeds capacity ({capacity}); it could never complete"
);
ContainerGetRequest {
state: Rc::clone(&self.state),
amount,
registered: false,
done: Rc::new(Cell::new(false)),
canceled: Rc::new(Cell::new(false)),
}
}
}
pub struct ContainerPutRequest {
state: Rc<RefCell<ContainerState>>,
amount: f64,
registered: bool,
done: Rc<Cell<bool>>,
canceled: Rc<Cell<bool>>,
}
impl std::fmt::Debug for ContainerPutRequest {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ContainerPutRequest")
.field("amount", &self.amount)
.field("registered", &self.registered)
.field("done", &self.done.get())
.finish_non_exhaustive()
}
}
impl Future for ContainerPutRequest {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
if self.done.get() {
return Poll::Ready(());
}
{
let mut state = self.state.borrow_mut();
if !self.registered
&& state.level + self.amount <= state.capacity
&& !has_live_put_waiter(&state)
{
state.level += self.amount;
trigger_cascade(&mut state);
return Poll::Ready(());
}
if !self.registered {
state.put_waiters.push_back(PutWaiter {
amount: self.amount,
waker: cx.waker().clone(),
done: Rc::clone(&self.done),
canceled: Rc::clone(&self.canceled),
});
}
}
self.registered = true;
Poll::Pending
}
}
impl Drop for ContainerPutRequest {
fn drop(&mut self) {
if self.registered && !self.done.get() {
self.canceled.set(true);
}
}
}
pub struct ContainerGetRequest {
state: Rc<RefCell<ContainerState>>,
amount: f64,
registered: bool,
done: Rc<Cell<bool>>,
canceled: Rc<Cell<bool>>,
}
impl std::fmt::Debug for ContainerGetRequest {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ContainerGetRequest")
.field("amount", &self.amount)
.field("registered", &self.registered)
.field("done", &self.done.get())
.finish_non_exhaustive()
}
}
impl Future for ContainerGetRequest {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
if self.done.get() {
return Poll::Ready(());
}
{
let mut state = self.state.borrow_mut();
if !self.registered && state.level >= self.amount && !has_live_get_waiter(&state) {
state.level -= self.amount;
trigger_cascade(&mut state);
return Poll::Ready(());
}
if !self.registered {
state.get_waiters.push_back(GetWaiter {
amount: self.amount,
waker: cx.waker().clone(),
done: Rc::clone(&self.done),
canceled: Rc::clone(&self.canceled),
});
}
}
self.registered = true;
Poll::Pending
}
}
impl Drop for ContainerGetRequest {
fn drop(&mut self) {
if self.registered && !self.done.get() {
self.canceled.set(true);
}
}
}