use std::net::SocketAddr;
use std::time::Duration;
use tokio::net::UdpSocket;
use crate::codec::{parse_binding_response, StunError, BINDING_SUCCESS};
use crate::credential::request::encode_identity_request;
use crate::credential::response::parse_challenge;
use crate::credential::signature::StunSigner;
use crate::credential::wire::{
is_valid_spki_der, BINDING_ERROR, ERR_STALE_NONCE, ERR_UNAUTHENTICATED, REALM,
};
use crate::scope::{scope_of, Scope};
use crate::transaction_id::new_transaction_id;
use super::request::encode_signed_request;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum SignedQueryError {
#[error("{0}")]
Stun(#[from] StunError),
#[error("server refused with code {code}")]
Refused {
code: u16,
},
#[error("credential exchange cannot proceed")]
BadChallenge,
}
enum Stage {
Initial,
AfterFirstChallenge,
AfterReSign,
}
pub async fn query_reflexive_address_signed(
socket: &UdpSocket,
server: SocketAddr,
timeout: Duration,
signer: &dyn StunSigner,
) -> Result<SocketAddr, SignedQueryError> {
let spki = signer.spki_der();
if !is_valid_spki_der(spki) {
return Err(SignedQueryError::BadChallenge);
}
let deadline = tokio::time::Instant::now() + timeout;
let mut buf = [0u8; 512];
let mut expected_txid = new_transaction_id();
send(
socket,
server,
&encode_identity_request(&expected_txid, spki),
)
.await?;
let mut stage = Stage::Initial;
loop {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
return Err(SignedQueryError::Stun(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(SignedQueryError::Stun(StunError::Io(e.to_string()))),
Err(_) => return Err(SignedQueryError::Stun(StunError::Timeout)),
};
if from != server {
continue; }
let msg = &buf[..n];
if msg.len() < 2 {
return Err(SignedQueryError::Stun(StunError::Truncated));
}
let msg_type = u16::from_be_bytes([msg[0], msg[1]]);
if msg_type == BINDING_SUCCESS {
match parse_binding_response(msg, Some(&expected_txid)) {
Ok(addr) if scope_of(addr) == Scope::NeverDialable => {
return Err(SignedQueryError::Stun(StunError::NoMappedAddress));
}
Ok(addr) => return Ok(addr),
Err(StunError::TransactionIdMismatch) => continue, Err(e) => return Err(SignedQueryError::Stun(e)),
}
} else if msg_type == BINDING_ERROR {
let challenge = match parse_challenge(msg, &expected_txid) {
Ok(c) => c,
Err(StunError::TransactionIdMismatch) => continue, Err(e) => return Err(SignedQueryError::Stun(e)),
};
match stage {
Stage::Initial => {
let has_dig_realm = challenge.realm.as_deref() == Some(REALM);
match (challenge.code, has_dig_realm, &challenge.nonce) {
(ERR_UNAUTHENTICATED, true, Some(nonce)) => {
expected_txid = new_transaction_id();
let req = encode_signed_request(&expected_txid, nonce, signer);
send(socket, server, &req).await?;
stage = Stage::AfterFirstChallenge;
}
_ => {
return Err(SignedQueryError::Refused {
code: challenge.code,
})
}
}
}
Stage::AfterFirstChallenge => match (challenge.code, &challenge.nonce) {
(ERR_STALE_NONCE, Some(nonce)) => {
expected_txid = new_transaction_id();
let req = encode_signed_request(&expected_txid, nonce, signer);
send(socket, server, &req).await?;
stage = Stage::AfterReSign;
}
_ => {
return Err(SignedQueryError::Refused {
code: challenge.code,
})
}
},
Stage::AfterReSign => {
return Err(SignedQueryError::Refused {
code: challenge.code,
});
}
}
} else {
return Err(SignedQueryError::Stun(StunError::UnexpectedType(msg_type)));
}
}
}
async fn send(
socket: &UdpSocket,
server: SocketAddr,
datagram: &[u8],
) -> Result<(), SignedQueryError> {
socket
.send_to(datagram, server)
.await
.map(|_| ())
.map_err(|e| SignedQueryError::Stun(StunError::Io(e.to_string())))
}