Skip to main content

rustpython_common/
float_ops.rs

1use core::f64::consts::LOG10_2;
2use malachite_bigint::{BigInt, ToBigInt};
3use num_traits::{Signed, ToPrimitive};
4
5#[must_use]
6pub const fn decompose_float(value: f64) -> (f64, i32) {
7    if value == 0.0 {
8        return (0.0, 0);
9    }
10    let bits = value.to_bits();
11    // Subnormals carry a biased exponent of 0 and no implicit leading mantissa
12    // bit, so the normal decomposition below would misread them. Scale them up
13    // into the normal range first (exact, since it is a power-of-two shift) and
14    // fold the scale back into the returned exponent.
15    let (bits, exponent_adjust) = if (bits >> 52) & 0x7ff == 0 {
16        ((value * (1u64 << 54) as f64).to_bits(), -54)
17    } else {
18        (bits, 0)
19    };
20    let exponent: i32 = ((bits >> 52) & 0x7ff) as i32 - 1022 + exponent_adjust;
21    let mantissa_bits = bits & (0x000f_ffff_ffff_ffff) | (1022 << 52);
22    (f64::from_bits(mantissa_bits), exponent)
23}
24
25/// Equate an integer to a float.
26///
27/// Returns true if and only if, when converted to each others types, both are equal.
28///
29/// # Examples
30///
31/// ```
32/// use malachite_bigint::BigInt;
33/// use rustpython_common::float_ops::eq_int;
34/// let a = 1.0f64;
35/// let b = BigInt::from(1);
36/// let c = 2.0f64;
37/// assert!(eq_int(a, &b));
38/// assert!(!eq_int(c, &b));
39/// ```
40///
41#[must_use]
42pub fn eq_int(value: f64, other: &BigInt) -> bool {
43    if let (Some(self_int), Some(other_float)) = (value.to_bigint(), other.to_f64()) {
44        value == other_float && self_int == *other
45    } else {
46        false
47    }
48}
49
50#[must_use]
51pub fn lt_int(value: f64, other_int: &BigInt) -> bool {
52    match (value.to_bigint(), other_int.to_f64()) {
53        (Some(self_int), Some(other_float)) => value < other_float || self_int < *other_int,
54        // finite float, other_int too big for float,
55        // the result depends only on other_int’s sign
56        (Some(_), None) => other_int.is_positive(),
57        // infinite float must be bigger or lower than any int, depending on its sign
58        _ if value.is_infinite() => value.is_sign_negative(),
59        // NaN, always false
60        _ => false,
61    }
62}
63
64#[must_use]
65pub fn gt_int(value: f64, other_int: &BigInt) -> bool {
66    match (value.to_bigint(), other_int.to_f64()) {
67        (Some(self_int), Some(other_float)) => value > other_float || self_int > *other_int,
68        // finite float, other_int too big for float,
69        // the result depends only on other_int’s sign
70        (Some(_), None) => other_int.is_negative(),
71        // infinite float must be bigger or lower than any int, depending on its sign
72        _ if value.is_infinite() => value.is_sign_positive(),
73        // NaN, always false
74        _ => false,
75    }
76}
77
78#[must_use]
79pub const fn div(v1: f64, v2: f64) -> Option<f64> {
80    if v2 != 0.0 { Some(v1 / v2) } else { None }
81}
82
83#[must_use]
84pub fn mod_(v1: f64, v2: f64) -> Option<f64> {
85    divmod(v1, v2).map(|(_, m)| m)
86}
87
88#[must_use]
89pub fn floordiv(v1: f64, v2: f64) -> Option<f64> {
90    divmod(v1, v2).map(|(d, _)| d)
91}
92
93// Canonical (floordiv, mod) for floats matching CPython's _float_div_mod
94// (Objects/floatobject.c). `mod_` and `floordiv` delegate here so that
95// `divmod(a, b) == (a // b, a % b)` holds by construction.
96#[must_use]
97pub fn divmod(v1: f64, v2: f64) -> Option<(f64, f64)> {
98    if v2 == 0.0 {
99        return None;
100    }
101    let mut m = v1 % v2;
102    let mut d = (v1 - m) / v2;
103    if m != 0.0 {
104        // Non-zero remainder must have the sign of the divisor.
105        if v2.is_sign_negative() != m.is_sign_negative() {
106            m += v2;
107            d -= 1.0;
108        }
109    } else {
110        // Zero remainder: sign matches divisor (IEEE 754 / CPython contract).
111        m = (0.0_f64).copysign(v2);
112    }
113    let d = if d != 0.0 {
114        let f = d.floor();
115        // Snap up if (v1 - m) / v2 undershot the true integer quotient by
116        // more than half an ULP (mirrors CPython's `if (div - *floordiv > 0.5)`).
117        if d - f > 0.5 { f + 1.0 } else { f }
118    } else {
119        // Zero quotient: take the sign of the true quotient v1 / v2.
120        (0.0_f64).copysign(v1 / v2)
121    };
122    Some((d, m))
123}
124
125// nextafter algorithm based off of https://gitlab.com/bronsonbdevost/next_afterf
126#[allow(clippy::float_cmp)]
127#[must_use]
128pub fn nextafter(x: f64, y: f64) -> f64 {
129    if x == y {
130        y
131    } else if x.is_nan() || y.is_nan() {
132        f64::NAN
133    } else if x >= f64::INFINITY {
134        f64::MAX
135    } else if x <= f64::NEG_INFINITY {
136        f64::MIN
137    } else if x == 0.0 {
138        f64::from_bits(1).copysign(y)
139    } else {
140        // next x after 0 if y is farther from 0 than x, otherwise next towards 0
141        // the sign is a separate bit in floats, so bits+1 moves away from 0 no matter the float
142        let b = x.to_bits();
143        let bits = if (y > x) == (x > 0.0) { b + 1 } else { b - 1 };
144        let ret = f64::from_bits(bits);
145        if ret == 0.0 { ret.copysign(x) } else { ret }
146    }
147}
148
149#[allow(clippy::float_cmp)]
150#[must_use]
151pub fn nextafter_with_steps(x: f64, y: f64, steps: u64) -> f64 {
152    if x == y {
153        y
154    } else if x.is_nan() || y.is_nan() {
155        f64::NAN
156    } else if x >= f64::INFINITY {
157        f64::MAX
158    } else if x <= f64::NEG_INFINITY {
159        f64::MIN
160    } else if x == 0.0 {
161        f64::from_bits(1).copysign(y)
162    } else {
163        if steps == 0 {
164            return x;
165        }
166
167        if x.is_nan() {
168            return x;
169        }
170
171        if y.is_nan() {
172            return y;
173        }
174
175        let sign_bit: u64 = 1 << 63;
176
177        let mut ux = x.to_bits();
178        let uy = y.to_bits();
179
180        let ax = ux & !sign_bit;
181        let ay = uy & !sign_bit;
182
183        // If signs are different
184        if ((ux ^ uy) & sign_bit) != 0 {
185            return if ax + ay <= steps {
186                f64::from_bits(uy)
187            } else if ax < steps {
188                let result = (uy & sign_bit) | (steps - ax);
189                f64::from_bits(result)
190            } else {
191                ux -= steps;
192                f64::from_bits(ux)
193            };
194        }
195
196        // If signs are the same
197        if ax > ay {
198            if ax - ay >= steps {
199                ux -= steps;
200                f64::from_bits(ux)
201            } else {
202                f64::from_bits(uy)
203            }
204        } else if ay - ax >= steps {
205            ux += steps;
206            f64::from_bits(ux)
207        } else {
208            f64::from_bits(uy)
209        }
210    }
211}
212
213#[must_use]
214pub fn ulp(x: f64) -> f64 {
215    if x.is_nan() {
216        return x;
217    }
218    let x = x.abs();
219    let x2 = nextafter(x, f64::INFINITY);
220    if x2.is_infinite() {
221        // special case: x is the largest positive representable float
222        let x2 = nextafter(x, f64::NEG_INFINITY);
223        x - x2
224    } else {
225        x2 - x
226    }
227}
228
229#[must_use]
230pub fn round_float_digits(x: f64, ndigits: i32) -> Option<f64> {
231    // Mirror CPython's `float.__round__` (Objects/floatobject.c), which uses
232    // `_Py_dg_dtoa` to round at the decimal level. Multiplying by 10**ndigits
233    // and rounding at the IEEE 754 binary level diverges for values that
234    // aren't exactly representable: 2.675 stores as 2.67499..., which dtoa
235    // correctly rounds down to 2.67, but `(2.675 * 100.0).round() / 100.0`
236    // lands on 2.68 because the multiplication produces a phantom 267.5 tie.
237    // Rust's `{:.*}` float formatting uses dtoa-style algorithms and matches
238    // CPython's `_Py_dg_dtoa` byte-for-byte.
239    if !x.is_finite() {
240        return Some(x);
241    }
242
243    const NDIGITS_MAX: i32 = ((f64::MANTISSA_DIGITS as i32 - f64::MIN_EXP) as f64 * LOG10_2) as i32;
244    const NDIGITS_MIN: i32 = -(((f64::MAX_EXP + 1) as f64 * LOG10_2) as i32);
245
246    if ndigits > NDIGITS_MAX {
247        return Some(x);
248    }
249    if ndigits < NDIGITS_MIN {
250        return Some(0.0f64.copysign(x));
251    }
252
253    let result: f64 = if ndigits >= 0 {
254        let s = format!("{:.*}", ndigits as usize, x);
255        s.parse().ok()?
256    } else {
257        round_at_power_of_ten(x, (-ndigits) as usize)?
258    };
259
260    if !result.is_finite() {
261        return None;
262    }
263    Some(result)
264}
265
266/// Round `x` half-to-even at the `10.pow(place)` digit, the way
267/// `_Py_dg_dtoa` in mode 3 does for a negative `ndigits`.
268///
269/// The digits come from the exact decimal expansion rather than from a divide
270/// by `10.pow(place)`, which rounds twice: once when the quotient is stored
271/// and again at the tie.
272fn round_at_power_of_ten(x: f64, place: usize) -> Option<f64> {
273    // `trunc` is exact, so formatting it writes the integer digits themselves.
274    let digits = format!("{:.0}", x.trunc().abs());
275    let has_fraction = x.fract() != 0.0;
276
277    let padded = format!("{digits:0>width$}", width = place + 1);
278    let (kept, dropped) = padded.split_at(padded.len() - place);
279    let mut kept: Vec<u8> = kept.bytes().collect();
280
281    // The dropped digits read as a fraction of the rounding place: "5" then
282    // zeros is the tie, which the fraction below them, if any, tips over.
283    let round_up = match dropped.as_bytes().split_first() {
284        None => false,
285        Some((&first, rest)) => {
286            first > b'5'
287                || (first == b'5'
288                    && (rest.iter().any(|&digit| digit != b'0')
289                        || has_fraction
290                        || kept.last().is_some_and(|digit| (digit - b'0') % 2 == 1)))
291        }
292    };
293
294    if round_up {
295        // The carry stops at the first digit that does not wrap; if none of
296        // them stops it, the number has grown a digit.
297        let carried = kept.iter_mut().rev().all(|digit| {
298            *digit = if *digit == b'9' { b'0' } else { *digit + 1 };
299            *digit == b'0'
300        });
301        if carried {
302            kept.insert(0, b'1');
303        }
304    }
305
306    let mut rounded = String::from_utf8(kept).ok()?;
307    rounded.extend(core::iter::repeat_n('0', place));
308    let magnitude: f64 = rounded.parse().ok()?;
309    Some(magnitude.copysign(x))
310}
311
312/// Error from [`from_hex`], mapping to the exception the caller should raise.
313#[derive(Debug, Clone, Copy, PartialEq, Eq)]
314pub enum HexFloatError {
315    /// ValueError "invalid hexadecimal floating-point string"
316    Invalid,
317    /// ValueError "hexadecimal string too long to convert"
318    TooLong,
319    /// OverflowError "hexadecimal value too large to represent as a float"
320    Overflow,
321}
322
323const DBL_MANT_DIG: i64 = 53;
324const DBL_MIN_EXP: i64 = -1021;
325const DBL_MAX_EXP: i64 = 1024;
326
327/// Read byte at `i`, returning `None` past the end so that digit/sign/space
328/// scans stop at the string boundary.
329#[inline]
330fn byte_at(bytes: &[u8], i: usize) -> Option<u8> {
331    bytes.get(i).copied()
332}
333
334/// '0'-'9' -> 0..9, 'a'-'f'/'A'-'F' -> 10..15, else `None`.
335#[inline]
336const fn hex_from_char(c: u8) -> Option<u8> {
337    match c {
338        b'0'..=b'9' => Some(c - b'0'),
339        b'a'..=b'f' => Some(c - b'a' + 10),
340        b'A'..=b'F' => Some(c - b'A' + 10),
341        _ => None,
342    }
343}
344
345/// Peek at byte `i` and decode it as a hex digit, or `None` if it is out of
346/// range or not a hex digit.
347#[inline]
348fn hex_digit_at(bytes: &[u8], i: usize) -> Option<u8> {
349    byte_at(bytes, i).and_then(hex_from_char)
350}
351
352/// `t` must be an ASCII-lowercase literal. Returns true if every byte of `t`
353/// matched case-insensitively starting at `s`.
354fn case_insensitive_match(bytes: &[u8], s: usize, t: &[u8]) -> bool {
355    let mut si = s;
356    let mut ti = 0;
357    while ti < t.len() && byte_at(bytes, si).is_some_and(|b| b.to_ascii_lowercase() == t[ti]) {
358        si += 1;
359        ti += 1;
360    }
361    ti == t.len()
362}
363
364/// Returns `Some((value, endptr))` when the text at `p` parses as inf/nan,
365/// otherwise `None`.
366fn parse_inf_or_nan(bytes: &[u8], p: usize) -> Option<(f64, usize)> {
367    let mut s = p;
368    let mut negate = false;
369    if byte_at(bytes, s) == Some(b'-') {
370        negate = true;
371        s += 1;
372    } else if byte_at(bytes, s) == Some(b'+') {
373        s += 1;
374    }
375    if case_insensitive_match(bytes, s, b"inf") {
376        s += 3;
377        if case_insensitive_match(bytes, s, b"inity") {
378            s += 5;
379        }
380        let value = if negate {
381            f64::NEG_INFINITY
382        } else {
383            f64::INFINITY
384        };
385        Some((value, s))
386    } else if case_insensitive_match(bytes, s, b"nan") {
387        s += 3;
388        let value = if negate {
389            f64::from_bits(0xfff8_0000_0000_0000)
390        } else {
391            f64::from_bits(0x7ff8_0000_0000_0000)
392        };
393        Some((value, s))
394    } else {
395        None
396    }
397}
398
399/// Correctly-rounded scalbn. Every call site scales an already-representable
400/// value, so the result is exact.
401const fn ldexp(x: f64, mut n: i32) -> f64 {
402    let x1p1023 = f64::from_bits(0x7fe0000000000000);
403    let x1p53 = f64::from_bits(0x4340000000000000);
404    let x1p_1022 = f64::from_bits(0x0010000000000000);
405    let mut y = x;
406    if n > 1023 {
407        y *= x1p1023;
408        n -= 1023;
409        if n > 1023 {
410            y *= x1p1023;
411            n -= 1023;
412            if n > 1023 {
413                n = 1023;
414            }
415        }
416    } else if n < -1022 {
417        y *= x1p_1022 * x1p53;
418        n += 1022 - 53;
419        if n < -1022 {
420            y *= x1p_1022 * x1p53;
421            n += 1022 - 53;
422            if n < -1022 {
423                n = -1022;
424            }
425        }
426    }
427    y * f64::from_bits(((0x3ff + n) as u64) << 52)
428}
429
430/// Parse the already-validated `[+-]?[0-9]+` slice `bytes[start..end]` as base-10
431/// signed, saturating to i64::MIN/MAX on overflow like strtol.
432fn strtol_saturating(bytes: &[u8], start: usize, end: usize) -> i64 {
433    let mut i = start;
434    let mut neg = false;
435    if i < end && (bytes[i] == b'+' || bytes[i] == b'-') {
436        neg = bytes[i] == b'-';
437        i += 1;
438    }
439    let mut val: i64 = 0;
440    let mut overflowed = false;
441    while i < end {
442        let d = (bytes[i] - b'0') as i64;
443        match val.checked_mul(10).and_then(|v| v.checked_add(d)) {
444            Some(v) => val = v,
445            None => {
446                overflowed = true;
447                break;
448            }
449        }
450        i += 1;
451    }
452    if overflowed {
453        if neg { i64::MIN } else { i64::MAX }
454    } else if neg {
455        -val
456    } else {
457        val
458    }
459}
460
461/// Parse a hexadecimal floating-point string (the `float.fromhex` grammar).
462///
463/// The raw string is consumed as-is: leading and trailing whitespace are handled
464/// internally using the ASCII space set, so callers must not trim first.
465pub fn from_hex(s: &str) -> Result<f64, HexFloatError> {
466    let bytes = s.as_bytes();
467    let s_end = bytes.len();
468
469    let mut negate = false;
470    let mut idx = 0usize;
471    let mut x;
472
473    // leading whitespace
474    while byte_at(bytes, idx).is_some_and(rustpython_wtf8::is_py_ascii_whitespace) {
475        idx += 1;
476    }
477
478    // infinities and nans
479    if let Some((value, end)) = parse_inf_or_nan(bytes, idx) {
480        idx = end;
481        return finish_hex(bytes, s_end, idx, negate, value);
482    }
483
484    // optional sign
485    if byte_at(bytes, idx) == Some(b'-') {
486        idx += 1;
487        negate = true;
488    } else if byte_at(bytes, idx) == Some(b'+') {
489        idx += 1;
490    }
491
492    // [0x]
493    let s_store = idx;
494    if byte_at(bytes, idx) == Some(b'0') {
495        idx += 1;
496        if matches!(byte_at(bytes, idx), Some(b'x' | b'X')) {
497            idx += 1;
498        } else {
499            idx = s_store;
500        }
501    }
502
503    // coefficient: <integer> [. <fraction>]
504    let coeff_start = idx;
505    while hex_digit_at(bytes, idx).is_some() {
506        idx += 1;
507    }
508    let s_store = idx;
509    let coeff_end = if byte_at(bytes, idx) == Some(b'.') {
510        idx += 1;
511        while hex_digit_at(bytes, idx).is_some() {
512            idx += 1;
513        }
514        idx - 1
515    } else {
516        idx
517    };
518
519    // ndigits = total # of hex digits; fdigits = # after point
520    let ndigits_total = (coeff_end - coeff_start) as i64;
521    let fdigits = (coeff_end - s_store) as i64;
522    if ndigits_total == 0 {
523        return Err(HexFloatError::Invalid);
524    }
525    let insane_bound = core::cmp::min(
526        DBL_MIN_EXP - DBL_MANT_DIG - i64::MIN / 2,
527        i64::MAX / 2 + 1 - DBL_MAX_EXP,
528    ) / 4;
529    if ndigits_total > insane_bound {
530        return Err(HexFloatError::TooLong);
531    }
532
533    // [p <exponent>]
534    let exp = if matches!(byte_at(bytes, idx), Some(b'p' | b'P')) {
535        idx += 1;
536        let exp_start = idx;
537        if matches!(byte_at(bytes, idx), Some(b'-' | b'+')) {
538            idx += 1;
539        }
540        if !matches!(byte_at(bytes, idx), Some(b'0'..=b'9')) {
541            return Err(HexFloatError::Invalid);
542        }
543        idx += 1;
544        while matches!(byte_at(bytes, idx), Some(b'0'..=b'9')) {
545            idx += 1;
546        }
547        strtol_saturating(bytes, exp_start, idx)
548    } else {
549        0
550    };
551
552    // HEX_DIGIT(j): jth hex digit counting from the least significant.
553    let hex_digit = |j: i64| -> i32 {
554        let byte_idx = if j < fdigits {
555            coeff_end as i64 - j
556        } else {
557            coeff_end as i64 - 1 - j
558        };
559        hex_digit_at(bytes, byte_idx as usize).expect("hex digit within coefficient") as i32
560    };
561
562    // Discard leading zeros, and catch extreme overflow and underflow.
563    let mut ndigits = ndigits_total;
564    while ndigits > 0 && hex_digit(ndigits - 1) == 0 {
565        ndigits -= 1;
566    }
567    if ndigits == 0 || exp < i64::MIN / 2 {
568        x = 0.0;
569        return finish_hex(bytes, s_end, idx, negate, x);
570    }
571    if exp > i64::MAX / 2 {
572        return Err(HexFloatError::Overflow);
573    }
574
575    // Adjust exponent for fractional part.
576    let exp = exp - 4 * fdigits;
577
578    // top_exp = 1 more than exponent of most significant bit of coefficient.
579    let mut top_exp = exp + 4 * (ndigits - 1);
580    let mut digit = hex_digit(ndigits - 1);
581    while digit != 0 {
582        top_exp += 1;
583        digit /= 2;
584    }
585
586    // catch almost all nonextreme cases of overflow and underflow here
587    if top_exp < DBL_MIN_EXP - DBL_MANT_DIG {
588        x = 0.0;
589        return finish_hex(bytes, s_end, idx, negate, x);
590    }
591    if top_exp > DBL_MAX_EXP {
592        return Err(HexFloatError::Overflow);
593    }
594
595    // lsb = exponent of least significant bit of the rounded value.
596    let lsb = core::cmp::max(top_exp, DBL_MIN_EXP) - DBL_MANT_DIG;
597
598    x = 0.0;
599    if exp >= lsb {
600        // no rounding required
601        let mut i = ndigits - 1;
602        while i >= 0 {
603            x = 16.0 * x + hex_digit(i) as f64;
604            i -= 1;
605        }
606        x = ldexp(x, exp as i32);
607        return finish_hex(bytes, s_end, idx, negate, x);
608    }
609
610    // rounding required. key_digit is the index of the hex digit
611    // containing the first bit to be rounded away.
612    let half_eps: i32 = 1 << ((lsb - exp - 1) % 4) as i32;
613    let key_digit = (lsb - exp - 1) / 4;
614    let mut i = ndigits - 1;
615    while i > key_digit {
616        x = 16.0 * x + hex_digit(i) as f64;
617        i -= 1;
618    }
619    let digit = hex_digit(key_digit);
620    x = 16.0 * x + (digit & (16 - 2 * half_eps)) as f64;
621
622    // round-half-even
623    if (digit & half_eps) != 0 {
624        let round_up = if (digit & (3 * half_eps - 1)) != 0
625            || (half_eps == 8 && key_digit + 1 < ndigits && (hex_digit(key_digit + 1) & 1) != 0)
626        {
627            true
628        } else {
629            let mut r = false;
630            let mut i = key_digit - 1;
631            while i >= 0 {
632                if hex_digit(i) != 0 {
633                    r = true;
634                    break;
635                }
636                i -= 1;
637            }
638            r
639        };
640        if round_up {
641            x += (2 * half_eps) as f64;
642            if top_exp == DBL_MAX_EXP && x == ldexp((2 * half_eps) as f64, DBL_MANT_DIG as i32) {
643                // overflow corner case
644                return Err(HexFloatError::Overflow);
645            }
646        }
647    }
648    x = ldexp(x, (exp + 4 * key_digit) as i32);
649
650    finish_hex(bytes, s_end, idx, negate, x)
651}
652
653/// Skip trailing whitespace, require the whole string was consumed, and apply
654/// the sign.
655fn finish_hex(
656    bytes: &[u8],
657    s_end: usize,
658    mut idx: usize,
659    negate: bool,
660    x: f64,
661) -> Result<f64, HexFloatError> {
662    while byte_at(bytes, idx).is_some_and(rustpython_wtf8::is_py_ascii_whitespace) {
663        idx += 1;
664    }
665    if idx != s_end {
666        return Err(HexFloatError::Invalid);
667    }
668    Ok(if negate { -x } else { x })
669}
670
671#[cfg(test)]
672mod from_hex_tests {
673    use super::{HexFloatError, from_hex};
674
675    fn bits(s: &str) -> u64 {
676        from_hex(s).unwrap().to_bits()
677    }
678
679    #[test]
680    fn from_hex_exact_bits() {
681        assert_eq!(bits("0x1p-1074"), 0x0000000000000001);
682        assert_eq!(bits("0x1.fffffffffffffp+1023"), 0x7fefffffffffffff);
683        // round-half-even ties
684        assert_eq!(bits("0x1.00000000000008p0"), 0x3ff0000000000000);
685        assert_eq!(bits("0x1.00000000000018p0"), 0x3ff0000000000002);
686        assert_eq!(bits("-0x1p0"), 0xbff0000000000000);
687        assert_eq!(bits("0x0p0"), 0x0000000000000000);
688        assert_eq!(bits("-0x0p0"), 0x8000000000000000);
689    }
690
691    #[test]
692    fn from_hex_inf_nan() {
693        assert_eq!(bits("inf"), 0x7ff0000000000000);
694        assert_eq!(bits("-inf"), 0xfff0000000000000);
695        assert_eq!(bits("Infinity"), 0x7ff0000000000000);
696
697        let n = from_hex("nan").unwrap();
698        assert!(n.is_nan());
699        assert_eq!(n.to_bits(), 0x7ff8000000000000);
700        let neg = from_hex("-nan").unwrap();
701        assert!(neg.is_nan());
702        assert_eq!(neg.to_bits(), 0xfff8000000000000);
703    }
704
705    #[test]
706    fn from_hex_whitespace() {
707        assert_eq!(bits("  0x1p0  "), 0x3ff0000000000000);
708        assert_eq!(bits("\t0x1p0\n"), 0x3ff0000000000000);
709    }
710
711    #[test]
712    fn from_hex_errors() {
713        assert_eq!(from_hex("0x1p1024"), Err(HexFloatError::Overflow));
714        assert_eq!(from_hex("0x1z"), Err(HexFloatError::Invalid));
715        assert_eq!(from_hex(""), Err(HexFloatError::Invalid));
716        assert_eq!(from_hex("0x1 p0"), Err(HexFloatError::Invalid));
717    }
718}
719
720#[cfg(test)]
721mod tests {
722    use super::*;
723    use crate::hash::hash_float;
724
725    /// Exact `2**e` for `e` in `[-1074, 1023]`, built from bits so extreme
726    /// exponents don't overflow through an intermediate `2**|e|`.
727    fn pow2(e: i32) -> f64 {
728        if e >= -1022 {
729            f64::from_bits(((e + 1023) as u64) << 52)
730        } else {
731            f64::from_bits(1u64 << (e + 1074))
732        }
733    }
734
735    /// `decompose_float` is a frexp returning the *magnitude* mantissa: for a
736    /// nonzero `value`, `m` lies in `[0.5, 1)` and `m * 2**e == value.abs()`,
737    /// including for subnormals which have no implicit leading mantissa bit.
738    /// (Its sole caller reintroduces the sign via `value.signum()`.)
739    #[test]
740    fn decompose_float_frexp_contract() {
741        let mut values = alloc::vec![
742            0.0,
743            f64::from_bits(1), // smallest subnormal
744            f64::from_bits(2),
745            f64::from_bits(0x000f_ffff_ffff_ffff), // largest subnormal
746            f64::MIN_POSITIVE,                     // DBL_MIN, smallest normal
747            f64::from_bits(f64::MIN_POSITIVE.to_bits() - 1), // predecessor
748            1.0,
749            1.5,
750            0.1,
751            core::f64::consts::PI,
752        ];
753        for e in -1074..=1023 {
754            values.push(pow2(e));
755            values.push(-pow2(e));
756        }
757        for &v in &values {
758            let (m, e) = decompose_float(v);
759            if v == 0.0 {
760                assert_eq!((m, e), (0.0, 0));
761                continue;
762            }
763            assert!(
764                (0.5..1.0).contains(&m),
765                "mantissa {m} out of [0.5, 1) for value {v:e}"
766            );
767            // Reconstruct: m * 2**e must round-trip to the magnitude. Fold one
768            // power of two into the mantissa so `e` stays within `pow2`'s range
769            // (frexp yields e up to 1024 for 2**1023).
770            let reconstructed = (m * 2.0) * pow2(e - 1);
771            assert_eq!(
772                reconstructed.to_bits(),
773                v.abs().to_bits(),
774                "reconstruction failed for {v:e}: m={m}, e={e}"
775            );
776        }
777    }
778
779    /// Subnormal frexp regression: hash of the smallest positive subnormal.
780    #[test]
781    fn hash_float_smallest_subnormal() {
782        // hash(5e-324) == 16777216 (CPython 3.14 ground truth). The pre-fix
783        // bit-twiddling frexp returned 8404992 here.
784        assert_eq!(hash_float(f64::from_bits(1)), Some(16777216));
785    }
786
787    /// Differential float-hash table captured from CPython 3.14.5, spanning
788    /// subnormal boundaries, powers of two across the whole exponent range, and
789    /// a spread of normals.
790    #[test]
791    fn hash_float_matches_cpython() {
792        const HASH_CASES: &[(u64, i64)] = &[
793            (0x0000000000000001, 16777216),            // smallest subnormal 5e-324
794            (0x0000000000000002, 33554432),            // subnormal
795            (0x00000000deadbeef, 62678480394911744),   // subnormal midrange
796            (0x0008000000000000, 16384),               // subnormal high bit
797            (0x000fffffffffffff, 2305843009196949503), // largest subnormal
798            (0x0010000000000000, 32768),               // DBL_MIN smallest normal
799            (0x8000000000000001, -16777216),           // negative smallest subnormal
800            (0x0020000000000000, 65536),               // 2**-1021
801            (0x0170000000000000, 137438953472),        // 2**-1000
802            (0x39b0000000000000, 4194304),             // 2**-100
803            (0x3f50000000000000, 2251799813685248),    // 2**-10
804            (0x3fe0000000000000, 1152921504606846976), // 2**-1
805            (0x3ff0000000000000, 1),                   // 2**0
806            (0x4000000000000000, 2),                   // 2**1
807            (0x4090000000000000, 1024),                // 2**10
808            (0x4630000000000000, 549755813888),        // 2**100
809            (0x7e70000000000000, 16777216),            // 2**1000
810            (0x7fe0000000000000, 140737488355328),     // 2**1023
811            (0xffe0000000000000, -140737488355328),    // -2**1023
812            (0x3ff8000000000000, 1152921504606846977), // 1.5
813            (0x400921fb54442d18, 326490430436040707),  // 3.141592653589793
814            (0x7e37e43c8800759c, 1224995262755759164), // 1e+300
815            (0x01a56e1fc2f8f359, 482449582752280463),  // 1e-300
816            (0x40c81cd6c8b43958, 1563361560246628409), // 12345.678
817            (0x3fb999999999999a, 230584300921369408),  // 0.1
818            (0x4005666666666666, 1556444031219243010), // 2.675
819            (0x4132d68700000000, 1234567),             // 1234567.0
820            (0x44dfe154f457ea13, 1428027733287631914), // 6.022e+23
821            (0x3c07a42f549647fb, 851769299698974080),  // 1.602e-19
822            (0xbff0000000000000, -2),                  // -1.0
823            (0xbfb999999999999a, -230584300921369408), // -0.1
824        ];
825        for &(bits, expected) in HASH_CASES {
826            let v = f64::from_bits(bits);
827            assert_eq!(
828                hash_float(v),
829                Some(expected),
830                "hash mismatch for {v:e} (bits {bits:#018x})"
831            );
832        }
833    }
834}