use crate::sockets::SocketAddressFamily;
use core::fmt;
use core::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use core::str::FromStr as _;
use core::time::Duration;
use rustix::fd::AsFd;
use rustix::io::Errno;
use rustix::net::{bind, connect, connect_unspec, sockopt};
use tracing::debug;
#[derive(Debug)]
pub enum ErrorCode {
Unknown,
AccessDenied,
NotSupported,
InvalidArgument,
OutOfMemory,
Timeout,
InvalidState,
AddressNotBindable,
AddressInUse,
RemoteUnreachable,
ConnectionRefused,
ConnectionReset,
ConnectionAborted,
DatagramTooLarge,
NotInProgress,
ConcurrencyConflict,
}
impl fmt::Display for ErrorCode {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt::Debug::fmt(self, f)
}
}
impl std::error::Error for ErrorCode {}
fn is_deprecated_ipv4_compatible(addr: Ipv6Addr) -> bool {
matches!(addr.segments(), [0, 0, 0, 0, 0, 0, _, _])
&& addr != Ipv6Addr::UNSPECIFIED
&& addr != Ipv6Addr::LOCALHOST
}
pub fn is_valid_address_family(addr: IpAddr, socket_family: SocketAddressFamily) -> bool {
match (socket_family, addr) {
(SocketAddressFamily::Ipv4, IpAddr::V4(..)) => true,
(SocketAddressFamily::Ipv6, IpAddr::V6(ipv6)) => {
!is_deprecated_ipv4_compatible(ipv6) && ipv6.to_ipv4_mapped().is_none()
}
_ => false,
}
}
pub fn is_valid_remote_address(addr: SocketAddr) -> bool {
!addr.ip().to_canonical().is_unspecified() && addr.port() != 0
}
pub fn is_valid_unicast_address(addr: IpAddr) -> bool {
match addr.to_canonical() {
IpAddr::V4(ipv4) => !ipv4.is_multicast() && !ipv4.is_broadcast(),
IpAddr::V6(ipv6) => !ipv6.is_multicast(),
}
}
pub fn to_ipv4_addr(addr: (u8, u8, u8, u8)) -> Ipv4Addr {
let (x0, x1, x2, x3) = addr;
Ipv4Addr::new(x0, x1, x2, x3)
}
pub fn from_ipv4_addr(addr: Ipv4Addr) -> (u8, u8, u8, u8) {
let [x0, x1, x2, x3] = addr.octets();
(x0, x1, x2, x3)
}
pub fn to_ipv6_addr(addr: (u16, u16, u16, u16, u16, u16, u16, u16)) -> Ipv6Addr {
let (x0, x1, x2, x3, x4, x5, x6, x7) = addr;
Ipv6Addr::new(x0, x1, x2, x3, x4, x5, x6, x7)
}
pub fn from_ipv6_addr(addr: Ipv6Addr) -> (u16, u16, u16, u16, u16, u16, u16, u16) {
let [x0, x1, x2, x3, x4, x5, x6, x7] = addr.segments();
(x0, x1, x2, x3, x4, x5, x6, x7)
}
pub fn normalize_get_buffer_size(value: usize) -> usize {
if cfg!(target_os = "linux") {
value / 2
} else {
value
}
}
pub fn normalize_set_buffer_size(value: usize) -> usize {
value.clamp(1, i32::MAX as usize)
}
impl From<std::io::Error> for ErrorCode {
fn from(value: std::io::Error) -> Self {
(&value).into()
}
}
impl From<&std::io::Error> for ErrorCode {
fn from(value: &std::io::Error) -> Self {
if let Some(errno) = Errno::from_io_error(value) {
return errno.into();
}
match value.kind() {
std::io::ErrorKind::AddrInUse => Self::AddressInUse,
std::io::ErrorKind::AddrNotAvailable => Self::AddressNotBindable,
std::io::ErrorKind::ConnectionAborted => Self::ConnectionAborted,
std::io::ErrorKind::ConnectionRefused => Self::ConnectionRefused,
std::io::ErrorKind::ConnectionReset => Self::ConnectionReset,
std::io::ErrorKind::InvalidInput => Self::InvalidArgument,
std::io::ErrorKind::NotConnected => Self::InvalidState,
std::io::ErrorKind::OutOfMemory => Self::OutOfMemory,
std::io::ErrorKind::PermissionDenied => Self::AccessDenied,
std::io::ErrorKind::TimedOut => Self::Timeout,
std::io::ErrorKind::Unsupported => Self::NotSupported,
_ => {
debug!("unknown I/O error: {value}");
Self::Unknown
}
}
}
}
impl From<Errno> for ErrorCode {
fn from(value: Errno) -> Self {
(&value).into()
}
}
impl From<&Errno> for ErrorCode {
fn from(value: &Errno) -> Self {
match *value {
#[cfg(not(windows))]
Errno::PERM => Self::AccessDenied,
Errno::ACCESS => Self::AccessDenied,
Errno::ADDRINUSE => Self::AddressInUse,
Errno::ADDRNOTAVAIL => Self::AddressNotBindable,
Errno::TIMEDOUT => Self::Timeout,
Errno::CONNREFUSED => Self::ConnectionRefused,
Errno::CONNRESET => Self::ConnectionReset,
Errno::CONNABORTED => Self::ConnectionAborted,
Errno::INVAL => Self::InvalidArgument,
Errno::HOSTUNREACH => Self::RemoteUnreachable,
Errno::HOSTDOWN => Self::RemoteUnreachable,
Errno::NETDOWN => Self::RemoteUnreachable,
Errno::NETUNREACH => Self::RemoteUnreachable,
#[cfg(target_os = "linux")]
Errno::NONET => Self::RemoteUnreachable,
Errno::ISCONN => Self::InvalidState,
Errno::NOTCONN => Self::InvalidState,
Errno::DESTADDRREQ => Self::InvalidState,
Errno::MSGSIZE => Self::DatagramTooLarge,
#[cfg(not(windows))]
Errno::NOMEM => Self::OutOfMemory,
Errno::NOBUFS => Self::OutOfMemory,
Errno::OPNOTSUPP => Self::NotSupported,
Errno::NOPROTOOPT => Self::NotSupported,
Errno::PFNOSUPPORT => Self::NotSupported,
Errno::PROTONOSUPPORT => Self::NotSupported,
Errno::PROTOTYPE => Self::NotSupported,
Errno::SOCKTNOSUPPORT => Self::NotSupported,
Errno::AFNOSUPPORT => Self::NotSupported,
_ => {
debug!("unknown I/O error: {value}");
Self::Unknown
}
}
}
}
pub fn get_ip_ttl(fd: impl AsFd) -> Result<u8, ErrorCode> {
let v = sockopt::ip_ttl(fd)?;
let Ok(v) = v.try_into() else {
return Err(ErrorCode::NotSupported);
};
Ok(v)
}
pub fn get_ipv6_unicast_hops(fd: impl AsFd) -> Result<u8, ErrorCode> {
let v = sockopt::ipv6_unicast_hops(fd)?;
Ok(v)
}
pub fn get_unicast_hop_limit(fd: impl AsFd, family: SocketAddressFamily) -> Result<u8, ErrorCode> {
match family {
SocketAddressFamily::Ipv4 => get_ip_ttl(fd),
SocketAddressFamily::Ipv6 => get_ipv6_unicast_hops(fd),
}
}
pub fn set_unicast_hop_limit(
fd: impl AsFd,
family: SocketAddressFamily,
value: u8,
) -> Result<(), ErrorCode> {
if value == 0 {
return Err(ErrorCode::InvalidArgument);
}
match family {
SocketAddressFamily::Ipv4 => {
sockopt::set_ip_ttl(fd, value.into())?;
}
SocketAddressFamily::Ipv6 => {
sockopt::set_ipv6_unicast_hops(fd, Some(value))?;
}
}
Ok(())
}
pub fn receive_buffer_size(fd: impl AsFd) -> Result<u64, ErrorCode> {
let v = sockopt::socket_recv_buffer_size(fd)?;
Ok(normalize_get_buffer_size(v).try_into().unwrap_or(u64::MAX))
}
pub fn set_receive_buffer_size(fd: impl AsFd, value: u64) -> Result<usize, ErrorCode> {
if value == 0 {
return Err(ErrorCode::InvalidArgument);
}
let value = value.try_into().unwrap_or(usize::MAX);
let value = normalize_set_buffer_size(value);
match sockopt::set_socket_recv_buffer_size(fd, value) {
Err(Errno::NOBUFS) => {}
Err(err) => return Err(err.into()),
_ => {}
};
Ok(value)
}
pub fn send_buffer_size(fd: impl AsFd) -> Result<u64, ErrorCode> {
let v = sockopt::socket_send_buffer_size(fd)?;
Ok(normalize_get_buffer_size(v).try_into().unwrap_or(u64::MAX))
}
pub fn set_send_buffer_size(fd: impl AsFd, value: u64) -> Result<usize, ErrorCode> {
if value == 0 {
return Err(ErrorCode::InvalidArgument);
}
let value = value.try_into().unwrap_or(usize::MAX);
let value = normalize_set_buffer_size(value);
match sockopt::set_socket_send_buffer_size(fd, value) {
Err(Errno::NOBUFS) => {}
Err(err) => return Err(err.into()),
_ => {}
};
Ok(value)
}
pub fn set_keep_alive_idle_time(fd: impl AsFd, value: u64) -> Result<u64, ErrorCode> {
const NANOS_PER_SEC: u64 = 1_000_000_000;
const MIN: u64 = NANOS_PER_SEC;
const MAX: u64 = (i16::MAX as u64) * NANOS_PER_SEC;
if value <= 0 {
return Err(ErrorCode::InvalidArgument);
}
let value = value.clamp(MIN, MAX);
sockopt::set_tcp_keepidle(fd, Duration::from_nanos(value))?;
Ok(value)
}
pub fn set_keep_alive_interval(fd: impl AsFd, value: Duration) -> Result<(), ErrorCode> {
const MIN: Duration = Duration::from_secs(1);
const MAX: Duration = Duration::from_secs(i16::MAX as u64);
if value <= Duration::ZERO {
return Err(ErrorCode::InvalidArgument);
}
sockopt::set_tcp_keepintvl(fd, value.clamp(MIN, MAX))?;
Ok(())
}
pub fn set_keep_alive_count(fd: impl AsFd, value: u32) -> Result<(), ErrorCode> {
const MIN_CNT: u32 = 1;
const MAX_CNT: u32 = i8::MAX as u32;
if value == 0 {
return Err(ErrorCode::InvalidArgument);
}
sockopt::set_tcp_keepcnt(fd, value.clamp(MIN_CNT, MAX_CNT))?;
Ok(())
}
pub fn tcp_bind(
socket: &tokio::net::TcpSocket,
local_address: SocketAddr,
) -> Result<(), ErrorCode> {
#[cfg(not(windows))]
{
_ = sockopt::set_socket_reuseaddr(&socket, true);
}
socket
.bind(local_address)
.map_err(|err| match Errno::from_io_error(&err) {
Some(Errno::AFNOSUPPORT) => ErrorCode::InvalidArgument,
#[cfg(windows)]
Some(Errno::NOBUFS) => ErrorCode::AddressInUse,
_ => err.into(),
})
}
pub fn udp_socket(family: SocketAddressFamily) -> std::io::Result<rustix::fd::OwnedFd> {
#[cfg(windows)]
static INIT: std::sync::Once = std::sync::Once::new();
#[cfg(windows)]
INIT.call_once(|| {
let _ = std::net::TcpStream::connect(std::net::SocketAddrV4::new(
std::net::Ipv4Addr::UNSPECIFIED,
0,
));
});
#[cfg(not(any(windows, target_vendor = "apple")))]
let flags = rustix::net::SocketFlags::CLOEXEC | rustix::net::SocketFlags::NONBLOCK;
#[cfg(any(windows, target_vendor = "apple"))]
let flags = rustix::net::SocketFlags::empty();
let socket = rustix::net::socket_with(
match family {
SocketAddressFamily::Ipv4 => rustix::net::AddressFamily::INET,
SocketAddressFamily::Ipv6 => rustix::net::AddressFamily::INET6,
},
rustix::net::SocketType::DGRAM,
flags,
None,
)?;
#[cfg(target_vendor = "apple")]
rustix::io::ioctl_fioclex(&socket)?;
#[cfg(any(windows, target_vendor = "apple"))]
rustix::io::ioctl_fionbio(&socket, true)?;
if family == SocketAddressFamily::Ipv6 {
rustix::net::sockopt::set_ipv6_v6only(&socket, true)?;
}
Ok(socket)
}
pub fn udp_bind(sockfd: impl AsFd, addr: SocketAddr) -> Result<(), ErrorCode> {
bind(sockfd, &addr).map_err(|err| match err {
#[cfg(windows)]
Errno::NOBUFS => ErrorCode::AddressInUse,
Errno::AFNOSUPPORT => ErrorCode::InvalidArgument,
_ => err.into(),
})
}
pub fn udp_connect(sockfd: impl AsFd, addr: SocketAddr) -> Result<(), Errno> {
match connect(sockfd.as_fd(), &addr) {
#[cfg(target_os = "linux")]
Err(Errno::INVAL) => {
_ = udp_disconnect(sockfd.as_fd());
return connect(sockfd.as_fd(), &addr);
}
r => r,
}
}
pub fn udp_disconnect(sockfd: impl AsFd) -> Result<(), Errno> {
match connect_unspec(sockfd) {
#[cfg(target_os = "macos")]
Err(Errno::INVAL | Errno::AFNOSUPPORT) => Ok(()),
r => r,
}
}
pub fn parse_host(name: &str) -> Result<url::Host, ErrorCode> {
match url::Host::parse(&name) {
Ok(host) => Ok(host),
Err(_) => {
if let Ok(addr) = Ipv6Addr::from_str(name) {
Ok(url::Host::Ipv6(addr))
} else {
Err(ErrorCode::InvalidArgument)
}
}
}
}
#[cfg(feature = "p3")]
pub fn implicit_bind_addr(family: SocketAddressFamily) -> SocketAddr {
let ip = match family {
SocketAddressFamily::Ipv4 => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
SocketAddressFamily::Ipv6 => IpAddr::V6(Ipv6Addr::UNSPECIFIED),
};
SocketAddr::new(ip, 0)
}