Skip to main content

rustpython_literal/
float.rs

1use crate::format::Case;
2use alloc::borrow::ToOwned;
3use alloc::format;
4use alloc::string::{String, ToString};
5use num_traits::Zero;
6
7#[must_use]
8pub fn parse_str(literal: &str) -> Option<f64> {
9    parse_inner(literal.trim().as_bytes())
10}
11
12#[must_use]
13pub fn parse_bytes(literal: &[u8]) -> Option<f64> {
14    parse_inner(literal.trim_ascii())
15}
16
17fn parse_inner(literal: &[u8]) -> Option<f64> {
18    use lexical_parse_float::{
19        FromLexicalWithOptions, NumberFormatBuilder, Options, format::PYTHON3_LITERAL,
20    };
21
22    // lexical-core's format::PYTHON_STRING is inaccurate
23    const PYTHON_STRING: u128 = NumberFormatBuilder::rebuild(PYTHON3_LITERAL)
24        .no_special(false)
25        .build_unchecked();
26    f64::from_lexical_with_options::<PYTHON_STRING>(literal, &Options::new()).ok()
27}
28
29#[must_use]
30pub fn is_integer(v: f64) -> bool {
31    v.is_finite() && v.fract() == 0.0
32}
33
34fn format_nan(case: Case) -> String {
35    let nan = match case {
36        Case::Lower => "nan",
37        Case::Upper => "NAN",
38    };
39
40    nan.to_string()
41}
42
43fn format_inf(case: Case) -> String {
44    let inf = match case {
45        Case::Lower => "inf",
46        Case::Upper => "INF",
47    };
48
49    inf.to_string()
50}
51
52#[must_use]
53pub const fn decimal_point_or_empty(precision: usize, alternate_form: bool) -> &'static str {
54    match (precision, alternate_form) {
55        (0, true) => ".",
56        _ => "",
57    }
58}
59
60/// Rust's `format!("{:.*}", n, x)` panics when `n` exceeds the fmt runtime's
61/// internal precision limit. User-supplied precision can legally reach far
62/// higher values (e.g. `f"{1.5:.1000000}"`) — clamp here so we produce a
63/// (truncated-but-valid) output instead of aborting the interpreter. Harmless
64/// in practice: f64 carries only ~17 significant digits, so precision beyond
65/// 65K is padding zeros at best.
66///
67/// The two caps differ by 1: `{:.*}` (plain) accepts `u16::MAX`, but `{:.*e}`
68/// (exponential) hits a tighter assertion (`ndigits > 0` in
69/// `core::num::flt2dec`) at exactly `u16::MAX`. Keeping plain at the higher
70/// cap preserves byte-identical output with CPython up through
71/// `precision == u16::MAX` for fixed / percent / general-non-scientific paths.
72pub const FMT_MAX_PRECISION: usize = u16::MAX as usize;
73pub const FMT_MAX_EXP_PRECISION: usize = u16::MAX as usize - 1;
74
75#[inline]
76#[must_use]
77pub fn clamp_fmt_precision(precision: usize) -> usize {
78    core::cmp::min(precision, FMT_MAX_PRECISION)
79}
80
81#[inline]
82#[must_use]
83pub fn clamp_exp_precision(precision: usize) -> usize {
84    core::cmp::min(precision, FMT_MAX_EXP_PRECISION)
85}
86
87#[must_use]
88pub fn format_fixed(precision: usize, magnitude: f64, case: Case, alternate_form: bool) -> String {
89    match magnitude {
90        magnitude if magnitude.is_finite() => {
91            let point = decimal_point_or_empty(precision, alternate_form);
92            let capped = clamp_fmt_precision(precision);
93            let mut out = format!("{magnitude:.capped$}");
94            // Pad with '0's up to the requested precision to match CPython
95            // byte-identically. `f64` has at most ~767 significant decimal
96            // digits, so any digit past `capped` is deterministically '0'.
97            let missing = precision.saturating_sub(capped);
98            if missing > 0 {
99                out.extend(core::iter::repeat_n('0', missing));
100            }
101            out.push_str(point);
102            out
103        }
104        magnitude if magnitude.is_nan() => format_nan(case),
105        magnitude if magnitude.is_infinite() => format_inf(case),
106        _ => "".to_string(),
107    }
108}
109
110// Formats floats into Python style exponent notation, by first formatting in Rust style
111// exponent notation (`1.0000e0`), then convert to Python style (`1.0000e+00`).
112#[must_use]
113pub fn format_exponent(
114    precision: usize,
115    magnitude: f64,
116    case: Case,
117    alternate_form: bool,
118) -> String {
119    match magnitude {
120        magnitude if magnitude.is_finite() => {
121            let capped = clamp_exp_precision(precision);
122            let r_exp = format!("{magnitude:.capped$e}");
123            let mut parts = r_exp.splitn(2, 'e');
124            let base = parts.next().unwrap();
125            let exponent = parts.next().unwrap().parse::<i64>().unwrap();
126            let e = match case {
127                Case::Lower => 'e',
128                Case::Upper => 'E',
129            };
130            let point = decimal_point_or_empty(precision, alternate_form);
131            // Pad with '0's up to the requested precision to match CPython
132            // byte-identically past our internal cap; see `format_fixed`.
133            let missing = precision.saturating_sub(capped);
134            let mut mantissa = String::with_capacity(base.len() + missing);
135            mantissa.push_str(base);
136            if missing > 0 {
137                mantissa.extend(core::iter::repeat_n('0', missing));
138            }
139            format!("{mantissa}{point}{e}{exponent:+#03}")
140        }
141        magnitude if magnitude.is_nan() => format_nan(case),
142        magnitude if magnitude.is_infinite() => format_inf(case),
143        _ => "".to_string(),
144    }
145}
146
147/// If s represents a floating point value, trailing zeros and a possibly trailing
148/// decimal point will be removed.
149/// This function does NOT work with decimal commas.
150fn maybe_remove_trailing_redundant_chars(s: String, alternate_form: bool) -> String {
151    if !alternate_form && s.contains('.') {
152        // only truncate floating point values when not in alternate form
153        let s = remove_trailing_zeros(s);
154        remove_trailing_decimal_point(s)
155    } else {
156        s
157    }
158}
159
160fn remove_trailing_zeros(s: String) -> String {
161    let mut s = s;
162    while s.ends_with('0') {
163        s.pop();
164    }
165    s
166}
167
168fn remove_trailing_decimal_point(s: String) -> String {
169    let mut s = s;
170    if s.ends_with('.') {
171        s.pop();
172    }
173    s
174}
175
176#[must_use]
177pub fn format_general(
178    precision: usize,
179    magnitude: f64,
180    case: Case,
181    alternate_form: bool,
182    always_shows_fract: bool,
183) -> String {
184    match magnitude {
185        magnitude if magnitude.is_finite() => {
186            let exp_precision = clamp_exp_precision(precision.saturating_sub(1));
187            let r_exp = format!("{magnitude:.exp_precision$e}");
188            let mut parts = r_exp.splitn(2, 'e');
189            let base = parts.next().unwrap();
190            let exponent = parts.next().unwrap().parse::<i64>().unwrap();
191            if exponent < -4 || exponent + (always_shows_fract as i64) >= (precision as i64) {
192                let e = match case {
193                    Case::Lower => 'e',
194                    Case::Upper => 'E',
195                };
196                // `base` is already produced at the clamped precision via
197                // `r_exp`. The previous `format!("{:.*}", precision + 1, base)`
198                // call was a no-op (magnitude is `.abs()`-ed at the caller, so
199                // base has no sign and its length was exactly `precision + 1`)
200                // — reuse `base` directly to avoid double-clamping that would
201                // drop the last 1-2 chars at high precision.
202                let base = maybe_remove_trailing_redundant_chars(base.to_owned(), alternate_form);
203                let point = decimal_point_or_empty(exp_precision, alternate_form);
204                format!("{base}{point}{e}{exponent:+#03}")
205            } else {
206                let precision =
207                    clamp_fmt_precision(((precision as i64) - 1 - exponent).max(0) as usize);
208                let magnitude = format!("{magnitude:.precision$}");
209                let base = maybe_remove_trailing_redundant_chars(magnitude, alternate_form);
210                let point = decimal_point_or_empty(precision, alternate_form);
211                format!("{base}{point}")
212            }
213        }
214        magnitude if magnitude.is_nan() => format_nan(case),
215        magnitude if magnitude.is_infinite() => format_inf(case),
216        _ => "".to_string(),
217    }
218}
219
220pub(crate) fn prefer_cpython_tie_repr(s: String, value: f64) -> String {
221    // Rust's shortest float formatter can land on the odd-digit neighbour of a
222    // rounding tie where round-half-to-even (what `repr` uses) picks the even
223    // one. When the last significant digit is odd and its even neighbour still
224    // round-trips and is no further from the value, prefer the even neighbour.
225    let boundary = s.find('e').unwrap_or(s.len());
226    let Some(digit_pos) = s[..boundary].bytes().rposition(|b| b.is_ascii_digit()) else {
227        return s;
228    };
229
230    let digit = s.as_bytes()[digit_pos];
231    if digit == b'0' {
232        return s;
233    }
234    let decremented = digit - 1;
235    if !(decremented - b'0').is_multiple_of(2) {
236        return s;
237    }
238
239    let mut candidate = s.clone();
240    candidate.replace_range(
241        digit_pos..=digit_pos,
242        core::str::from_utf8(&[decremented]).unwrap(),
243    );
244    if parse_str(&candidate).is_none_or(|parsed| parsed.to_bits() != value.to_bits()) {
245        return s;
246    }
247
248    let Some(current_distance) = decimal_distance_to_f64(&s, value) else {
249        return s;
250    };
251    let Some(candidate_distance) = decimal_distance_to_f64(&candidate, value) else {
252        return s;
253    };
254
255    if candidate_distance <= current_distance {
256        candidate
257    } else {
258        s
259    }
260}
261
262fn checked_pow_u128(base: u128, exp: u32) -> Option<u128> {
263    let mut result = 1u128;
264    for _ in 0..exp {
265        result = result.checked_mul(base)?;
266    }
267    Some(result)
268}
269
270fn parse_decimal_rational(s: &str) -> Option<(u128, u32)> {
271    let (mantissa, exponent) = match s.find('e') {
272        Some(pos) => (&s[..pos], s[pos + 1..].parse::<i32>().ok()?),
273        None => (s, 0),
274    };
275    let significand = mantissa.strip_prefix('-').unwrap_or(mantissa);
276    let dot_pos = significand.find('.');
277    let frac_digits = dot_pos.map_or(0, |pos| significand.len().saturating_sub(pos + 1));
278    let mut digits = String::with_capacity(significand.len());
279    for ch in significand.chars() {
280        if ch != '.' {
281            digits.push(ch);
282        }
283    }
284    let mut int = digits.parse::<u128>().ok()?;
285    let mut scale = i32::try_from(frac_digits).ok()? - exponent;
286    if scale < 0 {
287        int = int.checked_mul(checked_pow_u128(10, (-scale) as u32)?)?;
288        scale = 0;
289    }
290    Some((int, scale as u32))
291}
292
293fn f64_mantissa_exponent(value: f64) -> Option<(u128, i32)> {
294    let bits = value.abs().to_bits();
295    let exponent = ((bits >> 52) & 0x7ff) as i32;
296    let fraction = bits & ((1u64 << 52) - 1);
297    if exponent == 0 {
298        Some((u128::from(fraction), 1 - 1023 - 52))
299    } else if exponent < 0x7ff {
300        Some((u128::from((1u64 << 52) | fraction), exponent - 1023 - 52))
301    } else {
302        None
303    }
304}
305
306fn decimal_distance_to_f64(s: &str, value: f64) -> Option<u128> {
307    let (decimal_int, decimal_scale) = parse_decimal_rational(s)?;
308    let (mantissa, binary_exponent) = f64_mantissa_exponent(value)?;
309    if binary_exponent >= 0 || decimal_scale > 38 {
310        return None;
311    }
312
313    let binary_scale = u32::try_from(-binary_exponent).ok()?;
314    let common_twos = decimal_scale.max(binary_scale);
315    let decimal_scaled =
316        decimal_int.checked_mul(checked_pow_u128(2, common_twos - decimal_scale)?)?;
317    let five_power = checked_pow_u128(5, decimal_scale)?;
318    let binary_scaled = mantissa
319        .checked_mul(checked_pow_u128(2, common_twos - binary_scale)?)?
320        .checked_mul(five_power)?;
321
322    Some(decimal_scaled.abs_diff(binary_scaled))
323}
324
325// TODO: rewrite using format_general
326#[must_use]
327pub fn to_string(value: f64) -> String {
328    let lit = format!("{value:e}");
329    if let Some(position) = lit.find('e') {
330        let significand = &lit[..position];
331        let exponent = &lit[position + 1..];
332        let exponent = exponent.parse::<i32>().unwrap();
333        if exponent < 16 && exponent > -5 {
334            if is_integer(value) {
335                format!("{value:.1?}")
336            } else {
337                prefer_cpython_tie_repr(value.to_string(), value)
338            }
339        } else {
340            prefer_cpython_tie_repr(format!("{significand}e{exponent:+#03}"), value)
341        }
342    } else {
343        let mut s = value.to_string();
344        s.make_ascii_lowercase();
345        s
346    }
347}
348
349#[must_use]
350pub fn from_hex(s: &str) -> Option<f64> {
351    if let Ok(f) = hexf_parse::parse_hexf64(s, false) {
352        return Some(f);
353    }
354    match s.to_ascii_lowercase().as_str() {
355        "nan" | "+nan" | "-nan" => Some(f64::NAN),
356        "inf" | "infinity" | "+inf" | "+infinity" => Some(f64::INFINITY),
357        "-inf" | "-infinity" => Some(f64::NEG_INFINITY),
358        value => {
359            let mut hex = String::with_capacity(value.len());
360            let has_0x = value.contains("0x");
361            let has_p = value.contains('p');
362            let has_dot = value.contains('.');
363            let mut start = 0;
364
365            if !has_0x && value.starts_with('-') {
366                hex.push_str("-0x");
367                start += 1;
368            } else if !has_0x {
369                hex.push_str("0x");
370                if value.starts_with('+') {
371                    start += 1;
372                }
373            }
374
375            for (index, ch) in value.chars().enumerate() {
376                if ch == 'p' {
377                    if has_dot {
378                        hex.push('p');
379                    } else {
380                        hex.push_str(".p");
381                    }
382                } else if index >= start {
383                    hex.push(ch);
384                }
385            }
386
387            if !has_p && has_dot {
388                hex.push_str("p0");
389            } else if !has_p && !has_dot {
390                hex.push_str(".p0")
391            }
392
393            hexf_parse::parse_hexf64(hex.as_str(), false).ok()
394        }
395    }
396}
397
398#[must_use]
399pub fn to_hex(value: f64) -> String {
400    let bits = value.to_bits();
401    let sign_fmt = if bits >> 63 != 0 { "-" } else { "" };
402    match value {
403        value if value.is_zero() => format!("{sign_fmt}0x0.0p+0"),
404        value if value.is_infinite() => format!("{sign_fmt}inf"),
405        value if value.is_nan() => "nan".to_owned(),
406        _ => {
407            const FRACT_MASK: u64 = (1u64 << 52) - 1;
408            const EXP_MASK: u64 = 0x7ff;
409            let exponent = (bits >> 52) & EXP_MASK;
410            let fraction = bits & FRACT_MASK;
411            if exponent == 0 {
412                format!("{sign_fmt}0x0.{fraction:013x}p-1022")
413            } else {
414                let exponent = i32::try_from(exponent).unwrap() - 1023;
415                format!("{sign_fmt}0x1.{fraction:013x}p{exponent:+}")
416            }
417        }
418    }
419}
420
421#[cfg(test)]
422mod tests {
423    use super::*;
424
425    #[test]
426    fn repr_uses_cpython_tie_digit_for_power_of_two() {
427        assert_eq!(to_string(2.0f64.powi(-25)), "2.9802322387695312e-08");
428        assert_eq!(to_string((-2.0f64).powi(-25)), "-2.9802322387695312e-08");
429        assert_eq!(to_string(2.0f64.powi(-26)), "1.4901161193847656e-08");
430        assert_eq!(
431            to_string(2.0f64.powi(-14) - 2.0f64.powi(-25)),
432            "6.1005353927612305e-05"
433        );
434    }
435
436    #[test]
437    fn repr_normal_range_uses_cpython_tie_digit() {
438        // Rust's shortest formatter yields "161852602146008.13" for this
439        // value; round-half-to-even (what `repr` uses) picks "…08.12".
440        assert_eq!(
441            to_string(f64::from_bits(0x42e26687db6b9b04)),
442            "161852602146008.12"
443        );
444        // Non-tie values are left untouched.
445        assert_eq!(to_string(1.5), "1.5");
446        assert_eq!(to_string(0.1), "0.1");
447        assert_eq!(to_string(12.34), "12.34");
448        assert_eq!(to_string(100.0), "100.0");
449    }
450
451    #[test]
452    fn to_hex_works() {
453        use rand::RngExt;
454        assert_eq!(to_hex(f64::from_bits(1)), "0x0.0000000000001p-1022");
455        assert_eq!(to_hex(f64::from_bits(2)), "0x0.0000000000002p-1022");
456        assert_eq!(to_hex(-f64::from_bits(1)), "-0x0.0000000000001p-1022");
457        assert_eq!(to_hex(f64::MIN_POSITIVE), "0x1.0000000000000p-1022");
458        for _ in 0..20000 {
459            let bytes = rand::rng().random::<u64>();
460            let f = f64::from_bits(bytes);
461            if !f.is_finite() {
462                continue;
463            }
464            let hex = to_hex(f);
465            // println!("{} -> {}", f, hex);
466            let roundtrip = hexf_parse::parse_hexf64(&hex, false).unwrap();
467            // println!("  -> {}", roundtrip);
468            assert!(f == roundtrip, "{f} {hex} {roundtrip}");
469        }
470    }
471
472    #[test]
473    fn remove_trailing_zeros_works() {
474        assert!(remove_trailing_zeros(String::from("100")) == *"1");
475        assert!(remove_trailing_zeros(String::from("100.00")) == *"100.");
476
477        // leave leading zeros untouched
478        assert!(remove_trailing_zeros(String::from("001")) == *"001");
479
480        // leave strings untouched if they don't end with 0
481        assert!(remove_trailing_zeros(String::from("101")) == *"101");
482    }
483
484    #[test]
485    fn remove_trailing_decimal_point_works() {
486        assert!(remove_trailing_decimal_point(String::from("100.")) == *"100");
487        assert!(remove_trailing_decimal_point(String::from("1.")) == *"1");
488
489        // leave leading decimal points untouched
490        assert!(remove_trailing_decimal_point(String::from(".5")) == *".5");
491    }
492
493    #[test]
494    fn maybe_remove_trailing_redundant_chars_works() {
495        assert!(maybe_remove_trailing_redundant_chars(String::from("100."), true) == *"100.");
496        assert!(maybe_remove_trailing_redundant_chars(String::from("100."), false) == *"100");
497        assert!(maybe_remove_trailing_redundant_chars(String::from("1."), false) == *"1");
498        assert!(maybe_remove_trailing_redundant_chars(String::from("10.0"), false) == *"10");
499
500        // don't truncate integers
501        assert!(maybe_remove_trailing_redundant_chars(String::from("1000"), false) == *"1000");
502    }
503}