Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -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
Expand Down
4 changes: 4 additions & 0 deletions src/auto.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
111 changes: 88 additions & 23 deletions src/buffer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}
Expand All @@ -47,79 +51,89 @@ impl fmt::Debug for BufferPool {

/// Growable buffer wrapper used throughout dimpl for efficient memory management.
///
/// This is a newtype around `Vec<u8>` that provides convenient access to byte buffers
/// This wrapper around `Vec<u8>` 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<u8>);
pub struct Buf {
data: Vec<u8>,
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<u8>`.
pub fn into_vec(mut self) -> Vec<u8> {
std::mem::take(&mut self.0)
pub fn into_vec(self) -> Vec<u8> {
self.data
}
}

// aws-lc-rs AEAD operations require Extend<&u8> for appending authentication tags
impl<'a> Extend<&'a u8> for Buf {
fn extend<T: IntoIterator<Item = &'a u8>>(&mut self, iter: T) {
self.0.extend(iter.into_iter().copied());
self.data.extend(iter.into_iter().copied());
}
}

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()
}
}

Expand All @@ -131,7 +145,10 @@ pub trait ToBuf {

impl ToBuf for Vec<u8> {
fn to_buf(self) -> Buf {
Buf(self)
Buf {
data: self,
from_pool: false,
}
}
}

Expand Down Expand Up @@ -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);
}
}
20 changes: 12 additions & 8 deletions src/dtls12/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -92,18 +92,20 @@ 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,
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,
Expand Down Expand Up @@ -156,18 +158,20 @@ 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,
engine,
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,
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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);

Expand Down
20 changes: 10 additions & 10 deletions src/dtls12/context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(())
}
Expand All @@ -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);
Expand All @@ -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
Expand All @@ -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()
}

Expand All @@ -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
Expand All @@ -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(())
}
Expand Down
Loading
Loading