use std::collections::VecDeque;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex, Weak};
use std::task::{Context, Poll, Waker};
use weida_core::Limits;
use crate::conn::ConnCtx;
pub(crate) type Receipt =
Pin<Box<dyn Future<Output = Result<Option<u64>, weida_core::Error>> + Send + Sync>>;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Drained {
pub delivered: u64,
pub outstanding: u64,
}
pub(crate) struct DrainState {
draining: AtomicBool,
connections: Mutex<Vec<Weak<ConnCtx>>>,
evicted: AtomicU64,
}
impl DrainState {
pub(crate) fn new() -> DrainState {
DrainState {
draining: AtomicBool::new(false),
connections: Mutex::new(Vec::new()),
evicted: AtomicU64::new(0),
}
}
pub(crate) fn is_draining(&self) -> bool {
self.draining.load(Ordering::Relaxed)
}
pub(crate) fn begin(&self) {
self.draining.store(true, Ordering::Relaxed);
}
pub(crate) fn register(&self, conn: &Arc<ConnCtx>) {
let mut connections = self.connections.lock().expect("drain state poisoned");
if connections.len() == connections.capacity() {
connections.retain(|conn| conn.strong_count() > 0);
}
connections.push(Arc::downgrade(conn));
}
pub(crate) fn close_all(&self, code: u64, reason: &str) {
let connections = self.connections.lock().expect("drain state poisoned");
for conn in connections.iter().filter_map(|conn| conn.upgrade()) {
conn.conn.close(code, reason);
}
}
pub(crate) fn evict(&self) {
self.evicted.fetch_add(1, Ordering::Relaxed);
}
pub(crate) fn take(&self) -> (Vec<Receipt>, u64) {
let live: Vec<Arc<ConnCtx>> = {
let mut connections = self.connections.lock().expect("drain state poisoned");
connections.retain(|conn| conn.strong_count() > 0);
connections
.iter()
.filter_map(|conn| conn.upgrade())
.collect()
};
let mut receipts = Vec::new();
for conn in live {
receipts.append(&mut conn.parked.take());
}
(receipts, self.evicted.swap(0, Ordering::Relaxed))
}
}
pub(crate) struct ConnDrain {
parked: Mutex<VecDeque<Receipt>>,
max_parked: usize,
}
impl ConnDrain {
pub(crate) fn new(limits: &Limits, streams_are_local: bool) -> ConnDrain {
let by_stream_budget = limits.max_concurrent_uni_streams as usize
+ limits.max_concurrent_bidi_streams as usize;
let max_parked = if streams_are_local {
by_stream_budget.min(limits.max_local_streams / 2).max(1)
} else {
by_stream_budget
};
ConnDrain {
parked: Mutex::new(VecDeque::new()),
max_parked,
}
}
pub(crate) fn park(&self, receipt: Receipt) -> bool {
let mut parked = self.parked.lock().unwrap_or_else(|e| e.into_inner());
let mut lost = false;
if parked.len() >= self.max_parked {
parked.retain_mut(|receipt| settled(receipt).is_none());
if parked.len() >= self.max_parked {
parked.pop_front();
lost = true;
}
}
parked.push_back(receipt);
lost
}
pub(crate) fn reap(&self) -> usize {
let mut parked = self.parked.lock().unwrap_or_else(|e| e.into_inner());
let before = parked.len();
parked.retain_mut(|receipt| settled(receipt).is_none());
before - parked.len()
}
pub(crate) fn take(&self) -> Vec<Receipt> {
let mut parked = self.parked.lock().unwrap_or_else(|e| e.into_inner());
std::mem::take(&mut *parked).into()
}
}
fn settled(receipt: &mut Receipt) -> Option<Result<Option<u64>, weida_core::Error>> {
let mut cx = Context::from_waker(Waker::noop());
match receipt.as_mut().poll(&mut cx) {
Poll::Ready(outcome) => Some(outcome),
Poll::Pending => None,
}
}
pub(crate) async fn wait_for(
receipts: Vec<Receipt>,
evicted: u64,
deadline: impl Future<Output = ()>,
) -> Drained {
let mut delivered = 0u64;
let mut refused = 0u64;
let mut pending: Vec<Option<Receipt>> = receipts.into_iter().map(Some).collect();
{
let settle = std::future::poll_fn(|cx: &mut Context<'_>| {
let mut left = 0usize;
for slot in pending.iter_mut() {
let Some(receipt) = slot else { continue };
match receipt.as_mut().poll(cx) {
Poll::Ready(Ok(None)) => {
delivered += 1;
*slot = None;
}
Poll::Ready(_) => {
refused += 1;
*slot = None;
}
Poll::Pending => left += 1,
}
}
if left == 0 {
Poll::Ready(())
} else {
Poll::Pending
}
});
tokio::select! {
() = settle => {}
() = deadline => {}
}
}
let unfinished = pending.iter().filter(|slot| slot.is_some()).count() as u64;
Drained {
delivered,
outstanding: unfinished + refused + evicted,
}
}