use crate::types::StreamFrame;
use spin_lock::Mutex;
use std::cell::RefCell;
use std::collections::VecDeque;
use std::future::poll_fn;
use std::sync::Arc;
use std::task::{Poll, Waker};
mod spin_lock {
use std::cell::UnsafeCell;
use std::ops::{Deref, DerefMut};
use std::sync::atomic::{AtomicBool, Ordering};
const SPINS_BEFORE_YIELD: u32 = 64;
#[derive(Debug, Default)]
pub(crate) struct Mutex<T> {
locked: AtomicBool,
value: UnsafeCell<T>,
}
unsafe impl<T: Send> Send for Mutex<T> {}
unsafe impl<T: Send> Sync for Mutex<T> {}
pub(crate) struct Guard<'a, T> {
lock: &'a Mutex<T>,
}
impl<T> Mutex<T> {
pub(crate) fn new(value: T) -> Self {
Self {
locked: AtomicBool::new(false),
value: UnsafeCell::new(value),
}
}
pub(crate) fn lock(&self) -> Guard<'_, T> {
let mut spins = 0u32;
loop {
if !self.locked.load(Ordering::Relaxed)
&& self
.locked
.compare_exchange_weak(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
{
return Guard { lock: self };
}
spins += 1;
if spins < SPINS_BEFORE_YIELD {
std::hint::spin_loop();
} else {
std::thread::yield_now();
}
}
}
pub(crate) fn get_mut(&mut self) -> &mut T {
self.value.get_mut()
}
}
impl<T> Deref for Guard<'_, T> {
type Target = T;
fn deref(&self) -> &T {
unsafe { &*self.lock.value.get() }
}
}
impl<T> DerefMut for Guard<'_, T> {
fn deref_mut(&mut self) -> &mut T {
unsafe { &mut *self.lock.value.get() }
}
}
impl<T> Drop for Guard<'_, T> {
fn drop(&mut self) {
self.lock.locked.store(false, Ordering::Release);
}
}
}
const MAX_POOLED: usize = 16;
#[derive(Debug)]
struct ChannelState {
queue: VecDeque<StreamFrame>,
waker: Option<Waker>,
closed: bool,
}
#[derive(Debug)]
#[repr(align(128))]
pub(crate) struct StreamChannel {
state: Mutex<ChannelState>,
capacity: usize,
}
#[derive(Debug)]
pub(crate) struct StreamSender {
chan: Arc<StreamChannel>,
}
impl StreamChannel {
pub(crate) fn try_send(&self, frame: StreamFrame) -> Result<(), StreamFrame> {
let mut state = self.state.lock();
if state.closed || state.queue.len() >= self.capacity {
drop(state);
return Err(frame);
}
state.queue.push_back(frame);
let waker = state.waker.take();
drop(state);
if let Some(waker) = waker {
waker.wake();
}
Ok(())
}
}
impl StreamSender {
pub(crate) fn try_send(&self, frame: StreamFrame) -> Result<(), StreamFrame> {
self.chan.try_send(frame)
}
pub(crate) fn channel(&self) -> Arc<StreamChannel> {
self.chan.clone()
}
}
impl Drop for StreamSender {
fn drop(&mut self) {
let mut state = self.chan.state.lock();
state.closed = true;
let waker = state.waker.take();
drop(state);
if let Some(waker) = waker {
waker.wake();
}
}
}
#[derive(Debug)]
pub(crate) struct StreamChannelReceiver {
chan: Option<Arc<StreamChannel>>,
}
impl StreamChannelReceiver {
pub(crate) async fn recv(&mut self) -> Option<StreamFrame> {
tokio::task::consume_budget().await;
let chan = self.chan.as_ref().expect("receiver used after recycle");
poll_fn(|cx| {
let mut replacement = None;
loop {
let mut state = chan.state.lock();
if let Some(frame) = state.queue.pop_front() {
drop(state);
drop(replacement);
return Poll::Ready(Some(frame));
}
if state.closed {
drop(state);
drop(replacement);
return Poll::Ready(None);
}
if state
.waker
.as_ref()
.is_some_and(|waker| waker.will_wake(cx.waker()))
{
drop(state);
drop(replacement);
return Poll::Pending;
}
if let Some(replacement) = replacement.take() {
let previous = state.waker.replace(replacement);
drop(state);
drop(previous);
return Poll::Pending;
}
drop(state);
replacement = Some(cx.waker().clone());
}
})
.await
}
#[cfg(test)]
pub(crate) fn try_recv(&mut self) -> Option<StreamFrame> {
let chan = self.chan.as_ref().expect("receiver used after recycle");
chan.state.lock().queue.pop_front()
}
pub(crate) fn recycle(&mut self) {
let Some(mut chan) = self.chan.take() else {
return;
};
let Some(chan_mut) = Arc::get_mut(&mut chan) else {
return;
};
let state = chan_mut.state.get_mut();
state.queue.clear();
state.waker = None;
state.closed = false;
POOL.with(|pool| {
let mut pool = pool.borrow_mut();
if pool.len() < MAX_POOLED {
pool.push(chan);
}
});
}
}
thread_local! {
static POOL: RefCell<Vec<Arc<StreamChannel>>> = const { RefCell::new(Vec::new()) };
}
pub(crate) fn acquire(capacity: usize) -> (StreamSender, StreamChannelReceiver) {
let pooled = POOL.with(|pool| {
let mut pool = pool.borrow_mut();
match pool.last() {
Some(chan) if chan.capacity == capacity => pool.pop(),
_ => None,
}
});
let chan = pooled.unwrap_or_else(|| {
Arc::new(StreamChannel {
state: Mutex::new(ChannelState {
queue: VecDeque::with_capacity(capacity),
waker: None,
closed: false,
}),
capacity,
})
});
(
StreamSender { chan: chan.clone() },
StreamChannelReceiver { chan: Some(chan) },
)
}
#[cfg(test)]
mod tests {
use super::*;
use nylon_ring::NrStatus;
use std::future::Future;
struct ThreadWaker(std::thread::Thread);
impl std::task::Wake for ThreadWaker {
fn wake(self: Arc<Self>) {
self.0.unpark();
}
}
fn block_on<F: Future>(fut: F) -> F::Output {
let mut fut = std::pin::pin!(fut);
let waker = Waker::from(Arc::new(ThreadWaker(std::thread::current())));
let mut cx = std::task::Context::from_waker(&waker);
loop {
match fut.as_mut().poll(&mut cx) {
Poll::Ready(value) => return value,
Poll::Pending => std::thread::park(),
}
}
}
fn frame(byte: u8) -> StreamFrame {
StreamFrame {
status: NrStatus::Ok,
data: vec![byte],
}
}
fn drain_pool() {
POOL.with(|pool| pool.borrow_mut().clear());
}
#[test]
fn sender_drop_closes_after_buffered_frames() {
drain_pool();
let (tx, mut rx) = acquire(4);
assert!(tx.try_send(frame(1)).is_ok());
assert!(tx.try_send(frame(2)).is_ok());
drop(tx);
block_on(async {
assert_eq!(rx.recv().await.unwrap().data, vec![1]);
assert_eq!(rx.recv().await.unwrap().data, vec![2]);
assert!(rx.recv().await.is_none());
assert!(rx.recv().await.is_none(), "closed channel must stay closed");
});
}
#[test]
fn full_queue_reports_backpressure_and_returns_frame() {
drain_pool();
let (tx, mut rx) = acquire(1);
assert!(tx.try_send(frame(1)).is_ok());
assert_eq!(
tx.try_send(frame(2))
.expect_err("second frame must not fit")
.data,
vec![2]
);
assert_eq!(rx.try_recv().unwrap().data, vec![1]);
assert!(tx.try_send(frame(3)).is_ok());
}
#[test]
fn recycle_clears_stale_state_and_reuses_allocation() {
drain_pool();
let (tx, mut rx) = acquire(4);
assert!(tx.try_send(frame(9)).is_ok());
drop(tx);
let first_ptr = Arc::as_ptr(rx.chan.as_ref().unwrap());
rx.recycle();
let (tx, mut rx) = acquire(4);
assert_eq!(Arc::as_ptr(rx.chan.as_ref().unwrap()), first_ptr);
assert!(rx.try_recv().is_none(), "stale frame leaked across reuse");
assert!(tx.try_send(frame(1)).is_ok());
block_on(async {
assert_eq!(rx.recv().await.unwrap().data, vec![1]);
});
drop(tx);
block_on(async {
assert!(rx.recv().await.is_none());
});
}
#[test]
fn recycle_refuses_shared_channel_and_capacity_mismatch() {
drain_pool();
let (tx, mut rx) = acquire(4);
rx.recycle();
assert_eq!(
POOL.with(|pool| pool.borrow().len()),
0,
"shared channel must not be pooled"
);
drop(tx);
let (tx, mut rx) = acquire(4);
drop(tx);
rx.recycle();
let (_tx, rx3) = acquire(8);
assert_eq!(rx3.chan.as_ref().unwrap().capacity, 8);
}
#[test]
fn cross_thread_send_wakes_receiver() {
drain_pool();
let (tx, mut rx) = acquire(4);
let sender = std::thread::spawn(move || {
for value in 0..3u8 {
let mut pending = frame(value);
loop {
match tx.try_send(pending) {
Ok(()) => break,
Err(rejected) => {
pending = rejected;
std::thread::yield_now();
}
}
}
}
});
block_on(async {
for value in 0..3u8 {
assert_eq!(rx.recv().await.unwrap().data, vec![value]);
}
assert!(rx.recv().await.is_none());
});
sender.join().unwrap();
}
}