Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,10 @@ All significant changes to this project will be documented in this file.

* Add KLL sketches behind the `kll` feature, including rank, quantile, PMF, and CDF queries, custom item ordering, merging, and C++/Java-compatible serialization.

### 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
Expand Down
101 changes: 81 additions & 20 deletions datasketches/src/tdigest/sketch.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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::<f32>() + size_of::<u32>(), size_of::<f32>())
} else {
Expand Down Expand Up @@ -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 {
(
Expand All @@ -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,
Expand Down Expand Up @@ -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!(
Expand All @@ -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,
Expand All @@ -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!(
Expand All @@ -804,18 +837,31 @@ impl TDigestMut {
cursor.read_u32_be().map_err(make_error("<unused>"))?;
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,
Expand Down Expand Up @@ -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, Error> {
NonZeroU64::new(value)
.ok_or_else(|| Error::deserial(format!("malformed data: {tag} cannot be zero")))
Expand Down
54 changes: 54 additions & 0 deletions tests-integration/tests/serde_tests/tdigest.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<u8> {
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(&centroid_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();
Expand Down