use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::time::Duration;
use dig_ip::LocalStack;
use dig_nat::stun::{discover_reflexive_address, BINDING_SUCCESS, MAGIC_COOKIE};
use tokio::net::UdpSocket;
fn build_xor_response(addr: SocketAddr, txid: &[u8; 12]) -> Vec<u8> {
let cookie_be = MAGIC_COOKIE.to_be_bytes();
let mut value = vec![0u8]; let port = addr.port() ^ ((MAGIC_COOKIE >> 16) as u16);
match addr.ip() {
IpAddr::V4(v4) => {
value.push(0x01); value.extend_from_slice(&port.to_be_bytes());
let mut octets = v4.octets();
for (i, o) in octets.iter_mut().enumerate() {
*o ^= cookie_be[i];
}
value.extend_from_slice(&octets);
}
IpAddr::V6(v6) => {
value.push(0x02); value.extend_from_slice(&port.to_be_bytes());
let mut octets = v6.octets();
let mut key = [0u8; 16];
key[..4].copy_from_slice(&cookie_be);
key[4..].copy_from_slice(txid);
for (o, k) in octets.iter_mut().zip(key.iter()) {
*o ^= *k;
}
value.extend_from_slice(&octets);
}
}
let mut attr = 0x0020u16.to_be_bytes().to_vec(); attr.extend_from_slice(&(value.len() as u16).to_be_bytes());
attr.extend_from_slice(&value);
while attr.len() % 4 != 0 {
attr.push(0);
}
let mut msg = BINDING_SUCCESS.to_be_bytes().to_vec();
msg.extend_from_slice(&(attr.len() as u16).to_be_bytes());
msg.extend_from_slice(&cookie_be);
msg.extend_from_slice(txid);
msg.extend_from_slice(&attr);
msg
}
async fn spawn_responder(bind_ip: IpAddr, reflexive: SocketAddr) -> SocketAddr {
let socket = UdpSocket::bind(SocketAddr::new(bind_ip, 0))
.await
.expect("bind loopback STUN responder");
let addr = socket.local_addr().expect("responder local addr");
tokio::spawn(async move {
let mut buf = [0u8; 512];
loop {
let Ok((n, from)) = socket.recv_from(&mut buf).await else {
return;
};
if n < 20 {
continue;
}
let txid: [u8; 12] = buf[8..20].try_into().unwrap();
let resp = build_xor_response(reflexive, &txid);
let _ = socket.send_to(&resp, from).await;
}
});
addr
}
fn dead_v6() -> SocketAddr {
SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 1)
}
const SHORT: Duration = Duration::from_millis(300);
#[tokio::test]
async fn falls_back_to_ipv4_when_ipv6_stun_is_dead() {
let reflexive = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 1234);
let v4 = spawn_responder(IpAddr::V4(Ipv4Addr::LOCALHOST), reflexive).await;
let servers = [dead_v6(), v4];
let got = discover_reflexive_address(&servers, LocalStack::from_flags(true, true), SHORT).await;
assert_eq!(
got,
Some(reflexive),
"must fall back to the reachable IPv4 STUN, not null out"
);
}
#[tokio::test]
async fn ipv4_only_host_uses_ipv4_stun() {
let reflexive = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)), 9000);
let v4 = spawn_responder(IpAddr::V4(Ipv4Addr::LOCALHOST), reflexive).await;
let servers = [dead_v6(), v4];
let got =
discover_reflexive_address(&servers, LocalStack::from_flags(false, true), SHORT).await;
assert_eq!(got, Some(reflexive));
}
#[tokio::test]
async fn prefers_ipv6_when_it_answers() {
let v6_reflexive = SocketAddr::new(
IpAddr::V6(Ipv6Addr::new(0x2606, 0x4700, 0x4700, 0, 0, 0, 0, 0x1111)),
4321,
);
let v4_reflexive = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(9, 9, 9, 9)), 5555);
let v6 = spawn_responder(IpAddr::V6(Ipv6Addr::LOCALHOST), v6_reflexive).await;
let v4 = spawn_responder(IpAddr::V4(Ipv4Addr::LOCALHOST), v4_reflexive).await;
let servers = [v6, v4];
let got = discover_reflexive_address(&servers, LocalStack::from_flags(true, true), SHORT).await;
assert_eq!(
got,
Some(v6_reflexive),
"IPv6 must win when its STUN server answers"
);
}
#[tokio::test]
async fn empty_input_returns_none() {
let got = discover_reflexive_address(&[], LocalStack::from_flags(true, true), SHORT).await;
assert_eq!(got, None);
}
#[tokio::test]
async fn all_dead_returns_none() {
let servers = [
dead_v6(),
SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 1),
];
let got = discover_reflexive_address(&servers, LocalStack::from_flags(true, true), SHORT).await;
assert_eq!(got, None);
}