diff --git a/CHANGELOG.md b/CHANGELOG.md index 82d4fbb..65cb692 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,7 @@ # Unreleased + * Restore DTLS buffer reuse and reject unpooled buffer returns #167 + # 0.7.4 * Preserve DTLS 1.3 client key shares on cookie-only retries #163 diff --git a/src/auto.rs b/src/auto.rs index c32219c..f220a3f 100644 --- a/src/auto.rs +++ b/src/auto.rs @@ -50,6 +50,10 @@ const EXT_RENEGOTIATION_INFO: u16 = 0xFF01; /// into a `Client13` via `new_from_hybrid`. /// - **DTLS 1.2 fork**: all state is discarded (HelloVerifyRequest clears /// the transcript), and a fresh `Client12` is created. +/// +/// These buffers are temporary probe state created before an engine exists. +/// The selected engine copies the needed bytes into its own buffers; probe +/// buffers are dropped and never returned to the engine's pool. pub(crate) struct HybridClientHello { /// Client random used in the ClientHello. pub random: Random, diff --git a/src/buffer.rs b/src/buffer.rs index b1856c1..79e5a8c 100644 --- a/src/buffer.rs +++ b/src/buffer.rs @@ -23,15 +23,19 @@ impl BufferPool { /// /// Creates a new buffer if none is free. pub fn pop(&mut self) -> Buf { - if self.free.is_empty() { - self.free.push_back(Buf::new()); - } - // Unwrap is OK see above handling of empty. - self.free.pop_front().unwrap() + let mut buffer = self.free.pop_front().unwrap_or_default(); + buffer.from_pool = true; + buffer } /// Return a buffer to the pool. + /// + /// Panics if the buffer was not acquired from a pool. pub fn push(&mut self, mut buffer: Buf) { + assert!( + buffer.from_pool, + "only buffers acquired from a pool may be returned" + ); buffer.clear(); self.free.push_front(buffer); } @@ -47,47 +51,55 @@ impl fmt::Debug for BufferPool { /// Growable buffer wrapper used throughout dimpl for efficient memory management. /// -/// This is a newtype around `Vec` that provides convenient access to byte buffers +/// This wrapper around `Vec` provides convenient access to byte buffers /// and integrates with dimpl's buffer pooling system. +/// When filling a borrowed `&mut Buf`, mutate its contents instead of replacing +/// the buffer so its pool provenance is preserved. #[derive(Default)] -pub struct Buf(Vec); +pub struct Buf { + data: Vec, + from_pool: bool, +} impl Buf { - /// Create a new empty buffer. + /// Create an empty buffer for temporary or independently owned data. + /// + /// This buffer must not be returned to a pool. Recyclable buffers must be + /// acquired from their pool. pub fn new() -> Self { Self::default() } /// Clear the buffer, removing all data. pub fn clear(&mut self) { - self.0.clear(); + self.data.clear(); } /// Extend the buffer with a slice of bytes. pub fn extend_from_slice(&mut self, other: &[u8]) { - self.0.extend_from_slice(other); + self.data.extend_from_slice(other); } /// Push a single byte onto the buffer. pub fn push(&mut self, byte: u8) { - self.0.push(byte); + self.data.push(byte); } /// Resize the buffer to the specified length, filling with the given value. pub fn resize(&mut self, len: usize, value: u8) { - self.0.resize(len, value); + self.data.resize(len, value); } /// Convert the buffer into the underlying `Vec`. - pub fn into_vec(mut self) -> Vec { - std::mem::take(&mut self.0) + pub fn into_vec(self) -> Vec { + self.data } } // aws-lc-rs AEAD operations require Extend<&u8> for appending authentication tags impl<'a> Extend<&'a u8> for Buf { fn extend>(&mut self, iter: T) { - self.0.extend(iter.into_iter().copied()); + self.data.extend(iter.into_iter().copied()); } } @@ -95,31 +107,33 @@ impl Deref for Buf { type Target = [u8]; fn deref(&self) -> &Self::Target { - &self.0 + &self.data } } impl DerefMut for Buf { fn deref_mut(&mut self) -> &mut Self::Target { - &mut self.0 + &mut self.data } } impl AsRef<[u8]> for Buf { fn as_ref(&self) -> &[u8] { - &self.0 + &self.data } } impl AsMut<[u8]> for Buf { fn as_mut(&mut self) -> &mut [u8] { - &mut self.0 + &mut self.data } } impl fmt::Debug for Buf { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.debug_struct("Buf").field("len", &self.0.len()).finish() + f.debug_struct("Buf") + .field("len", &self.data.len()) + .finish() } } @@ -131,7 +145,10 @@ pub trait ToBuf { impl ToBuf for Vec { fn to_buf(self) -> Buf { - Buf(self) + Buf { + data: self, + from_pool: false, + } } } @@ -210,11 +227,59 @@ impl<'a> aes_gcm::aead::Buffer for TmpBuf<'a> { #[cfg(feature = "rust-crypto")] impl aes_gcm::aead::Buffer for Buf { fn extend_from_slice(&mut self, other: &[u8]) -> Result<(), aes_gcm::aead::Error> { - self.0.extend_from_slice(other); + self.data.extend_from_slice(other); Ok(()) } fn truncate(&mut self, len: usize) { - self.0.truncate(len); + self.data.truncate(len); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn pooled_buffers_keep_capacity_and_are_cleared() { + let mut pool = BufferPool::default(); + let mut buffer = pool.pop(); + buffer.extend_from_slice(&[0xAA; 256]); + let allocation = buffer.as_ptr(); + pool.push(buffer); + + let reused = pool.pop(); + assert!(reused.is_empty()); + assert_eq!(reused.as_ptr(), allocation); + assert!(reused.into_vec().capacity() >= 256); + } + + #[test] + #[should_panic(expected = "only buffers acquired from a pool may be returned")] + fn rejects_unpooled_buffer() { + BufferPool::default().push(Buf::new()); + } + + #[test] + #[should_panic(expected = "only buffers acquired from a pool may be returned")] + fn rejects_converted_buffer() { + BufferPool::default().push(vec![0xAA; 256].to_buf()); + } + + #[test] + #[should_panic(expected = "only buffers acquired from a pool may be returned")] + fn converting_to_vec_does_not_keep_pool_provenance() { + let mut pool = BufferPool::default(); + let buffer = pool.pop(); + pool.push(buffer.into_vec().to_buf()); + } + + #[test] + #[should_panic(expected = "only buffers acquired from a pool may be returned")] + fn taking_buffer_does_not_make_replacement_pooled() { + let mut pool = BufferPool::default(); + let mut buffer = pool.pop(); + pool.push(std::mem::take(&mut buffer)); + pool.push(buffer); } } diff --git a/src/dtls12/client.rs b/src/dtls12/client.rs index 4ec92a6..cf7150d 100644 --- a/src/dtls12/client.rs +++ b/src/dtls12/client.rs @@ -92,6 +92,8 @@ pub(crate) enum LocalEvent { impl Client { pub(crate) fn new_with_engine(mut engine: Engine, now: Instant) -> Client { engine.set_client(true); + let extension_data = engine.pop_buffer(); + let defragment_buffer = engine.pop_buffer(); Client { state: State::SendClientHello, @@ -99,11 +101,11 @@ impl Client { random: None, session_id: None, cookie: None, - extension_data: Buf::new(), + extension_data, negotiated_srtp_profile: None, server_random: None, server_certificates: Vec::with_capacity(3), - defragment_buffer: Buf::new(), + defragment_buffer, certificate_verify: false, captured_session_hash: None, last_now: now, @@ -156,6 +158,8 @@ impl Client { engine.transcript.extend_from_slice(handshake_fragment); // Advance epoch-0 record sequence past the hybrid CH record. engine.advance_epoch_0_sequence(); + let extension_data = engine.pop_buffer(); + let defragment_buffer = engine.pop_buffer(); let mut client = Client { state: State::AwaitHelloVerifyRequest, @@ -163,11 +167,11 @@ impl Client { random: Some(random), session_id: None, cookie: None, - extension_data: Buf::new(), + extension_data, negotiated_srtp_profile: None, server_random: None, server_certificates: Vec::with_capacity(3), - defragment_buffer: Buf::new(), + defragment_buffer, certificate_verify: false, captured_session_hash: None, last_now: now, @@ -740,13 +744,13 @@ impl State { // Process the server key exchange parameters // We already have the curve and public key extracted - let mut kx_buf = client.engine.pop_buffer(); + let kx_buf = client.engine.pop_buffer(); + let shared_secret = client.engine.pop_buffer(); client .engine .crypto_context_mut() - .process_ecdh_params(named_group, public_key_vec, &mut kx_buf) + .process_ecdh_params(named_group, public_key_vec, kx_buf, shared_secret) .map_err(Error::CryptoError)?; - client.engine.push_buffer(kx_buf); Ok(Self::AwaitCertificateRequest) } @@ -930,7 +934,7 @@ impl State { ))?; let suite_hash = cipher_suite.hash_algorithm(); - let mut buf = Buf::new(); + let mut buf = client.engine.pop_buffer(); client.engine.transcript_hash(suite_hash, &mut buf); client.captured_session_hash = Some(buf); diff --git a/src/dtls12/context.rs b/src/dtls12/context.rs index 30502c8..f4c455b 100644 --- a/src/dtls12/context.rs +++ b/src/dtls12/context.rs @@ -138,14 +138,14 @@ impl CryptoContext { pub fn compute_shared_secret( &mut self, peer_public_key: &[u8], - buf: &mut Buf, + mut buf: Buf, ) -> Result<(), CryptoError> { let ke = self .key_exchange .take() .ok_or(CryptoError::KeyExchangeNotInitialized)?; - ke.complete(peer_public_key, buf)?; - self.pre_master_secret = Some(core::mem::take(buf)); + ke.complete(peer_public_key, &mut buf)?; + self.pre_master_secret = Some(buf); // Note: we keep key_exchange_public_key since it may be needed later Ok(()) } @@ -163,6 +163,7 @@ impl CryptoContext { let psk = self.psk.as_ref().ok_or(CryptoError::PskNotSet)?; let n = psk.len(); // Total: 2 + N + 2 + N = 2N + 4 + // This secret is dropped after deriving the master secret, never pooled. let mut pms = Buf::new(); pms.extend_from_slice(&(n as u16).to_be_bytes()); pms.resize(pms.len() + n, 0); @@ -176,7 +177,7 @@ impl CryptoContext { pub fn init_ecdh_server( &mut self, named_group: NamedGroup, - kx_buf: &mut Buf, + kx_buf: Buf, ) -> Result<&[u8], CryptoError> { // Find the matching key exchange group from the provider let kx_group = self @@ -185,8 +186,7 @@ impl CryptoContext { .find(|g| g.name() == named_group) .ok_or(CryptoError::UnsupportedEcdheNamedGroup(named_group))?; - kx_buf.clear(); - self.key_exchange = Some(kx_group.start_exchange(core::mem::take(kx_buf))?); + self.key_exchange = Some(kx_group.start_exchange(kx_buf)?); self.maybe_init_key_exchange() } @@ -195,7 +195,8 @@ impl CryptoContext { &mut self, group: NamedGroup, server_public: &[u8], - kx_buf: &mut Buf, + kx_buf: Buf, + shared_secret: Buf, ) -> Result<(), CryptoError> { // Find the matching key exchange group from the provider let kx_group = self @@ -205,14 +206,13 @@ impl CryptoContext { .ok_or(CryptoError::UnsupportedEcdheNamedGroup(group))?; // Create a new ECDH key exchange - kx_buf.clear(); - self.key_exchange = Some(kx_group.start_exchange(core::mem::take(kx_buf))?); + self.key_exchange = Some(kx_group.start_exchange(kx_buf)?); // Generate our keypair let _our_public = self.maybe_init_key_exchange()?; // Compute shared secret with the server's public key - self.compute_shared_secret(server_public, kx_buf)?; + self.compute_shared_secret(server_public, shared_secret)?; Ok(()) } diff --git a/src/dtls12/engine.rs b/src/dtls12/engine.rs index 9310eb7..b60fd54 100644 --- a/src/dtls12/engine.rs +++ b/src/dtls12/engine.rs @@ -142,11 +142,13 @@ impl Engine { ExponentialBackoff::new(config.flight_start_rto(), config.flight_retries(), &mut rng); let crypto_context = CryptoContext::new(auth, Arc::clone(&config)); + let mut buffers_free = BufferPool::default(); + let transcript = buffers_free.pop(); Self { config, rng, - buffers_free: BufferPool::default(), + buffers_free, sequence_epoch_0: Sequence::new(0), sequence_epoch_n: Sequence::new(1), queue_rx: QueueRx::new(), @@ -159,7 +161,7 @@ impl Engine { is_client: false, peer_handshake_seq_no: 0, next_handshake_seq_no: 0, - transcript: Buf::new(), + transcript, replay: ReplayWindow::new(), flight_saved_records: Vec::new(), flight_backoff, @@ -244,6 +246,7 @@ impl Engine { self.config.max_queue_rx(), self.queue_rx ); + self.recycle_incoming(incoming); return Err(Error::ReceiveQueueFull); } @@ -280,12 +283,16 @@ impl Engine { // drive a resend. if let Some(dupe_seq) = maybe_dupe_seq { if dupe_seq < self.peer_handshake_seq_no && !self.peer_handshake_confirmed { - self.flight_resend("dupe triggers resend")?; + if let Err(error) = self.flight_resend("dupe triggers resend") { + self.recycle_incoming(incoming); + return Err(error); + } } } // Drop old duplicates we've already processed - don't let them block newer messages. if handshake.header.message_seq < self.peer_handshake_seq_no { + self.recycle_incoming(incoming); return Ok(()); } @@ -293,11 +300,13 @@ impl Engine { // Keep old plaintext handshake records available long enough to // trigger flight resends above, but never queue or process them as // new messages after peer encryption is enabled. + self.recycle_incoming(incoming); return Ok(()); } // Reject new handshakes after initial handshake is complete (renegotiation not supported). if self.release_app_data && handshake.header.message_seq >= self.peer_handshake_seq_no { + self.recycle_incoming(incoming); return Err(Error::RenegotiationAttempt); } @@ -318,6 +327,7 @@ impl Engine { } Ok(_) => { // Exact duplicate handshake fragment + self.recycle_incoming(incoming); } } @@ -332,6 +342,7 @@ impl Engine { && seq_current.epoch == 0 && first.record().content_type == ContentType::Handshake { + self.recycle_incoming(incoming); return Ok(()); } @@ -362,12 +373,19 @@ impl Engine { // For epoch 1, we have the replay window and there should // be no duplicates. assert_eq!(seq_current.epoch, 0); + self.recycle_incoming(incoming); } } Ok(()) } + fn recycle_incoming(&mut self, incoming: Incoming) { + for record in incoming.into_records() { + self.push_buffer(record.into_buffer()); + } + } + pub fn handle_timeout(&mut self, now: Instant) -> Result<(), Error> { if self.connect_timeout == Timeout::Unarmed { debug!( @@ -1255,6 +1273,14 @@ impl Engine { } impl RecordHandler for Engine { + fn pop_buffer(&mut self) -> Buf { + Engine::pop_buffer(self) + } + + fn push_buffer(&mut self, buffer: Buf) { + Engine::push_buffer(self, buffer); + } + fn classify_record(&mut self, record: Record) -> Result, Error> { let epoch = record.record().sequence.epoch; @@ -1398,3 +1424,42 @@ impl RecordHandler for Engine { self.release_app_data } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn duplicate_datagram_recycles_its_pooled_buffer() { + let mut engine = Engine::new(Arc::new(Config::default()), AuthMode::Psk); + let packet = [ + 0x14, 0xFE, 0xFD, // ChangeCipherSpec, DTLS 1.2. + 0, 0, // Epoch 0. + 0, 0, 0, 0, 0, 1, // Record sequence 1. + 0, 1, 1, // Fragment length 1 and ChangeCipherSpec body. + ]; + engine.parse_packet(&packet).expect("queue first datagram"); + let mut output = [0u8; 2048]; + assert!(matches!( + engine.poll_output(&mut output, Instant::now()), + Output::Timeout(_) + )); + + let mut buffer = engine.pop_buffer(); + buffer.resize(2048, 0xAA); + engine.push_buffer(buffer); + + engine + .parse_packet(&packet) + .expect("discard duplicate datagram"); + assert!(matches!( + engine.poll_output(&mut output, Instant::now()), + Output::Timeout(_) + )); + assert_eq!(engine.queue_rx.len(), 1); + assert!( + engine.pop_buffer().into_vec().capacity() >= 2048, + "discarding a duplicate must preserve its reusable receive allocation" + ); + } +} diff --git a/src/dtls12/incoming.rs b/src/dtls12/incoming.rs index 1083af4..8d7b45e 100644 --- a/src/dtls12/incoming.rs +++ b/src/dtls12/incoming.rs @@ -1,8 +1,8 @@ +use std::fmt; use std::ops::Deref; use std::sync::atomic::{AtomicBool, Ordering}; use arrayvec::ArrayVec; -use std::fmt; use crate::buffer::{Buf, TmpBuf}; use crate::crypto::{Aad, Nonce}; @@ -68,12 +68,36 @@ pub struct Records { impl Records { pub fn parse( - mut packet: &[u8], + packet: &[u8], decrypt: &mut dyn RecordHandler, cs: Option, ) -> Result { let mut parsed_records: ArrayVec = ArrayVec::new(); + if let Err(error) = Self::parse_records(packet, decrypt, cs, &mut parsed_records) { + for record in parsed_records { + decrypt.push_buffer(record.into_buffer()); + } + return Err(error); + } + + 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 }) + } + fn parse_records( + mut packet: &[u8], + decrypt: &mut dyn RecordHandler, + cs: Option, + parsed_records: &mut ArrayVec, + ) -> Result<(), InternalError> { // Find record boundaries and copy each record ONCE from the packet while !packet.is_empty() { if packet.len() < DTLSRecord::HEADER_LEN { @@ -92,7 +116,8 @@ impl Records { let record_slice = &packet[..record_end]; let record = Record::parse(record_slice, decrypt, cs)?; if let Some(record) = record { - if parsed_records.try_push(record).is_err() { + if let Err(error) = parsed_records.try_push(record) { + decrypt.push_buffer(error.element().into_buffer()); return Err(InternalError::too_many_records()); } } else { @@ -102,16 +127,7 @@ impl Records { 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 }) + Ok(()) } } @@ -139,9 +155,24 @@ impl Record { cs: Option, ) -> Result, InternalError> { // ONLY COPY: UDP packet slice -> pooled buffer - let mut buffer = Buf::new(); + let mut buffer = decrypt.pop_buffer(); buffer.extend_from_slice(record_slice); - let parsed = match ParsedRecord::parse(&buffer, cs, 0) { + + match Self::parse_buffer(&mut buffer, decrypt, cs) { + Ok(Some(parsed)) => Ok(Some(Record { buffer, parsed })), + result => { + decrypt.push_buffer(buffer); + result.map(|_| None) + } + } + } + + fn parse_buffer( + buffer: &mut Buf, + decrypt: &mut dyn RecordHandler, + cs: Option, + ) -> Result>, InternalError> { + let parsed = match ParsedRecord::parse(buffer, cs, 0) { Ok(p) => p, Err(e) => { // RFC 6347 §4.1.2.7: Invalid records SHOULD be silently discarded. @@ -151,18 +182,17 @@ impl Record { } }; let parsed = Box::new(parsed); - let record = Record { buffer, parsed }; // It is not enough to only look at the epoch, since to be able to decrypt the entire // preceeding set of flights sets up the cryptographic context. In a situation with // packet loss, we can end up seeing epoch 1 records before we can decrypt them. - let is_epoch_0 = record.record().sequence.epoch == 0; + let is_epoch_0 = parsed.record.sequence.epoch == 0; if is_epoch_0 || !decrypt.is_peer_encryption_enabled() { - return Ok(Some(record)); + return Ok(Some(parsed)); } // We need to decrypt the record and redo the parsing. - let dtls = record.record(); + let dtls = &parsed.record; let sequence = dtls.sequence; let content_type = dtls.content_type; @@ -177,10 +207,7 @@ impl Record { } // Get a reference to the buffer - let (aad, nonce) = decrypt.decryption_aad_and_nonce(dtls, &record.buffer); - - // Extract the buffer for decryption - let mut buffer = record.buffer; + let (aad, nonce) = decrypt.decryption_aad_and_nonce(dtls, buffer); // Local shorthand for where the encrypted ciphertext starts let ciph = DTLSRecord::HEADER_LEN + explicit_nonce_len; @@ -219,10 +246,10 @@ impl Record { buffer[11] = (new_len >> 8) as u8; buffer[12] = new_len as u8; - let parsed = ParsedRecord::parse(&buffer, cs, explicit_nonce_len)?; + let parsed = ParsedRecord::parse(buffer, cs, explicit_nonce_len)?; let parsed = Box::new(parsed); - Ok(Some(Record { buffer, parsed })) + Ok(Some(parsed)) } pub fn record(&self) -> &DTLSRecord { @@ -295,9 +322,13 @@ impl ParsedRecord { /// Trait abstracting record parsing-time handling for incoming records. /// /// This decouples the record parser from the full `Engine`, allowing the parse loop -/// to decrypt records, classify control records, and queue only the records that +/// to acquire buffers, decrypt records, classify control records, and queue only the records that /// should survive into `Incoming`. pub trait RecordHandler { + /// Acquire an empty buffer whose allocation can be reused for a record. + fn pop_buffer(&mut self) -> Buf; + /// Recycle a buffer when parsing fails or silently discards a record. + fn push_buffer(&mut self, buffer: Buf); fn classify_record(&mut self, record: Record) -> Result, Error>; fn is_peer_encryption_enabled(&self) -> bool; fn replay_check(&self, seq: Sequence) -> bool; @@ -394,19 +425,35 @@ impl std::panic::UnwindSafe for Incoming {} #[cfg(test)] mod tests { + use crate::buffer::BufferPool; + use super::*; #[derive(Default)] struct TestHandler { + buffers: BufferPool, + buffers_acquired: usize, + buffers_returned: usize, classify_calls: usize, dropped_alerts: usize, } impl RecordHandler for TestHandler { + fn pop_buffer(&mut self) -> Buf { + self.buffers_acquired += 1; + self.buffers.pop() + } + + fn push_buffer(&mut self, buffer: Buf) { + self.buffers_returned += 1; + self.buffers.push(buffer); + } + fn classify_record(&mut self, record: Record) -> Result, Error> { self.classify_calls += 1; if record.record().content_type == ContentType::Alert { self.dropped_alerts += 1; + self.push_buffer(record.into_buffer()); return Ok(None); } Ok(Some(record)) @@ -482,4 +529,77 @@ mod tests { ); assert_eq!(incoming.first().record().sequence.epoch, 1); } + + #[test] + fn receive_records_reuse_pooled_buffers() { + let mut handler = TestHandler::default(); + let mut buffer = handler.pop_buffer(); + buffer.resize(2048, 0xAA); + let allocation = buffer.as_ptr(); + handler.push_buffer(buffer); + + let packet = build_record(ContentType::ApplicationData, 1, 1, &[0x11; 1500]); + for _ in 0..10_000 { + let incoming = Incoming::parse_packet(&packet, &mut handler, None) + .expect("parse application data") + .expect("packet contains a record"); + assert_eq!(incoming.first().buffer().as_ptr(), allocation); + assert_eq!(incoming.first().buffer(), packet); + for record in incoming.into_records() { + handler.push_buffer(record.into_buffer()); + } + } + } + + #[test] + fn discarded_records_return_pooled_buffers() { + let mut handler = TestHandler::default(); + let mut buffer = handler.pop_buffer(); + buffer.resize(2048, 0xAA); + let allocation = buffer.as_ptr(); + handler.push_buffer(buffer); + + let packet = build_record(ContentType::ApplicationData, 0, 1, &[0x11; 1500]); + let incoming = Incoming::parse_packet(&packet, &mut handler, None) + .expect("invalid plaintext application data is discarded"); + assert!(incoming.is_none()); + let reused = handler.pop_buffer(); + assert!(reused.is_empty()); + assert_eq!(reused.as_ptr(), allocation); + } + + #[test] + fn truncated_datagram_recycles_all_pooled_records() { + let mut handler = TestHandler::default(); + let mut packet = build_record(ContentType::ApplicationData, 1, 1, &[0x11; 1500]); + packet.push(0xFF); + + assert!(Incoming::parse_packet(&packet, &mut handler, None).is_err()); + assert_eq!(handler.classify_calls, 0); + assert_eq!( + handler.buffers_returned, handler.buffers_acquired, + "discarding a malformed datagram must recycle its accepted records" + ); + } + + #[test] + fn oversized_datagram_recycles_all_pooled_records() { + let mut handler = TestHandler::default(); + let mut packet = Vec::new(); + for sequence in 0..9 { + packet.extend_from_slice(&build_record( + ContentType::ApplicationData, + 1, + sequence, + &[0x11; 32], + )); + } + + assert!(Incoming::parse_packet(&packet, &mut handler, None).is_err()); + assert_eq!(handler.classify_calls, 0); + assert_eq!( + handler.buffers_returned, handler.buffers_acquired, + "discarding a datagram with too many records must recycle every record" + ); + } } diff --git a/src/dtls12/server.rs b/src/dtls12/server.rs index 1e0a765..515a928 100644 --- a/src/dtls12/server.rs +++ b/src/dtls12/server.rs @@ -158,6 +158,8 @@ impl Server { engine.set_client(false); let cookie_secret: [u8; 32] = engine.rng.random(); + let extension_data = engine.pop_buffer(); + let defragment_buffer = engine.pop_buffer(); Server { state: State::AwaitClientHello, @@ -165,13 +167,13 @@ impl Server { random: None, session_id: None, cookie_secret, - extension_data: Buf::new(), + extension_data, negotiated_srtp_profile: None, client_supported_groups: None, client_signature_algorithms: None, client_random: None, client_certificates: Vec::with_capacity(3), - defragment_buffer: Buf::new(), + defragment_buffer, captured_session_hash: None, psk_valid: None, last_now: now, @@ -825,13 +827,12 @@ impl State { let client_pub = &server.defragment_buffer[public_key_range]; // Compute shared secret - let mut buf = server.engine.pop_buffer(); + let buf = server.engine.pop_buffer(); server .engine .crypto_context_mut() - .compute_shared_secret(client_pub, &mut buf) + .compute_shared_secret(client_pub, buf) .map_err(Error::CryptoError)?; - server.engine.push_buffer(buf); } // Capture session hash for EMS now (up to ClientKeyExchange) @@ -1252,10 +1253,10 @@ fn handshake_create_server_key_exchange( match key_exchange_algorithm { KeyExchangeAlgorithm::EECDH => { let (curve_type, named_group) = (CurveType::NamedCurve, named_group); - let mut kx_buf = engine.pop_buffer(); + let kx_buf = engine.pop_buffer(); let pubkey = engine .crypto_context_mut() - .init_ecdh_server(named_group, &mut kx_buf) + .init_ecdh_server(named_group, kx_buf) .map_err(Error::CryptoError)?; trace!( @@ -1274,8 +1275,6 @@ fn handshake_create_server_key_exchange( signed_data.push(pubkey.len() as u8); signed_data.extend_from_slice(pubkey); - engine.push_buffer(kx_buf); - let mut signature = engine.pop_buffer(); trace!("SKE signature hash: {:?}", hash_alg); diff --git a/src/dtls13/client.rs b/src/dtls13/client.rs index 47fdb4f..0407ece 100644 --- a/src/dtls13/client.rs +++ b/src/dtls13/client.rs @@ -137,16 +137,18 @@ pub(crate) enum LocalEvent { impl Client { pub(crate) fn new_with_engine(mut engine: Engine, now: Instant) -> Client { engine.set_client(true); + let extension_data = engine.pop_buffer(); + let defragment_buffer = engine.pop_buffer(); Client { state: State::SendClientHello, engine, random: None, session_id: None, - extension_data: Buf::new(), + extension_data, negotiated_srtp_profile: None, server_certificates: Vec::with_capacity(3), - defragment_buffer: Buf::new(), + defragment_buffer, client_auth_requested: false, cert_request_context: None, saved_cookie: None, @@ -182,16 +184,18 @@ impl Client { // Inject transcript + sequence state from the hybrid CH that was // already sent on the wire by ClientPending. engine.inject_hybrid_client_hello(&hybrid.transcript_bytes); + let extension_data = engine.pop_buffer(); + let defragment_buffer = engine.pop_buffer(); let mut client = Client { state: State::AwaitServerHello, engine, random: Some(hybrid.random), session_id: None, - extension_data: Buf::new(), + extension_data, negotiated_srtp_profile: None, server_certificates: Vec::with_capacity(3), - defragment_buffer: Buf::new(), + defragment_buffer, client_auth_requested: false, cert_request_context: None, saved_cookie: None, @@ -500,7 +504,7 @@ impl State { ExtensionType::Cookie => { let ext_data = ext.extension_data(&client.defragment_buffer); parse_cookie_extension(ext_data).map_err(InternalError::from)?; - let mut cookie = Buf::new(); + let mut cookie = client.engine.pop_buffer(); cookie.extend_from_slice(ext_data); client.saved_cookie = Some(cookie); } @@ -677,19 +681,15 @@ impl State { client.handshake_secret = Some(handshake_secret); client.engine.push_buffer(shared_secret); - // Save traffic secrets for Finished verification and client flight - let mut s_hs_copy = Buf::new(); - s_hs_copy.extend_from_slice(&s_hs_traffic); - client.server_hs_traffic_secret = Some(s_hs_copy); - let mut c_hs_copy = Buf::new(); - c_hs_copy.extend_from_slice(&c_hs_traffic); - client.client_hs_traffic_secret = Some(c_hs_copy); - // Install handshake keys (recv for server messages, send installed later) client .engine .install_handshake_keys(&c_hs_traffic, &s_hs_traffic)?; + // Retain the derived buffers for Finished verification and client flight. + client.server_hs_traffic_secret = Some(s_hs_traffic); + client.client_hs_traffic_secret = Some(c_hs_traffic); + // Enable peer encryption for server's epoch 2 messages client.engine.enable_peer_encryption()?; @@ -757,7 +757,7 @@ impl State { let cr_range = range.clone(); drop(maybe); let cr_data = &client.defragment_buffer[cr_range.clone()]; - let context = parse_certificate_request(cr_data, cr_range.start)?; + let context = parse_certificate_request(cr_data, cr_range.start, &mut client.engine)?; if let Some(ctx) = context { client.cert_request_context = Some(ctx); } @@ -813,7 +813,7 @@ impl State { for (i, cert_data) in cert_ranges.iter().enumerate() { trace!("Certificate #{} size: {} bytes", i + 1, cert_data.len()); - let mut buf = Buf::new(); + let mut buf = client.engine.pop_buffer(); buf.extend_from_slice(cert_data); client.server_certificates.push(buf); } @@ -1490,7 +1490,11 @@ pub(crate) fn verify_scheme_curve(scheme: SignatureScheme, cert_der: &[u8]) -> R /// /// Extracts the certificate_request_context and parses extensions including /// certificate_authorities. Returns the context if non-empty. -fn parse_certificate_request(cr_data: &[u8], base_offset: usize) -> Result, Error> { +fn parse_certificate_request( + cr_data: &[u8], + base_offset: usize, + engine: &mut Engine, +) -> Result, Error> { if cr_data.is_empty() { return Ok(None); } @@ -1505,7 +1509,7 @@ fn parse_certificate_request(cr_data: &[u8], base_offset: usize) -> Result= self.peer_handshake_seq_no && handshake.header.msg_type != MessageType::KeyUpdate { + self.recycle_incoming(incoming); return Err(Error::RenegotiationAttempt); } @@ -459,7 +467,10 @@ impl Engine { .try_push((seq.epoch as u64, seq.sequence_number)); } } - self.queue_rx[index] = incoming; + let replaced = std::mem::replace(&mut self.queue_rx[index], incoming); + self.recycle_incoming(replaced); + } else { + self.recycle_incoming(incoming); } } } @@ -481,12 +492,19 @@ impl Engine { // Duplicate - silently drop. For encrypted records (epoch >= 2) the replay // window filters most duplicates, but undecrypted ciphertext records can // reach here before enable_peer_encryption is called. + self.recycle_incoming(incoming); } } Ok(()) } + fn recycle_incoming(&mut self, incoming: Incoming) { + for record in incoming.into_records() { + self.push_buffer(record.into_buffer()); + } + } + pub fn handle_timeout(&mut self, now: Instant) -> Result<(), Error> { if self.connect_timeout == Timeout::Unarmed { debug!( @@ -1647,6 +1665,9 @@ impl Engine { ) -> Result<(Buf, Buf, Buf), Error> { // Call derive_early_secret first (needs &mut self) before borrowing hmac let early_secret = self.derive_early_secret()?; + let mut handshake_secret = self.buffers_free.pop(); + let mut c_hs_traffic = self.buffers_free.pop(); + let mut s_hs_traffic = self.buffers_free.pop(); let hash = self.hash_algorithm(); let hash_len = hash.output_len(); @@ -1667,7 +1688,6 @@ impl Engine { .map_err(Error::CryptoError)?; // handshake_secret = HKDF-Extract(derived, shared_secret) - let mut handshake_secret = Buf::new(); prf_hkdf::hkdf_extract(hmac, hash, &derived, shared_secret, &mut handshake_secret) .map_err(Error::CryptoError)?; @@ -1676,7 +1696,6 @@ impl Engine { self.transcript_hash(&mut transcript_hash); // client_handshake_traffic_secret - let mut c_hs_traffic = Buf::new(); prf_hkdf::hkdf_expand_label_dtls13( hmac, hash, @@ -1689,7 +1708,6 @@ impl Engine { .map_err(Error::CryptoError)?; // server_handshake_traffic_secret - let mut s_hs_traffic = Buf::new(); prf_hkdf::hkdf_expand_label_dtls13( hmac, hash, @@ -1701,6 +1719,7 @@ impl Engine { ) .map_err(Error::CryptoError)?; + self.buffers_free.push(early_secret); Ok((c_hs_traffic, s_hs_traffic, handshake_secret)) } @@ -1732,6 +1751,9 @@ impl Engine { &mut self, handshake_secret: &[u8], ) -> Result<(Buf, Buf), Error> { + let mut exp_master = self.buffers_free.pop(); + let mut c_ap_traffic = self.buffers_free.pop(); + let mut s_ap_traffic = self.buffers_free.pop(); let hash = self.hash_algorithm(); let hash_len = hash.output_len(); let hmac = self.hmac(); @@ -1762,7 +1784,6 @@ impl Engine { self.transcript_hash(&mut transcript_hash); // exporter_master_secret = Derive-Secret(master_secret, "exp master", transcript_hash) - let mut exp_master = Buf::new(); prf_hkdf::hkdf_expand_label_dtls13( hmac, hash, @@ -1775,7 +1796,6 @@ impl Engine { .map_err(Error::CryptoError)?; // client_application_traffic_secret_0 - let mut c_ap_traffic = Buf::new(); prf_hkdf::hkdf_expand_label_dtls13( hmac, hash, @@ -1788,7 +1808,6 @@ impl Engine { .map_err(Error::CryptoError)?; // server_application_traffic_secret_0 - let mut s_ap_traffic = Buf::new(); prf_hkdf::hkdf_expand_label_dtls13( hmac, hash, @@ -1967,7 +1986,9 @@ impl Engine { } /// Derive epoch keys (cipher + IV + sn_key) from a traffic secret. - fn derive_epoch_keys(&self, traffic_secret: &Buf) -> Result { + fn derive_epoch_keys(&mut self, traffic_secret: &Buf) -> Result { + let mut sn_key = self.buffers_free.pop(); + let mut secret = self.buffers_free.pop(); let hash = self.hash_algorithm(); let suite = self.suite_provider(); let hmac = self.hmac(); @@ -1999,7 +2020,6 @@ impl Engine { .map_err(Error::CryptoError)?; // sn_key = HKDF-Expand-Label(secret, "sn", "", key_length) - let mut sn_key = Buf::new(); prf_hkdf::hkdf_expand_label_dtls13( hmac, hash, @@ -2016,7 +2036,6 @@ impl Engine { let mut iv = [0u8; 12]; iv.copy_from_slice(&iv_buf); - let mut secret = Buf::new(); secret.extend_from_slice(traffic_secret); Ok(EpochKeys { @@ -2321,6 +2340,14 @@ fn reconstruct_sequence(partial: u64, expected: u64, bits: u32) -> u64 { // ========================================================================= impl RecordHandler for Engine { + fn pop_buffer(&mut self) -> Buf { + Engine::pop_buffer(self) + } + + fn push_buffer(&mut self, buffer: Buf) { + Engine::push_buffer(self, buffer); + } + fn classify_record(&mut self, record: Record) -> Result, Error> { if let Some(cn_seq) = self.close_notify_sequence { if record.record().sequence > cn_seq { @@ -2553,9 +2580,20 @@ mod tests { Engine::new(config, cert) } - struct PassthroughRecordHandler; + #[derive(Default)] + struct PassthroughRecordHandler { + buffers: BufferPool, + } impl RecordHandler for PassthroughRecordHandler { + fn pop_buffer(&mut self) -> Buf { + self.buffers.pop() + } + + fn push_buffer(&mut self, buffer: Buf) { + self.buffers.push(buffer); + } + fn classify_record(&mut self, record: Record) -> Result, Error> { Ok(Some(record)) } @@ -2641,10 +2679,40 @@ mod tests { packet } + #[test] + #[cfg(feature = "rcgen")] + fn duplicate_datagram_recycles_its_pooled_buffer() { + let mut engine = test_engine(); + let packet = encrypted_application_data_record(1, b"buffered ciphertext"); + engine.parse_packet(&packet).expect("queue first datagram"); + let mut output = [0u8; 2048]; + assert!(matches!( + engine.poll_output(&mut output, Instant::now()), + Output::Timeout(_) + )); + + let mut buffer = engine.pop_buffer(); + buffer.resize(2048, 0xAA); + engine.push_buffer(buffer); + + engine + .parse_packet(&packet) + .expect("discard duplicate datagram"); + assert!(matches!( + engine.poll_output(&mut output, Instant::now()), + Output::Timeout(_) + )); + assert_eq!(engine.queue_rx.len(), 1); + assert!( + engine.pop_buffer().into_vec().capacity() >= 2048, + "discarding a duplicate must preserve its reusable receive allocation" + ); + } + fn parsed_key_update(seq: u16) -> Incoming { Incoming::parse_packet( &encrypted_key_update_record(seq), - &mut PassthroughRecordHandler, + &mut PassthroughRecordHandler::default(), Some(Dtls13CipherSuite::AES_128_GCM_SHA256), ) .expect("parse key update packet") @@ -2656,7 +2724,7 @@ mod tests { packet.extend_from_slice(&encrypted_application_data_record(app_seq, b"app-data")); Incoming::parse_packet( &packet, - &mut PassthroughRecordHandler, + &mut PassthroughRecordHandler::default(), Some(Dtls13CipherSuite::AES_128_GCM_SHA256), ) .expect("parse coalesced packet") @@ -2747,7 +2815,7 @@ mod tests { // Pre-fill the pool with a buffer that has allocated capacity. // BufferPool::push clears contents but retains the allocation. - let mut marked = Buf::new(); + let mut marked = engine.pop_buffer(); marked.extend_from_slice(&[0xAA; 256]); engine.buffers_free.push(marked); @@ -2836,11 +2904,12 @@ mod tests { #[cfg(feature = "rcgen")] fn malformed_ack_record_number_vector_is_ignored() { let mut engine = test_engine(); + let fragment = engine.pop_buffer(); engine.flight_saved_records.push(Entry { content_type: ContentType::Handshake, epoch: 2, send_seq: 7, - fragment: Buf::new(), + fragment, acked: false, }); diff --git a/src/dtls13/incoming.rs b/src/dtls13/incoming.rs index 6dfeb6f..8940dac 100644 --- a/src/dtls13/incoming.rs +++ b/src/dtls13/incoming.rs @@ -1,8 +1,8 @@ +use std::fmt; use std::ops::Deref; use std::sync::atomic::{AtomicBool, Ordering}; use arrayvec::ArrayVec; -use std::fmt; use crate::buffer::{Buf, TmpBuf}; use crate::dtls13::message::{ContentType, Dtls13CipherSuite, Dtls13Record, Handshake, Sequence}; @@ -67,12 +67,36 @@ pub struct Records { impl Records { pub fn parse( - mut packet: &[u8], + packet: &[u8], decrypt: &mut dyn RecordHandler, cs: Option, ) -> Result { let mut parsed_records: ArrayVec = ArrayVec::new(); + if let Err(error) = Self::parse_records(packet, decrypt, cs, &mut parsed_records) { + for record in parsed_records { + decrypt.push_buffer(record.into_buffer()); + } + return Err(error); + } + + 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 }) + } + fn parse_records( + mut packet: &[u8], + decrypt: &mut dyn RecordHandler, + cs: Option, + parsed_records: &mut ArrayVec, + ) -> Result<(), InternalError> { // Find record boundaries and copy each record ONCE from the packet while !packet.is_empty() { let record_end = if Dtls13Record::is_ciphertext_header(packet[0]) { @@ -131,7 +155,8 @@ impl Records { let record_slice = &packet[..record_end]; let record = Record::parse(record_slice, decrypt, cs)?; if let Some(record) = record { - if parsed_records.try_push(record).is_err() { + if let Err(error) = parsed_records.try_push(record) { + decrypt.push_buffer(error.element().into_buffer()); return Err(InternalError::too_many_records()); } } else { @@ -141,16 +166,7 @@ impl Records { 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 }) + Ok(()) } } @@ -178,9 +194,23 @@ impl Record { cs: Option, ) -> Result, InternalError> { // ONLY COPY: UDP packet slice -> pooled buffer - let mut buffer = Buf::new(); + let mut buffer = decrypt.pop_buffer(); buffer.extend_from_slice(record_slice); + match Self::parse_buffer(&mut buffer, decrypt, cs) { + Ok(Some(parsed)) => Ok(Some(Record { buffer, parsed })), + result => { + decrypt.push_buffer(buffer); + result.map(|_| None) + } + } + } + + fn parse_buffer( + buffer: &mut Buf, + decrypt: &mut dyn RecordHandler, + cs: Option, + ) -> Result>, InternalError> { let is_ciphertext = Dtls13Record::is_ciphertext_header(buffer[0]); // Decrypt record number in-place before parsing (RFC 9147 Section 4.2.3) @@ -210,7 +240,7 @@ impl Record { } } - let parsed = match ParsedRecord::parse(&buffer, cs) { + let parsed = match ParsedRecord::parse(buffer, cs) { Ok(p) => p, Err(e) => { trace!("Discarding record: parse failed: {}", e); @@ -218,20 +248,19 @@ impl Record { } }; let parsed = Box::new(parsed); - let record = Record { buffer, parsed }; // Plaintext records (epoch 0) are not encrypted if !is_ciphertext || !decrypt.is_peer_encryption_enabled() { - return Ok(Some(record)); + return Ok(Some(parsed)); } // Resolve the full epoch from the 2-bit value in the unified header - let epoch_bits = record.record().sequence.epoch as u8; + let epoch_bits = parsed.record.sequence.epoch as u8; let full_epoch = decrypt.resolve_epoch(epoch_bits); // Resolve the full sequence number from the (now decrypted) partial value - let seq_bits = record.record().sequence.sequence_number; - let s_flag = record_slice[0] & 0b0000_1000 != 0; + let seq_bits = parsed.record.sequence.sequence_number; + let s_flag = buffer[0] & 0b0000_1000 != 0; let full_seq = decrypt.resolve_sequence(full_epoch, seq_bits, s_flag); let full_sequence = Sequence { @@ -246,20 +275,17 @@ impl Record { // Save the raw header bytes for AAD before mutating the buffer. // Max unified header without CID: flags(1) + seq(2) + length(2) = 5 bytes. - let header_end = record.record().fragment_range.start; + let header_end = parsed.record.fragment_range.start; // Reject protected records whose encrypted fragment is shorter than // the per-suite minimum — they cannot hold a valid ciphertext + tag, // so decryption would necessarily fail. Catching it here keeps the // cipher impls' bounds-checking from being the only line of defence. - if record.buffer.len() - header_end < decrypt.min_protected_fragment_len() { + if 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]); - - // Extract the buffer for decryption - let mut buffer = record.buffer; + header_buf[..header_end].copy_from_slice(&buffer[..header_end]); // The encrypted part starts right after the unified header. let ciphertext = &mut buffer[header_end..]; @@ -302,12 +328,12 @@ impl Record { length: content_len as u16, fragment_range: header_end..(header_end + content_len), }, - &buffer, + buffer, cs, ); let parsed = Box::new(parsed); - Ok(Some(Record { buffer, parsed })) + Ok(Some(parsed)) } pub fn record(&self) -> &Dtls13Record { @@ -397,9 +423,13 @@ impl ParsedRecord { /// Trait abstracting record parsing-time handling for incoming records. /// /// This decouples the record parser from the full `Engine`, allowing the parse loop -/// to decrypt records, classify control records, and queue only the records that +/// to acquire buffers, decrypt records, classify control records, and queue only the records that /// should survive into `Incoming`. pub trait RecordHandler { + /// Acquire an empty buffer whose allocation can be reused for a record. + fn pop_buffer(&mut self) -> Buf; + /// Recycle a buffer when parsing fails or silently discards a record. + fn push_buffer(&mut self, buffer: Buf); fn classify_record(&mut self, record: Record) -> Result, Error>; fn is_peer_encryption_enabled(&self) -> bool; fn resolve_epoch(&self, epoch_bits: u8) -> u16; @@ -526,19 +556,35 @@ impl std::panic::UnwindSafe for Incoming {} #[cfg(test)] mod tests { + use crate::buffer::BufferPool; + use super::*; #[derive(Default)] struct TestHandler { + buffers: BufferPool, + buffers_acquired: usize, + buffers_returned: usize, classify_calls: usize, dropped_acks: usize, } impl RecordHandler for TestHandler { + fn pop_buffer(&mut self) -> Buf { + self.buffers_acquired += 1; + self.buffers.pop() + } + + fn push_buffer(&mut self, buffer: Buf) { + self.buffers_returned += 1; + self.buffers.push(buffer); + } + fn classify_record(&mut self, record: Record) -> Result, Error> { self.classify_calls += 1; if record.record().content_type == ContentType::Ack { self.dropped_acks += 1; + self.push_buffer(record.into_buffer()); return Ok(None); } Ok(Some(record)) @@ -630,4 +676,72 @@ mod tests { ); assert_eq!(incoming.first().record().sequence.epoch, 2); } + + #[test] + fn receive_records_reuse_pooled_buffers() { + let mut handler = TestHandler::default(); + let mut buffer = handler.pop_buffer(); + buffer.resize(2048, 0xAA); + let allocation = buffer.as_ptr(); + handler.push_buffer(buffer); + + let packet = build_ciphertext_record(2, 1, &[0x11; 1500]); + for _ in 0..10_000 { + let incoming = Incoming::parse_packet(&packet, &mut handler, None) + .expect("parse application data") + .expect("packet contains a record"); + assert_eq!(incoming.first().buffer().as_ptr(), allocation); + assert_eq!(incoming.first().buffer(), packet); + for record in incoming.into_records() { + handler.push_buffer(record.into_buffer()); + } + } + } + + #[test] + fn discarded_records_return_pooled_buffers() { + let mut handler = TestHandler::default(); + let mut buffer = handler.pop_buffer(); + buffer.resize(2048, 0xAA); + let allocation = buffer.as_ptr(); + handler.push_buffer(buffer); + + let packet = build_plaintext_record(ContentType::ApplicationData, 1, &[0x11; 1500]); + let incoming = Incoming::parse_packet(&packet, &mut handler, None) + .expect("invalid plaintext application data is discarded"); + assert!(incoming.is_none()); + let reused = handler.pop_buffer(); + assert!(reused.is_empty()); + assert_eq!(reused.as_ptr(), allocation); + } + + #[test] + fn truncated_datagram_recycles_all_pooled_records() { + let mut handler = TestHandler::default(); + let mut packet = build_ciphertext_record(2, 1, &[0x11; 1500]); + packet.push(0xFF); + + assert!(Incoming::parse_packet(&packet, &mut handler, None).is_err()); + assert_eq!(handler.classify_calls, 0); + assert_eq!( + handler.buffers_returned, handler.buffers_acquired, + "discarding a malformed datagram must recycle its accepted records" + ); + } + + #[test] + fn oversized_datagram_recycles_all_pooled_records() { + let mut handler = TestHandler::default(); + let mut packet = Vec::new(); + for sequence in 0..17 { + packet.extend_from_slice(&build_ciphertext_record(2, sequence, &[0x11; 32])); + } + + assert!(Incoming::parse_packet(&packet, &mut handler, None).is_err()); + assert_eq!(handler.classify_calls, 0); + assert_eq!( + handler.buffers_returned, handler.buffers_acquired, + "discarding a datagram with too many records must recycle every record" + ); + } } diff --git a/src/dtls13/server.rs b/src/dtls13/server.rs index ebe3ed5..eee5c9c 100644 --- a/src/dtls13/server.rs +++ b/src/dtls13/server.rs @@ -185,16 +185,18 @@ impl Server { pub fn new_with_engine(mut engine: Engine, now: Instant, auto_mode: bool) -> Server { let cookie_secret = engine.random_arr(); + let extension_data = engine.pop_buffer(); + let defragment_buffer = engine.pop_buffer(); Server { state: State::AwaitClientHello, engine, random: None, client_session_id: None, - extension_data: Buf::new(), + extension_data, negotiated_srtp_profile: None, client_certificates: Vec::with_capacity(3), - defragment_buffer: Buf::new(), + defragment_buffer, last_now: now, local_events: VecDeque::new(), queued_data: Vec::new(), @@ -839,19 +841,15 @@ impl State { server.handshake_secret = Some(handshake_secret); server.engine.push_buffer(shared_secret); - // Save traffic secrets - let mut s_hs_copy = Buf::new(); - s_hs_copy.extend_from_slice(&s_hs_traffic); - server.server_hs_traffic_secret = Some(s_hs_copy); - let mut c_hs_copy = Buf::new(); - c_hs_copy.extend_from_slice(&c_hs_traffic); - server.client_hs_traffic_secret = Some(c_hs_copy); - // Install handshake keys server .engine .install_handshake_keys(&c_hs_traffic, &s_hs_traffic)?; + // Retain the derived buffers for Finished verification. + server.server_hs_traffic_secret = Some(s_hs_traffic); + server.client_hs_traffic_secret = Some(c_hs_traffic); + Ok(Self::SendEncryptedExtensions) } @@ -1003,7 +1001,7 @@ impl State { i + 1, cert_data.len() ); - let mut buf = Buf::new(); + let mut buf = server.engine.pop_buffer(); buf.extend_from_slice(cert_data); server.client_certificates.push(buf); }