From 5394d720c6f83b6465220d8d7c6ce5e0ed7c9fdf Mon Sep 17 00:00:00 2001 From: Kacy Fortner Date: Thu, 12 Feb 2026 11:21:47 -0500 Subject: [PATCH] harden: guard u32 length casts in persistence format replace bare `len() as u32` casts with checked conversions: - add format::write_len() helper that uses u32::try_from() and returns an error if the length exceeds u32::MAX - update all collection length writes in aof.rs and snapshot.rs to use the safe helper (10 call sites in AOF, 5 in snapshot) - guard encrypted ciphertext length with try_from in snapshot writer - also guard format::write_bytes() the same way while >4B items is unrealistic for in-memory collections, this prevents silent data corruption if it ever happened, and the error path is essentially free since the check always succeeds in practice. --- crates/ember-persistence/src/aof.rs | 20 ++++++++++---------- crates/ember-persistence/src/format.rs | 21 ++++++++++++++++++++- crates/ember-persistence/src/snapshot.rs | 20 +++++++++++++------- 3 files changed, 43 insertions(+), 18 deletions(-) diff --git a/crates/ember-persistence/src/aof.rs b/crates/ember-persistence/src/aof.rs index 794e93f0..d25cb9fd 100644 --- a/crates/ember-persistence/src/aof.rs +++ b/crates/ember-persistence/src/aof.rs @@ -216,7 +216,7 @@ impl AofRecord { AofRecord::LPush { key, values } => { format::write_u8(&mut buf, TAG_LPUSH)?; format::write_bytes(&mut buf, key.as_bytes())?; - format::write_u32(&mut buf, values.len() as u32)?; + format::write_len(&mut buf, values.len())?; for v in values { format::write_bytes(&mut buf, v)?; } @@ -224,7 +224,7 @@ impl AofRecord { AofRecord::RPush { key, values } => { format::write_u8(&mut buf, TAG_RPUSH)?; format::write_bytes(&mut buf, key.as_bytes())?; - format::write_u32(&mut buf, values.len() as u32)?; + format::write_len(&mut buf, values.len())?; for v in values { format::write_bytes(&mut buf, v)?; } @@ -240,7 +240,7 @@ impl AofRecord { AofRecord::ZAdd { key, members } => { format::write_u8(&mut buf, TAG_ZADD)?; format::write_bytes(&mut buf, key.as_bytes())?; - format::write_u32(&mut buf, members.len() as u32)?; + format::write_len(&mut buf, members.len())?; for (score, member) in members { format::write_f64(&mut buf, *score)?; format::write_bytes(&mut buf, member.as_bytes())?; @@ -249,7 +249,7 @@ impl AofRecord { AofRecord::ZRem { key, members } => { format::write_u8(&mut buf, TAG_ZREM)?; format::write_bytes(&mut buf, key.as_bytes())?; - format::write_u32(&mut buf, members.len() as u32)?; + format::write_len(&mut buf, members.len())?; for member in members { format::write_bytes(&mut buf, member.as_bytes())?; } @@ -274,7 +274,7 @@ impl AofRecord { AofRecord::HSet { key, fields } => { format::write_u8(&mut buf, TAG_HSET)?; format::write_bytes(&mut buf, key.as_bytes())?; - format::write_u32(&mut buf, fields.len() as u32)?; + format::write_len(&mut buf, fields.len())?; for (field, value) in fields { format::write_bytes(&mut buf, field.as_bytes())?; format::write_bytes(&mut buf, value)?; @@ -283,7 +283,7 @@ impl AofRecord { AofRecord::HDel { key, fields } => { format::write_u8(&mut buf, TAG_HDEL)?; format::write_bytes(&mut buf, key.as_bytes())?; - format::write_u32(&mut buf, fields.len() as u32)?; + format::write_len(&mut buf, fields.len())?; for field in fields { format::write_bytes(&mut buf, field.as_bytes())?; } @@ -297,7 +297,7 @@ impl AofRecord { AofRecord::SAdd { key, members } => { format::write_u8(&mut buf, TAG_SADD)?; format::write_bytes(&mut buf, key.as_bytes())?; - format::write_u32(&mut buf, members.len() as u32)?; + format::write_len(&mut buf, members.len())?; for member in members { format::write_bytes(&mut buf, member.as_bytes())?; } @@ -305,7 +305,7 @@ impl AofRecord { AofRecord::SRem { key, members } => { format::write_u8(&mut buf, TAG_SREM)?; format::write_bytes(&mut buf, key.as_bytes())?; - format::write_u32(&mut buf, members.len() as u32)?; + format::write_len(&mut buf, members.len())?; for member in members { format::write_bytes(&mut buf, member.as_bytes())?; } @@ -343,7 +343,7 @@ impl AofRecord { format::write_u8(&mut buf, TAG_VADD)?; format::write_bytes(&mut buf, key.as_bytes())?; format::write_bytes(&mut buf, element.as_bytes())?; - format::write_u32(&mut buf, vector.len() as u32)?; + format::write_len(&mut buf, vector.len())?; for &v in vector { format::write_f32(&mut buf, v)?; } @@ -653,7 +653,7 @@ impl AofWriter { if let Some(ref key) = self.encryption_key { let (nonce, ciphertext) = crate::encryption::encrypt_record(key, &payload)?; self.writer.write_all(&nonce)?; - format::write_u32(&mut self.writer, ciphertext.len() as u32)?; + format::write_len(&mut self.writer, ciphertext.len())?; self.writer.write_all(&ciphertext)?; return Ok(()); } diff --git a/crates/ember-persistence/src/format.rs b/crates/ember-persistence/src/format.rs index 6d3fc115..7c946f99 100644 --- a/crates/ember-persistence/src/format.rs +++ b/crates/ember-persistence/src/format.rs @@ -94,9 +94,28 @@ pub fn write_f64(w: &mut impl Write, val: f64) -> io::Result<()> { w.write_all(&val.to_le_bytes()) } +/// Writes a collection length as u32, returning an error if it exceeds `u32::MAX`. +pub fn write_len(w: &mut impl Write, len: usize) -> io::Result<()> { + let len = u32::try_from(len).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("collection length {len} exceeds u32::MAX"), + ) + })?; + write_u32(w, len) +} + /// Writes a length-prefixed byte slice: `[len: u32][data]`. +/// +/// Returns an error if the data length exceeds `u32::MAX`. pub fn write_bytes(w: &mut impl Write, data: &[u8]) -> io::Result<()> { - write_u32(w, data.len() as u32)?; + let len = u32::try_from(data.len()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("data length {} exceeds u32::MAX", data.len()), + ) + })?; + write_u32(w, len)?; w.write_all(data) } diff --git a/crates/ember-persistence/src/snapshot.rs b/crates/ember-persistence/src/snapshot.rs index 9c840421..30b88ecf 100644 --- a/crates/ember-persistence/src/snapshot.rs +++ b/crates/ember-persistence/src/snapshot.rs @@ -280,14 +280,14 @@ impl SnapshotWriter { } SnapValue::List(deque) => { format::write_u8(&mut buf, TYPE_LIST)?; - format::write_u32(&mut buf, deque.len() as u32)?; + format::write_len(&mut buf, deque.len())?; for item in deque { format::write_bytes(&mut buf, item)?; } } SnapValue::SortedSet(members) => { format::write_u8(&mut buf, TYPE_SORTED_SET)?; - format::write_u32(&mut buf, members.len() as u32)?; + format::write_len(&mut buf, members.len())?; for (score, member) in members { format::write_f64(&mut buf, *score)?; format::write_bytes(&mut buf, member.as_bytes())?; @@ -295,7 +295,7 @@ impl SnapshotWriter { } SnapValue::Hash(map) => { format::write_u8(&mut buf, TYPE_HASH)?; - format::write_u32(&mut buf, map.len() as u32)?; + format::write_len(&mut buf, map.len())?; for (field, value) in map { format::write_bytes(&mut buf, field.as_bytes())?; format::write_bytes(&mut buf, value)?; @@ -303,7 +303,7 @@ impl SnapshotWriter { } SnapValue::Set(set) => { format::write_u8(&mut buf, TYPE_SET)?; - format::write_u32(&mut buf, set.len() as u32)?; + format::write_len(&mut buf, set.len())?; for member in set { format::write_bytes(&mut buf, member.as_bytes())?; } @@ -323,7 +323,7 @@ impl SnapshotWriter { format::write_u32(&mut buf, *connectivity)?; format::write_u32(&mut buf, *expansion_add)?; format::write_u32(&mut buf, *dim)?; - format::write_u32(&mut buf, elements.len() as u32)?; + format::write_len(&mut buf, elements.len())?; for (name, vector) in elements { format::write_bytes(&mut buf, name.as_bytes())?; for &v in vector { @@ -343,13 +343,19 @@ impl SnapshotWriter { #[cfg(feature = "encryption")] if let Some(ref key) = self.encryption_key { let (nonce, ciphertext) = crate::encryption::encrypt_record(key, &buf)?; + let ct_len = u32::try_from(ciphertext.len()).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "encrypted record exceeds u32::MAX bytes", + ) + })?; // footer CRC covers the encrypted envelope self.hasher.update(&nonce); - let ct_len_bytes = (ciphertext.len() as u32).to_le_bytes(); + let ct_len_bytes = ct_len.to_le_bytes(); self.hasher.update(&ct_len_bytes); self.hasher.update(&ciphertext); self.writer.write_all(&nonce)?; - format::write_u32(&mut self.writer, ciphertext.len() as u32)?; + format::write_u32(&mut self.writer, ct_len)?; self.writer.write_all(&ciphertext)?; self.count += 1; return Ok(());