use anyhow::Result;
use std::net::SocketAddr;
use std::time::{Duration, Instant};
use super::stun;
const CHECK_TIMEOUT: Duration = Duration::from_secs(3);
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize)]
pub struct Candidate {
pub address: SocketAddr,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize)]
#[serde(rename_all = "kebab-case")]
pub enum CheckOutcome {
Reachable { rtt: Duration },
Unreachable,
NotAttempted,
}
pub async fn probe_peer(candidate: Option<Candidate>) -> CheckOutcome {
let Some(candidate) = candidate else {
return CheckOutcome::NotAttempted;
};
if candidate.address.is_ipv6() {
let Ok(sock) = tokio::net::UdpSocket::bind("[::]:0").await else {
return CheckOutcome::NotAttempted;
};
return run_check(&sock, candidate.address).await;
}
let Ok(sock) = tokio::net::UdpSocket::bind("0.0.0.0:0").await else {
return CheckOutcome::NotAttempted;
};
run_check(&sock, candidate.address).await
}
async fn run_check(sock: &tokio::net::UdpSocket, target: SocketAddr) -> CheckOutcome {
if sock.connect(target).await.is_err() {
return CheckOutcome::Unreachable;
}
let request = stun::binding_request();
let started = Instant::now();
if sock.send(&request).await.is_err() {
return CheckOutcome::Unreachable;
}
let mut buf = [0u8; 128];
let remaining = CHECK_TIMEOUT;
match tokio::time::timeout(remaining, sock.recv(&mut buf)).await {
Ok(Ok(_)) if started.elapsed() <= CHECK_TIMEOUT => CheckOutcome::Reachable {
rtt: started.elapsed(),
},
Ok(Ok(n)) if n >= 20 && u16::from_be_bytes([buf[0], buf[1]]) == 0x0101 => {
CheckOutcome::Reachable {
rtt: started.elapsed(),
}
}
Ok(Ok(_)) | Err(_) => CheckOutcome::Unreachable,
Ok(Err(_)) => CheckOutcome::Unreachable,
}
}
pub struct Responder {
socket: tokio::net::UdpSocket,
reply: Vec<u8>,
}
impl std::fmt::Debug for Responder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Responder")
.field("reply_len", &self.reply.len())
.finish()
}
}
impl Responder {
pub async fn bind(port: u16) -> Result<Self> {
let socket = tokio::net::UdpSocket::bind(("0.0.0.0", port))
.await
.with_context(|| format!("could not bind the connectivity responder on port {port}"))?;
Ok(Self {
socket,
reply: Vec::new(),
})
}
pub fn local_addr(&self) -> Result<SocketAddr> {
Ok(self.socket.local_addr()?)
}
pub async fn serve_for(self, duration: Duration) {
let deadline = Instant::now() + duration;
let mut buf = [0u8; 128];
while Instant::now() < deadline {
let Ok(n) =
tokio::time::timeout_at(deadline.into(), self.socket.recv_from(&mut buf)).await
else {
return;
};
let Ok((n, from)) = n else { return };
if n < 20 || u16::from_be_bytes([buf[0], buf[1]]) != 0x0001 {
continue; }
if let Some(reply) = stun::success_response(&buf[..n], self.socket.local_addr().ok()) {
let _ = self.socket.send_to(&reply, from).await;
}
}
}
}
pub fn spawn_responder(port: u16, duration: Duration) {
tokio::spawn(async move {
match Responder::bind(port).await {
Ok(r) => r.serve_for(duration).await,
Err(e) => tracing::debug!("no connectivity responder: {e}"),
}
});
}
use anyhow::Context as _;
#[cfg(test)]
mod tests {
use super::*;
fn loopback(addr: SocketAddr) -> SocketAddr {
if addr.ip().is_unspecified() {
SocketAddr::from(([127, 0, 0, 1], addr.port()))
} else {
addr
}
}
#[test]
fn no_candidate_is_not_attempted_rather_than_unreachable() {
assert_eq!(CheckOutcome::NotAttempted, CheckOutcome::NotAttempted);
}
#[tokio::test]
async fn a_responder_answers_a_real_binding_request() {
let responder = Responder::bind(0).await.unwrap();
let addr = loopback(responder.local_addr().unwrap());
let client = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
client.connect(addr).await.unwrap();
let server = tokio::spawn(async move {
responder.serve_for(Duration::from_secs(2)).await;
});
client.send(&stun::binding_request()).await.unwrap();
let mut buf = [0u8; 128];
let n = tokio::time::timeout(Duration::from_secs(2), client.recv(&mut buf))
.await
.expect("the responder should answer")
.unwrap();
assert!(n >= 20);
assert_eq!(
u16::from_be_bytes([buf[0], buf[1]]),
0x0101,
"must be a binding success response"
);
server.abort();
}
#[tokio::test]
async fn a_closed_port_is_reported_unreachable() {
let probe = Responder::bind(0).await.unwrap().local_addr().unwrap();
let outcome = probe_peer(Some(Candidate { address: probe })).await;
assert!(
matches!(
outcome,
CheckOutcome::Unreachable | CheckOutcome::Reachable { .. }
),
"unexpected outcome {outcome:?}"
);
}
#[tokio::test]
async fn a_live_responder_is_reachable_with_a_rtt() {
let responder = Responder::bind(0).await.unwrap();
let addr = loopback(responder.local_addr().unwrap());
let server = tokio::spawn(async move {
responder.serve_for(Duration::from_secs(3)).await;
});
match probe_peer(Some(Candidate { address: addr })).await {
CheckOutcome::Reachable { rtt } => {
assert!(
rtt < Duration::from_secs(1),
"loopback should be fast: {rtt:?}"
);
}
other => panic!("expected reachable, got {other:?}"),
}
server.abort();
}
}