use crate::msg::{addr::Addr, cmsg};
use core::fmt;
use s2n_quic_core::inet::ExplicitCongestionNotification;
use std::{
io::IoSliceMut,
marker::PhantomData,
ptr::NonNull,
sync::{
atomic::{AtomicUsize, Ordering},
Arc,
},
};
use tracing::trace;
pub(super) trait FreeList: 'static + Send + Sync {
fn free(&self, descriptor: Descriptor) -> Option<Box<dyn 'static + Send>>;
}
pub(super) struct Descriptor {
ptr: NonNull<DescriptorInner>,
phantom: PhantomData<DescriptorInner>,
}
impl Descriptor {
#[inline]
pub(super) fn new(ptr: NonNull<DescriptorInner>) -> Self {
Self {
ptr,
phantom: PhantomData,
}
}
#[inline]
pub(super) unsafe fn drop_in_place(&self) {
core::ptr::drop_in_place(self.ptr.as_ptr());
}
#[cfg(debug_assertions)]
pub(super) fn as_usize(&self) -> usize {
self.ptr.as_ptr().addr()
}
#[inline]
pub(super) fn id(&self) -> u32 {
self.inner().id
}
#[inline]
fn inner(&self) -> &DescriptorInner {
unsafe { self.ptr.as_ref() }
}
#[inline]
fn addr(&self) -> &Addr {
unsafe { self.inner().address.as_ref() }
}
#[inline]
fn data(&self) -> NonNull<u8> {
self.inner().payload
}
#[inline]
unsafe fn into_filled(self, len: u16, ecn: ExplicitCongestionNotification) -> Filled {
let inner = self.inner();
trace!(fill = inner.id, len, ?ecn);
debug_assert!(len <= inner.capacity);
inner.references.store(1, Ordering::Relaxed);
Filled {
desc: self,
offset: 0,
len,
ecn,
}
}
#[inline]
fn clone_filled(&self) -> Self {
let inner = self.inner();
inner.references.fetch_add(1, Ordering::Relaxed);
trace!(clone = inner.id);
Self {
ptr: self.ptr,
phantom: PhantomData,
}
}
#[inline]
unsafe fn drop_filled(&self) {
let inner = self.inner();
let desc_ref = inner.references.fetch_sub(1, Ordering::Release);
debug_assert_ne!(desc_ref, 0, "reference count underflow");
if desc_ref != 1 {
trace!(drop_desc_ref = inner.id);
return;
}
core::sync::atomic::fence(Ordering::Acquire);
let storage = inner.free(self);
trace!(free_desc = inner.id, state = %"filled");
drop(storage);
}
#[inline]
unsafe fn drop_unfilled(&self) {
let inner = self.inner();
let storage = inner.free(self);
trace!(free_desc = inner.id, state = %"unfilled");
let _ = inner;
drop(storage);
}
}
unsafe impl Send for Descriptor {}
unsafe impl Sync for Descriptor {}
pub(super) struct DescriptorInner {
id: u32,
capacity: u16,
address: NonNull<Addr>,
payload: NonNull<u8>,
references: AtomicUsize,
free_list: Arc<dyn FreeList>,
}
impl DescriptorInner {
pub(super) unsafe fn new(
id: u32,
capacity: u16,
address: NonNull<Addr>,
payload: NonNull<u8>,
free_list: Arc<dyn FreeList>,
) -> Self {
Self {
id,
capacity,
address,
payload,
references: AtomicUsize::new(0),
free_list,
}
}
#[inline]
unsafe fn free(&self, desc: &Descriptor) -> Option<Box<dyn 'static + Send>> {
debug_assert_eq!(desc.inner().references.load(Ordering::Relaxed), 0);
self.free_list.free(Descriptor {
ptr: desc.ptr,
phantom: PhantomData,
})
}
}
pub struct Unfilled {
desc: Option<Descriptor>,
}
impl fmt::Debug for Unfilled {
#[expect(
clippy::expect_used,
reason = "desc is always Some except transiently during recv_with/drop, so it is present here"
)]
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let desc = self.desc.as_ref().expect("invalid state");
f.debug_struct("Unfilled").field("id", &desc.id()).finish()
}
}
impl Unfilled {
#[inline]
pub(super) fn from_descriptor(desc: Descriptor) -> Self {
Self { desc: Some(desc) }
}
#[inline]
#[allow(
clippy::unwrap_in_result,
reason = "desc is always Some except transiently during recv_with/drop, so it is present here"
)]
pub fn recv_with<F, E>(mut self, f: F) -> Result<Segments, (Self, E)>
where
F: FnOnce(&mut Addr, &mut cmsg::Receiver, IoSliceMut) -> Result<usize, E>,
{
let desc = self.desc.take().expect("invalid state");
let inner = desc.inner();
let addr = unsafe { &mut *inner.address.as_ptr() };
let capacity = inner.capacity as usize;
let data = unsafe {
core::slice::from_raw_parts_mut(inner.payload.as_ptr(), capacity)
};
let iov = IoSliceMut::new(data);
let mut cmsg = cmsg::Receiver::default();
let len = match f(addr, &mut cmsg, iov) {
Ok(len) => {
debug_assert!(len <= capacity);
len.min(capacity) as u16
}
Err(err) => {
let unfilled = Self { desc: Some(desc) };
return Err((unfilled, err));
}
};
let desc = unsafe {
desc.into_filled(len, cmsg.ecn())
};
let segments = Segments {
descriptor: Some(desc),
segment_len: cmsg.segment_len(),
};
Ok(segments)
}
}
impl Drop for Unfilled {
#[inline]
fn drop(&mut self) {
if let Some(desc) = self.desc.take() {
unsafe {
desc.drop_unfilled();
}
}
}
}
pub struct Filled {
desc: Descriptor,
offset: u16,
len: u16,
ecn: ExplicitCongestionNotification,
}
impl fmt::Debug for Filled {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let alt = f.alternate();
let mut s = f.debug_struct("Filled");
s.field("id", &self.desc.id())
.field("remote_address", &self.remote_address().get())
.field("ecn", &self.ecn);
if alt {
s.field("payload", &self.payload());
} else {
s.field("payload_len", &self.len);
}
s.finish()
}
}
impl Filled {
#[inline]
pub fn ecn(&self) -> ExplicitCongestionNotification {
self.ecn
}
#[inline]
pub fn len(&self) -> u16 {
self.len
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len == 0
}
#[inline]
pub fn remote_address(&self) -> &Addr {
self.desc.addr()
}
#[inline]
pub fn payload(&self) -> &[u8] {
unsafe {
let ptr = self.desc.data().as_ptr().add(self.offset as _);
let len = self.len as usize;
core::slice::from_raw_parts(ptr, len)
}
}
#[inline]
pub fn payload_mut(&mut self) -> &mut [u8] {
unsafe {
let ptr = self.desc.data().as_ptr().add(self.offset as _);
let len = self.len as usize;
core::slice::from_raw_parts_mut(ptr, len)
}
}
#[must_use = "consider Filled::advance if you don't need the other half"]
#[inline]
pub fn split_to(&mut self, at: u16) -> Self {
assert!(at <= self.len);
let offset = self.offset;
let ecn = self.ecn;
self.offset += at;
self.len -= at;
let desc = self.desc.clone_filled();
Self {
desc,
offset,
len: at,
ecn,
}
}
#[inline]
pub fn truncate(&mut self, len: u16) {
self.len = len.min(self.len);
}
#[inline]
pub fn advance(&mut self, len: u16) {
assert!(len <= self.len);
self.offset += len;
self.len -= len;
}
}
impl Drop for Filled {
#[inline]
fn drop(&mut self) {
unsafe {
self.desc.drop_filled()
}
}
}
pub struct Segments {
descriptor: Option<Filled>,
segment_len: u16,
}
impl Iterator for Segments {
type Item = Filled;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
if self.segment_len == 0 {
return self.descriptor.take();
}
let descriptor = self.descriptor.as_mut()?;
if descriptor.len() > self.segment_len {
return Some(descriptor.split_to(self.segment_len as _));
}
self.descriptor.take()
}
}