use std::ops::Deref;
use std::sync::atomic::{AtomicBool, Ordering};
use arrayvec::ArrayVec;
use std::fmt;
use crate::Error;
use crate::buffer::{Buf, TmpBuf};
use crate::dtls13::message::{ContentType, Dtls13CipherSuite, Dtls13Record, Handshake, Sequence};
pub struct Incoming {
records: Box<Records>,
}
impl Incoming {
pub fn records(&self) -> &Records {
&self.records
}
pub fn first(&self) -> &Record {
&self.records()[0]
}
pub fn into_records(self) -> impl Iterator<Item = Record> {
self.records.records.into_iter()
}
}
impl Incoming {
pub fn parse_packet(
packet: &[u8],
decrypt: &mut dyn RecordHandler,
cs: Option<Dtls13CipherSuite>,
) -> Result<Option<Self>, Error> {
let records = Records::parse(packet, decrypt, cs)?;
if records.records.is_empty() {
return Ok(None);
}
let records = Box::new(records);
Ok(Some(Incoming { records }))
}
}
#[derive(Debug)]
pub struct Records {
pub records: ArrayVec<Record, 16>,
}
impl Records {
pub fn parse(
mut packet: &[u8],
decrypt: &mut dyn RecordHandler,
cs: Option<Dtls13CipherSuite>,
) -> Result<Records, Error> {
let mut parsed_records: ArrayVec<Record, 16> = ArrayVec::new();
while !packet.is_empty() {
let record_end = if Dtls13Record::is_ciphertext_header(packet[0]) {
if packet[0] & 0x10 != 0 {
break;
}
if packet.len() < 2 {
return Err(Error::ParseIncomplete);
}
let flags = packet[0];
let s_flag = flags & 0b0000_1000 != 0;
let l_flag = flags & 0b0000_0100 != 0;
let seq_len = if s_flag { 2 } else { 1 };
let len_len = if l_flag { 2 } else { 0 };
let header_len = 1 + seq_len + len_len;
if packet.len() < header_len {
return Err(Error::ParseIncomplete);
}
if l_flag {
let len_offset = 1 + seq_len;
let length_bytes: [u8; 2] =
packet[len_offset..len_offset + 2].try_into().unwrap();
let length = u16::from_be_bytes(length_bytes) as usize;
header_len + length
} else {
packet.len()
}
} else {
if packet.len() < Dtls13Record::PLAINTEXT_HEADER_LEN {
return Err(Error::ParseIncomplete);
}
let length_bytes: [u8; 2] = packet[Dtls13Record::PLAINTEXT_LENGTH_OFFSET]
.try_into()
.unwrap();
let length = u16::from_be_bytes(length_bytes) as usize;
Dtls13Record::PLAINTEXT_HEADER_LEN + length
};
if packet.len() < record_end {
return Err(Error::ParseIncomplete);
}
let record_slice = &packet[..record_end];
match Record::parse(record_slice, decrypt, cs) {
Ok(record) => {
if let Some(record) = record {
if parsed_records.try_push(record).is_err() {
return Err(Error::TooManyRecords);
}
} else {
trace!("Discarding replayed rec");
}
}
Err(e) => return Err(e),
}
packet = &packet[record_end..];
}
let mut records = ArrayVec::new();
for record in parsed_records {
if let Some(record) = decrypt.classify_record(record)? {
records
.try_push(record)
.expect("filtered records cannot exceed parsed records");
}
}
Ok(Records { records })
}
}
impl Deref for Records {
type Target = [Record];
fn deref(&self) -> &Self::Target {
&self.records
}
}
pub struct Record {
buffer: Buf,
parsed: Box<ParsedRecord>,
}
impl Record {
pub fn parse(
record_slice: &[u8],
decrypt: &mut dyn RecordHandler,
cs: Option<Dtls13CipherSuite>,
) -> Result<Option<Record>, Error> {
let mut buffer = Buf::new();
buffer.extend_from_slice(record_slice);
let is_ciphertext = Dtls13Record::is_ciphertext_header(buffer[0]);
if is_ciphertext && decrypt.is_peer_encryption_enabled() {
let flags = buffer[0];
let s_flag = flags & 0b0000_1000 != 0;
let l_flag = flags & 0b0000_0100 != 0;
let seq_len: usize = if s_flag { 2 } else { 1 };
let len_len: usize = if l_flag { 2 } else { 0 };
let header_len = 1 + seq_len + len_len;
if buffer.len() >= header_len + 16 {
let ciphertext_sample: [u8; 16] =
buffer[header_len..header_len + 16].try_into().unwrap();
let epoch_bits = flags & 0x03;
let full_epoch = decrypt.resolve_epoch(epoch_bits);
decrypt.decrypt_sequence_number(
full_epoch,
&mut buffer[1..1 + seq_len],
&ciphertext_sample,
);
}
}
let parsed = match ParsedRecord::parse(&buffer, cs) {
Ok(p) => p,
Err(e) => {
trace!("Discarding record: parse failed: {}", e);
return Ok(None);
}
};
let parsed = Box::new(parsed);
let record = Record { buffer, parsed };
if !is_ciphertext || !decrypt.is_peer_encryption_enabled() {
return Ok(Some(record));
}
let epoch_bits = record.record().sequence.epoch as u8;
let full_epoch = decrypt.resolve_epoch(epoch_bits);
let seq_bits = record.record().sequence.sequence_number;
let s_flag = record_slice[0] & 0b0000_1000 != 0;
let full_seq = decrypt.resolve_sequence(full_epoch, seq_bits, s_flag);
let full_sequence = Sequence {
epoch: full_epoch,
sequence_number: full_seq,
};
if !decrypt.replay_check(full_sequence) {
return Ok(None);
}
let header_end = record.record().fragment_range.start;
if record.buffer.len() - header_end < decrypt.min_protected_fragment_len() {
return Ok(None);
}
let mut header_buf = [0u8; 5];
header_buf[..header_end].copy_from_slice(&record.buffer[..header_end]);
let mut buffer = record.buffer;
let ciphertext = &mut buffer[header_end..];
let new_len = {
let mut buffer = TmpBuf::new(ciphertext);
match decrypt.decrypt_record(&header_buf[..header_end], full_sequence, &mut buffer) {
Ok(()) => {}
Err(e) => {
trace!("Discarding ciphertext record: decryption failed: {}", e);
return Ok(None);
}
}
buffer.len()
};
decrypt.replay_update(full_sequence);
let decrypted = &buffer[header_end..header_end + new_len];
let (inner_content_type, content_len) = match recover_inner_content_type(decrypted) {
Ok(v) => v,
Err(e) => {
trace!("Discarding record: invalid inner content type: {}", e);
return Ok(None);
}
};
let parsed = ParsedRecord::parse_decrypted(
Dtls13Record {
content_type: inner_content_type,
sequence: full_sequence,
length: content_len as u16,
fragment_range: header_end..(header_end + content_len),
},
&buffer,
cs,
);
let parsed = Box::new(parsed);
Ok(Some(Record { buffer, parsed }))
}
pub fn record(&self) -> &Dtls13Record {
&self.parsed.record
}
pub fn handshakes(&self) -> &[Handshake] {
&self.parsed.handshakes
}
pub fn first_handshake(&self) -> Option<&Handshake> {
self.parsed.handshakes.first()
}
pub fn is_handled(&self) -> bool {
if self.parsed.handshakes.is_empty() {
self.parsed.handled.load(Ordering::Relaxed)
} else {
self.parsed.handshakes.iter().all(|h| h.is_handled())
}
}
pub fn set_handled(&self) {
assert!(self.parsed.handshakes.is_empty());
self.parsed.handled.store(true, Ordering::Relaxed);
}
pub fn buffer(&self) -> &[u8] {
&self.buffer
}
pub(crate) fn into_buffer(self) -> Buf {
self.buffer
}
}
pub struct ParsedRecord {
record: Dtls13Record,
handshakes: ArrayVec<Handshake, 8>,
handled: AtomicBool,
}
impl ParsedRecord {
pub fn parse(
input: &[u8],
cipher_suite: Option<Dtls13CipherSuite>,
) -> Result<ParsedRecord, Error> {
let (_, record) = Dtls13Record::parse(input, 0)?;
let handshakes = if record.content_type == ContentType::Handshake {
let fragment_offset = record.fragment_range.start;
parse_handshakes(record.fragment(input), fragment_offset, cipher_suite)
} else {
ArrayVec::new()
};
Ok(ParsedRecord {
record,
handshakes,
handled: AtomicBool::new(false),
})
}
pub fn parse_decrypted(
record: Dtls13Record,
input: &[u8],
cipher_suite: Option<Dtls13CipherSuite>,
) -> ParsedRecord {
let handshakes = if record.content_type == ContentType::Handshake {
let fragment_offset = record.fragment_range.start;
parse_handshakes(record.fragment(input), fragment_offset, cipher_suite)
} else {
ArrayVec::new()
};
ParsedRecord {
record,
handshakes,
handled: AtomicBool::new(false),
}
}
}
pub trait RecordHandler {
fn classify_record(&mut self, record: Record) -> Result<Option<Record>, Error>;
fn is_peer_encryption_enabled(&self) -> bool;
fn resolve_epoch(&self, epoch_bits: u8) -> u16;
fn resolve_sequence(&self, epoch: u16, seq_bits: u64, s_flag: bool) -> u64;
fn replay_check(&self, seq: Sequence) -> bool;
fn replay_update(&mut self, seq: Sequence);
fn min_protected_fragment_len(&self) -> usize;
fn decrypt_record(
&mut self,
header: &[u8],
seq: Sequence,
ciphertext: &mut TmpBuf,
) -> Result<(), Error>;
fn decrypt_sequence_number(
&self,
epoch: u16,
seq_bytes: &mut [u8],
ciphertext_sample: &[u8; 16],
);
}
fn parse_handshakes(
mut input: &[u8],
mut base_offset: usize,
cipher_suite: Option<Dtls13CipherSuite>,
) -> ArrayVec<Handshake, 8> {
let mut handshakes = ArrayVec::new();
while !input.is_empty() {
if let Ok((remaining, handshake)) = Handshake::parse(input, base_offset, cipher_suite, true)
{
let len = input.len() - remaining.len();
base_offset += len;
input = remaining;
if handshakes.try_push(handshake).is_err() {
break;
}
} else {
break;
}
}
handshakes
}
fn recover_inner_content_type(decrypted: &[u8]) -> Result<(ContentType, usize), Error> {
let mut i = decrypted.len();
while i > 0 && decrypted[i - 1] == 0 {
i -= 1;
}
if i == 0 {
return Err(Error::ParseError(nom::error::ErrorKind::Fail));
}
i -= 1;
let content_type = ContentType::from_u8(decrypted[i]);
Ok((content_type, i))
}
impl fmt::Debug for Incoming {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Incoming")
.field("records", &self.records())
.finish()
}
}
impl fmt::Debug for Record {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Record")
.field("record", &self.parsed.record)
.field("handshakes", &self.parsed.handshakes)
.finish()
}
}
impl std::panic::UnwindSafe for Incoming {}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Default)]
struct TestHandler {
classify_calls: usize,
dropped_acks: usize,
}
impl RecordHandler for TestHandler {
fn classify_record(&mut self, record: Record) -> Result<Option<Record>, Error> {
self.classify_calls += 1;
if record.record().content_type == ContentType::Ack {
self.dropped_acks += 1;
return Ok(None);
}
Ok(Some(record))
}
fn is_peer_encryption_enabled(&self) -> bool {
false
}
fn resolve_epoch(&self, _epoch_bits: u8) -> u16 {
panic!("resolve_epoch should not be called when peer encryption is disabled");
}
fn resolve_sequence(&self, _epoch: u16, _seq_bits: u64, _s_flag: bool) -> u64 {
panic!("resolve_sequence should not be called when peer encryption is disabled");
}
fn replay_check(&self, _seq: Sequence) -> bool {
panic!("replay_check should not be called when peer encryption is disabled");
}
fn replay_update(&mut self, _seq: Sequence) {
panic!("replay_update should not be called when peer encryption is disabled");
}
fn min_protected_fragment_len(&self) -> usize {
panic!(
"min_protected_fragment_len should not be called when peer encryption is disabled"
);
}
fn decrypt_record(
&mut self,
_header: &[u8],
_seq: Sequence,
_ciphertext: &mut TmpBuf,
) -> Result<(), Error> {
panic!("decrypt_record should not be called when peer encryption is disabled");
}
fn decrypt_sequence_number(
&self,
_epoch: u16,
_seq_bytes: &mut [u8],
_ciphertext_sample: &[u8; 16],
) {
panic!("decrypt_sequence_number should not be called when peer encryption is disabled");
}
}
fn build_plaintext_record(content_type: ContentType, seq: u64, fragment: &[u8]) -> Vec<u8> {
let mut out = Vec::new();
out.push(content_type.as_u8());
out.extend_from_slice(&[0xFE, 0xFD]);
out.extend_from_slice(&0u16.to_be_bytes());
out.extend_from_slice(&seq.to_be_bytes()[2..]);
out.extend_from_slice(&(fragment.len() as u16).to_be_bytes());
out.extend_from_slice(fragment);
out
}
fn build_ciphertext_record(epoch: u16, seq: u16, fragment: &[u8]) -> Vec<u8> {
let mut out = Vec::new();
let flags = 0b0010_0000 | 0b0000_1000 | 0b0000_0100 | (epoch as u8 & 0x03);
out.push(flags);
out.extend_from_slice(&seq.to_be_bytes());
out.extend_from_slice(&(fragment.len() as u16).to_be_bytes());
out.extend_from_slice(fragment);
out
}
#[test]
fn parse_packet_filters_control_records_after_packet_validation() {
let mut packet = Vec::new();
packet.extend_from_slice(&build_plaintext_record(ContentType::Ack, 1, &[0xAA, 0xBB]));
packet.extend_from_slice(&build_ciphertext_record(2, 2, &[0x11, 0x22, 0x33]));
let mut handler = TestHandler::default();
let incoming = Incoming::parse_packet(&packet, &mut handler, None)
.unwrap()
.expect("ciphertext application data record should remain");
assert_eq!(handler.classify_calls, 2);
assert_eq!(handler.dropped_acks, 1);
assert_eq!(incoming.records().len(), 1);
assert_eq!(
incoming.first().record().content_type,
ContentType::ApplicationData
);
assert_eq!(incoming.first().record().sequence.epoch, 2);
}
}