use std::error::Error as StdError;
use std::fmt::Display;
use std::io::IoSliceMut;
use std::net::{IpAddr, SocketAddr};
use std::task::{Context, Poll};
pub trait UdpBind {
type Socket: UdpDatagrams;
fn bind(&self, local: SocketAddr) -> std::io::Result<Self::Socket>;
}
pub trait UdpAdoptStd: UdpBind {
fn adopt(&self, s: std::net::UdpSocket) -> std::io::Result<Self::Socket>;
}
pub trait UdpDatagrams {
fn try_send(&self, t: &Datagrams<'_>) -> std::io::Result<()>;
fn poll_writable(&self, cx: &mut Context<'_>) -> Poll<std::io::Result<()>>;
fn poll_recv(
&self,
cx: &mut Context<'_>,
bufs: &mut [IoSliceMut<'_>],
meta: &mut [RecvMeta],
) -> Poll<std::io::Result<usize>>;
fn local_addr(&self) -> std::io::Result<SocketAddr>;
fn caps(&self) -> UdpCaps {
UdpCaps::NONE
}
}
#[derive(Debug, Clone, Copy)]
pub struct Datagrams<'a> {
pub destination: SocketAddr,
pub src_ip: Option<IpAddr>,
pub ecn: Option<EcnCodepoint>,
pub segment_size: Option<usize>,
pub contents: &'a [u8],
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecvMeta {
pub addr: SocketAddr,
pub len: usize,
pub stride: usize,
pub ecn: Option<EcnCodepoint>,
pub dst_ip: Option<IpAddr>,
}
impl Default for RecvMeta {
fn default() -> Self {
Self {
addr: SocketAddr::from(([0, 0, 0, 0], 0)),
len: 0,
stride: 0,
ecn: None,
dst_ip: None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum EcnCodepoint {
Ect1 = 0b01,
Ect0 = 0b10,
Ce = 0b11,
}
impl EcnCodepoint {
pub fn from_bits(bits: u8) -> Option<Self> {
match bits & 0b11 {
0b01 => Some(Self::Ect1),
0b10 => Some(Self::Ect0),
0b11 => Some(Self::Ce),
_ => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UdpCaps {
pub max_send_segments: usize,
pub max_recv_segments: usize,
pub ecn: bool,
pub may_fragment: bool,
}
impl UdpCaps {
pub const NONE: Self = Self {
max_send_segments: 1,
max_recv_segments: 1,
ecn: false,
may_fragment: true,
};
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UnsupportedUdpOffload {
gso: bool,
ecn: bool,
}
impl UnsupportedUdpOffload {
pub fn names(&self) -> impl Iterator<Item = &'static str> {
[("gso", self.gso), ("ecn", self.ecn)]
.into_iter()
.filter_map(|(name, bad)| bad.then_some(name))
}
}
impl Display for UnsupportedUdpOffload {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(
"this socket does not have these UDP offloads, and does not silently drop them:",
)?;
for (i, name) in self.names().enumerate() {
f.write_str(if i > 0 { ", " } else { " " })?;
f.write_str(name)?;
}
Ok(())
}
}
impl StdError for UnsupportedUdpOffload {}
impl Datagrams<'_> {
pub fn segments(&self) -> usize {
match self.segment_size {
None => 1,
Some(0) => 1,
Some(n) => self.contents.len().div_ceil(n),
}
}
pub fn reject_unsupported(&self, caps: UdpCaps) -> std::io::Result<()> {
let gso = self.segments() > caps.max_send_segments;
let ecn = false;
if !gso && !ecn {
return Ok(());
}
Err(std::io::Error::new(
std::io::ErrorKind::Unsupported,
UnsupportedUdpOffload { gso, ecn },
))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn to(port: u16) -> SocketAddr {
SocketAddr::from(([127, 0, 0, 1], port))
}
fn plain(contents: &[u8]) -> Datagrams<'_> {
Datagrams {
destination: to(1),
src_ip: None,
ecn: None,
segment_size: None,
contents,
}
}
#[test]
fn none_is_the_conservative_base() {
let c = UdpCaps::NONE;
assert_eq!(c.max_send_segments, 1, "1 == no GSO, not 0 and not 64");
assert_eq!(c.max_recv_segments, 1, "1 == no GRO");
assert!(!c.ecn);
assert!(
c.may_fragment,
"the pessimistic answer: a socket that says nothing must not \
claim path MTU discovery is reliable"
);
}
#[test]
fn a_default_caps_impl_reports_nothing() {
struct Forgetful;
impl UdpDatagrams for Forgetful {
fn try_send(&self, _: &Datagrams<'_>) -> std::io::Result<()> {
unreachable!("this socket never sends")
}
fn poll_writable(&self, _: &mut Context<'_>) -> Poll<std::io::Result<()>> {
unreachable!("this socket never sends")
}
fn poll_recv(
&self,
_: &mut Context<'_>,
_: &mut [IoSliceMut<'_>],
_: &mut [RecvMeta],
) -> Poll<std::io::Result<usize>> {
unreachable!("this socket never receives")
}
fn local_addr(&self) -> std::io::Result<SocketAddr> {
unreachable!("this socket is never bound")
}
}
assert_eq!(Forgetful.caps(), UdpCaps::NONE);
}
#[test]
fn segments_counts_datagrams_not_bytes() {
assert_eq!(plain(&[0u8; 3600]).segments(), 1, "no GSO asked for");
let g = Datagrams {
segment_size: Some(1200),
..plain(&[0u8; 3600])
};
assert_eq!(g.segments(), 3);
let g = Datagrams {
segment_size: Some(1200),
..plain(&[0u8; 2401])
};
assert_eq!(g.segments(), 3);
}
#[test]
fn asking_for_nothing_is_never_an_offence() {
assert!(plain(b"hello").reject_unsupported(UdpCaps::NONE).is_ok());
let marked = Datagrams {
ecn: Some(EcnCodepoint::Ect0),
..plain(b"hello")
};
assert!(
marked.reject_unsupported(UdpCaps::NONE).is_ok(),
"a socket that declares no ECN is allowed to drop the marking — \
refusing here would make QUIC unusable on such a kernel"
);
}
#[test]
fn gso_beyond_the_declared_batch_is_refused_by_name() {
let g = Datagrams {
segment_size: Some(1200),
..plain(&[0u8; 3600])
};
assert!(
g.reject_unsupported(UdpCaps {
max_send_segments: 3,
..UdpCaps::NONE
})
.is_ok()
);
let err = g
.reject_unsupported(UdpCaps {
max_send_segments: 2,
..UdpCaps::NONE
})
.expect_err("three datagrams asked of a two-datagram socket");
assert_eq!(err.kind(), std::io::ErrorKind::Unsupported);
let payload = err
.get_ref()
.and_then(|e| e.downcast_ref::<UnsupportedUdpOffload>())
.expect("the typed payload survives the trip through io::Error");
assert_eq!(payload.names().collect::<Vec<_>>(), ["gso"]);
}
#[test]
fn an_unfilled_recv_slot_reports_no_ecn_rather_than_a_plausible_one() {
assert_eq!(RecvMeta::default().ecn, None);
assert_eq!(RecvMeta::default().stride, 0);
assert_eq!(RecvMeta::default().len, 0);
}
#[test]
fn ecn_bits_round_trip_and_zero_is_not_a_codepoint() {
for c in [EcnCodepoint::Ect1, EcnCodepoint::Ect0, EcnCodepoint::Ce] {
assert_eq!(EcnCodepoint::from_bits(c as u8), Some(c));
}
assert_eq!(
EcnCodepoint::from_bits(0b00),
None,
"`00` is not-ECN-capable, which is an absence and not a fourth variant"
);
assert_eq!(
EcnCodepoint::from_bits(0b1011_1110),
Some(EcnCodepoint::Ect0)
);
}
}