use core::ops::Add;
use crate::{RangeSet, Tnum};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BitRange<const K: usize = 2> {
ranges: RangeSet<K>,
bits: Tnum,
}
impl<const K: usize> BitRange<K> {
pub fn new(ranges: RangeSet<K>, bits: Tnum) -> Self {
let mut result = Self { ranges, bits };
result.reduce();
result
}
pub fn from_value(value: u64) -> Self {
Self::new(RangeSet::from_value(value), Tnum::from_value(value))
}
pub fn from_ranges(ranges: RangeSet<K>) -> Self {
Self::new(ranges, Tnum::default())
}
pub fn from_bits(bits: Tnum) -> Self {
Self::new(RangeSet::default(), bits)
}
pub fn ranges(&self) -> RangeSet<K> {
self.ranges
}
pub fn bits(&self) -> Tnum {
self.bits
}
pub fn contains_value(&self, value: u64) -> bool {
self.ranges.contains_value(value) && self.bits.contains_value(value)
}
pub fn is_empty(&self) -> bool {
let Some((value, mask)) = self.bits.parts() else {
return true;
};
!self
.ranges
.ranges()
.iter()
.any(|&(low, high)| tnum_intersects_range(value, mask, low, high))
}
pub fn union(self, other: Self) -> Self {
if self.is_empty() {
return other;
}
if other.is_empty() {
return self;
}
Self::new(self.ranges.union(other.ranges), self.bits.union(other.bits))
}
pub fn intersection(self, other: Self) -> Self {
Self::new(
self.ranges.intersection(other.ranges),
self.bits.intersection(other.bits),
)
}
fn reduce(&mut self) {
if self.ranges.is_empty() || !self.bits.has_value() {
self.ranges = RangeSet::empty();
self.bits = Tnum::empty();
return;
}
let mut range_bits = Tnum::empty();
for &(low, high) in self.ranges.ranges() {
let differing = low ^ high;
let unknown = if differing == 0 {
0
} else {
u64::MAX >> differing.leading_zeros()
};
let piece = Tnum::from_parts(low & !unknown, unknown);
range_bits = range_bits.union(piece);
}
self.bits = self.bits.intersection(range_bits);
if !self.bits.has_value() {
self.ranges = RangeSet::empty();
return;
}
let (low, high) = self.bits.unsigned_bounds();
self.ranges = self.ranges.intersection(RangeSet::from_range(low, high));
if self.ranges.is_empty() {
self.bits = Tnum::empty();
}
}
}
fn tnum_intersects_range(value: u64, mask: u64, low: u64, high: u64) -> bool {
let mut equal = Some(0_u64);
let mut greater: Option<u64> = None;
for bit in (0..64).rev() {
let bit_mask = 1_u64 << bit;
let may_zero = value & bit_mask == 0;
let may_one = value & bit_mask != 0 || mask & bit_mask != 0;
if let Some(prefix) = greater {
greater = if may_zero {
Some(prefix)
} else if may_one {
Some(prefix | bit_mask)
} else {
None
};
}
if let Some(prefix) = equal {
if low & bit_mask == 0 {
if may_one {
let candidate = prefix | bit_mask;
greater = Some(greater.map_or(candidate, |current| current.min(candidate)));
}
equal = may_zero.then_some(prefix);
} else {
equal = may_one.then_some(prefix | bit_mask);
}
}
}
equal.or(greater).is_some_and(|candidate| candidate <= high)
}
impl<const K: usize> Default for BitRange<K> {
fn default() -> Self {
Self::new(RangeSet::default(), Tnum::default())
}
}
impl<const K: usize> Add for BitRange<K> {
type Output = Self;
fn add(self, other: Self) -> Self {
if self.is_empty() || other.is_empty() {
return Self::new(RangeSet::empty(), Tnum::empty());
}
Self::new(self.ranges + other.ranges, self.bits + other.bits)
}
}