use std::{
net::SocketAddr,
time::{Duration, SystemTime},
};
use boring::sha::Sha256;
use quiche::ConnectionId;
use crate::Result;
pub trait AddressValidator {
fn mint_retry_token(
&self,
scid: &ConnectionId<'_>,
dcid: &ConnectionId<'_>,
new_scid: &ConnectionId<'_>,
src: &SocketAddr,
) -> Result<Vec<u8>>;
fn validate_address<'a>(
&self,
scid: &ConnectionId<'_>,
dcid: &ConnectionId<'_>,
src: &SocketAddr,
token: &'a [u8],
) -> Option<ConnectionId<'a>>;
}
pub struct SimpleAddressValidator([u8; 20], Duration);
impl SimpleAddressValidator {
pub fn new(expiration_interval: Duration) -> Self {
let mut seed = [0; 20];
boring::rand::rand_bytes(&mut seed).unwrap();
Self(seed, expiration_interval)
}
}
impl AddressValidator for SimpleAddressValidator {
fn mint_retry_token(
&self,
_scid: &ConnectionId<'_>,
dcid: &ConnectionId<'_>,
new_scid: &ConnectionId<'_>,
src: &SocketAddr,
) -> Result<Vec<u8>> {
let mut token = vec![];
match src.ip() {
std::net::IpAddr::V4(ipv4_addr) => token.extend_from_slice(&ipv4_addr.octets()),
std::net::IpAddr::V6(ipv6_addr) => token.extend_from_slice(&ipv6_addr.octets()),
};
let timestamp = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap()
.as_secs();
token.extend_from_slice(×tamp.to_be_bytes());
token.extend_from_slice(dcid);
let mut hasher = Sha256::new();
hasher.update(&self.0);
hasher.update(&token);
hasher.update(&new_scid);
token.extend_from_slice(&hasher.finish());
Ok(token)
}
fn validate_address<'a>(
&self,
_: &ConnectionId<'_>,
dcid: &ConnectionId<'_>,
src: &SocketAddr,
token: &'a [u8],
) -> Option<ConnectionId<'a>> {
let addr = match src.ip() {
std::net::IpAddr::V4(a) => a.octets().to_vec(),
std::net::IpAddr::V6(a) => a.octets().to_vec(),
};
if addr.len() + 40 > token.len() {
return None;
}
if addr != &token[..addr.len()] {
return None;
}
let timestamp = Duration::from_secs(u64::from_be_bytes(
token[addr.len()..addr.len() + 8].try_into().unwrap(),
));
let now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap();
if now - timestamp > self.1 {
return None;
}
let sha256 = &token[token.len() - 32..];
let mut hasher = Sha256::new();
hasher.update(&self.0);
hasher.update(&token[..token.len() - 32]);
hasher.update(&dcid);
if sha256 != hasher.finish() {
return None;
}
Some(ConnectionId::from_ref(
&token[addr.len() + 8..token.len() - 32],
))
}
}
#[cfg(test)]
mod tests {
use std::{net::SocketAddr, thread::sleep, time::Duration};
use crate::utils::random_conn_id;
use super::*;
#[test]
fn test_default_address_validator() {
let _validator = SimpleAddressValidator::new(Duration::from_secs(100));
let scid = random_conn_id();
let dcid = random_conn_id();
let new_scid = random_conn_id();
let src: SocketAddr = "127.0.0.1:1234".parse().unwrap();
let token = _validator
.mint_retry_token(&scid, &dcid, &new_scid, &src)
.unwrap();
assert_eq!(
_validator.validate_address(&scid, &new_scid, &src, &token),
Some(dcid.clone())
);
assert_eq!(
_validator.validate_address(&scid, &dcid, &src, &token),
None
);
assert_eq!(
_validator.validate_address(&scid, &new_scid, &src, &token),
Some(dcid.clone())
);
let src: SocketAddr = "0.0.0.0:1234".parse().unwrap();
assert_eq!(
_validator.validate_address(&scid, &new_scid, &src, &token),
None
);
let _validator = SimpleAddressValidator::new(Duration::from_secs(1));
let token = _validator
.mint_retry_token(&scid, &dcid, &new_scid, &src)
.unwrap();
assert_eq!(
_validator.validate_address(&scid, &new_scid, &src, &token),
Some(dcid.clone())
);
sleep(Duration::from_secs(2));
assert_eq!(
_validator.validate_address(&scid, &new_scid, &src, &token),
None
);
}
}