use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::time::Duration;
use tokio::net::UdpSocket;
pub const MAGIC_COOKIE: u32 = 0x2112_A442;
pub const BINDING_REQUEST: u16 = 0x0001;
pub const BINDING_SUCCESS: u16 = 0x0101;
pub const ATTR_XOR_MAPPED_ADDRESS: u16 = 0x0020;
pub const ATTR_MAPPED_ADDRESS: u16 = 0x0001;
const FAMILY_IPV4: u8 = 0x01;
const FAMILY_IPV6: u8 = 0x02;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum StunError {
#[error("STUN message truncated")]
Truncated,
#[error("bad STUN magic cookie")]
BadMagicCookie,
#[error("STUN transaction id mismatch")]
TransactionIdMismatch,
#[error("no usable mapped address in STUN response")]
NoMappedAddress,
#[error("unexpected STUN message type: {0:#06x}")]
UnexpectedType(u16),
#[error("STUN io: {0}")]
Io(String),
#[error("STUN request timed out")]
Timeout,
}
pub fn encode_binding_request(transaction_id: &[u8; 12]) -> Vec<u8> {
let mut msg = Vec::with_capacity(20);
msg.extend_from_slice(&BINDING_REQUEST.to_be_bytes()); msg.extend_from_slice(&0u16.to_be_bytes()); msg.extend_from_slice(&MAGIC_COOKIE.to_be_bytes()); msg.extend_from_slice(transaction_id); msg
}
fn is_usable_reflexive_addr(addr: &SocketAddr) -> bool {
if addr.port() == 0 {
return false;
}
match addr.ip() {
IpAddr::V4(v4) => is_usable_reflexive_v4(v4),
IpAddr::V6(v6) => match v6.to_ipv4() {
Some(v4) => is_usable_reflexive_v4(v4),
None => is_usable_reflexive_v6(v6),
},
}
}
fn is_usable_reflexive_v4(v4: Ipv4Addr) -> bool {
let [a, b, _, _] = v4.octets();
let is_this_network = a == 0;
let is_6to4_relay_anycast = v4.octets()[..3] == [192, 88, 99];
let is_benchmarking = a == 198 && (b & 0xfe) == 18;
let is_reserved_class_e = a >= 240;
!v4.is_unspecified()
&& !v4.is_loopback()
&& !v4.is_link_local()
&& !v4.is_multicast()
&& !v4.is_broadcast()
&& !v4.is_documentation()
&& !is_this_network
&& !is_6to4_relay_anycast
&& !is_benchmarking
&& !is_reserved_class_e
}
fn is_usable_reflexive_v6(v6: Ipv6Addr) -> bool {
let seg = v6.segments();
let is_link_local = (seg[0] & 0xffc0) == 0xfe80;
let is_documentation = seg[0] == 0x2001 && seg[1] == 0x0db8;
!v6.is_unspecified()
&& !v6.is_loopback()
&& !v6.is_multicast()
&& !is_link_local
&& !is_documentation
}
pub fn parse_binding_response(
msg: &[u8],
expected_txid: Option<&[u8; 12]>,
) -> Result<SocketAddr, StunError> {
if msg.len() < 20 {
return Err(StunError::Truncated);
}
let msg_type = u16::from_be_bytes([msg[0], msg[1]]);
let msg_len = u16::from_be_bytes([msg[2], msg[3]]) as usize;
let cookie = u32::from_be_bytes([msg[4], msg[5], msg[6], msg[7]]);
if cookie != MAGIC_COOKIE {
return Err(StunError::BadMagicCookie);
}
if msg_type != BINDING_SUCCESS {
return Err(StunError::UnexpectedType(msg_type));
}
let txid: [u8; 12] = msg[8..20].try_into().map_err(|_| StunError::Truncated)?;
if let Some(expected) = expected_txid {
if &txid != expected {
return Err(StunError::TransactionIdMismatch);
}
}
if msg.len() < 20 + msg_len {
return Err(StunError::Truncated);
}
let mut fallback: Option<SocketAddr> = None;
let mut off = 20usize;
let end = 20 + msg_len;
while off + 4 <= end {
let attr_type = u16::from_be_bytes([msg[off], msg[off + 1]]);
let attr_len = u16::from_be_bytes([msg[off + 2], msg[off + 3]]) as usize;
let val_start = off + 4;
let val_end = val_start + attr_len;
if val_end > end {
return Err(StunError::Truncated);
}
let value = &msg[val_start..val_end];
match attr_type {
ATTR_XOR_MAPPED_ADDRESS => {
return decode_mapped_address(value, &txid, true);
}
ATTR_MAPPED_ADDRESS if fallback.is_none() => {
fallback = decode_mapped_address(value, &txid, false).ok();
}
_ => {}
}
off = val_end + ((4 - (attr_len % 4)) % 4);
}
fallback.ok_or(StunError::NoMappedAddress)
}
fn decode_mapped_address(
value: &[u8],
txid: &[u8; 12],
xor: bool,
) -> Result<SocketAddr, StunError> {
if value.len() < 4 {
return Err(StunError::Truncated);
}
let family = value[1];
let raw_port = u16::from_be_bytes([value[2], value[3]]);
let cookie_be = MAGIC_COOKIE.to_be_bytes();
let port = if xor {
raw_port ^ ((MAGIC_COOKIE >> 16) as u16)
} else {
raw_port
};
match family {
FAMILY_IPV4 => {
if value.len() < 8 {
return Err(StunError::Truncated);
}
let mut octets = [value[4], value[5], value[6], value[7]];
if xor {
for (i, o) in octets.iter_mut().enumerate() {
*o ^= cookie_be[i];
}
}
Ok(SocketAddr::new(IpAddr::V4(Ipv4Addr::from(octets)), port))
}
FAMILY_IPV6 => {
if value.len() < 20 {
return Err(StunError::Truncated);
}
let mut octets = [0u8; 16];
octets.copy_from_slice(&value[4..20]);
if xor {
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;
}
}
Ok(SocketAddr::new(IpAddr::V6(Ipv6Addr::from(octets)), port))
}
other => Err(StunError::UnexpectedType(other as u16)),
}
}
pub async fn query_reflexive_address(
socket: &UdpSocket,
server: SocketAddr,
timeout: Duration,
) -> Result<SocketAddr, StunError> {
let txid = new_transaction_id();
let req = encode_binding_request(&txid);
socket
.send_to(&req, server)
.await
.map_err(|e| StunError::Io(e.to_string()))?;
let mut buf = [0u8; 512];
let deadline = tokio::time::Instant::now() + timeout;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
return Err(StunError::Timeout);
}
let (n, from) = match tokio::time::timeout(remaining, socket.recv_from(&mut buf)).await {
Ok(Ok(x)) => x,
Ok(Err(e)) => return Err(StunError::Io(e.to_string())),
Err(_) => return Err(StunError::Timeout),
};
if from != server {
continue;
}
let addr = parse_binding_response(&buf[..n], Some(&txid))?;
if !is_usable_reflexive_addr(&addr) {
return Err(StunError::NoMappedAddress);
}
return Ok(addr);
}
}
pub async fn discover_reflexive_address(
stun_servers: &[SocketAddr],
local: dig_ip::LocalStack,
timeout: Duration,
) -> Option<SocketAddr> {
if stun_servers.is_empty() {
return None;
}
let mut candidates = dig_ip::PeerCandidates::new();
candidates.extend(
stun_servers.iter().copied(),
dig_ip::CandidateSource::StunReflexive,
);
let config = dig_ip::DialConfig {
per_attempt_timeout: timeout,
..Default::default()
};
let winner = dig_ip::connect(&local, &candidates, config, |stun_addr| async move {
let bind: SocketAddr = if stun_addr.is_ipv6() {
(Ipv6Addr::UNSPECIFIED, 0).into()
} else {
(Ipv4Addr::UNSPECIFIED, 0).into()
};
let socket = UdpSocket::bind(bind)
.await
.map_err(|e| format!("bind {bind}: {e}"))?;
query_reflexive_address(&socket, stun_addr, timeout)
.await
.map_err(|e| e.to_string())
})
.await;
match winner {
Ok(w) => Some(w.conn),
Err(_) => None,
}
}
pub fn new_transaction_id() -> [u8; 12] {
use ring::rand::{SecureRandom, SystemRandom};
let mut id = [0u8; 12];
SystemRandom::new()
.fill(&mut id)
.expect("OS CSPRNG must be available to generate a STUN transaction id");
id
}
#[cfg(test)]
mod reflexive_guard_tests {
use super::is_usable_reflexive_addr;
use std::net::SocketAddr;
fn addr(s: &str) -> SocketAddr {
s.parse().expect("valid SocketAddr literal")
}
#[test]
fn accepts_genuinely_global_addresses() {
assert!(is_usable_reflexive_addr(&addr("1.1.1.1:443")));
assert!(is_usable_reflexive_addr(&addr("8.8.8.8:53")));
assert!(is_usable_reflexive_addr(&addr(
"[2606:4700:4700::1111]:443"
)));
}
#[test]
fn accepts_private_cgnat_and_ula() {
assert!(is_usable_reflexive_addr(&addr("192.168.1.5:9000")));
assert!(is_usable_reflexive_addr(&addr("10.0.0.7:9000")));
assert!(is_usable_reflexive_addr(&addr("172.16.5.5:9000")));
assert!(is_usable_reflexive_addr(&addr("100.64.0.1:9000"))); assert!(is_usable_reflexive_addr(&addr("[fd00::1]:9000"))); }
#[test]
fn rejects_port_zero() {
assert!(!is_usable_reflexive_addr(&addr("1.1.1.1:0")));
assert!(!is_usable_reflexive_addr(&addr("[2606:4700:4700::1111]:0")));
}
#[test]
fn rejects_reserved_ipv4() {
assert!(!is_usable_reflexive_addr(&addr("0.0.0.0:1234"))); assert!(!is_usable_reflexive_addr(&addr("127.0.0.1:1234"))); assert!(!is_usable_reflexive_addr(&addr("169.254.1.1:1234"))); assert!(!is_usable_reflexive_addr(&addr("224.0.0.1:1234"))); assert!(!is_usable_reflexive_addr(&addr("255.255.255.255:1234"))); assert!(!is_usable_reflexive_addr(&addr("192.0.2.1:1234"))); assert!(!is_usable_reflexive_addr(&addr("198.51.100.1:1234"))); assert!(!is_usable_reflexive_addr(&addr("203.0.113.1:1234"))); }
#[test]
fn rejects_reserved_ipv6() {
assert!(!is_usable_reflexive_addr(&addr("[::]:1234"))); assert!(!is_usable_reflexive_addr(&addr("[::1]:1234"))); assert!(!is_usable_reflexive_addr(&addr("[fe80::1]:1234"))); assert!(!is_usable_reflexive_addr(&addr("[febf::1]:1234"))); assert!(!is_usable_reflexive_addr(&addr("[ff02::1]:1234"))); assert!(!is_usable_reflexive_addr(&addr("[2001:db8::1]:1234"))); }
#[test]
fn rejects_ipv4_mapped_and_compat_smuggling_reserved_ranges() {
assert!(!is_usable_reflexive_addr(&addr("[::ffff:127.0.0.1]:1234"))); assert!(!is_usable_reflexive_addr(&addr(
"[::ffff:169.254.1.1]:1234"
))); assert!(!is_usable_reflexive_addr(&addr("[::ffff:224.0.0.1]:1234"))); assert!(!is_usable_reflexive_addr(&addr("[::ffff:192.0.2.1]:1234"))); assert!(!is_usable_reflexive_addr(&addr(
"[::ffff:255.255.255.255]:1234"
))); assert!(!is_usable_reflexive_addr(&addr("[::ffff:0.0.0.0]:1234"))); assert!(!is_usable_reflexive_addr(&addr("[::7f00:1]:1234"))); }
#[test]
fn accepts_ipv4_mapped_private() {
assert!(is_usable_reflexive_addr(&addr("[::ffff:10.0.0.1]:9000")));
}
#[test]
fn rejects_never_dialable_ipv4_ranges() {
assert!(!is_usable_reflexive_addr(&addr("198.18.0.1:1234"))); assert!(!is_usable_reflexive_addr(&addr("198.19.0.1:1234"))); assert!(!is_usable_reflexive_addr(&addr("240.0.0.1:1234"))); assert!(!is_usable_reflexive_addr(&addr("0.1.2.3:1234"))); assert!(!is_usable_reflexive_addr(&addr("192.88.99.1:1234"))); }
}