use core::f64::consts::LOG10_2;
use malachite_bigint::{BigInt, ToBigInt};
use num_traits::{Signed, ToPrimitive};
#[must_use]
pub const fn decompose_float(value: f64) -> (f64, i32) {
if value == 0.0 {
return (0.0, 0);
}
let bits = value.to_bits();
let (bits, exponent_adjust) = if (bits >> 52) & 0x7ff == 0 {
((value * (1u64 << 54) as f64).to_bits(), -54)
} else {
(bits, 0)
};
let exponent: i32 = ((bits >> 52) & 0x7ff) as i32 - 1022 + exponent_adjust;
let mantissa_bits = bits & (0x000f_ffff_ffff_ffff) | (1022 << 52);
(f64::from_bits(mantissa_bits), exponent)
}
#[must_use]
pub fn eq_int(value: f64, other: &BigInt) -> bool {
if let (Some(self_int), Some(other_float)) = (value.to_bigint(), other.to_f64()) {
value == other_float && self_int == *other
} else {
false
}
}
#[must_use]
pub fn lt_int(value: f64, other_int: &BigInt) -> bool {
match (value.to_bigint(), other_int.to_f64()) {
(Some(self_int), Some(other_float)) => value < other_float || self_int < *other_int,
(Some(_), None) => other_int.is_positive(),
_ if value.is_infinite() => value.is_sign_negative(),
_ => false,
}
}
#[must_use]
pub fn gt_int(value: f64, other_int: &BigInt) -> bool {
match (value.to_bigint(), other_int.to_f64()) {
(Some(self_int), Some(other_float)) => value > other_float || self_int > *other_int,
(Some(_), None) => other_int.is_negative(),
_ if value.is_infinite() => value.is_sign_positive(),
_ => false,
}
}
#[must_use]
pub const fn div(v1: f64, v2: f64) -> Option<f64> {
if v2 != 0.0 { Some(v1 / v2) } else { None }
}
#[must_use]
pub fn mod_(v1: f64, v2: f64) -> Option<f64> {
divmod(v1, v2).map(|(_, m)| m)
}
#[must_use]
pub fn floordiv(v1: f64, v2: f64) -> Option<f64> {
divmod(v1, v2).map(|(d, _)| d)
}
#[must_use]
pub fn divmod(v1: f64, v2: f64) -> Option<(f64, f64)> {
if v2 == 0.0 {
return None;
}
let mut m = v1 % v2;
let mut d = (v1 - m) / v2;
if m != 0.0 {
if v2.is_sign_negative() != m.is_sign_negative() {
m += v2;
d -= 1.0;
}
} else {
m = (0.0_f64).copysign(v2);
}
let d = if d != 0.0 {
let f = d.floor();
if d - f > 0.5 { f + 1.0 } else { f }
} else {
(0.0_f64).copysign(v1 / v2)
};
Some((d, m))
}
#[allow(clippy::float_cmp)]
#[must_use]
pub fn nextafter(x: f64, y: f64) -> f64 {
if x == y {
y
} else if x.is_nan() || y.is_nan() {
f64::NAN
} else if x >= f64::INFINITY {
f64::MAX
} else if x <= f64::NEG_INFINITY {
f64::MIN
} else if x == 0.0 {
f64::from_bits(1).copysign(y)
} else {
let b = x.to_bits();
let bits = if (y > x) == (x > 0.0) { b + 1 } else { b - 1 };
let ret = f64::from_bits(bits);
if ret == 0.0 { ret.copysign(x) } else { ret }
}
}
#[allow(clippy::float_cmp)]
#[must_use]
pub fn nextafter_with_steps(x: f64, y: f64, steps: u64) -> f64 {
if x == y {
y
} else if x.is_nan() || y.is_nan() {
f64::NAN
} else if x >= f64::INFINITY {
f64::MAX
} else if x <= f64::NEG_INFINITY {
f64::MIN
} else if x == 0.0 {
f64::from_bits(1).copysign(y)
} else {
if steps == 0 {
return x;
}
if x.is_nan() {
return x;
}
if y.is_nan() {
return y;
}
let sign_bit: u64 = 1 << 63;
let mut ux = x.to_bits();
let uy = y.to_bits();
let ax = ux & !sign_bit;
let ay = uy & !sign_bit;
if ((ux ^ uy) & sign_bit) != 0 {
return if ax + ay <= steps {
f64::from_bits(uy)
} else if ax < steps {
let result = (uy & sign_bit) | (steps - ax);
f64::from_bits(result)
} else {
ux -= steps;
f64::from_bits(ux)
};
}
if ax > ay {
if ax - ay >= steps {
ux -= steps;
f64::from_bits(ux)
} else {
f64::from_bits(uy)
}
} else if ay - ax >= steps {
ux += steps;
f64::from_bits(ux)
} else {
f64::from_bits(uy)
}
}
}
#[must_use]
pub fn ulp(x: f64) -> f64 {
if x.is_nan() {
return x;
}
let x = x.abs();
let x2 = nextafter(x, f64::INFINITY);
if x2.is_infinite() {
let x2 = nextafter(x, f64::NEG_INFINITY);
x - x2
} else {
x2 - x
}
}
#[must_use]
pub fn round_float_digits(x: f64, ndigits: i32) -> Option<f64> {
if !x.is_finite() {
return Some(x);
}
const NDIGITS_MAX: i32 = ((f64::MANTISSA_DIGITS as i32 - f64::MIN_EXP) as f64 * LOG10_2) as i32;
const NDIGITS_MIN: i32 = -(((f64::MAX_EXP + 1) as f64 * LOG10_2) as i32);
if ndigits > NDIGITS_MAX {
return Some(x);
}
if ndigits < NDIGITS_MIN {
return Some(0.0f64.copysign(x));
}
let result: f64 = if ndigits >= 0 {
let s = format!("{:.*}", ndigits as usize, x);
s.parse().ok()?
} else {
round_at_power_of_ten(x, (-ndigits) as usize)?
};
if !result.is_finite() {
return None;
}
Some(result)
}
fn round_at_power_of_ten(x: f64, place: usize) -> Option<f64> {
let digits = format!("{:.0}", x.trunc().abs());
let has_fraction = x.fract() != 0.0;
let padded = format!("{digits:0>width$}", width = place + 1);
let (kept, dropped) = padded.split_at(padded.len() - place);
let mut kept: Vec<u8> = kept.bytes().collect();
let round_up = match dropped.as_bytes().split_first() {
None => false,
Some((&first, rest)) => {
first > b'5'
|| (first == b'5'
&& (rest.iter().any(|&digit| digit != b'0')
|| has_fraction
|| kept.last().is_some_and(|digit| (digit - b'0') % 2 == 1)))
}
};
if round_up {
let carried = kept.iter_mut().rev().all(|digit| {
*digit = if *digit == b'9' { b'0' } else { *digit + 1 };
*digit == b'0'
});
if carried {
kept.insert(0, b'1');
}
}
let mut rounded = String::from_utf8(kept).ok()?;
rounded.extend(core::iter::repeat_n('0', place));
let magnitude: f64 = rounded.parse().ok()?;
Some(magnitude.copysign(x))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HexFloatError {
Invalid,
TooLong,
Overflow,
}
const DBL_MANT_DIG: i64 = 53;
const DBL_MIN_EXP: i64 = -1021;
const DBL_MAX_EXP: i64 = 1024;
#[inline]
fn byte_at(bytes: &[u8], i: usize) -> Option<u8> {
bytes.get(i).copied()
}
#[inline]
const fn hex_from_char(c: u8) -> Option<u8> {
match c {
b'0'..=b'9' => Some(c - b'0'),
b'a'..=b'f' => Some(c - b'a' + 10),
b'A'..=b'F' => Some(c - b'A' + 10),
_ => None,
}
}
#[inline]
fn hex_digit_at(bytes: &[u8], i: usize) -> Option<u8> {
byte_at(bytes, i).and_then(hex_from_char)
}
fn case_insensitive_match(bytes: &[u8], s: usize, t: &[u8]) -> bool {
let mut si = s;
let mut ti = 0;
while ti < t.len() && byte_at(bytes, si).is_some_and(|b| b.to_ascii_lowercase() == t[ti]) {
si += 1;
ti += 1;
}
ti == t.len()
}
fn parse_inf_or_nan(bytes: &[u8], p: usize) -> Option<(f64, usize)> {
let mut s = p;
let mut negate = false;
if byte_at(bytes, s) == Some(b'-') {
negate = true;
s += 1;
} else if byte_at(bytes, s) == Some(b'+') {
s += 1;
}
if case_insensitive_match(bytes, s, b"inf") {
s += 3;
if case_insensitive_match(bytes, s, b"inity") {
s += 5;
}
let value = if negate {
f64::NEG_INFINITY
} else {
f64::INFINITY
};
Some((value, s))
} else if case_insensitive_match(bytes, s, b"nan") {
s += 3;
let value = if negate {
f64::from_bits(0xfff8_0000_0000_0000)
} else {
f64::from_bits(0x7ff8_0000_0000_0000)
};
Some((value, s))
} else {
None
}
}
const fn ldexp(x: f64, mut n: i32) -> f64 {
let x1p1023 = f64::from_bits(0x7fe0000000000000);
let x1p53 = f64::from_bits(0x4340000000000000);
let x1p_1022 = f64::from_bits(0x0010000000000000);
let mut y = x;
if n > 1023 {
y *= x1p1023;
n -= 1023;
if n > 1023 {
y *= x1p1023;
n -= 1023;
if n > 1023 {
n = 1023;
}
}
} else if n < -1022 {
y *= x1p_1022 * x1p53;
n += 1022 - 53;
if n < -1022 {
y *= x1p_1022 * x1p53;
n += 1022 - 53;
if n < -1022 {
n = -1022;
}
}
}
y * f64::from_bits(((0x3ff + n) as u64) << 52)
}
fn strtol_saturating(bytes: &[u8], start: usize, end: usize) -> i64 {
let mut i = start;
let mut neg = false;
if i < end && (bytes[i] == b'+' || bytes[i] == b'-') {
neg = bytes[i] == b'-';
i += 1;
}
let mut val: i64 = 0;
let mut overflowed = false;
while i < end {
let d = (bytes[i] - b'0') as i64;
match val.checked_mul(10).and_then(|v| v.checked_add(d)) {
Some(v) => val = v,
None => {
overflowed = true;
break;
}
}
i += 1;
}
if overflowed {
if neg { i64::MIN } else { i64::MAX }
} else if neg {
-val
} else {
val
}
}
pub fn from_hex(s: &str) -> Result<f64, HexFloatError> {
let bytes = s.as_bytes();
let s_end = bytes.len();
let mut negate = false;
let mut idx = 0usize;
let mut x;
while byte_at(bytes, idx).is_some_and(rustpython_wtf8::is_py_ascii_whitespace) {
idx += 1;
}
if let Some((value, end)) = parse_inf_or_nan(bytes, idx) {
idx = end;
return finish_hex(bytes, s_end, idx, negate, value);
}
if byte_at(bytes, idx) == Some(b'-') {
idx += 1;
negate = true;
} else if byte_at(bytes, idx) == Some(b'+') {
idx += 1;
}
let s_store = idx;
if byte_at(bytes, idx) == Some(b'0') {
idx += 1;
if matches!(byte_at(bytes, idx), Some(b'x' | b'X')) {
idx += 1;
} else {
idx = s_store;
}
}
let coeff_start = idx;
while hex_digit_at(bytes, idx).is_some() {
idx += 1;
}
let s_store = idx;
let coeff_end = if byte_at(bytes, idx) == Some(b'.') {
idx += 1;
while hex_digit_at(bytes, idx).is_some() {
idx += 1;
}
idx - 1
} else {
idx
};
let ndigits_total = (coeff_end - coeff_start) as i64;
let fdigits = (coeff_end - s_store) as i64;
if ndigits_total == 0 {
return Err(HexFloatError::Invalid);
}
let insane_bound = core::cmp::min(
DBL_MIN_EXP - DBL_MANT_DIG - i64::MIN / 2,
i64::MAX / 2 + 1 - DBL_MAX_EXP,
) / 4;
if ndigits_total > insane_bound {
return Err(HexFloatError::TooLong);
}
let exp = if matches!(byte_at(bytes, idx), Some(b'p' | b'P')) {
idx += 1;
let exp_start = idx;
if matches!(byte_at(bytes, idx), Some(b'-' | b'+')) {
idx += 1;
}
if !matches!(byte_at(bytes, idx), Some(b'0'..=b'9')) {
return Err(HexFloatError::Invalid);
}
idx += 1;
while matches!(byte_at(bytes, idx), Some(b'0'..=b'9')) {
idx += 1;
}
strtol_saturating(bytes, exp_start, idx)
} else {
0
};
let hex_digit = |j: i64| -> i32 {
let byte_idx = if j < fdigits {
coeff_end as i64 - j
} else {
coeff_end as i64 - 1 - j
};
hex_digit_at(bytes, byte_idx as usize).expect("hex digit within coefficient") as i32
};
let mut ndigits = ndigits_total;
while ndigits > 0 && hex_digit(ndigits - 1) == 0 {
ndigits -= 1;
}
if ndigits == 0 || exp < i64::MIN / 2 {
x = 0.0;
return finish_hex(bytes, s_end, idx, negate, x);
}
if exp > i64::MAX / 2 {
return Err(HexFloatError::Overflow);
}
let exp = exp - 4 * fdigits;
let mut top_exp = exp + 4 * (ndigits - 1);
let mut digit = hex_digit(ndigits - 1);
while digit != 0 {
top_exp += 1;
digit /= 2;
}
if top_exp < DBL_MIN_EXP - DBL_MANT_DIG {
x = 0.0;
return finish_hex(bytes, s_end, idx, negate, x);
}
if top_exp > DBL_MAX_EXP {
return Err(HexFloatError::Overflow);
}
let lsb = core::cmp::max(top_exp, DBL_MIN_EXP) - DBL_MANT_DIG;
x = 0.0;
if exp >= lsb {
let mut i = ndigits - 1;
while i >= 0 {
x = 16.0 * x + hex_digit(i) as f64;
i -= 1;
}
x = ldexp(x, exp as i32);
return finish_hex(bytes, s_end, idx, negate, x);
}
let half_eps: i32 = 1 << ((lsb - exp - 1) % 4) as i32;
let key_digit = (lsb - exp - 1) / 4;
let mut i = ndigits - 1;
while i > key_digit {
x = 16.0 * x + hex_digit(i) as f64;
i -= 1;
}
let digit = hex_digit(key_digit);
x = 16.0 * x + (digit & (16 - 2 * half_eps)) as f64;
if (digit & half_eps) != 0 {
let round_up = if (digit & (3 * half_eps - 1)) != 0
|| (half_eps == 8 && key_digit + 1 < ndigits && (hex_digit(key_digit + 1) & 1) != 0)
{
true
} else {
let mut r = false;
let mut i = key_digit - 1;
while i >= 0 {
if hex_digit(i) != 0 {
r = true;
break;
}
i -= 1;
}
r
};
if round_up {
x += (2 * half_eps) as f64;
if top_exp == DBL_MAX_EXP && x == ldexp((2 * half_eps) as f64, DBL_MANT_DIG as i32) {
return Err(HexFloatError::Overflow);
}
}
}
x = ldexp(x, (exp + 4 * key_digit) as i32);
finish_hex(bytes, s_end, idx, negate, x)
}
fn finish_hex(
bytes: &[u8],
s_end: usize,
mut idx: usize,
negate: bool,
x: f64,
) -> Result<f64, HexFloatError> {
while byte_at(bytes, idx).is_some_and(rustpython_wtf8::is_py_ascii_whitespace) {
idx += 1;
}
if idx != s_end {
return Err(HexFloatError::Invalid);
}
Ok(if negate { -x } else { x })
}
#[cfg(test)]
mod from_hex_tests {
use super::{HexFloatError, from_hex};
fn bits(s: &str) -> u64 {
from_hex(s).unwrap().to_bits()
}
#[test]
fn from_hex_exact_bits() {
assert_eq!(bits("0x1p-1074"), 0x0000000000000001);
assert_eq!(bits("0x1.fffffffffffffp+1023"), 0x7fefffffffffffff);
assert_eq!(bits("0x1.00000000000008p0"), 0x3ff0000000000000);
assert_eq!(bits("0x1.00000000000018p0"), 0x3ff0000000000002);
assert_eq!(bits("-0x1p0"), 0xbff0000000000000);
assert_eq!(bits("0x0p0"), 0x0000000000000000);
assert_eq!(bits("-0x0p0"), 0x8000000000000000);
}
#[test]
fn from_hex_inf_nan() {
assert_eq!(bits("inf"), 0x7ff0000000000000);
assert_eq!(bits("-inf"), 0xfff0000000000000);
assert_eq!(bits("Infinity"), 0x7ff0000000000000);
let n = from_hex("nan").unwrap();
assert!(n.is_nan());
assert_eq!(n.to_bits(), 0x7ff8000000000000);
let neg = from_hex("-nan").unwrap();
assert!(neg.is_nan());
assert_eq!(neg.to_bits(), 0xfff8000000000000);
}
#[test]
fn from_hex_whitespace() {
assert_eq!(bits(" 0x1p0 "), 0x3ff0000000000000);
assert_eq!(bits("\t0x1p0\n"), 0x3ff0000000000000);
}
#[test]
fn from_hex_errors() {
assert_eq!(from_hex("0x1p1024"), Err(HexFloatError::Overflow));
assert_eq!(from_hex("0x1z"), Err(HexFloatError::Invalid));
assert_eq!(from_hex(""), Err(HexFloatError::Invalid));
assert_eq!(from_hex("0x1 p0"), Err(HexFloatError::Invalid));
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hash::hash_float;
fn pow2(e: i32) -> f64 {
if e >= -1022 {
f64::from_bits(((e + 1023) as u64) << 52)
} else {
f64::from_bits(1u64 << (e + 1074))
}
}
#[test]
fn decompose_float_frexp_contract() {
let mut values = alloc::vec![
0.0,
f64::from_bits(1), f64::from_bits(2),
f64::from_bits(0x000f_ffff_ffff_ffff), f64::MIN_POSITIVE, f64::from_bits(f64::MIN_POSITIVE.to_bits() - 1), 1.0,
1.5,
0.1,
core::f64::consts::PI,
];
for e in -1074..=1023 {
values.push(pow2(e));
values.push(-pow2(e));
}
for &v in &values {
let (m, e) = decompose_float(v);
if v == 0.0 {
assert_eq!((m, e), (0.0, 0));
continue;
}
assert!(
(0.5..1.0).contains(&m),
"mantissa {m} out of [0.5, 1) for value {v:e}"
);
let reconstructed = (m * 2.0) * pow2(e - 1);
assert_eq!(
reconstructed.to_bits(),
v.abs().to_bits(),
"reconstruction failed for {v:e}: m={m}, e={e}"
);
}
}
#[test]
fn hash_float_smallest_subnormal() {
assert_eq!(hash_float(f64::from_bits(1)), Some(16777216));
}
#[test]
fn hash_float_matches_cpython() {
const HASH_CASES: &[(u64, i64)] = &[
(0x0000000000000001, 16777216), (0x0000000000000002, 33554432), (0x00000000deadbeef, 62678480394911744), (0x0008000000000000, 16384), (0x000fffffffffffff, 2305843009196949503), (0x0010000000000000, 32768), (0x8000000000000001, -16777216), (0x0020000000000000, 65536), (0x0170000000000000, 137438953472), (0x39b0000000000000, 4194304), (0x3f50000000000000, 2251799813685248), (0x3fe0000000000000, 1152921504606846976), (0x3ff0000000000000, 1), (0x4000000000000000, 2), (0x4090000000000000, 1024), (0x4630000000000000, 549755813888), (0x7e70000000000000, 16777216), (0x7fe0000000000000, 140737488355328), (0xffe0000000000000, -140737488355328), (0x3ff8000000000000, 1152921504606846977), (0x400921fb54442d18, 326490430436040707), (0x7e37e43c8800759c, 1224995262755759164), (0x01a56e1fc2f8f359, 482449582752280463), (0x40c81cd6c8b43958, 1563361560246628409), (0x3fb999999999999a, 230584300921369408), (0x4005666666666666, 1556444031219243010), (0x4132d68700000000, 1234567), (0x44dfe154f457ea13, 1428027733287631914), (0x3c07a42f549647fb, 851769299698974080), (0xbff0000000000000, -2), (0xbfb999999999999a, -230584300921369408), ];
for &(bits, expected) in HASH_CASES {
let v = f64::from_bits(bits);
assert_eq!(
hash_float(v),
Some(expected),
"hash mismatch for {v:e} (bits {bits:#018x})"
);
}
}
}