use rmp_serde::{decode::Error as RmpDecodeError, encode::Error as RmpEncodeError, from_slice};
use serde::{Deserialize, Serialize};
use std::marker::PhantomData;
use crate::message_pack_format::envelope;
use crate::{CommonHeap, HashProfile, KeepLargest, SketchHasher};
use super::KMV;
const KMV_KIND: &[u8] = &[0x0e, 0x00];
#[derive(Debug, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
struct KmvMetadata {
metadata_version: u8,
hash_profile_id: String,
hash_algorithm: String,
seed_derivation: String,
input_encoding: String,
seed_list: Vec<u64>,
canonical_seed_index: u32,
k: u32,
}
fn kmv_metadata<H: HashProfile>(k: u32) -> KmvMetadata {
KmvMetadata {
metadata_version: 1,
hash_profile_id: H::PROFILE_ID.to_string(),
hash_algorithm: H::ALGORITHM.to_string(),
seed_derivation: H::SEED_DERIVATION.to_string(),
input_encoding: H::INPUT_ENCODING.to_string(),
seed_list: H::seed_list(),
canonical_seed_index: H::CANONICAL_SEED_INDEX,
k,
}
}
#[derive(Debug, Serialize, Deserialize)]
struct KmvPayload {
hashes: Vec<u64>,
}
#[derive(Serialize)]
struct HeapSeed {
data: Vec<u64>,
size: usize,
order: KeepLargest,
}
fn rebuild_heap(
k: usize,
ascending: Vec<u64>,
) -> Result<CommonHeap<u64, KeepLargest>, RmpDecodeError> {
let len = ascending.len();
let mut data = ascending;
data.reverse();
let seed = rmp_serde::to_vec_named(&HeapSeed {
data,
size: k,
order: KeepLargest,
})
.map_err(|err| RmpDecodeError::Uncategorized(err.to_string()))?;
let heap: CommonHeap<u64, KeepLargest> = from_slice(&seed)?;
if heap.len() != len || heap.capacity() != k {
return Err(RmpDecodeError::Uncategorized(format!(
"KMV heap rebuild: {} of {len} hashes under a bound of {} against k {k}",
heap.len(),
heap.capacity()
)));
}
Ok(heap)
}
impl<H: SketchHasher + HashProfile> KMV<H> {
pub fn serialize_to_bytes(&self) -> Result<Vec<u8>, RmpEncodeError> {
let k = u32::try_from(self.k).map_err(|_| {
RmpEncodeError::Syntax(format!("KMV k {} exceeds the u32 metadata field", self.k))
})?;
if k == 0 {
return Err(RmpEncodeError::Syntax(
"KMV k must be at least 1".to_string(),
));
}
let hashes = self.wire_hashes();
if hashes.len() > self.k {
return Err(RmpEncodeError::Syntax(format!(
"KMV holds {} hashes over a k of {}",
hashes.len(),
self.k
)));
}
if hashes.windows(2).any(|pair| pair[0] == pair[1]) {
return Err(RmpEncodeError::Syntax(
"KMV holds the same hash twice".to_string(),
));
}
if self.k_vals.capacity() != self.k {
return Err(RmpEncodeError::Syntax(format!(
"KMV k {} disagrees with the retained bound {}",
self.k,
self.k_vals.capacity()
)));
}
let metadata = rmp_serde::to_vec_named(&kmv_metadata::<H>(k))?;
let payload = rmp_serde::to_vec(&KmvPayload { hashes })?;
Ok(envelope::encode(KMV_KIND, &metadata, &payload))
}
pub fn deserialize_from_bytes(bytes: &[u8]) -> Result<Self, RmpDecodeError> {
let (kind_id, metadata, payload) =
envelope::split(bytes).map_err(RmpDecodeError::Uncategorized)?;
if kind_id != KMV_KIND {
return Err(RmpDecodeError::Uncategorized(format!(
"KMV kind_id mismatch: stored {kind_id:?}, expected {KMV_KIND:?}"
)));
}
let meta: KmvMetadata = from_slice(metadata)?;
if meta != kmv_metadata::<H>(meta.k) {
return Err(RmpDecodeError::Uncategorized(
"ASAPv1 KMV envelope: metadata mismatch".to_string(),
));
}
if meta.k == 0 {
return Err(RmpDecodeError::Uncategorized(
"KMV k must be at least 1".to_string(),
));
}
let k = meta.k as usize;
let decoded: KmvPayload = from_slice(payload)?;
if decoded.hashes.len() > k {
return Err(RmpDecodeError::Uncategorized(format!(
"KMV payload carries {} hashes over a k of {k}",
decoded.hashes.len()
)));
}
if decoded.hashes.windows(2).any(|pair| pair[0] >= pair[1]) {
return Err(RmpDecodeError::Uncategorized(
"KMV hashes are not strictly ascending".to_string(),
));
}
Ok(KMV {
k,
k_vals: rebuild_heap(k, decoded.hashes)?,
_hasher: PhantomData,
})
}
fn wire_hashes(&self) -> Vec<u64> {
let mut hashes: Vec<u64> = self.k_vals.iter().copied().collect();
hashes.sort_unstable();
hashes
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{CANONICAL_HASH_SEED, DataInput, DefaultXxHasher, HeapItem};
fn populated(k: usize, keys: u64) -> KMV {
let mut sketch: KMV = KMV::new(k);
for value in 0..keys {
sketch.insert(&DataInput::U64(value));
}
sketch
}
fn metadata_of(bytes: &[u8]) -> KmvMetadata {
let (_, metadata, _) = envelope::split(bytes).expect("split");
from_slice(metadata).expect("metadata")
}
fn payload_of(bytes: &[u8]) -> KmvPayload {
let (_, _, payload) = envelope::split(bytes).expect("split");
from_slice(payload).expect("payload")
}
#[test]
fn kmv_round_trip_serialization() {
let mut sketch = populated(64, 5_000);
assert_eq!(sketch.k_vals.len(), 64, "the bound was never reached");
let encoded = sketch.serialize_to_bytes().expect("serialize KMV");
assert!(encoded.starts_with(b"ASAPv1"));
assert_eq!(&encoded[7..10], &[2u8, 0x0e, 0x00]);
let meta = metadata_of(&encoded);
assert_eq!(meta.metadata_version, 1);
assert_eq!(meta.k, 64);
assert_eq!(meta.canonical_seed_index, CANONICAL_HASH_SEED as u32);
let emitted = payload_of(&encoded);
assert_eq!(emitted.hashes.len(), 64);
for pair in emitted.hashes.windows(2) {
assert!(pair[0] < pair[1], "the emitted hashes are not ascending");
}
let mut decoded = KMV::<DefaultXxHasher>::deserialize_from_bytes(&encoded).expect("decode");
assert_eq!(decoded.k, sketch.k);
assert_eq!(decoded.k_vals.capacity(), sketch.k_vals.capacity());
assert_eq!(decoded.wire_hashes(), sketch.wire_hashes());
assert_eq!(
decoded.k_vals.peek(),
emitted.hashes.last(),
"the rebuilt heap does not hold the largest hash at its root"
);
assert_eq!(decoded.estimate(), sketch.estimate());
assert_eq!(
decoded.serialize_to_bytes().expect("re-serialize"),
encoded,
"KMV bytes differed after a round trip"
);
}
#[test]
fn kmv_emitted_order_is_independent_of_insertion_order() {
let keys: Vec<u64> = (0..2_000).collect();
let mut forward: KMV = KMV::new(32);
let mut backward: KMV = KMV::new(32);
for value in &keys {
forward.insert(&DataInput::U64(*value));
}
for value in keys.iter().rev() {
backward.insert(&DataInput::U64(*value));
}
assert_eq!(forward.k_vals.len(), 32);
assert_ne!(
forward.k_vals.as_slice(),
backward.k_vals.as_slice(),
"the fixture cannot tell the emitted order from the heap order"
);
let bytes = forward.serialize_to_bytes().expect("serialize");
assert_eq!(
bytes,
backward.serialize_to_bytes().expect("serialize"),
"the emitted order followed the insertion order"
);
let decoded = KMV::<DefaultXxHasher>::deserialize_from_bytes(&bytes).expect("decode");
assert_eq!(decoded.serialize_to_bytes().expect("re-serialize"), bytes);
}
#[test]
fn kmv_empty_round_trip() {
let sketch: KMV = KMV::new(16);
let bytes = sketch.serialize_to_bytes().expect("serialize");
assert_eq!(metadata_of(&bytes).k, 16);
assert!(payload_of(&bytes).hashes.is_empty());
let mut decoded = KMV::<DefaultXxHasher>::deserialize_from_bytes(&bytes).expect("decode");
assert_eq!(decoded.k, 16);
assert_eq!(decoded.k_vals.len(), 0);
assert_eq!(decoded.k_vals.capacity(), 16);
assert_eq!(decoded.estimate(), 0.0);
assert_eq!(decoded.serialize_to_bytes().expect("re-serialize"), bytes);
}
#[test]
fn kmv_carries_hashes_at_full_u64_width() {
let mut sketch: KMV = KMV::new(4);
for value in [0u64, 1, u64::MAX / 2, u64::MAX] {
sketch.insert_by_hash(value);
}
let bytes = sketch.serialize_to_bytes().expect("serialize");
assert_eq!(
payload_of(&bytes).hashes,
vec![0, 1, u64::MAX / 2, u64::MAX]
);
let decoded = KMV::<DefaultXxHasher>::deserialize_from_bytes(&bytes).expect("decode");
assert_eq!(decoded.wire_hashes(), vec![0, 1, u64::MAX / 2, u64::MAX]);
assert_eq!(decoded.serialize_to_bytes().expect("re-serialize"), bytes);
}
#[test]
fn kmv_rejects_foreign_kind_id() {
let cms =
crate::CountMin::<crate::Vector2D<i64>, crate::RegularPath>::with_dimensions(3, 8);
let cms_bytes = cms.serialize_to_bytes().expect("serialize CMS");
let problem = KMV::<DefaultXxHasher>::deserialize_from_bytes(&cms_bytes)
.expect_err("CMS bytes must not decode as a KMV")
.to_string();
assert!(problem.contains("kind_id mismatch"), "got {problem}");
}
#[test]
fn kmv_metadata_rejects_unknown_keys() {
#[derive(Serialize)]
struct WithExtra {
metadata_version: u8,
hash_profile_id: String,
hash_algorithm: String,
seed_derivation: String,
input_encoding: String,
seed_list: Vec<u64>,
canonical_seed_index: u32,
k: u32,
bogus_field: u8, }
let m = kmv_metadata::<DefaultXxHasher>(64);
let extra = WithExtra {
metadata_version: m.metadata_version,
hash_profile_id: m.hash_profile_id.clone(),
hash_algorithm: m.hash_algorithm.clone(),
seed_derivation: m.seed_derivation.clone(),
input_encoding: m.input_encoding.clone(),
seed_list: m.seed_list.clone(),
canonical_seed_index: m.canonical_seed_index,
k: m.k,
bogus_field: 7,
};
let bytes = rmp_serde::to_vec_named(&extra).expect("encode");
assert!(
from_slice::<KmvMetadata>(&bytes).is_err(),
"an unexpected metadata key must be rejected"
);
}
#[test]
fn kmv_metadata_rejects_a_missing_k_key() {
#[derive(Serialize)]
struct WithoutK {
metadata_version: u8,
hash_profile_id: String,
hash_algorithm: String,
seed_derivation: String,
input_encoding: String,
seed_list: Vec<u64>,
canonical_seed_index: u32,
}
let m = kmv_metadata::<DefaultXxHasher>(64);
let without = WithoutK {
metadata_version: m.metadata_version,
hash_profile_id: m.hash_profile_id.clone(),
hash_algorithm: m.hash_algorithm.clone(),
seed_derivation: m.seed_derivation.clone(),
input_encoding: m.input_encoding.clone(),
seed_list: m.seed_list.clone(),
canonical_seed_index: m.canonical_seed_index,
};
let bytes = rmp_serde::to_vec_named(&without).expect("encode");
assert!(
from_slice::<KmvMetadata>(&bytes).is_err(),
"a missing metadata key must be rejected"
);
}
fn crafted(k: u32, hashes: Vec<u64>) -> Vec<u8> {
let metadata = rmp_serde::to_vec_named(&kmv_metadata::<DefaultXxHasher>(k)).expect("meta");
let payload = rmp_serde::to_vec(&KmvPayload { hashes }).expect("payload");
envelope::encode(KMV_KIND, &metadata, &payload)
}
#[test]
fn kmv_rejects_a_crafted_envelope() {
let cases: Vec<(Vec<u8>, &str)> = vec![
(crafted(0, Vec::new()), "k must be at least 1"),
(crafted(2, vec![1, 2, 3]), "3 hashes over a k of 2"),
(crafted(4, vec![3, 1, 2]), "not strictly ascending"),
(crafted(4, vec![1, 1, 2]), "not strictly ascending"),
];
for (bytes, expected) in cases {
let problem = KMV::<DefaultXxHasher>::deserialize_from_bytes(&bytes)
.expect_err("a crafted envelope must be rejected, not decoded")
.to_string();
assert!(
problem.contains(expected),
"expected a complaint about {expected}, got {problem}"
);
}
}
#[test]
fn kmv_rejects_a_payload_declaring_more_hashes_than_it_carries() {
let metadata = rmp_serde::to_vec_named(&kmv_metadata::<DefaultXxHasher>(4)).expect("meta");
let payload = vec![0x91, 0xdd, 0x40, 0x00, 0x00, 0x00, 0x01, 0x02];
let bytes = envelope::encode(KMV_KIND, &metadata, &payload);
assert!(
KMV::<DefaultXxHasher>::deserialize_from_bytes(&bytes).is_err(),
"an over-declared hash count must be rejected, not allocated"
);
}
#[test]
fn kmv_does_not_allocate_a_declared_k() {
let bytes = crafted(u32::MAX, vec![7, 9]);
let mut decoded = KMV::<DefaultXxHasher>::deserialize_from_bytes(&bytes).expect("decode");
assert_eq!(decoded.k, u32::MAX as usize);
assert_eq!(decoded.k_vals.len(), 2);
assert_eq!(decoded.estimate(), 2.0);
decoded.insert_by_hash(11);
assert_eq!(decoded.k_vals.len(), 3, "the bound evicted instead");
assert_eq!(decoded.wire_hashes(), vec![7, 9, 11]);
}
#[test]
fn kmv_refuses_a_k_the_metadata_cannot_carry() {
let sketch: KMV = KMV {
k: 1 << 40,
k_vals: CommonHeap::new_max(1),
_hasher: PhantomData,
};
let problem = sketch
.serialize_to_bytes()
.expect_err("an oversized k must not serialize")
.to_string();
assert!(
problem.contains("exceeds the u32 metadata field"),
"got {problem}"
);
}
#[test]
fn kmv_refuses_to_serialize_a_state_decode_would_reject() {
let empty: KMV = KMV::new(0);
assert!(
empty.serialize_to_bytes().is_err(),
"a k of zero must not serialize"
);
let mut over: KMV = KMV {
k: 1,
k_vals: CommonHeap::new_max(4),
_hasher: PhantomData,
};
over.insert_by_hash(3);
over.insert_by_hash(5);
let problem = over
.serialize_to_bytes()
.expect_err("a retained set past k must not serialize")
.to_string();
assert!(problem.contains("2 hashes over a k of 1"), "got {problem}");
}
#[test]
fn kmv_refuses_a_k_that_disagrees_with_the_retained_bound() {
let mut wider: KMV = KMV::new(4);
wider.insert_by_hash(7);
wider.k = 100;
let problem = wider
.serialize_to_bytes()
.expect_err("a k over the retained bound must not serialize")
.to_string();
assert!(
problem.contains("KMV k 100 disagrees with the retained bound 4"),
"got {problem}"
);
let narrower: KMV = KMV {
k: 1,
k_vals: CommonHeap::new_max(4),
_hasher: PhantomData,
};
let problem = narrower
.serialize_to_bytes()
.expect_err("a k under the retained bound must not serialize")
.to_string();
assert!(
problem.contains("KMV k 1 disagrees with the retained bound 4"),
"got {problem}"
);
}
#[derive(Clone, Debug)]
struct AltHasher;
impl SketchHasher for AltHasher {
type HashType = <DefaultXxHasher as SketchHasher>::HashType;
fn hash64_seeded(d: usize, key: &DataInput) -> u64 {
DefaultXxHasher::hash64_seeded(d, key)
}
fn hash128_seeded(d: usize, key: &DataInput) -> u128 {
DefaultXxHasher::hash128_seeded(d, key)
}
fn hash_item64_seeded(d: usize, key: &HeapItem) -> u64 {
DefaultXxHasher::hash_item64_seeded(d, key)
}
fn hash_item128_seeded(d: usize, key: &HeapItem) -> u128 {
DefaultXxHasher::hash_item128_seeded(d, key)
}
fn hash_for_matrix_seeded(
seed_idx: usize,
rows: usize,
cols: usize,
key: &DataInput,
) -> Self::HashType {
DefaultXxHasher::hash_for_matrix_seeded(seed_idx, rows, cols, key)
}
}
impl HashProfile for AltHasher {
const PROFILE_ID: &'static str = "test.alt.profile.v1";
const ALGORITHM: &'static str = "xxh3_64_128";
const SEED_DERIVATION: &'static str = "seed_list_index_wrap";
const INPUT_ENCODING: &'static str = "projectasap.input.v1";
fn seed_list() -> Vec<u64> {
vec![1, 2, 3, 4, 5]
}
const CANONICAL_SEED_INDEX: u32 = CANONICAL_HASH_SEED as u32;
const MATRIX_SEED_INDEX: u32 = 0;
}
#[test]
fn kmv_custom_hasher_profile_round_trips_and_is_self_describing() {
let mut alt: KMV<AltHasher> = KMV::new(16);
let mut std: KMV = KMV::new(16);
for value in 0..500u64 {
alt.insert(&DataInput::U64(value));
std.insert(&DataInput::U64(value));
}
let alt_bytes = alt.serialize_to_bytes().expect("alt serialize");
let decoded = KMV::<AltHasher>::deserialize_from_bytes(&alt_bytes).expect("alt decode");
assert_eq!(decoded.wire_hashes(), alt.wire_hashes());
assert_eq!(
decoded.serialize_to_bytes().expect("re-serialize"),
alt_bytes
);
let std_bytes = std.serialize_to_bytes().expect("std serialize");
assert_eq!(std.wire_hashes(), alt.wire_hashes());
assert_ne!(alt_bytes, std_bytes);
assert!(
KMV::<DefaultXxHasher>::deserialize_from_bytes(&alt_bytes).is_err(),
"standard-profile decode must reject custom-profile bytes"
);
}
}