use std::collections::{HashMap, VecDeque};
use std::sync::{Arc, Mutex};
use serde::Deserialize;
use tokio::sync::oneshot;
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct ClientId(String);
impl ClientId {
pub(crate) const MAX_LEN: usize = 64;
pub(crate) const DEFAULT: &'static str = "default";
pub(crate) fn from_header(value: Option<&str>) -> ClientId {
value.map_or_else(|| ClientId(Self::DEFAULT.to_owned()), Self::parse)
}
pub(crate) fn parse(raw: &str) -> ClientId {
let trimmed = raw.trim();
let valid = !trimmed.is_empty()
&& trimmed.len() <= Self::MAX_LEN
&& trimmed
.bytes()
.all(|b| b.is_ascii_alphanumeric() || matches!(b, b'-' | b'_' | b'.' | b':'));
if valid {
ClientId(trimmed.to_owned())
} else {
ClientId(Self::DEFAULT.to_owned())
}
}
pub(crate) fn as_str(&self) -> &str {
&self.0
}
}
#[derive(Debug, Clone, Deserialize)]
#[serde(deny_unknown_fields)]
#[non_exhaustive]
pub(crate) struct QueueConfig {
#[serde(default = "default_max_depth")]
pub max_depth: usize,
#[serde(default = "default_fair_scheduling")]
pub fair_scheduling: bool,
}
fn default_max_depth() -> usize {
100
}
fn default_fair_scheduling() -> bool {
true
}
impl Default for QueueConfig {
fn default() -> Self {
Self {
max_depth: default_max_depth(),
fair_scheduling: default_fair_scheduling(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub(crate) enum AdmitError {
#[error("queue full")]
QueueFull,
#[error("endpoint lane unavailable")]
Unavailable,
}
#[derive(Debug)]
#[must_use = "dropping the permit releases the concurrency slot"]
pub(crate) struct Permit {
limited: Option<Arc<LimitedLane>>,
}
impl Drop for Permit {
fn drop(&mut self) {
if let Some(lane) = self.limited.take() {
lane.release_slot();
}
}
}
#[derive(Debug, Clone)]
pub(crate) struct EndpointLane {
inner: LaneInner,
}
#[derive(Debug, Clone)]
enum LaneInner {
Unlimited,
Limited(Arc<LimitedLane>),
}
struct LimitedLane {
max_inflight: usize,
max_depth: usize,
fair: bool,
state: Mutex<WaitState>,
}
impl std::fmt::Debug for LimitedLane {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LimitedLane")
.field("max_inflight", &self.max_inflight)
.field("max_depth", &self.max_depth)
.field("fair", &self.fair)
.finish_non_exhaustive()
}
}
struct WaitState {
inflight: usize,
waiter_count: usize,
next_id: u64,
fifo: VecDeque<Waiter>,
by_client: HashMap<String, VecDeque<Waiter>>,
client_order: VecDeque<String>,
}
struct Waiter {
id: u64,
reply: oneshot::Sender<()>,
}
impl EndpointLane {
#[must_use]
pub(crate) fn unlimited() -> EndpointLane {
EndpointLane {
inner: LaneInner::Unlimited,
}
}
#[must_use]
pub(crate) fn new(concurrency: usize, queue: &QueueConfig) -> EndpointLane {
let Some(max_inflight) = std::num::NonZeroUsize::new(concurrency) else {
return EndpointLane::unlimited();
};
EndpointLane {
inner: LaneInner::Limited(Arc::new(LimitedLane {
max_inflight: max_inflight.get(),
max_depth: queue.max_depth.max(1),
fair: queue.fair_scheduling,
state: Mutex::new(WaitState {
inflight: 0,
waiter_count: 0,
next_id: 1,
fifo: VecDeque::new(),
by_client: HashMap::new(),
client_order: VecDeque::new(),
}),
})),
}
}
#[cfg(test)]
pub(crate) fn waiter_count(&self) -> usize {
match &self.inner {
LaneInner::Unlimited => 0,
LaneInner::Limited(lane) => {
lane.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.waiter_count
}
}
}
#[cfg(test)]
pub(crate) fn distinct_clients(&self) -> usize {
match &self.inner {
LaneInner::Unlimited => 0,
LaneInner::Limited(lane) => lane
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.client_order
.len(),
}
}
pub(crate) async fn admit(&self, client_key: &str) -> Result<Permit, AdmitError> {
let LaneInner::Limited(lane) = &self.inner else {
return Ok(Permit { limited: None });
};
let outcome = {
let mut state = lane
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if state.inflight < lane.max_inflight {
state.inflight += 1;
AdmitOutcome::Ready
} else if state.waiter_count >= lane.max_depth {
AdmitOutcome::Full
} else {
let (tx, rx) = oneshot::channel();
let id = state.next_id;
state.next_id = state.next_id.wrapping_add(1);
state.waiter_count += 1;
let waiter = Waiter { id, reply: tx };
let effective_key = if lane.fair {
enqueue_fair(&mut state, client_key, waiter)
} else {
state.fifo.push_back(waiter);
client_key.to_string()
};
AdmitOutcome::Queued {
id,
rx,
client_key: effective_key,
}
}
};
match outcome {
AdmitOutcome::Ready => Ok(Permit {
limited: Some(Arc::clone(lane)),
}),
AdmitOutcome::Full => Err(AdmitError::QueueFull),
AdmitOutcome::Queued { id, rx, client_key } => {
let cancel = CancelOnDrop {
lane: Arc::clone(lane),
client_key,
id,
armed: true,
};
match rx.await {
Ok(()) => {
cancel.disarm();
Ok(Permit {
limited: Some(Arc::clone(lane)),
})
}
Err(_) => Err(AdmitError::Unavailable),
}
}
}
}
}
enum AdmitOutcome {
Ready,
Full,
Queued {
id: u64,
rx: oneshot::Receiver<()>,
client_key: String,
},
}
struct CancelOnDrop {
lane: Arc<LimitedLane>,
client_key: String,
id: u64,
armed: bool,
}
impl CancelOnDrop {
fn disarm(mut self) {
self.armed = false;
}
}
impl Drop for CancelOnDrop {
fn drop(&mut self) {
if !self.armed {
return;
}
let removed = {
let mut state = self
.lane
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if remove_waiter(&mut state, self.lane.fair, &self.client_key, self.id) {
state.waiter_count = state.waiter_count.saturating_sub(1);
true
} else {
false
}
};
if !removed {
self.lane.release_slot();
}
}
}
impl LimitedLane {
fn release_slot(&self) {
let mut state = self
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
loop {
let next = if self.fair {
dequeue_fair(&mut state)
} else {
state.fifo.pop_front()
};
let Some(waiter) = next else {
state.inflight = state.inflight.saturating_sub(1);
return;
};
state.waiter_count = state.waiter_count.saturating_sub(1);
if waiter.reply.send(()).is_ok() {
return;
}
}
}
}
const MAX_DISTINCT_CLIENTS: usize = 32;
fn enqueue_fair(state: &mut WaitState, client_key: &str, waiter: Waiter) -> String {
let key = if state.by_client.contains_key(client_key)
|| state.client_order.len() < MAX_DISTINCT_CLIENTS
{
client_key.to_string()
} else {
ClientId::DEFAULT.to_string()
};
let queue = state.by_client.entry(key.clone()).or_default();
let was_empty = queue.is_empty();
queue.push_back(waiter);
if was_empty {
state.client_order.push_back(key.clone());
}
key
}
fn dequeue_fair(state: &mut WaitState) -> Option<Waiter> {
let client = state.client_order.pop_front()?;
let queue = state.by_client.get_mut(&client)?;
let waiter = queue.pop_front()?;
if queue.is_empty() {
state.by_client.remove(&client);
} else {
state.client_order.push_back(client);
}
Some(waiter)
}
const _: fn() = || {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<EndpointLane>();
assert_send_sync::<Permit>();
assert_send_sync::<QueueConfig>();
assert_send_sync::<AdmitError>();
assert_send_sync::<ClientId>();
};
#[cfg(test)]
mod tests;
fn remove_waiter(state: &mut WaitState, fair: bool, client_key: &str, id: u64) -> bool {
if fair {
let Some(queue) = state.by_client.get_mut(client_key) else {
return false;
};
let Some(pos) = queue.iter().position(|w| w.id == id) else {
return false;
};
queue.remove(pos);
if queue.is_empty() {
state.by_client.remove(client_key);
state.client_order.retain(|c| c != client_key);
}
true
} else {
let Some(pos) = state.fifo.iter().position(|w| w.id == id) else {
return false;
};
state.fifo.remove(pos);
true
}
}