use heapless::Vec;
use crate::consts::{HS_ENCRYPTED_EXTENSIONS, HS_FINISHED};
use crate::server_flight::FlightError;
const EXT_RECORD_SIZE_LIMIT: u16 = 0x001C;
pub struct ServerFlightReassembler<const N: usize> {
buf: Vec<u8, N>,
}
#[derive(Debug, PartialEq, Eq, Clone, Copy, thiserror::Error)]
pub enum ReassemblyError {
#[error("server-flight reassembler capacity exceeded")]
Overflow,
}
impl<const N: usize> ServerFlightReassembler<N> {
pub const fn new() -> Self {
Self { buf: Vec::new() }
}
pub fn push_content(&mut self, content: &[u8]) -> Result<(), ReassemblyError> {
self.buf
.extend_from_slice(content)
.map_err(|_| ReassemblyError::Overflow)
}
pub fn is_complete(&self) -> bool {
self.flight_end_offset().is_some()
}
pub fn flight_bytes(&self) -> Option<&[u8]> {
let end = self.flight_end_offset()?;
Some(&self.buf[..end])
}
fn flight_end_offset(&self) -> Option<usize> {
let buf: &[u8] = &self.buf;
let mut i = 0;
while i + 4 <= buf.len() {
let msg_type = buf[i];
let len = usize::try_from(u32::from_be_bytes([0, buf[i + 1], buf[i + 2], buf[i + 3]]))
.ok()?;
if len > buf.len() - i - 4 {
return None;
}
i += 4 + len;
if msg_type == HS_FINISHED {
return Some(i);
}
}
None
}
pub fn peek_ee_record_size_limit(&self) -> Result<Option<u16>, FlightError> {
let buf = self.flight_bytes().ok_or(FlightError::Truncated)?;
if buf.len() < 4 {
return Err(FlightError::Truncated);
}
if buf[0] != HS_ENCRYPTED_EXTENSIONS {
return Err(FlightError::UnexpectedHandshakeType {
expected: HS_ENCRYPTED_EXTENSIONS,
got: buf[0],
});
}
let ee_body_len = u32::from_be_bytes([0, buf[1], buf[2], buf[3]]) as usize;
if buf.len() < 4 + ee_body_len {
return Err(FlightError::Truncated);
}
let ee_body = &buf[4..4 + ee_body_len];
if ee_body.len() < 2 {
return Err(FlightError::Truncated);
}
let ext_total = u16::from_be_bytes([ee_body[0], ee_body[1]]) as usize;
if ee_body.len() != 2 + ext_total {
return Err(FlightError::Truncated);
}
let mut rest = &ee_body[2..];
let mut found: Option<u16> = None;
while !rest.is_empty() {
if rest.len() < 4 {
return Err(FlightError::Truncated);
}
let ext_type = u16::from_be_bytes([rest[0], rest[1]]);
let ext_len = u16::from_be_bytes([rest[2], rest[3]]) as usize;
if rest.len() < 4 + ext_len {
return Err(FlightError::Truncated);
}
let body = &rest[4..4 + ext_len];
if ext_type == EXT_RECORD_SIZE_LIMIT {
if body.len() != 2 {
return Err(FlightError::Truncated);
}
if found.is_some() {
return Err(FlightError::DuplicateExtension { ext_type });
}
found = Some(u16::from_be_bytes([body[0], body[1]]));
}
rest = &rest[4 + ext_len..];
}
Ok(found)
}
}
impl<const N: usize> Default for ServerFlightReassembler<N> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::consts::{HS_CERTIFICATE, HS_CERTIFICATE_VERIFY};
fn ee(body_len: usize) -> alloc_helper::Msg {
alloc_helper::msg(HS_ENCRYPTED_EXTENSIONS, body_len)
}
fn cert(body_len: usize) -> alloc_helper::Msg {
alloc_helper::msg(HS_CERTIFICATE, body_len)
}
fn cv(body_len: usize) -> alloc_helper::Msg {
alloc_helper::msg(HS_CERTIFICATE_VERIFY, body_len)
}
fn fin(body_len: usize) -> alloc_helper::Msg {
alloc_helper::msg(HS_FINISHED, body_len)
}
mod alloc_helper {
pub struct Msg(pub [u8; 1024], pub usize);
pub fn msg(ty: u8, body_len: usize) -> Msg {
assert!(body_len + 4 <= 1024);
let mut buf = [0u8; 1024];
buf[0] = ty;
buf[1] = ((body_len >> 16) & 0xff) as u8;
buf[2] = ((body_len >> 8) & 0xff) as u8;
buf[3] = (body_len & 0xff) as u8;
Msg(buf, 4 + body_len)
}
impl Msg {
pub fn as_slice(&self) -> &[u8] {
&self.0[..self.1]
}
}
}
#[test]
fn empty_buffer_is_not_complete() {
let r: ServerFlightReassembler<128> = ServerFlightReassembler::new();
assert!(!r.is_complete());
assert!(r.buf.is_empty());
}
#[test]
fn single_record_full_flight_is_complete() {
let mut r: ServerFlightReassembler<512> = ServerFlightReassembler::new();
let mut combined = heapless::Vec::<u8, 512>::new();
combined.extend_from_slice(ee(2).as_slice()).unwrap();
combined.extend_from_slice(cert(40).as_slice()).unwrap();
combined.extend_from_slice(cv(70).as_slice()).unwrap();
combined.extend_from_slice(fin(32).as_slice()).unwrap();
r.push_content(&combined).unwrap();
assert!(r.is_complete());
}
#[test]
fn partial_finished_is_not_complete() {
let mut r: ServerFlightReassembler<512> = ServerFlightReassembler::new();
r.push_content(ee(2).as_slice()).unwrap();
r.push_content(cert(40).as_slice()).unwrap();
let cv_full = cv(70);
let cv_slice = cv_full.as_slice();
r.push_content(&cv_slice[..cv_slice.len() - 5]).unwrap();
assert!(!r.is_complete());
}
#[test]
fn dangling_partial_header_before_finished_is_not_complete() {
let mut r: ServerFlightReassembler<512> = ServerFlightReassembler::new();
r.push_content(ee(2).as_slice()).unwrap();
r.push_content(&[HS_CERTIFICATE, 0]).unwrap();
assert!(!r.is_complete());
assert!(r.flight_bytes().is_none());
}
#[test]
fn multi_record_concat_completes() {
let mut r: ServerFlightReassembler<512> = ServerFlightReassembler::new();
let cert_full = cert(60);
let cert_slice = cert_full.as_slice();
let mut rec1 = heapless::Vec::<u8, 128>::new();
rec1.extend_from_slice(ee(2).as_slice()).unwrap();
rec1.extend_from_slice(&cert_slice[..30]).unwrap();
r.push_content(&rec1).unwrap();
assert!(!r.is_complete());
let mut rec2 = heapless::Vec::<u8, 256>::new();
rec2.extend_from_slice(&cert_slice[30..]).unwrap();
rec2.extend_from_slice(cv(70).as_slice()).unwrap();
r.push_content(&rec2).unwrap();
assert!(!r.is_complete());
r.push_content(fin(32).as_slice()).unwrap();
assert!(r.is_complete());
}
#[test]
fn trailing_bytes_after_finished_still_complete() {
let mut r: ServerFlightReassembler<512> = ServerFlightReassembler::new();
let mut combined = heapless::Vec::<u8, 512>::new();
combined.extend_from_slice(ee(2).as_slice()).unwrap();
combined.extend_from_slice(cert(40).as_slice()).unwrap();
combined.extend_from_slice(cv(70).as_slice()).unwrap();
combined.extend_from_slice(fin(32).as_slice()).unwrap();
let flight_only_len = combined.len();
let nst = alloc_helper::msg(4, 8);
combined.extend_from_slice(nst.as_slice()).unwrap();
r.push_content(&combined).unwrap();
assert!(
r.is_complete(),
"is_complete must tolerate trailing post-handshake bytes"
);
let flight = r.flight_bytes().expect("flight_bytes after Finished");
assert_eq!(flight.len(), flight_only_len);
assert!(r.buf.len() > flight.len());
assert_eq!(&r.buf[flight.len()..], nst.as_slice());
}
#[test]
fn flight_bytes_returns_none_until_finished() {
let mut r: ServerFlightReassembler<512> = ServerFlightReassembler::new();
r.push_content(ee(2).as_slice()).unwrap();
r.push_content(cert(40).as_slice()).unwrap();
assert!(r.flight_bytes().is_none());
r.push_content(cv(70).as_slice()).unwrap();
assert!(r.flight_bytes().is_none());
r.push_content(fin(32).as_slice()).unwrap();
assert!(r.flight_bytes().is_some());
}
#[test]
fn overflow_returns_error() {
let mut r: ServerFlightReassembler<32> = ServerFlightReassembler::new();
let big = [0u8; 64];
assert_eq!(r.push_content(&big), Err(ReassemblyError::Overflow));
}
#[test]
fn clear_resets_state() {
let mut r: ServerFlightReassembler<128> = ServerFlightReassembler::new();
r.push_content(ee(2).as_slice()).unwrap();
r.push_content(fin(32).as_slice()).unwrap();
assert!(r.is_complete());
r.buf.clear();
assert!(r.buf.is_empty());
assert!(!r.is_complete());
}
fn ee_with_extensions(extensions: &[u8]) -> heapless::Vec<u8, 512> {
let mut body = heapless::Vec::<u8, 512>::new();
let total = extensions.len() as u16;
body.extend_from_slice(&total.to_be_bytes()).unwrap();
body.extend_from_slice(extensions).unwrap();
let mut msg = heapless::Vec::<u8, 512>::new();
msg.push(HS_ENCRYPTED_EXTENSIONS).unwrap();
let len = body.len() as u32;
msg.push(((len >> 16) & 0xff) as u8).unwrap();
msg.push(((len >> 8) & 0xff) as u8).unwrap();
msg.push((len & 0xff) as u8).unwrap();
msg.extend_from_slice(&body).unwrap();
msg
}
fn record_size_limit_ext(value: u16) -> [u8; 6] {
let v = value.to_be_bytes();
[
(EXT_RECORD_SIZE_LIMIT >> 8) as u8,
(EXT_RECORD_SIZE_LIMIT & 0xff) as u8,
0x00,
0x02, v[0],
v[1],
]
}
fn complete_flight_with_ee(ee_bytes: &[u8]) -> ServerFlightReassembler<1024> {
let mut r: ServerFlightReassembler<1024> = ServerFlightReassembler::new();
let mut combined = heapless::Vec::<u8, 1024>::new();
combined.extend_from_slice(ee_bytes).unwrap();
combined.extend_from_slice(cert(40).as_slice()).unwrap();
combined.extend_from_slice(cv(70).as_slice()).unwrap();
combined.extend_from_slice(fin(32).as_slice()).unwrap();
r.push_content(&combined).unwrap();
r
}
#[test]
fn peek_ee_record_size_limit_absent_returns_none() {
let ee_bytes = ee_with_extensions(&[]);
let r = complete_flight_with_ee(&ee_bytes);
assert_eq!(r.peek_ee_record_size_limit(), Ok(None));
}
#[test]
fn peek_ee_record_size_limit_present_returns_value() {
let ext = record_size_limit_ext(8192);
let ee_bytes = ee_with_extensions(&ext);
let r = complete_flight_with_ee(&ee_bytes);
assert_eq!(r.peek_ee_record_size_limit(), Ok(Some(8192)));
}
#[test]
fn peek_ee_record_size_limit_skips_unrelated_extensions() {
let mut exts = heapless::Vec::<u8, 32>::new();
exts.extend_from_slice(&[0x00, 0x2A, 0x00, 0x04, 0xab, 0xab, 0xab, 0xab])
.unwrap();
exts.extend_from_slice(&record_size_limit_ext(2048))
.unwrap();
let ee_bytes = ee_with_extensions(&exts);
let r = complete_flight_with_ee(&ee_bytes);
assert_eq!(r.peek_ee_record_size_limit(), Ok(Some(2048)));
}
#[test]
fn peek_ee_record_size_limit_duplicate_rejected() {
let mut exts = heapless::Vec::<u8, 32>::new();
exts.extend_from_slice(&record_size_limit_ext(2048))
.unwrap();
exts.extend_from_slice(&record_size_limit_ext(4096))
.unwrap();
let ee_bytes = ee_with_extensions(&exts);
let r = complete_flight_with_ee(&ee_bytes);
assert_eq!(
r.peek_ee_record_size_limit(),
Err(FlightError::DuplicateExtension {
ext_type: EXT_RECORD_SIZE_LIMIT
})
);
}
#[test]
fn peek_ee_record_size_limit_wrong_first_msg_type() {
let mut r: ServerFlightReassembler<512> = ServerFlightReassembler::new();
let mut combined = heapless::Vec::<u8, 512>::new();
combined.extend_from_slice(cert(40).as_slice()).unwrap();
combined.extend_from_slice(fin(32).as_slice()).unwrap();
r.push_content(&combined).unwrap();
assert_eq!(
r.peek_ee_record_size_limit(),
Err(FlightError::UnexpectedHandshakeType {
expected: HS_ENCRYPTED_EXTENSIONS,
got: 11
})
);
}
#[test]
fn peek_ee_record_size_limit_returns_truncated_before_flight_complete() {
let r: ServerFlightReassembler<128> = ServerFlightReassembler::new();
assert_eq!(r.peek_ee_record_size_limit(), Err(FlightError::Truncated));
}
#[test]
fn peek_ee_record_size_limit_ext_with_wrong_len_rejected() {
let bad_ext = [
(EXT_RECORD_SIZE_LIMIT >> 8) as u8,
(EXT_RECORD_SIZE_LIMIT & 0xff) as u8,
0x00,
0x03,
0x10,
0x00,
0x00,
];
let ee_bytes = ee_with_extensions(&bad_ext);
let r = complete_flight_with_ee(&ee_bytes);
assert_eq!(r.peek_ee_record_size_limit(), Err(FlightError::Truncated));
}
#[test]
fn peek_ee_record_size_limit_extensions_total_inconsistent() {
let mut body = heapless::Vec::<u8, 32>::new();
body.extend_from_slice(&20u16.to_be_bytes()).unwrap();
body.extend_from_slice(&record_size_limit_ext(1024))
.unwrap();
let mut msg = heapless::Vec::<u8, 32>::new();
msg.push(HS_ENCRYPTED_EXTENSIONS).unwrap();
let len = body.len() as u32;
msg.push(((len >> 16) & 0xff) as u8).unwrap();
msg.push(((len >> 8) & 0xff) as u8).unwrap();
msg.push((len & 0xff) as u8).unwrap();
msg.extend_from_slice(&body).unwrap();
let r = complete_flight_with_ee(&msg);
assert_eq!(r.peek_ee_record_size_limit(), Err(FlightError::Truncated));
}
#[test]
fn peek_ee_record_size_limit_ext_header_truncated() {
let mut body = heapless::Vec::<u8, 16>::new();
body.extend_from_slice(&3u16.to_be_bytes()).unwrap();
body.extend_from_slice(&[0x00, 0x2A, 0x00]).unwrap();
let mut msg = heapless::Vec::<u8, 16>::new();
msg.push(HS_ENCRYPTED_EXTENSIONS).unwrap();
let len = body.len() as u32;
msg.push(((len >> 16) & 0xff) as u8).unwrap();
msg.push(((len >> 8) & 0xff) as u8).unwrap();
msg.push((len & 0xff) as u8).unwrap();
msg.extend_from_slice(&body).unwrap();
let r = complete_flight_with_ee(&msg);
assert_eq!(r.peek_ee_record_size_limit(), Err(FlightError::Truncated));
}
}