use libc::{cmsghdr, msghdr};
#[cfg(any(has_ip_pktinfo, has_ipv6_pktinfo))]
use core::net::IpAddr;
#[cfg(has_ip_pktinfo)]
use core::net::Ipv4Addr;
#[cfg(has_ipv6_pktinfo)]
use core::net::Ipv6Addr;
use super::RecvMeta;
pub(crate) struct CMsgRef<'a> {
inner: &'a cmsghdr,
}
impl CMsgRef<'_> {
#[inline]
pub(crate) fn level(&self) -> libc::c_int {
self.inner.cmsg_level
}
#[inline]
pub(crate) fn ty(&self) -> libc::c_int {
self.inner.cmsg_type
}
#[inline]
#[allow(clippy::unnecessary_cast)]
pub(crate) fn data_len(&self) -> usize {
let base = unsafe { libc::CMSG_LEN(0) } as usize;
(self.inner.cmsg_len as usize).saturating_sub(base)
}
#[inline]
pub(crate) unsafe fn data<T>(&self) -> *const T {
unsafe { libc::CMSG_DATA(self.inner as *const cmsghdr) as *const T }
}
}
pub(crate) struct CMsgIter<'a> {
msg: msghdr,
next: *const cmsghdr,
_lt: core::marker::PhantomData<&'a [u8]>,
}
impl<'a> CMsgIter<'a> {
pub(crate) fn new(buf: &'a [u8]) -> Self {
assert!(
buf.as_ptr().cast::<cmsghdr>().is_aligned(),
"control buffer is not aligned for cmsghdr"
);
let mut msg: msghdr = unsafe { core::mem::zeroed() };
msg.msg_control = buf.as_ptr() as *mut _;
msg.msg_controllen = buf.len() as _;
let next = unsafe { libc::CMSG_FIRSTHDR(&msg) } as *const cmsghdr;
Self {
msg,
next,
_lt: core::marker::PhantomData,
}
}
}
impl<'a> Iterator for CMsgIter<'a> {
type Item = CMsgRef<'a>;
fn next(&mut self) -> Option<Self::Item> {
if self.next.is_null() {
return None;
}
let cur = unsafe { &*self.next };
self.next = unsafe { libc::CMSG_NXTHDR(&self.msg, self.next) } as *const cmsghdr;
Some(CMsgRef { inner: cur })
}
}
pub(crate) struct CMsgBuilder<'a> {
buf: &'a mut [u8],
cursor: usize,
}
impl<'a> CMsgBuilder<'a> {
pub(crate) fn new(buf: &'a mut [u8]) -> Self {
assert!(
buf.as_ptr().cast::<cmsghdr>().is_aligned(),
"control buffer is not aligned for cmsghdr"
);
for b in buf.iter_mut() {
*b = 0;
}
Self { buf, cursor: 0 }
}
pub(crate) fn push<T: Copy>(
&mut self,
level: libc::c_int,
ty: libc::c_int,
value: &T,
) -> Result<(), ()> {
let payload_bytes = core::mem::size_of::<T>();
let space = unsafe { libc::CMSG_SPACE(payload_bytes as u32) } as usize;
let end = self.cursor.checked_add(space).ok_or(())?;
if end > self.buf.len() {
return Err(());
}
unsafe {
let hdr = self.buf.as_mut_ptr().add(self.cursor) as *mut cmsghdr;
(*hdr).cmsg_len = libc::CMSG_LEN(payload_bytes as u32) as _;
(*hdr).cmsg_level = level;
(*hdr).cmsg_type = ty;
let data = libc::CMSG_DATA(hdr) as *mut T;
core::ptr::write_unaligned(data, *value);
}
self.cursor = end;
Ok(())
}
#[inline]
pub(crate) fn finish(self) -> usize {
self.cursor
}
}
pub(super) struct AlignedCtrlBuf {
storage: Box<AlignedCtrlStorage>,
init_len: usize,
}
const CMSG_CAP: usize = 256;
#[repr(align(8))]
struct AlignedCtrlStorage([u8; CMSG_CAP]);
impl AlignedCtrlBuf {
pub fn new() -> Self {
Self {
storage: Box::new(AlignedCtrlStorage([0u8; CMSG_CAP])),
init_len: 0,
}
}
pub fn from_slice(src: &[u8]) -> Self {
assert!(
src.len() <= CMSG_CAP,
"outbound cmsg payload {} exceeds CMSG_CAP={CMSG_CAP}",
src.len()
);
let mut buf = Self::new();
buf.storage.0[..src.len()].copy_from_slice(src);
buf.init_len = src.len();
buf
}
pub fn filled(&self, kernel_len: usize) -> &[u8] {
let n = kernel_len.min(CMSG_CAP);
&self.storage.0[..n]
}
}
impl compio_buf::IoBuf for AlignedCtrlBuf {
fn as_init(&self) -> &[u8] {
&self.storage.0[..self.init_len]
}
}
impl compio_buf::IoBufMut for AlignedCtrlBuf {
fn as_uninit(&mut self) -> &mut [core::mem::MaybeUninit<u8>] {
let ptr = self.storage.0.as_mut_ptr() as *mut core::mem::MaybeUninit<u8>;
unsafe { core::slice::from_raw_parts_mut(ptr, CMSG_CAP) }
}
}
impl compio_buf::SetLen for AlignedCtrlBuf {
unsafe fn set_len(&mut self, len: usize) {
debug_assert!(len <= CMSG_CAP);
self.init_len = len.min(CMSG_CAP);
}
}
pub(super) fn enable_recv_cmsgs(sock: &std::net::UdpSocket) -> std::io::Result<()> {
use std::os::fd::AsRawFd;
let fd = sock.as_raw_fd();
let on: libc::c_int = 1;
let is_v6 = matches!(sock.local_addr()?, std::net::SocketAddr::V6(_));
let _ = (fd, on, is_v6);
if is_v6 {
#[cfg(has_ipv6_pktinfo)]
set_int(fd, libc::IPPROTO_IPV6, libc::IPV6_RECVPKTINFO, on)?;
#[cfg(has_recv_hoplimit)]
set_int(fd, libc::IPPROTO_IPV6, libc::IPV6_RECVHOPLIMIT, on)?;
} else {
#[cfg(has_ip_pktinfo)]
set_int(fd, libc::IPPROTO_IP, libc::IP_PKTINFO, on)?;
#[cfg(has_recv_hoplimit)]
set_int(fd, libc::IPPROTO_IP, libc::IP_RECVTTL, on)?;
}
#[cfg(all(has_recv_timestamp, recv_timestamp_ns))]
set_int(fd, libc::SOL_SOCKET, libc::SO_TIMESTAMPNS, on).ok();
#[cfg(all(has_recv_timestamp, not(recv_timestamp_ns)))]
set_int(fd, libc::SOL_SOCKET, libc::SO_TIMESTAMP, on).ok();
Ok(())
}
fn set_int(
fd: std::os::fd::RawFd,
level: libc::c_int,
optname: libc::c_int,
val: libc::c_int,
) -> std::io::Result<()> {
let rc = unsafe {
libc::setsockopt(
fd,
level,
optname,
&val as *const _ as *const _,
core::mem::size_of::<libc::c_int>() as libc::socklen_t,
)
};
if rc != 0 {
Err(std::io::Error::last_os_error())
} else {
Ok(())
}
}
pub(super) fn decode_unix_cmsgs(ctrl: &[u8], meta: &mut RecvMeta) {
if ctrl.is_empty() {
return;
}
if !ctrl.as_ptr().cast::<libc::cmsghdr>().is_aligned() {
return;
}
for c in CMsgIter::new(ctrl) {
match (c.level(), c.ty()) {
#[cfg(has_ip_pktinfo)]
(libc::IPPROTO_IP, libc::IP_PKTINFO) => {
if c.data_len() < core::mem::size_of::<libc::in_pktinfo>() {
continue;
}
let pi = unsafe { core::ptr::read_unaligned(c.data::<libc::in_pktinfo>()) };
meta.local_ip = IpAddr::V4(Ipv4Addr::from(u32::from_be(pi.ipi_spec_dst.s_addr)));
meta.interface_index = pi.ipi_ifindex as u32;
}
#[cfg(has_ipv6_pktinfo)]
(libc::IPPROTO_IPV6, libc::IPV6_PKTINFO) => {
if c.data_len() < core::mem::size_of::<libc::in6_pktinfo>() {
continue;
}
let pi = unsafe { core::ptr::read_unaligned(c.data::<libc::in6_pktinfo>()) };
meta.local_ip = IpAddr::V6(Ipv6Addr::from(pi.ipi6_addr.s6_addr));
meta.interface_index = pi.ipi6_ifindex as u32;
}
#[cfg(has_recv_hoplimit)]
(libc::IPPROTO_IP, libc::IP_TTL) | (libc::IPPROTO_IP, libc::IP_RECVTTL) => {
if c.data_len() < core::mem::size_of::<libc::c_int>() {
continue;
}
let v = unsafe { core::ptr::read_unaligned(c.data::<libc::c_int>()) };
meta.hop_limit = Some(v as u8);
}
#[cfg(has_recv_hoplimit)]
(libc::IPPROTO_IPV6, libc::IPV6_HOPLIMIT) => {
if c.data_len() < core::mem::size_of::<libc::c_int>() {
continue;
}
let v = unsafe { core::ptr::read_unaligned(c.data::<libc::c_int>()) };
meta.hop_limit = Some(v as u8);
}
#[cfg(all(has_recv_timestamp, recv_timestamp_ns))]
(libc::SOL_SOCKET, libc::SCM_TIMESTAMPNS) => {
if c.data_len() < core::mem::size_of::<libc::timespec>() {
continue;
}
let ts = unsafe { core::ptr::read_unaligned(c.data::<libc::timespec>()) };
let nanos = u32::try_from(ts.tv_nsec).unwrap_or(0).min(999_999_999);
if let Ok(secs) = u64::try_from(ts.tv_sec) {
meta.kernel_rx_time =
std::time::UNIX_EPOCH.checked_add(std::time::Duration::new(secs, nanos));
}
}
#[cfg(all(has_recv_timestamp, not(recv_timestamp_ns)))]
(libc::SOL_SOCKET, libc::SCM_TIMESTAMP) => {
if c.data_len() < core::mem::size_of::<libc::timeval>() {
continue;
}
let tv = unsafe { core::ptr::read_unaligned(c.data::<libc::timeval>()) };
let micros = u32::try_from(tv.tv_usec).unwrap_or(0).min(999_999);
if let Ok(secs) = u64::try_from(tv.tv_sec) {
meta.kernel_rx_time =
std::time::UNIX_EPOCH.checked_add(std::time::Duration::new(secs, micros * 1000));
}
}
_ => {}
}
}
}
#[cfg(all(unix, test))]
#[compio::test]
async fn from_std_enables_cmsgs_on_v6_socket() {
use crate::socket::Socket;
use std::net::{Ipv6Addr, UdpSocket};
let sock = match UdpSocket::bind((Ipv6Addr::LOCALHOST, 0)) {
Ok(s) => s,
Err(_) => return, };
let wrapped = Socket::from_std(sock).await;
assert!(
wrapped.is_ok(),
"from_std must enable cmsgs on a v6 socket without EINVAL, got {:?}",
wrapped.err()
);
}
#[cfg(all(unix, test))]
#[compio::test]
async fn from_std_enables_cmsgs_on_v4_socket() {
use crate::socket::Socket;
use std::net::{Ipv4Addr, UdpSocket};
let sock = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).expect("bind v4");
let wrapped = Socket::from_std(sock).await;
assert!(
wrapped.is_ok(),
"from_std must still succeed on a v4 socket, got {:?}",
wrapped.err()
);
}
#[cfg(all(unix, has_recv_timestamp, test))]
#[test]
fn cmsg_iter_walks_a_single_timestamp_cmsg() {
#[cfg(not(recv_timestamp_ns))]
use libc::{SCM_TIMESTAMP as TS_TYPE, timeval as TsPayload};
#[cfg(recv_timestamp_ns)]
use libc::{SCM_TIMESTAMPNS as TS_TYPE, timespec as TsPayload};
use libc::{SOL_SOCKET, cmsghdr};
#[cfg(recv_timestamp_ns)]
let payload = TsPayload {
tv_sec: 1234,
tv_nsec: 56,
};
#[cfg(not(recv_timestamp_ns))]
let payload = TsPayload {
tv_sec: 1234,
tv_usec: 56,
};
let payload_bytes = core::mem::size_of::<TsPayload>();
let total = unsafe { libc::CMSG_SPACE(payload_bytes as u32) } as usize;
assert!(core::mem::align_of::<cmsghdr>() <= core::mem::align_of::<u64>());
let words = total.div_ceil(core::mem::size_of::<u64>());
let mut backing: Vec<u64> = vec![0u64; words.max(1)];
let buf: &mut [u8] =
unsafe { core::slice::from_raw_parts_mut(backing.as_mut_ptr().cast::<u8>(), total) };
unsafe {
let hdr = buf.as_mut_ptr() as *mut cmsghdr;
(*hdr).cmsg_len = libc::CMSG_LEN(payload_bytes as u32) as _;
(*hdr).cmsg_level = SOL_SOCKET;
(*hdr).cmsg_type = TS_TYPE;
let data = libc::CMSG_DATA(hdr) as *mut TsPayload;
core::ptr::write_unaligned(data, payload);
}
let mut iter = CMsgIter::new(buf);
let first = iter.next().expect("one cmsg");
assert_eq!(first.level(), SOL_SOCKET);
assert_eq!(first.ty(), TS_TYPE);
let got = unsafe { core::ptr::read_unaligned(first.data::<TsPayload>()) };
assert_eq!(got.tv_sec, 1234);
#[cfg(recv_timestamp_ns)]
assert_eq!(got.tv_nsec, 56);
#[cfg(not(recv_timestamp_ns))]
assert_eq!(got.tv_usec, 56);
assert!(iter.next().is_none(), "no second cmsg");
}
#[cfg(all(unix, test))]
#[test]
fn cmsg_builder_emits_a_round_trippable_pktinfo() {
use libc::{IP_PKTINFO, IPPROTO_IP, in_addr, in_pktinfo};
let pktinfo = in_pktinfo {
ipi_ifindex: 7,
ipi_spec_dst: in_addr {
s_addr: u32::from_be_bytes([127, 0, 0, 1]).to_be(),
},
ipi_addr: in_addr {
s_addr: u32::from_be_bytes([127, 0, 0, 1]).to_be(),
},
};
assert!(core::mem::align_of::<cmsghdr>() <= core::mem::align_of::<u64>());
let mut backing: Vec<u64> = vec![0u64; 128 / core::mem::size_of::<u64>()];
let buf: &mut [u8] =
unsafe { core::slice::from_raw_parts_mut(backing.as_mut_ptr().cast::<u8>(), 128) };
let written = {
let mut b = CMsgBuilder::new(buf);
b.push(IPPROTO_IP, IP_PKTINFO, &pktinfo).expect("fits");
b.finish()
};
assert!(written > 0);
let mut iter = CMsgIter::new(&buf[..written]);
let cmsg = iter.next().expect("round-tripped one cmsg");
assert_eq!(cmsg.level(), IPPROTO_IP);
assert_eq!(cmsg.ty(), IP_PKTINFO);
let got = unsafe { core::ptr::read_unaligned(cmsg.data::<in_pktinfo>()) };
assert_eq!(got.ipi_ifindex, 7);
assert!(iter.next().is_none(), "no second cmsg");
}
#[cfg(all(unix, test))]
#[test]
fn truncated_pktinfo_cmsg_is_skipped_not_read() {
use libc::{IP_PKTINFO, IPPROTO_IP, cmsghdr};
let payload_bytes = core::mem::size_of::<libc::in_pktinfo>();
let total = unsafe { libc::CMSG_SPACE(payload_bytes as u32) } as usize;
assert!(core::mem::align_of::<cmsghdr>() <= core::mem::align_of::<u64>());
let words = total.div_ceil(core::mem::size_of::<u64>());
let mut backing: Vec<u64> = vec![0u64; words.max(1)];
let buf: &mut [u8] =
unsafe { core::slice::from_raw_parts_mut(backing.as_mut_ptr().cast::<u8>(), total) };
unsafe {
let hdr = buf.as_mut_ptr() as *mut cmsghdr;
(*hdr).cmsg_len = libc::CMSG_LEN(2) as _;
(*hdr).cmsg_level = IPPROTO_IP;
(*hdr).cmsg_type = IP_PKTINFO;
}
let mut meta = RecvMeta::empty(([0u8, 0, 0, 0], 0).into());
decode_unix_cmsgs(buf, &mut meta);
assert!(
meta.local_ip.is_unspecified(),
"truncated PKTINFO populated local_ip from a short cmsg"
);
assert_eq!(
meta.interface_index, 0,
"truncated PKTINFO populated interface_index from a short cmsg"
);
}
#[cfg(all(unix, has_recv_timestamp, test))]
#[test]
fn absurd_timestamp_does_not_panic() {
use libc::{SOL_SOCKET, cmsghdr};
#[cfg(not(recv_timestamp_ns))]
use libc::{SCM_TIMESTAMP as TS_TYPE, timeval as TsPayload};
#[cfg(recv_timestamp_ns)]
use libc::{SCM_TIMESTAMPNS as TS_TYPE, timespec as TsPayload};
#[cfg(recv_timestamp_ns)]
let payload = TsPayload {
tv_sec: i64::MAX,
tv_nsec: i64::MAX,
};
#[cfg(not(recv_timestamp_ns))]
let payload = TsPayload {
tv_sec: i64::MAX as _,
tv_usec: i64::MAX as _,
};
let payload_bytes = core::mem::size_of::<TsPayload>();
let total = unsafe { libc::CMSG_SPACE(payload_bytes as u32) } as usize;
assert!(core::mem::align_of::<cmsghdr>() <= core::mem::align_of::<u64>());
let words = total.div_ceil(core::mem::size_of::<u64>());
let mut backing: Vec<u64> = vec![0u64; words.max(1)];
let buf: &mut [u8] =
unsafe { core::slice::from_raw_parts_mut(backing.as_mut_ptr().cast::<u8>(), total) };
unsafe {
let hdr = buf.as_mut_ptr() as *mut cmsghdr;
(*hdr).cmsg_len = libc::CMSG_LEN(payload_bytes as u32) as _;
(*hdr).cmsg_level = SOL_SOCKET;
(*hdr).cmsg_type = TS_TYPE;
let data = libc::CMSG_DATA(hdr) as *mut TsPayload;
core::ptr::write_unaligned(data, payload);
}
let mut meta = RecvMeta::empty(([0u8, 0, 0, 0], 0).into());
decode_unix_cmsgs(buf, &mut meta);
if let Some(t) = meta.kernel_rx_time {
let _ = t.duration_since(std::time::UNIX_EPOCH);
}
}