use crate::buffer::Buf;
use crate::types::ProtocolVersion;
use arrayvec::ArrayVec;
use nom::Err;
use nom::IResult;
use nom::bytes::complete::take;
use nom::error::{Error, ErrorKind};
use nom::number::complete::be_u8;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SupportedVersionsClientHello {
pub versions: ArrayVec<ProtocolVersion, 3>,
}
impl SupportedVersionsClientHello {
pub fn parse(input: &[u8]) -> IResult<&[u8], SupportedVersionsClientHello> {
let (input, list_len) = be_u8(input)?;
if list_len == 0 || list_len % 2 != 0 {
return Err(Err::Failure(Error::new(input, ErrorKind::LengthValue)));
}
let (input, versions_data) = take(list_len)(input)?;
if !input.is_empty() {
return Err(Err::Failure(Error::new(input, ErrorKind::LengthValue)));
}
let mut versions = ArrayVec::new();
let mut rest = versions_data;
while !rest.is_empty() {
let (r, version) = ProtocolVersion::parse(rest)?;
if !matches!(version, ProtocolVersion::Unknown(_)) {
versions
.try_push(version)
.map_err(|_| Err::Failure(Error::new(rest, ErrorKind::LengthValue)))?;
}
rest = r;
}
Ok((input, SupportedVersionsClientHello { versions }))
}
pub fn serialize(&self, output: &mut Buf) {
output.push((self.versions.len() * 2) as u8);
for version in &self.versions {
output.extend_from_slice(&version.as_u16().to_be_bytes());
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SupportedVersionsServerHello {
pub selected_version: ProtocolVersion,
}
impl SupportedVersionsServerHello {
pub fn parse(input: &[u8]) -> IResult<&[u8], SupportedVersionsServerHello> {
let (input, selected_version) = ProtocolVersion::parse(input)?;
if !input.is_empty() {
return Err(Err::Failure(Error::new(input, ErrorKind::LengthValue)));
}
Ok((input, SupportedVersionsServerHello { selected_version }))
}
pub fn serialize(&self, output: &mut Buf) {
output.extend_from_slice(&self.selected_version.as_u16().to_be_bytes());
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::buffer::Buf;
#[test]
fn client_hello_roundtrip() {
let message: &[u8] = &[
0x04, 0xFE, 0xFC, 0xFE, 0xFD, ];
let (rest, parsed) = SupportedVersionsClientHello::parse(message).unwrap();
assert!(rest.is_empty());
let mut serialized = Buf::new();
parsed.serialize(&mut serialized);
assert_eq!(&*serialized, message);
}
#[test]
fn server_hello_roundtrip() {
let message: &[u8] = &[
0xFE, 0xFC, ];
let (rest, parsed) = SupportedVersionsServerHello::parse(message).unwrap();
assert!(rest.is_empty());
let mut serialized = Buf::new();
parsed.serialize(&mut serialized);
assert_eq!(&*serialized, message);
}
#[test]
fn unknown_client_hello_versions_are_ignored() {
let message: &[u8] = &[
0x08, 0xFE, 0xFC, 0xFE, 0xFD, 0xFE, 0xFF, 0xFE, 0xFE, ];
let (rest, parsed) = SupportedVersionsClientHello::parse(message).unwrap();
assert!(rest.is_empty());
assert_eq!(
parsed.versions.as_slice(),
&[
ProtocolVersion::DTLS1_3,
ProtocolVersion::DTLS1_2,
ProtocolVersion::DTLS1_0,
]
);
}
#[test]
fn too_many_duplicate_client_hello_versions_are_rejected() {
let message: &[u8] = &[
0x08, 0xFE, 0xFC, 0xFE, 0xFC, 0xFE, 0xFC, 0xFE, 0xFC, ];
let result = SupportedVersionsClientHello::parse(message);
assert!(
matches!(
result,
Err(nom::Err::Failure(error))
if error.code == nom::error::ErrorKind::LengthValue
),
"too many supported versions should fail with LengthValue"
);
}
#[test]
fn odd_client_hello_versions_are_rejected() {
let result = SupportedVersionsClientHello::parse(&[
0x03, 0xFE, 0xFC, 0xFF,
]);
assert!(
matches!(
result,
Err(nom::Err::Failure(error))
if error.code == nom::error::ErrorKind::LengthValue
),
"odd supported_versions vector should fail with LengthValue"
);
}
#[test]
fn trailing_client_hello_versions_bytes_are_rejected() {
let result = SupportedVersionsClientHello::parse(&[
0x02, 0xFE, 0xFC, 0xFF, ]);
assert!(
matches!(
result,
Err(nom::Err::Failure(error))
if error.code == nom::error::ErrorKind::LengthValue
),
"trailing supported_versions bytes should fail with LengthValue"
);
}
#[test]
fn trailing_server_hello_version_bytes_are_rejected() {
let result = SupportedVersionsServerHello::parse(&[
0xFE, 0xFC, 0xFF, ]);
assert!(
matches!(
result,
Err(nom::Err::Failure(error))
if error.code == nom::error::ErrorKind::LengthValue
),
"trailing selected_version bytes should fail with LengthValue"
);
}
}