use rmp_serde::{decode::Error as RmpDecodeError, encode::Error as RmpEncodeError, from_slice};
use serde::{Deserialize, Serialize};
use crate::common::numerical::NumericalValue;
use crate::message_pack_format::envelope;
use super::{
CAPACITY_CACHE_LEN, Coin, KLL, MAX_CACHEABLE_K, MAX_LEVELS, checked_weighted_count,
compute_max_capacity,
};
const KLL_KIND_FAMILY: u8 = 0x06;
pub(crate) const KLL_KIND_COMPACT: &[u8] = &[KLL_KIND_FAMILY, 0x00];
pub(crate) const KLL_KIND_DYNAMIC: &[u8] = &[KLL_KIND_FAMILY, 0x01];
pub trait KllWireItem: Copy {
const ITEM_TYPE: &'static str;
}
impl KllWireItem for f64 {
const ITEM_TYPE: &'static str = "f64";
}
impl KllWireItem for i64 {
const ITEM_TYPE: &'static str = "i64";
}
#[derive(Debug, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub(crate) struct KllMetadata {
pub(crate) metadata_version: u8,
pub(crate) k: u32,
pub(crate) m: u32,
pub(crate) item_type: String,
#[serde(skip_serializing_if = "Option::is_none", default)]
pub(crate) seed: Option<u64>,
}
pub(crate) fn kll_metadata<T: KllWireItem>(k: u32, m: u32, seed: Option<u64>) -> KllMetadata {
KllMetadata {
metadata_version: 1,
k,
m,
item_type: T::ITEM_TYPE.to_string(),
seed,
}
}
#[derive(Debug, Serialize, Deserialize)]
pub(crate) struct KllCoinWire {
pub(crate) state: u64,
pub(crate) bit_cache: u64,
pub(crate) remaining_bits: u32,
}
#[derive(Debug, Serialize, Deserialize)]
pub(crate) struct KllPayload<T> {
pub(crate) levels: Vec<u32>,
pub(crate) items: Vec<T>,
pub(crate) coin: KllCoinWire,
}
pub(crate) fn validate_kll_payload<T>(
levels: &[u32],
items: &[T],
coin: &KllCoinWire,
) -> Result<usize, RmpDecodeError> {
if levels.len() < 2 {
return Err(RmpDecodeError::Uncategorized(format!(
"KLL payload: levels too short (len {}, need >= 2)",
levels.len()
)));
}
let num_levels = levels.len() - 1;
if num_levels > MAX_LEVELS {
return Err(RmpDecodeError::Uncategorized(format!(
"KLL payload: num_levels {num_levels} exceeds MAX_LEVELS {MAX_LEVELS}"
)));
}
if levels[0] != 0 {
return Err(RmpDecodeError::Uncategorized(format!(
"KLL payload: levels[0] must be 0, got {}",
levels[0]
)));
}
if levels.windows(2).any(|w| w[0] > w[1]) {
return Err(RmpDecodeError::Uncategorized(
"KLL payload: levels must be non-decreasing".to_string(),
));
}
if *levels.last().unwrap() as usize != items.len() {
return Err(RmpDecodeError::Uncategorized(format!(
"KLL payload: levels[last] {} != items.len() {}",
levels.last().unwrap(),
items.len()
)));
}
if coin.remaining_bits > u64::BITS {
return Err(RmpDecodeError::Uncategorized(format!(
"KLL payload: coin remaining_bits {} exceeds 64",
coin.remaining_bits
)));
}
let sizes: Vec<usize> = (0..num_levels)
.map(|h| {
let i = num_levels - 1 - h;
(levels[i + 1] - levels[i]) as usize
})
.collect();
if checked_weighted_count(&sizes).is_none() {
return Err(RmpDecodeError::Uncategorized(
"KLL payload: level layout overflows weighted count".to_string(),
));
}
Ok(num_levels)
}
pub(crate) fn split_and_validate_meta<'a, T: KllWireItem>(
bytes: &'a [u8],
expected_kind_id: &[u8],
) -> Result<(KllMetadata, &'a [u8]), RmpDecodeError> {
let (kind_id, metadata, payload) =
envelope::split(bytes).map_err(RmpDecodeError::Uncategorized)?;
if kind_id != expected_kind_id {
return Err(RmpDecodeError::Uncategorized(format!(
"KLL kind_id mismatch: stored {kind_id:?}, expected {expected_kind_id:?}"
)));
}
let meta: KllMetadata = from_slice(metadata)?;
if meta != kll_metadata::<T>(meta.k, meta.m, meta.seed) {
return Err(RmpDecodeError::Uncategorized(
"ASAPv1 KLL envelope: metadata mismatch".to_string(),
));
}
if meta.m < 2 || meta.m > meta.k || meta.k > MAX_CACHEABLE_K as u32 {
return Err(RmpDecodeError::Uncategorized(format!(
"ASAPv1 KLL envelope: k={}, m={} outside valid range (2 <= m <= k <= {MAX_CACHEABLE_K})",
meta.k, meta.m
)));
}
Ok((meta, payload))
}
impl<T> KLL<T>
where
T: NumericalValue + KllWireItem + Serialize + for<'de> Deserialize<'de>,
{
pub fn serialize_to_bytes(&self) -> Result<Vec<u8>, RmpEncodeError> {
let metadata =
rmp_serde::to_vec_named(&kll_metadata::<T>(self.k as u32, self.m as u32, self.seed))?;
let (state, bit_cache, remaining_bits) = self.co.to_wire();
let payload = rmp_serde::to_vec(&KllPayload {
levels: self.wire_levels(),
items: self.wire_items(),
coin: KllCoinWire {
state,
bit_cache,
remaining_bits,
},
})?;
Ok(envelope::encode(KLL_KIND_COMPACT, &metadata, &payload))
}
pub fn deserialize_from_bytes(bytes: &[u8]) -> Result<Self, RmpDecodeError> {
let (meta, payload_bytes) = split_and_validate_meta::<T>(bytes, KLL_KIND_COMPACT)?;
let payload: KllPayload<T> = from_slice(payload_bytes)?;
let num_levels = validate_kll_payload(&payload.levels, &payload.items, &payload.coin)?;
Self::from_wire_top_first(
meta.k as usize,
meta.m as usize,
meta.seed,
num_levels,
payload,
)
}
fn from_wire_top_first(
k: usize,
m: usize,
seed: Option<u64>,
num_levels: usize,
payload: KllPayload<T>,
) -> Result<Self, RmpDecodeError> {
let KllPayload {
levels,
items,
coin,
} = payload;
let max_cap = compute_max_capacity(k, m);
let total = items.len();
if total > max_cap {
return Err(RmpDecodeError::Uncategorized(format!(
"KLL payload: {total} items exceed max_capacity {max_cap} for k={k}, m={m}"
)));
}
let offset = max_cap - total;
let mut buf = vec![T::default(); max_cap].into_boxed_slice();
let mut internal_levels = vec![0usize; MAX_LEVELS + 1].into_boxed_slice();
let mut cursor = offset;
for h in 0..num_levels {
let top_i = num_levels - 1 - h;
let s = levels[top_i] as usize;
let e = levels[top_i + 1] as usize;
internal_levels[h] = cursor;
if h == 0 {
for (j, &v) in items[s..e].iter().rev().enumerate() {
buf[cursor + j] = v;
}
} else {
buf[cursor..cursor + (e - s)].copy_from_slice(&items[s..e]);
}
cursor += e - s;
}
internal_levels[num_levels] = cursor;
debug_assert_eq!(cursor, max_cap);
let mut sketch = KLL {
items: buf,
levels: internal_levels,
k,
m,
num_levels,
max_capacity: max_cap,
co: Coin::from_wire(coin.state, coin.bit_cache, coin.remaining_bits as u8),
seed,
capacity_cache: [0; CAPACITY_CACHE_LEN],
top_height: 0,
level0_capacity: 0,
merge_buf: Vec::with_capacity(k),
cdf_cache: None,
};
sketch.rebuild_capacity_cache();
sketch.ensure_levels_sorted();
Ok(sketch)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sketches::kll::KLL;
fn build_kll(k: i32, seed: u64, n: u64) -> KLL<f64> {
let mut sketch = KLL::<f64>::init_kll_with_seed(k, seed);
for v in 1..=n {
sketch.update(&(v as f64));
}
sketch
}
#[test]
fn kll_envelope_structure_and_round_trip() {
let sketch = build_kll(200, 42, 200_000);
let bytes = sketch.serialize_to_bytes().expect("serialize");
assert!(bytes.starts_with(envelope::MAGIC));
assert_eq!(bytes[6], envelope::VERSION);
assert_eq!(bytes[7], 2, "kind_id_len");
assert_eq!(&bytes[8..10], KLL_KIND_COMPACT);
let decoded = KLL::<f64>::deserialize_from_bytes(&bytes).expect("decode");
assert_eq!(
decoded.serialize_to_bytes().expect("re-serialize"),
bytes,
"KLL serialized bytes differed after round trip"
);
for &q in &[0.0, 0.01, 0.25, 0.5, 0.75, 0.99, 1.0] {
assert_eq!(
decoded.quantile(q),
sketch.quantile(q),
"quantile mismatch at q={q} after round trip"
);
}
}
#[test]
fn kll_empty_round_trip() {
let sketch = KLL::<f64>::init_kll_with_seed(200, 7);
let bytes = sketch.serialize_to_bytes().expect("serialize");
let decoded = KLL::<f64>::deserialize_from_bytes(&bytes).expect("decode");
assert_eq!(decoded.count(), 0);
assert_eq!(decoded.serialize_to_bytes().expect("re-serialize"), bytes);
}
#[test]
fn kll_i64_round_trip() {
let mut sketch = KLL::<i64>::init_kll_with_seed(200, 5);
for v in 1..=50_000i64 {
sketch.update(&v);
}
let bytes = sketch.serialize_to_bytes().expect("serialize");
assert_eq!(&bytes[8..10], KLL_KIND_COMPACT);
let decoded = KLL::<i64>::deserialize_from_bytes(&bytes).expect("decode");
assert_eq!(decoded.serialize_to_bytes().expect("re-serialize"), bytes);
assert_eq!(decoded.count(), sketch.count());
}
#[test]
fn kll_item_type_cross_rejection() {
let sketch = build_kll(200, 1, 1000);
let bytes = sketch.serialize_to_bytes().expect("serialize");
assert!(
KLL::<i64>::deserialize_from_bytes(&bytes).is_err(),
"f64 KLL bytes must be rejected by an i64 decoder"
);
}
#[test]
fn kll_metadata_rejects_unknown_keys() {
#[derive(Serialize)]
struct WithExtra {
metadata_version: u8,
k: u32,
m: u32,
item_type: String,
bogus_field: u8,
}
let extra = WithExtra {
metadata_version: 1,
k: 200,
m: 8,
item_type: "f64".to_string(),
bogus_field: 7,
};
let bytes = rmp_serde::to_vec_named(&extra).expect("encode");
assert!(
rmp_serde::from_slice::<KllMetadata>(&bytes).is_err(),
"an unexpected metadata key must be rejected"
);
}
#[test]
fn kll_rejects_inconsistent_levels() {
let metadata = rmp_serde::to_vec_named(&kll_metadata::<f64>(200, 8, None)).unwrap();
let payload = rmp_serde::to_vec(&KllPayload::<f64> {
levels: vec![0, 3],
items: vec![1.0, 2.0], coin: KllCoinWire {
state: 1,
bit_cache: 0,
remaining_bits: 0,
},
})
.unwrap();
let bytes = envelope::encode(KLL_KIND_COMPACT, &metadata, &payload);
assert!(
KLL::<f64>::deserialize_from_bytes(&bytes).is_err(),
"inconsistent level layout must be rejected, not panic"
);
}
#[test]
fn kll_rejects_out_of_range_k_m() {
let empty_payload = || {
rmp_serde::to_vec(&KllPayload::<f64> {
levels: vec![0, 0],
items: Vec::new(),
coin: KllCoinWire {
state: 1,
bit_cache: 0,
remaining_bits: 0,
},
})
.unwrap()
};
for (k, m) in [
(u32::MAX, u32::MAX),
(MAX_CACHEABLE_K as u32 + 1, 8),
(200, 1),
] {
let metadata = rmp_serde::to_vec_named(&kll_metadata::<f64>(k, m, None)).unwrap();
let bytes = envelope::encode(KLL_KIND_COMPACT, &metadata, &empty_payload());
assert!(
KLL::<f64>::deserialize_from_bytes(&bytes).is_err(),
"k={k}, m={m} must be rejected, not allocated"
);
}
}
#[test]
fn kll_seed_present_when_seeded_omitted_when_unseeded() {
let seeded = KLL::<f64>::init_kll_with_seed(200, 42);
let bytes = seeded.serialize_to_bytes().expect("serialize");
let (_k, meta_bytes, _p) = envelope::split(&bytes).expect("split");
let meta: KllMetadata = rmp_serde::from_slice(meta_bytes).expect("meta");
assert_eq!(meta.seed, Some(42), "seeded KLL must record its seed");
let mut unseeded = KLL::<f64>::init_kll(200); unseeded.update(&1.0);
let bytes = unseeded.serialize_to_bytes().expect("serialize");
let (_k, meta_bytes, _p) = envelope::split(&bytes).expect("split");
let meta: KllMetadata = rmp_serde::from_slice(meta_bytes).expect("meta");
assert_eq!(meta.seed, None, "unseeded KLL must omit the seed key");
}
#[test]
fn kll_unseeded_round_trip_byte_stable() {
let mut s = KLL::<f64>::init_kll(200); for v in 1..=2000u64 {
s.update(&(v as f64));
}
let bytes = s.serialize_to_bytes().expect("serialize");
let decoded = KLL::<f64>::deserialize_from_bytes(&bytes).expect("decode");
assert_eq!(
decoded.serialize_to_bytes().expect("re-serialize"),
bytes,
"unseeded KLL round trip not byte-stable"
);
assert_eq!(decoded.count(), s.count());
}
#[test]
fn kll_seed_survives_round_trip_so_clear_stays_deterministic() {
let mut src = KLL::<f64>::init_kll_with_seed(200, 42);
for v in 1..=5000u64 {
src.update(&(v as f64));
}
let mut a = KLL::<f64>::deserialize_from_bytes(&src.serialize_to_bytes().unwrap()).unwrap();
let mut b = KLL::<f64>::init_kll_with_seed(200, 42);
a.clear();
b.clear();
for v in 1..=3000u64 {
a.update(&(v as f64));
b.update(&(v as f64));
}
assert_eq!(
a.serialize_to_bytes().unwrap(),
b.serialize_to_bytes().unwrap(),
"decoded sketch lost its seed: clear() diverged from a fresh seeded sketch"
);
}
#[test]
fn kll_rejects_weighted_count_overflow() {
let num_levels = 61usize;
let mut levels = vec![16u32; num_levels + 1];
levels[0] = 0;
let metadata = rmp_serde::to_vec_named(&kll_metadata::<f64>(200, 8, None)).unwrap();
let payload = rmp_serde::to_vec(&KllPayload::<f64> {
levels,
items: vec![1.0; 16],
coin: KllCoinWire {
state: 1,
bit_cache: 0,
remaining_bits: 0,
},
})
.unwrap();
let bytes = envelope::encode(KLL_KIND_COMPACT, &metadata, &payload);
assert!(
KLL::<f64>::deserialize_from_bytes(&bytes).is_err(),
"a level layout that overflows the weighted count must be rejected"
);
}
#[test]
fn kll_dynamic_kind_id_rejected_by_compact() {
let metadata = rmp_serde::to_vec_named(&kll_metadata::<f64>(200, 8, None)).unwrap();
let payload = rmp_serde::to_vec(&KllPayload::<f64> {
levels: vec![0, 0],
items: Vec::<f64>::new(),
coin: KllCoinWire {
state: 1,
bit_cache: 0,
remaining_bits: 0,
},
})
.unwrap();
let bytes = envelope::encode(KLL_KIND_DYNAMIC, &metadata, &payload);
assert!(
KLL::<f64>::deserialize_from_bytes(&bytes).is_err(),
"dynamic kind_id must be rejected by the compact decoder"
);
}
}