use std::{
convert::TryFrom,
fmt::Display,
hash::{Hash, Hasher},
net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6},
path::Path,
};
use base64::Engine;
use bytes::BufMut;
use dhttp_identity::certificate::{CertificateChainKey, CertificateSequence};
use dquic::qbase::net::addr::EndpointAddr as DquicEndpointAddr;
use nom::{
IResult, Parser,
bytes::streaming::take,
combinator::{flat_map, map},
error::{ErrorKind, make_error},
number::streaming::{be_u8, be_u16, be_u32, be_u128},
};
use rustls::{SignatureScheme, pki_types::SubjectPublicKeyInfoDer};
use snafu::{ResultExt, Snafu};
use crate::core::parser::{
sigin,
varint::{VarInt, WriteVarInt, be_varint},
};
#[derive(Debug, Snafu)]
#[snafu(module)]
pub enum SignEndpointError {
#[snafu(display("failed to sign endpoint address"))]
Sign {
source: dhttp_identity::identity::SignError,
},
#[snafu(display("no supported signature scheme for endpoint address"))]
NoSupportedScheme,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct EndpointSignature {
scheme: u16,
signature: Vec<u8>,
}
#[derive(Debug, Clone)]
pub struct EndpointAddr {
flags: u8,
sequence: Option<CertificateSequence>,
load: Option<f32>,
signature: Option<EndpointSignature>,
pub primary: SocketAddr,
pub agent: Option<SocketAddr>,
}
impl PartialEq for EndpointAddr {
fn eq(&self, other: &Self) -> bool {
self.flags == other.flags
&& self.sequence == other.sequence
&& self.load.map(f32::to_bits) == other.load.map(f32::to_bits)
&& self.signature == other.signature
&& self.primary == other.primary
&& self.agent == other.agent
}
}
impl Eq for EndpointAddr {}
impl Hash for EndpointAddr {
fn hash<H: Hasher>(&self, state: &mut H) {
self.flags.hash(state);
self.sequence.hash(state);
self.load.map(f32::to_bits).hash(state);
self.signature.hash(state);
self.primary.hash(state);
self.agent.hash(state);
}
}
impl EndpointAddr {
const FLAG_FAMILY: u8 = 0b1000_0000; const FLAG_MAIN: u8 = 0b0100_0000; const FLAG_CLUSTERED: u8 = 0b0010_0000; const FLAG_NAT: u8 = 0b0001_0000; const FLAG_LOAD: u8 = 0b0000_1000; const FLAG_SIGNED: u8 = 0b0000_0001;
pub fn direct_v4(addr: SocketAddrV4) -> Self {
Self {
flags: 0, sequence: None,
load: None,
signature: None,
primary: addr.into(),
agent: None,
}
}
pub fn direct_v6(addr: SocketAddrV6) -> Self {
Self {
flags: Self::FLAG_FAMILY, sequence: None,
load: None,
signature: None,
primary: addr.into(),
agent: None,
}
}
pub fn nat_v4(outer: SocketAddrV4, agent: SocketAddrV4) -> Self {
Self {
flags: Self::FLAG_NAT, sequence: None,
load: None,
signature: None,
primary: outer.into(),
agent: Some(agent.into()),
}
}
pub fn nat_v6(outer: SocketAddrV6, agent: SocketAddrV6) -> Self {
Self {
flags: Self::FLAG_FAMILY | Self::FLAG_NAT, sequence: None,
load: None,
signature: None,
primary: outer.into(),
agent: Some(agent.into()),
}
}
pub fn is_ipv6(&self) -> bool {
self.flags & Self::FLAG_FAMILY != 0
}
pub fn is_nat(&self) -> bool {
self.flags & Self::FLAG_NAT != 0
}
pub fn is_clustered(&self) -> bool {
self.flags & Self::FLAG_CLUSTERED != 0
}
pub fn is_load(&self) -> bool {
self.flags & Self::FLAG_LOAD != 0
}
pub fn set_clustered(&mut self, clustered: bool) {
if clustered {
self.flags |= Self::FLAG_CLUSTERED;
} else {
self.flags &= !Self::FLAG_CLUSTERED;
self.sequence = None; }
}
pub fn set_load(&mut self, load: Option<f32>) {
self.load = load;
if self.load.is_some() {
self.flags |= Self::FLAG_LOAD;
} else {
self.flags &= !Self::FLAG_LOAD;
}
}
pub async fn sign_with_authority(
&mut self,
authority: &(impl dhttp_identity::identity::LocalAuthority + ?Sized),
) -> Result<(), SignEndpointError> {
self.set_signed(true);
let data = self.signed_data();
let scheme = authority
.cert_chain()
.first()
.and_then(|_| sigin::canonical_scheme_for_spki(authority.public_key()))
.ok_or(SignEndpointError::NoSupportedScheme)?;
let signature = authority
.sign(&data)
.await
.context(sign_endpoint_error::SignSnafu)?;
self.signature = Some(EndpointSignature {
scheme: u16::from(scheme),
signature,
});
Ok(())
}
pub fn verify_signature(
&self,
spki: SubjectPublicKeyInfoDer<'_>,
) -> Result<bool, sigin::VerifyError> {
let Some(sig) = &self.signature else {
return Ok(false);
};
let data = self.signed_data();
sigin::verify(
spki,
SignatureScheme::from(sig.scheme),
&data,
&sig.signature,
)
}
pub fn verify_signature_from_der(&self, cert_der: &[u8]) -> Result<bool, sigin::VerifyError> {
let (_, cert) = x509_parser::parse_x509_certificate(cert_der).map_err(|e| {
sigin::VerifyError::InvalidCertificate {
details: e.to_string(),
}
})?;
let spki = SubjectPublicKeyInfoDer::from(cert.tbs_certificate.subject_pki.raw);
self.verify_signature(spki)
}
pub fn verify_signature_from_pem(&self, cert_pem: &[u8]) -> Result<bool, sigin::VerifyError> {
let mut reader = std::io::Cursor::new(cert_pem);
if let Some(item) = rustls_pemfile::certs(&mut reader).next() {
let cert_der = item.map_err(|e| sigin::VerifyError::InvalidPem { source: e })?;
return self.verify_signature_from_der(&cert_der);
}
Err(sigin::VerifyError::InvalidCertificate {
details: "No certificate found in PEM".to_string(),
})
}
pub fn verify_signature_from_base64(
&self,
cert_base64: &str,
) -> Result<bool, sigin::VerifyError> {
let cert_base64 = cert_base64.trim();
let cert_der = base64::engine::general_purpose::STANDARD
.decode(cert_base64)
.map_err(|e| sigin::VerifyError::InvalidBase64 { source: e })?;
self.verify_signature_from_der(&cert_der)
}
pub fn verify_signature_from_file(
&self,
path: impl AsRef<Path>,
) -> Result<bool, sigin::VerifyError> {
let contents = std::fs::read(path).map_err(|e| sigin::VerifyError::Io { source: e })?;
if let Ok(res) = self.verify_signature_from_pem(&contents) {
return Ok(res);
}
self.verify_signature_from_der(&contents)
}
pub fn is_main(&self) -> bool {
self.flags() & Self::FLAG_MAIN == Self::FLAG_MAIN
}
pub fn set_main(&mut self, is_main: bool) {
let flags = self.flags_mut();
if is_main {
*flags |= Self::FLAG_MAIN;
} else {
*flags &= !Self::FLAG_MAIN;
}
}
pub fn is_signed(&self) -> bool {
self.flags() & Self::FLAG_SIGNED == Self::FLAG_SIGNED
}
pub fn set_signed(&mut self, is_signed: bool) {
let flags = self.flags_mut();
if is_signed {
*flags |= Self::FLAG_SIGNED;
} else {
*flags &= !Self::FLAG_SIGNED;
}
}
pub fn encpding_size(&self) -> usize {
let mut meta_len = 1;
if let Some(seq) = &self.sequence {
meta_len += VarInt::from_u32(seq.get()).encoding_size();
}
if self.load.is_some() {
meta_len += 4; }
if self.is_signed()
&& let Some(sig) = &self.signature
{
let sig_len =
VarInt::try_from(sig.signature.len() as u64).unwrap_or(VarInt::from_u32(0));
meta_len += 2 + sig_len.encoding_size() + sig.signature.len();
}
let addr_len = match (self.is_ipv6(), self.is_nat()) {
(false, false) => 2 + 4, (false, true) => (2 + 4) * 2, (true, false) => 2 + 16, (true, true) => (2 + 16) * 2, };
meta_len + addr_len
}
pub fn addr(&self) -> SocketAddr {
self.primary
}
pub fn agent_addr(&self) -> Option<SocketAddr> {
self.agent
}
pub fn sequence(&self) -> Option<CertificateSequence> {
self.sequence
}
pub fn normalized_sequence(&self) -> CertificateSequence {
self.sequence
.unwrap_or_else(|| CertificateSequence::from(0u8))
}
pub fn set_sequence(&mut self, sequence: CertificateSequence) {
if sequence.get() > 0 {
self.sequence = Some(sequence);
self.set_clustered(true);
} else {
self.sequence = None;
self.set_clustered(false);
}
}
pub fn certificate_chain_key(&self) -> CertificateChainKey {
if self.is_main() {
crate::core::certificate::primary_chain_key(self.normalized_sequence())
} else {
crate::core::certificate::secondary_chain_key(self.normalized_sequence())
}
}
pub fn load(&self) -> Option<f32> {
self.load
}
fn flags(&self) -> u8 {
self.flags
}
fn flags_mut(&mut self) -> &mut u8 {
&mut self.flags
}
pub fn signature(&self) -> Option<&EndpointSignature> {
self.signature.as_ref()
}
pub fn signature_base64(&self) -> Option<String> {
self.signature
.as_ref()
.map(|sig| base64::engine::general_purpose::STANDARD.encode(&sig.signature))
}
fn write_base<B: BufMut>(&self, buf: &mut B) {
buf.put_u8(self.flags);
if let Some(seq) = &self.sequence {
buf.put_varint(VarInt::from_u32(seq.get()));
}
match self.primary {
SocketAddr::V4(addr) => buf.put_socket_addr_v4(&addr),
SocketAddr::V6(addr) => buf.put_socket_addr_v6(&addr),
}
if let Some(agent_addr) = &self.agent {
match agent_addr {
SocketAddr::V4(addr) => buf.put_socket_addr_v4(addr),
SocketAddr::V6(addr) => buf.put_socket_addr_v6(addr),
}
}
if let Some(load) = self.load {
buf.put_u32(load.to_bits());
}
}
fn signed_data(&self) -> Vec<u8> {
let mut unsigned = self.clone();
unsigned.set_signed(true);
unsigned.signature = None;
let mut buf = bytes::BytesMut::with_capacity(unsigned.encpding_size());
unsigned.write_base(&mut buf);
buf.to_vec()
}
}
pub(crate) trait WriteEndpointAddr {
fn put_endpoint_addr(&mut self, endpoint: &EndpointAddr);
}
impl<B: BufMut> WriteEndpointAddr for B {
fn put_endpoint_addr(&mut self, endpoint: &EndpointAddr) {
endpoint.write_base(self);
if endpoint.is_signed()
&& let Some(sig) = endpoint.signature()
{
self.put_u16(sig.scheme);
let len = VarInt::try_from(sig.signature.len() as u64).unwrap_or(VarInt::from_u32(0));
self.put_varint(len);
self.put_slice(&sig.signature);
}
}
}
pub fn be_endpoint_addr(input: &[u8]) -> nom::IResult<&[u8], EndpointAddr> {
let (remain, flags) = be_u8(input)?;
let is_clustered = flags & EndpointAddr::FLAG_CLUSTERED != 0;
let is_ipv6 = flags & EndpointAddr::FLAG_FAMILY != 0;
let is_nat = flags & EndpointAddr::FLAG_NAT != 0;
let has_load = flags & EndpointAddr::FLAG_LOAD != 0;
let (remain, sequence) = if is_clustered {
let (remain, seq) = be_varint(remain)?;
let sequence = match CertificateSequence::try_from(seq.into_inner()) {
Ok(sequence) => sequence,
Err(_error) => {
return Err(nom::Err::Failure(make_error(remain, ErrorKind::TooLarge)));
}
};
(remain, Some(sequence))
} else {
(remain, None)
};
let (remain, primary) = if is_ipv6 {
let (remain, addr) = be_socket_addr_v6(remain)?;
(remain, SocketAddr::V6(addr))
} else {
let (remain, addr) = be_socket_addr_v4(remain)?;
(remain, SocketAddr::V4(addr))
};
let (remain, agent) = if is_nat {
let agent_addr = if is_ipv6 {
let (remain, addr) = be_socket_addr_v6(remain)?;
(remain, SocketAddr::V6(addr))
} else {
let (remain, addr) = be_socket_addr_v4(remain)?;
(remain, SocketAddr::V4(addr))
};
let (remain, addr) = agent_addr;
(remain, Some(addr))
} else {
(remain, None)
};
let (remain, load) = if has_load {
let (remain, load) = be_u32(remain)?;
(remain, Some(f32::from_bits(load)))
} else {
(remain, None)
};
let (remain, signature) = be_endpoint_signature(remain, flags)?;
Ok((
remain,
EndpointAddr {
flags,
sequence,
load,
signature,
primary,
agent,
},
))
}
pub(crate) fn be_endpoint_addr_compat(
input: &[u8],
rdlen: u16,
) -> nom::IResult<&[u8], EndpointAddr> {
let legacy_lengths = [
6, 12, 18, 36, ];
if legacy_lengths.contains(&(rdlen as usize)) {
if let Ok((remaining, endpoint)) = be_endpoint_addr(input)
&& remaining.is_empty()
{
return Ok((remaining, endpoint));
}
return be_legacy_endpoint_addr_by_length(input, rdlen);
}
be_endpoint_addr(input)
}
fn be_legacy_endpoint_addr_by_length(
input: &[u8],
rdlen: u16,
) -> nom::IResult<&[u8], EndpointAddr> {
match rdlen {
6 => {
let (remain, addr) = be_socket_addr_v4(input)?;
Ok((
remain,
EndpointAddr {
flags: 0,
sequence: None,
load: None,
signature: None,
primary: addr.into(),
agent: None,
},
))
}
12 => {
let (remain, primary) = be_socket_addr_v4(input)?;
let (remain, agent) = be_socket_addr_v4(remain)?;
Ok((
remain,
EndpointAddr {
flags: EndpointAddr::FLAG_NAT,
sequence: None,
load: None,
signature: None,
primary: primary.into(),
agent: Some(agent.into()),
},
))
}
18 => {
let (remain, addr) = be_socket_addr_v6(input)?;
Ok((
remain,
EndpointAddr {
flags: EndpointAddr::FLAG_FAMILY,
sequence: None,
load: None,
signature: None,
primary: addr.into(),
agent: None,
},
))
}
36 => {
let (remain, primary) = be_socket_addr_v6(input)?;
let (remain, agent) = be_socket_addr_v6(remain)?;
Ok((
remain,
EndpointAddr {
flags: EndpointAddr::FLAG_FAMILY | EndpointAddr::FLAG_NAT,
sequence: None,
load: None,
signature: None,
primary: primary.into(),
agent: Some(agent.into()),
},
))
}
_ => Err(nom::Err::Error(nom::error::make_error(
input,
nom::error::ErrorKind::LengthValue,
))),
}
}
fn be_endpoint_signature(input: &[u8], flags: u8) -> IResult<&[u8], Option<EndpointSignature>> {
if (flags & EndpointAddr::FLAG_SIGNED) != EndpointAddr::FLAG_SIGNED {
if !input.is_empty() {
return Err(nom::Err::Error(make_error(input, ErrorKind::Eof)));
}
return Ok((input, None));
}
let (remain, scheme_u16) = be_u16(input)?;
let (remain, sig_len) = be_varint(remain)?;
let sig_len = usize::try_from(sig_len.into_inner())
.map_err(|_| nom::Err::Error(make_error(remain, ErrorKind::TooLarge)))?;
let (remain, sig) = take(sig_len)(remain)?;
Ok((
remain,
Some(EndpointSignature {
scheme: scheme_u16,
signature: sig.to_vec(),
}),
))
}
pub trait WriteSocketAddr {
fn put_socket_addr_v4(&mut self, addr: &SocketAddrV4);
fn put_socket_addr_v6(&mut self, addr: &SocketAddrV6);
fn put_socket_addr(&mut self, addr: &SocketAddr) {
match addr {
SocketAddr::V4(v4) => self.put_socket_addr_v4(v4),
SocketAddr::V6(v6) => self.put_socket_addr_v6(v6),
}
}
}
impl<T: BufMut> WriteSocketAddr for T {
fn put_socket_addr_v4(&mut self, addr: &SocketAddrV4) {
self.put_u16(addr.port());
self.put_u32(u32::from(*addr.ip()));
}
fn put_socket_addr_v6(&mut self, addr: &SocketAddrV6) {
self.put_u16(addr.port());
self.put_u128(u128::from(*addr.ip()));
}
}
pub fn be_socket_addr_v4(input: &[u8]) -> IResult<&[u8], SocketAddrV4> {
flat_map(be_u16, |port| {
map(be_ipv4_addr, move |ip| SocketAddrV4::new(ip, port))
})
.parse(input)
}
pub fn be_socket_addr_v6(input: &[u8]) -> IResult<&[u8], SocketAddrV6> {
flat_map(be_u16, |port| {
map(be_ipv6_addr, move |ip| SocketAddrV6::new(ip, port, 0, 0))
})
.parse(input)
}
pub fn be_ipv4_addr(input: &[u8]) -> IResult<&[u8], Ipv4Addr> {
map(be_u32, Ipv4Addr::from).parse(input)
}
pub fn be_ipv6_addr(input: &[u8]) -> IResult<&[u8], Ipv6Addr> {
map(be_u128, Ipv6Addr::from).parse(input)
}
pub fn be_ip_addr(is_v6: bool) -> impl Fn(&[u8]) -> IResult<&[u8], IpAddr> {
move |input| match is_v6 {
true => map(be_u128, |ip| IpAddr::V6(Ipv6Addr::from(ip))).parse(input),
false => map(be_u32, |ip| IpAddr::V4(Ipv4Addr::from(ip))).parse(input),
}
}
impl Display for EndpointAddr {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
if let Some(agent_addr) = &self.agent {
write!(f, "{}-{agent_addr}", self.primary)
} else {
write!(f, "{}", self.primary)
}
}
}
impl TryFrom<DquicEndpointAddr> for EndpointAddr {
type Error = ();
fn try_from(value: DquicEndpointAddr) -> Result<Self, Self::Error> {
match value {
DquicEndpointAddr::Direct {
addr: SocketAddr::V4(addr),
} => Ok(Self::direct_v4(addr)),
DquicEndpointAddr::Direct {
addr: SocketAddr::V6(addr),
} => Ok(Self::direct_v6(addr)),
DquicEndpointAddr::Mediate {
agent: SocketAddr::V4(agent),
outer: SocketAddr::V4(outer),
} => Ok(Self::nat_v4(outer, agent)),
DquicEndpointAddr::Mediate {
agent: SocketAddr::V6(agent),
outer: SocketAddr::V6(outer),
} => Ok(Self::nat_v6(outer, agent)),
_ => Err(()),
}
}
}
impl TryFrom<EndpointAddr> for DquicEndpointAddr {
type Error = ();
fn try_from(value: EndpointAddr) -> Result<Self, Self::Error> {
if let Some(agent_addr) = value.agent {
match (value.primary, agent_addr) {
(SocketAddr::V4(outer), SocketAddr::V4(agent)) => Ok(DquicEndpointAddr::Mediate {
outer: SocketAddr::V4(outer),
agent: SocketAddr::V4(agent),
}),
(SocketAddr::V6(outer), SocketAddr::V6(agent)) => Ok(DquicEndpointAddr::Mediate {
outer: SocketAddr::V6(outer),
agent: SocketAddr::V6(agent),
}),
_ => Err(()),
}
} else {
match value.primary {
SocketAddr::V4(addr) => Ok(DquicEndpointAddr::Direct {
addr: SocketAddr::V4(addr),
}),
SocketAddr::V6(addr) => Ok(DquicEndpointAddr::Direct {
addr: SocketAddr::V6(addr),
}),
}
}
}
}
pub async fn sign_endponit_address(
server_id: u8,
authority: Option<&(impl dhttp_identity::identity::LocalAuthority + ?Sized)>,
endpoint: DquicEndpointAddr,
) -> Option<EndpointAddr> {
let mut ep: EndpointAddr = endpoint.try_into().ok()?;
ep.set_main(server_id == 0);
ep.set_sequence(CertificateSequence::from(server_id));
if let Some(authority) = authority {
let _ = ep.sign_with_authority(authority).await;
}
Some(ep)
}
#[cfg(test)]
mod tests {
use std::{
net::{Ipv4Addr, Ipv6Addr},
sync::Arc,
};
use bytes::BytesMut;
use futures::future::BoxFuture;
use ring::signature::KeyPair;
use rustls::sign::{Signer, SigningKey};
use super::*;
fn v4_outer() -> SocketAddrV4 {
SocketAddrV4::new(Ipv4Addr::new(203, 0, 113, 10), 4433)
}
#[test]
fn endpoint_certificate_chain_key_normalizes_missing_sequence() {
let mut endpoint = EndpointAddr::direct_v4(v4_outer());
endpoint.set_main(true);
let key = endpoint.certificate_chain_key();
assert_eq!(key.usage().kind_flag(), "0");
assert_eq!(key.sequence().get(), 0);
}
#[test]
fn endpoint_certificate_chain_key_uses_present_sequence() {
let mut endpoint = EndpointAddr::direct_v4(v4_outer());
endpoint.set_main(false);
endpoint.set_sequence(
dhttp_identity::certificate::CertificateSequence::try_from(7u32).unwrap(),
);
let key = endpoint.certificate_chain_key();
assert_eq!(key.usage().kind_flag(), "1");
assert_eq!(key.sequence().get(), 7);
}
#[test]
fn endpoint_parser_rejects_over_range_certificate_sequence() {
let sequence = crate::core::parser::varint::VarInt::from_u64(
dhttp_identity::certificate::CertificateSequence::MAX as u64 + 1,
)
.unwrap();
let mut packet = BytesMut::new();
packet.put_u8(EndpointAddr::FLAG_MAIN | EndpointAddr::FLAG_CLUSTERED);
packet.put_varint(sequence);
packet.put_u16(v4_outer().port());
packet.put_slice(&v4_outer().ip().octets());
assert!(be_endpoint_addr(&packet).is_err());
}
#[test]
fn legacy_endpoint_v4_direct_without_meta() {
let port = 5353u16;
let ip = Ipv4Addr::new(10, 0, 0, 1);
let mut buf = BytesMut::new();
buf.extend_from_slice(&port.to_be_bytes());
buf.extend_from_slice(&u32::from(ip).to_be_bytes());
let (remain, decoded) = be_endpoint_addr_compat(&buf, 6).unwrap();
assert!(remain.is_empty());
assert_eq!(
decoded,
EndpointAddr::direct_v4(SocketAddrV4::new(ip, port))
);
}
#[test]
fn legacy_endpoint_v4_nat_without_meta() {
let outer = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 1000);
let agent = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 2), 2000);
let mut buf = BytesMut::new();
buf.extend_from_slice(&outer.port().to_be_bytes());
buf.extend_from_slice(&u32::from(*outer.ip()).to_be_bytes());
buf.extend_from_slice(&agent.port().to_be_bytes());
buf.extend_from_slice(&u32::from(*agent.ip()).to_be_bytes());
let (remain, decoded) = be_endpoint_addr_compat(&buf, 12).unwrap();
assert!(remain.is_empty());
assert_eq!(decoded, EndpointAddr::nat_v4(outer, agent));
}
#[test]
fn flag_bit_ops_work() {
let addr = SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 1), 5353);
let mut ep = EndpointAddr {
flags: 0b0011_0000,
sequence: None,
load: None,
signature: None,
primary: addr.into(),
agent: None,
};
assert!(!ep.is_main());
assert!(!ep.is_signed());
ep.set_main(true);
assert!(ep.is_main());
assert_eq!(ep.flags, 0b0111_0000);
ep.set_signed(true);
assert!(ep.is_signed());
assert_eq!(ep.flags, 0b0111_0001);
ep.set_main(false);
assert!(!ep.is_main());
assert!(ep.is_signed());
assert_eq!(ep.flags, 0b0011_0001);
ep.set_signed(false);
assert!(!ep.is_signed());
assert_eq!(ep.flags, 0b0011_0000);
}
#[test]
fn varint_roundtrip_and_len() {
fn roundtrip(v: u64) {
let v = VarInt::from_u64(v).unwrap();
let mut buf = BytesMut::new();
buf.put_varint(v);
assert_eq!(buf.len(), v.encoding_size());
let (remain, decoded) = be_varint(&buf).unwrap();
assert!(remain.is_empty());
assert_eq!(decoded, v);
}
for v in [
0u64,
1,
63,
64,
16383,
16384,
(1 << 30) - 1,
1 << 30,
(1 << 62) - 1,
] {
roundtrip(v);
}
}
#[test]
fn varint_rejects_overflow_and_incomplete() {
assert!(VarInt::from_u64((1 << 62) + 1).is_err());
let incomplete = [0b01_000000u8];
match be_varint(&incomplete) {
Err(nom::Err::Incomplete(_)) => {}
other => panic!("expected Incomplete, got {other:?}"),
}
}
#[test]
fn endpoint_encode_decode_roundtrip() {
let v4_outer = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 1000);
let v4_agent = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 2), 2000);
let v6_outer = SocketAddrV6::new(Ipv6Addr::LOCALHOST, 3000, 0, 0);
let v6_agent = SocketAddrV6::new(Ipv6Addr::LOCALHOST, 4000, 0, 0);
let mut with_load = EndpointAddr::direct_v4(v4_outer);
with_load.set_load(Some(0.42_f32));
let cases = vec![
EndpointAddr {
flags: EndpointAddr::FLAG_MAIN | EndpointAddr::FLAG_CLUSTERED,
sequence: Some(CertificateSequence::from(0u8)),
load: None,
signature: None,
primary: v4_outer.into(),
agent: None,
},
EndpointAddr {
flags: EndpointAddr::FLAG_NAT | EndpointAddr::FLAG_CLUSTERED,
sequence: Some(CertificateSequence::try_from(127u32).unwrap()),
load: None,
signature: None,
primary: v4_outer.into(),
agent: Some(v4_agent.into()),
},
EndpointAddr {
flags: EndpointAddr::FLAG_FAMILY
| EndpointAddr::FLAG_MAIN
| EndpointAddr::FLAG_CLUSTERED,
sequence: Some(CertificateSequence::try_from(128u32).unwrap()),
load: None,
signature: None,
primary: v6_outer.into(),
agent: None,
},
EndpointAddr {
flags: EndpointAddr::FLAG_FAMILY
| EndpointAddr::FLAG_NAT
| EndpointAddr::FLAG_CLUSTERED,
sequence: Some(CertificateSequence::try_from(16_384u32).unwrap()),
load: None,
signature: None,
primary: v6_outer.into(),
agent: Some(v6_agent.into()),
},
with_load,
];
for ep in cases {
let mut buf = BytesMut::new();
buf.put_endpoint_addr(&ep);
assert_eq!(buf.len(), ep.encpding_size());
let (remain, decoded) = be_endpoint_addr(&buf).unwrap();
assert!(remain.is_empty());
assert_eq!(decoded, ep);
}
}
#[test]
fn compat_parser_does_not_misclassify_modern_lengths_as_legacy() {
let mut direct = EndpointAddr::direct_v4("203.0.113.10:4433".parse().unwrap());
direct.set_main(true);
direct.set_sequence(CertificateSequence::try_from(10u32).unwrap());
direct.set_load(Some(1.0));
let mut nat = EndpointAddr::nat_v4(
"198.51.100.10:4433".parse().unwrap(),
"192.0.2.10:4433".parse().unwrap(),
);
nat.set_main(true);
nat.set_sequence(CertificateSequence::from(1u8));
nat.set_load(Some(2.0));
for endpoint in [direct, nat] {
let mut buf = BytesMut::new();
buf.put_endpoint_addr(&endpoint);
assert!([12, 18].contains(&buf.len()));
let (remaining, decoded) =
be_endpoint_addr_compat(&buf, u16::try_from(buf.len()).unwrap()).unwrap();
assert!(remaining.is_empty());
assert_eq!(decoded, endpoint);
}
}
#[test]
fn endpoint_signature_roundtrip_and_verify() {
#[derive(Debug)]
struct Ed25519Key {
keypair: Arc<ring::signature::Ed25519KeyPair>,
cert_chain: Vec<rustls::pki_types::CertificateDer<'static>>,
}
#[derive(Debug)]
struct Ed25519Signer(Arc<ring::signature::Ed25519KeyPair>);
impl Signer for Ed25519Signer {
fn sign(&self, message: &[u8]) -> Result<Vec<u8>, rustls::Error> {
Ok(self.0.sign(message).as_ref().to_vec())
}
fn scheme(&self) -> SignatureScheme {
SignatureScheme::ED25519
}
}
impl SigningKey for Ed25519Key {
fn choose_scheme(&self, offered: &[SignatureScheme]) -> Option<Box<dyn Signer>> {
offered
.contains(&SignatureScheme::ED25519)
.then(|| Box::new(Ed25519Signer(self.keypair.clone())) as Box<dyn Signer>)
}
fn algorithm(&self) -> rustls::SignatureAlgorithm {
rustls::SignatureAlgorithm::ED25519
}
}
impl dhttp_identity::identity::LocalAuthority for Ed25519Key {
fn name(&self) -> &str {
"authority.example"
}
fn cert_chain(&self) -> &[rustls::pki_types::CertificateDer<'static>] {
&self.cert_chain
}
fn sign(
&self,
data: &[u8],
) -> BoxFuture<'_, Result<Vec<u8>, dhttp_identity::identity::SignError>> {
let result = dhttp_identity::identity::sign_with_key(self, data);
Box::pin(std::future::ready(result))
}
}
let rng = ring::rand::SystemRandom::new();
let pkcs8 = ring::signature::Ed25519KeyPair::generate_pkcs8(&rng).unwrap();
let keypair =
Arc::new(ring::signature::Ed25519KeyPair::from_pkcs8(pkcs8.as_ref()).unwrap());
let mut spki = Vec::with_capacity(44);
spki.extend_from_slice(&[
0x30, 0x2a, 0x30, 0x05, 0x06, 0x03, 0x2b, 0x65, 0x70, 0x03, 0x21, 0x00,
]);
spki.extend_from_slice(keypair.public_key().as_ref());
let key = Ed25519Key {
keypair: keypair.clone(),
cert_chain: vec![rustls::pki_types::CertificateDer::from(spki.clone())],
};
let addr = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 5353);
let mut ep = EndpointAddr::direct_v4(addr);
ep.set_main(true);
futures::executor::block_on(ep.sign_with_authority(&key)).unwrap();
let mut buf = BytesMut::new();
buf.put_endpoint_addr(&ep);
assert_eq!(buf.len(), ep.encpding_size());
let (remain, decoded) = be_endpoint_addr(&buf).unwrap();
assert!(remain.is_empty());
assert!(decoded.is_signed());
assert!(decoded.signature().is_some());
assert!(
decoded
.verify_signature(SubjectPublicKeyInfoDer::from(spki.as_slice()))
.unwrap()
);
let mut tampered = decoded.clone();
tampered.set_main(false);
assert!(
!tampered
.verify_signature(SubjectPublicKeyInfoDer::from(spki.as_slice()))
.unwrap()
);
}
#[test]
fn sign_with_authority_uses_canonical_scheme_from_public_key() {
#[derive(Debug)]
struct Ed25519Authority {
cert_chain: Vec<rustls::pki_types::CertificateDer<'static>>,
}
impl dhttp_identity::identity::LocalAuthority for Ed25519Authority {
fn name(&self) -> &str {
"authority.example"
}
fn cert_chain(&self) -> &[rustls::pki_types::CertificateDer<'static>] {
&self.cert_chain
}
fn sign(
&self,
_data: &[u8],
) -> BoxFuture<'_, Result<Vec<u8>, dhttp_identity::identity::SignError>> {
Box::pin(async move { Ok(vec![1, 2, 3]) })
}
}
let cert_chain = vec![rustls::pki_types::CertificateDer::from(vec![
0x30, 0x2a, 0x30, 0x05, 0x06, 0x03, 0x2b, 0x65, 0x70, 0x03, 0x21, 0x00, 0, 0, 0, 0, 0,
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
])];
let mut ep = EndpointAddr::direct_v4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5353));
futures::executor::block_on(ep.sign_with_authority(&Ed25519Authority { cert_chain }))
.unwrap();
let signature = ep.signature().unwrap();
assert_eq!(
SignatureScheme::from(signature.scheme),
SignatureScheme::ED25519
);
assert_eq!(signature.signature, vec![1, 2, 3]);
}
#[test]
fn optional_fields_flags_follow_values() {
let addr = SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 1), 5353);
let mut ep = EndpointAddr::direct_v4(addr);
assert!(!ep.is_load());
ep.set_load(Some(0.5_f32));
assert!(ep.is_load());
assert_eq!(ep.load(), Some(0.5_f32));
ep.set_load(None);
assert!(!ep.is_load());
assert_eq!(ep.load(), None);
}
}