#[derive(Clone, Debug, Default, PartialEq)]
pub struct Bitmap(roaring::RoaringBitmap);
impl Bitmap {
pub fn new() -> Self {
Self::default()
}
pub fn len(&self) -> u64 {
self.0.len()
}
pub fn clear(&mut self) {
self.0.clear()
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
pub fn insert(&mut self, value: u32) -> bool {
self.0.insert(value)
}
pub fn insert_range<R>(&mut self, range: R) -> u64
where
R: std::ops::RangeBounds<u32>,
{
self.0.insert_range(range)
}
pub fn remove(&mut self, value: u32) -> bool {
self.0.remove(value)
}
pub fn contains(&self, value: u32) -> bool {
self.0.contains(value)
}
pub fn iter(&self) -> impl Iterator<Item = u32> + use<'_> {
self.0.iter()
}
#[cfg(feature = "serde")]
#[cfg_attr(doc_cfg, doc(cfg(feature = "serde")))]
pub fn deserialize_from<R: std::io::Read>(reader: R) -> std::io::Result<Self> {
roaring::RoaringBitmap::deserialize_from(reader).map(Self)
}
#[cfg(feature = "serde")]
#[cfg_attr(doc_cfg, doc(cfg(feature = "serde")))]
pub fn serialize_into<W: std::io::Write>(&self, writer: W) -> std::io::Result<()> {
self.0.serialize_into(writer)
}
}
impl FromIterator<u32> for Bitmap {
fn from_iter<I: IntoIterator<Item = u32>>(iterator: I) -> Bitmap {
let mut bitmap = Self::new();
bitmap.extend(iterator);
bitmap
}
}
impl<'a> FromIterator<&'a u32> for Bitmap {
fn from_iter<I: IntoIterator<Item = &'a u32>>(iterator: I) -> Bitmap {
let mut bitmap = Self::new();
bitmap.extend(iterator);
bitmap
}
}
impl Extend<u32> for Bitmap {
fn extend<I: IntoIterator<Item = u32>>(&mut self, values: I) {
self.0.extend(values)
}
}
impl<'a> Extend<&'a u32> for Bitmap {
fn extend<I: IntoIterator<Item = &'a u32>>(&mut self, values: I) {
self.extend(values.into_iter().copied());
}
}
#[cfg(feature = "serde")]
#[cfg_attr(doc_cfg, doc(cfg(feature = "serde")))]
impl serde::Serialize for Bitmap {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
use serde_with::SerializeAs;
let mut bytes = vec![];
self.serialize_into(&mut bytes)
.map_err(serde::ser::Error::custom)?;
if serializer.is_human_readable() {
let b64 = <base64ct::Base64 as base64ct::Encoding>::encode_string(&bytes);
serde::Serialize::serialize(&b64, serializer)
} else {
serde_with::Bytes::serialize_as(&bytes, serializer)
}
}
}
#[cfg(feature = "serde")]
#[cfg_attr(doc_cfg, doc(cfg(feature = "serde")))]
impl<'de> serde::Deserialize<'de> for Bitmap {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
use serde_with::DeserializeAs;
if deserializer.is_human_readable() {
let b64: std::borrow::Cow<'de, str> = serde::Deserialize::deserialize(deserializer)?;
let bytes = <base64ct::Base64 as base64ct::Encoding>::decode_vec(&b64)
.map_err(serde::de::Error::custom)?;
deserialize_canonical_bitmap(&bytes).map_err(serde::de::Error::custom)
} else {
let bytes: std::borrow::Cow<'de, [u8]> =
serde_with::Bytes::deserialize_as(deserializer)?;
deserialize_canonical_bitmap(&bytes).map_err(serde::de::Error::custom)
}
}
}
#[cfg(feature = "serde")]
fn deserialize_canonical_bitmap(bytes: &[u8]) -> std::io::Result<Bitmap> {
let mut reader: &[u8] = bytes;
let bitmap = Bitmap::deserialize_from(&mut reader)?;
if !reader.is_empty() {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"trailing bytes after roaring bitmap encoding: {} byte(s) remaining",
reader.len(),
),
));
}
Ok(bitmap)
}
#[cfg(feature = "proptest")]
#[cfg_attr(doc_cfg, doc(cfg(feature = "proptest")))]
impl proptest::arbitrary::Arbitrary for Bitmap {
type Parameters = ();
type Strategy = proptest::strategy::BoxedStrategy<Self>;
fn arbitrary_with(_args: Self::Parameters) -> Self::Strategy {
use proptest::collection::vec;
use proptest::prelude::*;
vec(any::<u32>(), 0..32).prop_map(Self::from_iter).boxed()
}
}
#[cfg(test)]
mod test {
use super::*;
use base64ct::Encoding;
#[test]
fn test_unique_deserialize() {
let raw = "OjAAAAoAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAWAAAAFoAAABcAAAAXgAAAGAAAABiAAAAZAAAAGYAAABoAAAAagAAAAEAAQABAAEAAQABAAEAAQABAAEA";
let bytes = base64ct::Base64::decode_vec(raw).unwrap();
roaring::RoaringBitmap::deserialize_from(&bytes[..]).unwrap_err();
Bitmap::deserialize_from(&bytes[..]).unwrap_err();
}
#[test]
fn bcs_deserialize_rejects_trailing_bytes() {
let bitmap: Bitmap = (1..4).collect();
let mut canonical = Vec::new();
bitmap.serialize_into(&mut canonical).unwrap();
let canonical_bcs = bcs::to_bytes::<Vec<u8>>(&canonical).unwrap();
let decoded: Bitmap = bcs::from_bytes(&canonical_bcs).unwrap();
assert_eq!(decoded, bitmap);
let mut padded = canonical.clone();
padded.extend_from_slice(&[0xff, 0xff, 0xff, 0xff]);
let padded_bcs = bcs::to_bytes::<Vec<u8>>(&padded).unwrap();
let err = bcs::from_bytes::<Bitmap>(&padded_bcs).unwrap_err();
assert!(
err.to_string().contains("trailing bytes"),
"unexpected error: {err}"
);
}
}