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 6f3f09b  refactor(req): clarify value comparison contract (#243)
6f3f09b is described below

commit 6f3f09bb199484f31ee2599b41da6e29a99fb413
Author: tison <[email protected]>
AuthorDate: Sat Aug 29 10:02:31 2026 +0800

    refactor(req): clarify value comparison contract (#243)
---
 datasketches/src/req/compactor.rs                  |  14 +-
 datasketches/src/req/sketch.rs                     |  14 +-
 datasketches/src/req/sorted_view.rs                |  16 +-
 datasketches/src/req/value.rs                      | 252 +++++++++++++--------
 tests-integration/tests/req_test/query.rs          |  16 ++
 .../tests/req_test/sorted_view_api.rs              |  22 +-
 6 files changed, 223 insertions(+), 111 deletions(-)

diff --git a/datasketches/src/req/compactor.rs 
b/datasketches/src/req/compactor.rs
index da79181..07a6f38 100644
--- a/datasketches/src/req/compactor.rs
+++ b/datasketches/src/req/compactor.rs
@@ -35,7 +35,7 @@ fn normalized_sort_state<T: ReqValue>(items: &[T], 
claimed_sorted: bool) -> Resu
         if item.is_nan() {
             return Err(Error::deserial("REQ compactor contains a NaN item"));
         }
-        if sorted && previous.is_some_and(|previous| 
previous.total_cmp(item).is_gt()) {
+        if sorted && previous.is_some_and(|previous| 
previous.compare(item).is_gt()) {
             sorted = false;
         }
         previous = Some(item);
@@ -140,7 +140,7 @@ where
                 self.merge_sorted(&other.items);
             } else {
                 let mut other_items = other.items.clone();
-                other_items.sort_unstable_by(|a, b| a.total_cmp(b));
+                other_items.sort_unstable_by(|a, b| a.compare(b));
                 self.merge_sorted(&other_items);
             }
         }
@@ -158,15 +158,15 @@ where
     pub(super) fn count_below(&self, item: &T, inclusive: bool) -> usize {
         if self.is_sorted {
             if inclusive {
-                self.items.partition_point(|x| x.total_cmp(item).is_le())
+                self.items.partition_point(|x| x.compare(item).is_le())
             } else {
-                self.items.partition_point(|x| x.total_cmp(item).is_lt())
+                self.items.partition_point(|x| x.compare(item).is_lt())
             }
         } else {
             self.items
                 .iter()
                 .filter(|x| {
-                    let ord = x.total_cmp(item);
+                    let ord = x.compare(item);
                     if inclusive { ord.is_le() } else { ord.is_lt() }
                 })
                 .count()
@@ -201,7 +201,7 @@ where
 
         // Two-pointer merge into scratch buffer
         while i < a.len() && j < b.len() {
-            if a[i].total_cmp(&b[j]).is_le() {
+            if a[i].compare(&b[j]).is_le() {
                 self.scratch_buffer.push(a[i].clone());
                 i += 1;
             } else {
@@ -229,7 +229,7 @@ where
     pub(super) fn sort(&mut self) {
         if !self.is_sorted {
             // Use unstable sort for better performance (stable not needed for 
REQ sketch)
-            self.items.sort_unstable_by(|a, b| a.total_cmp(b));
+            self.items.sort_unstable_by(|a, b| a.compare(b));
             self.is_sorted = true;
         }
     }
diff --git a/datasketches/src/req/sketch.rs b/datasketches/src/req/sketch.rs
index 3ff0a53..16fe895 100644
--- a/datasketches/src/req/sketch.rs
+++ b/datasketches/src/req/sketch.rs
@@ -133,12 +133,12 @@ impl<T: ReqValue> ReqSketch<T> {
         }
         match &mut self.min_item {
             None => self.min_item = Some(item.clone()),
-            Some(cur) if item.total_cmp(cur).is_lt() => *cur = item.clone(),
+            Some(cur) if item.compare(cur).is_lt() => *cur = item.clone(),
             _ => {}
         }
         match &mut self.max_item {
             None => self.max_item = Some(item.clone()),
-            Some(cur) if item.total_cmp(cur).is_gt() => *cur = item.clone(),
+            Some(cur) if item.compare(cur).is_gt() => *cur = item.clone(),
             _ => {}
         }
 
@@ -303,14 +303,14 @@ impl<T: ReqValue> ReqSketch<T> {
         if let Some(m) = &other.min_item {
             match &self.min_item {
                 None => self.min_item = Some(m.clone()),
-                Some(cur) if m.total_cmp(cur).is_lt() => self.min_item = 
Some(m.clone()),
+                Some(cur) if m.compare(cur).is_lt() => self.min_item = 
Some(m.clone()),
                 _ => {}
             }
         }
         if let Some(m) = &other.max_item {
             match &self.max_item {
                 None => self.max_item = Some(m.clone()),
-                Some(cur) if m.total_cmp(cur).is_gt() => self.max_item = 
Some(m.clone()),
+                Some(cur) if m.compare(cur).is_gt() => self.max_item = 
Some(m.clone()),
                 _ => {}
             }
         }
@@ -649,7 +649,7 @@ impl<T: ReqValue> ReqSketch<T> {
             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() {
+            if min.compare(max).is_gt() {
                 return Err(Error::deserial(
                     "REQ sketch min item is greater than max item",
                 ));
@@ -685,10 +685,10 @@ impl<T: ReqValue> ReqSketch<T> {
                 let mut mn = first.clone();
                 let mut mx = first.clone();
                 for x in iter {
-                    if x.total_cmp(&mn).is_lt() {
+                    if x.compare(&mn).is_lt() {
                         mn = x.clone();
                     }
-                    if x.total_cmp(&mx).is_gt() {
+                    if x.compare(&mx).is_gt() {
                         mx = x.clone();
                     }
                 }
diff --git a/datasketches/src/req/sorted_view.rs 
b/datasketches/src/req/sorted_view.rs
index a8c1d98..5c2d5a1 100644
--- a/datasketches/src/req/sorted_view.rs
+++ b/datasketches/src/req/sorted_view.rs
@@ -60,7 +60,7 @@ where
         }
 
         // Sort by item value - use unstable sort for better performance
-        weighted_items.sort_unstable_by(|a, b| a.0.total_cmp(&b.0));
+        weighted_items.sort_unstable_by(|a, b| a.0.compare(&b.0));
 
         let mut items: Vec<T> = Vec::with_capacity(weighted_items.len());
         let mut cumulative_weights = Vec::with_capacity(weighted_items.len());
@@ -68,7 +68,7 @@ where
 
         for (item, weight) in weighted_items {
             if let Some(last) = items.last() {
-                if matches!(last.total_cmp(&item), std::cmp::Ordering::Equal) {
+                if matches!(last.compare(&item), std::cmp::Ordering::Equal) {
                     cumulative_weight += weight;
                     let last_idx = cumulative_weights.len() - 1;
                     cumulative_weights[last_idx] = cumulative_weight;
@@ -122,7 +122,7 @@ where
             SearchCriteria::Inclusive => {
                 // Find the last position where items[i] <= item
                 // partition_point finds first index where predicate is false
-                let pos = self.items.partition_point(|x| 
x.total_cmp(item).is_le());
+                let pos = self.items.partition_point(|x| 
x.compare(item).is_le());
                 if pos == 0 {
                     Ok(0.0)
                 } else {
@@ -131,7 +131,7 @@ where
             }
             SearchCriteria::Exclusive => {
                 // Find the last position where items[i] < item
-                let pos = self.items.partition_point(|x| 
x.total_cmp(item).is_lt());
+                let pos = self.items.partition_point(|x| 
x.compare(item).is_lt());
                 if pos == 0 {
                     Ok(0.0)
                 } else {
@@ -259,9 +259,11 @@ where
     // Private helper methods
 
     fn validate_split_points(&self, split_points: &[T]) -> Result<(), Error> {
-        // Check that split points are monotonically increasing
-        for i in 1..split_points.len() {
-            if split_points[i - 1].total_cmp(&split_points[i]).is_ge() {
+        for (i, split_point) in split_points.iter().enumerate() {
+            if split_point.is_nan() {
+                return Err(Error::invalid_argument("Split points must not be 
NaN"));
+            }
+            if i > 0 && split_points[i - 1].compare(split_point).is_ge() {
                 return Err(Error::invalid_argument(
                     "Split points must be unique and monotonically 
increasing".to_string(),
                 ));
diff --git a/datasketches/src/req/value.rs b/datasketches/src/req/value.rs
index b065ae3..a2c5ed4 100644
--- a/datasketches/src/req/value.rs
+++ b/datasketches/src/req/value.rs
@@ -18,23 +18,20 @@
 //! Trait for types storable in a [`ReqSketch`](crate::req::ReqSketch).
 
 use std::cmp::Ordering;
+use std::mem::size_of;
 
 use crate::codec::SketchBytes;
 use crate::codec::SketchSlice;
+use crate::codec::assert::insufficient_data;
 use crate::error::Error;
 
 /// Trait for types that can be stored in a 
[`ReqSketch`](crate::req::ReqSketch).
 ///
-/// Provides total ordering (so floating-point types with NaN are well-defined 
under
-/// sketch operations) and binary serialization compatible with the Apache 
DataSketches
+/// Provides ordering and binary serialization compatible with the Apache 
DataSketches
 /// REQ wire format used by the C++ and Java reference implementations.
 pub trait ReqValue: Sized + Clone + PartialOrd {
-    /// Total ordering used for sketch operations (sort, compaction, rank, 
quantile).
-    ///
-    /// For integer types this is equivalent to [`Ord::cmp`]. For 
floating-point types
-    /// this delegates to [`f32::total_cmp`] / [`f64::total_cmp`] so NaN 
comparisons are
-    /// deterministic.
-    fn total_cmp(&self, other: &Self) -> Ordering;
+    /// Compares two values. See each implementation for its ordering 
semantics.
+    fn compare(&self, other: &Self) -> Ordering;
 
     /// Returns true if this value is the floating-point NaN sentinel.
     ///
@@ -55,78 +52,141 @@ pub trait ReqValue: Sized + Clone + PartialOrd {
     fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, Error>;
 }
 
-macro_rules! impl_req_value_primitive {
-    // Form with explicit is_nan body (for float types).
-    ($t:ty, $read:ident, $write:ident, $cmp:expr, nan: $nan:expr) => {
-        impl ReqValue for $t {
-            #[inline(always)]
-            fn total_cmp(&self, other: &Self) -> Ordering {
-                $cmp(self, other)
-            }
-
-            fn serialize_size(_item: &Self) -> usize {
-                std::mem::size_of::<$t>()
-            }
-
-            fn serialize_value(&self, bytes: &mut SketchBytes) {
-                bytes.$write(*self);
-            }
-
-            fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, 
Error> {
-                cursor.$read().map_err(|_| {
-                    Error::insufficient_data(concat!(
-                        "failed to read ",
-                        stringify!($t),
-                        " from REQ sketch"
-                    ))
-                })
-            }
-
-            #[inline(always)]
-            fn is_nan(&self) -> bool {
-                $nan(self)
-            }
-        }
-    };
-    // Form without is_nan (for integer types — default returns false).
-    ($t:ty, $read:ident, $write:ident, $cmp:expr) => {
-        impl ReqValue for $t {
-            #[inline(always)]
-            fn total_cmp(&self, other: &Self) -> Ordering {
-                $cmp(self, other)
-            }
-
-            fn serialize_size(_item: &Self) -> usize {
-                std::mem::size_of::<$t>()
-            }
-
-            fn serialize_value(&self, bytes: &mut SketchBytes) {
-                bytes.$write(*self);
-            }
-
-            fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, 
Error> {
-                cursor.$read().map_err(|_| {
-                    Error::insufficient_data(concat!(
-                        "failed to read ",
-                        stringify!($t),
-                        " from REQ sketch"
-                    ))
-                })
-            }
-        }
-    };
+impl ReqValue for i32 {
+    #[inline(always)]
+    fn compare(&self, other: &Self) -> Ordering {
+        self.cmp(other)
+    }
+
+    fn serialize_size(_item: &Self) -> usize {
+        size_of::<Self>()
+    }
+
+    fn serialize_value(&self, bytes: &mut SketchBytes) {
+        bytes.write_i32_le(*self);
+    }
+
+    fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, Error> {
+        cursor
+            .read_i32_le()
+            .map_err(insufficient_data("failed to read i32 from REQ sketch"))
+    }
+}
+
+impl ReqValue for i64 {
+    #[inline(always)]
+    fn compare(&self, other: &Self) -> Ordering {
+        self.cmp(other)
+    }
+
+    fn serialize_size(_item: &Self) -> usize {
+        size_of::<Self>()
+    }
+
+    fn serialize_value(&self, bytes: &mut SketchBytes) {
+        bytes.write_i64_le(*self);
+    }
+
+    fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, Error> {
+        cursor
+            .read_i64_le()
+            .map_err(insufficient_data("failed to read i64 from REQ sketch"))
+    }
+}
+
+impl ReqValue for u32 {
+    #[inline(always)]
+    fn compare(&self, other: &Self) -> Ordering {
+        self.cmp(other)
+    }
+
+    fn serialize_size(_item: &Self) -> usize {
+        size_of::<Self>()
+    }
+
+    fn serialize_value(&self, bytes: &mut SketchBytes) {
+        bytes.write_u32_le(*self);
+    }
+
+    fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, Error> {
+        cursor
+            .read_u32_le()
+            .map_err(insufficient_data("failed to read u32 from REQ sketch"))
+    }
+}
+
+impl ReqValue for u64 {
+    #[inline(always)]
+    fn compare(&self, other: &Self) -> Ordering {
+        self.cmp(other)
+    }
+
+    fn serialize_size(_item: &Self) -> usize {
+        size_of::<Self>()
+    }
+
+    fn serialize_value(&self, bytes: &mut SketchBytes) {
+        bytes.write_u64_le(*self);
+    }
+
+    fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, Error> {
+        cursor
+            .read_u64_le()
+            .map_err(insufficient_data("failed to read u64 from REQ sketch"))
+    }
+}
+
+impl ReqValue for f32 {
+    #[inline(always)]
+    fn compare(&self, other: &Self) -> Ordering {
+        self.partial_cmp(other).unwrap()
+    }
+
+    #[inline(always)]
+    fn is_nan(&self) -> bool {
+        f32::is_nan(*self)
+    }
+
+    fn serialize_size(_item: &Self) -> usize {
+        size_of::<Self>()
+    }
+
+    fn serialize_value(&self, bytes: &mut SketchBytes) {
+        bytes.write_f32_le(*self);
+    }
+
+    fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, Error> {
+        cursor
+            .read_f32_le()
+            .map_err(insufficient_data("failed to read f32 from REQ sketch"))
+    }
 }
 
-impl_req_value_primitive!(i32, read_i32_le, write_i32_le, Ord::cmp);
-impl_req_value_primitive!(i64, read_i64_le, write_i64_le, Ord::cmp);
-impl_req_value_primitive!(u32, read_u32_le, write_u32_le, Ord::cmp);
-impl_req_value_primitive!(u64, read_u64_le, write_u64_le, Ord::cmp);
-impl_req_value_primitive!(f32, read_f32_le, write_f32_le,
-    |a: &f32, b: &f32| if let Some(o) = a.partial_cmp(b) { o } else { 
f32::total_cmp(a, b) },
-    nan: |x: &f32| f32::is_nan(*x));
-impl_req_value_primitive!(f64, read_f64_le, write_f64_le,
-    |a: &f64, b: &f64| if let Some(o) = a.partial_cmp(b) { o } else { 
f64::total_cmp(a, b) },
-    nan: |x: &f64| f64::is_nan(*x));
+impl ReqValue for f64 {
+    #[inline(always)]
+    fn compare(&self, other: &Self) -> Ordering {
+        self.partial_cmp(other).unwrap()
+    }
+
+    #[inline(always)]
+    fn is_nan(&self) -> bool {
+        f64::is_nan(*self)
+    }
+
+    fn serialize_size(_item: &Self) -> usize {
+        size_of::<Self>()
+    }
+
+    fn serialize_value(&self, bytes: &mut SketchBytes) {
+        bytes.write_f64_le(*self);
+    }
+
+    fn deserialize_value(cursor: &mut SketchSlice<'_>) -> Result<Self, Error> {
+        cursor
+            .read_f64_le()
+            .map_err(insufficient_data("failed to read f64 from REQ sketch"))
+    }
+}
 
 #[cfg(test)]
 mod tests {
@@ -173,19 +233,33 @@ mod tests {
     }
 
     #[test]
-    fn total_cmp_handles_nan_for_floats() {
-        // Pure NaN comparisons under PartialOrd return None; total_cmp must 
give a definite
-        // Ordering.
-        let nan = f64::NAN;
-        let one = 1.0_f64;
-        assert_ne!(<f64 as ReqValue>::total_cmp(&nan, &one), Ordering::Equal);
-        assert_eq!(<f64 as ReqValue>::total_cmp(&nan, &nan), Ordering::Equal);
+    fn compare_for_f32_uses_numeric_order() {
+        assert_eq!(<f32 as ReqValue>::compare(&-0.0, &0.0), Ordering::Equal);
+        assert_eq!(
+            <f32 as ReqValue>::compare(&f32::NEG_INFINITY, &f32::INFINITY),
+            Ordering::Less
+        );
+    }
+
+    #[test]
+    fn compare_for_f64_uses_numeric_order() {
+        assert_eq!(<f64 as ReqValue>::compare(&-0.0, &0.0), Ordering::Equal);
+        assert_eq!(
+            <f64 as ReqValue>::compare(&f64::NEG_INFINITY, &f64::INFINITY),
+            Ordering::Less
+        );
+    }
+
+    #[test]
+    #[should_panic]
+    fn compare_for_floats_rejects_nan() {
+        <f64 as ReqValue>::compare(&f64::NAN, &0.0);
     }
 
     #[test]
-    fn total_cmp_for_integers_matches_ord() {
-        assert_eq!(<i64 as ReqValue>::total_cmp(&3, &5), Ordering::Less);
-        assert_eq!(<i64 as ReqValue>::total_cmp(&5, &5), Ordering::Equal);
-        assert_eq!(<i64 as ReqValue>::total_cmp(&7, &5), Ordering::Greater);
+    fn compare_for_integers_matches_ord() {
+        assert_eq!(<i64 as ReqValue>::compare(&3, &5), Ordering::Less);
+        assert_eq!(<i64 as ReqValue>::compare(&5, &5), Ordering::Equal);
+        assert_eq!(<i64 as ReqValue>::compare(&7, &5), Ordering::Greater);
     }
 }
diff --git a/tests-integration/tests/req_test/query.rs 
b/tests-integration/tests/req_test/query.rs
index c07bff9..30f3092 100644
--- a/tests-integration/tests/req_test/query.rs
+++ b/tests-integration/tests/req_test/query.rs
@@ -191,3 +191,19 @@ fn search_criteria_rank_consistency() -> Result<(), Error> 
{
 
     Ok(())
 }
+
+#[test]
+fn signed_zeros_share_rank_and_cannot_be_distinct_splits() -> Result<(), 
Error> {
+    let mut sketch = ReqSketch::default();
+    sketch.update(-0.0_f64);
+    sketch.update(0.0_f64);
+
+    for value in [-0.0, 0.0] {
+        assert_eq!(sketch.rank(&value, SearchCriteria::Exclusive)?, 0.0);
+        assert_eq!(sketch.rank(&value, SearchCriteria::Inclusive)?, 1.0);
+    }
+
+    assert!(sketch.pmf(&[-0.0, 0.0], SearchCriteria::Inclusive).is_err());
+
+    Ok(())
+}
diff --git a/tests-integration/tests/req_test/sorted_view_api.rs 
b/tests-integration/tests/req_test/sorted_view_api.rs
index b5ed000..32a518b 100644
--- a/tests-integration/tests/req_test/sorted_view_api.rs
+++ b/tests-integration/tests/req_test/sorted_view_api.rs
@@ -119,18 +119,38 @@ fn view_rank_is_primary_query_name() {
 }
 
 #[test]
-fn nan_query_items_are_rejected() {
+fn nan_query_items_and_split_points_are_rejected() {
     let sketch = populated_sketch(100);
     let error = sketch
         .rank(&f64::NAN, SearchCriteria::Inclusive)
         .unwrap_err();
     assert_eq!(error.kind(), ErrorKind::InvalidArgument);
+    assert_eq!(
+        sketch
+            .pmf(&[f64::NAN], SearchCriteria::Inclusive)
+            .unwrap_err()
+            .kind(),
+        ErrorKind::InvalidArgument
+    );
+    assert_eq!(
+        sketch
+            .cdf(&[50.0, f64::NAN], SearchCriteria::Inclusive)
+            .unwrap_err()
+            .kind(),
+        ErrorKind::InvalidArgument
+    );
 
     let view = sketch.sorted_view();
     assert_that!(
         view.rank(&f64::NAN, SearchCriteria::Inclusive),
         err(anything())
     );
+    assert_eq!(
+        view.pmf(&[f64::NAN], SearchCriteria::Inclusive)
+            .unwrap_err()
+            .kind(),
+        ErrorKind::InvalidArgument
+    );
 }
 
 #[test]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to