use std::io;
use std::net::{SocketAddr, TcpStream, UdpSocket};
use std::time::Duration;
#[cfg(any(
target_os = "android",
target_os = "fuchsia",
target_os = "linux",
target_os = "ios",
target_os = "visionos",
target_os = "macos",
target_os = "tvos",
target_os = "watchos",
target_os = "illumos",
target_os = "solaris",
))]
pub use imp::{install, requested, tcp_connect_timeout, udp_connect};
#[cfg(not(any(
target_os = "android",
target_os = "fuchsia",
target_os = "linux",
target_os = "ios",
target_os = "visionos",
target_os = "macos",
target_os = "tvos",
target_os = "watchos",
target_os = "illumos",
target_os = "solaris",
)))]
pub use stub::{install, requested, tcp_connect_timeout, udp_connect};
#[cfg(any(
target_os = "android",
target_os = "fuchsia",
target_os = "linux",
target_os = "ios",
target_os = "visionos",
target_os = "macos",
target_os = "tvos",
target_os = "watchos",
target_os = "illumos",
target_os = "solaris",
))]
mod imp {
use super::*;
use std::num::NonZeroU32;
use std::os::fd::{AsRawFd, OwnedFd};
use std::sync::OnceLock;
use std::time::Instant;
use socket2::{Domain, Protocol, Socket, Type};
struct Bound {
name: String,
#[allow(dead_code)] index: NonZeroU32,
}
static IFACE: OnceLock<Bound> = OnceLock::new();
pub fn install(name: &str) -> Result<(), String> {
let index =
name_to_index(name).ok_or_else(|| format!("no such network interface: '{name}'"))?;
let bound = Bound {
name: name.to_string(),
index,
};
let probe = Socket::new(Domain::IPV4, Type::STREAM, Some(Protocol::TCP))
.map_err(|e| format!("could not open a socket to test interface binding: {e}"))?;
apply(&probe, &bound, true)
.map_err(|e| format!("cannot bind to interface '{name}': {e}"))?;
let _ = IFACE.set(bound);
Ok(())
}
pub fn requested() -> Option<&'static str> {
IFACE.get().map(|b| b.name.as_str())
}
pub fn tcp_connect_timeout(addr: SocketAddr, timeout: Duration) -> io::Result<TcpStream> {
let socket = Socket::new(Domain::for_address(addr), Type::STREAM, Some(Protocol::TCP))?;
if let Some(bound) = IFACE.get() {
apply(&socket, bound, addr.is_ipv4())?;
}
connect_timeout(&socket, addr, timeout)?;
Ok(TcpStream::from(OwnedFd::from(socket)))
}
pub fn udp_connect(target: SocketAddr) -> io::Result<UdpSocket> {
let socket = Socket::new(
Domain::for_address(target),
Type::DGRAM,
Some(Protocol::UDP),
)?;
if let Some(bound) = IFACE.get() {
apply(&socket, bound, target.is_ipv4())?;
}
let local: SocketAddr = if target.is_ipv4() {
SocketAddr::from(([0u8; 4], 0))
} else {
SocketAddr::from(([0u16; 8], 0))
};
socket.bind(&local.into())?;
socket.connect(&target.into())?;
Ok(UdpSocket::from(OwnedFd::from(socket)))
}
fn connect_timeout(socket: &Socket, addr: SocketAddr, timeout: Duration) -> io::Result<()> {
socket.set_nonblocking(true)?;
let started = socket.connect(&addr.into());
match started {
Ok(()) => {
socket.set_nonblocking(false)?;
return Ok(());
}
Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => {}
Err(ref e) if e.raw_os_error() == Some(libc::EINPROGRESS) => {}
Err(e) => {
let _ = socket.set_nonblocking(false);
return Err(e);
}
}
let poll_result = poll_writable(socket, timeout);
socket.set_nonblocking(false)?;
poll_result?;
match socket.take_error()? {
Some(err) => Err(err),
None => Ok(()),
}
}
fn poll_writable(socket: &Socket, timeout: Duration) -> io::Result<()> {
let start = Instant::now();
let mut pollfd = libc::pollfd {
fd: socket.as_raw_fd(),
events: libc::POLLOUT,
revents: 0,
};
loop {
let elapsed = start.elapsed();
if elapsed >= timeout {
return Err(io::ErrorKind::TimedOut.into());
}
let remaining = (timeout - elapsed)
.as_millis()
.clamp(1, libc::c_int::MAX as u128) as libc::c_int;
let rv = unsafe { libc::poll(&mut pollfd, 1, remaining) };
if rv < 0 {
let e = io::Error::last_os_error();
if e.kind() == io::ErrorKind::Interrupted {
continue;
}
return Err(e);
}
if rv == 0 {
return Err(io::ErrorKind::TimedOut.into());
}
return Ok(());
}
}
fn apply(socket: &Socket, bound: &Bound, ipv4: bool) -> io::Result<()> {
#[cfg(any(target_os = "android", target_os = "fuchsia", target_os = "linux"))]
{
let _ = ipv4;
socket.bind_device(Some(bound.name.as_bytes()))
}
#[cfg(not(any(target_os = "android", target_os = "fuchsia", target_os = "linux")))]
{
if ipv4 {
socket.bind_device_by_index_v4(Some(bound.index))
} else {
socket.bind_device_by_index_v6(Some(bound.index))
}
}
}
fn name_to_index(name: &str) -> Option<NonZeroU32> {
let cname = std::ffi::CString::new(name).ok()?;
let index = unsafe { libc::if_nametoindex(cname.as_ptr()) };
NonZeroU32::new(index)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn unknown_interface_has_no_index() {
assert!(name_to_index("definitely-not-an-iface-9999").is_none());
}
#[test]
fn install_rejects_an_unknown_interface() {
let err = install("definitely-not-an-iface-9999").unwrap_err();
assert!(err.contains("no such network interface"), "got: {err}");
}
#[test]
fn loopback_resolves_to_an_index() {
let resolved = name_to_index("lo").is_some() || name_to_index("lo0").is_some();
assert!(resolved, "loopback interface should resolve to an index");
}
}
}
#[cfg(not(any(
target_os = "android",
target_os = "fuchsia",
target_os = "linux",
target_os = "ios",
target_os = "visionos",
target_os = "macos",
target_os = "tvos",
target_os = "watchos",
target_os = "illumos",
target_os = "solaris",
)))]
mod stub {
use super::*;
pub fn install(_name: &str) -> Result<(), String> {
Err("binding to a network interface is not supported on this platform".to_string())
}
pub fn requested() -> Option<&'static str> {
None
}
pub fn tcp_connect_timeout(addr: SocketAddr, timeout: Duration) -> io::Result<TcpStream> {
TcpStream::connect_timeout(&addr, timeout)
}
pub fn udp_connect(target: SocketAddr) -> io::Result<UdpSocket> {
let local = if target.is_ipv4() {
"0.0.0.0:0"
} else {
"[::]:0"
};
let socket = UdpSocket::bind(local)?;
socket.connect(target)?;
Ok(socket)
}
}