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(&section_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(&section_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]

Reply via email to