use std::borrow::Cow;
use std::io;
use std::path::Path;
use roaring::RoaringBitmap;
use super::format::{BitmaskContent, BitmaskHeader, Encoding, HEADER_SIZE, MAGIC, VERSION};
use crate::common::bitvec::{BitSlice, BitVec};
use crate::common::generic_consts::Sequential;
use crate::common::universal_io::{
OpenOptions, ReadRange, TypedStorage, UioResult, UniversalIoError, UniversalRead,
UniversalReadFs,
};
#[derive(Debug)]
pub struct StoredBitmask<S> {
storage: TypedStorage<S, u8>,
logical_len: u64,
pub(super) encoding: Encoding,
pub(super) payload_len: u64,
}
fn dense_to_bitvec(bytes: &[u8], len: usize) -> BitVec {
let mut words = vec![0u64; bytes.len().div_ceil(size_of::<u64>())];
bytemuck::cast_slice_mut::<u64, u8>(&mut words)[..bytes.len()].copy_from_slice(bytes);
let mut bits = BitVec::from_vec(words);
bits.truncate(len);
bits
}
fn invalid_data(path: &Path, message: impl std::fmt::Display) -> UniversalIoError {
UniversalIoError::Io(io::Error::new(
io::ErrorKind::InvalidData,
format!("{}: {message}", path.display()),
))
}
impl<S: UniversalRead> StoredBitmask<S> {
pub fn open<Fs: UniversalReadFs<File = S>>(
fs: &Fs,
path: impl AsRef<Path>,
options: OpenOptions,
extra: Fs::OpenExtra,
) -> UioResult<Self> {
let path = path.as_ref();
let storage = TypedStorage::open(fs, path, options, extra)?;
let file_len = storage.len()?;
if file_len < HEADER_SIZE as u64 {
return Err(invalid_data(
path,
format_args!("file of {file_len} bytes is too short to be a stored bitmask"),
));
}
let header_bytes = storage.read(ReadRange::new(0, HEADER_SIZE as u64), Sequential)?;
let header: BitmaskHeader = bytemuck::pod_read_unaligned(&header_bytes);
if header.magic != MAGIC {
return Err(invalid_data(path, "not a stored bitmask file (bad magic)"));
}
if header.version != VERSION {
return Err(invalid_data(
path,
format_args!(
"unsupported stored bitmask version {} (supported: {VERSION})",
header.version,
),
));
}
let Some(encoding) = Encoding::from_u32(header.encoding) else {
return Err(invalid_data(
path,
format_args!("unknown stored bitmask encoding {}", header.encoding),
));
};
if header.logical_len > u64::from(u32::MAX) + 1 {
return Err(invalid_data(
path,
format_args!(
"bitmask of {} bits exceeds the u32 position space",
header.logical_len,
),
));
}
if header.payload_len > file_len - HEADER_SIZE as u64 {
return Err(invalid_data(
path,
format_args!(
"payload of {} bytes exceeds file of {file_len} bytes",
header.payload_len,
),
));
}
if encoding == Encoding::Dense
&& header.payload_len < header.logical_len.div_ceil(u64::from(u8::BITS))
{
return Err(invalid_data(
path,
format_args!(
"dense payload of {} bytes is too short for {} bits",
header.payload_len, header.logical_len,
),
));
}
Ok(Self {
storage,
logical_len: header.logical_len,
encoding,
payload_len: header.payload_len,
})
}
pub fn bit_len(&self) -> u64 {
self.logical_len
}
pub fn read(&self) -> UioResult<BitmaskContent<'_>> {
let payload = self.storage.read(
ReadRange::new(HEADER_SIZE as u64, self.payload_len),
Sequential,
)?;
match self.encoding {
Encoding::Dense => {
let len = self.logical_len as usize;
let bits = match payload {
Cow::Borrowed(bytes) => match bytemuck::try_cast_slice::<u8, u64>(bytes) {
Ok(words) => Cow::Borrowed(&BitSlice::from_slice(words)[..len]),
Err(_) => Cow::Owned(dense_to_bitvec(bytes, len)),
},
Cow::Owned(bytes) => Cow::Owned(dense_to_bitvec(&bytes, len)),
};
Ok(BitmaskContent::Dense(bits))
}
Encoding::RoaringOnes => Ok(BitmaskContent::Ones(self.decode_roaring(&payload)?)),
Encoding::RoaringZeros => Ok(BitmaskContent::Zeros(self.decode_roaring(&payload)?)),
}
}
pub fn read_ones(&self) -> UioResult<RoaringBitmap> {
match self.read()? {
BitmaskContent::Dense(bits) => Ok(bits
.iter_ones()
.map(|idx| idx as u32)
.collect::<RoaringBitmap>()),
BitmaskContent::Ones(ones) => Ok(ones),
BitmaskContent::Zeros(zeros) => {
let mut ones = RoaringBitmap::new();
if let Some(last) = self.logical_len.checked_sub(1) {
ones.insert_range(0..=last as u32);
ones -= zeros;
}
Ok(ones)
}
}
}
fn decode_roaring(&self, payload: &[u8]) -> UioResult<RoaringBitmap> {
let bitmap = RoaringBitmap::deserialize_from(payload)?;
if let Some(max) = bitmap.max()
&& u64::from(max) >= self.logical_len
{
return Err(UniversalIoError::Io(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"stored bitmask position {max} out of {} bits",
self.logical_len,
),
)));
}
Ok(bitmap)
}
}