use std::net::{IpAddr, SocketAddr};
#[derive(Debug, Clone, Eq, PartialEq, Hash)]
pub struct AddrBand {
pub(crate) necessary_ip: IpAddr,
pub(crate) anticipated_ports: Vec<u16>,
}
impl Iterator for AddrBand {
type Item = SocketAddr;
fn next(&mut self) -> Option<Self::Item> {
self.anticipated_ports
.pop()
.map(|port| SocketAddr::new(self.necessary_ip, port))
}
}
#[cfg(not(target_family = "wasm"))]
mod native {
use super::*;
use crate::nat_identification::NatType;
use citadel_io::tokio::net::UdpSocket;
use itertools::Itertools;
use std::net::{Ipv4Addr, Ipv6Addr};
#[derive(Debug)]
pub struct HolePunchConfig {
pub bands: Vec<Vec<AddrBand>>,
pub(crate) locally_bound_sockets: Option<Vec<UdpSocket>>,
}
impl IntoIterator for HolePunchConfig {
type Item = Vec<SocketAddr>;
type IntoIter = std::vec::IntoIter<Vec<SocketAddr>>;
fn into_iter(mut self) -> Self::IntoIter {
let mut ret = vec![];
for band_set in self.bands.drain(..) {
let mut this_set = vec![];
for mut band in band_set {
for next in band.by_ref() {
this_set.push(next);
}
}
ret.push(this_set);
}
ret.into_iter()
}
}
impl HolePunchConfig {
pub fn new(
peer_nat: &NatType,
peer_internal_addrs: &[SocketAddr],
peer_reflexive_addrs: &[SocketAddr],
local_sockets: Vec<UdpSocket>,
) -> Self {
assert_eq!(peer_internal_addrs.len(), local_sockets.len());
let mut this = HolePunchConfig {
bands: Vec::new(),
locally_bound_sockets: Some(local_sockets),
};
for peer_internal_addr in peer_internal_addrs {
let mut bands = if let Some(bands) = peer_nat.predict(peer_internal_addr) {
bands
} else if cfg!(feature = "localhost-testing") {
log::info!(target: "citadel", "Will revert to localhost testing mode (not recommended for production use (peer addr: {peer_internal_addr:?}))");
get_localhost_bands(peer_internal_addr)
} else {
vec![AddrBand {
necessary_ip: peer_internal_addr.ip(),
anticipated_ports: vec![peer_internal_addr.port()],
}]
};
bands.extend(peer_reflexive_addrs.iter().map(|addr| AddrBand {
necessary_ip: addr.ip(),
anticipated_ports: vec![addr.port()],
}));
bands.extend(get_localhost_bands(peer_internal_addr));
let bands = bands.into_iter().unique().collect();
this.bands.push(bands)
}
this
}
}
fn get_localhost_bands(peer_internal_addr: &SocketAddr) -> Vec<AddrBand> {
vec![
AddrBand {
necessary_ip: IpAddr::from(Ipv4Addr::LOCALHOST),
anticipated_ports: vec![peer_internal_addr.port()],
},
AddrBand {
necessary_ip: IpAddr::V6(Ipv6Addr::LOCALHOST),
anticipated_ports: vec![peer_internal_addr.port()],
},
]
}
}
#[cfg(not(target_family = "wasm"))]
pub use native::HolePunchConfig;
#[cfg(all(test, not(target_family = "wasm")))]
mod tests {
use super::HolePunchConfig;
use crate::nat_identification::NatType;
use citadel_io::tokio;
use std::net::SocketAddr;
use std::str::FromStr;
#[tokio::test]
async fn reflexive_addr_becomes_candidate() {
let socket = citadel_io::tokio::net::UdpSocket::bind("127.0.0.1:0")
.await
.unwrap();
let peer_internal = SocketAddr::from_str("192.168.1.50:40000").unwrap();
let reflexive = SocketAddr::from_str("203.0.113.7:55555").unwrap();
let nat = NatType::default();
let config = HolePunchConfig::new(&nat, &[peer_internal], &[reflexive], vec![socket]);
let candidates: Vec<SocketAddr> = config.into_iter().flatten().collect();
assert!(
candidates.contains(&reflexive),
"observed reflexive candidate must be pinged, got: {candidates:?}"
);
}
}