use std::{
collections::VecDeque,
num::NonZeroUsize,
panic::{AssertUnwindSafe, catch_unwind},
sync::{
Arc, Mutex,
atomic::{AtomicU8, Ordering as AtomicOrdering},
},
};
use tokio::sync::Notify;
use super::error::DropCapacityError;
const CAPACITY_WAITING: u8 = 0;
const CAPACITY_GRANTED: u8 = 1;
const CAPACITY_CLOSED: u8 = 2;
const CAPACITY_TAKEN: u8 = 3;
const CAPACITY_CANCELED: u8 = 4;
struct CapacitySignal {
status: AtomicU8,
changed: Notify,
}
impl CapacitySignal {
fn waiting() -> Self {
Self {
status: AtomicU8::new(CAPACITY_WAITING),
changed: Notify::new(),
}
}
}
struct LimitedCapacityState {
available: usize,
effective_capacity: usize,
waiters: VecDeque<Arc<CapacitySignal>>,
}
enum CapacityMode {
Limited(LimitedCapacityState),
Unlimited,
}
struct CapacityState {
closed: bool,
mode: CapacityMode,
}
pub(super) struct CapacityBroker {
limit: Option<NonZeroUsize>,
state: Mutex<CapacityState>,
retirement_reporter: Mutex<Option<RetirementReporter>>,
}
#[derive(Clone, Copy)]
pub(super) struct CapacityRetirement {
pub(super) configured_capacity: usize,
pub(super) effective_capacity: usize,
pub(super) retired_units: usize,
}
pub(super) type RetirementReporter = Arc<dyn Fn(CapacityRetirement) + Send + Sync + 'static>;
pub(super) struct CapacitySnapshot {
pub(super) configured_limit: Option<usize>,
pub(super) effective_limit: Option<usize>,
pub(super) available: Option<usize>,
pub(super) waiters: usize,
pub(super) open: bool,
}
enum CapacityStart {
Ready(OwnershipPermit),
Waiting(CapacityRequest),
}
struct CapacityRequest {
broker: Arc<CapacityBroker>,
signal: Arc<CapacitySignal>,
active: bool,
}
impl CapacityRequest {
async fn wait(mut self) -> Result<OwnershipPermit, DropCapacityError> {
loop {
let changed = self.signal.changed.notified();
tokio::pin!(changed);
changed.as_mut().enable();
match self.signal.status.load(AtomicOrdering::Acquire) {
CAPACITY_WAITING => changed.await,
CAPACITY_GRANTED => {
self.signal
.status
.store(CAPACITY_TAKEN, AtomicOrdering::Release);
self.active = false;
return Ok(OwnershipPermit::new(Arc::clone(&self.broker), 1));
}
CAPACITY_CLOSED => {
self.active = false;
return Err(self.broker.error());
}
_invalid => {
self.broker.close();
self.active = false;
return Err(self.broker.error());
}
}
}
}
}
impl Drop for CapacityRequest {
fn drop(&mut self) {
if self.active {
self.broker.cancel(&self.signal);
}
}
}
impl CapacityBroker {
pub(super) fn new(limit: Option<NonZeroUsize>) -> Arc<Self> {
let mode = match limit {
Some(limit) => CapacityMode::Limited(LimitedCapacityState {
available: limit.get(),
effective_capacity: limit.get(),
waiters: VecDeque::new(),
}),
None => CapacityMode::Unlimited,
};
Arc::new(Self {
limit,
state: Mutex::new(CapacityState {
closed: false,
mode,
}),
retirement_reporter: Mutex::new(None),
})
}
pub(super) fn set_retirement_reporter(&self, reporter: RetirementReporter) {
*self
.retirement_reporter
.lock()
.unwrap_or_else(|error| error.into_inner()) = Some(reporter);
}
fn error(&self) -> DropCapacityError {
DropCapacityError::new(self.limit)
}
pub(super) async fn acquire_one(
self: &Arc<Self>,
) -> Result<OwnershipPermit, DropCapacityError> {
match self.start_acquire_one()? {
CapacityStart::Ready(permit) => Ok(permit),
CapacityStart::Waiting(request) => request.wait().await,
}
}
pub(super) fn try_acquire(
self: &Arc<Self>,
units: usize,
) -> Result<OwnershipPermit, DropCapacityError> {
if units == 0 || self.limit.is_some_and(|limit| units > limit.get()) {
return Err(self.error());
}
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
if state.closed {
return Err(self.error());
}
let CapacityMode::Limited(limited) = &mut state.mode else {
return Ok(OwnershipPermit::new(Arc::clone(self), units));
};
if units > limited.effective_capacity
|| !limited.waiters.is_empty()
|| limited.available < units
{
return Err(self.error());
}
limited.available -= units;
Ok(OwnershipPermit::new(Arc::clone(self), units))
}
fn start_acquire_one(self: &Arc<Self>) -> Result<CapacityStart, DropCapacityError> {
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
if state.closed {
return Err(self.error());
}
let CapacityMode::Limited(limited) = &mut state.mode else {
return Ok(CapacityStart::Ready(OwnershipPermit::new(
Arc::clone(self),
1,
)));
};
if limited.effective_capacity == 0 {
return Err(self.error());
}
if limited.waiters.is_empty() && limited.available > 0 {
limited.available -= 1;
return Ok(CapacityStart::Ready(OwnershipPermit::new(
Arc::clone(self),
1,
)));
}
let capacity = self
.limit
.expect("limited capacity mode has a configured limit")
.get();
if limited.waiters.len() >= capacity {
return Err(self.error());
}
let signal = Arc::new(CapacitySignal::waiting());
limited.waiters.push_back(Arc::clone(&signal));
drop(state);
Ok(CapacityStart::Waiting(CapacityRequest {
broker: Arc::clone(self),
signal,
active: true,
}))
}
fn release(&self, units: usize) {
if units == 0 {
return;
}
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
let ready = if Self::return_capacity_locked(&mut state, units) {
Self::dispatch_locked(&mut state)
} else {
Self::close_waiters_locked(&mut state)
};
drop(state);
Self::notify(ready);
}
fn retire(&self, units: usize) {
if units == 0 {
return;
}
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
if state.closed {
return;
}
let CapacityMode::Limited(limited) = &mut state.mode else {
return;
};
let mut signals = Vec::new();
let mut retirement = None;
let close = if let Some(effective_capacity) = limited.effective_capacity.checked_sub(units)
{
if limited.available > effective_capacity {
true
} else {
limited.effective_capacity = effective_capacity;
retirement = Some(CapacityRetirement {
configured_capacity: self
.limit
.expect("limited capacity mode has a configured limit")
.get(),
effective_capacity,
retired_units: units,
});
effective_capacity == 0
}
} else {
true
};
if close {
state.closed = true;
signals.extend(Self::close_waiters_locked(&mut state));
} else {
signals.extend(Self::dispatch_locked(&mut state));
}
drop(state);
Self::notify(signals);
if let Some(retirement) = retirement {
self.report_retirement(retirement);
}
}
fn report_retirement(&self, retirement: CapacityRetirement) {
let reporter = self
.retirement_reporter
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone();
if let Some(reporter) = reporter
&& let Err(payload) = catch_unwind(AssertUnwindSafe(|| reporter(retirement)))
{
std::mem::forget(payload);
}
}
fn cancel(&self, signal: &Arc<CapacitySignal>) {
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
match signal.status.load(AtomicOrdering::Acquire) {
CAPACITY_WAITING => {
let CapacityMode::Limited(limited) = &mut state.mode else {
state.closed = true;
signal
.status
.store(CAPACITY_CANCELED, AtomicOrdering::Release);
return;
};
if let Some(position) = limited
.waiters
.iter()
.position(|waiter| Arc::ptr_eq(waiter, signal))
{
limited.waiters.remove(position);
}
}
CAPACITY_GRANTED => {
if !Self::return_capacity_locked(&mut state, 1) {
signal
.status
.store(CAPACITY_CANCELED, AtomicOrdering::Release);
let closed = Self::close_waiters_locked(&mut state);
drop(state);
Self::notify(closed);
return;
}
}
CAPACITY_CLOSED => {}
CAPACITY_TAKEN | CAPACITY_CANCELED => return,
_invalid => {
state.closed = true;
signal
.status
.store(CAPACITY_CANCELED, AtomicOrdering::Release);
let closed = Self::close_waiters_locked(&mut state);
drop(state);
Self::notify(closed);
return;
}
}
signal
.status
.store(CAPACITY_CANCELED, AtomicOrdering::Release);
let ready = Self::dispatch_locked(&mut state);
drop(state);
Self::notify(ready);
}
pub(super) fn close(&self) {
let signals = {
let mut state = self.state.lock().unwrap_or_else(|error| error.into_inner());
if state.closed {
return;
}
state.closed = true;
Self::close_waiters_locked(&mut state)
};
Self::notify(signals);
}
fn return_capacity_locked(state: &mut CapacityState, units: usize) -> bool {
let CapacityMode::Limited(limited) = &mut state.mode else {
return true;
};
let Some(available) = limited.available.checked_add(units) else {
state.closed = true;
return false;
};
if available > limited.effective_capacity {
state.closed = true;
return false;
}
limited.available = available;
true
}
fn close_waiters_locked(state: &mut CapacityState) -> Vec<Arc<CapacitySignal>> {
let CapacityMode::Limited(limited) = &mut state.mode else {
return Vec::new();
};
limited
.waiters
.drain(..)
.inspect(|signal| {
signal
.status
.store(CAPACITY_CLOSED, AtomicOrdering::Release);
})
.collect()
}
fn dispatch_locked(state: &mut CapacityState) -> Vec<Arc<CapacitySignal>> {
if state.closed {
return Vec::new();
}
let CapacityMode::Limited(limited) = &mut state.mode else {
return Vec::new();
};
Self::dispatch_limited(limited)
}
fn dispatch_limited(limited: &mut LimitedCapacityState) -> Vec<Arc<CapacitySignal>> {
let mut ready = Vec::new();
while limited.available > 0 {
let Some(signal) = limited.waiters.pop_front() else {
break;
};
limited.available -= 1;
signal
.status
.store(CAPACITY_GRANTED, AtomicOrdering::Release);
ready.push(signal);
}
ready
}
fn notify(signals: Vec<Arc<CapacitySignal>>) {
for signal in signals {
signal.changed.notify_waiters();
}
}
pub(super) fn snapshot(&self) -> CapacitySnapshot {
let state = self.state.lock().unwrap_or_else(|error| error.into_inner());
match &state.mode {
CapacityMode::Limited(limited) => CapacitySnapshot {
configured_limit: self.limit.map(NonZeroUsize::get),
effective_limit: Some(limited.effective_capacity),
available: Some(limited.available),
waiters: limited.waiters.len(),
open: !state.closed,
},
CapacityMode::Unlimited => CapacitySnapshot {
configured_limit: None,
effective_limit: None,
available: None,
waiters: 0,
open: !state.closed,
},
}
}
#[cfg(test)]
pub(super) fn available(&self) -> usize {
let state = self.state.lock().unwrap_or_else(|error| error.into_inner());
let CapacityMode::Limited(limited) = &state.mode else {
panic!("unlimited capacity has no available-unit count");
};
limited.available
}
#[cfg(test)]
pub(super) fn effective_capacity(&self) -> usize {
let state = self.state.lock().unwrap_or_else(|error| error.into_inner());
let CapacityMode::Limited(limited) = &state.mode else {
panic!("unlimited capacity has no effective-unit count");
};
limited.effective_capacity
}
#[cfg(test)]
pub(super) fn waiter_count(&self) -> usize {
let state = self.state.lock().unwrap_or_else(|error| error.into_inner());
match &state.mode {
CapacityMode::Limited(limited) => limited.waiters.len(),
CapacityMode::Unlimited => 0,
}
}
}
pub(super) struct OwnershipPermit {
broker: Arc<CapacityBroker>,
units: usize,
}
impl OwnershipPermit {
fn new(broker: Arc<CapacityBroker>, units: usize) -> Self {
Self { broker, units }
}
pub(super) fn split_one(&mut self) -> Option<Self> {
if self.units == 0 {
return None;
}
self.units -= 1;
Some(Self::new(Arc::clone(&self.broker), 1))
}
pub(super) fn retire(mut self) {
let units = std::mem::take(&mut self.units);
self.broker.retire(units);
}
pub(super) fn close_without_release(mut self) {
self.units = 0;
self.broker.close();
}
}
impl Drop for OwnershipPermit {
fn drop(&mut self) {
let units = std::mem::take(&mut self.units);
self.broker.release(units);
}
}