Skip to main content

smtlib/theories/
floating_point.rs

1#![doc = concat!("```ignore\n", include_str!("./FloatingPoint.smt2"), "```")]
2
3use smtlib_lowlevel::{
4    ast::{self, Identifier, Index, QualIdentifier, Term},
5    lexicon::{self, Numeral},
6    Storage,
7};
8
9use crate::{
10    sorts::Sort,
11    terms::{
12        app, qual_ident, ApplicationArgs, Const, Dynamic, IntoWithStorage, STerm, Sorted,
13        StaticSorted,
14    },
15    theories::fixed_size_bit_vectors::BitVec,
16    Bool, Real,
17};
18
19/// The SMT-LIB sort for rounding modes in floating-point operations.
20#[derive(Debug, Clone, Copy)]
21pub struct RoundingMode<'st>(STerm<'st>);
22
23impl<'st> From<Const<'st, RoundingMode<'st>>> for RoundingMode<'st> {
24    fn from(c: Const<'st, RoundingMode<'st>>) -> Self {
25        c.1
26    }
27}
28
29impl<'st> std::fmt::Display for RoundingMode<'st> {
30    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31        self.term().fmt(f)
32    }
33}
34
35impl<'st> From<RoundingMode<'st>> for Dynamic<'st> {
36    fn from(i: RoundingMode<'st>) -> Self {
37        i.into_dynamic()
38    }
39}
40
41impl<'st> From<RoundingMode<'st>> for STerm<'st> {
42    fn from(i: RoundingMode<'st>) -> Self {
43        i.0
44    }
45}
46
47impl<'st> From<STerm<'st>> for RoundingMode<'st> {
48    fn from(t: STerm<'st>) -> Self {
49        RoundingMode(t)
50    }
51}
52
53impl<'st> From<(STerm<'st>, Sort<'st>)> for RoundingMode<'st> {
54    fn from((t, _s): (STerm<'st>, Sort<'st>)) -> Self {
55        // TODO: consider checking sort compatibility if _s is not None
56        t.into()
57    }
58}
59
60impl<'st> StaticSorted<'st> for RoundingMode<'st> {
61    type Inner = Self;
62    const AST_SORT: ast::Sort<'static> = ast::Sort::new_simple("RoundingMode");
63
64    fn static_st(&self) -> &'st Storage {
65        self.st()
66    }
67
68    fn sort() -> Sort<'st> {
69        Self::AST_SORT.into()
70    }
71
72    fn new_const(st: &'st Storage, name: &str) -> Const<'st, Self> {
73        let name = st.alloc_str(name);
74        let rm = Term::Identifier(qual_ident(name, Some(Self::AST_SORT)));
75        let rm = STerm::new(st, rm);
76        Const(name, rm.into())
77    }
78}
79
80impl<'st> RoundingMode<'st> {
81    fn new_mode_val(st: &'st Storage, name: &'static str) -> Self {
82        STerm::new(st, Term::Identifier(qual_ident(st.alloc_str(name), None))).into()
83    }
84
85    /// Round nearest ties to even
86    pub fn rne(st: &'st Storage) -> Self {
87        Self::new_mode_val(st, "RNE")
88    }
89    /// Round nearest ties to away
90    pub fn rna(st: &'st Storage) -> Self {
91        Self::new_mode_val(st, "RNA")
92    }
93    /// Round toward positive
94    pub fn rtp(st: &'st Storage) -> Self {
95        Self::new_mode_val(st, "RTP")
96    }
97    /// Round toward negative
98    pub fn rtn(st: &'st Storage) -> Self {
99        Self::new_mode_val(st, "RTN")
100    }
101    /// Round toward zero
102    pub fn rtz(st: &'st Storage) -> Self {
103        Self::new_mode_val(st, "RTZ")
104    }
105}
106
107/// A floating-point number, parameterized by exponent bits (EB) and significand
108/// bits (SB). Significand bits include the sign bit.
109#[derive(Debug, Clone, Copy)]
110pub struct FloatingPoint<'st, const EB: usize, const SB: usize>(STerm<'st>);
111
112/// Alias for (_ FloatingPoint 5 11) - IEEE binary16
113pub type Float16<'st> = FloatingPoint<'st, 5, 11>;
114/// Alias for (_ FloatingPoint 8 24) - IEEE binary32
115pub type Float32<'st> = FloatingPoint<'st, 8, 24>;
116/// Alias for (_ FloatingPoint 11 53) - IEEE binary64
117pub type Float64<'st> = FloatingPoint<'st, 11, 53>;
118/// Alias for (_ FloatingPoint 15 113) - IEEE binary128
119pub type Float128<'st> = FloatingPoint<'st, 15, 113>;
120
121impl<'st, const EB: usize, const SB: usize> From<Const<'st, FloatingPoint<'st, EB, SB>>>
122    for FloatingPoint<'st, EB, SB>
123{
124    fn from(c: Const<'st, FloatingPoint<'st, EB, SB>>) -> Self {
125        c.1
126    }
127}
128
129impl<const EB: usize, const SB: usize> std::fmt::Display for FloatingPoint<'_, EB, SB> {
130    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
131        self.term().fmt(f)
132    }
133}
134
135impl<'st, const EB: usize, const SB: usize> From<FloatingPoint<'st, EB, SB>> for Dynamic<'st> {
136    fn from(i: FloatingPoint<'st, EB, SB>) -> Self {
137        i.into_dynamic()
138    }
139}
140
141impl<'st, const EB: usize, const SB: usize> From<FloatingPoint<'st, EB, SB>> for STerm<'st> {
142    fn from(i: FloatingPoint<'st, EB, SB>) -> Self {
143        i.0
144    }
145}
146
147impl<'st, const EB: usize, const SB: usize> From<STerm<'st>> for FloatingPoint<'st, EB, SB> {
148    fn from(t: STerm<'st>) -> Self {
149        FloatingPoint(t)
150    }
151}
152
153impl<'st, const EB: usize, const SB: usize> From<(STerm<'st>, Sort<'st>)>
154    for FloatingPoint<'st, EB, SB>
155{
156    fn from((t, _): (STerm<'st>, Sort<'st>)) -> Self {
157        t.into()
158    }
159}
160
161impl<'st, const EB: usize, const SB: usize> StaticSorted<'st> for FloatingPoint<'st, EB, SB> {
162    type Inner = Self;
163    const AST_SORT: ast::Sort<'static> = ast::Sort::new_indexed(
164        "FloatingPoint",
165        &[
166            Index::Numeral(lexicon::Numeral::from_usize(EB)),
167            Index::Numeral(lexicon::Numeral::from_usize(SB)),
168        ],
169    );
170
171    fn static_st(&self) -> &'st Storage {
172        self.sterm().st()
173    }
174
175    fn sort() -> Sort<'st> {
176        Self::AST_SORT.into()
177    }
178
179    fn new_const(st: &'st Storage, name: &str) -> Const<'st, Self> {
180        let name = st.alloc_str(name);
181        let fp = Term::Identifier(qual_ident(name, Some(Self::AST_SORT)));
182        let fp = STerm::new(st, fp);
183        Const(name, fp.into())
184    }
185}
186
187trait FloatHelper: Sized {
188    const EXPONENT_BITS: u32;
189    const SIGNIFICAND_BITS: u32;
190    const SIGN_SHIFT: u32 = Self::EXPONENT_BITS + Self::SIGNIFICAND_BITS;
191    const EXPONENT_SHIFT: u32 = Self::SIGNIFICAND_BITS;
192    fn from_bits_u64(bits: u64) -> Self;
193    fn parse_special(symbol: &str) -> Self;
194
195    fn from_bits_parts(sign: i64, exponent: i64, significand: i64) -> Self {
196        let bits = ((sign as u64) << Self::SIGN_SHIFT)
197            | ((exponent as u64) << Self::EXPONENT_SHIFT)
198            | (significand as u64);
199        Self::from_bits_u64(bits)
200    }
201}
202
203impl FloatHelper for f32 {
204    const EXPONENT_BITS: u32 = 8;
205    const SIGNIFICAND_BITS: u32 = 23;
206
207    fn from_bits_u64(bits: u64) -> Self {
208        f32::from_bits(bits as u32)
209    }
210    fn parse_special(symbol: &str) -> Self {
211        match symbol {
212            "+zero" => 0.0,
213            "-zero" => -0.0,
214            "+oo" => f32::INFINITY,
215            "-oo" => f32::NEG_INFINITY,
216            "NaN" => f32::NAN,
217            _ => panic!("Unknown floating-point constant: {}", symbol),
218        }
219    }
220}
221
222impl FloatHelper for f64 {
223    const EXPONENT_BITS: u32 = 11;
224    const SIGNIFICAND_BITS: u32 = 52;
225
226    fn from_bits_u64(bits: u64) -> Self {
227        f64::from_bits(bits)
228    }
229    fn parse_special(symbol: &str) -> Self {
230        match symbol {
231            "+zero" => 0.0,
232            "-zero" => -0.0,
233            "+oo" => f64::INFINITY,
234            "-oo" => f64::NEG_INFINITY,
235            "NaN" => f64::NAN,
236            _ => panic!("Unknown floating-point constant: {}", symbol),
237        }
238    }
239}
240
241impl<'st> IntoWithStorage<'st, Float32<'st>> for f32 {
242    fn into_with_storage(self, st: &'st Storage) -> Float32<'st> {
243        Float32::to_fp_from_bits(st, BitVec::new_prim(st, self.to_bits()))
244    }
245}
246
247impl<'st> IntoWithStorage<'st, Float64<'st>> for f64 {
248    fn into_with_storage(self, st: &'st Storage) -> Float64<'st> {
249        Float64::to_fp_from_bits(st, BitVec::new_prim(st, self.to_bits()))
250    }
251}
252
253fn spec_constant_to_i64<'st>(value: &ast::SpecConstant<'st>) -> i64 {
254    match value {
255        ast::SpecConstant::Numeral(n) => n.into_u128().unwrap().try_into().unwrap(),
256        ast::SpecConstant::Hexadecimal(h) => h.parse().unwrap(),
257        ast::SpecConstant::Binary(b) => b.parse().unwrap(),
258        _ => panic!("Unsupported constant type for bit conversion: {:?}", value),
259    }
260}
261
262fn term_to_i64<'st>(value: &Term<'st>) -> i64 {
263    match value {
264        Term::SpecConstant(spec_constant) => spec_constant_to_i64(spec_constant),
265        _ => panic!("Expected spec constant, got: {:?}", value),
266    }
267}
268
269fn try_float_from_term<F: FloatHelper>(term: &Term) -> Result<F, std::num::ParseIntError> {
270    Ok(match term {
271        Term::Identifier(QualIdentifier::Identifier(Identifier::Indexed(symbol, _))) => {
272            F::parse_special(symbol.0)
273        }
274        Term::Application(QualIdentifier::Identifier(Identifier::Simple(symbol)), args) => {
275            assert_eq!(symbol.0, "fp");
276            let sign = term_to_i64(args[0]);
277            let exponent = term_to_i64(args[1]);
278            let significand = term_to_i64(args[2]);
279            F::from_bits_parts(sign, exponent, significand)
280        }
281        _ => panic!("Unexpected term: {:?}", term),
282    })
283}
284
285impl<'st> TryFrom<Float32<'st>> for f32 {
286    type Error = std::num::ParseIntError;
287
288    fn try_from(value: Float32<'st>) -> Result<Self, Self::Error> {
289        try_float_from_term(value.term())
290    }
291}
292
293impl<'st> TryFrom<Float64<'st>> for f64 {
294    type Error = std::num::ParseIntError;
295
296    fn try_from(value: Float64<'st>) -> Result<Self, Self::Error> {
297        try_float_from_term(value.term())
298    }
299}
300
301impl<'st, const EB: usize, const SB: usize> FloatingPoint<'st, EB, SB> {
302    /// Construct a new bit-vec.
303    pub fn new(
304        st: &'st Storage,
305        value: impl IntoWithStorage<'st, FloatingPoint<'st, EB, SB>>,
306    ) -> FloatingPoint<'st, EB, SB> {
307        value.into_with_storage(st)
308    }
309
310    fn st(&self) -> &'st Storage {
311        self.0.st()
312    }
313    fn term(&self) -> STerm<'st> {
314        self.0
315    }
316
317    fn app_fn_indexed_args<T: From<STerm<'st>>>(
318        st: &'st Storage,
319        op: &'st str,
320        indices: impl IntoIterator<Item = usize>,
321        args: impl ApplicationArgs<'st>,
322    ) -> T {
323        let index_nodes = indices
324            .into_iter()
325            .map(|i| Index::Numeral(Numeral::from_usize(i)))
326            .collect::<Vec<_>>();
327        let qual_id = QualIdentifier::Identifier(Identifier::indexed(
328            st.alloc_str(op),
329            st.alloc_slice(&index_nodes),
330        ));
331        STerm::new(st, Term::Application(qual_id, args.into_args(st))).into()
332    }
333
334    fn unop_rm_indexed<T: From<STerm<'st>>>(
335        self,
336        op: &'st str,
337        rm: RoundingMode<'st>,
338        index_val: usize,
339    ) -> T {
340        Self::app_fn_indexed_args(self.st(), op, [index_val], (rm.term(), self.term()))
341    }
342
343    fn op_const_indexed<T: From<STerm<'st>>>(
344        st: &'st Storage,
345        op: &'st str,
346        indices: impl IntoIterator<Item = usize>,
347    ) -> T {
348        let index_nodes = indices
349            .into_iter()
350            .map(|i| Index::Numeral(Numeral::from_usize(i)))
351            .collect::<Vec<_>>();
352        let qual_id = QualIdentifier::Identifier(Identifier::indexed(
353            st.alloc_str(op),
354            st.alloc_slice(&index_nodes),
355        ));
356        STerm::new(st, Term::Identifier(qual_id)).into()
357    }
358
359    /// Creates a floating-point value from sign, exponent, and significand
360    /// bit-vectors. `i = sb - 1`
361    ///
362    /// Note: The SB_1 parameter is a workaround for
363    /// `#![feature(generic_const_exprs)]`.
364    pub fn fp<const SB_1: usize>(
365        st: &'st Storage,
366        sign: BitVec<'st, 1>,
367        exponent: BitVec<'st, EB>,
368        significand: BitVec<'st, SB_1>,
369    ) -> Self {
370        assert_eq!(SB_1, SB - 1);
371        app(st, "fp", (sign.term(), exponent.term(), significand.term())).into()
372    }
373
374    /// Positive infinity
375    pub fn plus_oo(st: &'st Storage) -> Self {
376        Self::op_const_indexed(st, "+oo", [EB, SB])
377    }
378
379    /// Negative infinity
380    pub fn minus_oo(st: &'st Storage) -> Self {
381        Self::op_const_indexed(st, "-oo", [EB, SB])
382    }
383
384    /// Positive zero
385    pub fn plus_zero(st: &'st Storage) -> Self {
386        Self::op_const_indexed(st, "+zero", [EB, SB])
387    }
388
389    /// Negative zero
390    pub fn minus_zero(st: &'st Storage) -> Self {
391        Self::op_const_indexed(st, "-zero", [EB, SB])
392    }
393
394    /// Not a Number
395    pub fn nan(st: &'st Storage) -> Self {
396        Self::op_const_indexed(st, "NaN", [EB, SB])
397    }
398
399    // Operators
400    fn unop<T: From<STerm<'st>>>(self, op: &'st str) -> T {
401        app(self.st(), op, self.term()).into()
402    }
403
404    fn binop<T: From<STerm<'st>>>(self, op: &'st str, other: Self) -> T {
405        app(self.st(), op, (self.term(), other.term())).into()
406    }
407
408    fn ternop_rm<T: From<STerm<'st>>>(
409        self,
410        op: &'st str,
411        rm: RoundingMode<'st>,
412        other1: Self,
413        other2: Self,
414    ) -> T {
415        app(
416            self.st(),
417            op,
418            [rm.term(), self.term(), other1.term(), other2.term()],
419        )
420        .into()
421    }
422
423    fn binop_rm<T: From<STerm<'st>>>(self, op: &'st str, rm: RoundingMode<'st>, other: Self) -> T {
424        app(self.st(), op, (rm.term(), self.term(), other.term())).into()
425    }
426
427    fn unop_rm<T: From<STerm<'st>>>(self, op: &'st str, rm: RoundingMode<'st>) -> T {
428        app(self.st(), op, (rm.term(), self.term())).into()
429    }
430
431    /// Absolute value (`fp.abs`)
432    pub fn fp_abs(self) -> Self {
433        self.unop("fp.abs")
434    }
435
436    /// Negation (`fp.neg`)
437    pub fn fp_neg(self) -> Self {
438        self.unop("fp.neg")
439    }
440
441    /// Addition (`fp.add`)
442    pub fn fp_add(self, rm: RoundingMode<'st>, other: Self) -> Self {
443        self.binop_rm("fp.add", rm, other)
444    }
445
446    /// Subtraction (`fp.sub`)
447    pub fn fp_sub(self, rm: RoundingMode<'st>, other: Self) -> Self {
448        self.binop_rm("fp.sub", rm, other)
449    }
450
451    /// Multiplication (`fp.mul`)
452    pub fn fp_mul(self, rm: RoundingMode<'st>, other: Self) -> Self {
453        self.binop_rm("fp.mul", rm, other)
454    }
455
456    /// Division (`fp.div`)
457    pub fn fp_div(self, rm: RoundingMode<'st>, other: Self) -> Self {
458        self.binop_rm("fp.div", rm, other)
459    }
460
461    /// Fused multiplication and addition: `(self * other1) + other2` (`fp.fma`)
462    pub fn fp_fma(self, rm: RoundingMode<'st>, other1: Self, other2: Self) -> Self {
463        self.ternop_rm("fp.fma", rm, other1, other2)
464    }
465
466    /// Square root (`fp.sqrt`)
467    pub fn fp_sqrt(self, rm: RoundingMode<'st>) -> Self {
468        self.unop_rm("fp.sqrt", rm)
469    }
470
471    /// Remainder: `self - other * n`, where `n` in Z is nearest to `self/other`
472    /// (`fp.rem`)
473    pub fn fp_rem(self, other: Self) -> Self {
474        self.binop("fp.rem", other)
475    }
476
477    /// Rounding to integral (`fp.roundToIntegral`)
478    pub fn fp_round_to_integral(self, rm: RoundingMode<'st>) -> Self {
479        self.unop_rm("fp.roundToIntegral", rm)
480    }
481
482    /// Minimum (`fp.min`)
483    pub fn fp_min(self, other: Self) -> Self {
484        self.binop("fp.min", other)
485    }
486
487    /// Maximum (`fp.max`)
488    pub fn fp_max(self, other: Self) -> Self {
489        self.binop("fp.max", other)
490    }
491
492    /// Less than or equal (`fp.leq`)
493    pub fn fp_leq(self, other: Self) -> Bool<'st> {
494        self.binop("fp.leq", other)
495    }
496
497    /// Less than (`fp.lt`)
498    pub fn fp_lt(self, other: Self) -> Bool<'st> {
499        self.binop("fp.lt", other)
500    }
501
502    /// Greater than or equal (`fp.geq`)
503    pub fn fp_geq(self, other: Self) -> Bool<'st> {
504        self.binop("fp.geq", other)
505    }
506
507    /// Greater than (`fp.gt`)
508    pub fn fp_gt(self, other: Self) -> Bool<'st> {
509        self.binop("fp.gt", other)
510    }
511
512    /// IEEE 754-2008 equality (`fp.eq`)
513    pub fn fp_eq(self, other: Self) -> Bool<'st> {
514        self.binop("fp.eq", other)
515    }
516
517    /// Is normal (`fp.isNormal`)
518    pub fn fp_is_normal(self) -> Bool<'st> {
519        self.unop("fp.isNormal")
520    }
521    /// Is subnormal (`fp.isSubnormal`)
522    pub fn fp_is_subnormal(self) -> Bool<'st> {
523        self.unop("fp.isSubnormal")
524    }
525    /// Is zero (`fp.isZero`)
526    pub fn fp_is_zero(self) -> Bool<'st> {
527        self.unop("fp.isZero")
528    }
529    /// Is infinite (`fp.isInfinite`)
530    pub fn fp_is_infinite(self) -> Bool<'st> {
531        self.unop("fp.isInfinite")
532    }
533    /// Is NaN (`fp.isNaN`)
534    pub fn fp_is_nan(self) -> Bool<'st> {
535        self.unop("fp.isNaN")
536    }
537    /// Is negative (`fp.isNegative`)
538    pub fn fp_is_negative(self) -> Bool<'st> {
539        self.unop("fp.isNegative")
540    }
541    /// Is positive (`fp.isPositive`)
542    pub fn fp_is_positive(self) -> Bool<'st> {
543        self.unop("fp.isPositive")
544    }
545
546    // Conversions
547
548    /// From single bitstring representation in IEEE 754-2008 interchange
549    /// format. `M = EB + SB` (`to_fp`)
550    pub fn to_fp_from_bits<const M: usize>(st: &'st Storage, bv: BitVec<'st, M>) -> Self {
551        assert_eq!(M, EB + SB, "BitVec size M must be EB + SB");
552        Self::app_fn_indexed_args(st, "to_fp", [EB, SB], bv.term())
553    }
554
555    /// From another floating point sort (`to_fp`)
556    pub fn to_fp_from_fp<const MB: usize, const NB: usize>(
557        st: &'st Storage,
558        rm: RoundingMode<'st>,
559        fp_other: FloatingPoint<'st, MB, NB>,
560    ) -> Self {
561        Self::app_fn_indexed_args(st, "to_fp", [EB, SB], (rm.term(), fp_other.term()))
562    }
563
564    /// From real (`to_fp`)
565    pub fn to_fp_from_real(st: &'st Storage, rm: RoundingMode<'st>, real: Real<'st>) -> Self {
566        Self::app_fn_indexed_args(st, "to_fp", [EB, SB], (rm.term(), real.term()))
567    }
568
569    /// From signed machine integer, represented as a 2's complement bit vector
570    /// (`to_fp`)
571    pub fn to_fp_from_signed_bit_vec<const M: usize>(
572        st: &'st Storage,
573        rm: RoundingMode<'st>,
574        bv: BitVec<'st, M>,
575    ) -> Self {
576        Self::app_fn_indexed_args(st, "to_fp", [EB, SB], (rm.term(), bv.term()))
577    }
578
579    /// From unsigned machine integer, represented as bit vector
580    /// (`to_fp_unsigned`)
581    pub fn to_fp_from_unsigned_bit_vec<const M: usize>(
582        st: &'st Storage,
583        rm: RoundingMode<'st>,
584        bv: BitVec<'st, M>,
585    ) -> Self {
586        Self::app_fn_indexed_args(st, "to_fp_unsigned", [EB, SB], (rm.term(), bv.term()))
587    }
588
589    /// To unsigned machine integer, represented as a bit vector (`fp.to_ubv`)
590    pub fn fp_to_ubv<const M: usize>(self, rm: RoundingMode<'st>) -> BitVec<'st, M> {
591        self.unop_rm_indexed("fp.to_ubv", rm, M)
592    }
593
594    /// To signed machine integer, represented as a 2's complement bit vector
595    /// (`fp.to_sbv`)
596    pub fn fp_to_sbv<const M: usize>(self, rm: RoundingMode<'st>) -> BitVec<'st, M> {
597        self.unop_rm_indexed("fp.to_sbv", rm, M)
598    }
599
600    /// To real (`fp.to_real`)
601    pub fn fp_to_real(self) -> Real<'st> {
602        self.unop("fp.to_real")
603    }
604}
605
606#[cfg(test)]
607mod tests {
608    use smtlib_lowlevel::{backend::z3_binary::Z3Binary, StderrLogger, Storage};
609
610    use super::*;
611    use crate::{terms::Sorted, theories::fixed_size_bit_vectors::BitVec, SatResult, Solver};
612
613    fn test_solver<'a>(st: &'a Storage) -> Solver<'a, Z3Binary> {
614        let mut res = Solver::new(st, Z3Binary::new("z3").unwrap()).unwrap();
615        res.set_logger(StderrLogger);
616        res
617    }
618
619    #[test]
620    fn test_fp_convert_rust_floats() {
621        let st = Storage::new();
622        let mut solver = test_solver(&st);
623
624        let f32_const = Float32::new_const(&st, "f32_const");
625        let f64_const = Float64::new_const(&st, "f64_const");
626
627        for f in [
628            0.0,
629            -0.0,
630            1.0,
631            1. / 3.,
632            -1. / 3.,
633            123456.789,
634            f64::MIN,
635            f64::MAX,
636            f64::MIN_POSITIVE,
637            f64::EPSILON,
638            f64::INFINITY,
639            f64::NEG_INFINITY,
640        ] {
641            let model = solver
642                .scope(|solver| {
643                    solver.assert(f64_const._eq(f))?;
644                    solver.assert(f32_const._eq(f as f32))?;
645                    let f64_bits = BitVec::new_prim(&st, f.to_bits());
646                    solver.assert(f64_const._eq(Float64::to_fp_from_bits(&st, f64_bits)))?;
647                    solver.check_sat()?;
648                    solver.get_model()
649                })
650                .unwrap();
651            let f_model: f64 = model.eval(f64_const).unwrap().try_into().unwrap();
652            let f_model_32: f32 = model.eval(f32_const).unwrap().try_into().unwrap();
653            assert_eq!(f_model, f);
654            assert_eq!(f_model_32, f as f32);
655        }
656    }
657
658    #[test]
659    fn test_fp_constants_and_classification() -> Result<(), Box<dyn std::error::Error>> {
660        let st = Storage::new();
661        let mut solver = test_solver(&st);
662
663        let p_zero = Float32::plus_zero(&st);
664        let n_zero = Float32::minus_zero(&st);
665        let p_inf = Float32::plus_oo(&st);
666        let n_inf = Float32::minus_oo(&st);
667        let nan = Float32::nan(&st);
668
669        solver.assert(p_zero.fp_is_zero())?;
670        solver.assert(n_zero.fp_is_zero())?;
671        solver.assert(p_inf.fp_is_infinite())?;
672        solver.assert(p_inf.fp_is_positive())?;
673        solver.assert(n_inf.fp_is_infinite())?;
674        solver.assert(n_inf.fp_is_negative())?;
675        solver.assert(nan.fp_is_nan())?;
676
677        solver.assert(!p_zero.fp_is_nan())?;
678        solver.assert(!p_inf.fp_is_nan())?;
679
680        assert_eq!(solver.check_sat()?, SatResult::Sat);
681        Ok(())
682    }
683
684    #[test]
685    fn test_fp_abs_neg() -> Result<(), Box<dyn std::error::Error>> {
686        let st = Storage::new();
687        let mut solver = test_solver(&st);
688
689        let neg_two = Float32::new(&st, -2.0f32);
690        let abs_neg_two = neg_two.fp_abs();
691        let neg_neg_two = neg_two.fp_neg();
692
693        solver.assert(abs_neg_two._eq(2.0))?;
694        solver.assert(neg_neg_two._eq(2.0))?;
695
696        solver.assert(neg_two.fp_is_negative())?;
697        solver.assert(!neg_two.fp_is_positive())?;
698        solver.assert(!abs_neg_two.fp_is_negative())?;
699        solver.assert(abs_neg_two.fp_is_positive())?;
700
701        assert_eq!(solver.check_sat()?, SatResult::Sat);
702        Ok(())
703    }
704
705    #[test]
706    fn test_fp_add() -> Result<(), Box<dyn std::error::Error>> {
707        let st = Storage::new();
708        let mut solver = test_solver(&st);
709        let rne = RoundingMode::rne(&st);
710        let one = Float32::new(&st, 1.0f32);
711        let sum_one_one = one.fp_add(rne, one);
712        solver.assert(sum_one_one._eq(2.0))?;
713
714        assert_eq!(solver.check_sat()?, SatResult::Sat);
715        Ok(())
716    }
717
718    #[test]
719    fn test_fp_fma() -> Result<(), Box<dyn std::error::Error>> {
720        let st = Storage::new();
721        let mut solver = test_solver(&st);
722        let rne = RoundingMode::rne(&st);
723
724        let one = Float32::new(&st, 1.0f32);
725        let two = Float32::new(&st, 2.0f32);
726        let three = Float32::new(&st, 3.0f32);
727
728        let result = one.fp_fma(rne, two, three); // (one * two) + three
729        solver.assert(result._eq(5.0f32))?;
730
731        assert_eq!(solver.check_sat()?, SatResult::Sat);
732        Ok(())
733    }
734
735    #[test]
736    fn test_to_fp_from_bit_vec() -> Result<(), Box<dyn std::error::Error>> {
737        const EXP_BITS_F32: usize = 8;
738        const SIG_BITS_F32: usize = 24;
739
740        let st = Storage::new();
741        let mut solver = test_solver(&st);
742
743        let ieee_1_0_val: i64 = 0x3f800000;
744        let ieee_1_0_bv: BitVec<{ EXP_BITS_F32 + SIG_BITS_F32 }> = BitVec::new(&st, ieee_1_0_val);
745
746        let fp_val = Float32::to_fp_from_bits(&st, ieee_1_0_bv);
747        solver.assert(fp_val._eq(1.0))?;
748
749        assert_eq!(solver.check_sat()?, SatResult::Sat);
750        Ok(())
751    }
752
753    #[test]
754    fn test_fp_to_ubv() -> Result<(), Box<dyn std::error::Error>> {
755        let st = Storage::new();
756        let mut solver = test_solver(&st);
757        let rtz = RoundingMode::rtz(&st);
758        let rne = RoundingMode::rne(&st);
759
760        let two_fp = Float32::new(&st, 2.0f32);
761        solver.assert(two_fp.fp_to_ubv::<42>(rtz)._eq(2i64))?;
762        solver.assert(two_fp.fp_to_ubv::<42>(rne)._eq(2i64))?;
763
764        let two_point_75_fp = Float32::new(&st, 2.75f32);
765        solver.assert(two_point_75_fp.fp_to_ubv::<42>(rtz)._eq(2i64))?;
766        solver.assert(two_point_75_fp.fp_to_ubv::<42>(rne)._eq(3i64))?;
767
768        assert_eq!(solver.check_sat()?, SatResult::Sat);
769        Ok(())
770    }
771
772    #[test]
773    fn test_usage() {
774        let st = Storage::new();
775        let mut solver = test_solver(&st);
776
777        let inp_3 = BitVec::<8>::new_const(&st, "inp-3");
778        let inp_2 = BitVec::<8>::new_const(&st, "inp-2");
779        let inp_1 = BitVec::<8>::new_const(&st, "inp-1");
780        let inp_0 = BitVec::<8>::new_const(&st, "inp-0");
781
782        let concat_3_2 = inp_3.concat_::<8, 16>(inp_2);
783        let concat_1_0 = inp_1.concat_::<8, 16>(inp_0);
784        let full_bits = concat_3_2.concat_::<16, 32>(concat_1_0);
785
786        let fp_from_bits = Float32::to_fp_from_bits(&st, full_bits);
787        let fp_constant = Float32::to_fp_from_bits(&st, BitVec::<32>::new(&st, 0x4974_2400i64));
788
789        let fp_lt = fp_from_bits.fp_lt(fp_constant);
790        let one_bv = BitVec::<32>::new(&st, 1i64);
791        let zero_bv = BitVec::<32>::new(&st, 0i64);
792
793        let ite_result = fp_lt.ite(one_bv, zero_bv);
794        let not_eq = !ite_result._eq(zero_bv);
795
796        solver.assert(not_eq).unwrap();
797        solver.check_sat().unwrap();
798        let _ = solver.get_model().unwrap();
799    }
800}