use std::alloc::{Layout, alloc_zeroed, dealloc, handle_alloc_error};
use std::cell::{Cell, RefCell};
use std::collections::VecDeque;
use std::io;
use std::net::{IpAddr, SocketAddr, SocketAddrV6, UdpSocket};
use std::os::fd::AsRawFd;
use std::ptr::NonNull;
use std::rc::Rc;
use std::sync::atomic::{AtomicU16, Ordering};
use std::task::Poll;
use io_uring::{cqueue, opcode, types};
use crate::Error;
use crate::metrics::Counters;
use crate::shared::{Cqe, Op, Shared};
use crate::worker::Owner;
const CONTROL_LEN: usize = 64;
const NAME_LEN: usize = std::mem::size_of::<libc::sockaddr_storage>();
const RECV_OVERHEAD: usize = 16 + NAME_LEN + CONTROL_LEN;
const MAX_RECV: usize = 64 * 1024;
const MAX_GSO_SEGMENTS: usize = 64;
const MAX_RX_BUFFERS: u16 = 1 << 15;
const INITIAL_RX_BUFFERS: u16 = 16;
const INITIAL_TX_BUFFERS: u16 = 64;
fn grown(len: usize, max: u16) -> Option<u16> {
let max = usize::from(max);
match len < max {
true => Some(len.saturating_mul(2).clamp(1, max) as u16),
false => None,
}
}
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct Config {
pub gro: bool,
pub gso: bool,
pub multishot: bool,
pub rx_buffers_max: u16,
pub rx_buffer_len: usize,
pub tx_buffers_max: u16,
pub tx_buffer_len: usize,
}
impl Default for Config {
fn default() -> Self {
Self {
gro: true,
gso: true,
multishot: true,
rx_buffers_max: 256,
rx_buffer_len: MAX_RECV + RECV_OVERHEAD,
tx_buffers_max: 1024,
tx_buffer_len: 64 * 1024,
}
}
}
struct RxBuf {
data: Box<[u8]>,
outstanding: usize,
kernel_done: bool,
claimed: bool,
}
struct Queued {
bid: u16,
start: usize,
len: usize,
from: SocketAddr,
stride: usize,
ecn: Option<Ecn>,
}
#[repr(u8)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Ecn {
Ect0 = 0b10,
Ect1 = 0b01,
Ce = 0b11,
}
impl Ecn {
fn from_bits(bits: u8) -> Option<Self> {
match bits & 0b11 {
0b10 => Some(Self::Ect0),
0b01 => Some(Self::Ect1),
0b11 => Some(Self::Ce),
_ => None,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct Transmit {
pub to: SocketAddr,
pub len: usize,
pub segment: usize,
pub ecn: Option<Ecn>,
}
fn recycle_if_idle(rx: &mut Rx, bid: u16) -> bool {
let buf = &mut rx.bufs[bid as usize];
if buf.outstanding > 0 {
return false;
}
if buf.claimed {
buf.claimed = false;
return true;
}
if !buf.kernel_done {
return false;
}
buf.kernel_done = false;
let addr = buf.data.as_mut_ptr();
let len = buf.data.len();
if let Some(ring) = &mut rx.ring {
ring.add(bid, addr, len);
ring.publish();
}
true
}
fn should_grow(rx: &Rx, multishot: bool) -> bool {
if rx.starved {
return true;
}
match multishot {
true => !rx.bufs.iter().any(|buf| !buf.kernel_done),
false => !rx.bufs.iter().any(|buf| !buf.claimed && buf.outstanding == 0),
}
}
fn grow_rx(rx: &mut Rx, config: &Config) -> bool {
let Some(target) = grown(rx.bufs.len(), config.rx_buffers_max) else {
return false;
};
while rx.bufs.len() < usize::from(target) {
let bid = rx.bufs.len() as u16;
rx.bufs.push(RxBuf {
data: vec![0u8; config.rx_buffer_len].into_boxed_slice(),
outstanding: 0,
kernel_done: false,
claimed: false,
});
if let Some(ring) = &mut rx.ring {
let buf = &mut rx.bufs[bid as usize];
let addr = buf.data.as_mut_ptr();
let len = buf.data.len();
ring.add(bid, addr, len);
}
}
if let Some(ring) = &mut rx.ring {
ring.publish();
}
true
}
fn grow_tx(tx: &mut Tx, config: &Config) -> bool {
let Some(target) = grown(tx.bufs.len(), config.tx_buffers_max) else {
return false;
};
while tx.bufs.len() < usize::from(target) {
tx.free.push(tx.bufs.len() as u16);
tx.bufs.push(TxSlot::new(config.tx_buffer_len));
}
true
}
struct BufRing {
ptr: NonNull<types::BufRingEntry>,
layout: Layout,
mask: u16,
tail: u16,
}
impl BufRing {
fn new(entries: u16) -> Self {
let layout = Layout::from_size_align(entries as usize * std::mem::size_of::<types::BufRingEntry>(), 4096)
.expect("buffer ring layout");
let ptr = unsafe { alloc_zeroed(layout) };
let ptr = NonNull::new(ptr.cast::<types::BufRingEntry>()).unwrap_or_else(|| handle_alloc_error(layout));
Self {
ptr,
layout,
mask: entries - 1,
tail: 0,
}
}
fn add(&mut self, bid: u16, addr: *mut u8, len: usize) {
let index = (self.tail & self.mask) as usize;
let entry = unsafe { &mut *self.ptr.as_ptr().add(index) };
entry.set_addr(addr as u64);
entry.set_len(len as u32);
entry.set_bid(bid);
self.tail = self.tail.wrapping_add(1);
}
fn publish(&self) {
let tail = unsafe { types::BufRingEntry::tail(self.ptr.as_ptr()) }.cast_mut();
unsafe { AtomicU16::from_ptr(tail) }.store(self.tail, Ordering::Release);
}
}
impl Drop for BufRing {
fn drop(&mut self) {
unsafe { dealloc(self.ptr.as_ptr().cast(), self.layout) };
}
}
struct Rx {
bufs: Vec<RxBuf>,
ring: Option<BufRing>,
hdr: Box<libc::msghdr>,
queue: VecDeque<Queued>,
waiters: kio::WaiterList,
armed: Option<u64>,
starved: bool,
error: Option<i32>,
}
struct Tx {
bufs: Vec<TxSlot>,
free: Vec<u16>,
waiters: kio::WaiterList,
stalled: bool,
error: Option<i32>,
}
struct TxSlot {
data: Box<[u8]>,
headers: Vec<SendHdr>,
in_flight: usize,
}
impl TxSlot {
fn new(len: usize) -> Self {
Self {
data: vec![0u8; len].into_boxed_slice(),
headers: Vec::new(),
in_flight: 0,
}
}
}
pub enum Bound {
Lone(UdpSocket),
Member(moq_sock::shard::Socket),
}
impl From<UdpSocket> for Bound {
fn from(socket: UdpSocket) -> Self {
Self::Lone(socket)
}
}
impl From<moq_sock::shard::Socket> for Bound {
fn from(member: moq_sock::shard::Socket) -> Self {
Self::Member(member)
}
}
pub(crate) struct SockShared {
io: UdpSocket,
owner: Owner,
shard: Option<moq_sock::shard::Shard>,
metrics: std::sync::Arc<Counters>,
config: Config,
bgid: u16,
closed: Cell<bool>,
rx: RefCell<Rx>,
tx: RefCell<Tx>,
}
impl SockShared {
fn worker_gone(&self) -> bool {
self.owner.handle().is_none()
}
fn release_rx(self: &Rc<Self>, bid: u16) {
let mut rx = self.rx.borrow_mut();
rx.bufs[bid as usize].outstanding -= 1;
if !recycle_if_idle(&mut rx, bid) {
return;
}
if rx.armed.is_none() && rx.error.is_none() {
drop(rx);
if let Some(shared) = self.owner.upgrade() {
arm_recv(&shared, self);
}
}
}
fn release_tx(&self, id: u16) {
let mut tx = self.tx.borrow_mut();
debug_assert_eq!(tx.bufs[id as usize].in_flight, 0);
tx.free.push(id);
tx.stalled = false;
tx.waiters.wake();
}
fn stage_tx(&self, id: u16) {
self.tx.borrow_mut().bufs[id as usize].in_flight += 1;
}
fn complete_tx(&self, id: u16) {
let mut tx = self.tx.borrow_mut();
let slot = &mut tx.bufs[id as usize];
debug_assert!(slot.in_flight > 0);
slot.in_flight -= 1;
if slot.in_flight == 0 {
tx.free.push(id);
tx.stalled = false;
tx.waiters.wake();
}
}
fn fail_rx(&self, code: i32) {
let mut rx = self.rx.borrow_mut();
rx.error.get_or_insert(code);
rx.waiters.wake();
}
fn fail_tx(&self, code: i32) {
let mut tx = self.tx.borrow_mut();
tx.error.get_or_insert(code);
tx.waiters.wake();
}
}
impl Drop for SockShared {
fn drop(&mut self) {
if self.rx.borrow().ring.is_some()
&& let Some(shared) = self.owner.upgrade()
{
let ring = shared.ring.borrow_mut();
let _ = ring.submitter().unregister_buf_ring(self.bgid);
}
}
}
pub struct Socket {
shared: Rc<SockShared>,
}
impl Socket {
#[cfg(test)]
pub(crate) fn downgrade(&self) -> std::rc::Weak<SockShared> {
Rc::downgrade(&self.shared)
}
pub(crate) fn bind(shared: &Rc<Shared>, bound: Bound, config: Config) -> Result<Self, Error> {
let (io, shard) = match bound {
Bound::Lone(io) => (io, None),
Bound::Member(member) => {
let shard = member.shard();
(member.into_inner(), Some(shard))
}
};
let floor = if config.gro { MAX_RECV + RECV_OVERHEAD } else { 2048 };
if config.rx_buffer_len < floor || config.rx_buffers_max == 0 || config.tx_buffers_max == 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"receive buffers must hold one worst-case receive ({floor} bytes) and both pools need at least one buffer"
),
)
.into());
}
if config.rx_buffers_max > MAX_RX_BUFFERS {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"receive pool holds at most {MAX_RX_BUFFERS} buffers, got {}",
config.rx_buffers_max
),
)
.into());
}
if config.gro {
set_option(&io, libc::SOL_UDP, libc::UDP_GRO)?;
}
if io.local_addr()?.is_ipv6() {
set_option(&io, libc::IPPROTO_IPV6, libc::IPV6_RECVTCLASS)?;
}
set_option(&io, libc::IPPROTO_IP, libc::IP_RECVTOS)?;
let rx_cap = config.rx_buffers_max.next_power_of_two();
let rx_count = INITIAL_RX_BUFFERS.min(config.rx_buffers_max);
let mut bufs = Vec::with_capacity(rx_count as usize);
for _ in 0..rx_count {
bufs.push(RxBuf {
data: vec![0u8; config.rx_buffer_len].into_boxed_slice(),
outstanding: 0,
kernel_done: false,
claimed: false,
});
}
let bgid = shared.next_bgid.get();
shared
.next_bgid
.set(bgid.checked_add(1).expect("buffer group ids exhausted"));
let ring = if config.multishot {
let mut ring = BufRing::new(rx_cap);
{
let io_ring = shared.ring.borrow_mut();
unsafe {
io_ring
.submitter()
.register_buf_ring_with_flags(ring.ptr.as_ptr() as u64, rx_cap, bgid, 0)
.map_err(Error::ring)?;
}
}
for (bid, buf) in bufs.iter_mut().enumerate() {
let addr = buf.data.as_mut_ptr();
let len = buf.data.len();
ring.add(bid as u16, addr, len);
}
ring.publish();
Some(ring)
} else {
None
};
let mut hdr: Box<libc::msghdr> = Box::new(unsafe { std::mem::zeroed() });
hdr.msg_namelen = NAME_LEN as libc::socklen_t;
hdr.msg_controllen = CONTROL_LEN;
let tx_count = INITIAL_TX_BUFFERS.min(config.tx_buffers_max);
let tx = Tx {
bufs: (0..tx_count).map(|_| TxSlot::new(config.tx_buffer_len)).collect(),
free: (0..tx_count).collect(),
waiters: kio::WaiterList::new(),
stalled: false,
error: None,
};
let sock = Rc::new(SockShared {
io,
owner: Owner::new(shared),
shard,
metrics: shared.metrics.clone(),
config,
bgid,
closed: Cell::new(false),
rx: RefCell::new(Rx {
bufs,
ring,
hdr,
queue: VecDeque::new(),
waiters: kio::WaiterList::new(),
armed: None,
starved: false,
error: None,
}),
tx: RefCell::new(tx),
});
arm_recv(shared, &sock);
if let Some(code) = sock.rx.borrow().error {
return Err(io::Error::from_raw_os_error(code).into());
}
Ok(Self { shared: sock })
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.shared.io.local_addr()
}
pub(crate) fn owner(&self) -> Owner {
self.shared.owner.clone()
}
pub(crate) fn shard(&self) -> Option<moq_sock::shard::Shard> {
self.shared.shard
}
pub fn poll_recv(&self, waiter: &kio::Waiter) -> Poll<io::Result<Packet>> {
let mut rx = self.shared.rx.borrow_mut();
if let Some(queued) = rx.queue.pop_front() {
let buf = &rx.bufs[queued.bid as usize];
let ptr = unsafe { NonNull::new_unchecked(buf.data.as_ptr().cast_mut().add(queued.start)) };
return Poll::Ready(Ok(Packet {
sock: self.shared.clone(),
bid: queued.bid,
ptr,
len: queued.len,
stride: queued.stride,
from: queued.from,
ecn: queued.ecn,
}));
}
if let Some(code) = rx.error {
return Poll::Ready(Err(io::Error::from_raw_os_error(code)));
}
if self.shared.worker_gone() {
return Poll::Ready(Err(Shared::gone_error()));
}
waiter.register(&mut rx.waiters);
Poll::Pending
}
pub async fn recv(&self) -> io::Result<Packet> {
kio::wait(|waiter| self.poll_recv(waiter)).await
}
pub fn poll_acquire(&self, waiter: &kio::Waiter) -> Poll<io::Result<TxBuf>> {
let mut tx = self.shared.tx.borrow_mut();
if let Some(code) = tx.error {
return Poll::Ready(Err(io::Error::from_raw_os_error(code)));
}
if self.shared.worker_gone() {
return Poll::Ready(Err(Shared::gone_error()));
}
if tx.free.is_empty() {
grow_tx(&mut tx, &self.shared.config);
}
if let Some(id) = tx.free.pop() {
let slot = &mut tx.bufs[id as usize];
let ptr = unsafe { NonNull::new_unchecked(slot.data.as_mut_ptr()) };
let cap = slot.data.len();
return Poll::Ready(Ok(TxBuf {
sock: self.shared.clone(),
id,
ptr,
cap,
armed: false,
}));
}
if !tx.stalled {
tx.stalled = true;
self.shared.metrics.tx_stalls.add(1);
}
waiter.register(&mut tx.waiters);
Poll::Pending
}
pub async fn acquire(&self) -> io::Result<TxBuf> {
kio::wait(|waiter| self.poll_acquire(waiter)).await
}
}
impl Drop for Socket {
fn drop(&mut self) {
self.shared.closed.set(true);
let rx = self.shared.rx.borrow();
if let (Some(key), Some(shared)) = (rx.armed, self.shared.owner.upgrade()) {
drop(rx);
let _ = shared.cancel(key);
}
}
}
impl std::fmt::Debug for Socket {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Socket")
.field("addr", &self.shared.io.local_addr())
.finish()
}
}
pub struct Packet {
sock: Rc<SockShared>,
bid: u16,
ptr: NonNull<u8>,
len: usize,
stride: usize,
from: SocketAddr,
ecn: Option<Ecn>,
}
impl Packet {
pub fn from(&self) -> SocketAddr {
self.from
}
pub fn ecn(&self) -> Option<Ecn> {
self.ecn
}
pub fn stride(&self) -> usize {
self.stride
}
pub fn payload(&self) -> &[u8] {
unsafe { std::slice::from_raw_parts(self.ptr.as_ptr(), self.len) }
}
pub fn payload_mut(&mut self) -> &mut [u8] {
unsafe { std::slice::from_raw_parts_mut(self.ptr.as_ptr(), self.len) }
}
pub fn segments(&mut self) -> impl Iterator<Item = &mut [u8]> {
let stride = self.stride;
self.payload_mut().chunks_mut(stride)
}
}
impl Drop for Packet {
fn drop(&mut self) {
self.sock.release_rx(self.bid);
}
}
impl std::fmt::Debug for Packet {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Packet")
.field("from", &self.from)
.field("len", &self.len)
.field("stride", &self.stride)
.field("ecn", &self.ecn)
.finish()
}
}
pub struct TxBuf {
sock: Rc<SockShared>,
id: u16,
ptr: NonNull<u8>,
cap: usize,
armed: bool,
}
impl TxBuf {
pub fn send(mut self, transmit: Transmit) -> io::Result<()> {
let Transmit { to, len, segment, ecn } = transmit;
if len == 0 || len > self.cap || segment == 0 || segment > usize::from(u16::MAX) {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"invalid send: {len} bytes in {segment} byte segments from a {} byte buffer",
self.cap
),
));
}
let shared = match self.sock.owner.upgrade() {
Some(shared) if !shared.stopped.get() => shared,
_ => return Err(Shared::gone_error()),
};
let segments = len.div_ceil(segment);
let limit = match self.sock.config.gso {
true => MAX_GSO_SEGMENTS,
false => shared.ring.borrow().params().sq_entries() as usize,
};
if segments > limit {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!("send of {segments} datagrams exceeds the {limit} one call may stage"),
));
}
self.armed = true;
let sock = self.sock.clone();
let base = self.ptr.as_ptr();
let headers = {
let mut tx = sock.tx.borrow_mut();
let headers = &mut tx.bufs[self.id as usize].headers;
if headers.len() < segments {
headers.resize_with(segments, SendHdr::zeroed);
}
unsafe { NonNull::new_unchecked(headers.as_mut_ptr()) }
};
let staging = Staging {
sock: sock.clone(),
id: self.id,
headers,
};
let one = SendOne {
to,
ecn,
segment: sock.config.gso.then_some(segment as u16),
};
if sock.config.gso {
send_one(&shared, &staging, 0, base, len, &one)?;
} else {
for index in 0..segments {
let offset = index * segment;
let chunk = segment.min(len - offset);
send_one(&shared, &staging, index, unsafe { base.add(offset) }, chunk, &one)?;
}
}
sock.metrics.tx_datagrams.add(segments as u64);
Ok(())
}
}
impl std::ops::Deref for TxBuf {
type Target = [u8];
fn deref(&self) -> &[u8] {
unsafe { std::slice::from_raw_parts(self.ptr.as_ptr(), self.cap) }
}
}
impl std::ops::DerefMut for TxBuf {
fn deref_mut(&mut self) -> &mut [u8] {
unsafe { std::slice::from_raw_parts_mut(self.ptr.as_ptr(), self.cap) }
}
}
impl Drop for TxBuf {
fn drop(&mut self) {
if !self.armed {
self.sock.release_tx(self.id);
}
}
}
impl std::fmt::Debug for TxBuf {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TxBuf").field("cap", &self.cap).finish()
}
}
#[repr(C, align(8))]
struct Control([u8; CONTROL_LEN]);
struct SendHdr {
hdr: libc::msghdr,
iov: libc::iovec,
name: libc::sockaddr_storage,
control: Control,
}
impl SendHdr {
fn zeroed() -> Self {
unsafe { std::mem::zeroed() }
}
}
struct Staging {
sock: Rc<SockShared>,
id: u16,
headers: NonNull<SendHdr>,
}
pub(crate) struct SendOp {
sock: Rc<SockShared>,
id: u16,
expect: usize,
}
impl Drop for SendOp {
fn drop(&mut self) {
self.sock.complete_tx(self.id);
}
}
struct SendOne {
to: SocketAddr,
ecn: Option<Ecn>,
segment: Option<u16>,
}
fn send_one(
shared: &Rc<Shared>,
staging: &Staging,
index: usize,
base: *mut u8,
len: usize,
one: &SendOne,
) -> io::Result<()> {
let SendOne { to, ecn, segment } = *one;
let hdr = unsafe { &mut *staging.headers.as_ptr().add(index) };
*hdr = SendHdr::zeroed();
hdr.iov = libc::iovec {
iov_base: base.cast(),
iov_len: len,
};
let name_len = encode_addr(to, &mut hdr.name);
hdr.hdr.msg_name = (&raw mut hdr.name).cast();
hdr.hdr.msg_namelen = name_len;
hdr.hdr.msg_iov = &raw mut hdr.iov;
hdr.hdr.msg_iovlen = 1;
unsafe {
let mut space = 0;
if segment.is_some() {
space += libc::CMSG_SPACE(std::mem::size_of::<u16>() as _) as usize;
}
if ecn.is_some() {
space += libc::CMSG_SPACE(std::mem::size_of::<libc::c_int>() as _) as usize;
}
if space > 0 {
hdr.hdr.msg_control = hdr.control.0.as_mut_ptr().cast();
hdr.hdr.msg_controllen = space;
}
let mut cmsg = libc::CMSG_FIRSTHDR(&hdr.hdr);
if let Some(segment) = segment {
(*cmsg).cmsg_level = libc::SOL_UDP;
(*cmsg).cmsg_type = libc::UDP_SEGMENT;
(*cmsg).cmsg_len = libc::CMSG_LEN(std::mem::size_of::<u16>() as _) as usize;
std::ptr::write_unaligned(libc::CMSG_DATA(cmsg).cast::<u16>(), segment);
cmsg = libc::CMSG_NXTHDR(&hdr.hdr, cmsg);
}
if let Some(ecn) = ecn {
let is_ipv4 = match to.ip() {
IpAddr::V4(_) => true,
IpAddr::V6(v6) => v6.to_ipv4_mapped().is_some(),
};
let (level, kind) = match is_ipv4 {
true => (libc::IPPROTO_IP, libc::IP_TOS),
false => (libc::IPPROTO_IPV6, libc::IPV6_TCLASS),
};
(*cmsg).cmsg_level = level;
(*cmsg).cmsg_type = kind;
(*cmsg).cmsg_len = libc::CMSG_LEN(std::mem::size_of::<libc::c_int>() as _) as usize;
std::ptr::write_unaligned(libc::CMSG_DATA(cmsg).cast::<libc::c_int>(), ecn as libc::c_int);
}
}
let hdr_ptr = &raw const hdr.hdr;
staging.sock.stage_tx(staging.id);
let key = shared.insert(Op::Send(SendOp {
sock: staging.sock.clone(),
id: staging.id,
expect: len,
}));
let entry = opcode::SendMsg::new(types::Fd(staging.sock.io.as_raw_fd()), hdr_ptr)
.build()
.user_data(key);
if let Err(err) = shared.push(&entry) {
shared.ops.borrow_mut().remove(key as usize);
return Err(err);
}
staging.sock.metrics.tx_sends.add(1);
Ok(())
}
pub(crate) fn arm_recv(shared: &Rc<Shared>, sock: &Rc<SockShared>) {
if sock.closed.get() || shared.stopped.get() {
return;
}
let mut rx = sock.rx.borrow_mut();
if rx.armed.is_some() || rx.error.is_some() {
return;
}
if should_grow(&rx, sock.config.multishot) {
rx.starved = false;
grow_rx(&mut rx, &sock.config);
}
let entry = if sock.config.multishot {
if !rx.bufs.iter().any(|buf| !buf.kernel_done) {
sock.metrics.rx_exhausted.add(1);
return;
}
let key = shared.insert(Op::Recv {
sock: sock.clone(),
one: None,
});
rx.armed = Some(key);
opcode::RecvMsgMulti::new(types::Fd(sock.io.as_raw_fd()), &*rx.hdr, sock.bgid)
.build()
.user_data(key)
} else {
let Some(bid) = rx
.bufs
.iter()
.position(|buf| !buf.claimed && buf.outstanding == 0)
.map(|bid| bid as u16)
else {
sock.metrics.rx_exhausted.add(1);
return;
};
rx.bufs[bid as usize].claimed = true;
let mut one: Box<OneshotRecv> = Box::new(unsafe { std::mem::zeroed() });
one.bid = bid;
one.iov = libc::iovec {
iov_base: rx.bufs[bid as usize].data.as_mut_ptr().cast(),
iov_len: rx.bufs[bid as usize].data.len(),
};
one.hdr.msg_name = (&raw mut one.name).cast();
one.hdr.msg_namelen = NAME_LEN as libc::socklen_t;
one.hdr.msg_iov = &raw mut one.iov;
one.hdr.msg_iovlen = 1;
one.hdr.msg_control = one.control.0.as_mut_ptr().cast();
one.hdr.msg_controllen = CONTROL_LEN;
let hdr_ptr = &raw mut one.hdr;
let key = shared.insert(Op::Recv {
sock: sock.clone(),
one: Some(one),
});
rx.armed = Some(key);
opcode::RecvMsg::new(types::Fd(sock.io.as_raw_fd()), hdr_ptr)
.build()
.user_data(key)
};
drop(rx);
if let Err(err) = shared.push(&entry) {
let key = sock.rx.borrow_mut().armed.take().expect("just armed");
shared.ops.borrow_mut().remove(key as usize);
sock.fail_rx(err.raw_os_error().unwrap_or(libc::EIO));
}
}
pub(crate) struct OneshotRecv {
hdr: libc::msghdr,
iov: libc::iovec,
name: libc::sockaddr_storage,
control: Control,
bid: u16,
}
pub(crate) fn on_recv(
shared: &Rc<Shared>,
sock: &Rc<SockShared>,
one: Option<Box<OneshotRecv>>,
cqe: Cqe,
terminal: bool,
) {
if terminal {
sock.rx.borrow_mut().armed = None;
}
if cqe.result < 0 {
let code = -cqe.result;
if let Some(one) = &one {
let mut rx = sock.rx.borrow_mut();
rx.bufs[one.bid as usize].claimed = false;
}
match code {
libc::ENOBUFS => {
sock.metrics.rx_enobufs.add(1);
sock.rx.borrow_mut().starved = true;
}
libc::ECANCELED => return,
_ => {
sock.fail_rx(code);
return;
}
}
arm_recv(shared, sock);
return;
}
let received = match one {
None => on_recv_multi(sock, cqe),
Some(one) => on_recv_oneshot(*one, cqe),
};
match received {
Ok((_, Some(queued))) => {
sock.metrics.rx_receives.add(1);
sock.metrics
.rx_datagrams
.add(queued.len.div_ceil(queued.stride.max(1)) as u64);
let mut rx = sock.rx.borrow_mut();
rx.bufs[queued.bid as usize].outstanding += 1;
rx.queue.push_back(queued);
rx.waiters.wake();
}
Ok((bid, None)) => {
recycle_if_idle(&mut sock.rx.borrow_mut(), bid);
}
Err(code) => {
sock.fail_rx(code);
return;
}
}
if terminal {
arm_recv(shared, sock);
}
}
fn on_recv_multi(sock: &Rc<SockShared>, cqe: Cqe) -> Result<(u16, Option<Queued>), i32> {
let mut rx = sock.rx.borrow_mut();
let rx = &mut *rx;
let Some(bid) = cqueue::buffer_select(cqe.flags) else {
return Err(libc::EPROTO);
};
let len = cqe.result as usize;
let buf = &mut rx.bufs[bid as usize];
if len > buf.data.len() {
return Err(libc::EPROTO);
}
buf.kernel_done = true;
let slice = &buf.data[..len];
let Ok(out) = types::RecvMsgOut::parse(slice, &rx.hdr) else {
tracing::warn!("dropping malformed multishot recvmsg completion");
return Ok((bid, None));
};
if out.is_payload_truncated() || out.is_control_data_truncated() {
tracing::warn!("dropping truncated receive (buffer tail too small for a full coalesce)");
return Ok((bid, None));
}
let Some(from) = decode_addr(out.name_data()) else {
tracing::warn!("dropping receive with an unparseable source address");
return Ok((bid, None));
};
let payload = out.payload_data();
if payload.is_empty() {
return Ok((bid, None));
}
let meta = RecvMeta::parse(out.control_data());
let payload_start = payload.as_ptr() as usize - buf.data.as_ptr() as usize;
Ok((
bid,
Some(Queued {
bid,
start: payload_start,
len: payload.len(),
from,
stride: meta.stride.unwrap_or(payload.len()),
ecn: meta.ecn,
}),
))
}
fn on_recv_oneshot(one: OneshotRecv, cqe: Cqe) -> Result<(u16, Option<Queued>), i32> {
let bid = one.bid;
let len = cqe.result as usize;
if one.hdr.msg_flags & (libc::MSG_TRUNC | libc::MSG_CTRUNC) != 0 {
tracing::warn!("dropping truncated oneshot receive");
return Ok((bid, None));
}
let name = {
let ptr = (&raw const one.name).cast::<u8>();
unsafe { std::slice::from_raw_parts(ptr, (one.hdr.msg_namelen as usize).min(NAME_LEN)) }
};
let Some(from) = decode_addr(name) else {
tracing::warn!("dropping receive with an unparseable source address");
return Ok((bid, None));
};
if len == 0 {
return Ok((bid, None));
}
let control = &one.control.0[..one.hdr.msg_controllen.min(CONTROL_LEN)];
let meta = RecvMeta::parse(control);
Ok((
bid,
Some(Queued {
bid,
start: 0,
len,
from,
stride: meta.stride.unwrap_or(len),
ecn: meta.ecn,
}),
))
}
pub(crate) fn on_send(op: SendOp, cqe: Cqe) {
if cqe.result < 0 {
let code = -cqe.result;
if code == libc::ECONNREFUSED {
tracing::debug!("send completed with ECONNREFUSED");
return;
}
if code != libc::ECANCELED {
op.sock.fail_tx(code);
}
} else if cqe.result as usize != op.expect {
tracing::warn!(sent = cqe.result, expected = op.expect, "short UDP send");
op.sock.fail_tx(libc::EIO);
}
}
#[derive(Default)]
struct RecvMeta {
stride: Option<usize>,
ecn: Option<Ecn>,
}
impl RecvMeta {
fn parse(control: &[u8]) -> Self {
let mut meta = Self::default();
let header_len = unsafe { libc::CMSG_LEN(0) as usize };
let mut offset = 0;
while offset + header_len <= control.len() {
let header = unsafe { control.as_ptr().add(offset).cast::<libc::cmsghdr>().read_unaligned() };
let message_len = header.cmsg_len;
if message_len < header_len || offset + message_len > control.len() {
return meta;
}
let data = &control[offset + header_len..offset + message_len];
match (header.cmsg_level, header.cmsg_type) {
(libc::SOL_UDP, libc::UDP_GRO) => {
meta.stride = read_int(data).and_then(|value| usize::try_from(value).ok());
}
(libc::IPPROTO_IP, libc::IP_TOS) => {
meta.ecn = data.first().and_then(|bits| Ecn::from_bits(*bits));
}
(libc::IPPROTO_IPV6, libc::IPV6_TCLASS) => {
meta.ecn = read_int(data).and_then(|value| Ecn::from_bits(value as u8));
}
_ => {}
}
let aligned = unsafe { libc::CMSG_SPACE((message_len - header_len) as _) as usize };
offset = offset.saturating_add(aligned.max(header_len));
}
meta
}
}
fn read_int(data: &[u8]) -> Option<libc::c_int> {
let bytes = data.get(..std::mem::size_of::<libc::c_int>())?;
Some(libc::c_int::from_ne_bytes(bytes.try_into().ok()?))
}
fn set_option(io: &UdpSocket, level: libc::c_int, name: libc::c_int) -> io::Result<()> {
let on: libc::c_int = 1;
let ret = unsafe {
libc::setsockopt(
io.as_raw_fd(),
level,
name,
(&raw const on).cast(),
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
)
};
match ret {
0 => Ok(()),
_ => Err(io::Error::last_os_error()),
}
}
fn encode_addr(addr: SocketAddr, out: &mut libc::sockaddr_storage) -> libc::socklen_t {
match addr {
SocketAddr::V4(v4) => {
let sin = libc::sockaddr_in {
sin_family: libc::AF_INET as libc::sa_family_t,
sin_port: v4.port().to_be(),
sin_addr: libc::in_addr {
s_addr: u32::from_ne_bytes(v4.ip().octets()),
},
sin_zero: [0; 8],
};
unsafe { (&raw mut *out).cast::<libc::sockaddr_in>().write(sin) };
std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t
}
SocketAddr::V6(v6) => {
let sin6 = libc::sockaddr_in6 {
sin6_family: libc::AF_INET6 as libc::sa_family_t,
sin6_port: v6.port().to_be(),
sin6_flowinfo: v6.flowinfo(),
sin6_addr: libc::in6_addr {
s6_addr: v6.ip().octets(),
},
sin6_scope_id: v6.scope_id(),
};
unsafe { (&raw mut *out).cast::<libc::sockaddr_in6>().write(sin6) };
std::mem::size_of::<libc::sockaddr_in6>() as libc::socklen_t
}
}
}
fn decode_addr(name: &[u8]) -> Option<SocketAddr> {
if name.len() < std::mem::size_of::<libc::sa_family_t>() {
return None;
}
const FAMILY_LEN: usize = std::mem::size_of::<libc::sa_family_t>();
let mut family = [0u8; FAMILY_LEN];
family.copy_from_slice(&name[..FAMILY_LEN]);
match libc::sa_family_t::from_ne_bytes(family) as libc::c_int {
libc::AF_INET if name.len() >= std::mem::size_of::<libc::sockaddr_in>() => {
let sin = unsafe { name.as_ptr().cast::<libc::sockaddr_in>().read_unaligned() };
Some(SocketAddr::from((
sin.sin_addr.s_addr.to_ne_bytes(),
u16::from_be(sin.sin_port),
)))
}
libc::AF_INET6 if name.len() >= std::mem::size_of::<libc::sockaddr_in6>() => {
let sin6 = unsafe { name.as_ptr().cast::<libc::sockaddr_in6>().read_unaligned() };
Some(SocketAddr::V6(SocketAddrV6::new(
sin6.sin6_addr.s6_addr.into(),
u16::from_be(sin6.sin6_port),
sin6.sin6_flowinfo,
sin6.sin6_scope_id,
)))
}
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddrV4};
fn roundtrip(addr: SocketAddr) -> Option<SocketAddr> {
let mut storage: libc::sockaddr_storage = unsafe { std::mem::zeroed() };
let len = encode_addr(addr, &mut storage) as usize;
let name = unsafe { std::slice::from_raw_parts((&raw const storage).cast::<u8>(), len) };
decode_addr(name)
}
#[test]
fn addr_roundtrip_v4() {
let addr = SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(192, 0, 2, 7), 4443));
assert_eq!(roundtrip(addr), Some(addr));
}
#[test]
fn addr_roundtrip_v6_keeps_scope_and_flow() {
let ip = Ipv6Addr::new(0xfe80, 0, 0, 0, 0, 0, 0, 1);
let addr = SocketAddr::V6(SocketAddrV6::new(ip, 4443, 0x12345, 3));
assert_eq!(roundtrip(addr), Some(addr));
}
fn empty_rx() -> Rx {
Rx {
bufs: Vec::new(),
ring: Some(BufRing::new(64)),
hdr: Box::new(unsafe { std::mem::zeroed() }),
queue: VecDeque::new(),
waiters: kio::WaiterList::new(),
armed: None,
starved: false,
error: None,
}
}
#[test]
fn a_recycled_buffer_does_not_mask_a_recorded_starvation() {
let config = Config::default();
let mut rx = empty_rx();
grow_rx(&mut rx, &config);
assert!(!should_grow(&rx, true), "a buffer is in the ring");
rx.starved = true;
assert!(should_grow(&rx, true), "the kernel ran dry, recycle or not");
assert!(should_grow(&rx, false), "and the oneshot path reads it too");
}
#[test]
fn the_receive_pool_doubles_to_its_ceiling() {
let config = Config {
rx_buffers_max: 40,
..Default::default()
};
let mut rx = empty_rx();
for expected in [1u16, 2, 4, 8, 16, 32, 40] {
assert!(grow_rx(&mut rx, &config), "growth stopped short of {expected}");
assert_eq!(rx.bufs.len(), usize::from(expected));
assert_eq!(rx.ring.as_ref().expect("ring").tail, expected);
}
assert!(!grow_rx(&mut rx, &config), "grew past the ceiling");
}
}