#![allow(dead_code)]
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll, Waker};
use bytes::Buf;
use futures::task::AtomicWaker;
use crate::quiche;
pub(crate) const PKT_BUF_LEN: usize = 64 * 1024;
pub(crate) const MAX_CHUNK: usize = 16 * 1024;
pub(crate) fn send_from_buf<B, F>(buf: &mut B, mut sink: F) -> Result<usize, quiche::Error>
where
B: Buf,
F: FnMut(&[u8]) -> Result<usize, quiche::Error>,
{
if !buf.has_remaining() {
return Ok(0);
}
let chunk = buf.chunk();
debug_assert!(!chunk.is_empty(), "Buf::has_remaining but empty chunk");
let written = sink(chunk)?;
debug_assert!(written <= chunk.len(), "sink accepted more than offered");
buf.advance(written);
Ok(written)
}
pub(crate) struct TerminalCell<T> {
inner: Arc<TerminalInner<T>>,
}
struct TerminalInner<T> {
value: Mutex<Option<T>>,
waker: AtomicWaker,
}
impl<T> Clone for TerminalCell<T> {
fn clone(&self) -> Self {
TerminalCell {
inner: Arc::clone(&self.inner),
}
}
}
impl<T: Clone> Default for TerminalCell<T> {
fn default() -> Self {
Self::new()
}
}
impl<T: Clone> TerminalCell<T> {
pub(crate) fn new() -> Self {
TerminalCell {
inner: Arc::new(TerminalInner {
value: Mutex::new(None),
waker: AtomicWaker::new(),
}),
}
}
pub(crate) fn set(&self, value: T) -> bool {
{
let mut slot = self.inner.value.lock().expect("TerminalCell poisoned");
if slot.is_some() {
return false;
}
*slot = Some(value);
}
self.inner.waker.wake();
true
}
pub(crate) fn get(&self) -> Option<T> {
self.inner
.value
.lock()
.expect("TerminalCell poisoned")
.clone()
}
pub(crate) fn poll(&self, cx: &mut Context<'_>) -> Poll<T> {
if let Some(v) = self.get() {
return Poll::Ready(v);
}
self.inner.waker.register(cx.waker());
match self.get() {
Some(v) => Poll::Ready(v),
None => Poll::Pending,
}
}
}
#[cfg_attr(test, derive(Debug))]
pub(crate) enum WriteOutcome<E> {
Done(Result<(), E>),
Cancelled,
}
pub(crate) struct WriteCompletion<E> {
inner: Arc<WriteCompletionInner<E>>,
}
struct WriteCompletionInner<E> {
generation: AtomicU64,
slot: Mutex<Option<(u64, WriteOutcome<E>)>>,
waker: AtomicWaker,
}
impl<E> Clone for WriteCompletion<E> {
fn clone(&self) -> Self {
WriteCompletion {
inner: Arc::clone(&self.inner),
}
}
}
impl<E> Default for WriteCompletion<E> {
fn default() -> Self {
Self::new()
}
}
impl<E> WriteCompletion<E> {
pub(crate) fn new() -> Self {
WriteCompletion {
inner: Arc::new(WriteCompletionInner {
generation: AtomicU64::new(0),
slot: Mutex::new(None),
waker: AtomicWaker::new(),
}),
}
}
pub(crate) fn begin(&self) -> u64 {
let generation = self.inner.generation.fetch_add(1, Ordering::AcqRel) + 1;
*self.inner.slot.lock().expect("WriteCompletion poisoned") = None;
generation
}
pub(crate) fn completer(&self, generation: u64) -> WriteCompleter<E> {
WriteCompleter {
cell: self.clone(),
generation,
completed: false,
}
}
fn set_if_current(&self, generation: u64, outcome: WriteOutcome<E>) {
{
let mut slot = self.inner.slot.lock().expect("WriteCompletion poisoned");
if self.inner.generation.load(Ordering::Acquire) != generation || slot.is_some() {
return; }
*slot = Some((generation, outcome));
}
self.inner.waker.wake();
}
pub(crate) fn poll(&self, generation: u64, cx: &mut Context<'_>) -> Poll<WriteOutcome<E>> {
self.inner.waker.register(cx.waker());
let mut slot = self.inner.slot.lock().expect("WriteCompletion poisoned");
if matches!(slot.as_ref(), Some((g, _)) if *g == generation) {
let (_, outcome) = slot.take().expect("slot just matched");
return Poll::Ready(outcome);
}
Poll::Pending
}
#[cfg(test)]
pub(crate) fn try_take(&self, generation: u64) -> Option<WriteOutcome<E>> {
let mut slot = self.inner.slot.lock().expect("WriteCompletion poisoned");
if matches!(slot.as_ref(), Some((g, _)) if *g == generation) {
return Some(slot.take().expect("slot just matched").1);
}
None
}
#[cfg(test)]
pub(crate) fn generation(&self) -> u64 {
self.inner.generation.load(Ordering::Acquire)
}
#[cfg(test)]
pub(crate) fn same_cell(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.inner, &other.inner)
}
}
pub(crate) struct WriteCompleter<E> {
cell: WriteCompletion<E>,
generation: u64,
completed: bool,
}
impl<E> WriteCompleter<E> {
pub(crate) fn complete(mut self, result: Result<(), E>) {
self.cell
.set_if_current(self.generation, WriteOutcome::Done(result));
self.completed = true;
}
}
impl<E> Drop for WriteCompleter<E> {
fn drop(&mut self) {
if !self.completed {
self.cell
.set_if_current(self.generation, WriteOutcome::Cancelled);
}
}
}
pub(crate) struct SendAccounting {
resident: AtomicUsize,
cap: Option<usize>,
waiters: Mutex<HashMap<u64, Waker>>,
}
impl SendAccounting {
pub(crate) fn new(cap: Option<usize>) -> Arc<Self> {
Arc::new(SendAccounting {
resident: AtomicUsize::new(0),
cap,
waiters: Mutex::new(HashMap::new()),
})
}
pub(crate) fn resident(&self) -> usize {
self.resident.load(Ordering::Acquire)
}
pub(crate) fn cap(&self) -> Option<usize> {
self.cap
}
pub(crate) fn try_reserve(self: &Arc<Self>, bytes: usize) -> Option<SendBytesPermit> {
let outcome = self
.resident
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |cur| {
let next = cur.checked_add(bytes)?;
match self.cap {
Some(cap) if bytes > 0 && next > cap && cur != 0 => None,
_ => Some(next),
}
});
match outcome {
Ok(_) => Some(SendBytesPermit {
accounting: Arc::clone(self),
bytes,
}),
Err(_) => None,
}
}
pub(crate) fn register_waiter(&self, id: u64, waker: &Waker) {
if self.cap.is_none() {
return;
}
let mut waiters = self.waiters.lock().expect("send-accounting waiters lock");
match waiters.get(&id) {
Some(existing) if existing.will_wake(waker) => {}
_ => {
waiters.insert(id, waker.clone());
}
}
}
pub(crate) fn unregister_waiter(&self, id: u64) {
if self.cap.is_none() {
return;
}
let mut waiters = self.waiters.lock().expect("send-accounting waiters lock");
waiters.remove(&id);
}
fn release(&self, bytes: usize) {
if bytes != 0 {
let _ = self
.resident
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |cur| {
Some(cur.saturating_sub(bytes))
});
}
if self.cap.is_none() {
return;
}
let wakers: Vec<Waker> = {
let mut waiters = self.waiters.lock().expect("send-accounting waiters lock");
waiters.drain().map(|(_id, w)| w).collect()
};
for w in wakers {
w.wake();
}
}
}
pub(crate) struct SendBytesPermit {
accounting: Arc<SendAccounting>,
bytes: usize,
}
impl SendBytesPermit {
pub(crate) fn bytes(&self) -> usize {
self.bytes
}
}
impl Drop for SendBytesPermit {
fn drop(&mut self) {
self.accounting.release(self.bytes);
}
}
#[cfg(test)]
mod tests {
use super::*;
use bytes::Bytes;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{RawWaker, RawWakerVTable, Waker};
#[test]
fn send_from_buf_partial_consume_contiguous() {
let mut buf = Bytes::from_static(b"hello world"); let n = send_from_buf(&mut buf, |chunk| {
assert_eq!(chunk, b"hello world");
Ok(5)
})
.unwrap();
assert_eq!(n, 5);
assert_eq!(buf.remaining(), 6);
let n = send_from_buf(&mut buf, |chunk| {
assert_eq!(chunk, b" world");
Ok(chunk.len())
})
.unwrap();
assert_eq!(n, 6);
assert_eq!(buf.remaining(), 0);
let n = send_from_buf(&mut buf, |_| panic!("sink must not be called")).unwrap();
assert_eq!(n, 0);
}
#[test]
fn send_from_buf_walks_noncontiguous_segments() {
let mut buf = Bytes::from_static(b"AAA").chain(Bytes::from_static(b"BBBB"));
let mut sent = Vec::new();
while buf.has_remaining() {
send_from_buf(&mut buf, |chunk| {
sent.extend_from_slice(chunk);
Ok(chunk.len())
})
.unwrap();
}
assert_eq!(sent, b"AAABBBB");
}
#[test]
fn send_from_buf_propagates_quiche_error() {
let mut buf = Bytes::from_static(b"x");
let err = send_from_buf(&mut buf, |_| Err(quiche::Error::Done)).unwrap_err();
assert!(matches!(err, quiche::Error::Done));
assert_eq!(buf.remaining(), 1);
}
#[test]
fn terminal_cell_first_writer_wins() {
let cell = TerminalCell::<u32>::new();
assert!(cell.set(1));
assert!(!cell.set(2)); assert_eq!(cell.get(), Some(1));
}
#[test]
fn terminal_cell_fast_path_ready() {
let cell = TerminalCell::<u32>::new();
cell.set(7);
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
assert_eq!(cell.poll(&mut cx), Poll::Ready(7));
}
#[test]
fn terminal_cell_pending_then_woken() {
let cell = TerminalCell::<u32>::new();
let woken = Arc::new(AtomicBool::new(false));
let waker = flag_waker(woken.clone());
let mut cx = Context::from_waker(&waker);
assert_eq!(cell.poll(&mut cx), Poll::Pending);
assert!(!woken.load(Ordering::SeqCst));
assert!(cell.set(99));
assert!(
woken.load(Ordering::SeqCst),
"set must wake the registered waker"
);
assert_eq!(cell.poll(&mut cx), Poll::Ready(99));
}
#[test]
fn terminal_cell_set_between_check_and_register() {
let cell = TerminalCell::<u32>::new();
let hook = Box::new(RaceHook {
cell: cell.clone(),
value: 55,
});
let hook_ptr = Box::into_raw(hook);
let waker = unsafe { Waker::from_raw(RawWaker::new(hook_ptr as *const (), &RACE_VTABLE)) };
let mut cx = Context::from_waker(&waker);
assert_eq!(cell.poll(&mut cx), Poll::Ready(55));
drop(waker);
unsafe { drop(Box::from_raw(hook_ptr)) };
}
#[test]
fn write_completion_reuses_one_cell_across_generations() {
let cell = WriteCompletion::<i32>::new();
assert_eq!(cell.generation(), 0);
for expected_gen in 1..=8u64 {
let generation = cell.begin();
assert_eq!(generation, expected_gen);
let completer = cell.completer(generation);
assert!(completer.cell.same_cell(&cell), "completer reuses the cell");
completer.complete(Ok(()));
assert!(matches!(
cell.try_take(generation),
Some(WriteOutcome::Done(Ok(())))
));
assert!(cell.try_take(generation).is_none(), "consumed exactly once");
}
assert_eq!(cell.generation(), 8, "one generation per write, one cell");
}
#[test]
fn write_completion_set_if_current_drops_stale_generation() {
let cell = WriteCompletion::<i32>::new();
let g1 = cell.begin();
let stale = cell.completer(g1);
let g2 = cell.begin();
assert_eq!(g2, g1 + 1);
let current = cell.completer(g2);
stale.complete(Ok(()));
assert!(cell.try_take(g1).is_none(), "stale g1 store dropped");
assert!(cell.try_take(g2).is_none(), "stale store must not fill g2");
current.complete(Err(-7));
assert!(matches!(
cell.try_take(g2),
Some(WriteOutcome::Done(Err(-7)))
));
assert!(cell.try_take(g2).is_none(), "g2 consumed exactly once");
}
#[test]
fn write_completion_drop_signals_cancelled() {
let cell = WriteCompletion::<i32>::new();
let generation = cell.begin();
let completer = cell.completer(generation);
drop(completer);
assert!(matches!(
cell.try_take(generation),
Some(WriteOutcome::Cancelled)
));
}
#[test]
fn write_completion_stale_drop_does_not_cancel_new_generation() {
let cell = WriteCompletion::<i32>::new();
let g1 = cell.begin();
let stale = cell.completer(g1);
let g2 = cell.begin();
let current = cell.completer(g2);
drop(stale); assert!(cell.try_take(g2).is_none(), "stale drop must not fill g2");
current.complete(Ok(()));
assert!(matches!(
cell.try_take(g2),
Some(WriteOutcome::Done(Ok(())))
));
}
#[test]
fn write_completion_poll_registers_and_wakes() {
let cell = WriteCompletion::<i32>::new();
let generation = cell.begin();
let flag = Arc::new(AtomicBool::new(false));
let waker = flag_waker(flag.clone());
let mut cx = Context::from_waker(&waker);
assert!(matches!(cell.poll(generation, &mut cx), Poll::Pending));
assert!(!flag.load(Ordering::SeqCst));
cell.completer(generation).complete(Ok(()));
assert!(flag.load(Ordering::SeqCst), "completion woke the poller");
match cell.poll(generation, &mut cx) {
Poll::Ready(WriteOutcome::Done(Ok(()))) => {}
other => panic!("expected Ready(Done(Ok)), got {other:?}"),
}
}
#[test]
fn send_accounting_unlimited_increments_and_releases_once() {
let acct = SendAccounting::new(None);
assert_eq!(acct.resident(), 0);
assert_eq!(acct.cap(), None);
let p1 = acct.try_reserve(1000).expect("unlimited reserve");
let p2 = acct.try_reserve(2500).expect("unlimited reserve");
assert_eq!(acct.resident(), 3500);
assert_eq!(p1.bytes(), 1000);
drop(p1);
assert_eq!(acct.resident(), 2500);
drop(p2);
assert_eq!(acct.resident(), 0);
}
#[test]
fn send_accounting_capped_rejects_over_cap_then_admits_after_release() {
let acct = SendAccounting::new(Some(100));
let p1 = acct.try_reserve(60).expect("fits under cap");
assert_eq!(acct.resident(), 60);
assert!(acct.try_reserve(60).is_none());
assert_eq!(
acct.resident(),
60,
"rejected reserve must not mutate residency"
);
let p2 = acct.try_reserve(40).expect("40 fits (100 total == cap)");
assert_eq!(acct.resident(), 100);
assert!(acct.try_reserve(1).is_none(), "at cap, nothing more admits");
drop(p1);
let _p3 = acct.try_reserve(60).expect("fits after release");
assert_eq!(acct.resident(), 100);
drop(p2);
}
#[test]
fn send_accounting_oversize_admits_one_unit_and_bounds_at_cap_plus_unit() {
let acct = SendAccounting::new(Some(100));
let big = acct
.try_reserve(250)
.expect("oversize admits when nothing resident");
assert_eq!(acct.resident(), 250);
assert!(acct.try_reserve(1).is_none());
assert!(acct.try_reserve(250).is_none());
assert_eq!(acct.resident(), 250);
drop(big);
assert_eq!(acct.resident(), 0);
let acct0 = SendAccounting::new(Some(0));
let unit = acct0
.try_reserve(10)
.expect("cap==0 admits one in-flight unit");
assert!(acct0.try_reserve(1).is_none());
drop(unit);
assert!(acct0.try_reserve(10).is_some());
}
#[test]
fn send_accounting_release_wakes_parked_waiter() {
let acct = SendAccounting::new(Some(100));
let p = acct.try_reserve(100).expect("fills cap");
let flag = Arc::new(AtomicBool::new(false));
let waker = flag_waker(flag.clone());
assert!(acct.try_reserve(50).is_none());
acct.register_waiter(7, &waker);
assert!(acct.try_reserve(50).is_none());
assert!(!flag.load(Ordering::SeqCst));
drop(p);
assert!(
flag.load(Ordering::SeqCst),
"release must wake the parked admission"
);
assert!(acct.try_reserve(50).is_some());
}
#[test]
fn send_accounting_register_waiter_dedups_equal_waker() {
let acct = SendAccounting::new(Some(10));
let _p = acct.try_reserve(10).expect("fills cap");
let flag = Arc::new(AtomicBool::new(false));
let waker = flag_waker(flag.clone());
acct.register_waiter(3, &waker);
acct.register_waiter(3, &waker);
acct.register_waiter(3, &waker);
assert_eq!(acct.waiters.lock().unwrap().len(), 1);
}
#[test]
fn send_accounting_unregister_drops_parked_waiter() {
let acct = SendAccounting::new(Some(100));
let _p = acct.try_reserve(100).expect("fills cap");
let flag_a = Arc::new(AtomicBool::new(false));
let flag_b = Arc::new(AtomicBool::new(false));
acct.register_waiter(1, &flag_waker(flag_a.clone()));
acct.register_waiter(2, &flag_waker(flag_b.clone()));
assert_eq!(acct.waiters.lock().unwrap().len(), 2);
acct.unregister_waiter(1);
assert_eq!(acct.waiters.lock().unwrap().len(), 1);
acct.unregister_waiter(1);
assert_eq!(acct.waiters.lock().unwrap().len(), 1);
drop(_p);
assert!(
!flag_a.load(Ordering::SeqCst),
"unregistered stream not woken"
);
assert!(flag_b.load(Ordering::SeqCst), "still-parked stream woken");
}
#[test]
fn send_accounting_unlimited_never_registers_waiters() {
let acct = SendAccounting::new(None);
let flag = Arc::new(AtomicBool::new(false));
acct.register_waiter(1, &flag_waker(flag));
assert!(acct.waiters.lock().unwrap().is_empty());
acct.unregister_waiter(1);
assert!(acct.waiters.lock().unwrap().is_empty());
}
#[test]
fn send_accounting_zero_byte_reserve_is_free() {
let acct = SendAccounting::new(Some(100));
let p = acct
.try_reserve(0)
.expect("zero-byte reserve always admits");
assert_eq!(acct.resident(), 0);
assert_eq!(p.bytes(), 0);
drop(p);
assert_eq!(acct.resident(), 0);
}
#[test]
fn send_accounting_zero_byte_reserve_admits_even_when_over_cap() {
let acct = SendAccounting::new(Some(100));
let big = acct
.try_reserve(250)
.expect("oversize admits as the one unit");
assert!(acct.resident() > acct.cap().unwrap());
assert!(acct.try_reserve(1).is_none());
let z = acct
.try_reserve(0)
.expect("zero-byte reserve admits over cap");
assert_eq!(z.bytes(), 0);
drop(z);
drop(big);
assert_eq!(acct.resident(), 0);
}
#[test]
fn send_accounting_concurrent_admissions_never_exceed_cap() {
use std::thread;
const UNIT: usize = 10;
const CAP: usize = 100;
let acct = Arc::new(SendAccounting::new(Some(CAP)));
let peak = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..16 {
let acct = Arc::clone(&acct);
let peak = Arc::clone(&peak);
handles.push(thread::spawn(move || {
let mut held: Vec<SendBytesPermit> = Vec::new();
for _ in 0..256 {
if let Some(p) = acct.try_reserve(UNIT) {
peak.fetch_max(acct.resident(), Ordering::SeqCst);
held.push(p);
if held.len() > 3 {
held.remove(0); }
}
}
}));
}
for h in handles {
h.join().unwrap();
}
assert!(
peak.load(Ordering::SeqCst) <= CAP,
"peak residency {} oversubscribed cap {CAP}",
peak.load(Ordering::SeqCst),
);
assert_eq!(acct.resident(), 0, "all reservations released");
}
fn flag_waker(flag: Arc<AtomicBool>) -> Waker {
let ptr = Arc::into_raw(flag) as *const ();
unsafe { Waker::from_raw(RawWaker::new(ptr, &FLAG_VTABLE)) }
}
static FLAG_VTABLE: RawWakerVTable = RawWakerVTable::new(
|p| unsafe {
let arc = Arc::from_raw(p as *const AtomicBool);
let cloned = arc.clone();
std::mem::forget(arc);
RawWaker::new(Arc::into_raw(cloned) as *const (), &FLAG_VTABLE)
},
|p| unsafe {
let arc = Arc::from_raw(p as *const AtomicBool);
arc.store(true, Ordering::SeqCst);
},
|p| unsafe {
let arc = Arc::from_raw(p as *const AtomicBool);
arc.store(true, Ordering::SeqCst);
std::mem::forget(arc);
},
|p| unsafe {
drop(Arc::from_raw(p as *const AtomicBool));
},
);
struct RaceHook {
cell: TerminalCell<u32>,
value: u32,
}
static RACE_VTABLE: RawWakerVTable = RawWakerVTable::new(
|p| unsafe {
let hook = &*(p as *const RaceHook);
hook.cell.set(hook.value);
RawWaker::new(std::ptr::null(), &NOOP_VTABLE)
},
|_| {},
|_| {},
|_| {},
);
static NOOP_VTABLE: RawWakerVTable = RawWakerVTable::new(
|_| RawWaker::new(std::ptr::null(), &NOOP_VTABLE),
|_| {},
|_| {},
|_| {},
);
}