Skip to main content

cedar_policy_symcc/symcc/
bitvec.rs

1/*
2 * Copyright Cedar Contributors
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *      https://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17//! Implementation of [`BitVec`].
18
19use std::sync::LazyLock;
20
21use crate::symcc::type_abbrevs::{Int, Nat, Width};
22use miette::Diagnostic;
23use num_bigint::{BigInt, BigUint, ToBigInt};
24use num_traits::cast::ToPrimitive;
25use thiserror::Error;
26
27/// Implementation of the Lean BitVec in Rust. The Lean version is a wrapper around a `Fin`,
28/// a finite natural number that is guaranteed to be less than 2^width. In our implementation
29/// we use a BigUint and enforce the invariant that it is less than 2^width. Trying to
30/// create a bit-vector from a value greater than 2^width will truncate the value.
31#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
32pub struct BitVec {
33    width: Width,
34    v: BigUint,
35}
36
37static TWO: LazyLock<BigUint> = LazyLock::new(|| BigUint::from(2u128));
38
39/// Errors in [`BitVec`] operations.
40#[derive(Debug, Diagnostic, Error)]
41pub enum BitVecError {
42    /// Extract out of bounds.
43    #[error("extract out of bounds")]
44    ExtractOutOfBounds,
45    /// Mismatched bit-vector widths in various operations.
46    #[error("mismatched bit-vector widths in {0}")]
47    MismatchedWidths(String),
48    /// Shift amount too large to fit in u32.
49    #[error("shift amount too large to fit in u32")]
50    ShiftAmountTooLarge,
51}
52
53type Result<T> = std::result::Result<T, BitVecError>;
54
55impl BitVec {
56    /// Converts an (unsigned) [`Nat`] into a [`BitVec`] of the given width.
57    ///
58    /// Silently wraps if `v` does not fit.
59    pub fn of_nat(width: Width, v: Nat) -> Self {
60        BitVec::new(width, v)
61    }
62
63    /// Converts a (signed) [`Int`] into a [`BitVec`] of the given width.
64    ///
65    /// Silently wraps if `v` does not fit.
66    pub fn of_int(width: Width, v: Int) -> Self {
67        if v >= BigInt::ZERO {
68            #[expect(
69                clippy::unwrap_used,
70                reason = "already checked that `v` is nonnegative, so `to_biguint()` will return `Some`."
71            )]
72            BitVec::new(width, v.to_biguint().unwrap())
73        } else {
74            // Do 2's complement encoding for the given bit-width.
75            #[expect(
76                clippy::unwrap_used,
77                reason = "Safe because -v is guaranteed to be positive now."
78            )]
79            let pos = BitVec::new(width, (-v).to_biguint().unwrap());
80            #[expect(
81                clippy::expect_used,
82                reason = "Both arguments have width equal to `width`"
83            )]
84            BitVec::add(&pos.not(), &BitVec::of_nat(width, BigUint::from(1u128)))
85                .expect("both arguments have width equal to `width`")
86        }
87    }
88
89    /// Converts a [`u128`] into a [`BitVec`] of the given width.
90    ///
91    /// Silently wraps if `val` does not fit.
92    pub fn of_u128(width: Width, val: u128) -> Self {
93        BitVec::of_nat(width, BigUint::from(val))
94    }
95
96    /// Converts an [`i128`] into a [`BitVec`] of the given width.
97    ///
98    /// Silently wraps if `val` does not fit.
99    pub fn of_i128(width: Width, val: i128) -> Self {
100        BitVec::of_int(width, BigInt::from(val))
101    }
102
103    /// Interprets a [`BitVec`] as a [`Nat`].
104    pub fn to_nat(&self) -> Nat {
105        self.v.clone()
106    }
107
108    /// Interprets a [`BitVec`] as a [`Nat`].
109    pub fn as_nat(&self) -> &Nat {
110        &self.v
111    }
112
113    /// Interprets a [`BitVec`] as an [`Int`].
114    pub fn to_int(&self) -> Int {
115        let sign_bit = self.msb();
116        if self.width.get() < 2 {
117            if sign_bit {
118                BigInt::from(-1)
119            } else {
120                BigInt::ZERO
121            }
122        } else {
123            // extract_bits follows the SMT-LIB semantics of returning bits from [i:j].
124            // Val returns the 2's complement value without the sign bit.
125            #[expect(clippy::unwrap_used, reason = "Already checked that width is not < 2")]
126            let val = self.extract_bits(0, self.width.get() - 2).unwrap();
127            #[expect(
128                clippy::unwrap_used,
129                reason = "The implementation of BigUint::to_bigint always returns Some"
130            )]
131            let val_bigint = BigUint::to_bigint(&val.v).unwrap();
132            if !sign_bit {
133                val_bigint
134            } else {
135                #[expect(
136                    clippy::unwrap_used,
137                    reason = "The implementation of BigUint::to_bigint always returns Some"
138                )]
139                let res =
140                    -1 * BigUint::to_bigint(&TWO.pow(self.width.get() - 1)).unwrap() + val_bigint;
141                res
142            }
143        }
144    }
145
146    /// Returns an integer representing the extracted bits from low to high, inclusive.
147    ///
148    /// `low` and/or `high` may be 0 without causing errors. However, we must
149    /// have `low <= high < self.width`, else we'll get `Err` (not panic).
150    pub fn extract_bits(&self, low: u32, high: u32) -> Result<Self> {
151        if low <= high && high < self.width.get() {
152            let rem = &self.v % TWO.pow(high + 1);
153            let quotient = rem / TWO.pow(low);
154            #[expect(
155                clippy::expect_used,
156                reason = "Because we add 1, the value cannot be 0"
157            )]
158            Ok(BitVec::of_nat(
159                Width::new(high - low + 1).expect("because we add 1, the value cannot be 0"),
160                quotient,
161            ))
162        } else {
163            Err(BitVecError::ExtractOutOfBounds)
164        }
165    }
166
167    //// Helper functions
168
169    /// Synonym for `of_nat()`
170    fn new(width: Width, val: Nat) -> Self {
171        // Unlike in Lean, we do not need to check if `width` is 0 because `width` is a `NonZeroU32` and thus cannot be 0 by construction
172        let v = val % TWO.pow(width.get());
173        BitVec { width, v }
174    }
175
176    /// Returns a bit-vector with all bits set to 1 of the given width
177    fn all_ones(width: Width) -> Self {
178        let all_ones = TWO.pow(width.get() + 1) - 1u32;
179        BitVec::of_nat(width, all_ones)
180    }
181
182    /// Returns whether the most significant bit is set
183    fn msb(&self) -> bool {
184        #[expect(
185            clippy::unwrap_used,
186            reason = "these arguments to extract_bits must always satisfy low <= high < self.width. note that self.width is a NonZeroU32 and thus cannot be 0"
187        )]
188        let bit = self
189            .extract_bits(self.width.get() - 1, self.width.get() - 1)
190            .unwrap()
191            .v;
192        bit != BigUint::ZERO
193    }
194
195    /// Returns whether the bit-vector is zero.
196    fn is_zero(&self) -> bool {
197        self.v == BigUint::ZERO
198    }
199
200    ////
201    // Functions from SymCC/Data.lean
202    ////
203
204    /// Returns the bit-width of the bit-vector.
205    pub const fn width(&self) -> Width {
206        self.width
207    }
208
209    /// Returns (as Int) the minimum signed value that fits in the given bit-width.
210    ///
211    /// Compare to `int_min()`, which returns the minimum signed value as a `BitVec`.
212    pub fn signed_min(n: Width) -> Int {
213        // Unlike in Lean, we do not need to check if `width` is 0 because `width` is a `NonZeroU32` and thus cannot be 0 by construction
214        #[expect(
215            clippy::unwrap_used,
216            reason = "The implementation of BigUint::to_bigint always returns Some"
217        )]
218        let two_to_n_minus_1 = BigUint::to_bigint(&TWO.pow(n.get() - 1)).unwrap();
219        -two_to_n_minus_1
220    }
221
222    /// Returns the maximum signed value that fits in the given bit-width.
223    pub fn signed_max(n: Width) -> Int {
224        // Unlike in Lean, we do not need to check if `width` is 0 because `width` is a `NonZeroU32` and thus cannot be 0 by construction
225        #[expect(
226            clippy::unwrap_used,
227            reason = "The implementation of BigUint::to_bigint always returns Some"
228        )]
229        let two_to_n_minus_1 = BigUint::to_bigint(&TWO.pow(n.get() - 1)).unwrap();
230        two_to_n_minus_1 - 1
231    }
232
233    /// Checks if the given [`Int`] fits in the bit-width.
234    pub fn overflows(n: Width, i: &Int) -> bool {
235        i < &BitVec::signed_min(n) || i > &BitVec::signed_max(n)
236    }
237
238    ////
239    // Functions from Lean BitVec standard library
240    ////
241
242    /// Bitwise not.
243    pub fn not(&self) -> Self {
244        BitVec::of_nat(self.width, &self.v ^ BitVec::all_ones(self.width).v)
245    }
246
247    /// Bit-vector negation.
248    pub fn neg(&self) -> Self {
249        let one = BitVec::of_u128(self.width, 1);
250        #[expect(
251            clippy::unwrap_used,
252            reason = "`self.not()` and `one` have width equal to `self.width`"
253        )]
254        BitVec::add(&self.not(), &one).unwrap()
255    }
256
257    /// Minimum signed value of the given bit-width, encoded as a [`BitVec`].
258    ///
259    /// Compare to `signed_min()`, which returns the minimum signed value as an `Int`.
260    pub fn int_min(width: Width) -> Self {
261        BitVec::of_nat(width, TWO.pow(width.get() - 1))
262    }
263
264    /// Bit-vector signed less-than.
265    pub fn slt(lhs: &Self, rhs: &Self) -> Result<bool> {
266        if lhs.width != rhs.width {
267            Err(BitVecError::MismatchedWidths("slt".into()))
268        } else {
269            Ok(lhs.to_int() < rhs.to_int())
270        }
271    }
272
273    /// Bit-vector signed less-than-or-equal.
274    pub fn sle(lhs: &Self, rhs: &Self) -> Result<bool> {
275        if lhs.width != rhs.width {
276            Err(BitVecError::MismatchedWidths("sle".into()))
277        } else {
278            Ok(lhs.to_int() <= rhs.to_int())
279        }
280    }
281
282    /// Bit-vector unsigned less-than-or-equal.
283    pub fn ule(lhs: &Self, rhs: &Self) -> Result<bool> {
284        if lhs.width != rhs.width {
285            Err(BitVecError::MismatchedWidths("ule".into()))
286        } else {
287            Ok(lhs.v <= rhs.v)
288        }
289    }
290
291    /// Bit-vector unsigned less-than.
292    pub fn ult(lhs: &Self, rhs: &Self) -> Result<bool> {
293        if lhs.width != rhs.width {
294            Err(BitVecError::MismatchedWidths("ult".into()))
295        } else {
296            Ok(lhs.v < rhs.v)
297        }
298    }
299
300    /// Bit-vector addition.
301    ///
302    /// Only returns `Err` if the `lhs` and `rhs` widths mismatch.
303    /// In particular, overflow is not an `Err`.
304    pub fn add(lhs: &Self, rhs: &Self) -> Result<Self> {
305        if lhs.width != rhs.width {
306            Err(BitVecError::MismatchedWidths("add".into()))
307        } else {
308            Ok(BitVec::of_nat(lhs.width, &lhs.v + &rhs.v))
309        }
310    }
311
312    /// Bit-vector subtraction.
313    ///
314    /// Only returns `Err` if the `lhs` and `rhs` widths mismatch.
315    /// In particular, overflow is not an `Err`.
316    pub fn sub(lhs: &Self, rhs: &Self) -> Result<Self> {
317        if lhs.width != rhs.width {
318            Err(BitVecError::MismatchedWidths("sub".into()))
319        } else {
320            BitVec::add(lhs, &rhs.neg())
321        }
322    }
323
324    /// Bit-vector multiplication.
325    ///
326    /// Only returns `Err` if the `lhs` and `rhs` widths mismatch.
327    /// In particular, overflow is not an `Err`.
328    pub fn mul(lhs: &Self, rhs: &Self) -> Result<Self> {
329        if lhs.width != rhs.width {
330            Err(BitVecError::MismatchedWidths("mul".into()))
331        } else {
332            Ok(BitVec::of_nat(lhs.width, &lhs.v * &rhs.v))
333        }
334    }
335
336    /// Bit-vector unsigned division.
337    ///
338    /// Semantics to match SMT bit-vector theory here: <https://smt-lib.org/theories-FixedSizeBitVectors.shtml>
339    ///
340    /// Only returns `Err` if the `lhs` and `rhs` widths mismatch.
341    /// In particular, overflow is not an `Err`.
342    pub fn udiv(lhs: &Self, rhs: &Self) -> Result<Self> {
343        if lhs.width != rhs.width {
344            return Err(BitVecError::MismatchedWidths("udiv".into()));
345        };
346        if rhs.v == BigUint::ZERO {
347            Ok(BitVec::all_ones(lhs.width))
348        } else {
349            Ok(BitVec::of_nat(lhs.width, &lhs.v / &rhs.v))
350        }
351    }
352
353    /// Bit-vector unsigned remainder.
354    ///
355    /// Semantics to match SMT bit-vector theory here: <https://smt-lib.org/theories-FixedSizeBitVectors.shtml>
356    ///
357    /// Only returns `Err` if the `lhs` and `rhs` widths mismatch.
358    /// In particular, overflow is not an `Err`.
359    pub fn urem(lhs: &Self, rhs: &Self) -> Result<Self> {
360        if lhs.width != rhs.width {
361            return Err(BitVecError::MismatchedWidths("urem".into()));
362        };
363        if rhs.v == BigUint::ZERO {
364            Ok(lhs.clone())
365        } else {
366            Ok(BitVec::of_nat(lhs.width, &lhs.v % &rhs.v))
367        }
368    }
369
370    /// Bit-vector signed division.
371    ///
372    /// Semantics to match SMT bit-vector logic here: <https://smt-lib.org/logics-all.shtml>
373    ///
374    /// Only returns `Err` if the `lhs` and `rhs` widths mismatch.
375    /// In particular, overflow is not an `Err`.
376    pub fn sdiv(lhs: &Self, rhs: &Self) -> Result<Self> {
377        if lhs.width != rhs.width {
378            return Err(BitVecError::MismatchedWidths("sdiv".into()));
379        };
380        let lhs_msb = lhs.msb();
381        let rhs_msb = rhs.msb();
382
383        if !lhs_msb && !rhs_msb {
384            BitVec::udiv(lhs, rhs)
385        } else if lhs_msb && !rhs_msb {
386            Ok(BitVec::neg(&BitVec::udiv(&BitVec::neg(lhs), rhs)?))
387        } else if !lhs_msb && rhs_msb {
388            Ok(BitVec::neg(&BitVec::udiv(lhs, &BitVec::neg(rhs))?))
389        } else {
390            BitVec::udiv(&BitVec::neg(lhs), &BitVec::neg(rhs))
391        }
392    }
393
394    /// Bit-vector signed remainder.
395    ///
396    /// Semantics to match SMT bit-vector logic here: <https://smt-lib.org/logics-all.shtml>
397    ///
398    /// Only returns `Err` if the `lhs` and `rhs` widths mismatch.
399    /// In particular, overflow is not an `Err`.
400    pub fn srem(lhs: &Self, rhs: &Self) -> Result<Self> {
401        if lhs.width != rhs.width {
402            return Err(BitVecError::MismatchedWidths("srem".into()));
403        };
404        let lhs_msb = lhs.msb();
405        let rhs_msb = rhs.msb();
406
407        if !lhs_msb && !rhs_msb {
408            BitVec::urem(lhs, rhs)
409        } else if lhs_msb && !rhs_msb {
410            Ok(BitVec::neg(&BitVec::urem(&BitVec::neg(lhs), rhs)?))
411        } else if !lhs_msb && rhs_msb {
412            BitVec::urem(lhs, &BitVec::neg(rhs))
413        } else {
414            Ok(BitVec::neg(&BitVec::urem(
415                &BitVec::neg(lhs),
416                &BitVec::neg(rhs),
417            )?))
418        }
419    }
420
421    /// Bit-vector signed modulus.
422    ///
423    /// Semantics to match SMT bit-vector logic here: <https://smt-lib.org/logics-all.shtml>
424    ///
425    /// Only returns `Err` if the `lhs` and `rhs` widths mismatch.
426    /// In particular, overflow is not an `Err`.
427    pub fn smod(lhs: &Self, rhs: &Self) -> Result<Self> {
428        if lhs.width != rhs.width {
429            return Err(BitVecError::MismatchedWidths("smod".into()));
430        };
431        let lhs_msb = lhs.msb();
432        let rhs_msb = rhs.msb();
433
434        let abs_lhs = if !lhs_msb { lhs } else { &BitVec::neg(lhs) };
435
436        let abs_rhs = if !rhs_msb { rhs } else { &BitVec::neg(rhs) };
437
438        let u = BitVec::urem(abs_lhs, abs_rhs)?;
439        if u.is_zero() || (!lhs_msb && !rhs_msb) {
440            Ok(u)
441        } else if lhs_msb && !rhs_msb {
442            BitVec::add(&BitVec::neg(&u), rhs)
443        } else if !lhs_msb && rhs_msb {
444            BitVec::add(&u, rhs)
445        } else {
446            Ok(BitVec::neg(&u))
447        }
448    }
449
450    /// Bit-vector left shift.
451    ///
452    /// Returns `Err` if the `lhs` and `rhs` widths mismatch, or if the shift
453    /// amount does not fit in a `u32`.
454    pub fn shl(lhs: &Self, rhs: &Self) -> Result<Self> {
455        if lhs.width != rhs.width {
456            return Err(BitVecError::MismatchedWidths("shl".into()));
457        };
458        let shift_amount = rhs.v.to_u32().ok_or(BitVecError::ShiftAmountTooLarge)?;
459        let val = &lhs.v * TWO.pow(shift_amount);
460        Ok(BitVec::of_nat(lhs.width, val))
461    }
462
463    /// Bit-vector logical right shift.
464    ///
465    /// Returns `Err` if the `lhs` and `rhs` widths mismatch, or if the shift
466    /// amount does not fit in a `u32`.
467    pub fn lshr(lhs: &Self, rhs: &Self) -> Result<Self> {
468        if lhs.width != rhs.width {
469            return Err(BitVecError::MismatchedWidths("lshr".into()));
470        };
471        let shift_amount = rhs.v.to_u32().ok_or(BitVecError::ShiftAmountTooLarge)?;
472        let val = &lhs.v / TWO.pow(shift_amount);
473        Ok(BitVec::of_nat(lhs.width, val))
474    }
475
476    /// Bit-vector concatenation.
477    ///
478    /// Panics if the total width exceeds u32::MAX.
479    /// As of this writing, we shouldn't ever construct any bitvector longer than 128,
480    /// which is an extremely long way from u32::MAX.
481    pub fn concat(lhs: &Self, rhs: &Self) -> Result<Self> {
482        #[expect(
483            clippy::expect_used,
484            reason = "Function is documented to panic if total width exceeds u32::MAX"
485        )]
486        let width = lhs
487            .width
488            .checked_add(rhs.width.get())
489            .expect("width will not overflow u32");
490        let new_val = (&lhs.v << rhs.width().get()) + &rhs.v;
491        Ok(BitVec::of_nat(width, new_val))
492    }
493
494    /// Bit-vector unsigned (zero) extension.
495    ///
496    /// This matches the Lean implementation that just adjusts the length of the
497    /// bit-vector to match n (and not the SMT-LIB implementation that zero extends
498    /// the bit-vector by n bits). If n is less than the current bit-width it will
499    /// truncate
500    pub fn zero_extend(bv: &Self, n: Width) -> Self {
501        BitVec::of_nat(n, bv.to_nat())
502    }
503}
504
505impl std::fmt::Display for BitVec {
506    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
507        write!(f, "(bv{} {})", self.width(), self.as_nat())
508    }
509}
510
511#[cfg(test)]
512mod tests {
513    use super::*;
514
515    /// Creates a BitVec from a binary string of '0's and '1's
516    /// e.g. "0010" creates a bit vector of width 4 with value 2
517    fn from_bin_str(s: &str) -> BitVec {
518        assert!(!s.is_empty(), "Cannot create bitvector from empty string.");
519        // Validate that the string only contains '0' and '1'
520        for c in s.chars() {
521            assert!(
522                c == '0' || c == '1',
523                "Binary string must only contain '0' or '1'"
524            );
525        }
526
527        // Parse the binary string into a BigInt
528        let val = BigUint::parse_bytes(s.as_bytes(), 2).unwrap();
529        BitVec::of_nat(
530            Width::new(s.len().try_into().unwrap())
531                .expect("already checked that the string length is not 0"),
532            val,
533        )
534    }
535
536    /// Panics if `width` is 0
537    #[track_caller]
538    fn bitvec(width: u32, val: u128) -> BitVec {
539        BitVec::of_u128(Width::new(width).unwrap(), val)
540    }
541
542    /// Panics if `width` is 0
543    #[track_caller]
544    fn bitvec_i(width: u32, val: i128) -> BitVec {
545        BitVec::of_i128(Width::new(width).unwrap(), val)
546    }
547
548    #[track_caller]
549    fn assert_eq_int(rhs: BigInt, lhs: i32) {
550        assert_eq!(rhs, BigInt::from(lhs));
551    }
552
553    #[track_caller]
554    fn assert_eq_nat(rhs: BigUint, lhs: u32) {
555        assert_eq!(rhs, BigUint::from(lhs));
556    }
557
558    #[test]
559    fn test_from_bin_str() {
560        // Test basic binary string conversion
561        let bv = from_bin_str("0010");
562        assert_eq!(bv.width().get(), 4);
563        assert_eq_nat(bv.to_nat(), 2);
564
565        // Test different values
566        assert_eq_nat(from_bin_str("101").to_nat(), 5);
567        assert_eq_nat(from_bin_str("1111").to_nat(), 15);
568        assert_eq_nat(from_bin_str("10000").to_nat(), 16);
569        assert_eq_nat(from_bin_str("00000").to_nat(), 0);
570    }
571
572    #[test]
573    fn test_constructors() {
574        // Test regular constructor
575        let bv1 = bitvec(4, 2);
576        assert_eq!(bv1.width().get(), 4);
577        assert_eq_nat(bv1.to_nat(), 2);
578
579        // Test value overflow wrapping
580        let bv2 = bitvec(3, 10); // 10 in binary is 1010, truncated to 3 bits: 010
581        assert_eq!(bv2.width().get(), 3);
582        assert_eq_nat(bv2.to_nat(), 2);
583
584        // Test of_nat
585        let bv3 = bitvec(5, 10);
586        assert_eq!(bv3.width().get(), 5);
587        assert_eq_nat(bv3.to_nat(), 10);
588
589        // Test of_int
590        let bv4 = bitvec_i(6, -1); // -1 in two's complement 6-bit is 111111
591        assert_eq!(bv4.width().get(), 6);
592        // A negative number will be represented as its two's complement form
593        assert_eq_nat(bv4.to_nat(), 63); // 111111 in binary is 63 as a natural number
594        assert_eq_int(bv4.to_int(), -1); // 111111 in binary is -1 as a signed number
595
596        assert_eq!(bv4, BitVec::all_ones(bv4.width()));
597    }
598
599    #[test]
600    fn test_extract_bits() {
601        let bv = from_bin_str("110101");
602
603        // Extract middle bits
604        let extracted = bv.extract_bits(1, 3).unwrap(); // Extract bits 1,2,3 (010)
605        assert_eq!(extracted.width().get(), 3);
606        assert_eq_nat(extracted.to_nat(), 2);
607
608        // Extract single bit
609        let bit = bv.extract_bits(5, 5).unwrap(); // Extract the most significant bit
610        assert_eq!(bit.width().get(), 1);
611        assert_eq_nat(bit.to_nat(), 1);
612
613        // Extract all bits
614        let all = bv.extract_bits(0, 5).unwrap(); // Extract all bits
615        assert_eq!(all.width().get(), 6);
616        assert_eq_nat(all.to_nat(), 53); // 110101 = 53
617    }
618
619    #[test]
620    fn test_bitwise_ops() {
621        // Test NOT operation
622        let bv = from_bin_str("1010");
623        let not_bv = bv.not();
624        // In a 4-bit representation, NOT 1010 = 0101
625        assert_eq_nat(not_bv.to_nat(), 5);
626        assert_eq!(not_bv.not(), bv);
627    }
628
629    #[test]
630    fn test_arithmetic_ops() {
631        let bv = from_bin_str("1010");
632        assert_eq_int(bv.to_int(), -6);
633        // Test negation (two's complement)
634        let neg_bv = bv.neg();
635        // NEG 1010 = NOT 1010 + 1 = 0101 + 1 = 0110 = 6
636        assert_eq_nat(neg_bv.to_nat(), 6);
637
638        let bv_min = from_bin_str("1000");
639        assert_eq!(bv_min, bv_min.neg());
640
641        let bv1 = from_bin_str("0101"); // 5
642        let bv2: BitVec = from_bin_str("0011"); // 3
643        let bv3: BitVec = from_bin_str("1011"); // 11
644
645        // Test addition without overflow
646        let sum = BitVec::add(&bv1, &bv2).unwrap();
647        assert_eq_nat(sum.to_nat(), 8); // 5 + 3 = 8
648
649        // Test addition with overflow
650        let sum = BitVec::add(&bv1, &bv3).unwrap();
651        assert_eq_nat(sum.to_nat(), 0); // 5 + 11 = 16 mod 2^4 = 0
652
653        // Test subtraction
654        let diff = BitVec::sub(&bv1, &bv2).unwrap();
655        assert_eq_nat(diff.to_nat(), 2); // 5 - 3 = 2
656
657        // Test subtraction negative value
658        let diff = BitVec::sub(&bv2, &bv1).unwrap();
659        assert_eq_int(diff.to_int(), -2); // 3 - 5 = -2
660
661        // Test multiplication no overflow
662        let prod = BitVec::mul(&bv1, &bv2).unwrap();
663        assert_eq_nat(prod.to_nat(), 15); // 5 * 3 = 15
664
665        // Test multiplication overflow
666        let prod = BitVec::mul(&bv2, &bv3).unwrap();
667        assert_eq_nat(prod.to_nat(), 1); // 3 * 11 mod 2^4 = 33 mod 2^4 = 1
668    }
669
670    #[test]
671    fn test_division() {
672        let bv1 = from_bin_str("0101"); // 5
673        let bv2: BitVec = from_bin_str("0011"); // 3
674
675        // Test division unsigned
676        let quot = BitVec::udiv(&bv1, &bv2).unwrap();
677        assert_eq_nat(quot.to_nat(), 1); // 5 / 3 = 1
678
679        let rem = BitVec::urem(&bv1, &bv2).unwrap();
680        assert_eq_nat(rem.to_nat(), 2); // 5 mod 3 = 2
681
682        // Test division by zero unsigned
683        let zero = bitvec(4, 0);
684        let div_by_zero = BitVec::udiv(&bv1, &zero).unwrap();
685        assert_eq_nat(div_by_zero.to_nat(), 15); // All ones for division by 0
686
687        let rem_by_zero = BitVec::urem(&bv1, &zero).unwrap();
688        assert_eq!(rem_by_zero.to_nat(), bv1.to_nat());
689
690        // Test division signed
691
692        // Both values positive
693        let quot = BitVec::sdiv(&bv1, &bv2).unwrap();
694        assert_eq_nat(quot.to_nat(), 1); // 5 / 3 = 1
695        let srem = BitVec::srem(&bv1, &bv2).unwrap();
696        assert_eq_int(srem.to_int(), 2);
697        let smod = BitVec::smod(&bv1, &bv2).unwrap();
698        assert_eq_int(smod.to_int(), 2);
699
700        // Divident is negative, divisor positive
701        let neg_bv1 = from_bin_str("1011"); // -5 in 4-bit two's complement
702        let pos_bv2 = from_bin_str("0011"); // 3
703
704        // -5 / 3 = -1 (truncated towards zero)
705        let quot_neg_pos = BitVec::sdiv(&neg_bv1, &pos_bv2).unwrap();
706        assert_eq_int(quot_neg_pos.to_int(), -1);
707
708        // -5 srem 3 = -2 (remainder has same sign as dividend)
709        let srem_neg_pos = BitVec::srem(&neg_bv1, &pos_bv2).unwrap();
710        assert_eq_int(srem_neg_pos.to_int(), -2);
711
712        // -5 smod 3 = 1 (result has same sign as divisor when non-zero)
713        let smod_neg_pos = BitVec::smod(&neg_bv1, &pos_bv2).unwrap();
714        assert_eq_int(smod_neg_pos.to_int(), 1);
715
716        // Test with -8 (INT_MIN) / 3
717        let int_min = from_bin_str("1000"); // -8 in 4-bit two's complement
718        let quot_min_pos = BitVec::sdiv(&int_min, &pos_bv2).unwrap();
719        assert_eq_int(quot_min_pos.to_int(), -2); // -8 / 3 = -2
720
721        let srem_min_pos = BitVec::srem(&int_min, &pos_bv2).unwrap();
722        assert_eq_int(srem_min_pos.to_int(), -2); // -8 srem 3 = -2
723
724        let smod_min_pos = BitVec::smod(&int_min, &pos_bv2).unwrap();
725        assert_eq_int(smod_min_pos.to_int(), 1); // -8 smod 3 = 1
726
727        // Divident is positive, divisor negative
728        let pos_bv1 = from_bin_str("0101"); // 5
729        let neg_bv2 = from_bin_str("1101"); // -3 in 4-bit two's complement
730
731        // 5 / (-3) = -1 (truncated towards zero)
732        let quot_pos_neg = BitVec::sdiv(&pos_bv1, &neg_bv2).unwrap();
733        assert_eq_int(quot_pos_neg.to_int(), -1);
734
735        // 5 srem (-3) = 2 (remainder has same sign as dividend)
736        let srem_pos_neg = BitVec::srem(&pos_bv1, &neg_bv2).unwrap();
737        assert_eq_int(srem_pos_neg.to_int(), 2);
738
739        // 5 smod (-3) = -1 (result has same sign as divisor when non-zero)
740        let smod_pos_neg = BitVec::smod(&pos_bv1, &neg_bv2).unwrap();
741        assert_eq_int(smod_pos_neg.to_int(), -1);
742
743        // Test with 7 / (-3)
744        let pos_bv7 = from_bin_str("0111"); // 7
745        let quot_7_neg3 = BitVec::sdiv(&pos_bv7, &neg_bv2).unwrap();
746        assert_eq_int(quot_7_neg3.to_int(), -2); // 7 / (-3) = -2
747
748        let srem_7_neg3 = BitVec::srem(&pos_bv7, &neg_bv2).unwrap();
749        assert_eq_int(srem_7_neg3.to_int(), 1); // 7 srem (-3) = 1
750
751        let smod_7_neg3 = BitVec::smod(&pos_bv7, &neg_bv2).unwrap();
752        assert_eq_int(smod_7_neg3.to_int(), -2); // 7 smod (-3) = -2
753
754        // Divident and divisor negative
755        let neg_bv5 = from_bin_str("1011"); // -5 in 4-bit two's complement
756        let neg_bv3 = from_bin_str("1101"); // -3 in 4-bit two's complement
757
758        // (-5) / (-3) = 1 (positive result)
759        let quot_neg_neg = BitVec::sdiv(&neg_bv5, &neg_bv3).unwrap();
760        assert_eq_int(quot_neg_neg.to_int(), 1);
761
762        // (-5) srem (-3) = -2 (remainder has same sign as dividend)
763        let srem_neg_neg = BitVec::srem(&neg_bv5, &neg_bv3).unwrap();
764        assert_eq_int(srem_neg_neg.to_int(), -2);
765
766        // (-5) smod (-3) = -2 (result has same sign as divisor)
767        let smod_neg_neg = BitVec::smod(&neg_bv5, &neg_bv3).unwrap();
768        assert_eq_int(smod_neg_neg.to_int(), -2);
769
770        // Test with (-8) / (-3)
771        let quot_min_neg = BitVec::sdiv(&int_min, &neg_bv3).unwrap();
772        assert_eq_int(quot_min_neg.to_int(), 2); // (-8) / (-3) = 2
773
774        let srem_min_neg = BitVec::srem(&int_min, &neg_bv3).unwrap();
775        assert_eq_int(srem_min_neg.to_int(), -2); // (-8) srem (-3) = -2
776
777        let smod_min_neg = BitVec::smod(&int_min, &neg_bv3).unwrap();
778        assert_eq_int(smod_min_neg.to_int(), -2); // (-8) smod (-3) = -2
779
780        // Test edge case: (-8) / (-1) - potential overflow case
781        let neg_one = from_bin_str("1111"); // -1 in 4-bit two's complement
782        let quot_min_neg1 = BitVec::sdiv(&int_min, &neg_one).unwrap();
783        assert_eq_int(quot_min_neg1.to_int(), -8); // (-8) / (-1) = 8, but wraps to -8 in 4-bit
784
785        let srem_min_neg1 = BitVec::srem(&int_min, &neg_one).unwrap();
786        assert_eq_int(srem_min_neg1.to_int(), 0); // (-8) srem (-1) = 0
787
788        let smod_min_neg1 = BitVec::smod(&int_min, &neg_one).unwrap();
789        assert_eq_int(smod_min_neg1.to_int(), 0); // (-8) smod (-1) = 0
790
791        // Test division by zero signed
792        let div_by_zero = BitVec::sdiv(&bv1, &zero).unwrap();
793        // Should follow SMT-LIB semantics
794        assert_eq_nat(div_by_zero.to_nat(), 15); // All ones for positive dividend
795
796        let div_by_zero_neg = BitVec::sdiv(&neg_bv1, &zero).unwrap();
797        assert_eq_nat(div_by_zero_neg.to_nat(), 1); // If the divident is negative, then division by 0 is one
798
799        let srem_by_zero = BitVec::srem(&bv1, &zero).unwrap();
800        assert_eq!(srem_by_zero.to_nat(), bv1.to_nat()); // Return dividend for srem by zero
801
802        let smod_by_zero = BitVec::smod(&bv1, &zero).unwrap();
803        assert_eq!(smod_by_zero.to_nat(), bv1.to_nat()); // Return dividend for smod by zero
804    }
805
806    #[test]
807    fn test_comparison_ops() {
808        let bv1 = from_bin_str("0101"); // 5
809        let bv2 = from_bin_str("0011"); // 3
810
811        // Test less than
812        assert!(!BitVec::slt(&bv1, &bv2).unwrap()); // 5 < 3 is false
813        assert!(BitVec::slt(&bv2, &bv1).unwrap()); // 3 < 5 is true
814
815        // Test less than or equal
816        assert!(!BitVec::sle(&bv1, &bv2).unwrap()); // 5 <= 3 is false
817        assert!(BitVec::sle(&bv2, &bv1).unwrap()); // 3 <= 5 is true
818        assert!(BitVec::sle(&bv1, &bv1).unwrap()); // 5 <= 5 is true
819
820        // Test for negative numbers
821        let zero = from_bin_str("0000"); // -5 in 4-bit two's complement
822        let neg_5 = from_bin_str("1011"); // -5 in 4-bit two's complement
823        let neg_3: BitVec = from_bin_str("1101"); // -3 in 4-bit two's complement
824        let neg_8: BitVec = from_bin_str("1000"); // -8 in 4-bit two's complement
825        let one: BitVec = from_bin_str("0001"); // 1 in 4-bit two's complement
826
827        assert!(BitVec::ult(&neg_5, &neg_3).unwrap());
828        assert!(BitVec::ult(&one, &neg_8).unwrap());
829        assert!(BitVec::slt(&neg_5, &neg_3).unwrap());
830        assert!(BitVec::slt(&neg_5, &zero).unwrap());
831    }
832
833    #[test]
834    fn test_shift_ops() {
835        let bv1 = from_bin_str("0101"); // 5
836        let shift1 = from_bin_str("0001"); // shift by 1
837        let shift2 = from_bin_str("0010"); // shift by 2
838
839        // Test left shift
840        let left_shift1 = BitVec::shl(&bv1, &shift1).unwrap(); // 5 << 1 = 10 (1010)
841        assert_eq_nat(left_shift1.to_nat(), 10);
842
843        let left_shift2 = BitVec::shl(&bv1, &shift2).unwrap(); // 5 << 2 = 20 (modulo 16 = 4)
844        assert_eq_nat(left_shift2.to_nat(), 4);
845
846        // Test logical right shift
847        let bv2 = from_bin_str("1100"); // 12
848        let right_shift1 = BitVec::lshr(&bv2, &shift1).unwrap(); // 12 >> 1 = 6
849        assert_eq_nat(right_shift1.to_nat(), 6);
850
851        let right_shift2 = BitVec::lshr(&bv2, &shift2).unwrap(); // 12 >> 2 = 3
852        assert_eq_nat(right_shift2.to_nat(), 3);
853    }
854
855    #[test]
856    fn test_concat() {
857        let bv1 = from_bin_str("101"); // 5
858        let bv2 = from_bin_str("11"); // 3
859
860        let concat1 = BitVec::concat(&bv1, &bv2).unwrap(); // 10111 (binary) = 23
861        assert_eq!(concat1.width().get(), 5); // 3 + 2 = 5 bits
862        assert_eq_nat(concat1.to_nat(), 23);
863
864        let concat2 = BitVec::concat(&bv2, &bv1).unwrap(); // 11101 (binary) = 29
865        assert_eq!(concat2.width().get(), 5); // 2 + 3 = 5 bits
866        assert_eq_nat(concat2.to_nat(), 29);
867    }
868
869    #[test]
870    fn test_zero_extend() {
871        let bv = from_bin_str("101"); // 5
872
873        // Extending to the same width should return the same value
874        let extended1 = BitVec::zero_extend(&bv, Width::new(3).unwrap());
875        assert_eq!(extended1.width().get(), 3);
876        assert_eq_nat(extended1.to_nat(), 5);
877
878        // Testing with a smaller width (allowed by implementation)
879        let extended2 = BitVec::zero_extend(&bv, Width::new(2).unwrap());
880        assert_eq!(extended2.width().get(), 2);
881        assert_eq_nat(extended2.to_nat(), 1); // 101 truncated to 2 bits is 01 = 1
882    }
883
884    #[test]
885    fn test_overflow() {
886        // Test signed_min and signed_max
887        assert_eq!(BitVec::signed_min(Width::new(4).unwrap()), BigInt::from(-8)); // -2^(4-1)
888        assert_eq!(BitVec::signed_max(Width::new(4).unwrap()), BigInt::from(7)); // 2^(4-1) - 1
889
890        // Test overflow detection
891        assert!(BitVec::overflows(Width::new(4).unwrap(), &BigInt::from(8))); // 8 overflows 4-bit signed
892        assert!(BitVec::overflows(Width::new(4).unwrap(), &BigInt::from(-9))); // -9 overflows 4-bit signed
893        assert!(!BitVec::overflows(Width::new(4).unwrap(), &BigInt::from(7))); // 7 doesn't overflow
894        assert!(!BitVec::overflows(
895            Width::new(4).unwrap(),
896            &BigInt::from(-8)
897        )); // -8 doesn't overflow
898
899        // Test int_min
900        let min = BitVec::int_min(Width::new(4).unwrap());
901        assert_eq!(min.width().get(), 4);
902        assert_eq_nat(min.to_nat(), 8); // -8 in 4-bit two's complement is 1000 (8)
903    }
904}