use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, TcpListener, TcpStream, UdpSocket};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc;
use std::time::{Duration, Instant};
use socket2::{Domain, Protocol, SockAddr, Socket, Type};
const ACCEPT_POLL: Duration = Duration::from_millis(20);
const ATTEMPT: Duration = Duration::from_secs(2);
pub struct Punch {
port: u16,
listeners: Vec<TcpListener>,
}
impl Punch {
pub fn bind(port: u16) -> Result<Punch, String> {
let v4 = listener(SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), port))?;
let port = v4.local_addr().map_err(|e| e.to_string())?.port();
let mut listeners = vec![v4];
if let Ok(v6) = listener(SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), port)) {
listeners.push(v6);
}
Ok(Punch { port, listeners })
}
pub fn port(&self) -> u16 {
self.port
}
pub fn punch(&self, targets: Vec<SocketAddr>, window: Duration) -> Option<TcpStream> {
let (tx, rx) = mpsc::channel();
let done = Arc::new(AtomicBool::new(false));
for target in ordered(targets) {
let (tx, done, port) = (tx.clone(), Arc::clone(&done), self.port);
std::thread::spawn(move || connect(port, target, window, &done, &tx));
}
drop(tx);
let started = Instant::now();
let mut landed = None;
while landed.is_none() && started.elapsed() < window {
landed = self
.listeners
.iter()
.find_map(|listener| listener.accept().ok().map(|(stream, _)| stream))
.or_else(|| rx.try_recv().ok());
if landed.is_none() {
std::thread::sleep(ACCEPT_POLL);
}
}
done.store(true, Ordering::Relaxed);
landed.and_then(|stream| stream.set_nonblocking(false).ok().map(|()| stream))
}
}
pub(crate) fn ordered(targets: Vec<SocketAddr>) -> Vec<SocketAddr> {
let (v6, v4): (Vec<_>, Vec<_>) = targets.into_iter().partition(SocketAddr::is_ipv6);
v6.into_iter().chain(v4).collect()
}
fn connect(
port: u16,
target: SocketAddr,
window: Duration,
done: &AtomicBool,
tx: &mpsc::Sender<TcpStream>,
) {
let started = Instant::now();
while !done.load(Ordering::Relaxed) && started.elapsed() < window {
let attempt = ATTEMPT.min(window.saturating_sub(started.elapsed()));
let sent = Instant::now();
if let Ok(socket) = reusable(target.is_ipv6(), port)
&& socket
.connect_timeout(&SockAddr::from(target), attempt)
.is_ok()
{
let _ = tx.send(TcpStream::from(socket));
return;
}
std::thread::sleep(attempt.saturating_sub(sent.elapsed()).max(ACCEPT_POLL));
}
}
fn listener(at: SocketAddr) -> Result<TcpListener, String> {
let socket = reusable(at.is_ipv6(), at.port()).map_err(|e| format!("punch {at}: {e}"))?;
socket.listen(8).map_err(|e| format!("punch {at}: {e}"))?;
socket
.set_nonblocking(true)
.map_err(|e| format!("punch {at}: {e}"))?;
Ok(TcpListener::from(socket))
}
fn reusable(v6: bool, port: u16) -> std::io::Result<Socket> {
let domain = if v6 { Domain::IPV6 } else { Domain::IPV4 };
let socket = Socket::new(domain, Type::STREAM, Some(Protocol::TCP))?;
socket.set_reuse_address(true)?;
reuse_port(&socket)?;
let ip = if v6 {
IpAddr::V6(Ipv6Addr::UNSPECIFIED)
} else {
IpAddr::V4(Ipv4Addr::UNSPECIFIED)
};
if v6 {
socket.set_only_v6(true)?;
}
socket.bind(&SockAddr::from(SocketAddr::new(ip, port)))?;
Ok(socket)
}
#[cfg(unix)]
fn reuse_port(socket: &Socket) -> std::io::Result<()> {
socket.set_reuse_port(true)
}
#[cfg(not(unix))]
fn reuse_port(_socket: &Socket) -> std::io::Result<()> {
Ok(())
}
pub(crate) fn local_ips() -> Vec<IpAddr> {
["[2001:db8::1]:53", "192.0.2.1:53"]
.iter()
.filter_map(|probe| {
let bind = if probe.starts_with('[') {
"[::]:0"
} else {
"0.0.0.0:0"
};
let socket = UdpSocket::bind(bind).ok()?;
socket.connect(probe).ok()?;
let ip = socket.local_addr().ok()?.ip();
(!ip.is_loopback()).then_some(ip)
})
.collect()
}
#[cfg(test)]
mod tests;