1use 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;
17pub 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 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 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 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 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 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 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 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 _ => 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()); }
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 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 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}