use std::io;
use super::bloom::{BloomFilter, decode_bloom_block};
const PREFIX_LEN_HEADER: usize = 8;
const MIN_BLOOM_BLOCK: usize = 4;
pub(crate) struct FilterBlock {
user: BloomFilter,
prefix: Option<BloomFilter>,
charge: usize,
}
impl FilterBlock {
pub(crate) fn decode(region: &[u8]) -> io::Result<Self> {
if region.len() < PREFIX_LEN_HEADER {
return Err(invalid_data(
"bloom region too short for prefix-bloom length header",
));
}
let prefix_len = usize::try_from(u64::from_le_bytes(
region[0..PREFIX_LEN_HEADER].try_into().unwrap(),
))
.map_err(|_| invalid_data("prefix bloom length is too large to address"))?;
let user_offset = PREFIX_LEN_HEADER
.checked_add(prefix_len)
.ok_or_else(|| invalid_data("prefix bloom length overflows"))?;
if user_offset > region.len() {
return Err(invalid_data("prefix bloom length exceeds bloom region"));
}
if prefix_len > 0 && prefix_len < MIN_BLOOM_BLOCK {
return Err(invalid_data("prefix bloom block too short"));
}
if region.len() - user_offset < MIN_BLOOM_BLOCK {
return Err(invalid_data("user bloom block too short"));
}
let prefix =
(prefix_len > 0).then(|| decode_bloom_block(®ion[PREFIX_LEN_HEADER..user_offset]));
let user = decode_bloom_block(®ion[user_offset..]);
Ok(Self {
user,
prefix,
charge: std::mem::size_of::<Self>() + region.len(),
})
}
pub(crate) fn may_contain(&self, user_key: &[u8]) -> bool {
self.user.may_contain(user_key)
}
pub(crate) fn may_have_prefix(&self, prefix: &[u8]) -> bool {
match &self.prefix {
Some(filter) => filter.may_contain(prefix),
None => true,
}
}
pub(crate) fn charge(&self) -> usize {
self.charge
}
}
impl std::fmt::Debug for FilterBlock {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FilterBlock")
.field("has_prefix_bloom", &self.prefix.is_some())
.field("charge", &self.charge)
.finish()
}
}
fn invalid_data(message: &'static str) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, message)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::engine::bloom::{BloomFilterBuilder, encode_bloom_block};
fn region(prefix_keys: Option<&[&[u8]]>, user_keys: &[&[u8]]) -> Vec<u8> {
let prefix_bytes = match prefix_keys {
Some(keys) => {
let mut builder = BloomFilterBuilder::new(10);
for key in keys {
builder.add_key(key);
}
encode_bloom_block(&builder.build())
}
None => Vec::new(),
};
let mut builder = BloomFilterBuilder::new(10);
for key in user_keys {
builder.add_key(key);
}
let user_bytes = encode_bloom_block(&builder.build());
let mut out = Vec::new();
out.extend_from_slice(&(prefix_bytes.len() as u64).to_le_bytes());
out.extend_from_slice(&prefix_bytes);
out.extend_from_slice(&user_bytes);
out
}
#[test]
fn decodes_region_without_prefix_bloom() {
let bytes = region(None, &[b"alpha", b"beta"]);
let filter = FilterBlock::decode(&bytes).unwrap();
assert!(filter.may_contain(b"alpha"));
assert!(filter.may_contain(b"beta"));
assert!(
filter.may_have_prefix(b"anything"),
"no prefix bloom means no negative information"
);
}
#[test]
fn decodes_region_with_prefix_bloom() {
let bytes = region(Some(&[b"pre"]), &[b"prefix-key"]);
let filter = FilterBlock::decode(&bytes).unwrap();
assert!(filter.may_contain(b"prefix-key"));
assert!(filter.may_have_prefix(b"pre"));
}
#[test]
fn rejects_short_region() {
for len in 0..PREFIX_LEN_HEADER {
let err = FilterBlock::decode(&vec![0u8; len]).expect_err("short region");
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
}
#[test]
fn rejects_prefix_length_past_region() {
let mut bytes = region(None, &[b"k"]);
let claimed = bytes.len() as u64 + 1;
bytes[0..8].copy_from_slice(&claimed.to_le_bytes());
let err = FilterBlock::decode(&bytes).expect_err("prefix past end");
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn rejects_truncated_prefix_bloom() {
let mut bytes = region(None, &[b"k"]);
bytes[0..8].copy_from_slice(&2u64.to_le_bytes());
let err = FilterBlock::decode(&bytes).expect_err("prefix bloom too short");
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn rejects_truncated_user_bloom() {
let bytes = vec![0u8; PREFIX_LEN_HEADER + 3];
let err = FilterBlock::decode(&bytes).expect_err("user bloom too short");
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn charge_covers_the_whole_region() {
let bytes = region(Some(&[b"pre"]), &[b"a", b"b", b"c"]);
let filter = FilterBlock::decode(&bytes).unwrap();
assert!(filter.charge() >= bytes.len());
}
}