use std::{
future::poll_fn,
sync::{Arc, Mutex},
task::{Context, Poll, Waker},
};
use crate::Error;
#[derive(Debug, Clone, Copy)]
struct CreditState {
used: u64,
max: u64,
released: u64,
closed: bool,
}
#[derive(Clone, Debug)]
pub struct Credit {
inner: Arc<Mutex<Inner>>,
}
#[derive(Debug)]
struct Inner {
state: CreditState,
wakers: Vec<Waker>,
}
impl Inner {
fn update<T>(&mut self, f: impl FnOnce(&mut CreditState) -> (T, bool)) -> (T, Vec<Waker>) {
let (out, notify) = f(&mut self.state);
let wakers = if notify {
std::mem::take(&mut self.wakers)
} else {
Vec::new()
};
(out, wakers)
}
fn park(&mut self, cx: &Context<'_>) {
if !self.wakers.iter().any(|w| w.will_wake(cx.waker())) {
self.wakers.push(cx.waker().clone());
}
}
}
fn wake_all(wakers: Vec<Waker>) {
for waker in wakers {
waker.wake();
}
}
impl Credit {
pub fn new(max: u64) -> Self {
Self {
inner: Arc::new(Mutex::new(Inner {
state: CreditState {
used: 0,
max,
released: 0,
closed: false,
},
wakers: Vec::new(),
})),
}
}
pub fn try_claim(&self, limit: u64) -> u64 {
let mut inner = self.inner.lock().unwrap();
let available = inner.state.max.saturating_sub(inner.state.used);
let claimed = limit.min(available);
inner.state.used += claimed;
claimed
}
pub async fn claim(&self, limit: u64) -> Result<u64, Error> {
poll_fn(|cx| self.poll_claim(cx, limit)).await
}
pub fn poll_claim(&self, cx: &mut Context<'_>, limit: u64) -> Poll<Result<u64, Error>> {
let mut inner = self.inner.lock().unwrap();
if inner.state.closed {
return Poll::Ready(Err(Error::Closed));
}
let available = inner.state.max.saturating_sub(inner.state.used);
let claimed = limit.min(available);
if claimed > 0 {
inner.state.used += claimed;
return Poll::Ready(Ok(claimed));
}
inner.park(cx);
Poll::Pending
}
pub async fn claim_index(&self) -> Result<u64, Error> {
poll_fn(|cx| self.poll_claim_index(cx)).await
}
pub fn poll_claim_index(&self, cx: &mut Context<'_>) -> Poll<Result<u64, Error>> {
let mut inner = self.inner.lock().unwrap();
if inner.state.closed {
return Poll::Ready(Err(Error::Closed));
}
if inner.state.used < inner.state.max {
let index = inner.state.used;
inner.state.used += 1;
return Poll::Ready(Ok(index));
}
inner.park(cx);
Poll::Pending
}
pub fn close(&self) {
let (_, wakers) = self.inner.lock().unwrap().update(|state| {
let changed = !state.closed;
state.closed = true;
((), changed)
});
wake_all(wakers);
}
pub fn release(&self, amount: u64) {
let (_, wakers) = self.inner.lock().unwrap().update(|state| {
let new = state.used.saturating_sub(amount);
let changed = new != state.used;
state.used = new;
((), changed)
});
wake_all(wakers);
}
pub fn increase_max(&self, new_max: u64) -> Result<(), Error> {
let (ok, wakers) = self.inner.lock().unwrap().update(|state| {
if new_max < state.max {
return (false, false);
}
let changed = new_max != state.max;
state.max = new_max;
(true, changed)
});
wake_all(wakers);
if ok {
Ok(())
} else {
Err(Error::FlowControlError)
}
}
pub fn receive_up_to(&self, value: u64) -> bool {
let mut inner = self.inner.lock().unwrap();
if value > inner.state.max {
return false;
}
inner.state.used = inner.state.used.max(value);
true
}
pub fn receive(&self, len: u64) -> bool {
let mut inner = self.inner.lock().unwrap();
if inner.state.used + len > inner.state.max {
return false;
}
inner.state.used += len;
true
}
pub fn consume(&self, len: u64) -> Option<u64> {
let (update, wakers) = self.inner.lock().unwrap().update(|state| {
state.released += len;
if state.used + 2 * state.released > state.max {
let new_max = state.max + state.released;
state.max = new_max;
state.released = 0;
(Some(new_max), true)
} else {
(None, false)
}
});
wake_all(wakers);
update
}
}