use std::io;
use std::net::{SocketAddr, UdpSocket};
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Shard {
index: u16,
count: u16,
}
impl Shard {
fn new(index: u16, count: u16) -> Option<Self> {
(count <= MAX_SHARDS && index < count).then_some(Self { index, count })
}
pub fn index(self) -> u16 {
self.index
}
pub fn count(self) -> u16 {
self.count
}
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum Error {
#[error("a reuseport group holds at most {max} members; {count} were asked for")]
Count {
count: u16,
max: u16,
},
#[error("another reuseport group already holds port {port}")]
Overlap {
port: u16,
},
#[error("failed to resolve an ephemeral port")]
Resolve(#[source] io::Error),
}
#[derive(Debug)]
pub struct Group {
count: u16,
next: u16,
state: Arc<State>,
}
impl Group {
pub fn acquire(addr: SocketAddr, count: u16) -> Result<Self, Error> {
let count = count.max(1);
if count > MAX_SHARDS {
return Err(Error::Count { count, max: MAX_SHARDS });
}
let addr = match addr.port() {
0 => crate::bind::udp(crate::bind::Udp::new(addr))
.and_then(|socket| socket.local_addr())
.map_err(Error::Resolve)?,
_ => addr,
};
let port = addr.port();
let lock = Lock::acquire(port).map_err(|_| Error::Overlap { port })?;
Ok(Self {
count,
next: 0,
state: Arc::new(State::new(addr, lock)),
})
}
pub fn count(&self) -> u16 {
self.count
}
pub fn addr(&self) -> SocketAddr {
self.state.addr
}
pub fn member(&mut self) -> Option<Member> {
let shard = Shard::new(self.next, self.count)?;
self.next += 1;
Some(Member {
shard,
state: self.state.clone(),
})
}
pub fn complete(self, claims: Vec<Claim>) -> io::Result<Bound> {
if claims.len() != usize::from(self.count) {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"reuseport group needs {} bound members; {} were provided",
self.count,
claims.len()
),
));
}
for (index, claim) in claims.iter().enumerate() {
if !Arc::ptr_eq(&self.state, &claim.state) || usize::from(claim.shard.index()) != index {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"reuseport claims must belong to this group in slot order",
));
}
}
let sockets: Vec<_> = claims.into_iter().map(|claim| (claim.shard, claim.socket)).collect();
let (_, last) = sockets.last().expect("a group always has at least one member");
attach(last, self.count)?;
Ok(Bound {
sockets,
next: 0,
state: self.state,
})
}
}
#[derive(Debug)]
pub struct Member {
shard: Shard,
state: Arc<State>,
}
impl Member {
pub fn shard(&self) -> Shard {
self.shard
}
pub fn bind(self) -> io::Result<Claim> {
let mut bound = self.state.bound();
if *bound != self.shard.index() {
return Err(io::Error::other(format!(
"reuseport member {} cannot bind while {} of {} are in: the kernel numbers a group by bind order",
self.shard.index(),
*bound,
self.shard.count(),
)));
}
let socket = bind(self.state.addr, self.shard)?;
*bound += 1;
drop(bound);
Ok(Claim {
shard: self.shard,
socket,
state: self.state,
})
}
}
#[derive(Debug)]
pub struct Claim {
shard: Shard,
socket: UdpSocket,
state: Arc<State>,
}
#[derive(Debug)]
pub struct Bound {
sockets: Vec<(Shard, UdpSocket)>,
next: usize,
state: Arc<State>,
}
impl Bound {
pub fn count(&self) -> u16 {
self.sockets.len().try_into().expect("group count fits u16")
}
pub fn addr(&self) -> SocketAddr {
self.state.addr
}
pub fn member(&mut self) -> io::Result<Option<Socket>> {
let Some((shard, socket)) = self.sockets.get(self.next) else {
return Ok(None);
};
let socket = socket.try_clone()?;
self.next += 1;
Ok(Some(Socket { shard: *shard, socket }))
}
}
#[derive(Debug)]
pub struct Socket {
shard: Shard,
socket: UdpSocket,
}
impl Socket {
pub fn shard(&self) -> Shard {
self.shard
}
pub fn into_inner(self) -> UdpSocket {
self.socket
}
}
#[derive(Debug)]
struct State {
addr: SocketAddr,
bound: Mutex<u16>,
_lock: Option<Lock>,
}
impl State {
fn new(addr: SocketAddr, lock: Option<Lock>) -> Self {
Self {
addr,
bound: Mutex::new(0),
_lock: lock,
}
}
fn bound(&self) -> MutexGuard<'_, u16> {
self.bound.lock().unwrap_or_else(PoisonError::into_inner)
}
}
fn bind(addr: SocketAddr, shard: Shard) -> io::Result<UdpSocket> {
if shard.index() == 0 {
drop(crate::bind::udp(crate::bind::Udp::new(addr))?);
}
crate::bind::udp(crate::bind::Udp::new(addr).with_reuse_port(true))
}
#[derive(Debug)]
struct Lock {
#[cfg(target_os = "linux")]
_file: std::fs::File,
}
impl Lock {
fn acquire(port: u16) -> io::Result<Option<Self>> {
#[cfg(target_os = "linux")]
{
use std::os::unix::fs::OpenOptionsExt;
let Some(dir) = dir() else {
return Ok(None);
};
let path = dir.join(format!("quic-workers-{port}.lock"));
let file = match std::fs::OpenOptions::new()
.write(true)
.create(true)
.truncate(false)
.mode(0o600)
.custom_flags(libc::O_NOFOLLOW | libc::O_NONBLOCK)
.open(&path)
{
Ok(file) => file,
Err(err) => {
tracing::warn!(?path, %err, "cannot open the lock file; group overlap detection falls back to the bind probe");
return Ok(None);
}
};
{
use std::os::unix::fs::{MetadataExt, PermissionsExt};
let euid = unsafe { libc::geteuid() };
let trusted = file.metadata().is_ok_and(|meta| meta.is_file() && meta.uid() == euid);
if !trusted || file.set_permissions(std::fs::Permissions::from_mode(0o600)).is_err() {
tracing::warn!(
?path,
"the lock file is not exclusively ours; group overlap detection falls back to the bind probe"
);
return Ok(None);
}
}
let res = unsafe { libc::flock(std::os::fd::AsRawFd::as_raw_fd(&file), libc::LOCK_EX | libc::LOCK_NB) };
if res == 0 {
return Ok(Some(Self { _file: file }));
}
let err = io::Error::last_os_error();
if err.kind() == io::ErrorKind::WouldBlock {
return Err(err);
}
tracing::warn!(?path, %err, "cannot lock the lock file; group overlap detection falls back to the bind probe");
Ok(None)
}
#[cfg(not(target_os = "linux"))]
{
let _ = port;
Ok(None)
}
}
}
#[cfg(target_os = "linux")]
fn dir() -> Option<std::path::PathBuf> {
if let Some(runtime) = std::env::var_os("XDG_RUNTIME_DIR").filter(|dir| !dir.is_empty())
&& let Some(dir) = prepare(std::path::PathBuf::from(runtime).join("moq"))
{
return Some(dir);
}
let euid = unsafe { libc::geteuid() };
prepare(std::env::temp_dir().join(format!("moq-{euid}")))
}
#[cfg(target_os = "linux")]
fn prepare(dir: std::path::PathBuf) -> Option<std::path::PathBuf> {
use std::os::unix::fs::{DirBuilderExt, MetadataExt, PermissionsExt};
let euid = unsafe { libc::geteuid() };
let parent = dir
.parent()
.and_then(|parent| std::fs::symlink_metadata(parent).ok())
.is_some_and(|meta| {
let mode = meta.permissions().mode();
let owner = meta.uid() == euid || meta.uid() == 0;
meta.is_dir() && owner && (mode & 0o1000 != 0 || mode & 0o022 == 0)
});
if !parent {
tracing::warn!(
?dir,
"the lock directory's parent cannot protect it; group overlap detection falls back to the bind probe"
);
return None;
}
let created = std::fs::DirBuilder::new().mode(0o700).create(&dir);
if let Err(err) = &created
&& err.kind() != io::ErrorKind::AlreadyExists
{
tracing::warn!(?dir, %err, "cannot create a lock directory; group overlap detection falls back to the bind probe");
return None;
}
let safe = std::fs::symlink_metadata(&dir)
.map(|meta| meta.is_dir() && meta.uid() == euid && meta.permissions().mode() & 0o077 == 0)
.unwrap_or(false);
if !safe {
tracing::warn!(
?dir,
"lock directory is not exclusively ours; group overlap detection falls back to the bind probe"
);
return None;
}
Some(dir)
}
pub const MAX_SHARDS: u16 = 256;
pub fn cid_prefix(shard: Shard) -> u8 {
use rand::RngExt;
let count = u32::from(shard.count());
let strides = 256 / count;
let stride = rand::rng().random_range(0..strides);
(stride * count + u32::from(shard.index())) as u8
}
#[cfg(target_os = "linux")]
fn attach(socket: &UdpSocket, count: u16) -> io::Result<()> {
use std::os::fd::AsRawFd;
let program = program(count);
let fprog = libc::sock_fprog {
len: program.len() as u16,
filter: program.as_ptr() as *mut libc::sock_filter,
};
let res = unsafe {
libc::setsockopt(
socket.as_raw_fd(),
libc::SOL_SOCKET,
libc::SO_ATTACH_REUSEPORT_CBPF,
std::ptr::from_ref(&fprog).cast(),
size_of::<libc::sock_fprog>() as libc::socklen_t,
)
};
if res != 0 {
return Err(io::Error::last_os_error());
}
tracing::debug!(count, "steering the reuseport group by connection ID");
Ok(())
}
#[cfg(not(target_os = "linux"))]
fn attach(_socket: &UdpSocket, _count: u16) -> io::Result<()> {
Err(io::Error::new(
io::ErrorKind::Unsupported,
"reuseport steering is Linux-only",
))
}
#[cfg(target_os = "linux")]
fn program(count: u16) -> [libc::sock_filter; 7] {
const LONG_HEADER: u32 = 0x80;
const LONG_DCID: u32 = 6;
const SHORT_DCID: u32 = 1;
fn insn(code: u32, jt: u8, jf: u8, k: u32) -> libc::sock_filter {
libc::sock_filter {
code: code as u16,
jt,
jf,
k,
}
}
[
insn(libc::BPF_LD | libc::BPF_B | libc::BPF_ABS, 0, 0, 0),
insn(libc::BPF_JMP | libc::BPF_JSET | libc::BPF_K, 2, 0, LONG_HEADER),
insn(libc::BPF_LD | libc::BPF_B | libc::BPF_ABS, 0, 0, SHORT_DCID),
insn(libc::BPF_JMP | libc::BPF_JA | libc::BPF_K, 0, 0, 1),
insn(libc::BPF_LD | libc::BPF_B | libc::BPF_ABS, 0, 0, LONG_DCID),
insn(libc::BPF_ALU | libc::BPF_MOD | libc::BPF_K, 0, 0, u32::from(count)),
insn(libc::BPF_RET | libc::BPF_A, 0, 0, 0),
]
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn shard_slots_are_bounded() {
assert_eq!(Shard::new(0, 1).map(|shard| shard.count()), Some(1));
assert_eq!(Shard::new(3, 4).map(|shard| shard.index()), Some(3));
assert!(Shard::new(4, 4).is_none());
assert!(Shard::new(0, 0).is_none());
assert!(Shard::new(0, MAX_SHARDS).is_some());
assert!(Shard::new(0, MAX_SHARDS + 1).is_none());
}
#[test]
fn a_group_hands_out_every_slot_in_order() {
const COUNT: u16 = 4;
let mut group = Group::acquire("127.0.0.1:0".parse().unwrap(), COUNT).unwrap();
assert_eq!(group.count(), COUNT);
for index in 0..COUNT {
let member = group.member().expect("a slot per member");
assert_eq!(member.shard().index(), index);
assert_eq!(member.shard().count(), COUNT);
}
assert!(group.member().is_none(), "the group cannot be resized");
}
#[test]
fn an_empty_group_holds_one_member() {
let mut group = Group::acquire("127.0.0.1:0".parse().unwrap(), 0).unwrap();
assert_eq!(group.count(), 1);
assert_eq!(group.member().map(|member| member.shard().count()), Some(1));
}
#[test]
fn an_unaddressable_group_is_refused() {
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
assert!(Group::acquire(addr, MAX_SHARDS).is_ok());
assert!(matches!(
Group::acquire(addr, MAX_SHARDS + 1),
Err(Error::Count { count: 257, max: 256 })
));
}
#[test]
fn a_prefix_reduces_to_its_own_shard() {
for count in [1u16, 2, 3, 4, 7, 8, 16, 64, 255, MAX_SHARDS] {
for index in 0..count {
let shard = Shard::new(index, count).unwrap();
for _ in 0..64 {
let prefix = cid_prefix(shard);
assert_eq!(
u16::from(prefix) % count,
index,
"prefix {prefix} of shard {index}/{count} steers elsewhere"
);
}
}
}
}
#[test]
fn prefixes_cover_every_shard() {
const COUNT: u16 = 8;
let mut seen = std::collections::HashSet::new();
for index in 0..COUNT {
let shard = Shard::new(index, COUNT).unwrap();
for _ in 0..256 {
seen.insert(u16::from(cid_prefix(shard)) % COUNT);
}
}
assert_eq!(seen.len(), usize::from(COUNT));
}
#[test]
#[cfg(target_os = "linux")]
fn a_group_shares_one_ephemeral_port() {
const COUNT: u16 = 3;
let mut group = Group::acquire("127.0.0.1:0".parse().unwrap(), COUNT).unwrap();
let claims: Vec<Claim> = (0..COUNT)
.map(|_| group.member().expect("a slot per member").bind().expect("bind member"))
.collect();
let mut group = group.complete(claims).expect("complete group");
let sockets: Vec<UdpSocket> = (0..COUNT)
.map(|_| group.member().unwrap().expect("a socket per member").into_inner())
.collect();
let addr = group.addr();
assert_ne!(addr.port(), 0, "the group holds the port its first member bound");
for socket in &sockets {
assert_eq!(
socket.local_addr().unwrap(),
addr,
"every member holds the group's port"
);
}
}
#[test]
#[cfg(target_os = "linux")]
fn binding_out_of_order_is_refused() {
let mut group = Group::acquire("127.0.0.1:0".parse().unwrap(), 2).unwrap();
let first = group.member().expect("first slot");
let second = group.member().expect("second slot");
second.bind().expect_err("the second member cannot bind first");
first.bind().expect("bind the first member");
assert_ne!(group.addr().port(), 0);
}
#[test]
#[cfg(target_os = "linux")]
fn a_partial_group_cannot_serve() {
let mut group = Group::acquire("127.0.0.1:0".parse().unwrap(), 2).unwrap();
let claim = group.member().unwrap().bind().expect("bind first member");
group
.complete(vec![claim])
.expect_err("a partial group must not complete");
}
#[test]
#[cfg(target_os = "linux")]
fn a_connection_id_reaches_its_own_member() {
const COUNT: u16 = 4;
let mut group = Group::acquire("127.0.0.1:0".parse().unwrap(), COUNT).unwrap();
let mut claims = Vec::new();
while let Some(member) = group.member() {
claims.push(member.bind().expect("bind group member"));
}
let mut group = group.complete(claims).expect("complete group");
let mut sockets = Vec::new();
let mut shards = Vec::new();
while let Some(member) = group.member().expect("clone retained socket") {
shards.push(member.shard());
let socket = member.into_inner();
socket.set_nonblocking(true).unwrap();
sockets.push(socket);
}
let addr = group.addr();
fn receiver(group: &[UdpSocket]) -> Option<usize> {
let mut buf = [0u8; 64];
group.iter().position(|socket| socket.recv_from(&mut buf).is_ok())
}
for (index, shard) in shards.into_iter().enumerate() {
let prefix = cid_prefix(shard);
let short = [0x40, prefix, 1, 2, 3, 4, 5, 6, 7, 8];
let long = [0xc0, 0, 0, 0, 1, 8, prefix, 1, 2, 3, 4, 5, 6, 7];
for (form, packet) in [("short", &short[..]), ("long", &long[..])] {
let sender = crate::bind::udp(crate::bind::Udp::new("127.0.0.1:0".parse().unwrap())).unwrap();
sender.send_to(packet, addr).unwrap();
std::thread::sleep(std::time::Duration::from_millis(50));
assert_eq!(
receiver(&sockets),
Some(index),
"{form} header for member {index} landed on the wrong socket"
);
}
}
}
#[test]
#[cfg(target_os = "linux")]
fn dropping_a_serving_handle_does_not_resize_the_group() {
let mut group = Group::acquire("127.0.0.1:0".parse().unwrap(), 2).unwrap();
let mut claims = Vec::new();
while let Some(member) = group.member() {
claims.push(member.bind().expect("bind group member"));
}
let mut group = group.complete(claims).expect("complete group");
let first = group.member().unwrap().expect("first serving socket");
let second = group.member().unwrap().expect("second serving socket");
let shard = second.shard();
let socket = second.into_inner();
socket
.set_read_timeout(Some(std::time::Duration::from_secs(1)))
.unwrap();
drop(first);
let packet = [0x40, cid_prefix(shard), 1, 2, 3, 4, 5, 6, 7, 8];
let sender = crate::bind::udp(crate::bind::Udp::new("127.0.0.1:0".parse().unwrap())).unwrap();
sender.send_to(&packet, group.addr()).unwrap();
let mut buf = [0; 16];
assert_eq!(socket.recv_from(&mut buf).unwrap().0, packet.len());
}
#[test]
#[cfg(target_os = "linux")]
fn a_second_group_cannot_take_the_port() {
let addr = {
let probe = crate::bind::udp(crate::bind::Udp::new("127.0.0.1:0".parse().unwrap())).unwrap();
probe.local_addr().unwrap()
};
let mut first = Group::acquire(addr, 1).unwrap();
let claim = first.member().unwrap().bind().expect("bind the first group");
let first = first.complete(vec![claim]).expect("complete the first group");
assert!(
matches!(Group::acquire(addr, 1), Err(Error::Overlap { .. })),
"a second group took a held port"
);
drop(first);
Group::acquire(addr, 1).expect("the released port must be takeable again");
}
#[test]
#[cfg(target_os = "linux")]
fn an_ephemeral_group_takes_its_port_up_front() {
let group = Group::acquire("127.0.0.1:0".parse().unwrap(), 1).unwrap();
let addr = group.addr();
assert_ne!(addr.port(), 0, "the port is resolved before any member binds");
assert!(
matches!(Group::acquire(addr, 1), Err(Error::Overlap { .. })),
"a second group took an ephemeral group's port"
);
}
#[test]
#[cfg(target_os = "linux")]
fn an_ephemeral_port_taken_before_binding_is_refused() {
let mut group = Group::acquire("127.0.0.1:0".parse().unwrap(), 1).unwrap();
let _intruder = crate::bind::udp(crate::bind::Udp::new(group.addr()).with_reuse_port(true)).unwrap();
let err = group.member().unwrap().bind().expect_err("joined the intruder's group");
assert_eq!(err.kind(), io::ErrorKind::AddrInUse);
}
#[test]
#[cfg(target_os = "linux")]
fn an_outstanding_member_keeps_the_port() {
let addr = {
let probe = crate::bind::udp(crate::bind::Udp::new("127.0.0.1:0".parse().unwrap())).unwrap();
probe.local_addr().unwrap()
};
let mut group = Group::acquire(addr, 1).unwrap();
let member = group.member().expect("the only slot");
drop(group);
assert!(
matches!(Group::acquire(addr, 1), Err(Error::Overlap { .. })),
"a second group took a port an unbound member still holds"
);
let claim = member.bind().expect("bind the outstanding member");
drop(claim);
Group::acquire(addr, 1).expect("the released port must be takeable again");
}
#[test]
#[cfg(target_os = "linux")]
fn the_program_branches_to_the_right_loads() {
let program = program(4);
assert_eq!(program.len(), 7);
assert_eq!(program[1].jt, 2, "long header must land on the long load");
assert_eq!(program[1].jf, 0, "short header must fall through");
assert_eq!(program[2].k, 1, "short header reads the byte after the first");
assert_eq!(program[3].k, 1, "the short path must skip the long load");
assert_eq!(program[4].k, 6, "long header reads past version and length");
assert_eq!(program[5].k, 4, "the modulus is the group size");
}
}