use std::borrow::Cow;
use std::path::Path;
use bitvec::mem::BitRegister;
use bitvec::order::Lsb0;
use itertools::Itertools;
use crate::common::bitvec::BitVec;
use crate::common::generic_consts::Random;
use crate::common::universal_io::{
Flusher, OpenOptions, ReadRange, Result, TypedStorage, UniversalIoError, UniversalRead,
UniversalReadFs, UniversalWrite,
};
const BITS_PER_ELEMENT: u32 = BitStore::BITS;
type BitStore = u64;
type BitSlice = bitvec::slice::BitSlice<BitStore, Lsb0>;
pub type MmapBitSlice = StoredBitSlice<crate::common::universal_io::MmapFile>;
#[derive(Debug)]
pub struct StoredBitSlice<S> {
storage: TypedStorage<S, BitStore>,
element_len: u64,
}
impl<S: UniversalRead> StoredBitSlice<S> {
pub fn open(
fs: &S::Fs,
path: impl AsRef<Path>,
options: OpenOptions,
extra: <S::Fs as UniversalReadFs>::OpenExtra,
) -> Result<Self> {
let storage = TypedStorage::open(fs, path, options, extra)?;
let element_len = storage.len()?;
Ok(Self {
storage,
element_len,
})
}
pub fn reopen(&mut self) -> Result<()> {
self.storage.reopen()?;
self.element_len = self.storage.len()?;
Ok(())
}
pub fn bit_len(&self) -> u64 {
self.element_len * u64::from(BITS_PER_ELEMENT)
}
pub fn element_len(&self) -> u64 {
self.element_len
}
#[inline(always)]
fn element_idx(bit_idx: u64) -> u64 {
bit_idx >> <BitStore as BitRegister>::INDX
}
#[inline(always)]
fn bit_within_element(bit_idx: u64) -> u8 {
bit_idx as u8 & <BitStore as BitRegister>::MASK
}
pub fn read_all(&self) -> Result<Cow<'_, BitSlice>> {
let elements = self.storage.read_whole()?;
match elements {
Cow::Borrowed(slice) => Ok(Cow::Borrowed(BitSlice::from_slice(slice))),
Cow::Owned(vec) => Ok(Cow::Owned(BitVec::from_vec(vec))),
}
}
pub fn read_bit_range(&self, range: std::ops::Range<u64>) -> Result<Cow<'_, BitSlice>> {
if range.is_empty() {
return Ok(Cow::Borrowed(BitSlice::empty()));
}
let elem_start = Self::element_idx(range.start);
let elem_end = range.end.div_ceil(u64::from(BITS_PER_ELEMENT));
let num_elements = elem_end - elem_start;
let elements = self.storage.read::<Random>(ReadRange {
byte_offset: elem_start * size_of::<BitStore>() as u64,
length: num_elements,
})?;
let bit_offset = Self::bit_within_element(range.start) as usize;
let bit_end = bit_offset + (range.end - range.start) as usize;
match elements {
Cow::Borrowed(slice) => {
let bits = BitSlice::from_slice(slice);
Ok(Cow::Borrowed(&bits[bit_offset..bit_end]))
}
Cow::Owned(vec) => {
let bits = BitVec::from_vec(vec);
Ok(Cow::Owned(bits[bit_offset..bit_end].to_bitvec()))
}
}
}
pub fn count_ones(&self) -> Result<usize> {
Ok(self.read_all()?.count_ones())
}
pub fn get_bit(&self, bit_index: u64) -> Result<Option<bool>> {
let element_index = Self::element_idx(bit_index);
let bit_within_element = Self::bit_within_element(bit_index);
if element_index >= self.element_len {
return Ok(None);
}
let element = self
.storage
.read::<Random>(ReadRange::one(element_index * size_of::<BitStore>() as u64))?[0];
let bitslice = BitSlice::from_element(&element);
Ok(bitslice
.get(bit_within_element as usize)
.as_deref()
.copied())
}
pub fn populate(&self) -> Result<()> {
self.storage.populate()
}
pub fn clear_ram_cache(&self) -> Result<()> {
self.storage.clear_ram_cache()
}
}
impl<S: UniversalWrite> StoredBitSlice<S> {
pub fn set_ascending_bits_batch(
&mut self,
updates: impl IntoIterator<Item = (u64, bool)>,
) -> Result<()> {
let mut prev_element: Option<u64> = None;
let mut run_start = 0u64;
let runs = updates.into_iter().chunk_by(move |(bit_idx, _)| {
let element_idx = Self::element_idx(*bit_idx);
if prev_element.is_none_or(|prev| element_idx > prev + 1) {
run_start = element_idx;
}
prev_element = Some(element_idx);
run_start
});
for (element_start, run_updates) in &runs {
let run_updates: Vec<_> = run_updates.collect();
let last_element = Self::element_idx(run_updates.last().unwrap().0);
let num_elements = last_element - element_start + 1;
if element_start + num_elements > self.element_len {
return Err(UniversalIoError::OutOfBounds {
start: element_start,
end: element_start + num_elements,
elements: self.element_len as usize,
});
}
let mut buf = self
.storage
.read::<Random>(ReadRange {
byte_offset: element_start * size_of::<BitStore>() as u64,
length: num_elements,
})?
.into_owned();
let bitslice = BitSlice::from_slice_mut(&mut buf);
for (bit_idx, value) in run_updates {
let bit_offset =
bit_idx as usize - (element_start as usize * BITS_PER_ELEMENT as usize);
bitslice.set(bit_offset, value);
}
self.storage
.write(element_start * size_of::<BitStore>() as u64, &buf)?;
}
Ok(())
}
pub fn write_bitslice<T2, O2>(&mut self, source: &bitvec::slice::BitSlice<T2, O2>) -> Result<()>
where
T2: bitvec::store::BitStore,
O2: bitvec::order::BitOrder,
{
let bit_count = source.len() as u64;
if bit_count == 0 {
return Ok(());
}
if bit_count > self.bit_len() {
return Err(UniversalIoError::OutOfBounds {
start: 0,
end: bit_count,
elements: self.bit_len() as usize,
});
}
let element_count = bit_count.div_ceil(u64::from(BITS_PER_ELEMENT));
let existing = self.storage.read::<Random>(ReadRange {
byte_offset: 0,
length: element_count,
})?;
let mut buf = existing.into_owned();
let buf_bits = BitSlice::from_slice_mut(&mut buf);
buf_bits[..bit_count as usize].clone_from_bitslice(source);
self.storage.write(0, &buf)
}
pub fn replace_bit(&mut self, bit_index: u64, value: bool) -> Result<bool> {
let element_index = Self::element_idx(bit_index);
let bit_within_element = Self::bit_within_element(bit_index);
if element_index >= self.element_len {
return Err(UniversalIoError::OutOfBounds {
start: bit_index,
end: bit_index + 1,
elements: self.bit_len() as usize,
});
}
let mut element = self
.storage
.read::<Random>(ReadRange::one(element_index * size_of::<BitStore>() as u64))?[0];
let element = &mut element;
let bitslice = BitSlice::from_element_mut(element);
let old_bit = bitslice.replace(bit_within_element as usize, value);
if old_bit != value {
self.storage
.write(element_index * size_of::<BitStore>() as u64, &[*element])?;
}
Ok(old_bit)
}
pub fn flusher(&self) -> Flusher {
self.storage.flusher()
}
}
#[cfg(test)]
mod tests {
use std::io::Write;
use tempfile::NamedTempFile;
use super::*;
use crate::common::universal_io::MmapFs;
fn create_temp_file(data: &[u8]) -> NamedTempFile {
let mut f = NamedTempFile::new().unwrap();
let aligned_len = if data.is_empty() {
0
} else {
data.len()
.next_multiple_of(std::mem::size_of::<BitStore>())
.max(std::mem::size_of::<BitStore>())
};
let mut buf = vec![0u8; aligned_len];
buf[..data.len()].copy_from_slice(data);
f.write_all(&buf).unwrap();
f.flush().unwrap();
f
}
#[test]
fn test_read_whole_bitslice() {
let data = [
0b10110010, 0b01001111, 0x00, 0x00, 0x00, 0x00, 0x00, 0b10000000,
0b00000001, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0b11111111,
];
let f = create_temp_file(&data);
let storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
assert_eq!(storage.element_len(), 2);
assert_eq!(storage.bit_len(), 128);
let bs = storage.read_all().unwrap();
assert!(!bs[0]); assert!(bs[1]); assert!(!bs[2]); assert!(!bs[3]); assert!(bs[4]); assert!(bs[5]); assert!(!bs[6]); assert!(bs[7]);
assert!(bs[8]); assert!(bs[9]); assert!(bs[10]); assert!(bs[11]); assert!(!bs[12]); assert!(!bs[13]); assert!(bs[14]); assert!(!bs[15]);
assert!(!bs[56]); assert!(bs[63]);
assert!(bs[64]); assert!(!bs[65]);
for i in 120..=127 {
assert!(bs[i], "bit {i} should be set");
}
for i in 72..120 {
assert!(!bs[i], "bit {i} should be clear");
}
}
#[test]
fn test_get_single_bit() {
let data = [0xB2u8]; let f = create_temp_file(&data);
let storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
assert_eq!(storage.get_bit(0).unwrap(), Some(false));
assert_eq!(storage.get_bit(1).unwrap(), Some(true));
assert_eq!(storage.get_bit(4).unwrap(), Some(true));
assert_eq!(storage.get_bit(7).unwrap(), Some(true));
assert_eq!(storage.get_bit(64).unwrap(), None);
}
#[test]
fn test_set_bit() {
let f = create_temp_file(&[0x00; 8]);
let mut storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
storage.replace_bit(3, true).unwrap();
assert_eq!(storage.get_bit(3).unwrap(), Some(true));
assert_eq!(storage.get_bit(0).unwrap(), Some(false));
storage.replace_bit(3, false).unwrap();
assert_eq!(storage.get_bit(3).unwrap(), Some(false));
}
#[test]
fn test_set_bit_out_of_bounds() {
let f = create_temp_file(&[0x00; 8]);
let mut storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
assert!(storage.replace_bit(storage.bit_len(), true).is_err());
}
#[test]
fn test_replace_bit() {
let f = create_temp_file(&[0xFF; 8]); let mut storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
let old = storage.replace_bit(2, false).unwrap();
assert!(old);
assert_eq!(storage.get_bit(2).unwrap(), Some(false));
let old = storage.replace_bit(2, true).unwrap();
assert!(!old);
assert_eq!(storage.get_bit(2).unwrap(), Some(true));
}
#[test]
fn test_set_bits_batch() {
const NUM_BITS: u64 = 8192; let f = create_temp_file(&[0x00; (NUM_BITS / 8) as usize]);
let mut storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
assert_eq!(storage.bit_len(), NUM_BITS);
fn assert_bits(storage: &MmapBitSlice, expected: impl Fn(u64) -> bool) {
let bs = storage.read_all().unwrap();
for i in 0..storage.bit_len() {
assert_eq!(bs[i as usize], expected(i), "mismatch at bit {i}",);
}
}
storage
.set_ascending_bits_batch((0..NUM_BITS).filter(|i| i % 2 == 1).map(|i| (i, true)))
.unwrap();
assert_bits(&storage, |i| i % 2 == 1);
assert_eq!(storage.count_ones().unwrap(), (NUM_BITS / 2) as usize);
storage
.set_ascending_bits_batch((0..NUM_BITS).filter(|i| i % 2 == 0).map(|i| (i, true)))
.unwrap();
assert_bits(&storage, |_| true);
assert_eq!(storage.count_ones().unwrap(), NUM_BITS as usize);
storage
.set_ascending_bits_batch((0..NUM_BITS).filter(|i| i % 3 == 0).map(|i| (i, false)))
.unwrap();
assert_bits(&storage, |i| i % 3 != 0);
storage
.set_ascending_bits_batch(
(0..NUM_BITS)
.filter(|i| i % 64 == 0 || i % 64 == 63)
.map(|i| (i, true)),
)
.unwrap();
assert_bits(&storage, |i| i % 3 != 0 || i % 64 == 0 || i % 64 == 63);
storage
.set_ascending_bits_batch((0..NUM_BITS).map(|i| (i, false)))
.unwrap();
assert_bits(&storage, |_| false);
assert_eq!(storage.count_ones().unwrap(), 0);
storage
.set_ascending_bits_batch(
[0, 3, 7, 111]
.into_iter()
.flat_map(|el: u64| (el * 64..el * 64 + 64).map(|i| (i, true))),
)
.unwrap();
assert_bits(&storage, |i| matches!(i / 64, 0 | 3 | 7 | 111));
assert!(
storage
.set_ascending_bits_batch([(NUM_BITS, true)])
.is_err()
);
storage
.set_ascending_bits_batch(std::iter::empty::<(u64, bool)>())
.unwrap();
assert_bits(&storage, |i| matches!(i / 64, 0 | 3 | 7 | 111));
}
#[test]
fn test_flusher() {
let f = create_temp_file(&[0x00; 8]);
let mut storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
storage.replace_bit(0, true).unwrap();
storage.flusher()().unwrap();
let storage2: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
assert_eq!(storage2.get_bit(0).unwrap(), Some(true));
}
#[test]
fn test_bit_len() {
let f = create_temp_file(&[0u8; 16]);
let storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
assert_eq!(storage.element_len(), 2);
assert_eq!(storage.bit_len(), 128);
}
#[test]
fn test_read_all_as_bitslice() {
let data = [0xAB, 0xCD, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00];
let f = create_temp_file(&data);
let storage: MmapBitSlice =
StoredBitSlice::open(&MmapFs, f.path(), OpenOptions::new_for_test(), ()).unwrap();
let bs = storage.read_all().unwrap();
assert_eq!(bs.len(), storage.bit_len() as usize);
assert!(matches!(bs, Cow::Borrowed(_)));
}
}