use anyhow::{anyhow, Context, Result};
use quinn::Endpoint;
use std::net::{SocketAddr, UdpSocket};
use std::sync::Arc;
use std::time::{Duration, Instant};
use crate::direct;
use crate::net::Transport;
pub fn holepunch_enabled() -> bool {
std::env::var("FILAMENT_HOLEPUNCH").map(|v| v == "1").unwrap_or(false)
}
pub const PUNCH_BUDGET: Duration = Duration::from_secs(3);
const PUNCH_INTERVAL: Duration = Duration::from_millis(75);
const PUNCH_MAGIC: &[u8] = b"FILAMENT-PUNCH-v1";
const STUN_MAGIC_COOKIE: u32 = 0x2112_A442;
pub fn stun_srflx(sock: &UdpSocket, stun_server: SocketAddr) -> Result<SocketAddr> {
let mut txid = [0u8; 12];
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
txid.copy_from_slice(&now.to_be_bytes()[..12]);
let mut req = Vec::with_capacity(20);
req.extend_from_slice(&0x0001u16.to_be_bytes()); req.extend_from_slice(&0x0000u16.to_be_bytes()); req.extend_from_slice(&STUN_MAGIC_COOKIE.to_be_bytes());
req.extend_from_slice(&txid);
let prev_timeout = sock.read_timeout().ok().flatten();
sock.set_read_timeout(Some(Duration::from_millis(700)))
.context("stun: set read timeout")?;
let mut last_err = anyhow!("stun: no response");
for _ in 0..4 {
if let Err(e) = sock.send_to(&req, stun_server) {
last_err = anyhow!("stun: send_to {stun_server}: {e}");
continue;
}
let mut buf = [0u8; 512];
match sock.recv_from(&mut buf) {
Ok((n, _from)) => {
if let Some(addr) = parse_xor_mapped_address(&buf[..n], &txid) {
let _ = sock.set_read_timeout(prev_timeout);
return Ok(addr);
}
last_err = anyhow!("stun: response had no XOR-MAPPED-ADDRESS");
}
Err(e) => last_err = anyhow!("stun: recv_from: {e}"),
}
}
let _ = sock.set_read_timeout(prev_timeout);
Err(last_err)
}
pub fn stun_srflx_any(sock: &UdpSocket, servers: &[SocketAddr]) -> Result<SocketAddr> {
match servers {
[] => return Err(anyhow!("stun: no servers")),
[one] => return stun_srflx(sock, *one),
_ => {}
}
let mut txid = [0u8; 12];
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
txid.copy_from_slice(&now.to_be_bytes()[..12]);
let mut req = Vec::with_capacity(20);
req.extend_from_slice(&0x0001u16.to_be_bytes());
req.extend_from_slice(&0x0000u16.to_be_bytes());
req.extend_from_slice(&STUN_MAGIC_COOKIE.to_be_bytes());
req.extend_from_slice(&txid);
let prev_timeout = sock.read_timeout().ok().flatten();
sock.set_read_timeout(Some(Duration::from_millis(700)))
.context("stun: set read timeout")?;
let mut last_err = anyhow!("stun: no response");
for _ in 0..4 {
for s in servers {
if let Err(e) = sock.send_to(&req, s) {
#[allow(unused_assignments)]
{
last_err = anyhow!("stun: send_to {s}: {e}");
}
}
}
let mut buf = [0u8; 512];
match sock.recv_from(&mut buf) {
Ok((n, _from)) => {
if let Some(addr) = parse_xor_mapped_address(&buf[..n], &txid) {
let _ = sock.set_read_timeout(prev_timeout);
return Ok(addr);
}
last_err = anyhow!("stun: response had no XOR-MAPPED-ADDRESS");
}
Err(e) => last_err = anyhow!("stun: recv_from: {e}"),
}
}
let _ = sock.set_read_timeout(prev_timeout);
Err(last_err)
}
fn parse_xor_mapped_address(buf: &[u8], _txid: &[u8; 12]) -> Option<SocketAddr> {
if buf.len() < 20 {
return None;
}
let mut i = 20usize;
while i + 4 <= buf.len() {
let attr_type = u16::from_be_bytes([buf[i], buf[i + 1]]);
let attr_len = u16::from_be_bytes([buf[i + 2], buf[i + 3]]) as usize;
let val_start = i + 4;
if val_start + attr_len > buf.len() {
break;
}
let val = &buf[val_start..val_start + attr_len];
match attr_type {
0x0020 => {
if val.len() >= 8 && val[1] == 0x01 {
let xport = u16::from_be_bytes([val[2], val[3]]);
let port = xport ^ ((STUN_MAGIC_COOKIE >> 16) as u16);
let cookie = STUN_MAGIC_COOKIE.to_be_bytes();
let ip = std::net::Ipv4Addr::new(
val[4] ^ cookie[0],
val[5] ^ cookie[1],
val[6] ^ cookie[2],
val[7] ^ cookie[3],
);
return Some(SocketAddr::new(ip.into(), port));
}
}
0x0001 => {
if val.len() >= 8 && val[1] == 0x01 {
let port = u16::from_be_bytes([val[2], val[3]]);
let ip = std::net::Ipv4Addr::new(val[4], val[5], val[6], val[7]);
return Some(SocketAddr::new(ip.into(), port));
}
}
_ => {}
}
i = val_start + ((attr_len + 3) & !3);
}
None
}
pub fn stun_server_addr(stun_urls: &[String]) -> Option<SocketAddr> {
if let Ok(v) = std::env::var("FILAMENT_STUN") {
if let Some(a) = resolve_host_port(v.trim()) {
return Some(a);
}
}
for url in stun_urls {
let rest = url.strip_prefix("stun:").or_else(|| url.strip_prefix("stuns:"))?;
let hostport = rest.split('?').next().unwrap_or(rest);
if let Some(a) = resolve_host_port(hostport) {
return Some(a);
}
}
None
}
pub fn stun_server_addrs(stun_urls: &[String]) -> Vec<SocketAddr> {
if let Ok(v) = std::env::var("FILAMENT_STUN") {
if let Some(a) = resolve_host_port(v.trim()) {
return vec![a];
}
}
let mut out: Vec<SocketAddr> = Vec::new();
for url in stun_urls {
let Some(rest) = url.strip_prefix("stun:").or_else(|| url.strip_prefix("stuns:")) else {
continue;
};
let hostport = rest.split('?').next().unwrap_or(rest);
if let Some(a) = resolve_host_port(hostport) {
if !out.contains(&a) {
out.push(a);
}
}
}
out
}
fn resolve_host_port(hostport: &str) -> Option<SocketAddr> {
use std::net::ToSocketAddrs;
let hp = if hostport.contains(':') {
hostport.to_string()
} else {
format!("{hostport}:3478")
};
hp.to_socket_addrs().ok()?.next()
}
pub fn bind_punch_socket() -> Result<UdpSocket> {
UdpSocket::bind("0.0.0.0:0").context("bind punch socket")
}
pub fn punch(sock: &UdpSocket, peer_srflx: SocketAddr) -> Result<()> {
sock.set_read_timeout(Some(Duration::from_millis(50)))
.context("punch: set read timeout")?;
let start = Instant::now();
let mut last_send = Instant::now()
.checked_sub(PUNCH_INTERVAL)
.unwrap_or_else(Instant::now);
let mut sent = 0u32;
let mut recvd = false;
let mut buf = [0u8; 256];
while start.elapsed() < PUNCH_BUDGET {
if last_send.elapsed() >= PUNCH_INTERVAL {
let _ = sock.send_to(PUNCH_MAGIC, peer_srflx);
sent += 1;
last_send = Instant::now();
}
match sock.recv_from(&mut buf) {
Ok((n, from)) => {
if n >= PUNCH_MAGIC.len() && &buf[..PUNCH_MAGIC.len()] == PUNCH_MAGIC {
let _ = from;
recvd = true;
}
}
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock
|| e.kind() == std::io::ErrorKind::TimedOut => {}
Err(_) => {}
}
if sent >= 1 && recvd {
for _ in 0..3 {
let _ = sock.send_to(PUNCH_MAGIC, peer_srflx);
std::thread::sleep(Duration::from_millis(20));
}
return Ok(());
}
}
Err(anyhow!(
"HOLEPUNCH-FAIL: no bidirectional punch in budget (sent={sent}, recvd={recvd}), symmetric NAT?"
))
}
pub fn endpoint_from_socket(sock: UdpSocket) -> Result<Endpoint> {
let mut ep = Endpoint::new(
quinn::EndpointConfig::default(),
Some(direct::server_config()?),
sock,
Arc::new(quinn::TokioRuntime),
)
.context("build quinn endpoint on punched socket")?;
ep.set_default_client_config(direct::client_config()?);
Ok(ep)
}
pub async fn connect(
punch_sock: UdpSocket,
peer_srflx: SocketAddr,
secret: &str,
peer_id: String,
tx: tokio::sync::mpsc::UnboundedSender<crate::net::Ev>,
answerer: bool,
) -> Option<Arc<dyn Transport>> {
let punch_result = tokio::task::spawn_blocking(move || {
let r = punch(&punch_sock, peer_srflx);
(r, punch_sock)
})
.await
.ok()?;
let (r, punch_sock) = punch_result;
if let Err(e) = r {
crate::ui::trace(&format!("filament: {e}"));
return None; }
crate::ui::debug(&format!("filament: HOLEPUNCH ok, NAT open toward {peer_srflx}, starting QUIC"));
let endpoint = match endpoint_from_socket(punch_sock) {
Ok(ep) => ep,
Err(e) => {
crate::ui::trace(&format!("filament: holepunch endpoint build failed: {e}"));
return None;
}
};
direct::race_connect_labeled(
endpoint,
vec![peer_srflx.to_string()],
secret,
peer_id,
tx,
"holepunched",
answerer,
)
.await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn stun_server_addr_parses_url() {
let urls = vec!["stun:198.18.0.1:3478".to_string()];
let a = stun_server_addr(&urls).unwrap();
assert_eq!(a.to_string(), "198.18.0.1:3478");
}
#[test]
fn stun_server_addr_strips_query() {
let urls = vec!["stun:198.18.0.1:3478?transport=udp".to_string()];
let a = stun_server_addr(&urls).unwrap();
assert_eq!(a.port(), 3478);
}
#[test]
fn stun_server_addrs_collects_and_dedups() {
let urls = vec![
"stun:198.18.0.1:3478".to_string(),
"stun:198.18.0.2:3478?transport=udp".to_string(),
"stun:198.18.0.1:3478".to_string(), "turn:198.18.0.9:3478".to_string(), ];
let addrs = stun_server_addrs(&urls);
assert_eq!(addrs.len(), 2, "two distinct stun servers, turn ignored, dup collapsed");
assert!(addrs.iter().any(|a| a.to_string() == "198.18.0.1:3478"));
assert!(addrs.iter().any(|a| a.to_string() == "198.18.0.2:3478"));
}
#[test]
fn xor_mapped_address_roundtrip() {
let ip = std::net::Ipv4Addr::new(203, 0, 113, 5);
let port: u16 = 50000;
let cookie = STUN_MAGIC_COOKIE.to_be_bytes();
let xport = port ^ ((STUN_MAGIC_COOKIE >> 16) as u16);
let octets = ip.octets();
let mut resp = vec![0u8; 20];
resp[0] = 0x01;
resp[1] = 0x01; resp.extend_from_slice(&0x0020u16.to_be_bytes());
resp.extend_from_slice(&0x0008u16.to_be_bytes());
resp.push(0x00);
resp.push(0x01); resp.extend_from_slice(&xport.to_be_bytes());
resp.extend_from_slice(&[
octets[0] ^ cookie[0],
octets[1] ^ cookie[1],
octets[2] ^ cookie[2],
octets[3] ^ cookie[3],
]);
let got = parse_xor_mapped_address(&resp, &[0u8; 12]).unwrap();
assert_eq!(got, SocketAddr::new(ip.into(), port));
}
#[test]
fn stun_and_punch_against_local_reflector() {
let srv = UdpSocket::bind("127.0.0.1:0").unwrap();
let srv_addr = srv.local_addr().unwrap();
let handle = std::thread::spawn(move || {
let mut buf = [0u8; 512];
let (_n, from) = srv.recv_from(&mut buf).unwrap();
let cookie = STUN_MAGIC_COOKIE.to_be_bytes();
let port = from.port() ^ ((STUN_MAGIC_COOKIE >> 16) as u16);
let octets = match from.ip() {
std::net::IpAddr::V4(v4) => v4.octets(),
_ => [127, 0, 0, 1],
};
let mut resp = vec![0u8; 20];
resp[0] = 0x01;
resp[1] = 0x01;
resp[8..20].copy_from_slice(&buf[8..20]);
resp.extend_from_slice(&0x0020u16.to_be_bytes());
resp.extend_from_slice(&0x0008u16.to_be_bytes());
resp.push(0x00);
resp.push(0x01);
resp.extend_from_slice(&port.to_be_bytes());
resp.extend_from_slice(&[
octets[0] ^ cookie[0],
octets[1] ^ cookie[1],
octets[2] ^ cookie[2],
octets[3] ^ cookie[3],
]);
srv.send_to(&resp, from).unwrap();
});
let client = UdpSocket::bind("127.0.0.1:0").unwrap();
let srflx = stun_srflx(&client, srv_addr).unwrap();
assert!(srflx.ip().is_loopback(), "srflx ip {} not loopback", srflx.ip());
assert_eq!(srflx.port(), client.local_addr().unwrap().port());
handle.join().unwrap();
}
}