use std::{
net::SocketAddr,
time::{Duration, Instant, SystemTime},
};
use boring::sha::Sha256;
use crossbeam_utils::sync::Unparker;
use quiche::{ConnectionId, Header, RecvInfo, SendInfo};
use crate::poll::{Error, Group, Result, utils::random_conn_id};
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],
))
}
}
pub enum Handshake {
Handshake(usize),
Accept(quiche::Connection),
}
pub struct Acceptor {
config: quiche::Config,
address_validator: Box<dyn AddressValidator + Send>,
}
impl Acceptor {
pub fn new<A: AddressValidator + Send + 'static>(
config: quiche::Config,
address_validator: A,
) -> Self {
Self {
config,
address_validator: Box::new(address_validator),
}
}
pub fn handshake(
&mut self,
header: &Header<'_>,
buf: &mut [u8],
read_size: usize,
recv_info: RecvInfo,
) -> Result<Handshake> {
if !quiche::version_is_supported(header.version) {
return self.negotiate_version(header, buf, read_size, recv_info);
}
let token = header.token.as_ref().unwrap();
if token.is_empty() {
return self.retry(header, buf, read_size, recv_info);
}
let odcid = match self.address_validator.validate_address(
&header.scid,
&header.dcid,
&recv_info.from,
token,
) {
Some(odcid) => odcid,
None => {
log::error!(
"failed to validate address, from={:?}, to={}, scid={:?}, dcid={:?}",
recv_info.from,
recv_info.to,
header.scid,
header.dcid
);
return Err(Error::ValidateAddress);
}
};
let quiche_conn = match quiche::accept(
&header.dcid,
Some(&odcid),
recv_info.to,
recv_info.from,
&mut self.config,
) {
Ok(conn) => {
log::trace!(
"QuicServer(initial) accept new conn, from={:?}, to={}, scid={:?}, dcid={:?}, odcid={:?}",
recv_info.from,
recv_info.to,
header.scid,
header.dcid,
odcid
);
conn
}
Err(err) => {
log::error!(
"failed to accept connection, from={:?}, to={}, scid={:?}, dcid={:?}, err={}",
recv_info.from,
recv_info.to,
header.scid,
header.dcid,
err
);
return Err(Error::Quiche(err));
}
};
Ok(Handshake::Accept(quiche_conn))
}
fn retry(
&self,
header: &Header<'_>,
buf: &mut [u8],
_recv_size: usize,
recv_info: RecvInfo,
) -> Result<Handshake> {
let new_scid = random_conn_id();
log::trace!(
"retry, from={:?}, to={}, scid={:?}, dcid={:?}, new_scid={:?}",
recv_info.from,
recv_info.to,
header.scid,
header.dcid,
new_scid
);
let token = self.address_validator.mint_retry_token(
&header.scid,
&header.dcid,
&new_scid,
&recv_info.from,
)?;
let send_size = match quiche::retry(
&header.scid,
&header.dcid,
&new_scid,
&token,
header.version,
buf,
) {
Ok(send_size) => send_size,
Err(err) => {
log::error!(
"failed to generate retry packet, from={:?}, to={}, scid={:?}, dcid={:?}, err={}",
recv_info.from,
recv_info.to,
header.scid,
header.dcid,
err
);
return Err(Error::Quiche(err));
}
};
Ok(Handshake::Handshake(send_size))
}
fn negotiate_version(
&self,
header: &Header<'_>,
buf: &mut [u8],
_recv_size: usize,
recv_info: RecvInfo,
) -> Result<Handshake> {
log::trace!(
"negotiate_version, from={:?}, to={}, scid={:?}, dcid={:?}",
recv_info.from,
recv_info.to,
header.scid,
header.dcid
);
let send_size = match quiche::negotiate_version(&header.scid, &header.dcid, buf) {
Ok(send_size) => send_size,
Err(err) => {
log::error!(
"failed to generate negotiation_version packet, from={:?}, to={}, scid={:?}, dcid={:?}, err={}",
recv_info.from,
recv_info.to,
header.scid,
header.dcid,
err
);
return Err(Error::Quiche(err));
}
};
Ok(Handshake::Handshake(send_size))
}
}
pub trait ServerGroup {
fn server_dispatch(
&self,
acceptor: &mut Acceptor,
buf: &mut [u8],
recv_size: usize,
recv_info: RecvInfo,
unparker: Option<&Unparker>,
) -> Result<(usize, SendInfo)>;
}
impl ServerGroup for Group {
fn server_dispatch(
&self,
acceptor: &mut Acceptor,
buf: &mut [u8],
recv_size: usize,
recv_info: RecvInfo,
unparker: Option<&Unparker>,
) -> Result<(usize, SendInfo)> {
let header = quiche::Header::from_slice(&mut buf[..recv_size], quiche::MAX_CONN_ID_LEN)
.map_err(Error::Quiche)?;
match self.recv_(&header.dcid, &mut buf[..recv_size], recv_info, unparker) {
Ok((token, _)) => match self.send(token, buf) {
Err(Error::Busy) | Err(Error::Retry) => Ok((
0,
SendInfo {
at: Instant::now(),
from: recv_info.to,
to: recv_info.from,
},
)),
r => r,
},
Err(Error::NotFound) => match acceptor.handshake(&header, buf, recv_size, recv_info) {
Ok(Handshake::Accept(conn)) => {
let token = self.register(conn)?;
match self.recv_(&header.dcid, &mut buf[..recv_size], recv_info, None) {
Ok(_) => {}
Err(Error::Busy) | Err(Error::Retry) => {
unreachable!("Newly registered connections should be idle");
}
Err(err) => return Err(err),
}
match self.send(token, buf) {
Err(Error::Busy) | Err(Error::Retry) => Ok((
0,
SendInfo {
at: Instant::now(),
from: recv_info.to,
to: recv_info.from,
},
)),
r => r,
}
}
Ok(Handshake::Handshake(send_size)) => Ok((
send_size,
SendInfo {
at: Instant::now(),
from: recv_info.to,
to: recv_info.from,
},
)),
Err(err) => Err(err),
},
Err(err) => Err(err),
}
}
}
#[cfg(test)]
mod tests {
use std::{net::SocketAddr, thread::sleep, time::Duration};
use crate::poll::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
);
}
}