This is an automated email from the ASF dual-hosted git repository.
tisonkun pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datasketches-rust.git
The following commit(s) were added to refs/heads/main by this push:
new e02c693 fix: reject invalid REQ compactor states on deserialize (#221)
e02c693 is described below
commit e02c693e627d348b5772709b635c354c43e449c6
Author: Cestercian <[email protected]>
AuthorDate: Tue Aug 25 08:15:17 2026 -0700
fix: reject invalid REQ compactor states on deserialize (#221)
Co-authored-by: Cursor Agent <[email protected]>
Co-authored-by: tison <[email protected]>
---
datasketches/src/req/compactor.rs | 121 ++++++++++--------
datasketches/src/req/mod.rs | 8 ++
datasketches/src/req/serialization.rs | 101 +++++++++++++++
datasketches/src/req/sketch.rs | 93 ++++++++++----
datasketches/tests/serde_tests/req.rs | 223 ++++++++++++++++++++++++++--------
5 files changed, 429 insertions(+), 117 deletions(-)
diff --git a/datasketches/src/req/compactor.rs
b/datasketches/src/req/compactor.rs
index ea5c4b0..43d8b23 100644
--- a/datasketches/src/req/compactor.rs
+++ b/datasketches/src/req/compactor.rs
@@ -22,11 +22,24 @@
use super::MIN_K;
use super::RankAccuracy;
+use super::nearest_even_section_size;
use super::value::ReqValue;
use crate::error::Error;
-fn nearest_even(value: f32) -> u32 {
- ((value / 2.0).round() as u32) << 1
+fn validate_deserialized_items<T: ReqValue>(items: &[T], sorted: bool) ->
Result<(), Error> {
+ if items.iter().any(ReqValue::is_nan) {
+ return Err(Error::deserial("REQ compactor contains a NaN item"));
+ }
+ if sorted
+ && !items
+ .windows(2)
+ .all(|items| items[0].total_cmp(&items[1]).is_le())
+ {
+ return Err(Error::deserial(
+ "REQ compactor claims to be sorted but its items are not",
+ ));
+ }
+ Ok(())
}
/// A compactor maintains items at a specific level of the REQ sketch.
@@ -71,8 +84,8 @@ where
/// * `rank_accuracy` - Rank accuracy configuration
pub(super) fn new(lg_weight: u8, k: u16, rank_accuracy: RankAccuracy) ->
Self {
let section_size_raw = k as f32;
- let section_size = nearest_even(section_size_raw);
- let num_sections = 3u8;
+ let section_size = nearest_even_section_size(section_size_raw);
+ let num_sections = super::INITIAL_SECTIONS_PER_COMPACTOR;
let nominal: usize = (2 * section_size * num_sections as u32) as usize;
@@ -289,23 +302,38 @@ where
1u64 << self.lg_weight
}
+ /// Returns the minimum stream length implied by this compactor's state.
+ ///
+ /// Each compaction removes at least two items of this level's weight.
Merging
+ /// states with bitwise OR cannot make the result larger than their sum,
so the
+ /// same lower bound holds for both updated and merged sketches.
+ pub(super) fn minimum_stream_length(&self) -> Option<u64> {
+ self.state.checked_mul(2)?.checked_mul(self.weight())
+ }
+
// Private helper methods
fn ensure_enough_sections(&mut self) -> bool {
- let ssr = self.section_size_raw / std::f32::consts::SQRT_2;
- let ne = nearest_even(ssr);
-
- if self.num_sections <= 64
- && self.state >= (1u64 << (self.num_sections - 1))
- && ne >= u32::from(MIN_K)
- {
- self.section_size_raw = ssr;
- self.section_size = ne;
- self.num_sections <<= 1; // Double the sections
- true
- } else {
- false
+ let Some(threshold) = self
+ .num_sections
+ .checked_sub(1)
+ .and_then(|shift| 1u64.checked_shl(u32::from(shift)))
+ else {
+ return false;
+ };
+ let Some(num_sections) = self.num_sections.checked_mul(2) else {
+ return false;
+ };
+ let section_size_raw = self.section_size_raw /
std::f32::consts::SQRT_2;
+ let section_size = nearest_even_section_size(section_size_raw);
+
+ if self.state >= threshold && section_size >= u32::from(MIN_K) {
+ self.section_size_raw = section_size_raw;
+ self.section_size = section_size;
+ self.num_sections = num_sections;
+ return true;
}
+ false
}
#[inline(always)]
@@ -362,8 +390,10 @@ where
/// Deserialize a compactor (preamble + items) from the byte cursor.
pub(super) fn deserialize(
cursor: &mut crate::codec::SketchSlice<'_>,
+ k: u16,
+ expected_lg_weight: u8,
rank_accuracy: super::RankAccuracy,
- is_level_zero_sorted: bool,
+ sorted: bool,
) -> Result<Self, crate::error::Error> {
use crate::codec::assert::insufficient_data;
let state = cursor
@@ -385,22 +415,14 @@ where
.read_u32_le()
.map_err(insufficient_data("compactor.num_items"))?;
- // Validate the wire-controlled fields before they feed capacity/weight
- // arithmetic. A legitimate compactor always satisfies these bounds
- // (`section_size` derives from k ≤ MAX_K and only shrinks;
`lg_weight` is the
- // level index), so rejecting anything else keeps `nominal_capacity`
and
- // `weight` from overflowing on crafted input.
- if !(0.0..=super::MAX_K as f32).contains(§ion_size_raw) {
- return Err(Error::invalid_argument(format!(
- "REQ compactor section_size {section_size_raw} out of range"
- )));
- }
- // `weight()` computes `1u64 << lg_weight`, which overflows once
lg_weight ≥ 64.
- if lg_weight >= 64 {
- return Err(Error::invalid_argument(format!(
- "REQ compactor lg_weight {lg_weight} exceeds maximum"
- )));
- }
+ super::serialization::validate_compactor_state(
+ k,
+ expected_lg_weight,
+ state,
+ section_size_raw,
+ lg_weight,
+ num_sections,
+ )?;
// Don't trust `num_items` for the allocation: a malformed length
could request
// a multi-gigabyte reservation before the per-item reads below fail.
The buffer
@@ -411,6 +433,7 @@ where
for _ in 0..num_items {
items.push(T::deserialize_value(cursor)?);
}
+ validate_deserialized_items(&items, sorted)?;
Ok(Compactor::from_serialized_state(
lg_weight,
@@ -418,7 +441,7 @@ where
num_sections,
state,
items,
- is_level_zero_sorted,
+ sorted,
rank_accuracy,
))
}
@@ -435,14 +458,15 @@ where
rank_accuracy: super::RankAccuracy,
items: Vec<T>,
is_sorted: bool,
- ) -> Self {
+ ) -> Result<Self, Error> {
+ validate_deserialized_items(&items, is_sorted)?;
let mut c = Self::new(0, k, rank_accuracy);
for item in items {
c.append(item);
}
// append() may have flipped is_sorted off; restore the wire flag
verbatim.
c.is_sorted = is_sorted;
- c
+ Ok(c)
}
/// Reconstruct a Compactor from deserialized state.
@@ -465,7 +489,7 @@ where
is_sorted,
state,
scratch_buffer: Vec::new(),
- section_size: nearest_even(section_size_raw),
+ section_size: nearest_even_section_size(section_size_raw),
num_sections,
lg_weight,
rank_accuracy,
@@ -527,15 +551,15 @@ mod tests {
}
#[test]
- fn test_nearest_even() {
- assert_eq!(nearest_even(0.0), 0); // 0/2=0, round(0)=0, 0<<1=0
- assert_eq!(nearest_even(1.0), 2); // 1/2=0.5, round(0.5)=1, 1<<1=2
- assert_eq!(nearest_even(2.0), 2); // 2/2=1, round(1)=1, 1<<1=2
- assert_eq!(nearest_even(3.0), 4); // 3/2=1.5, round(1.5)=2, 2<<1=4
- assert_eq!(nearest_even(4.0), 4); // 4/2=2, round(2)=2, 2<<1=4
- assert_eq!(nearest_even(4.6), 4); // 4.6/2=2.3, round(2.3)=2, 2<<1=4
- assert_eq!(nearest_even(5.6), 6); // 5.6/2=2.8, round(2.8)=3, 3<<1=6
- assert_eq!(nearest_even(13.0), 14); // 13/2=6.5, round(6.5)=7, 7<<1=14
+ fn test_nearest_even_section_size() {
+ assert_eq!(nearest_even_section_size(0.0), 0); // 0/2=0, round(0)=0,
0<<1=0
+ assert_eq!(nearest_even_section_size(1.0), 2); // 1/2=0.5,
round(0.5)=1, 1<<1=2
+ assert_eq!(nearest_even_section_size(2.0), 2); // 2/2=1, round(1)=1,
1<<1=2
+ assert_eq!(nearest_even_section_size(3.0), 4); // 3/2=1.5,
round(1.5)=2, 2<<1=4
+ assert_eq!(nearest_even_section_size(4.0), 4); // 4/2=2, round(2)=2,
2<<1=4
+ assert_eq!(nearest_even_section_size(4.6), 4); // 4.6/2=2.3,
round(2.3)=2, 2<<1=4
+ assert_eq!(nearest_even_section_size(5.6), 6); // 5.6/2=2.8,
round(2.8)=3, 3<<1=6
+ assert_eq!(nearest_even_section_size(13.0), 14); // 13/2=6.5,
round(6.5)=7, 7<<1=14
}
#[test]
@@ -571,7 +595,8 @@ mod tests {
let raw = bytes.into_bytes();
let mut cursor = SketchSlice::new(&raw);
- let c2 = Compactor::<f32>::deserialize(&mut cursor,
RankAccuracy::HighRank, true).unwrap();
+ let c2 = Compactor::<f32>::deserialize(&mut cursor, 12, 0,
RankAccuracy::HighRank, true)
+ .unwrap();
assert_eq!(c.num_items(), c2.num_items());
assert_eq!(c.lg_weight(), c2.lg_weight());
diff --git a/datasketches/src/req/mod.rs b/datasketches/src/req/mod.rs
index c17b441..a2b58f1 100644
--- a/datasketches/src/req/mod.rs
+++ b/datasketches/src/req/mod.rs
@@ -31,6 +31,14 @@ mod sorted_view;
mod union;
mod value;
+/// Number of sections in a newly created compactor. The section count and size
+/// determine its capacity and compaction range; the count doubles as its
state grows.
+const INITIAL_SECTIONS_PER_COMPACTOR: u8 = 3;
+
+fn nearest_even_section_size(value: f32) -> u32 {
+ ((value / 2.0).round() as u32) << 1
+}
+
pub use self::iter::ReqSketchIterator;
pub use self::sketch::ReqSketch;
pub use self::sketch::ReqSketchBuilder;
diff --git a/datasketches/src/req/serialization.rs
b/datasketches/src/req/serialization.rs
index 19061d8..e3ad461 100644
--- a/datasketches/src/req/serialization.rs
+++ b/datasketches/src/req/serialization.rs
@@ -17,6 +17,9 @@
//! REQ sketch wire format — constants and helpers shared by sketch +
compactor serdes.
+use super::INITIAL_SECTIONS_PER_COMPACTOR;
+use super::MIN_K;
+use super::nearest_even_section_size;
use crate::codec::assert::ensure_preamble_longs_in;
use crate::codec::assert::ensure_serial_version_is;
use crate::error::Error;
@@ -33,6 +36,104 @@ pub(super) const FLAG_IS_HIGH_RANK: u8 = 1 << 3;
pub(super) const FLAG_RAW_ITEMS: u8 = 1 << 4;
pub(super) const FLAG_IS_LEVEL_ZERO_SORTED: u8 = 1 << 5;
+fn section_growth_threshold(num_sections: u8) -> Option<u64> {
+ num_sections
+ .checked_sub(1)
+ .and_then(|shift| 1u64.checked_shl(u32::from(shift)))
+}
+
+fn reachable_section_doublings(state: u64, num_sections: u8) -> Option<u32> {
+ let mut sections = INITIAL_SECTIONS_PER_COMPACTOR;
+ let mut doublings = 0;
+ while sections < num_sections {
+ if state < section_growth_threshold(sections)? {
+ return None;
+ }
+ sections = sections.checked_mul(2)?;
+ doublings += 1;
+ }
+ (sections == num_sections).then_some(doublings)
+}
+
+fn compatible_next_section_sizes(section_size_raw: f32) -> [Option<f32>; 2] {
+ let single_precision = section_size_raw / std::f32::consts::SQRT_2;
+ let widened = (f64::from(section_size_raw) / std::f64::consts::SQRT_2) as
f32;
+ [
+ (nearest_even_section_size(single_precision) >= u32::from(MIN_K))
+ .then_some(single_precision),
+ (nearest_even_section_size(section_size_raw) > u32::from(MIN_K)
+ && nearest_even_section_size(widened) >= u32::from(MIN_K))
+ .then_some(widened),
+ ]
+}
+
+fn has_reachable_section_configuration(
+ k: u16,
+ state: u64,
+ section_size_raw: f32,
+ num_sections: u8,
+) -> bool {
+ let Some(doublings) = reachable_section_doublings(state, num_sections)
else {
+ return false;
+ };
+
+ // The wire format exposes the running floating-point section size. C++
+ // updates it in single precision, while Java divides in double precision
+ // before narrowing to f32. Accept either result at every step so a sketch
+ // can be deserialized and continued in another implementation.
+ let mut candidates = vec![k as f32];
+ for _ in 0..doublings {
+ let mut next = Vec::with_capacity(candidates.len() * 2);
+ for candidate in candidates {
+ for value in compatible_next_section_sizes(candidate)
+ .into_iter()
+ .flatten()
+ {
+ if !next.contains(&value) {
+ next.push(value);
+ }
+ }
+ }
+ candidates = next;
+ }
+ if !candidates
+ .iter()
+ .any(|candidate| candidate.to_bits() == section_size_raw.to_bits())
+ {
+ return false;
+ }
+
+ // If every compatible producer would already have doubled the section
+ // count at this state, the serialized configuration is stale.
+ section_growth_threshold(num_sections).is_none_or(|threshold| {
+ state < threshold
+ || compatible_next_section_sizes(section_size_raw)
+ .iter()
+ .any(Option::is_none)
+ })
+}
+
+pub(super) fn validate_compactor_state(
+ k: u16,
+ expected_lg_weight: u8,
+ state: u64,
+ section_size_raw: f32,
+ lg_weight: u8,
+ num_sections: u8,
+) -> Result<(), Error> {
+ if lg_weight != expected_lg_weight {
+ return Err(Error::deserial(format!(
+ "REQ compactor lg_weight {lg_weight} does not match level
{expected_lg_weight}"
+ )));
+ }
+ if !has_reachable_section_configuration(k, state, section_size_raw,
num_sections) {
+ return Err(Error::deserial(format!(
+ "REQ compactor section configuration is not reachable (k={k},
state={state}, section_size_raw={section_size_raw},
num_sections={num_sections})"
+ )));
+ }
+ Ok(())
+}
+
pub(super) fn check_serial_version(actual: u8) -> Result<(), Error> {
ensure_serial_version_is(SERIAL_VERSION, actual)
}
diff --git a/datasketches/src/req/sketch.rs b/datasketches/src/req/sketch.rs
index d1f1245..5af3ac7 100644
--- a/datasketches/src/req/sketch.rs
+++ b/datasketches/src/req/sketch.rs
@@ -165,7 +165,7 @@ impl<T: ReqValue> ReqSketch<T> {
self.n += 1;
self.num_retained += 1;
- if self.num_retained == self.max_nom_size {
+ if self.num_retained >= self.max_nom_size {
self.compress();
}
}
@@ -359,10 +359,8 @@ impl<T: ReqValue> ReqSketch<T> {
}
const FIXED_RSE_FACTOR: f64 = 0.084;
- const INIT_NUM_SECTIONS: u8 = 3;
-
fn relative_rse_factor() -> f64 {
- (0.0512 / Self::INIT_NUM_SECTIONS as f64).sqrt()
+ (0.0512 / super::INITIAL_SECTIONS_PER_COMPACTOR as f64).sqrt()
}
fn compute_rank_lower_bound(
@@ -411,7 +409,7 @@ impl<T: ReqValue> ReqSketch<T> {
n: u64,
hra: bool,
) -> bool {
- let base_cap = k as u64 * Self::INIT_NUM_SECTIONS as u64;
+ let base_cap = k as u64 * super::INITIAL_SECTIONS_PER_COMPACTOR as u64;
if num_levels == 1 || n <= base_cap {
return true;
}
@@ -563,6 +561,11 @@ impl<T: ReqValue> ReqSketch<T> {
/// Deserialize a sketch from bytes produced by [`Self::serialize`] or by
the
/// C++/Java reference implementations.
+ ///
+ /// # Errors
+ ///
+ /// Returns an error if the input is truncated or contains an inconsistent
+ /// REQ serialized state.
pub fn deserialize(bytes: &[u8]) -> Result<Self, Error> {
use super::compactor::Compactor;
use super::serialization::FLAG_IS_EMPTY;
@@ -606,19 +609,17 @@ impl<T: ReqValue> ReqSketch<T> {
RankAccuracy::LowRank
};
if !(MIN_K..=MAX_K).contains(&k) || k % 2 != 0 {
- return Err(Error::invalid_argument(format!(
- "k {k} is not a valid REQ k value"
- )));
+ return Err(Error::deserial(format!("k {k} is not a valid REQ k
value")));
}
if is_empty {
if num_levels != 0 {
- return Err(Error::invalid_argument(format!(
+ return Err(Error::deserial(format!(
"empty REQ sketch must have 0 levels, got {num_levels}"
)));
}
if num_raw_items != 0 {
- return Err(Error::invalid_argument(format!(
+ return Err(Error::deserial(format!(
"empty REQ sketch must have 0 raw items, got
{num_raw_items}"
)));
}
@@ -626,24 +627,29 @@ impl<T: ReqValue> ReqSketch<T> {
}
if num_levels == 0 {
- return Err(Error::invalid_argument(
+ return Err(Error::deserial(
"non-empty REQ sketch must have at least one level",
));
}
+ if num_levels > 64 {
+ return Err(Error::deserial(
+ "REQ sketch cannot have more than 64 levels",
+ ));
+ }
if raw_items {
if num_levels != 1 {
- return Err(Error::invalid_argument(format!(
+ return Err(Error::deserial(format!(
"raw-items REQ sketch must have exactly 1 level, got
{num_levels}"
)));
}
if num_raw_items == 0 || num_raw_items as u64 >
RAW_ITEMS_THRESHOLD {
- return Err(Error::invalid_argument(format!(
+ return Err(Error::deserial(format!(
"raw-items REQ sketch must contain
1..={RAW_ITEMS_THRESHOLD} items, got {num_raw_items}"
)));
}
} else if num_raw_items != 0 {
- return Err(Error::invalid_argument(format!(
+ return Err(Error::deserial(format!(
"non-raw REQ sketch must have 0 raw items, got {num_raw_items}"
)));
}
@@ -656,6 +662,16 @@ impl<T: ReqValue> ReqSketch<T> {
n = cursor.read_u64_le().map_err(insufficient_data("n"))?;
min_item = Some(T::deserialize_value(&mut cursor)?);
max_item = Some(T::deserialize_value(&mut cursor)?);
+ let min = min_item.as_ref().unwrap();
+ let max = max_item.as_ref().unwrap();
+ if min.is_nan() || max.is_nan() {
+ return Err(Error::deserial("REQ sketch min or max item is
NaN"));
+ }
+ if min.total_cmp(max).is_gt() {
+ return Err(Error::deserial(
+ "REQ sketch min item is greater than max item",
+ ));
+ }
}
let mut compactors: Vec<Compactor<T>> = Vec::with_capacity(num_levels
as usize);
@@ -667,12 +683,13 @@ impl<T: ReqValue> ReqSketch<T> {
items.push(T::deserialize_value(&mut cursor)?);
}
let c =
- Compactor::<T>::raw_items_compactor(k, rank_accuracy, items,
is_level_zero_sorted);
+ Compactor::<T>::raw_items_compactor(k, rank_accuracy, items,
is_level_zero_sorted)?;
compactors.push(c);
} else {
for i in 0..num_levels {
- let level_sorted = if i == 0 { is_level_zero_sorted } else {
true };
- let c = Compactor::<T>::deserialize(&mut cursor,
rank_accuracy, level_sorted)?;
+ let level_sorted = i > 0 || is_level_zero_sorted;
+ let c =
+ Compactor::<T>::deserialize(&mut cursor, k, i,
rank_accuracy, level_sorted)?;
compactors.push(c);
}
}
@@ -699,18 +716,52 @@ impl<T: ReqValue> ReqSketch<T> {
}
if n == 0 || min_item.is_none() || max_item.is_none() {
- return Err(Error::invalid_argument(
- "non-empty REQ sketch contains no items",
+ return Err(Error::deserial("non-empty REQ sketch contains no
items"));
+ }
+
+ let expected_raw_items = num_levels == 1 && n <= RAW_ITEMS_THRESHOLD;
+ if raw_items != expected_raw_items {
+ return Err(Error::deserial(
+ "REQ sketch RAW_ITEMS flag is inconsistent with num_levels and
n",
));
}
+ if compactors.iter().any(|compactor| {
+ compactor
+ .minimum_stream_length()
+ .is_none_or(|minimum_n| minimum_n > n)
+ }) {
+ return Err(Error::deserial("REQ compactor state exceeds stream
length"));
+ }
+
+ let (retained_count, nominal_capacity, weighted_count) = compactors
+ .iter()
+ .try_fold(
+ (0u32, 0u32, 0u64),
+ |(retained, capacity, weighted), compactor| {
+ Some((
+ retained.checked_add(compactor.num_items())?,
+ capacity.checked_add(compactor.nominal_capacity())?,
+ weighted.checked_add(
+ (compactor.num_items() as
u64).checked_mul(compactor.weight())?,
+ )?,
+ ))
+ },
+ )
+ .ok_or_else(|| Error::deserial("REQ compactor totals overflow"))?;
+ if weighted_count != n {
+ return Err(Error::deserial(format!(
+ "REQ retained weighted count {weighted_count} does not match n
{n}"
+ )));
+ }
+
let mut sketch = ReqSketch::try_new(k, rank_accuracy)?;
sketch.n = n;
sketch.min_item = min_item;
sketch.max_item = max_item;
sketch.compactors = compactors;
- sketch.update_max_nom_size();
- sketch.update_num_retained();
+ sketch.max_nom_size = nominal_capacity;
+ sketch.num_retained = retained_count;
Ok(sketch)
}
diff --git a/datasketches/tests/serde_tests/req.rs
b/datasketches/tests/serde_tests/req.rs
index 8539819..445fbf3 100644
--- a/datasketches/tests/serde_tests/req.rs
+++ b/datasketches/tests/serde_tests/req.rs
@@ -53,9 +53,9 @@ where
#[test]
fn round_trip_f64_matrix() {
- for &k in &[4u16, 12, 1024] {
+ for &k in &[4u16, 6, 10, 12, 1024] {
for &ra in &[RankAccuracy::HighRank, RankAccuracy::LowRank] {
- for &n in &[0u64, 1, 4, 5, 100, 10_000] {
+ for &n in &[0u64, 1, 4, 5, 100, 1_250, 2_562, 10_000, 100_000] {
round_trip_one::<f64>(k, ra, n, |i| i as f64);
}
}
@@ -258,67 +258,194 @@ fn merge_preserves_order_across_serde_round_trip() {
// ---------- Deserialize hardening: malformed compactor fields ----------
//
-// A non-empty, non-raw, single-level sketch carries a full 20-byte compactor
-// preamble whose `section_size_raw`, `lg_weight`, and `num_items` fields are
read
-// straight off the wire. Without bounds checks these crafted values either
panic
-// (arithmetic overflow) or trigger an unbounded allocation in
`Compactor::deserialize`.
-
-/// Builds a non-empty, non-raw, single-level (`num_levels = 1`) REQ sketch
image
-/// with a fully specified compactor preamble, so an individual field can be
made
-/// malformed in isolation. With valid inputs the result deserializes
successfully
-/// (see `single_level_image_is_valid_baseline`).
-fn single_level_image(
- section_size_raw: f32,
- lg_weight: u8,
- num_sections: u8,
- num_items: u32,
- items: &[f32],
-) -> Vec<u8> {
- // Preamble (8 bytes): preamble_ints = 2 (EXACT, since num_levels == 1),
- // serial_version = 1, family = 17 (REQ), flags = 8 (IS_HIGH_RANK: not
empty,
- // not raw), k = 12 (u16 LE), num_levels = 1, num_raw_items = 0.
- let mut b = vec![2u8, 1, 17, 8, 12, 0, 1, 0];
- // Compactor preamble (20 bytes).
- b.extend_from_slice(&0u64.to_le_bytes()); // state
- b.extend_from_slice(§ion_size_raw.to_le_bytes());
- b.push(lg_weight);
- b.push(num_sections);
- b.extend_from_slice(&0u16.to_le_bytes()); // padding
- b.extend_from_slice(&num_items.to_le_bytes());
- for &item in items {
- b.extend_from_slice(&item.to_le_bytes());
+const EXACT_COMPACTOR_OFFSET: usize = 8;
+const ESTIMATION_COMPACTOR_OFFSET: usize = 24;
+const STATE_OFFSET: usize = 0;
+const SECTION_SIZE_RAW_OFFSET: usize = 8;
+const LG_WEIGHT_OFFSET: usize = 12;
+const NUM_SECTIONS_OFFSET: usize = 13;
+const NUM_ITEMS_OFFSET: usize = 16;
+const FLAG_LEVEL_ZERO_SORTED: u8 = 1 << 5;
+
+fn exact_image(k: u16, items: &[f32]) -> Vec<u8> {
+ let mut bytes = vec![2u8, 1, 17, 8];
+ bytes.extend_from_slice(&k.to_le_bytes());
+ bytes.extend_from_slice(&[1, 0]);
+ bytes.extend_from_slice(&0u64.to_le_bytes());
+ bytes.extend_from_slice(&(k as f32).to_le_bytes());
+ bytes.extend_from_slice(&[0, 3, 0, 0]);
+ bytes.extend_from_slice(&(items.len() as u32).to_le_bytes());
+ for item in items {
+ bytes.extend_from_slice(&item.to_le_bytes());
}
- b
+ bytes
+}
+
+fn estimation_image(k: u16, n: u64) -> Vec<u8> {
+ let mut sketch = ReqSketch::<f32>::try_new(k,
RankAccuracy::HighRank).unwrap();
+ for item in 1..=n {
+ sketch.update(item as f32);
+ }
+ let bytes = sketch.serialize();
+ assert!(bytes[6] > 1);
+ bytes
+}
+
+fn read_u64(bytes: &[u8], offset: usize) -> u64 {
+ u64::from_le_bytes(bytes[offset..offset + 8].try_into().unwrap())
+}
+
+fn assert_invalid_data(bytes: &[u8]) {
+ let error = ReqSketch::<f32>::deserialize(bytes).unwrap_err();
+ assert_eq!(error.kind(), ErrorKind::InvalidData);
}
#[test]
-fn single_level_image_is_valid_baseline() {
- // Control: the builder with well-formed fields round-trips, so the
malformed
- // variants below isolate exactly one bad field.
- let bytes = single_level_image(12.0, 0, 3, 1, &[1.0]);
+fn canonical_exact_image_is_valid() {
+ let bytes = exact_image(12, &[1.0, 2.0, 3.0, 4.0, 5.0]);
assert_that!(ReqSketch::<f32>::deserialize(&bytes), ok(anything()));
}
#[test]
-fn deserialize_rejects_out_of_range_section_size() {
- // A garbage section_size_raw drives the `nominal_capacity` arithmetic to
overflow.
- let bytes = single_level_image(1e30, 0, 3, 1, &[1.0]);
- assert_that!(ReqSketch::<f32>::deserialize(&bytes), err(anything()));
+fn deserialize_rejects_issue_218_states() {
+ let mut zero_sections = exact_image(12, &[1.0, 2.0, 3.0, 4.0, 5.0]);
+ zero_sections[EXACT_COMPACTOR_OFFSET + NUM_SECTIONS_OFFSET] = 0;
+ assert_invalid_data(&zero_sections);
+
+ let mut wrong_weight = exact_image(12, &[1.0, 2.0, 3.0, 4.0, 5.0]);
+ wrong_weight[EXACT_COMPACTOR_OFFSET + LG_WEIGHT_OFFSET] = 63;
+ assert_invalid_data(&wrong_weight);
+}
+
+#[test]
+fn deserialize_rejects_inconsistent_weighted_count() {
+ let mut bytes = estimation_image(12, 1_000);
+ bytes[8..16].copy_from_slice(&1_001u64.to_le_bytes());
+ assert_invalid_data(&bytes);
+}
+
+#[test]
+fn deserialize_rejects_unreachable_section_configuration() {
+ let mut invalid_raw = exact_image(12, &[1.0, 2.0, 3.0, 4.0, 5.0]);
+ invalid_raw[EXACT_COMPACTOR_OFFSET + SECTION_SIZE_RAW_OFFSET
+ ..EXACT_COMPACTOR_OFFSET + SECTION_SIZE_RAW_OFFSET + 4]
+ .copy_from_slice(&0.0f32.to_le_bytes());
+ assert_invalid_data(&invalid_raw);
+
+ let mut invalid_sections = exact_image(12, &[1.0, 2.0, 3.0, 4.0, 5.0]);
+ invalid_sections[EXACT_COMPACTOR_OFFSET + NUM_SECTIONS_OFFSET] = 6;
+ assert_invalid_data(&invalid_sections);
+}
+
+#[test]
+fn deserialize_accepts_java_minimum_section_schedule() {
+ let mut bytes = estimation_image(6, 1_250);
+ let compactor = ESTIMATION_COMPACTOR_OFFSET;
+ assert_eq!(read_u64(&bytes, compactor + STATE_OFFSET), 32);
+
+ let java_raw = (6.0 / std::f64::consts::SQRT_2) as f32;
+ bytes[compactor + SECTION_SIZE_RAW_OFFSET..compactor +
SECTION_SIZE_RAW_OFFSET + 4]
+ .copy_from_slice(&java_raw.to_le_bytes());
+ bytes[compactor + NUM_SECTIONS_OFFSET] = 6;
+
+ let mut sketch = ReqSketch::<f32>::deserialize(&bytes).unwrap();
+ for item in 1_251..=2_500 {
+ sketch.update(item as f32);
+ }
+ let continued = sketch.serialize();
+ assert_that!(ReqSketch::<f32>::deserialize(&continued), ok(anything()));
+}
+
+#[test]
+fn deserialize_rejects_capacity_changing_float_drift() {
+ let mut bytes = estimation_image(10, 2_562);
+ let raw_offset = ESTIMATION_COMPACTOR_OFFSET + SECTION_SIZE_RAW_OFFSET;
+ let raw_bits = u32::from_le_bytes(bytes[raw_offset..raw_offset +
4].try_into().unwrap());
+ assert_eq!(read_u64(&bytes, ESTIMATION_COMPACTOR_OFFSET), 32);
+ assert_eq!(f32::from_bits(raw_bits), 5.0);
+ assert_eq!(bytes[ESTIMATION_COMPACTOR_OFFSET + NUM_SECTIONS_OFFSET], 12);
+
+ // One ULP below 5.0 rounds to a section size of 4 rather than 6.
+ bytes[raw_offset..raw_offset + 4].copy_from_slice(&(raw_bits -
1).to_le_bytes());
+ assert_invalid_data(&bytes);
+}
+
+#[test]
+fn deserialize_rejects_state_inconsistent_with_stream_length() {
+ let mut bytes = estimation_image(12, 1_000);
+ let compactor = ESTIMATION_COMPACTOR_OFFSET;
+ let state = 501u64;
+ bytes[compactor + STATE_OFFSET..compactor + STATE_OFFSET + 8]
+ .copy_from_slice(&state.to_le_bytes());
+ let mut raw = 12.0f32;
+ for _ in 0..2 {
+ raw /= std::f32::consts::SQRT_2;
+ }
+ bytes[compactor + SECTION_SIZE_RAW_OFFSET..compactor +
SECTION_SIZE_RAW_OFFSET + 4]
+ .copy_from_slice(&raw.to_le_bytes());
+ bytes[compactor + NUM_SECTIONS_OFFSET] = 12;
+ assert_invalid_data(&bytes);
+}
+
+#[test]
+fn deserialize_rejects_complementary_states_that_overflow_on_merge() {
+ let items: Vec<f32> = (0..192).map(|item| item as f32).collect();
+ let mut raw = 12.0f32;
+ for _ in 0..4 {
+ raw /= std::f32::consts::SQRT_2;
+ }
+
+ let states = [0xAAAA_AAAA_AAAA_AAAAu64, 0x5555_5555_5555_5555u64];
+ assert_eq!(states[0] | states[1], u64::MAX);
+ for state in states {
+ let mut bytes = exact_image(12, &items);
+ let compactor = EXACT_COMPACTOR_OFFSET;
+ bytes[compactor + STATE_OFFSET..compactor + STATE_OFFSET + 8]
+ .copy_from_slice(&state.to_le_bytes());
+ bytes[compactor + SECTION_SIZE_RAW_OFFSET..compactor +
SECTION_SIZE_RAW_OFFSET + 4]
+ .copy_from_slice(&raw.to_le_bytes());
+ bytes[compactor + NUM_SECTIONS_OFFSET] = 48;
+ assert_invalid_data(&bytes);
+ }
+}
+
+#[test]
+fn deserialize_rejects_false_sorted_claim_and_nan() {
+ let mut unsorted = exact_image(12, &[3.0, 4.0, 5.0, 1.0, 2.0]);
+ unsorted[3] |= FLAG_LEVEL_ZERO_SORTED;
+ assert_invalid_data(&unsorted);
+
+ let nan = exact_image(12, &[1.0, 2.0, f32::NAN, 4.0, 5.0]);
+ assert_invalid_data(&nan);
+}
+
+#[test]
+fn deserialize_rejects_invalid_extrema_and_raw_nan() {
+ let mut nan_min = estimation_image(12, 1_000);
+ nan_min[16..20].copy_from_slice(&f32::NAN.to_le_bytes());
+ assert_invalid_data(&nan_min);
+
+ let mut reversed = estimation_image(12, 1_000);
+ reversed[16..20].copy_from_slice(&2.0f32.to_le_bytes());
+ reversed[20..24].copy_from_slice(&1.0f32.to_le_bytes());
+ assert_invalid_data(&reversed);
+
+ let mut raw_nan = vec![2u8, 1, 17, 8 | 16, 12, 0, 1, 1];
+ raw_nan.extend_from_slice(&f32::NAN.to_le_bytes());
+ assert_invalid_data(&raw_nan);
}
#[test]
-fn deserialize_rejects_oversized_lg_weight() {
- // lg_weight >= 64 makes the per-item weight `1u64 << lg_weight` overflow.
- let bytes = single_level_image(12.0, 64, 3, 1, &[1.0]);
- assert_that!(ReqSketch::<f32>::deserialize(&bytes), err(anything()));
+fn deserialize_rejects_noncanonical_exact_mode() {
+ assert_invalid_data(&exact_image(12, &[1.0]));
}
#[test]
fn deserialize_rejects_oversized_compactor_num_items() {
- // num_items claims billions of items while only one is supplied:
deserialize
- // must fail gracefully without attempting a multi-gigabyte allocation.
- let bytes = single_level_image(12.0, 0, 3, u32::MAX, &[1.0]);
- assert_that!(ReqSketch::<f32>::deserialize(&bytes), err(anything()));
+ let mut bytes = exact_image(12, &[1.0, 2.0, 3.0, 4.0, 5.0]);
+ let offset = EXACT_COMPACTOR_OFFSET + NUM_ITEMS_OFFSET;
+ bytes[offset..offset + 4].copy_from_slice(&u32::MAX.to_le_bytes());
+ assert_invalid_data(&bytes);
}
// ---------- Cross-language compatibility ----------
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]