use crate::message::{self, Message};
use alloc::sync::Arc;
use core::{
mem::size_of,
ptr::NonNull,
sync::atomic::AtomicU32,
task::{Context, Poll},
};
use s2n_quic_core::{
assume,
sync::{
atomic_waker,
cursor::{self, Cursor},
CachePadded,
},
};
const CURSOR_SIZE: usize = size_of::<CachePadded<AtomicU32>>();
const PRODUCER_OFFSET: usize = 0;
const CONSUMER_OFFSET: usize = CURSOR_SIZE;
const DATA_OFFSET: usize = CURSOR_SIZE * 2;
pub fn pair<T: Message>(entries: u32, payload_len: u32) -> (Producer<T>, Consumer<T>) {
let storage = T::alloc(entries, payload_len, DATA_OFFSET);
let storage = Arc::new(storage);
let ptr = storage.as_ptr();
let wakers = atomic_waker::pair();
let drop_wakers = atomic_waker::pair();
let consumer = Consumer {
cursor: unsafe { builder(ptr, entries).build_consumer() },
wakers: wakers.0,
drop_waker: drop_wakers.0,
storage: storage.clone(),
};
let producer = Producer {
cursor: unsafe { builder(ptr, entries).build_producer() },
wakers: wakers.1,
drop_waker: drop_wakers.1,
storage,
};
(producer, consumer)
}
pub struct Consumer<T: Message> {
cursor: Cursor<T>,
wakers: atomic_waker::Handle,
drop_waker: atomic_waker::Handle,
#[allow(dead_code)]
storage: Arc<message::Storage>,
}
unsafe impl<T: Message> Send for Consumer<T> {}
unsafe impl<T: Message> Sync for Consumer<T> {}
impl<T: Message> Consumer<T> {
#[inline]
pub fn acquire(&mut self, watermark: u32) -> u32 {
self.cursor.acquire_consumer(watermark)
}
pub fn register_drop_waker(&mut self, cx: &mut Context) {
self.drop_waker.register(cx.waker())
}
#[inline]
pub fn poll_acquire(&mut self, watermark: u32, cx: &mut Context) -> Poll<u32> {
macro_rules! try_acquire {
() => {{
let count = self.acquire(watermark);
if count > 0 {
return Poll::Ready(count);
}
}};
}
try_acquire!();
self.wakers.register(cx.waker());
try_acquire!();
Poll::Pending
}
#[inline]
pub fn release(&mut self, release_len: u32) {
self.release_no_wake(release_len);
self.wake();
}
#[inline]
pub fn release_no_wake(&mut self, release_len: u32) {
if release_len == 0 {
return;
}
debug_assert!(
release_len <= self.cursor.cached_consumer_len(),
"cannot release more messages than acquired"
);
unsafe {
sync_ring_regions::<_, false>(&self.cursor, release_len, replicate_payload_len);
}
self.cursor.release_consumer(release_len);
}
#[inline]
pub fn wake(&self) {
self.wakers.wake()
}
#[inline]
pub fn data(&mut self) -> &mut [T] {
let idx = self.cursor.cached_consumer();
let len = self.cursor.cached_consumer_len();
let ptr = self.cursor.data_ptr();
unsafe {
let ptr = ptr.as_ptr().add(idx as _);
core::slice::from_raw_parts_mut(ptr, len as _)
}
}
#[inline]
pub fn is_open(&self) -> bool {
self.wakers.is_open()
}
}
pub struct Producer<T: Message> {
cursor: Cursor<T>,
wakers: atomic_waker::Handle,
drop_waker: atomic_waker::Handle,
#[allow(dead_code)]
storage: Arc<message::Storage>,
}
unsafe impl<T: Message> Send for Producer<T> {}
unsafe impl<T: Message> Sync for Producer<T> {}
impl<T: Message> Producer<T> {
#[inline]
pub fn acquire(&mut self, watermark: u32) -> u32 {
self.cursor.acquire_producer(watermark)
}
pub fn register_drop_waker(&mut self, cx: &mut Context) {
self.drop_waker.register(cx.waker())
}
#[inline]
pub fn poll_acquire(&mut self, watermark: u32, cx: &mut Context) -> Poll<u32> {
macro_rules! try_acquire {
() => {{
let count = self.acquire(watermark);
if count > 0 {
return Poll::Ready(count);
}
}};
}
try_acquire!();
self.wakers.register(cx.waker());
try_acquire!();
Poll::Pending
}
#[inline]
#[allow(dead_code)] pub fn release(&mut self, release_len: u32) {
if release_len == 0 {
return;
}
self.release_no_wake(release_len);
self.wake();
}
#[inline]
pub fn release_no_wake(&mut self, release_len: u32) {
if release_len == 0 {
return;
}
debug_assert!(
release_len <= self.cursor.cached_producer_len(),
"cannot release more messages than acquired"
);
unsafe {
sync_ring_regions::<_, true>(&self.cursor, release_len, replicate);
}
self.cursor.release_producer(release_len);
}
#[inline]
pub fn wake(&self) {
self.wakers.wake()
}
#[inline]
pub fn data(&mut self) -> &mut [T] {
let idx = self.cursor.cached_producer();
let len = self.cursor.cached_producer_len();
let ptr = self.cursor.data_ptr();
unsafe {
let ptr = ptr.as_ptr().add(idx as _);
core::slice::from_raw_parts_mut(ptr, len as _)
}
}
#[inline]
pub fn is_open(&self) -> bool {
self.wakers.is_open()
}
}
#[inline]
unsafe fn replicate<T: Message>(src: *mut T, dest: *mut T, len: usize) {
debug_assert_ne!(len, 0);
#[cfg(debug_assertions)]
{
let src_slice = core::slice::from_raw_parts(src, len as _);
let dest_slice = core::slice::from_raw_parts(dest, len as _);
for (src_message, dest_message) in src_slice.iter().zip(dest_slice) {
T::validate_replication(src_message, dest_message);
}
}
core::ptr::copy_nonoverlapping(src, dest, len as _);
}
#[inline]
unsafe fn replicate_payload_len<T: Message>(src: *mut T, dest: *mut T, len: usize) {
let src_slice = core::slice::from_raw_parts_mut(src, len as _);
let dest_slice = core::slice::from_raw_parts_mut(dest, len as _);
for (src_message, dest_message) in src_slice.iter_mut().zip(dest_slice) {
dest_message.set_payload_len(src_message.payload_len());
}
}
unsafe fn sync_ring_regions<T: Message, const PRODUCER: bool>(
cursor: &Cursor<T>,
release_len: u32,
f: unsafe fn(src: *mut T, dest: *mut T, len: usize),
) {
let idx = if PRODUCER {
cursor.cached_producer()
} else {
cursor.cached_consumer()
};
let ring_size = cursor.capacity();
assume!(ring_size > idx, "idx should never exceed the ring size");
let max_possible_replications = ring_size - idx;
let replication_count = max_possible_replications.min(release_len);
assume!(
replication_count != 0,
"we should always be releasing at least 1 item"
);
let primary = cursor.data_ptr().as_ptr().add(idx as _);
let secondary = primary.add(ring_size as _);
f(primary, secondary, replication_count as _);
assume!(
idx.checked_add(release_len).is_some(),
"overflow amount should not exceed u32::MAX"
);
assume!(
idx + release_len < ring_size * 2,
"overflow amount should not extend beyond the secondary replica"
);
let overflow_amount = (idx + release_len).checked_sub(ring_size).filter(|v| {
*v > 0
});
if let Some(replication_count) = overflow_amount {
let primary = cursor.data_ptr().as_ptr();
let secondary = primary.add(ring_size as _);
f(secondary, primary, replication_count as _);
}
}
#[inline]
unsafe fn builder<T: Message>(ptr: *mut u8, size: u32) -> cursor::Builder<T> {
let producer = ptr.add(PRODUCER_OFFSET) as *mut _;
let producer = NonNull::new(producer).unwrap();
let consumer = ptr.add(CONSUMER_OFFSET) as *mut _;
let consumer = NonNull::new(consumer).unwrap();
let data = ptr.add(DATA_OFFSET) as *mut _;
let data = NonNull::new(data).unwrap();
cursor::Builder {
producer,
consumer,
data,
size,
}
}
#[cfg(test)]
mod tests {
use super::*;
use bolero::check;
use s2n_quic_core::{
inet::{ExplicitCongestionNotification, SocketAddress},
path::{Handle as _, LocalAddress, RemoteAddress},
};
#[cfg(not(kani))]
type Counts = Vec<u32>;
#[cfg(kani)]
type Counts = s2n_quic_core::testing::InlineVec<u32, 2>;
macro_rules! replication_test {
($name:ident, $msg:ty) => {
#[test]
#[cfg_attr(kani, kani::proof, kani::solver(cadical), kani::unwind(3))]
#[cfg(any(not(kani), kani_slow))] fn $name() {
check!().with_type::<Counts>().for_each(|counts| {
let entries = if cfg!(kani) { 2 } else { 16 };
let payload_len = if cfg!(kani) { 2 } else { 128 };
let (mut producer, mut consumer) = pair::<$msg>(entries, payload_len);
let mut counter = 0;
for count in counts.iter().copied() {
let count = producer.acquire(count);
for entry in &mut producer.data()[..count as usize] {
unsafe {
entry.set_payload_len(counter);
}
counter += 1;
}
producer.release(count);
#[cfg(kani)]
let ids_to_check = {
let idx: u32 = kani::any();
kani::assume(idx < entries);
idx..idx + 1
};
#[cfg(not(kani))]
let ids_to_check = 0..entries;
for idx in ids_to_check {
let ptr = producer.cursor.data_ptr().as_ptr();
unsafe {
let primary = &*ptr.add(idx as _);
let secondary = &*ptr.add((idx + entries) as _);
assert_eq!(primary.payload_len(), secondary.payload_len());
}
}
let count = consumer.acquire(count);
consumer.release(count);
}
});
}
};
}
replication_test!(simple_replication, crate::message::simple::Message);
replication_test!(testing_replication, crate::io::testing::message::Message);
#[cfg(s2n_quic_platform_socket_msg)]
replication_test!(msg_replication, crate::message::msg::Message);
#[cfg(s2n_quic_platform_socket_mmsg)]
replication_test!(mmsg_replication, crate::message::mmsg::Message);
macro_rules! send_recv_test {
($name:ident, $msg:ty) => {
#[test]
fn $name() {
check!().with_type::<Counts>().for_each(|counts| {
let entries = if cfg!(miri) { 2 } else { 16 };
let payload_len = if cfg!(miri) { 4 } else { 128 };
let (mut producer, mut consumer) = pair::<$msg>(entries, payload_len);
let mut tx_counter = 0u32;
let mut rx_counter = 0u32;
let local_address = LocalAddress::from(SocketAddress::default());
for count in counts.iter().copied() {
let count = producer.acquire(count);
for entry in &mut producer.data()[..count as usize] {
unsafe {
entry.reset(payload_len as _);
}
let mut remote_address = SocketAddress::default();
remote_address.set_port(tx_counter as _);
let remote_address = RemoteAddress::from(remote_address);
let handle =
<$msg as Message>::Handle::from_remote_address(remote_address);
let ecn = ExplicitCongestionNotification::new(tx_counter as _);
let payload = tx_counter.to_le_bytes();
let msg = (handle, ecn, &payload[..]);
entry.tx_write(msg).unwrap();
tx_counter += 1;
}
producer.release(count);
let count = consumer.acquire(count);
for entry in consumer.data() {
let message = entry.rx_read(&local_address).unwrap();
message.for_each(|header, payload| {
if <$msg>::SUPPORTS_ECN {
let ecn = ExplicitCongestionNotification::new(rx_counter as _);
assert_eq!(header.ecn, ecn);
}
let counter: &[u8; 4] = (&*payload).try_into().unwrap();
let counter = u32::from_le_bytes(*counter);
assert_eq!(counter, rx_counter);
rx_counter += 1;
});
}
consumer.release(count);
}
});
}
};
}
send_recv_test!(simple_send_recv, crate::message::simple::Message);
send_recv_test!(testing_send_recv, crate::io::testing::message::Message);
#[cfg(s2n_quic_platform_socket_msg)]
send_recv_test!(msg_send_recv, crate::message::msg::Message);
#[cfg(s2n_quic_platform_socket_mmsg)]
send_recv_test!(mmsg_send_recv, crate::message::mmsg::Message);
macro_rules! consumer_modifications_test {
($name:ident, $msg:ty) => {
#[test]
fn $name() {
check!().with_type::<u32>().for_each(|&count| {
let entries = if cfg!(kani) { 2 } else { 16 };
let payload_len = if cfg!(kani) { 2 } else { 128 };
let count = count % entries;
let (mut producer, mut consumer) = pair::<$msg>(entries, payload_len);
producer.acquire(u32::MAX);
for entry in &mut producer.data()[..count as usize] {
unsafe {
entry.set_payload_len(100);
}
}
producer.release(count);
let count = consumer.acquire(u32::MAX);
for entry in &mut consumer.data()[..count as usize] {
unsafe {
entry.reset(payload_len as usize);
}
}
consumer.release(count);
producer.acquire(u32::MAX);
let s = producer.data();
for entry in s {
assert_eq!(entry.payload_len(), payload_len as usize);
}
});
}
};
}
consumer_modifications_test!(simple_rx_modifications, crate::message::simple::Message);
consumer_modifications_test!(
testing_rx_modifications,
crate::io::testing::message::Message
);
#[cfg(s2n_quic_platform_socket_msg)]
consumer_modifications_test!(msg_rx_modifications, crate::message::msg::Message);
#[cfg(s2n_quic_platform_socket_mmsg)]
consumer_modifications_test!(mmsg_rx_modifications, crate::message::mmsg::Message);
}