use alloc::{vec, vec::Vec};
use ax_sync::Mutex;
use hashbrown::HashMap;
use smoltcp::{
iface::{SocketHandle, SocketSet},
socket::AnySocket,
wire::IpAddress,
};
use crate::{NetError, NetResult, addr::listen_addrs_conflict};
#[derive(Clone, Debug)]
struct UdpBoundEntry {
addr: Option<IpAddress>,
reuse_port: bool,
handle: SocketHandle,
}
pub(crate) struct SocketSetWrapper<'a> {
pub inner: Mutex<SocketSet<'a>>,
udp_binds: Mutex<HashMap<u16, Vec<UdpBoundEntry>>>,
}
impl<'a> SocketSetWrapper<'a> {
pub fn new() -> Self {
Self {
inner: Mutex::new(SocketSet::new(vec![])),
udp_binds: Mutex::new(HashMap::new()),
}
}
pub fn add<T: AnySocket<'a>>(&self, socket: T) -> SocketHandle {
let handle = self.inner.lock().add(socket);
debug!("socket {}: created", handle);
handle
}
pub fn with_socket_mut<T: AnySocket<'a>, R, F>(&self, handle: SocketHandle, f: F) -> R
where
F: FnOnce(&mut T) -> R,
{
let mut set = self.inner.lock();
let socket = set.get_mut(handle);
f(socket)
}
pub fn udp_bind(
&self,
handle: SocketHandle,
addr: IpAddress,
port: u16,
reuse_port: bool,
) -> NetResult {
if port == 0 {
return Ok(());
}
let addr = (!addr.is_unspecified()).then_some(addr);
let mut binds = self.udp_binds.lock();
let entries = binds.entry(port).or_default();
if entries
.iter()
.any(|entry| udp_binds_conflict(entry, addr, reuse_port))
{
return Err(NetError::AddrInUse);
}
entries.push(UdpBoundEntry {
addr,
reuse_port,
handle,
});
Ok(())
}
pub fn udp_port_available(&self, addr: IpAddress, port: u16) -> bool {
if port == 0 {
return true;
}
let addr = (!addr.is_unspecified()).then_some(addr);
match self.udp_binds.lock().get(&port) {
None => true,
Some(entries) => !entries
.iter()
.any(|entry| listen_addrs_conflict(entry.addr, addr)),
}
}
pub fn udp_unbind(&self, handle: SocketHandle) {
self.udp_binds.lock().retain(|_, entries| {
entries.retain(|entry| entry.handle != handle);
!entries.is_empty()
});
}
pub fn remove(&self, handle: SocketHandle) {
self.udp_unbind(handle);
self.inner.lock().remove(handle);
debug!("socket {}: destroyed", handle);
}
}
fn udp_binds_conflict(entry: &UdpBoundEntry, addr: Option<IpAddress>, reuse_port: bool) -> bool {
listen_addrs_conflict(entry.addr, addr)
&& !(reuse_port && entry.reuse_port && entry.addr == addr)
}
#[cfg(test)]
mod tests {
use alloc::vec;
use smoltcp::{
iface::SocketSet,
socket::udp,
storage::PacketMetadata,
wire::{IpAddress, Ipv4Address},
};
use super::*;
fn addr(a: u8, b: u8, c: u8, d: u8) -> IpAddress {
IpAddress::Ipv4(Ipv4Address::new(a, b, c, d))
}
fn entry(addr: Option<IpAddress>, reuse_port: bool) -> UdpBoundEntry {
let mut sockets = SocketSet::new(vec![]);
let handle = sockets.add(udp::Socket::new(
udp::PacketBuffer::new(vec![PacketMetadata::EMPTY; 1], vec![0; 8]),
udp::PacketBuffer::new(vec![PacketMetadata::EMPTY; 1], vec![0; 8]),
));
UdpBoundEntry {
handle,
addr,
reuse_port,
}
}
#[test]
fn wildcard_and_identical_udp_bindings_conflict() {
let specific = Some(addr(192, 0, 2, 10));
let owner = entry(specific, false);
assert!(udp_binds_conflict(&owner, specific, false));
assert!(udp_binds_conflict(&owner, None, false));
assert!(!udp_binds_conflict(
&owner,
Some(addr(198, 51, 100, 20)),
false,
));
}
#[test]
fn udp_reuseport_requires_an_exact_reuseport_group() {
let specific = Some(addr(127, 0, 0, 1));
let plain = entry(specific, false);
let reuse = entry(specific, true);
assert!(udp_binds_conflict(&plain, specific, true));
assert!(!udp_binds_conflict(&reuse, specific, true));
assert!(udp_binds_conflict(&reuse, None, true));
}
}