From 9e970efb6e62eaf1e6ccb6a7a92b6c02698face0 Mon Sep 17 00:00:00 2001 From: tison Date: Wed, 2 Sep 2026 10:33:46 +0800 Subject: [PATCH] fix(tdigest): validate deserialized state invariants --- CHANGELOG.md | 4 + datasketches/src/tdigest/sketch.rs | 101 ++++++++++++++---- .../tests/serde_tests/tdigest.rs | 54 ++++++++++ 3 files changed, 139 insertions(+), 20 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 969a56a..85f13d1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,10 @@ All significant changes to this project will be documented in this file. ## Unreleased +### Bug fixes + +* T-Digest deserialization now rejects unknown or conflicting flags, reversed extrema, out-of-range values, unsorted centroids, and non-empty images without stored values. + ## v0.5.0 ### Breaking changes diff --git a/datasketches/src/tdigest/sketch.rs b/datasketches/src/tdigest/sketch.rs index 942c735..bedd640 100644 --- a/datasketches/src/tdigest/sketch.rs +++ b/datasketches/src/tdigest/sketch.rs @@ -595,8 +595,20 @@ impl TDigestMut { return Err(Error::deserial(format!("k must be at least 10, got {k}"))); } let flags = cursor.read_u8().map_err(insufficient_data("flags"))?; + let known_flags = FLAGS_IS_EMPTY | FLAGS_IS_SINGLE_VALUE | FLAGS_REVERSE_MERGE; + if flags & !known_flags != 0 { + return Err(Error::deserial(format!( + "malformed data: unknown TDigest flags 0x{:02x}", + flags & !known_flags + ))); + } let is_empty = (flags & FLAGS_IS_EMPTY) != 0; let is_single_value = (flags & FLAGS_IS_SINGLE_VALUE) != 0; + if is_empty && is_single_value { + return Err(Error::deserial( + "malformed data: empty and single-value flags are mutually exclusive", + )); + } let expected_preamble_longs = if is_empty || is_single_value { PREAMBLE_LONGS_EMPTY_OR_SINGLE } else { @@ -655,10 +667,7 @@ impl TDigestMut { cursor.read_f64_le().map_err(insufficient_data("max"))?, ) }; - check_non_nan(min, "min")?; - check_non_nan(max, "max")?; - check_finite(min, "min")?; - check_finite(max, "max")?; + check_extrema(min, max, "TDigest")?; let (centroid_bytes, buffered_value_bytes) = if is_f32 { (size_of::() + size_of::(), size_of::()) } else { @@ -686,8 +695,15 @@ impl TDigestMut { let stored_centroids = num_centroids.checked_add(num_buffered).ok_or_else(|| { Error::deserial("num_centroids and num_buffered exceed the supported size") })?; + if stored_centroids == 0 { + return Err(Error::deserial( + "malformed data: non-empty TDigest must contain a centroid or buffered value", + )); + } let mut centroids = Vec::with_capacity(stored_centroids); let mut compressed_weight = 0u64; + let mut previous_mean = min; + let mut centroid_means_valid = true; for bytes in centroid_payload.chunks_exact(centroid_bytes) { let (mean, weight) = if is_f32 { ( @@ -700,26 +716,36 @@ impl TDigestMut { u64::from_le_bytes(bytes[8..].try_into().unwrap()), ) }; - check_non_nan(mean, "centroid mean")?; - check_finite(mean, "centroid")?; + centroid_means_valid &= mean.is_finite() & (mean >= previous_mean) & (mean <= max); + previous_mean = mean; let weight = check_nonzero(weight, "centroid weight")?; compressed_weight = checked_weight_sum(compressed_weight, weight.get())?; centroids.push(Centroid { mean, weight }); } + if !centroid_means_valid { + return Err(Error::deserial( + "malformed data: centroid means must be finite, within extrema, and nondecreasing", + )); + } checked_weight_sum(compressed_weight, num_buffered as u64)?; + let mut buffered_values_valid = true; for bytes in buffered_payload.chunks_exact(buffered_value_bytes) { let value = if is_f32 { f32::from_le_bytes(bytes.try_into().unwrap()) as f64 } else { f64::from_le_bytes(bytes.try_into().unwrap()) }; - check_non_nan(value, "buffered_value mean")?; - check_finite(value, "buffered_value mean")?; + buffered_values_valid &= value.is_finite() & (value >= min) & (value <= max); centroids.push(Centroid { mean: value, weight: DEFAULT_WEIGHT, }); } + if !buffered_values_valid { + return Err(Error::deserial( + "malformed data: buffered values must be finite and within extrema", + )); + } Ok(TDigestMut::make( k, reverse_merge, @@ -748,10 +774,7 @@ impl TDigestMut { // compatibility with asBytes() let min = cursor.read_f64_be().map_err(make_error("min"))?; let max = cursor.read_f64_be().map_err(make_error("max"))?; - check_non_nan(min, "min in compat double format")?; - check_non_nan(max, "max in compat double format")?; - check_finite(min, "min in compat double format")?; - check_finite(max, "max in compat double format")?; + check_extrema(min, max, "compat double TDigest")?; let k = cursor.read_f64_be().map_err(make_error("k"))? as u16; if k < 10 { return Err(Error::deserial(format!( @@ -760,18 +783,31 @@ impl TDigestMut { } let num_centroids = cursor.read_u32_be().map_err(make_error("num_centroids"))? as usize; + if num_centroids == 0 { + return Err(Error::deserial( + "malformed data: compat double TDigest must contain a centroid", + )); + } let mut total_weight = 0u64; let mut centroids = Vec::with_capacity(num_centroids); + let mut previous_mean = min; + let mut centroid_means_valid = true; for _ in 0..num_centroids { let weight = cursor.read_f64_be().map_err(make_error("weight"))?; let mean = cursor.read_f64_be().map_err(make_error("mean"))?; let weight = check_compat_weight(weight, "centroid weight in compat double format")?; - check_non_nan(mean, "centroid mean in compat double format")?; - check_finite(mean, "centroid mean in compat double format")?; + centroid_means_valid &= + mean.is_finite() & (mean >= previous_mean) & (mean <= max); + previous_mean = mean; total_weight = checked_weight_sum(total_weight, weight.get())?; centroids.push(Centroid { mean, weight }); } + if !centroid_means_valid { + return Err(Error::deserial( + "malformed data: centroid means in compat double format must be finite, within extrema, and nondecreasing", + )); + } Ok(TDigestMut::make( k, false, @@ -789,10 +825,7 @@ impl TDigestMut { // reference implementation uses doubles for min and max let min = cursor.read_f64_be().map_err(make_error("min"))?; let max = cursor.read_f64_be().map_err(make_error("max"))?; - check_non_nan(min, "min in compat float format")?; - check_non_nan(max, "max in compat float format")?; - check_finite(min, "min in compat float format")?; - check_finite(max, "max in compat float format")?; + check_extrema(min, max, "compat float TDigest")?; let k = cursor.read_f32_be().map_err(make_error("k"))? as u16; if k < 10 { return Err(Error::deserial(format!( @@ -804,18 +837,31 @@ impl TDigestMut { cursor.read_u32_be().map_err(make_error(""))?; let num_centroids = cursor.read_u16_be().map_err(make_error("num_centroids"))? as usize; + if num_centroids == 0 { + return Err(Error::deserial( + "malformed data: compat float TDigest must contain a centroid", + )); + } let mut total_weight = 0u64; let mut centroids = Vec::with_capacity(num_centroids); + let mut previous_mean = min; + let mut centroid_means_valid = true; for _ in 0..num_centroids { let weight = cursor.read_f32_be().map_err(make_error("weight"))? as f64; let mean = cursor.read_f32_be().map_err(make_error("mean"))? as f64; let weight = check_compat_weight(weight, "centroid weight in compat float format")?; - check_non_nan(mean, "centroid mean in compat float format")?; - check_finite(mean, "centroid mean in compat float format")?; + centroid_means_valid &= + mean.is_finite() & (mean >= previous_mean) & (mean <= max); + previous_mean = mean; total_weight = checked_weight_sum(total_weight, weight.get())?; centroids.push(Centroid { mean, weight }); } + if !centroid_means_valid { + return Err(Error::deserial( + "malformed data: centroid means in compat float format must be finite, within extrema, and nondecreasing", + )); + } Ok(TDigestMut::make( k, false, @@ -1597,6 +1643,21 @@ fn check_finite(value: f64, tag: &'static str) -> Result<(), Error> { Ok(()) } +#[inline] +fn check_extrema(min: f64, max: f64, format: &'static str) -> Result<(), Error> { + if !min.is_finite() || !max.is_finite() { + return Err(Error::deserial(format!( + "malformed data: {format} extrema must be finite" + ))); + } + if min > max { + return Err(Error::deserial(format!( + "malformed data: {format} min {min} exceeds max {max}" + ))); + } + Ok(()) +} + fn check_nonzero(value: u64, tag: &'static str) -> Result { NonZeroU64::new(value) .ok_or_else(|| Error::deserial(format!("malformed data: {tag} cannot be zero"))) diff --git a/tests-integration/tests/serde_tests/tdigest.rs b/tests-integration/tests/serde_tests/tdigest.rs index c5bd0fe..1b871ef 100644 --- a/tests-integration/tests/serde_tests/tdigest.rs +++ b/tests-integration/tests/serde_tests/tdigest.rs @@ -350,6 +350,60 @@ fn test_updates_normalize_overfull_deserialized_mixed_buffer() { assert_eq!(roundtrip.max_value(), Some(1_000.0)); } +fn serialized_two_value_digest() -> Vec { + let mut tdigest = TDigestMut::new(100).unwrap(); + tdigest.update(0.0); + tdigest.update(1.0); + tdigest.serialize() +} + +fn assert_invalid_tdigest(bytes: &[u8]) { + let error = TDigestMut::deserialize(bytes).unwrap_err(); + assert_eq!(error.kind(), datasketches::error::ErrorKind::InvalidData); +} + +#[test] +fn test_deserialize_rejects_unknown_or_conflicting_flags() { + let mut unknown = serialized_two_value_digest(); + unknown[5] |= 0x80; + assert_invalid_tdigest(&unknown); + + let mut empty = TDigestMut::new(100).unwrap(); + let mut conflicting = empty.serialize(); + conflicting[5] |= 1 << 1; + assert_invalid_tdigest(&conflicting); +} + +#[test] +fn test_deserialize_rejects_invalid_extrema_and_centroid_ranges() { + let mut reversed_extrema = serialized_two_value_digest(); + reversed_extrema[16..24].copy_from_slice(&2_f64.to_le_bytes()); + assert_invalid_tdigest(&reversed_extrema); + + let mut centroid_outside_extrema = serialized_two_value_digest(); + centroid_outside_extrema[32..40].copy_from_slice(&(-1_f64).to_le_bytes()); + assert_invalid_tdigest(¢roid_outside_extrema); + + let mut buffered_outside_extrema = serialized_two_value_digest(); + buffered_outside_extrema[8..12].copy_from_slice(&1_u32.to_le_bytes()); + buffered_outside_extrema[12..16].copy_from_slice(&1_u32.to_le_bytes()); + buffered_outside_extrema[48..56].copy_from_slice(&2_f64.to_le_bytes()); + assert_invalid_tdigest(&buffered_outside_extrema); +} + +#[test] +fn test_deserialize_rejects_unsorted_or_missing_centroids() { + let mut unsorted = serialized_two_value_digest(); + unsorted[32..40].copy_from_slice(&1_f64.to_le_bytes()); + unsorted[48..56].copy_from_slice(&0_f64.to_le_bytes()); + assert_invalid_tdigest(&unsorted); + + let mut missing = serialized_two_value_digest(); + missing[8..12].copy_from_slice(&0_u32.to_le_bytes()); + missing[12..16].copy_from_slice(&0_u32.to_le_bytes()); + assert_invalid_tdigest(&missing); +} + #[test] fn test_deserialize_rejects_truncated_large_payload_before_allocation() { let mut tdigest = TDigestMut::new(10).unwrap();