Skip to main content

probl_number/
lib.rs

1//! One integer type with inline small values and shared, canonical large values.
2//! The hard size ceiling bounds parsing and individual allocations even before
3//! the runtime has a work budget. Hosts can impose a smaller runtime ceiling.
4
5use num_bigint::BigInt;
6use num_rational::BigRational;
7use num_traits::{FromPrimitive, Signed, ToPrimitive, Zero};
8use rustc_hash::FxHasher;
9use std::borrow::Cow;
10use std::cmp::Ordering;
11use std::fmt;
12use std::hash::{Hash, Hasher};
13use std::str::FromStr;
14use std::sync::Arc;
15
16pub const MAX_INTEGER_BITS: u64 = 65_536;
17// ceil(MAX_INTEGER_BITS * log10(2)); check bits after parsing the last digit.
18pub const MAX_INTEGER_DIGITS: usize = 19_729;
19
20#[derive(Clone)]
21pub struct Integer(Repr);
22
23#[derive(Clone)]
24enum Repr {
25    Small(i64),
26    Large(Arc<Large>),
27}
28
29struct Large {
30    value: BigInt,
31    hash: u64,
32}
33
34#[derive(Clone, Copy, Debug, PartialEq, Eq)]
35pub enum IntError {
36    Invalid,
37    TooLarge,
38    DivisionByZero,
39}
40
41impl fmt::Display for IntError {
42    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
43        match self {
44            Self::Invalid => f.write_str("invalid integer"),
45            Self::TooLarge => write!(f, "integer size exceeds the limit of {MAX_INTEGER_BITS} bits"),
46            Self::DivisionByZero => f.write_str("division by zero"),
47        }
48    }
49}
50
51impl Integer {
52    pub const ZERO: Self = Self(Repr::Small(0));
53    pub const ONE: Self = Self(Repr::Small(1));
54
55    /// Parse unsigned binary or hexadecimal digits (no prefix or separators).
56    /// Check the significant bit count before allocating a bigint. Decimal
57    /// parsing remains separate so data files keep their decimal-only contract.
58    pub fn from_radix_digits(digits: &str, radix: u32) -> Result<Self, IntError> {
59        let valid = match radix {
60            2 => digits.bytes().all(|b| matches!(b, b'0' | b'1')),
61            16 => digits.bytes().all(|b| b.is_ascii_hexdigit()),
62            _ => false,
63        };
64        if digits.is_empty() || !valid {
65            return Err(IntError::Invalid);
66        }
67        let digits = digits.trim_start_matches('0');
68        if digits.is_empty() {
69            return Ok(Self::ZERO);
70        }
71        let first = (digits.as_bytes()[0] as char).to_digit(radix).expect("validated digit");
72        let bits = (digits.len() as u64 - 1)
73            .saturating_mul(u64::from(radix.trailing_zeros()))
74            .saturating_add(u64::from(32 - first.leading_zeros()));
75        if bits > MAX_INTEGER_BITS {
76            return Err(IntError::TooLarge);
77        }
78        if let Ok(n) = i64::from_str_radix(digits, radix) {
79            return Ok(n.into());
80        }
81        Self::from_big(BigInt::parse_bytes(digits.as_bytes(), radix).ok_or(IntError::Invalid)?)
82    }
83
84    pub fn from_big(value: BigInt) -> Result<Self, IntError> {
85        if let Some(n) = value.to_i64() {
86            return Ok(n.into());
87        }
88        if value.bits() > MAX_INTEGER_BITS {
89            return Err(IntError::TooLarge);
90        }
91        let mut h = FxHasher::default();
92        value.hash(&mut h);
93        Ok(Self(Repr::Large(Arc::new(Large {
94            value,
95            hash: h.finish(),
96        }))))
97    }
98
99    pub fn big(&self) -> Cow<'_, BigInt> {
100        match &self.0 {
101            Repr::Small(n) => Cow::Owned(BigInt::from(*n)),
102            Repr::Large(n) => Cow::Borrowed(&n.value),
103        }
104    }
105
106    pub fn to_i64(&self) -> Option<i64> {
107        match self.0 {
108            Repr::Small(n) => Some(n),
109            _ => None,
110        }
111    }
112    pub fn to_u64(&self) -> Option<u64> {
113        match &self.0 {
114            Repr::Small(n) => u64::try_from(*n).ok(),
115            Repr::Large(n) => n.value.to_u64(),
116        }
117    }
118    pub fn to_u128(&self) -> Option<u128> {
119        match &self.0 {
120            Repr::Small(n) => u128::try_from(*n).ok(),
121            Repr::Large(n) => n.value.to_u128(),
122        }
123    }
124    pub fn to_f64(&self) -> Option<f64> {
125        match &self.0 {
126            Repr::Small(n) => Some(*n as f64),
127            Repr::Large(n) => n.value.to_f64().filter(|x| x.is_finite()),
128        }
129    }
130    pub fn from_f64(x: f64) -> Option<Self> {
131        if x >= i64::MIN as f64 && x < -(i64::MIN as f64) {
132            Some((x as i64).into())
133        } else {
134            BigInt::from_f64(x).and_then(|n| Self::from_big(n).ok())
135        }
136    }
137    pub fn bits(&self) -> u64 {
138        match &self.0 {
139            Repr::Small(n) => 64 - u64::from(n.unsigned_abs().leading_zeros()),
140            Repr::Large(n) => n.value.bits(),
141        }
142    }
143    /// Count set bits in the magnitude, ignoring the sign.
144    pub fn bit_count(&self) -> u64 {
145        match &self.0 {
146            Repr::Small(n) => u64::from(n.unsigned_abs().count_ones()),
147            Repr::Large(n) => n.value.magnitude().count_ones(),
148        }
149    }
150    /// Signed bit operations use infinite two's-complement sign extension.
151    pub fn bit_and(&self, rhs: &Self) -> Result<Self, IntError> {
152        if let (Some(a), Some(b)) = (self.to_i64(), rhs.to_i64()) {
153            return Ok((a & b).into());
154        }
155        Self::from_big(self.big().as_ref() & rhs.big().as_ref())
156    }
157    pub fn bit_or(&self, rhs: &Self) -> Result<Self, IntError> {
158        if let (Some(a), Some(b)) = (self.to_i64(), rhs.to_i64()) {
159            return Ok((a | b).into());
160        }
161        Self::from_big(self.big().as_ref() | rhs.big().as_ref())
162    }
163    pub fn bit_xor(&self, rhs: &Self) -> Result<Self, IntError> {
164        if let (Some(a), Some(b)) = (self.to_i64(), rhs.to_i64()) {
165            return Ok((a ^ b).into());
166        }
167        Self::from_big(self.big().as_ref() ^ rhs.big().as_ref())
168    }
169    pub fn bit_not(&self) -> Result<Self, IntError> {
170        if let Some(n) = self.to_i64() {
171            return Ok((!n).into());
172        }
173        Self::from_big(!self.big().as_ref())
174    }
175    pub fn is_negative(&self) -> bool {
176        match &self.0 {
177            Repr::Small(n) => *n < 0,
178            Repr::Large(n) => n.value.is_negative(),
179        }
180    }
181    pub fn is_zero(&self) -> bool {
182        matches!(self.0, Repr::Small(0))
183    }
184    pub fn is_odd(&self) -> bool {
185        match &self.0 {
186            Repr::Small(n) => n & 1 != 0,
187            Repr::Large(n) => n.value.bit(0),
188        }
189    }
190    pub fn magnitude_bit(&self, i: u64) -> bool {
191        match &self.0 {
192            Repr::Small(n) => i < 64 && (n.unsigned_abs() >> i) & 1 != 0,
193            Repr::Large(n) => n.value.magnitude().bit(i),
194        }
195    }
196    pub fn negated(&self) -> Self {
197        if let Some(n) = self.to_i64().and_then(i64::checked_neg) {
198            return n.into();
199        }
200        Self::from_big(-self.big().as_ref()).expect("negation preserves magnitude")
201    }
202    pub fn abs(&self) -> Self {
203        if self.is_negative() {
204            self.negated()
205        } else {
206            self.clone()
207        }
208    }
209
210    pub fn add(&self, rhs: &Self) -> Result<Self, IntError> {
211        if let (Some(a), Some(b)) = (self.to_i64(), rhs.to_i64()) {
212            if let Some(n) = a.checked_add(b) {
213                return Ok(n.into());
214            }
215        }
216        Self::from_big(self.big().as_ref() + rhs.big().as_ref())
217    }
218    pub fn sub(&self, rhs: &Self) -> Result<Self, IntError> {
219        if let (Some(a), Some(b)) = (self.to_i64(), rhs.to_i64()) {
220            if let Some(n) = a.checked_sub(b) {
221                return Ok(n.into());
222            }
223        }
224        Self::from_big(self.big().as_ref() - rhs.big().as_ref())
225    }
226    pub fn mul(&self, rhs: &Self) -> Result<Self, IntError> {
227        if let (Some(a), Some(b)) = (self.to_i64(), rhs.to_i64()) {
228            if let Some(n) = a.checked_mul(b) {
229                return Ok(n.into());
230            }
231        }
232        if !self.is_zero() && !rhs.is_zero() && self.bits() + rhs.bits() - 1 > MAX_INTEGER_BITS {
233            return Err(IntError::TooLarge);
234        }
235        Self::from_big(self.big().as_ref() * rhs.big().as_ref())
236    }
237    pub fn div_mod(&self, rhs: &Self) -> Result<(Self, Self), IntError> {
238        if rhs.is_zero() {
239            return Err(IntError::DivisionByZero);
240        }
241        if let (Some(a), Some(b)) = (self.to_i64(), rhs.to_i64()) {
242            if let Some(mut q) = a.checked_div(b) {
243                let mut r = a % b;
244                if r != 0 && (r < 0) != (b < 0) {
245                    q -= 1;
246                    r += b;
247                }
248                return Ok((q.into(), r.into()));
249            }
250        }
251        let (a, b) = (self.big(), rhs.big());
252        let mut q = a.as_ref() / b.as_ref();
253        let mut r = a.as_ref() % b.as_ref();
254        if !r.is_zero() && r.is_negative() != b.is_negative() {
255            q -= 1;
256            r += b.as_ref();
257        }
258        Ok((Self::from_big(q)?, Self::from_big(r)?))
259    }
260    pub fn pow(&self, exponent: u32) -> Result<Self, IntError> {
261        if let Some(n) = self.to_i64().and_then(|n| n.checked_pow(exponent)) {
262            return Ok(n.into());
263        }
264        if self.bits().saturating_sub(1).saturating_mul(u64::from(exponent)) >= MAX_INTEGER_BITS {
265            return Err(IntError::TooLarge);
266        }
267        Self::from_big(self.big().pow(exponent))
268    }
269    /// Convert the quotient jointly: individually converting huge operands
270    /// would turn a finite ratio into infinity / infinity.
271    pub fn ratio(&self, denominator: &Self) -> Option<f64> {
272        if denominator.is_zero() {
273            return None;
274        }
275        if let (Some(a), Some(b)) = (self.to_i64(), denominator.to_i64()) {
276            if a.unsigned_abs() <= 1 << 53 && b.unsigned_abs() <= 1 << 53 {
277                return Some(a as f64 / b as f64);
278            }
279        }
280        BigRational::new_raw(self.big().into_owned(), denominator.big().into_owned())
281            .to_f64()
282            .filter(|x| x.is_finite())
283    }
284
285    /// ceil(n*p) for a nonnegative n and a finite probability. The float is
286    /// interpreted exactly, so rank selection works beyond float precision.
287    /// The temporary product uses at most 53 extra bits above n's size.
288    pub fn probability_rank(&self, p: f64) -> Result<Self, IntError> {
289        if self.is_negative() || !p.is_finite() || !(0.0..=1.0).contains(&p) {
290            return Err(IntError::Invalid);
291        }
292        let fraction = BigRational::from_float(p).ok_or(IntError::Invalid)?;
293        let product = self.big().as_ref() * fraction.numer();
294        let denominator = fraction.denom();
295        let q = &product / denominator;
296        Self::from_big(if (&product % denominator).is_zero() { q } else { q + 1 })
297    }
298
299    /// Round an integer to a power of ten, halfway away from zero. The
300    /// temporary scale is bounded to at most a few bits above the value ceiling.
301    pub fn round_decimal(&self, places: u32) -> Result<Self, IntError> {
302        if places as usize > MAX_INTEGER_DIGITS || u64::from(places) > self.bits() * 30103 / 100000 + 1 {
303            return Ok(Self::ZERO);
304        }
305        let scale = BigInt::from(10).pow(places);
306        let magnitude = self.big().abs();
307        let rounded = (magnitude + &scale / 2u32) / &scale * &scale;
308        Self::from_big(if self.is_negative() { -rounded } else { rounded })
309    }
310    /// Compare to the exact value represented by a float, without converting
311    /// the integer to a rounded float. Finite floats need at most 1024 bits.
312    pub fn cmp_f64(&self, x: f64) -> Option<Ordering> {
313        if x.is_nan() {
314            return None;
315        }
316        if x == f64::INFINITY {
317            return Some(Ordering::Less);
318        }
319        if x == f64::NEG_INFINITY {
320            return Some(Ordering::Greater);
321        }
322        let truncated = Self::from_f64(x).expect("finite float fits the integer ceiling");
323        let cmp = self.cmp(&truncated);
324        Some(if cmp.is_eq() {
325            0.0f64.partial_cmp(&x.fract()).unwrap()
326        } else {
327            cmp
328        })
329    }
330}
331
332impl FromStr for Integer {
333    type Err = IntError;
334    fn from_str(s: &str) -> Result<Self, Self::Err> {
335        if let Ok(n) = s.parse::<i64>() {
336            return Ok(n.into());
337        }
338        let digits = s.strip_prefix(['+', '-']).unwrap_or(s);
339        if digits.is_empty() || !digits.bytes().all(|c| c.is_ascii_digit()) {
340            return Err(IntError::Invalid);
341        }
342        let digits = digits.trim_start_matches('0');
343        if digits.len() > MAX_INTEGER_DIGITS {
344            return Err(IntError::TooLarge);
345        }
346        if digits.is_empty() {
347            return Ok(Self::ZERO);
348        }
349        let n = BigInt::from_str(digits).map_err(|_| IntError::Invalid)?;
350        Self::from_big(if s.starts_with('-') { -n } else { n })
351    }
352}
353impl From<i64> for Integer {
354    fn from(n: i64) -> Self {
355        Self(Repr::Small(n))
356    }
357}
358macro_rules! from_integer {
359    ($($t:ty),*) => { $(impl From<$t> for Integer { fn from(n: $t) -> Self { match i64::try_from(n) { Ok(n) => Self::from(n), Err(_) => Self::from_big(BigInt::from(n)).expect("machine integer fits") } } })* };
360}
361impl From<i32> for Integer {
362    fn from(n: i32) -> Self {
363        Self::from(i64::from(n))
364    }
365}
366impl From<u32> for Integer {
367    fn from(n: u32) -> Self {
368        Self::from(i64::from(n))
369    }
370}
371from_integer!(u64, usize, i128, u128);
372impl PartialEq for Integer {
373    #[inline]
374    fn eq(&self, rhs: &Self) -> bool {
375        match (&self.0, &rhs.0) {
376            (Repr::Small(a), Repr::Small(b)) => a == b,
377            (Repr::Large(a), Repr::Large(b)) => Arc::ptr_eq(a, b) || (a.hash == b.hash && a.value == b.value),
378            // A large integer is outside the range of a small one (`from_big`).
379            _ => false,
380        }
381    }
382}
383impl Eq for Integer {}
384impl PartialOrd for Integer {
385    fn partial_cmp(&self, rhs: &Self) -> Option<Ordering> {
386        Some(self.cmp(rhs))
387    }
388}
389impl Ord for Integer {
390    fn cmp(&self, rhs: &Self) -> Ordering {
391        match (&self.0, &rhs.0) {
392            (Repr::Small(a), Repr::Small(b)) => a.cmp(b),
393            (Repr::Large(a), Repr::Large(b)) => a.value.cmp(&b.value),
394            (Repr::Large(a), Repr::Small(_)) => {
395                if a.value.is_negative() {
396                    Ordering::Less
397                } else {
398                    Ordering::Greater
399                }
400            }
401            (Repr::Small(_), Repr::Large(b)) => {
402                if b.value.is_negative() {
403                    Ordering::Greater
404                } else {
405                    Ordering::Less
406                }
407            }
408        }
409    }
410}
411impl PartialEq<i64> for Integer {
412    fn eq(&self, rhs: &i64) -> bool {
413        self.to_i64() == Some(*rhs)
414    }
415}
416impl PartialOrd<i64> for Integer {
417    fn partial_cmp(&self, rhs: &i64) -> Option<Ordering> {
418        Some(self.cmp(&Self::from(*rhs)))
419    }
420}
421impl Hash for Integer {
422    fn hash<H: Hasher>(&self, h: &mut H) {
423        match &self.0 {
424            Repr::Small(n) => {
425                0u8.hash(h);
426                n.hash(h);
427            }
428            Repr::Large(n) => {
429                1u8.hash(h);
430                n.hash.hash(h);
431            }
432        }
433    }
434}
435impl fmt::Display for Integer {
436    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
437        match &self.0 {
438            Repr::Small(n) => n.fmt(f),
439            Repr::Large(n) => n.value.fmt(f),
440        }
441    }
442}
443impl fmt::Debug for Integer {
444    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
445        fmt::Display::fmt(self, f)
446    }
447}
448
449#[cfg(test)]
450mod tests {
451    use super::*;
452
453    #[test]
454    fn radix_parsing_is_exact_and_checks_limits_before_allocation() {
455        for bits in [0, 1, 31, 32, 53, 63, 64, 65, 127, 1024, 65535] {
456            let n: BigInt = (BigInt::from(1) << bits) - 1u32;
457            let expected = Integer::from_big(n.clone()).unwrap();
458            for radix in [2, 16] {
459                let digits = n.to_str_radix(radix);
460                assert_eq!(Integer::from_radix_digits(&digits, radix).unwrap(), expected);
461                assert_eq!(
462                    Integer::from_radix_digits(&format!("000{digits}"), radix).unwrap(),
463                    expected
464                );
465            }
466        }
467        for (radix, digits) in [(2, "1".repeat(65536)), (16, "f".repeat(16384))] {
468            assert_eq!(
469                Integer::from_radix_digits(&digits, radix).unwrap().bits(),
470                MAX_INTEGER_BITS
471            );
472            assert_eq!(
473                Integer::from_radix_digits(&format!("1{digits}"), radix),
474                Err(IntError::TooLarge)
475            );
476        }
477        for radix in [2, 16] {
478            assert_eq!(
479                Integer::from_radix_digits(&"0".repeat(100_000), radix).unwrap(),
480                Integer::ZERO
481            );
482            for digits in ["", "_1", "1_", "-1", "+1", "1g"] {
483                assert_eq!(Integer::from_radix_digits(digits, radix), Err(IntError::Invalid));
484            }
485        }
486        assert_eq!(Integer::from_radix_digits("1", 0), Err(IntError::Invalid));
487        assert!("0xff".parse::<Integer>().is_err()); // data parsing stays decimal
488    }
489
490    #[test]
491    fn bit_operations_match_signed_machine_integers_across_storage_boundaries() {
492        let values = [
493            i128::MIN,
494            i64::MIN as i128 - 1,
495            i64::MIN as i128,
496            -101,
497            -1,
498            0,
499            1,
500            101,
501            i64::MAX as i128,
502            i64::MAX as i128 + 1,
503            i128::MAX,
504        ];
505        for a in values {
506            let x = Integer::from(a);
507            assert_eq!(x.bit_not().unwrap(), Integer::from(!a));
508            assert_eq!(x.bit_count(), u64::from(a.unsigned_abs().count_ones()));
509            for b in values {
510                let y = Integer::from(b);
511                assert_eq!(x.bit_and(&y).unwrap(), Integer::from(a & b));
512                assert_eq!(x.bit_or(&y).unwrap(), Integer::from(a | b));
513                assert_eq!(x.bit_xor(&y).unwrap(), Integer::from(a ^ b));
514            }
515            // Results that fit in i64 must regain the inline representation.
516            assert_eq!(x.bit_xor(&x).unwrap().to_i64(), Some(0));
517            assert_eq!(x.bit_or(&(-1).into()).unwrap().to_i64(), Some(-1));
518        }
519    }
520
521    #[test]
522    fn bit_operations_obey_the_magnitude_ceiling() {
523        let max = Integer::from_radix_digits(&"f".repeat(16384), 16).unwrap();
524        assert_eq!(max.bit_count(), MAX_INTEGER_BITS);
525        assert_eq!(max.negated().bit_count(), MAX_INTEGER_BITS);
526        assert_eq!(max.bit_not(), Err(IntError::TooLarge));
527        assert_eq!(max.bit_xor(&(-1).into()), Err(IntError::TooLarge));
528        assert_eq!(max.bit_and(&(-1).into()).unwrap(), max);
529        assert_eq!(max.bit_or(&(-1).into()).unwrap().to_i64(), Some(-1));
530        let below = max.sub(&Integer::ONE).unwrap();
531        assert_eq!(max.negated().bit_not().unwrap(), below);
532        assert_eq!(below.bit_not().unwrap(), max.negated());
533        // AND can also require one more magnitude bit for negative operands.
534        assert_eq!(max.negated().bit_and(&(-2).into()), Err(IntError::TooLarge));
535    }
536
537    #[test]
538    fn signed_arithmetic_matches_wider_machine_integers() {
539        let values = [
540            i64::MIN as i128 - 1,
541            i64::MIN as i128,
542            -101,
543            -1,
544            0,
545            1,
546            101,
547            i64::MAX as i128,
548            i64::MAX as i128 + 1,
549        ];
550        for a in values {
551            for b in values {
552                let (x, y) = (Integer::from(a), Integer::from(b));
553                assert_eq!(x.add(&y).unwrap().to_string(), (a + b).to_string());
554                assert_eq!(x.sub(&y).unwrap().to_string(), (a - b).to_string());
555                assert_eq!(x.mul(&y).unwrap().to_string(), (a * b).to_string());
556                if b != 0 {
557                    let (mut q, mut r) = (a / b, a % b);
558                    if r != 0 && (r < 0) != (b < 0) {
559                        q -= 1;
560                        r += b;
561                    }
562                    let actual = x.div_mod(&y).unwrap();
563                    assert_eq!(actual, (q.into(), r.into()));
564                    assert_eq!(actual.0.mul(&y).unwrap().add(&actual.1).unwrap(), x);
565                }
566            }
567        }
568    }
569
570    #[test]
571    fn comparisons_agree_with_exact_rationals() {
572        let integers = [
573            "-10000000000000000000000000000000001",
574            "-9223372036854775809",
575            "-9007199254740993",
576            "-1",
577            "0",
578            "1",
579            "9007199254740993",
580            "9223372036854775807",
581            "10000000000000000000000000000000001",
582        ];
583        let mut seed = 123456789u64;
584        for _ in 0..2000 {
585            seed ^= seed << 13;
586            seed ^= seed >> 7;
587            seed ^= seed << 17;
588            let f = f64::from_bits(seed);
589            if !f.is_finite() {
590                continue;
591            }
592            let rational = BigRational::from_float(f).unwrap();
593            for text in integers {
594                let n: Integer = text.parse().unwrap();
595                assert_eq!(
596                    n.cmp_f64(f),
597                    Some(BigRational::from_integer(n.big().into_owned()).cmp(&rational)),
598                    "{n}, {f}"
599                );
600            }
601        }
602    }
603
604    #[test]
605    fn parsing_and_operations_enforce_the_bit_ceiling() {
606        let max = Integer::from(2).pow(65535).unwrap();
607        assert_eq!(max.bits(), MAX_INTEGER_BITS);
608        assert_eq!(max.to_string().parse::<Integer>().unwrap(), max);
609        assert_eq!(max.mul(&2.into()), Err(IntError::TooLarge));
610        assert_eq!(Integer::from(2).pow(65536), Err(IntError::TooLarge));
611        assert_eq!(
612            "9".repeat(MAX_INTEGER_DIGITS + 1).parse::<Integer>(),
613            Err(IntError::TooLarge)
614        );
615        assert_eq!("0".repeat(100_000).parse::<Integer>().unwrap(), Integer::ZERO);
616        assert_eq!(
617            format!("-{}12", "0".repeat(100_000)).parse::<Integer>().unwrap(),
618            Integer::from(-12)
619        );
620        assert_eq!("+".parse::<Integer>(), Err(IntError::Invalid));
621    }
622}