use super::extension::ExtensionVec;
use super::extensions::use_srtp::{SrtpProfileId, UseSrtpExtension};
use super::{CompressionMethod, Dtls12CipherSuite, Extension, ExtensionType};
use super::{ProtocolVersion, Random, SessionId};
use crate::buffer::Buf;
use arrayvec::ArrayVec;
use nom::Err;
use nom::error::{Error, ErrorKind};
use nom::{IResult, bytes::complete::take, number::complete::be_u16};
#[derive(Debug, PartialEq, Eq)]
pub struct ServerHello {
pub server_version: ProtocolVersion,
pub random: Random,
pub session_id: SessionId,
pub cipher_suite: Dtls12CipherSuite,
pub compression_method: CompressionMethod,
pub extensions: Option<ExtensionVec>,
}
impl ServerHello {
pub fn new(
server_version: ProtocolVersion,
random: Random,
session_id: SessionId,
cipher_suite: Dtls12CipherSuite,
compression_method: CompressionMethod,
extensions: Option<ExtensionVec>,
) -> Self {
ServerHello {
server_version,
random,
session_id,
cipher_suite,
compression_method,
extensions,
}
}
pub fn with_extensions(mut self, buf: &mut Buf, srtp_profile: Option<SrtpProfileId>) -> Self {
buf.clear();
let mut ranges: ArrayVec<
(ExtensionType, usize, usize),
{ ExtensionType::supported().len() },
> = ArrayVec::new();
if let Some(pid) = srtp_profile {
let start = buf.len();
let mut profiles = ArrayVec::new();
profiles.push(pid);
let ext = UseSrtpExtension::new(profiles, ArrayVec::new());
ext.serialize(buf);
ranges.push((ExtensionType::UseSrtp, start, buf.len()));
}
let start = buf.len();
ranges.push((ExtensionType::ExtendedMasterSecret, start, start));
let start = buf.len();
buf.push(0); ranges.push((ExtensionType::RenegotiationInfo, start, buf.len()));
let mut extensions = ExtensionVec::new();
for (t, s, e) in ranges {
extensions.push(Extension {
extension_type: t,
extension_data_range: s..e,
});
}
self.extensions = Some(extensions);
self
}
pub fn parse(input: &[u8], base_offset: usize) -> IResult<&[u8], ServerHello> {
let original_input = input;
let (input, server_version) = ProtocolVersion::parse(input)?;
let (input, random) = Random::parse(input)?;
let (input, session_id) = SessionId::parse(input)?;
let (input, cipher_suite) = Dtls12CipherSuite::parse(input)?;
let (input, compression_method) = CompressionMethod::parse(input)?;
let (input, extensions) = if !input.is_empty() {
if input.len() < 2 {
return Err(Err::Failure(Error::new(input, ErrorKind::Eof)));
}
let (input, extensions_len) = be_u16(input)?;
if input.len() < extensions_len as usize {
return Err(Err::Failure(Error::new(input, ErrorKind::Eof)));
}
if extensions_len > 0 {
let (rest, input_ext) = take(extensions_len)(input)?;
if !rest.is_empty() {
return Err(Err::Failure(Error::new(rest, ErrorKind::LengthValue)));
}
let consumed_to_ext_data =
input_ext.as_ptr() as usize - original_input.as_ptr() as usize;
let ext_base_offset = base_offset + consumed_to_ext_data;
let mut extensions_vec = ExtensionVec::new();
let mut current_input = input_ext;
let mut current_offset = ext_base_offset;
while !current_input.is_empty() {
let before_len = current_input.len();
let (new_rest, ext) = Extension::parse(current_input, current_offset)?;
let parsed_len = before_len - new_rest.len();
current_offset += parsed_len;
if ext.extension_type.is_supported() {
if extensions_vec
.iter()
.any(|existing| existing.extension_type == ext.extension_type)
{
return Err(Err::Failure(Error::new(
current_input,
ErrorKind::LengthValue,
)));
}
extensions_vec.try_push(ext).map_err(|_| {
Err::Failure(Error::new(current_input, ErrorKind::LengthValue))
})?;
}
current_input = new_rest;
}
(rest, Some(extensions_vec))
} else {
(input, None)
}
} else {
(input, None)
};
Ok((
input,
ServerHello {
server_version,
random,
session_id,
cipher_suite,
compression_method,
extensions,
},
))
}
pub fn serialize(&self, source_buf: &[u8], output: &mut Buf) {
output.extend_from_slice(&self.server_version.as_u16().to_be_bytes());
self.random.serialize(output);
output.push(self.session_id.len() as u8);
output.extend_from_slice(&self.session_id);
output.extend_from_slice(&self.cipher_suite.as_u16().to_be_bytes());
output.push(self.compression_method.as_u8());
if let Some(extensions) = &self.extensions {
let mut extensions_len = 0;
for ext in extensions.iter() {
let ext_data = ext.extension_data(source_buf);
extensions_len += 2 + 2 + ext_data.len();
}
output.extend_from_slice(&(extensions_len as u16).to_be_bytes());
for ext in extensions {
ext.serialize(source_buf, output);
}
}
}
}
#[cfg(test)]
mod test {
use super::*;
use crate::buffer::Buf;
const MESSAGE: &[u8] = &[
0xFE, 0xFD, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0A, 0x0B, 0x0C, 0x0D, 0x0E, 0x0F,
0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1A, 0x1B, 0x1C, 0x1D, 0x1E,
0x1F, 0x20, 0x01, 0xAA, 0xC0, 0x2B, 0x00, 0x00, 0x0C, 0x00, 0x0A, 0x00, 0x08, 0x00, 0x06, 0x00, 0x17, 0x00, 0x18, 0x00, 0x19, ];
#[test]
fn roundtrip() {
let (rest, parsed) = ServerHello::parse(MESSAGE, 0).unwrap();
assert!(rest.is_empty());
let mut serialized = Buf::new();
parsed.serialize(MESSAGE, &mut serialized);
assert_eq!(&*serialized, MESSAGE);
}
#[test]
fn session_id_too_long() {
let mut message = MESSAGE.to_vec();
message[34] = 0x21;
let result = ServerHello::parse(&message, 0);
assert!(result.is_err());
}
#[test]
fn duplicate_supported_extensions_are_rejected() {
for count in [2, 9] {
let mut message = MESSAGE[..39].to_vec();
message.extend_from_slice(&(count as u16 * 4).to_be_bytes());
for _ in 0..count {
message
.extend_from_slice(&ExtensionType::ExtendedMasterSecret.as_u16().to_be_bytes());
message.extend_from_slice(&0u16.to_be_bytes());
}
let result = ServerHello::parse(&message, 0);
assert!(
matches!(
result,
Err(nom::Err::Failure(error))
if error.code == nom::error::ErrorKind::LengthValue
),
"duplicate supported extensions should fail with LengthValue"
);
}
}
}