use std::{io::Cursor, str::Utf8Error};
use roaring::{RoaringBitmap, RoaringTreemap};
use snafu::{Backtrace, ResultExt, Snafu};
use crate::coverage::{Coverage, EntityCoverage, EntityIdentity, EntityIdentityError, EntityValue};
const ENTITY_COVERAGE_MAGIC: &[u8; 8] = b"TSTECOV2";
const ENTITY_VALUE_UTF8: u8 = 1;
const ENTITY_VALUE_INT32: u8 = 2;
const ENTITY_VALUE_INT64: u8 = 3;
const ENTITY_VALUE_UINT64: u8 = 4;
#[derive(Debug, Snafu)]
#[non_exhaustive]
pub enum CoverageCodecError {
#[snafu(display("Failed to serialize roaring bitmap: {source}"))]
BitmapSerialization {
source: std::io::Error,
backtrace: Backtrace,
},
#[snafu(display("Failed to deserialize roaring bitmap: {source}"))]
BitmapDeserialization {
source: std::io::Error,
backtrace: Backtrace,
},
#[snafu(display("Entity coverage {field} is too large to serialize"))]
LengthOverflow {
field: &'static str,
backtrace: Backtrace,
},
#[snafu(display("Invalid entity coverage payload identifier"))]
InvalidEntityCoverageMagic {
backtrace: Backtrace,
},
#[snafu(display("Truncated entity coverage payload"))]
TruncatedPayload {
backtrace: Backtrace,
},
#[snafu(display("Invalid entity coverage {field}"))]
InvalidLength {
field: &'static str,
backtrace: Backtrace,
},
#[snafu(display("Invalid entity identity string: {source}"))]
InvalidEntityUtf8 {
source: Utf8Error,
backtrace: Backtrace,
},
#[snafu(display("Unknown entity identity value type tag: {tag}"))]
UnknownEntityValueTag {
tag: u8,
backtrace: Backtrace,
},
#[snafu(display("Invalid entity identity: {source}"))]
InvalidEntityIdentity {
source: EntityIdentityError,
backtrace: Backtrace,
},
#[snafu(display("Duplicate entity identity in coverage payload: {identity:?}"))]
DuplicateEntityIdentity {
identity: EntityIdentity,
backtrace: Backtrace,
},
#[snafu(display("Malformed nested entity coverage: {source}"))]
MalformedNestedCoverage {
#[snafu(source(from(CoverageCodecError, Box::new)), backtrace)]
source: Box<CoverageCodecError>,
},
#[snafu(display("Trailing bytes after coverage payload"))]
TrailingBytes {
backtrace: Backtrace,
},
}
pub fn coverage_to_bytes(cov: &Coverage) -> Result<Vec<u8>, CoverageCodecError> {
let mut out = Vec::new();
{
let mut w = Cursor::new(&mut out);
cov.present()
.serialize_into(&mut w)
.context(BitmapSerializationSnafu)?;
}
Ok(out)
}
pub fn coverage_from_bytes(bytes: &[u8]) -> Result<Coverage, CoverageCodecError> {
let mut r = Cursor::new(bytes);
let present = RoaringTreemap::deserialize_from(&mut r).context(BitmapDeserializationSnafu)?;
if r.position() != bytes.len() as u64 {
return TrailingBytesSnafu.fail();
}
Ok(Coverage::from_treemap(present))
}
pub fn entity_coverage_to_bytes(coverage: &EntityCoverage) -> Result<Vec<u8>, CoverageCodecError> {
let entity_count = u32::try_from(coverage.identity_count()).map_err(|_| {
LengthOverflowSnafu {
field: "entity count",
}
.build()
})?;
let mut out = Vec::new();
out.extend_from_slice(ENTITY_COVERAGE_MAGIC);
out.extend_from_slice(&entity_count.to_be_bytes());
for (identity, nested) in coverage.iter() {
let component_count = u32::try_from(identity.components().len()).map_err(|_| {
LengthOverflowSnafu {
field: "identity component count",
}
.build()
})?;
out.extend_from_slice(&component_count.to_be_bytes());
for component in identity.components() {
match component {
EntityValue::Utf8(value) => {
out.push(ENTITY_VALUE_UTF8);
let component_len = u64::try_from(value.len()).map_err(|_| {
LengthOverflowSnafu {
field: "identity component length",
}
.build()
})?;
out.extend_from_slice(&component_len.to_be_bytes());
out.extend_from_slice(value.as_bytes());
}
EntityValue::Int32(value) => {
out.push(ENTITY_VALUE_INT32);
out.extend_from_slice(&value.to_be_bytes());
}
EntityValue::Int64(value) => {
out.push(ENTITY_VALUE_INT64);
out.extend_from_slice(&value.to_be_bytes());
}
EntityValue::UInt64(value) => {
out.push(ENTITY_VALUE_UINT64);
out.extend_from_slice(&value.to_be_bytes());
}
}
}
let nested_bytes = canonical_nested_coverage_to_bytes(nested)?;
let nested_len = u64::try_from(nested_bytes.len()).map_err(|_| {
LengthOverflowSnafu {
field: "nested coverage length",
}
.build()
})?;
out.extend_from_slice(&nested_len.to_be_bytes());
out.extend_from_slice(&nested_bytes);
}
Ok(out)
}
pub fn entity_coverage_from_bytes(bytes: &[u8]) -> Result<EntityCoverage, CoverageCodecError> {
let mut remaining = bytes;
if take(&mut remaining, ENTITY_COVERAGE_MAGIC.len())? != ENTITY_COVERAGE_MAGIC {
return InvalidEntityCoverageMagicSnafu.fail();
}
let entity_count = read_u32(&mut remaining)? as usize;
let mut coverage = EntityCoverage::empty();
for _ in 0..entity_count {
let component_count = read_u32(&mut remaining)? as usize;
if component_count > remaining.len().saturating_sub(8) / 5 {
return InvalidLengthSnafu {
field: "identity component count",
}
.fail();
}
let mut components = Vec::new();
for _ in 0..component_count {
let tag = take(&mut remaining, 1)?[0];
let component = match tag {
ENTITY_VALUE_UTF8 => {
let component_len = read_u64(&mut remaining)?;
let component_bytes =
take_declared(&mut remaining, component_len, "identity component length")?;
let component =
std::str::from_utf8(component_bytes).context(InvalidEntityUtf8Snafu)?;
EntityValue::Utf8(component.to_owned())
}
ENTITY_VALUE_INT32 => EntityValue::Int32(read_i32(&mut remaining)?),
ENTITY_VALUE_INT64 => EntityValue::Int64(read_i64(&mut remaining)?),
ENTITY_VALUE_UINT64 => EntityValue::UInt64(read_u64(&mut remaining)?),
tag => return UnknownEntityValueTagSnafu { tag }.fail(),
};
components.push(component);
}
let identity = EntityIdentity::try_new(components).context(InvalidEntityIdentitySnafu)?;
if coverage.get(&identity).is_some() {
return DuplicateEntityIdentitySnafu { identity }.fail();
}
let nested_len = read_u64(&mut remaining)?;
let nested_bytes = take_declared(&mut remaining, nested_len, "nested coverage length")?;
let nested = coverage_from_bytes(nested_bytes).context(MalformedNestedCoverageSnafu)?;
coverage.union_coverage(identity, nested);
}
if !remaining.is_empty() {
return TrailingBytesSnafu.fail();
}
Ok(coverage)
}
fn canonical_nested_coverage_to_bytes(coverage: &Coverage) -> Result<Vec<u8>, CoverageCodecError> {
let present = RoaringTreemap::from_bitmaps(
coverage
.present()
.bitmaps()
.filter(|(_, bitmap)| !bitmap.is_empty())
.map(|(key, bitmap)| {
let mut canonical = RoaringBitmap::new();
let mut ranges = bitmap.iter();
while let Some(range) = ranges.next_range() {
canonical.insert_range(range);
}
canonical.optimize();
(key, canonical)
}),
);
coverage_to_bytes(&Coverage::from_treemap(present))
}
fn take<'a>(remaining: &mut &'a [u8], len: usize) -> Result<&'a [u8], CoverageCodecError> {
if remaining.len() < len {
return TruncatedPayloadSnafu.fail();
}
let (value, rest) = remaining.split_at(len);
*remaining = rest;
Ok(value)
}
fn take_declared<'a>(
remaining: &mut &'a [u8],
len: u64,
field: &'static str,
) -> Result<&'a [u8], CoverageCodecError> {
let len = usize::try_from(len).map_err(|_| InvalidLengthSnafu { field }.build())?;
if len > remaining.len() {
return InvalidLengthSnafu { field }.fail();
}
take(remaining, len)
}
fn read_u32(remaining: &mut &[u8]) -> Result<u32, CoverageCodecError> {
let mut encoded = [0; 4];
encoded.copy_from_slice(take(remaining, 4)?);
Ok(u32::from_be_bytes(encoded))
}
fn read_i32(remaining: &mut &[u8]) -> Result<i32, CoverageCodecError> {
let mut encoded = [0; 4];
encoded.copy_from_slice(take(remaining, 4)?);
Ok(i32::from_be_bytes(encoded))
}
fn read_i64(remaining: &mut &[u8]) -> Result<i64, CoverageCodecError> {
let mut encoded = [0; 8];
encoded.copy_from_slice(take(remaining, 8)?);
Ok(i64::from_be_bytes(encoded))
}
fn read_u64(remaining: &mut &[u8]) -> Result<u64, CoverageCodecError> {
let mut encoded = [0; 8];
encoded.copy_from_slice(take(remaining, 8)?);
Ok(u64::from_be_bytes(encoded))
}
#[cfg(test)]
mod tests {
use super::*;
const ROARING_ZERO: &[u8] = &[
1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0x3a, 0x30, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 16, 0, 0, 0, 0, 0, ];
const ROARING_MAX: &[u8] = &[
1, 0, 0, 0, 0, 0, 0, 0, 0xff, 0xff, 0xff, 0xff, 0x3a, 0x30, 0, 0, 1, 0, 0, 0, 0xff, 0xff, 0, 0, 16, 0, 0, 0, 0xff, 0xff, ];
fn identity(components: &[&str]) -> EntityIdentity {
EntityIdentity::try_new(
components
.iter()
.map(|component| EntityValue::from(*component))
.collect(),
)
.unwrap()
}
#[test]
fn round_trip_empty_and_non_empty() {
let cov_empty = Coverage::empty();
let bytes = coverage_to_bytes(&cov_empty).expect("serialize empty");
let restored = coverage_from_bytes(&bytes).expect("deserialize empty");
assert_eq!(cov_empty.cardinality(), restored.cardinality());
let cov = Coverage::from_iter(vec![1u64, 2, 3, u64::MAX]);
let bytes = coverage_to_bytes(&cov).expect("serialize non-empty");
let restored = coverage_from_bytes(&bytes).expect("deserialize non-empty");
assert_eq!(cov.present(), restored.present());
}
#[test]
fn deserialize_rejects_invalid_bytes() {
let bad = b"not a roaring bitmap";
let err = coverage_from_bytes(bad).unwrap_err();
match err {
CoverageCodecError::BitmapDeserialization { .. } => {}
_ => panic!("expected deserialize error"),
}
}
#[test]
fn deserialize_rejects_trailing_valid_payload() {
let mut bytes = coverage_to_bytes(&Coverage::empty()).unwrap();
bytes.extend_from_slice(&coverage_to_bytes(&Coverage::from_iter([1u64])).unwrap());
let err = coverage_from_bytes(&bytes).unwrap_err();
assert!(matches!(err, CoverageCodecError::TrailingBytes { .. }));
}
#[test]
fn serialize_reports_io_error() {
struct FailingWriter;
impl std::io::Write for FailingWriter {
fn write(&mut self, _buf: &[u8]) -> std::io::Result<usize> {
Err(std::io::Error::other("fail"))
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
let cov = Coverage::from_iter(vec![1u64]);
let err = {
let mut w = FailingWriter;
cov.present()
.serialize_into(&mut w)
.context(BitmapSerializationSnafu)
.unwrap_err()
};
match err {
CoverageCodecError::BitmapSerialization { .. } => {}
_ => panic!("expected serialize error"),
}
}
#[test]
fn entity_coverage_round_trips_empty_and_one_identity() {
let empty_bytes = entity_coverage_to_bytes(&EntityCoverage::empty()).unwrap();
assert_eq!(
empty_bytes,
[ENTITY_COVERAGE_MAGIC.as_slice(), &[0; 4]].concat()
);
assert_eq!(
entity_coverage_from_bytes(&empty_bytes).unwrap(),
EntityCoverage::empty()
);
let entity = identity(&["venue", "symbol"]);
let mut coverage = EntityCoverage::empty();
coverage.union_coverage(entity.clone(), Coverage::empty());
let empty_identity_bytes = entity_coverage_to_bytes(&coverage).unwrap();
assert_ne!(empty_identity_bytes, empty_bytes);
assert_eq!(
entity_coverage_from_bytes(&empty_identity_bytes).unwrap(),
coverage
);
coverage.union_coverage(entity, [0, u64::MAX].into_iter().collect());
let bytes = entity_coverage_to_bytes(&coverage).unwrap();
assert_eq!(entity_coverage_from_bytes(&bytes).unwrap(), coverage);
}
#[test]
fn entity_coverage_keeps_composite_identities_and_intervals_independent() {
let first = identity(&["a", "b:c"]);
let second = identity(&["a:b", "c"]);
let mut coverage = EntityCoverage::empty();
coverage.union_coverage(first.clone(), [7].into_iter().collect());
coverage.union_coverage(second.clone(), [7].into_iter().collect());
let restored =
entity_coverage_from_bytes(&entity_coverage_to_bytes(&coverage).unwrap()).unwrap();
assert_eq!(restored.get(&first).unwrap().cardinality(), 1);
assert_eq!(restored.get(&second).unwrap().cardinality(), 1);
assert_eq!(restored.cardinality(), 2);
}
#[test]
fn entity_coverage_serialization_uses_canonical_identity_order() {
let first = identity(&["A"]);
let second = identity(&["B"]);
let mut forward = EntityCoverage::empty();
forward.union_coverage(first.clone(), [1].into_iter().collect());
forward.union_coverage(second.clone(), [2].into_iter().collect());
let mut reverse = EntityCoverage::empty();
reverse.union_coverage(second, [2].into_iter().collect());
reverse.union_coverage(first, [1].into_iter().collect());
assert_eq!(
entity_coverage_to_bytes(&forward).unwrap(),
entity_coverage_to_bytes(&reverse).unwrap()
);
}
#[test]
fn entity_coverage_serialization_canonicalizes_roaring_storage() {
let partition = 1u64 << 32;
let inserted: Coverage = (1..=3).chain(partition + 1..=partition + 5_000).collect();
let mut ranged = RoaringTreemap::new();
ranged.insert_range(1..=3);
ranged.insert_range(partition + 1..=partition + 5_000);
let ranged = Coverage::from_treemap(ranged);
assert_eq!(inserted, ranged);
assert_ne!(
coverage_to_bytes(&inserted).unwrap(),
coverage_to_bytes(&ranged).unwrap()
);
let entity = identity(&["A"]);
let mut left = EntityCoverage::empty();
left.union_coverage(entity.clone(), inserted);
let mut right = EntityCoverage::empty();
right.union_coverage(entity, ranged);
assert_eq!(
entity_coverage_to_bytes(&left).unwrap(),
entity_coverage_to_bytes(&right).unwrap()
);
let empty_partition = Coverage::from_treemap(RoaringTreemap::from_bitmaps([(
7,
roaring::RoaringBitmap::new(),
)]));
let mut logically_empty = EntityCoverage::empty();
logically_empty.union_coverage(identity(&["empty"]), empty_partition);
let mut canonical_empty = EntityCoverage::empty();
canonical_empty.union_coverage(identity(&["empty"]), Coverage::empty());
assert_eq!(
entity_coverage_to_bytes(&logically_empty).unwrap(),
entity_coverage_to_bytes(&canonical_empty).unwrap()
);
}
#[test]
fn entity_coverage_v2_golden_payload_is_stable() {
let mut coverage = EntityCoverage::empty();
coverage.union_coverage(
EntityIdentity::try_new(vec![
EntityValue::from("\u{6771}\u{4eac}"),
EntityValue::Int32(-1),
EntityValue::Int64(i64::MIN),
EntityValue::UInt64(u64::MAX),
])
.unwrap(),
[0].into_iter().collect(),
);
let expected = [
b"TSTECOV2".as_slice(),
&[0, 0, 0, 1], &[0, 0, 0, 4], &[ENTITY_VALUE_UTF8],
&[0, 0, 0, 0, 0, 0, 0, 6],
&[0xe6, 0x9d, 0xb1, 0xe4, 0xba, 0xac],
&[ENTITY_VALUE_INT32],
&(-1i32).to_be_bytes(),
&[ENTITY_VALUE_INT64],
&i64::MIN.to_be_bytes(),
&[ENTITY_VALUE_UINT64],
&u64::MAX.to_be_bytes(),
&[0, 0, 0, 0, 0, 0, 0, 30],
ROARING_ZERO,
]
.concat();
assert_eq!(entity_coverage_to_bytes(&coverage).unwrap(), expected);
assert_eq!(entity_coverage_from_bytes(&expected).unwrap(), coverage);
}
#[test]
fn entity_coverage_decoder_rejects_every_truncated_prefix() {
let mut coverage = EntityCoverage::empty();
coverage.union_coverage(identity(&["A"]), [1].into_iter().collect());
let bytes = entity_coverage_to_bytes(&coverage).unwrap();
for end in 0..bytes.len() {
assert!(entity_coverage_from_bytes(&bytes[..end]).is_err());
}
}
#[test]
fn entity_coverage_decoder_rejects_invalid_magic_lengths_and_strings() {
let mut coverage = EntityCoverage::empty();
coverage.union_coverage(identity(&["A"]), [1].into_iter().collect());
let bytes = entity_coverage_to_bytes(&coverage).unwrap();
let mut invalid_magic = bytes.clone();
invalid_magic[0] ^= 0xff;
assert!(matches!(
entity_coverage_from_bytes(&invalid_magic),
Err(CoverageCodecError::InvalidEntityCoverageMagic { .. })
));
let mut version_one = bytes.clone();
version_one[..8].copy_from_slice(b"TSTECOV1");
assert!(matches!(
entity_coverage_from_bytes(&version_one),
Err(CoverageCodecError::InvalidEntityCoverageMagic { .. })
));
let mut invalid_count = bytes.clone();
invalid_count[8..12].copy_from_slice(&u32::MAX.to_be_bytes());
assert!(entity_coverage_from_bytes(&invalid_count).is_err());
let mut empty_identity = bytes.clone();
empty_identity[12..16].copy_from_slice(&0u32.to_be_bytes());
assert!(matches!(
entity_coverage_from_bytes(&empty_identity),
Err(CoverageCodecError::InvalidEntityIdentity { .. })
));
let mut invalid_length = bytes.clone();
invalid_length[17..25].copy_from_slice(&u64::MAX.to_be_bytes());
assert!(matches!(
entity_coverage_from_bytes(&invalid_length),
Err(CoverageCodecError::InvalidLength { .. })
));
let mut invalid_string = bytes.clone();
invalid_string[25] = 0xff;
assert!(matches!(
entity_coverage_from_bytes(&invalid_string),
Err(CoverageCodecError::InvalidEntityUtf8 { .. })
));
let mut unknown_type = bytes;
unknown_type[16] = u8::MAX;
assert!(matches!(
entity_coverage_from_bytes(&unknown_type),
Err(CoverageCodecError::UnknownEntityValueTag { tag: u8::MAX, .. })
));
}
#[test]
fn entity_coverage_decoder_rejects_duplicate_identities() {
let mut coverage = EntityCoverage::empty();
coverage.union_coverage(identity(&["A"]), Coverage::empty());
let mut bytes = entity_coverage_to_bytes(&coverage).unwrap();
let duplicate = bytes[12..].to_vec();
bytes[8..12].copy_from_slice(&2u32.to_be_bytes());
bytes.extend_from_slice(&duplicate);
assert!(matches!(
entity_coverage_from_bytes(&bytes),
Err(CoverageCodecError::DuplicateEntityIdentity { .. })
));
}
#[test]
fn entity_coverage_decoder_rejects_malformed_nested_and_trailing_bytes() {
let mut coverage = EntityCoverage::empty();
coverage.union_coverage(identity(&["A"]), [1].into_iter().collect());
let bytes = entity_coverage_to_bytes(&coverage).unwrap();
let mut malformed_nested = bytes.clone();
malformed_nested[26..34].copy_from_slice(&1u64.to_be_bytes());
assert!(matches!(
entity_coverage_from_bytes(&malformed_nested),
Err(CoverageCodecError::MalformedNestedCoverage { .. })
));
let mut trailing = bytes;
trailing.push(0);
assert!(matches!(
entity_coverage_from_bytes(&trailing),
Err(CoverageCodecError::TrailingBytes { .. })
));
}
#[test]
fn global_coverage_codec_golden_bytes_are_unchanged_and_distinct() {
let global_empty = coverage_to_bytes(&Coverage::empty()).unwrap();
assert_eq!(global_empty, vec![0; 8]);
assert!(matches!(
entity_coverage_from_bytes(&global_empty),
Err(CoverageCodecError::InvalidEntityCoverageMagic { .. })
| Err(CoverageCodecError::TruncatedPayload { .. })
));
let global_extremes: Coverage = [0, u64::MAX].into_iter().collect();
let global_extremes_bytes = [
&[2, 0, 0, 0, 0, 0, 0, 0],
&ROARING_ZERO[8..],
&ROARING_MAX[8..],
]
.concat();
assert_eq!(
coverage_to_bytes(&global_extremes).unwrap(),
global_extremes_bytes
);
assert_eq!(
coverage_from_bytes(&global_extremes_bytes)
.unwrap()
.present(),
global_extremes.present()
);
let entity_empty = entity_coverage_to_bytes(&EntityCoverage::empty()).unwrap();
assert!(coverage_from_bytes(&entity_empty).is_err());
}
}