Skip to main content

numeric_domains/
bit_range.rs

1use core::ops::Add;
2
3use crate::{RangeSet, Tnum};
4
5/// Reduced product of a bounded interval union and known/unknown bits.
6///
7/// Its concrete values satisfy both components. Reduction exchanges cheap
8/// unsigned-bound and common-bit facts after construction and operations.
9///
10/// This follows the reduced-product pattern used by the Linux eBPF verifier,
11/// which maintains signed bounds, unsigned bounds, and a tnum together:
12/// <https://docs.kernel.org/bpf/verifier.html#register-value-tracking>
13/// LLVM implements the analogous conversions between `ConstantRange` and
14/// `KnownBits`:
15/// <https://github.com/llvm/llvm-project/blob/main/llvm/lib/IR/ConstantRange.cpp>
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17pub struct BitRange<const K: usize = 2> {
18    ranges: RangeSet<K>,
19    bits: Tnum,
20}
21
22impl<const K: usize> BitRange<K> {
23    pub fn new(ranges: RangeSet<K>, bits: Tnum) -> Self {
24        let mut result = Self { ranges, bits };
25        result.reduce();
26        result
27    }
28
29    pub fn from_value(value: u64) -> Self {
30        Self::new(RangeSet::from_value(value), Tnum::from_value(value))
31    }
32
33    pub fn from_ranges(ranges: RangeSet<K>) -> Self {
34        Self::new(ranges, Tnum::default())
35    }
36
37    pub fn from_bits(bits: Tnum) -> Self {
38        Self::new(RangeSet::default(), bits)
39    }
40
41    pub fn ranges(&self) -> RangeSet<K> {
42        self.ranges
43    }
44    pub fn bits(&self) -> Tnum {
45        self.bits
46    }
47
48    pub fn contains_value(&self, value: u64) -> bool {
49        self.ranges.contains_value(value) && self.bits.contains_value(value)
50    }
51
52    pub fn is_empty(&self) -> bool {
53        let Some((value, mask)) = self.bits.parts() else {
54            return true;
55        };
56        !self
57            .ranges
58            .ranges()
59            .iter()
60            .any(|&(low, high)| tnum_intersects_range(value, mask, low, high))
61    }
62
63    pub fn union(self, other: Self) -> Self {
64        if self.is_empty() {
65            return other;
66        }
67        if other.is_empty() {
68            return self;
69        }
70        Self::new(self.ranges.union(other.ranges), self.bits.union(other.bits))
71    }
72
73    pub fn intersection(self, other: Self) -> Self {
74        Self::new(
75            self.ranges.intersection(other.ranges),
76            self.bits.intersection(other.bits),
77        )
78    }
79
80    fn reduce(&mut self) {
81        if self.ranges.is_empty() || !self.bits.has_value() {
82            self.ranges = RangeSet::empty();
83            self.bits = Tnum::empty();
84            return;
85        }
86
87        // Every value in a linear interval shares the prefix above the most
88        // significant bit on which its endpoints differ.
89        let mut range_bits = Tnum::empty();
90        for &(low, high) in self.ranges.ranges() {
91            let differing = low ^ high;
92            let unknown = if differing == 0 {
93                0
94            } else {
95                u64::MAX >> differing.leading_zeros()
96            };
97            let piece = Tnum::from_parts(low & !unknown, unknown);
98            range_bits = range_bits.union(piece);
99        }
100        self.bits = self.bits.intersection(range_bits);
101        if !self.bits.has_value() {
102            self.ranges = RangeSet::empty();
103            return;
104        }
105
106        // A tnum's unsigned extrema are exact, even if there are holes.
107        let (low, high) = self.bits.unsigned_bounds();
108        self.ranges = self.ranges.intersection(RangeSet::from_range(low, high));
109        if self.ranges.is_empty() {
110            self.bits = Tnum::empty();
111        }
112    }
113}
114
115/// Whether a canonical tnum has a concrete value in the inclusive interval.
116fn tnum_intersects_range(value: u64, mask: u64, low: u64, high: u64) -> bool {
117    // Build the smallest permitted value greater than or equal to `low`, one
118    // bit at a time. `equal` follows `low`; `greater` is already above it and
119    // therefore takes the smallest permitted suffix.
120    let mut equal = Some(0_u64);
121    let mut greater: Option<u64> = None;
122    for bit in (0..64).rev() {
123        let bit_mask = 1_u64 << bit;
124        let may_zero = value & bit_mask == 0;
125        let may_one = value & bit_mask != 0 || mask & bit_mask != 0;
126
127        if let Some(prefix) = greater {
128            greater = if may_zero {
129                Some(prefix)
130            } else if may_one {
131                Some(prefix | bit_mask)
132            } else {
133                None
134            };
135        }
136
137        if let Some(prefix) = equal {
138            if low & bit_mask == 0 {
139                if may_one {
140                    let candidate = prefix | bit_mask;
141                    greater = Some(greater.map_or(candidate, |current| current.min(candidate)));
142                }
143                equal = may_zero.then_some(prefix);
144            } else {
145                equal = may_one.then_some(prefix | bit_mask);
146            }
147        }
148    }
149
150    equal.or(greater).is_some_and(|candidate| candidate <= high)
151}
152
153impl<const K: usize> Default for BitRange<K> {
154    fn default() -> Self {
155        Self::new(RangeSet::default(), Tnum::default())
156    }
157}
158
159impl<const K: usize> Add for BitRange<K> {
160    type Output = Self;
161
162    fn add(self, other: Self) -> Self {
163        if self.is_empty() || other.is_empty() {
164            return Self::new(RangeSet::empty(), Tnum::empty());
165        }
166        Self::new(self.ranges + other.ranges, self.bits + other.bits)
167    }
168}