use socket2::{Domain, Protocol, SockAddr, Socket, Type};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, TcpListener, TcpStream};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc;
use std::time::{Duration, Instant};
const ACCEPT_POLL: Duration = Duration::from_millis(20);
const ATTEMPT: Duration = Duration::from_secs(2);
pub(crate) const LINGER: Duration = Duration::from_millis(300);
pub(crate) struct Punch {
port: u16,
listeners: Vec<TcpListener>,
}
impl Punch {
pub(crate) 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(crate) fn port(&self) -> u16 {
self.port
}
pub(crate) fn accept(&self) -> Vec<TcpStream> {
let accepted = self.listeners.iter().filter_map(|l| l.accept().ok());
accepted
.map(|(stream, _)| {
let _ = stream.set_nonblocking(false);
stream
})
.collect()
}
#[cfg(test)]
pub(crate) fn punch(&self, targets: Vec<SocketAddr>, window: Duration) -> Vec<TcpStream> {
self.toward(targets, window, &|| self.accept())
}
pub(crate) fn toward(
&self,
targets: Vec<SocketAddr>,
window: Duration,
landed: &dyn Fn() -> Vec<TcpStream>,
) -> Vec<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 streams = Vec::new();
let mut linger_until = None;
while started.elapsed() < window && linger_until.is_none_or(|at| Instant::now() < at) {
streams.extend(landed());
streams.extend(rx.try_iter());
if !streams.is_empty() && linger_until.is_none() {
linger_until = Some(Instant::now() + LINGER);
}
std::thread::sleep(ACCEPT_POLL);
}
done.store(true, Ordering::Relaxed);
streams
}
}
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()));
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(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(())
}
mod local;
pub(crate) use local::local_ips;
#[cfg(test)]
mod tests;