use crate::nat_identification::{NatType, IDENTIFY_TIMEOUT};
use crate::udp_traversal::hole_punch_config::HolePunchConfig;
use crate::udp_traversal::hole_punched_socket::HolePunchedUdpSocket;
use crate::udp_traversal::linear::encrypted_config_container::HolePunchConfigContainer;
use crate::udp_traversal::multi::DualStackUdpHolePuncher;
use citadel_io::tokio::net::UdpSocket;
use futures::Future;
use netbeam::reliable_conn::ReliableOrderedStreamToTargetExt;
use netbeam::sync::network_endpoint::NetworkEndpoint;
use netbeam::sync::subscription::Subscribable;
use serde::{Deserialize, Serialize};
use std::net::SocketAddr;
use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::Duration;
#[derive(Serialize, Deserialize)]
struct HolePunchCandidates {
internal_bind_addrs: Vec<SocketAddr>,
reflexive_addrs: Vec<SocketAddr>,
}
pub struct UdpHolePuncher<'a> {
driver: Pin<Box<dyn Future<Output = Result<HolePunchedUdpSocket, anyhow::Error>> + Send + 'a>>,
}
const DEFAULT_TIMEOUT: Duration =
Duration::from_millis((IDENTIFY_TIMEOUT.as_millis() + 17000) as u64);
impl<'a> UdpHolePuncher<'a> {
pub fn new(
conn: &'a NetworkEndpoint,
encrypted_config_container: HolePunchConfigContainer,
) -> Self {
Self::new_timeout(conn, encrypted_config_container, DEFAULT_TIMEOUT)
}
pub fn new_timeout(
conn: &'a NetworkEndpoint,
encrypted_config_container: HolePunchConfigContainer,
timeout: Duration,
) -> Self {
Self {
driver: Box::pin(
async move { driver(conn, encrypted_config_container, timeout).await },
),
}
}
}
impl Future for UdpHolePuncher<'_> {
type Output = Result<HolePunchedUdpSocket, anyhow::Error>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.driver.as_mut().poll(cx)
}
}
const MAX_RETRIES: usize = 3;
#[cfg_attr(
feature = "localhost-testing",
tracing::instrument(level = "trace", target = "citadel", skip_all, ret, err(Debug))
)]
async fn driver(
conn: &NetworkEndpoint,
encrypted_config_container: HolePunchConfigContainer,
timeout: Duration,
) -> Result<HolePunchedUdpSocket, anyhow::Error> {
let mut retries = 0;
loop {
log::trace!(target: "citadel", "[driver] Attempt {}/{} starting (timeout: {:?})", retries + 1, MAX_RETRIES, timeout);
let task = citadel_io::time::timeout(
timeout,
driver_inner(conn, encrypted_config_container.clone()),
);
match task.await {
Ok(Ok(res)) => {
log::trace!(target: "citadel", "[driver] Attempt {} succeeded!", retries + 1);
return Ok(res);
}
Ok(Err(err)) => {
log::warn!(target: "citadel", "[driver] Attempt {}/{} failed with error: {err:?}", retries + 1, MAX_RETRIES);
}
Err(_) => {
log::warn!(target: "citadel", "[driver] Attempt {}/{} timed-out after {:?}", retries + 1, MAX_RETRIES, timeout);
}
}
retries += 1;
if retries >= MAX_RETRIES {
log::error!(target: "citadel", "[driver] All {} attempts exhausted, giving up", MAX_RETRIES);
return Err(anyhow::Error::msg("Max retries reached for UDP Traversal"));
}
log::trace!(target: "citadel", "[driver] Retrying... ({} attempts remaining)", MAX_RETRIES - retries);
}
}
async fn driver_inner(
conn: &NetworkEndpoint,
mut encrypted_config_container: HolePunchConfigContainer,
) -> Result<HolePunchedUdpSocket, anyhow::Error> {
log::trace!(target: "citadel", "[driver] Starting hole puncher ...");
log::trace!(target: "citadel", "[driver] Step 1: Initiating subscription...");
let stream = match conn.initiate_subscription().await {
Ok(s) => {
log::trace!(target: "citadel", "[driver] Step 1: Subscription successful");
s
}
Err(e) => {
log::error!(target: "citadel", "[driver] Step 1 FAILED: Subscription error: {e:?}");
return Err(e);
}
};
let stream = &stream;
let stun_servers = encrypted_config_container.take_stun_servers();
log::trace!(target: "citadel", "[driver] Step 2: Identifying local NAT type...");
let local_nat_type = match NatType::identify(stun_servers.clone()).await {
Ok(nat) => {
log::trace!(target: "citadel", "[driver] Step 2: NAT identification successful: {nat:?}");
nat
}
Err(e) => {
log::error!(target: "citadel", "[driver] Step 2 FAILED: NAT identification error: {e:?}");
return Err(anyhow::Error::msg(e.to_string()));
}
};
let local_nat_type = &local_nat_type;
log::trace!(target: "citadel", "[driver] Step 3: Exchanging NAT types with peer...");
if let Err(e) = stream.send_serialized(local_nat_type).await {
log::error!(target: "citadel", "[driver] Step 3 FAILED: Send NAT type error: {e:?}");
return Err(anyhow::Error::from(e));
}
log::trace!(target: "citadel", "[driver] Step 3a: Sent local NAT type, waiting for peer...");
let peer_nat_type = match stream.recv_serialized::<NatType>().await {
Ok(nat) => {
log::trace!(target: "citadel", "[driver] Step 3: Received peer NAT type: {nat:?}");
nat
}
Err(e) => {
log::error!(target: "citadel", "[driver] Step 3 FAILED: Receive NAT type error: {e:?}");
return Err(e.into());
}
};
let peer_nat_type = &peer_nat_type;
let local_initial_socket = get_optimal_bind_socket(local_nat_type, peer_nat_type)?;
let internal_bind_addr_optimal = local_initial_socket.local_addr()?;
let mut sockets = vec![local_initial_socket];
let mut internal_addresses = vec![internal_bind_addr_optimal];
if internal_bind_addr_optimal.is_ipv6() {
let additional_socket = crate::socket_helpers::get_udp_socket("0.0.0.0:0")?;
internal_addresses.push(additional_socket.local_addr()?);
sockets.push(additional_socket);
}
let local_reflexive_addrs = probe_reflexive_addrs(&sockets, stun_servers.as_deref()).await;
log::info!(target: "citadel", "[driver] Local reflexive (srflx) addrs: {local_reflexive_addrs:?}");
log::trace!(target: "citadel", "[driver] Step 4: Exchanging address candidates...");
let local_candidates = HolePunchCandidates {
internal_bind_addrs: internal_addresses,
reflexive_addrs: local_reflexive_addrs,
};
let peer_candidates = match conn.sync_exchange_payload(local_candidates).await {
Ok(candidates) => {
log::trace!(target: "citadel", "[driver] Step 4: Sync exchange successful, peer internal: {:?}, peer reflexive: {:?}", candidates.internal_bind_addrs, candidates.reflexive_addrs);
candidates
}
Err(e) => {
log::error!(target: "citadel", "[driver] Step 4 FAILED: Sync exchange error: {e:?}");
return Err(e);
}
};
let peer_internal_bind_addrs = peer_candidates.internal_bind_addrs;
log::info!(target: "citadel", "\n~~~~~~~~~~~~\n [driver] Local NAT type: {local_nat_type:?}\n Peer NAT type: {peer_nat_type:?}");
log::info!(target: "citadel", "[driver] Local internal bind addr: {internal_bind_addr_optimal:?}\nPeer internal bind addr: {peer_internal_bind_addrs:?}");
log::info!(target: "citadel", "\n~~~~~~~~~~~~\n");
let hole_punch_config = HolePunchConfig::new(
peer_nat_type,
&peer_internal_bind_addrs,
&peer_candidates.reflexive_addrs,
sockets,
);
let conn = conn.clone();
log::trace!(target: "citadel", "[driver] Step 5: Creating and executing DualStackUdpHolePuncher...");
let hole_puncher = match DualStackUdpHolePuncher::new(
conn.node_type(),
encrypted_config_container,
hole_punch_config,
conn,
) {
Ok(hp) => {
log::trace!(target: "citadel", "[driver] Step 5a: DualStackUdpHolePuncher created successfully");
hp
}
Err(e) => {
log::error!(target: "citadel", "[driver] Step 5 FAILED: DualStackUdpHolePuncher creation error: {e:?}");
return Err(e);
}
};
log::trace!(target: "citadel", "[driver] Step 5b: Awaiting hole punch result...");
let res = hole_puncher.await;
log::info!(target: "citadel", "Hole Punch Status: {res:?}");
res.map_err(|err| {
anyhow::Error::msg(format!(
"**HOLE-PUNCH-ERR**: {err:?} | local_nat_type: {local_nat_type:?} | peer_nat_type: {peer_nat_type:?}",
))
})
}
async fn probe_reflexive_addrs(
sockets: &[UdpSocket],
stun_servers: Option<&[String]>,
) -> Vec<SocketAddr> {
if cfg!(feature = "localhost-testing") {
return Vec::new();
}
let probes = sockets
.iter()
.map(|socket| NatType::get_reflexive_addr(socket, stun_servers));
futures::future::join_all(probes)
.await
.into_iter()
.flatten()
.collect()
}
pub fn get_optimal_bind_socket(
local_nat_info: &NatType,
peer_nat_info: &NatType,
) -> Result<UdpSocket, anyhow::Error> {
let mut local_has_an_external_ipv6_addr = false;
let mut peer_has_an_external_ipv6_addr = false;
if let Some(other_info) = local_nat_info.ip_addr_info() {
if other_info.external_ipv6.is_some() {
local_has_an_external_ipv6_addr = true;
}
}
if let Some(other_info) = peer_nat_info.ip_addr_info() {
if other_info.external_ipv6.is_some() {
peer_has_an_external_ipv6_addr = true;
}
}
let local_allows_ipv6 = local_nat_info.is_ipv6_compatible();
let peer_allows_ipv6 = peer_nat_info.is_ipv6_compatible();
if local_allows_ipv6
&& local_has_an_external_ipv6_addr
&& peer_has_an_external_ipv6_addr
&& peer_allows_ipv6
{
crate::socket_helpers::get_udp_socket("[::]:0")
} else {
crate::socket_helpers::get_udp_socket("0.0.0.0:0")
}
}
pub trait EndpointHolePunchExt {
fn begin_udp_hole_punch(
&self,
encrypted_config_container: HolePunchConfigContainer,
) -> UdpHolePuncher<'_>;
}
impl EndpointHolePunchExt for NetworkEndpoint {
fn begin_udp_hole_punch(
&self,
encrypted_config_container: HolePunchConfigContainer,
) -> UdpHolePuncher<'_> {
UdpHolePuncher::new(self, encrypted_config_container)
}
}
#[cfg(test)]
mod tests {
use crate::udp_traversal::udp_hole_puncher::EndpointHolePunchExt;
use citadel_io::tokio;
use netbeam::sync::test_utils::create_streams_with_addrs_and_lag;
use rstest::rstest;
#[rstest]
#[case(0)]
#[case(50)]
#[case(70)]
#[tokio::test]
async fn test_dual_hole_puncher(#[case] lag: usize) {
citadel_logging::setup_log();
let (server_stream, client_stream) = create_streams_with_addrs_and_lag(lag).await;
let server = async move {
let res = server_stream.begin_udp_hole_punch(Default::default()).await;
log::trace!(target: "citadel", "Server res: {res:?}");
res.unwrap()
};
let client = async move {
let res = client_stream.begin_udp_hole_punch(Default::default()).await;
log::trace!(target: "citadel", "Client res: {res:?}");
res.unwrap()
};
let server = citadel_io::tokio::task::spawn(server);
let client = citadel_io::tokio::task::spawn(client);
let (res0, res1) = citadel_io::tokio::join!(server, client);
log::trace!(target: "citadel", "JOIN complete! {res0:?} | {res1:?}");
let (_res0, _res1) = (res0.unwrap(), res1.unwrap());
#[cfg(not(target_os = "windows"))]
{
let dummy_bytes = b"Hello, world!";
log::trace!(target: "citadel", "A");
_res0
.send_to(dummy_bytes as &[u8], _res0.addr.send_address)
.await
.unwrap();
log::trace!(target: "citadel", "B");
let buf = &mut [0u8; 4096];
let (len, _addr) = _res1.recv_from(buf).await.unwrap();
log::trace!(target: "citadel", "C");
assert_ne!(len, 0);
_res1
.send_to(dummy_bytes, _res1.addr.send_address)
.await
.unwrap();
let (len, _addr) = _res0.recv_from(buf).await.unwrap();
assert_ne!(len, 0);
log::trace!(target: "citadel", "D");
}
}
#[cfg(not(target_os = "windows"))]
#[cfg_attr(coverage, ignore)]
#[tokio::test]
async fn test_dual_hole_puncher_high_lag_consensus() {
use std::time::Duration;
citadel_logging::setup_log();
const LAG_MS: usize = 450;
const ITERS: usize = 5;
const PER_ITER_TIMEOUT: Duration = Duration::from_secs(75);
for i in 0..ITERS {
let (server_stream, client_stream) = create_streams_with_addrs_and_lag(LAG_MS).await;
let iteration = async move {
let server = citadel_io::tokio::task::spawn(async move {
server_stream
.begin_udp_hole_punch(Default::default())
.await
.map_err(|e| e.to_string())
});
let client = citadel_io::tokio::task::spawn(async move {
client_stream
.begin_udp_hole_punch(Default::default())
.await
.map_err(|e| e.to_string())
});
let (res0, res1) = citadel_io::tokio::join!(server, client);
let s0 = res0
.expect("server task panicked")
.expect("server punch err");
let s1 = res1
.expect("client task panicked")
.expect("client punch err");
let dummy = b"Hello, world!";
s0.send_to(dummy as &[u8], s0.addr.send_address)
.await
.unwrap();
let buf = &mut [0u8; 4096];
let (len, _) = s1.recv_from(buf).await.unwrap();
assert_ne!(len, 0);
s1.send_to(dummy as &[u8], s1.addr.send_address)
.await
.unwrap();
let (len, _) = s0.recv_from(buf).await.unwrap();
assert_ne!(len, 0);
};
match citadel_io::tokio::time::timeout(PER_ITER_TIMEOUT, iteration).await {
Ok(()) => {
log::info!(target: "citadel", "[hole-punch-consensus] iter {i}/{ITERS} OK")
}
Err(_) => {
panic!("iter {i}: hole-punch consensus deadlocked at lag {LAG_MS}ms (per-attempt timeout too short?)")
}
}
}
}
}