use std::mem::MaybeUninit;
use std::net::{SocketAddr, SocketAddrV4, SocketAddrV6};
use std::os::fd::AsRawFd;
use socket2::{Domain, MsgHdrMut, Protocol, Socket, Type};
use tokio::io::unix::{AsyncFd, AsyncFdReadyGuard};
use tokio::io::Interest;
use crate::addr::ToIpAddr;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum SocketType {
Raw,
Dgram,
}
struct NewSocket {
socket: Socket,
sock_type: SocketType,
}
fn new_icmp_socket(domain: Domain, protocol: Protocol) -> std::io::Result<NewSocket> {
#[cfg(any(target_os = "linux", target_os = "android",))]
{
let sock = Socket::new(domain, Type::DGRAM, Some(protocol));
match sock {
Ok(socket) => {
return Ok(NewSocket {
socket,
sock_type: SocketType::Dgram,
});
}
Err(e) => {
let fallback = matches!(
e.raw_os_error(),
Some(libc::EACCES | libc::EAFNOSUPPORT | libc::EPROTONOSUPPORT)
);
if fallback {
let raw = Socket::new(domain, Type::RAW, Some(protocol))?;
return Ok(NewSocket {
socket: raw,
sock_type: SocketType::Raw,
});
}
return Err(e);
}
}
}
#[cfg(any(
target_os = "macos",
target_os = "ios",
target_os = "tvos",
target_os = "watchos",
target_os = "visionos",
))]
if !is_root() {
return Ok(NewSocket {
socket: Socket::new(domain, Type::DGRAM, Some(protocol))?,
sock_type: SocketType::Dgram,
});
}
#[cfg_attr(
any(target_os = "linux", target_os = "android",),
allow(unreachable_code)
)]
{
Ok(NewSocket {
socket: Socket::new(domain, Type::RAW, Some(protocol))?,
sock_type: SocketType::Raw,
})
}
}
#[cfg(any(
target_os = "macos",
target_os = "ios",
target_os = "tvos",
target_os = "watchos",
target_os = "visionos",
))]
fn is_root() -> bool {
unsafe { libc::getuid() == 0 }
}
pub struct IcmpSocket {
io: AsyncFd<Socket>,
sock_type: SocketType,
#[cfg(any(target_os = "linux", target_os = "android"))]
dgram_ident: Option<u16>,
}
impl IcmpSocket {
pub async fn bind<A: ToIpAddr>(addr: A) -> std::io::Result<IcmpSocket> {
let ip_addr = addr.to_ip_addr().await?;
let (domain, protocol) = match ip_addr {
std::net::IpAddr::V4(_) => (Domain::IPV4, Protocol::ICMPV4),
std::net::IpAddr::V6(_) => (Domain::IPV6, Protocol::ICMPV6),
};
let NewSocket { socket, sock_type } = new_icmp_socket(domain, protocol)?;
socket.set_nonblocking(true)?;
#[cfg(any(target_os = "linux", target_os = "android"))]
let dgram_ident = if sock_type == SocketType::Dgram {
use std::sync::atomic::Ordering;
let ident = loop {
let candidate = crate::REQ_ID.fetch_add(1, Ordering::Relaxed);
if candidate != 0 {
break candidate;
}
};
Some(ident)
} else {
None
};
if sock_type == SocketType::Dgram && domain == Domain::IPV4 {
let hold: libc::c_int = 1;
let _ = unsafe {
libc::setsockopt(
socket.as_raw_fd(),
libc::IPPROTO_IP,
libc::IP_RECVTTL,
(&raw const hold).cast(),
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
)
};
}
if domain == Domain::IPV6 {
socket.set_recv_hoplimit_v6(true)?;
}
let skip_dontfrag = {
#[cfg(any(
target_os = "macos",
target_os = "ios",
target_os = "tvos",
target_os = "watchos",
target_os = "visionos",
))]
{
domain == Domain::IPV6 && !is_root()
}
#[cfg(not(any(
target_os = "macos",
target_os = "ios",
target_os = "tvos",
target_os = "watchos",
target_os = "visionos",
)))]
{
false
}
};
if !skip_dontfrag {
set_dont_fragment(&socket, domain, true)?;
}
#[cfg(any(target_os = "linux", target_os = "android"))]
let bind_port = dgram_ident.unwrap_or(0);
#[cfg(not(any(target_os = "linux", target_os = "android")))]
let bind_port = 0u16;
let sock_addr = match ip_addr {
std::net::IpAddr::V4(ipv4_addr) => {
SocketAddr::V4(SocketAddrV4::new(ipv4_addr, bind_port))
}
std::net::IpAddr::V6(ipv6_addr) => {
SocketAddr::V6(SocketAddrV6::new(ipv6_addr, bind_port, 0, 0))
}
};
socket.bind(&sock_addr.into())?;
let io = AsyncFd::new(socket)?;
Ok(Self {
io,
sock_type,
#[cfg(any(target_os = "linux", target_os = "android"))]
dgram_ident,
})
}
pub async fn connect<A: ToIpAddr>(&self, addr: A) -> std::io::Result<()> {
let ip_addr = addr.to_ip_addr().await?;
let socket_addr = match ip_addr {
std::net::IpAddr::V4(ipv4_addr) => SocketAddr::V4(SocketAddrV4::new(ipv4_addr, 0u16)),
std::net::IpAddr::V6(ipv6_addr) => {
SocketAddr::V6(SocketAddrV6::new(ipv6_addr, 0u16, 0, 0))
}
};
self.io.get_ref().connect(&socket_addr.into())
}
pub(crate) fn sock_type(&self) -> SocketType {
self.sock_type
}
#[cfg(any(target_os = "linux", target_os = "android"))]
pub(crate) fn dgram_ident(&self) -> Option<u16> {
self.dgram_ident
}
pub async fn ready(
&self,
interest: Interest,
) -> std::io::Result<AsyncFdReadyGuard<'_, Socket>> {
self.io.ready(interest).await
}
pub async fn writable(&self) -> std::io::Result<()> {
let _ = self.ready(Interest::WRITABLE).await?;
Ok(())
}
pub async fn send(&self, buf: &[u8]) -> std::io::Result<usize> {
self.io.async_io(Interest::WRITABLE, |s| s.send(buf)).await
}
pub async fn readable(&self) -> std::io::Result<()> {
let _ = self.ready(Interest::READABLE).await?;
Ok(())
}
pub async fn recv(&self, buf: &mut [MaybeUninit<u8>]) -> std::io::Result<usize> {
self.io.async_io(Interest::READABLE, |s| s.recv(buf)).await
}
pub(crate) async fn recvmsg(&self, msg: &mut MsgHdrMut<'_, '_, '_>) -> std::io::Result<usize> {
self.io
.async_io(Interest::READABLE, |s| s.recvmsg(msg, 0))
.await
}
}
#[cfg(any(
target_os = "linux",
target_os = "l4re",
target_os = "android",
target_os = "emscripten"
))]
fn set_dont_fragment(socket: &Socket, domain: Domain, dont_fragment: bool) -> std::io::Result<()> {
match domain {
Domain::IPV4 => {
let payload = if dont_fragment {
libc::IP_PMTUDISC_DO
} else {
libc::IP_PMTUDISC_DONT
};
unsafe { setsockopt(socket, libc::IPPROTO_IP, libc::IP_MTU_DISCOVER, payload) }
}
Domain::IPV6 => {
let payload = if dont_fragment {
libc::IPV6_PMTUDISC_DO
} else {
libc::IPV6_PMTUDISC_DONT
};
unsafe { setsockopt(socket, libc::IPPROTO_IPV6, libc::IPV6_MTU_DISCOVER, payload) }
}
_ => Ok(()),
}
}
#[cfg(any(
target_os = "macos",
target_os = "ios",
target_os = "tvos",
target_os = "watchos",
target_os = "visionos",
target_os = "freebsd",
target_os = "dragonfly",
target_os = "openbsd",
target_os = "netbsd"
))]
fn set_dont_fragment(socket: &Socket, domain: Domain, dont_fragment: bool) -> std::io::Result<()> {
match domain {
Domain::IPV4 => unsafe {
setsockopt(
socket,
libc::IPPROTO_IP,
libc::IP_DONTFRAG,
dont_fragment as libc::c_int,
)
},
Domain::IPV6 => unsafe {
setsockopt(
socket,
libc::IPPROTO_IPV6,
libc::IPV6_DONTFRAG,
dont_fragment as libc::c_int,
)
},
_ => Ok(()),
}
}
#[allow(clippy::needless_pass_by_value)]
unsafe fn setsockopt<T>(
socket: &Socket,
opt: libc::c_int,
val: libc::c_int,
payload: T,
) -> std::io::Result<()> {
let payload = (&raw const payload).cast();
let res = unsafe {
libc::setsockopt(
socket.as_raw_fd(),
opt,
val,
payload,
std::mem::size_of::<T>() as libc::socklen_t,
)
};
if res != 0 {
return Err(std::io::Error::last_os_error());
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use super::IcmpSocket;
#[tokio::test]
async fn bind_accepts_str_literal() {
IcmpSocket::bind("127.0.0.1").await.unwrap();
}
#[tokio::test]
async fn bind_accepts_owned_string() {
IcmpSocket::bind(String::from("127.0.0.1")).await.unwrap();
}
#[tokio::test]
async fn bind_accepts_ipv4addr() {
IcmpSocket::bind(Ipv4Addr::LOCALHOST).await.unwrap();
}
#[tokio::test]
async fn bind_accepts_ipv6addr() {
IcmpSocket::bind(Ipv6Addr::LOCALHOST).await.unwrap();
}
#[tokio::test]
async fn bind_accepts_ip_addr() {
IcmpSocket::bind(IpAddr::V4(Ipv4Addr::LOCALHOST))
.await
.unwrap();
}
#[tokio::test]
async fn connect_accepts_str_literal() {
let sock = IcmpSocket::bind(Ipv4Addr::LOCALHOST).await.unwrap();
sock.connect("127.0.0.1").await.unwrap();
}
#[tokio::test]
async fn connect_accepts_owned_string() {
let sock = IcmpSocket::bind(Ipv4Addr::LOCALHOST).await.unwrap();
sock.connect(String::from("127.0.0.1")).await.unwrap();
}
#[tokio::test]
async fn connect_accepts_ipv4addr() {
let sock = IcmpSocket::bind(Ipv4Addr::LOCALHOST).await.unwrap();
sock.connect(Ipv4Addr::LOCALHOST).await.unwrap();
}
#[tokio::test]
async fn connect_accepts_ipv6addr() {
let sock = IcmpSocket::bind(Ipv6Addr::LOCALHOST).await.unwrap();
sock.connect(Ipv6Addr::LOCALHOST).await.unwrap();
}
#[tokio::test]
async fn connect_accepts_ip_addr() {
let sock = IcmpSocket::bind(Ipv4Addr::LOCALHOST).await.unwrap();
sock.connect(IpAddr::V4(Ipv4Addr::LOCALHOST)).await.unwrap();
}
}