use crate::cache::CacheKey;
use crate::error::{Result, SomaError};
use ciborium::value::Value as Cbor;
use serde::Serialize;
pub fn canonical_bytes<T: Serialize + ?Sized>(value: &T) -> Result<Vec<u8>> {
let mut raw = Vec::new();
ciborium::ser::into_writer(value, &mut raw)
.map_err(|e| SomaError::Cache(format!("not canonically serializable: {e}")))?;
let decoded: Cbor = ciborium::de::from_reader(raw.as_slice())
.map_err(|e| SomaError::Cache(format!("canonical re-decode failed: {e}")))?;
let canon = canonicalize(decoded)?;
encode(&canon)
}
pub fn hash_canonical<T: Serialize + ?Sized>(value: &T) -> Result<CacheKey> {
Ok(CacheKey::hash_data(&canonical_bytes(value)?))
}
fn encode(v: &Cbor) -> Result<Vec<u8>> {
let mut buf = Vec::new();
ciborium::ser::into_writer(v, &mut buf)
.map_err(|e| SomaError::Cache(format!("canonical encode failed: {e}")))?;
Ok(buf)
}
fn canonicalize(v: Cbor) -> Result<Cbor> {
Ok(match v {
Cbor::Float(f) => Cbor::Float(canonical_float(f)),
Cbor::Array(items) => {
Cbor::Array(items.into_iter().map(canonicalize).collect::<Result<_>>()?)
}
Cbor::Map(entries) => {
let mut encoded: Vec<(Vec<u8>, Cbor, Cbor)> = Vec::with_capacity(entries.len());
for (key, value) in entries {
let key = canonicalize(key)?;
let value = canonicalize(value)?;
let key_bytes = encode(&key)?;
encoded.push((key_bytes, key, value));
}
encoded.sort_by(|a, b| a.0.cmp(&b.0));
for pair in encoded.windows(2) {
if pair[0].0 == pair[1].0 {
return Err(SomaError::Cache(
"duplicate map key in canonical encoding".into(),
));
}
}
Cbor::Map(encoded.into_iter().map(|(_, k, v)| (k, v)).collect())
}
Cbor::Tag(tag, inner) => Cbor::Tag(tag, Box::new(canonicalize(*inner)?)),
other => other,
})
}
fn canonical_float(f: f64) -> f64 {
if f.is_nan() {
f64::NAN
} else if f == 0.0 {
0.0
} else {
f
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn hashmap_encoding_is_order_independent() {
let mut reference: Option<Vec<u8>> = None;
for i in 0..100 {
let mut map = HashMap::new();
let keys = ["alpha", "beta", "gamma", "delta", "epsilon"];
for (j, _k) in keys.iter().enumerate() {
let idx = (i + j) % keys.len();
map.insert(keys[idx].to_string(), idx as i64);
}
let bytes = canonical_bytes(&map).unwrap();
match &reference {
None => reference = Some(bytes),
Some(r) => assert_eq!(r, &bytes, "iteration {i} diverged"),
}
}
}
#[test]
fn nested_structures_canonicalize() {
#[derive(serde::Serialize)]
struct Config {
name: String,
params: HashMap<String, f64>,
layers: Vec<u32>,
}
let mut params = HashMap::new();
params.insert("lr".into(), 0.001);
params.insert("momentum".into(), 0.9);
let a = Config {
name: "m".into(),
params: params.clone(),
layers: vec![64, 32],
};
let b = Config {
name: "m".into(),
params,
layers: vec![64, 32],
};
assert_eq!(hash_canonical(&a).unwrap(), hash_canonical(&b).unwrap());
}
#[test]
fn negative_zero_normalizes() {
assert_eq!(
canonical_bytes(&(-0.0f64)).unwrap(),
canonical_bytes(&0.0f64).unwrap()
);
}
#[test]
fn nan_payloads_collapse_to_one_encoding() {
let quiet = f64::NAN;
let payload = f64::from_bits(0x7ff8_0000_0000_0001);
assert!(payload.is_nan());
assert_eq!(
canonical_bytes(&quiet).unwrap(),
canonical_bytes(&payload).unwrap()
);
}
#[test]
fn distinct_values_distinct_hashes() {
assert_ne!(
hash_canonical(&vec![1.0f64, 2.0]).unwrap(),
hash_canonical(&vec![2.0f64, 1.0]).unwrap()
);
assert_ne!(hash_canonical(&"a").unwrap(), hash_canonical(&"b").unwrap());
}
#[test]
fn float_and_int_do_not_collide() {
assert_ne!(
canonical_bytes(&1u64).unwrap(),
canonical_bytes(&1.0f64).unwrap()
);
}
}