use rmp_serde::{decode::Error as RmpDecodeError, encode::Error as RmpEncodeError, from_slice};
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
use crate::common::hash::check_matrix_rows;
use crate::message_pack_format::envelope;
use crate::{DataInput, HashProfile, SketchHasher, Vector2D};
use super::{Coco, CocoBucket};
const COCO_KIND: &[u8] = &[0x0c, 0x00];
#[derive(Debug, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
struct CocoMetadata {
metadata_version: u8,
hash_profile_id: String,
hash_algorithm: String,
seed_derivation: String,
input_encoding: String,
seed_list: Vec<u64>,
rows: u32,
cols: u32,
}
fn coco_metadata<H: HashProfile>(rows: u32, cols: u32) -> CocoMetadata {
CocoMetadata {
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(),
rows,
cols,
}
}
#[derive(Debug, Serialize, Deserialize)]
struct CocoPayload {
keys: Vec<Option<String>>,
values: Vec<u64>,
}
fn check_geometry(rows: usize, cols: usize) -> Result<(), String> {
if rows == 0 || cols == 0 {
return Err(format!(
"Coco table dimensions must be non-zero: rows={rows}, cols={cols}"
));
}
check_matrix_rows("Coco table", rows)?;
Ok(())
}
fn check_placement<H: SketchHasher>(keys: &[Option<String>], cols: usize) -> Result<(), String> {
let mut seen: HashSet<&str> = HashSet::new();
for (i, slot) in keys.iter().enumerate() {
let Some(key) = slot.as_deref() else { continue };
let (row, col) = (i / cols, i % cols);
let mapped = H::hash64_seeded(row, &DataInput::Str(key)) as usize % cols;
if mapped != col {
return Err(format!(
"Coco bucket ({row}, {col}) holds a key that maps to column {mapped}"
));
}
if !seen.insert(key) {
return Err(format!(
"Coco bucket ({row}, {col}) repeats a key already stored elsewhere in the table"
));
}
}
Ok(())
}
impl<H: SketchHasher + HashProfile> Coco<H> {
pub fn serialize_to_bytes(&self) -> Result<Vec<u8>, RmpEncodeError> {
let (rows, cols) = (self.d, self.w);
check_geometry(rows, cols).map_err(RmpEncodeError::Syntax)?;
let buckets = self.table.as_slice();
if self.table.rows() != rows
|| self.table.cols() != cols
|| buckets.len() != rows.saturating_mul(cols)
{
return Err(RmpEncodeError::Syntax(format!(
"ASAPv1 Coco envelope: table {}x{} ({} buckets) != declared {rows}x{cols}",
self.table.rows(),
self.table.cols(),
buckets.len()
)));
}
let keys: Vec<Option<String>> = buckets.iter().map(|b| b.full_key.clone()).collect();
check_placement::<H>(&keys, cols).map_err(RmpEncodeError::Syntax)?;
let to_u32 = |v: usize, name: &str| {
u32::try_from(v).map_err(|_| {
RmpEncodeError::Syntax(format!("ASAPv1 Coco envelope: {name} {v} exceeds u32"))
})
};
let metadata = rmp_serde::to_vec_named(&coco_metadata::<H>(
to_u32(rows, "rows")?,
to_u32(cols, "cols")?,
))?;
let payload = rmp_serde::to_vec(&CocoPayload {
keys,
values: buckets.iter().map(|b| b.val).collect(),
})?;
Ok(envelope::encode(COCO_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 != COCO_KIND {
return Err(RmpDecodeError::Uncategorized(format!(
"Coco kind_id mismatch: stored {kind_id:?}, expected {COCO_KIND:?}"
)));
}
let meta: CocoMetadata = from_slice(metadata)?;
if meta != coco_metadata::<H>(meta.rows, meta.cols) {
return Err(RmpDecodeError::Uncategorized(
"ASAPv1 Coco envelope: metadata mismatch".to_string(),
));
}
let (rows, cols) = (meta.rows as usize, meta.cols as usize);
check_geometry(rows, cols).map_err(RmpDecodeError::Uncategorized)?;
let mut p: CocoPayload = from_slice(payload)?;
let expected = rows.saturating_mul(cols);
if p.keys.len() != expected || p.values.len() != expected {
return Err(RmpDecodeError::Uncategorized(format!(
"Coco payload lengths (keys {}, values {}) != rows*cols {expected}",
p.keys.len(),
p.values.len()
)));
}
if let Some(i) = (0..expected).find(|&i| p.keys[i].is_none() && p.values[i] != 0) {
return Err(RmpDecodeError::Uncategorized(format!(
"Coco bucket {i} is unoccupied but carries value {}",
p.values[i]
)));
}
check_placement::<H>(&p.keys, cols).map_err(RmpDecodeError::Uncategorized)?;
let table = Vector2D::from_fn(rows, cols, |r, c| {
let i = r * cols + c;
CocoBucket {
full_key: p.keys[i].take(),
val: p.values[i],
}
});
Ok(Coco {
w: cols,
d: rows,
table,
_hasher: std::marker::PhantomData,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{CANONICAL_HASH_SEED, DataInput, DefaultXxHasher, MATRIX_MAX_ROWS, Vector2D};
#[test]
fn coco_rejects_too_many_rows() {
let rows = MATRIX_MAX_ROWS + 1;
assert!(
Coco::<DefaultXxHasher>::init_with_size(8, rows)
.serialize_to_bytes()
.is_err(),
"a table past MATRIX_MAX_ROWS must not serialize"
);
let metadata =
rmp_serde::to_vec_named(&coco_metadata::<DefaultXxHasher>(rows as u32, 8)).unwrap();
let payload = rmp_serde::to_vec(&CocoPayload {
keys: vec![None; rows * 8],
values: vec![0; rows * 8],
})
.unwrap();
let bytes = envelope::encode(COCO_KIND, &metadata, &payload);
let problem = Coco::<DefaultXxHasher>::deserialize_from_bytes(&bytes)
.expect_err("rows past MATRIX_MAX_ROWS must be rejected")
.to_string();
assert!(problem.contains("MATRIX_MAX_ROWS"), "got {problem}");
assert!(
Coco::<DefaultXxHasher>::init_with_size(8, MATRIX_MAX_ROWS)
.serialize_to_bytes()
.is_ok()
);
}
fn cells<H: SketchHasher>(sketch: &Coco<H>) -> Vec<(Option<String>, u64)> {
sketch
.table
.as_slice()
.iter()
.map(|b| (b.full_key.clone(), b.val))
.collect()
}
#[test]
fn coco_round_trip_serialization() {
let mut sketch: Coco = Coco::init_with_size(8, 4);
sketch.insert("19.98.10.26|80", 521);
sketch.insert("34.52.73.17|118", 856);
let encoded = sketch.serialize_to_bytes().expect("serialize Coco");
assert!(encoded.starts_with(b"ASAPv1"));
assert_eq!(&encoded[7..10], &[2u8, 0x0c, 0x00]);
let decoded: Coco = Coco::deserialize_from_bytes(&encoded).expect("deserialize Coco");
assert_eq!(decoded.w, 8);
assert_eq!(decoded.d, 4);
assert_eq!(cells(&sketch), cells(&decoded));
assert_eq!(decoded.estimate_key("19.98.10.26|80"), 521);
}
#[test]
fn coco_rejects_foreign_kind_id() {
let cms = crate::CountMin::<Vector2D<i64>, crate::RegularPath>::with_dimensions(4, 8);
let cms_bytes = cms.serialize_to_bytes().expect("serialize CMS");
assert!(
Coco::<DefaultXxHasher>::deserialize_from_bytes(&cms_bytes).is_err(),
"CMS bytes must not decode as a Coco"
);
}
#[test]
fn coco_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>,
rows: u32,
cols: u32,
bogus_field: u8, }
let m = coco_metadata::<DefaultXxHasher>(4, 8);
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(),
rows: m.rows,
cols: m.cols,
bogus_field: 7,
};
let bytes = rmp_serde::to_vec_named(&extra).unwrap();
assert!(rmp_serde::from_slice::<CocoMetadata>(&bytes).is_err());
}
#[test]
fn coco_metadata_rejects_a_missing_cols_key() {
#[derive(Serialize)]
struct WithoutCols {
metadata_version: u8,
hash_profile_id: String,
hash_algorithm: String,
seed_derivation: String,
input_encoding: String,
seed_list: Vec<u64>,
rows: u32,
}
let m = coco_metadata::<DefaultXxHasher>(4, 8);
let without = WithoutCols {
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(),
rows: m.rows,
};
let bytes = rmp_serde::to_vec_named(&without).unwrap();
assert!(rmp_serde::from_slice::<CocoMetadata>(&bytes).is_err());
}
#[test]
fn coco_rejects_zero_dimension_payload() {
let metadata = rmp_serde::to_vec_named(&coco_metadata::<DefaultXxHasher>(4, 0)).unwrap();
let payload = rmp_serde::to_vec(&CocoPayload {
keys: Vec::new(),
values: Vec::new(),
})
.unwrap();
let bytes = envelope::encode(COCO_KIND, &metadata, &payload);
assert!(
Coco::<DefaultXxHasher>::deserialize_from_bytes(&bytes).is_err(),
"zero-dimension metadata must be rejected, not panic"
);
}
#[test]
fn coco_rejects_dimension_length_mismatch() {
let metadata = rmp_serde::to_vec_named(&coco_metadata::<DefaultXxHasher>(
MATRIX_MAX_ROWS as u32,
1 << 24,
))
.unwrap();
let payload = rmp_serde::to_vec(&CocoPayload {
keys: vec![None, None, None],
values: vec![0, 0, 0],
})
.unwrap();
let bytes = envelope::encode(COCO_KIND, &metadata, &payload);
assert!(
Coco::<DefaultXxHasher>::deserialize_from_bytes(&bytes).is_err(),
"payload lengths must match rows*cols"
);
}
#[test]
fn coco_rejects_mass_under_an_unoccupied_bucket() {
let metadata = rmp_serde::to_vec_named(&coco_metadata::<DefaultXxHasher>(1, 2)).unwrap();
let payload = rmp_serde::to_vec(&CocoPayload {
keys: vec![None, None],
values: vec![0, 9],
})
.unwrap();
let bytes = envelope::encode(COCO_KIND, &metadata, &payload);
assert!(
Coco::<DefaultXxHasher>::deserialize_from_bytes(&bytes).is_err(),
"a nil key with a non-zero value must be rejected"
);
}
fn mapped_col(row: usize, key: &str, cols: usize) -> usize {
DefaultXxHasher::hash64_seeded(row, &DataInput::Str(key)) as usize % cols
}
fn craft(rows: usize, cols: usize, cells: Vec<(Option<String>, u64)>) -> Vec<u8> {
let metadata =
rmp_serde::to_vec_named(&coco_metadata::<DefaultXxHasher>(rows as u32, cols as u32))
.unwrap();
let payload = rmp_serde::to_vec(&CocoPayload {
keys: cells.iter().map(|(k, _)| k.clone()).collect(),
values: cells.iter().map(|(_, v)| *v).collect(),
})
.unwrap();
envelope::encode(COCO_KIND, &metadata, &payload)
}
#[test]
fn coco_rejects_a_key_in_a_bucket_it_does_not_hash_to() {
let (rows, cols) = (1, 2);
let key = "flow::misplaced";
let wrong = 1 - mapped_col(0, key, cols);
let mut cells = vec![(None, 0), (None, 0)];
cells[wrong] = (Some(key.to_string()), 11);
let problem = Coco::<DefaultXxHasher>::deserialize_from_bytes(&craft(rows, cols, cells))
.expect_err("a misplaced key must be rejected")
.to_string();
assert!(problem.contains("maps to column"), "got {problem}");
}
#[test]
fn coco_rejects_a_key_stored_in_two_buckets() {
let (rows, cols) = (2, 2);
let key = "flow::twice";
let mut cells = vec![(None, 0); rows * cols];
for row in 0..rows {
cells[row * cols + mapped_col(row, key, cols)] = (Some(key.to_string()), u64::MAX / 2);
}
assert_eq!(
cells.iter().filter(|(k, _)| k.is_some()).count(),
2,
"the two rows must place the key in distinct buckets"
);
let problem = Coco::<DefaultXxHasher>::deserialize_from_bytes(&craft(rows, cols, cells))
.expect_err("a duplicated key must be rejected")
.to_string();
assert!(problem.contains("repeats a key"), "got {problem}");
}
#[test]
fn coco_rejects_serializing_a_geometry_mismatch() {
let mut sketch: Coco = Coco::init_with_size(8, 4);
sketch.w = 16;
assert!(
sketch.serialize_to_bytes().is_err(),
"a table whose bucket count disagrees with its geometry must not serialize"
);
let empty: Coco = Coco::init_with_size(8, 0);
assert!(
empty.serialize_to_bytes().is_err(),
"a zero-dimension geometry must not serialize"
);
}
#[test]
fn coco_rejects_serializing_a_hand_built_table() {
let key = "flow::misplaced";
let mut misplaced: Coco = Coco::init_with_size(2, 1);
let wrong = 1 - mapped_col(0, key, 2);
misplaced.table.as_mut_slice()[wrong] = CocoBucket {
full_key: Some(key.to_string()),
val: 11,
};
let problem = misplaced
.serialize_to_bytes()
.expect_err("a misplaced key must not serialize")
.to_string();
assert!(problem.contains("maps to column"), "got {problem}");
let (rows, cols) = (2, 2);
let mut twice: Coco = Coco::init_with_size(cols, rows);
for row in 0..rows {
twice.table.as_mut_slice()[row * cols + mapped_col(row, key, cols)] = CocoBucket {
full_key: Some(key.to_string()),
val: 7,
};
}
let problem = twice
.serialize_to_bytes()
.expect_err("a key in two buckets must not serialize")
.to_string();
assert!(problem.contains("repeats a key"), "got {problem}");
}
#[test]
fn coco_empty_buckets_are_nil_and_never_collide_with_an_empty_key() {
let all_empty: Coco = Coco::init_with_size(4, 2);
let encoded = all_empty.serialize_to_bytes().expect("serialize");
let (_, _, payload) = envelope::split(&encoded).expect("split");
assert_eq!(
payload.iter().filter(|&&b| b == 0xc0).count(),
8,
"nil per bucket"
);
assert!(!payload.contains(&0xa0), "no empty-string key is emitted");
let decoded: Coco = Coco::deserialize_from_bytes(&encoded).expect("deserialize");
assert_eq!(cells(&all_empty), cells(&decoded));
assert!(decoded.recorded_flows().next().is_none());
let mut with_empty_key: Coco = Coco::init_with_size(4, 2);
with_empty_key.insert("", 5);
let encoded = with_empty_key.serialize_to_bytes().expect("serialize");
let (_, _, payload) = envelope::split(&encoded).expect("split");
assert!(payload.contains(&0xa0), "the inserted empty key is a str");
let decoded: Coco = Coco::deserialize_from_bytes(&encoded).expect("deserialize");
assert_eq!(cells(&with_empty_key), cells(&decoded));
assert_eq!(decoded.recorded_flows().count(), 1);
assert_eq!(decoded.estimate_key(""), 5);
}
#[test]
fn coco_mixed_occupancy_round_trips() {
let mut sketch: Coco = Coco::init_with_size(16, 3);
for i in 0..10u64 {
sketch.insert(&format!("flow::{i}"), i + 1);
}
let occupied = sketch.recorded_flows().count();
assert!(occupied > 0 && occupied < 16 * 3, "mix both bucket states");
let encoded = sketch.serialize_to_bytes().expect("serialize");
let decoded: Coco = Coco::deserialize_from_bytes(&encoded).expect("deserialize");
assert_eq!(cells(&sketch), cells(&decoded));
for i in 0..10u64 {
let key = format!("flow::{i}");
assert_eq!(sketch.estimate_key(&key), decoded.estimate_key(&key));
}
}
#[test]
fn coco_decoded_sketch_reserializes_byte_identically() {
let mut sketch: Coco = Coco::init_with_size(16, 3);
for i in 0..12u64 {
sketch.insert(&format!("fam{}|item{i}", i % 4), i % 5 + 1);
}
let encoded = sketch.serialize_to_bytes().expect("serialize");
let decoded: Coco = Coco::deserialize_from_bytes(&encoded).expect("deserialize");
assert_eq!(encoded, decoded.serialize_to_bytes().expect("re-serialize"));
}
#[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: &crate::HeapItem) -> u64 {
DefaultXxHasher::hash_item64_seeded(d, key)
}
fn hash_item128_seeded(d: usize, key: &crate::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 coco_custom_hasher_profile_round_trips_and_is_self_describing() {
let mut alt: Coco<AltHasher> = Coco::init_with_size(8, 4);
let mut std: Coco = Coco::init_with_size(8, 4);
alt.insert("flow::42", 7);
alt.insert("flow::7", 3);
std.insert("flow::42", 7);
std.insert("flow::7", 3);
let alt_bytes = alt.serialize_to_bytes().expect("alt serialize");
let decoded = Coco::<AltHasher>::deserialize_from_bytes(&alt_bytes).expect("alt decode");
assert_eq!(cells(&alt), cells(&decoded));
let std_bytes = std.serialize_to_bytes().expect("std serialize");
assert_ne!(alt_bytes, std_bytes);
assert!(
Coco::<DefaultXxHasher>::deserialize_from_bytes(&alt_bytes).is_err(),
"standard-profile decode must reject custom-profile bytes"
);
}
}