diff --git a/src/analyze/annot.rs b/src/analyze/annot.rs index b7d7098b..2692687d 100644 --- a/src/analyze/annot.rs +++ b/src/analyze/annot.rs @@ -66,6 +66,30 @@ pub fn int_model_path() -> [Symbol; 3] { ] } +pub fn bit_vec_model_path() -> [Symbol; 3] { + [ + Symbol::intern("thrust"), + Symbol::intern("def"), + Symbol::intern("bit_vec_model"), + ] +} + +pub fn bit_vec_from_int_path() -> [Symbol; 3] { + [ + Symbol::intern("thrust"), + Symbol::intern("def"), + Symbol::intern("bit_vec_from_int"), + ] +} + +pub fn bit_vec_to_int_path() -> [Symbol; 3] { + [ + Symbol::intern("thrust"), + Symbol::intern("def"), + Symbol::intern("bit_vec_to_int"), + ] +} + pub fn mut_model_path() -> [Symbol; 3] { [ Symbol::intern("thrust"), diff --git a/src/analyze/annot_fn.rs b/src/analyze/annot_fn.rs index dba2e126..23218e26 100644 --- a/src/analyze/annot_fn.rs +++ b/src/analyze/annot_fn.rs @@ -394,6 +394,49 @@ impl<'a, 'tcx> AnnotFnTranslator<'a, 'tcx> { self.tcx.normalize_erasing_regions(typing_env, instantiated) } + fn bit_vec_ty(&self, expr: &'tcx rustc_hir::Expr<'tcx>) -> Option { + match self.type_builder.build(self.expr_ty(expr)) { + rty::Type::BitVec(ty) => Some(ty), + _ => None, + } + } + + fn bit_vec_binary_op( + &self, + op: rustc_hir::BinOpKind, + ty: rty::BitVecType, + lhs: &'tcx rustc_hir::Expr<'tcx>, + rhs: &'tcx rustc_hir::Expr<'tcx>, + ) -> FormulaOrTerm { + use rustc_hir::BinOpKind; + + let lhs = self.to_term(lhs); + let rhs = self.to_term(rhs); + let fun = match (op, ty.signed) { + (BinOpKind::Eq, _) => return FormulaOrTerm::BinOp(lhs, AmbiguousBinOp::Eq, rhs), + (BinOpKind::Ne, _) => return FormulaOrTerm::BinOp(lhs, AmbiguousBinOp::Ne, rhs), + (BinOpKind::Add, _) => chc::Function::BVADD, + (BinOpKind::Sub, _) => chc::Function::BVSUB, + (BinOpKind::Mul, _) => chc::Function::BVMUL, + (BinOpKind::BitAnd, _) => chc::Function::BVAND, + (BinOpKind::BitOr, _) => chc::Function::BVOR, + (BinOpKind::BitXor, _) => chc::Function::BVXOR, + (BinOpKind::Shl, _) => chc::Function::BVSHL, + (BinOpKind::Shr, true) => chc::Function::BVASHR, + (BinOpKind::Shr, false) => chc::Function::BVLSHR, + (BinOpKind::Lt, true) => chc::Function::BVSLT, + (BinOpKind::Lt, false) => chc::Function::BVULT, + (BinOpKind::Le, true) => chc::Function::BVSLE, + (BinOpKind::Le, false) => chc::Function::BVULE, + (BinOpKind::Gt, true) => chc::Function::BVSGT, + (BinOpKind::Gt, false) => chc::Function::BVUGT, + (BinOpKind::Ge, true) => chc::Function::BVSGE, + (BinOpKind::Ge, false) => chc::Function::BVUGE, + _ => unimplemented!("unsupported BitVec operator in formula: {:?}", op), + }; + FormulaOrTerm::Term(chc::Term::App(fun, vec![lhs, rhs])) + } + fn pat_ty(&self, pat: &'tcx rustc_hir::Pat<'tcx>) -> mir_ty::Ty<'tcx> { let ty = self.typeck.pat_ty(pat); let instantiated = mir_ty::EarlyBinder::bind(ty).instantiate(self.tcx, self.generic_args); @@ -659,6 +702,9 @@ impl<'a, 'tcx> AnnotFnTranslator<'a, 'tcx> { match hir.kind { ExprKind::Binary(op, lhs, rhs) => { + if let Some(ty) = self.bit_vec_ty(lhs) { + return self.bit_vec_binary_op(op.node, ty, lhs, rhs); + } match op.node { rustc_hir::BinOpKind::Or => { let lhs = self.to_formula_or_term(lhs); @@ -703,8 +749,13 @@ impl<'a, 'tcx> AnnotFnTranslator<'a, 'tcx> { } ExprKind::Unary(op, operand) => match op { rustc_hir::UnOp::Neg => { + let is_bit_vec = self.bit_vec_ty(operand).is_some(); let operand = self.to_term(operand); - FormulaOrTerm::Term(operand.neg()) + if is_bit_vec { + FormulaOrTerm::Term(chc::Term::App(chc::Function::BVNEG, vec![operand])) + } else { + FormulaOrTerm::Term(operand.neg()) + } } rustc_hir::UnOp::Not => { let operand_ty = self.expr_ty(operand); @@ -713,6 +764,10 @@ impl<'a, 'tcx> AnnotFnTranslator<'a, 'tcx> { let operand = self.to_term(operand); FormulaOrTerm::Term(operand.mut_final()) } + Some(adt) if Some(adt.did()) == self.def_ids.bit_vec_model() => { + let operand = self.to_term(operand); + FormulaOrTerm::Term(chc::Term::App(chc::Function::BVNOT, vec![operand])) + } _ => { let operand = self.to_formula_or_term(operand); FormulaOrTerm::Not(Box::new(operand)) @@ -847,6 +902,20 @@ impl<'a, 'tcx> AnnotFnTranslator<'a, 'tcx> { let t = self.to_term(receiver); return FormulaOrTerm::Term(t); } + if Some(def_id) == self.def_ids.bit_vec_to_int() { + assert!( + args.is_empty(), + "BitVec::to_int does not take any arguments" + ); + let ty = self.bit_vec_ty(receiver).unwrap(); + let fun = if ty.signed { + chc::Function::SBV_TO_INT + } else { + chc::Function::UBV_TO_INT + }; + let t = self.to_term(receiver); + return FormulaOrTerm::Term(chc::Term::App(fun, vec![t])); + } if Some(def_id) == self.def_ids.seq_len() { assert!(args.is_empty(), "Seq::len does not take any arguments"); let t = self.to_term(receiver); @@ -954,6 +1023,12 @@ impl<'a, 'tcx> AnnotFnTranslator<'a, 'tcx> { let t = self.to_term(&args[0]); return FormulaOrTerm::Term(chc::Term::box_(t)); } + if Some(def_id) == self.def_ids.bit_vec_from_int() { + assert_eq!(args.len(), 1, "BitVec::from_int takes exactly 1 argument"); + let width = self.bit_vec_ty(hir).unwrap().width; + let t = self.to_term(&args[0]); + return FormulaOrTerm::Term(t.int_to_bit_vec(width)); + } if Some(def_id) == self.def_ids.seq_empty() { assert!(args.is_empty(), "Seq::empty does not take any arguments"); let elem_sort = self.node_arg_type_at(func_expr.hir_id, 0).to_sort(); diff --git a/src/analyze/did_cache.rs b/src/analyze/did_cache.rs index 55aa1f40..ea02c777 100644 --- a/src/analyze/did_cache.rs +++ b/src/analyze/did_cache.rs @@ -15,6 +15,7 @@ struct DefIds { model_ty: OnceCell>, int_model: OnceCell>, + bit_vec_model: OnceCell>, mut_model: OnceCell>, box_model: OnceCell>, array_model: OnceCell>, @@ -23,6 +24,8 @@ struct DefIds { mut_model_new: OnceCell>, box_model_new: OnceCell>, array_model_store: OnceCell>, + bit_vec_from_int: OnceCell>, + bit_vec_to_int: OnceCell>, seq_model: OnceCell>, seq_empty: OnceCell>, @@ -140,6 +143,27 @@ impl<'tcx> DefIdCache<'tcx> { .get_or_init(|| self.annotated_def(&crate::analyze::annot::int_model_path())) } + pub fn bit_vec_model(&self) -> Option { + *self + .def_ids + .bit_vec_model + .get_or_init(|| self.annotated_def(&crate::analyze::annot::bit_vec_model_path())) + } + + pub fn bit_vec_from_int(&self) -> Option { + *self + .def_ids + .bit_vec_from_int + .get_or_init(|| self.annotated_def(&crate::analyze::annot::bit_vec_from_int_path())) + } + + pub fn bit_vec_to_int(&self) -> Option { + *self + .def_ids + .bit_vec_to_int + .get_or_init(|| self.annotated_def(&crate::analyze::annot::bit_vec_to_int_path())) + } + pub fn mut_model(&self) -> Option { *self .def_ids diff --git a/src/chc.rs b/src/chc.rs index 13882b2d..5f12e598 100644 --- a/src/chc.rs +++ b/src/chc.rs @@ -90,6 +90,7 @@ impl DatatypeSort { pub enum Sort { Null, Int, + BitVec { width: u32 }, Bool, String, Param(usize), @@ -116,6 +117,7 @@ where match self { Sort::Null => allocator.text("null"), Sort::Int => allocator.text("int"), + Sort::BitVec { width } => allocator.text(format!("bv{width}")), Sort::Bool => allocator.text("bool"), Sort::String => allocator.text("string"), Sort::Param(i) => allocator.text("T").append(allocator.as_string(i)), @@ -185,7 +187,12 @@ impl Sort { fn walk_impl<'a, 'b>(&'a self, mut f: Box) { f(self); match self { - Sort::Null | Sort::Int | Sort::Bool | Sort::String | Sort::Param(_) => {} + Sort::Null + | Sort::Int + | Sort::BitVec { .. } + | Sort::Bool + | Sort::String + | Sort::Param(_) => {} Sort::Box(s) | Sort::Mut(s) | Sort::Seq(s) => s.walk(Box::new(&mut f)), Sort::Tuple(ss) => { for s in ss { @@ -234,6 +241,10 @@ impl Sort { Sort::Int } + pub fn bit_vec(width: u32) -> Self { + Sort::BitVec { width } + } + pub fn string() -> Self { Sort::String } @@ -417,6 +428,26 @@ impl Function { Self::OR => Sort::bool(), Self::NOT => Sort::bool(), Self::NEG => Sort::int(), + Self::BVADD + | Self::BVSUB + | Self::BVMUL + | Self::BVAND + | Self::BVOR + | Self::BVXOR + | Self::BVSHL + | Self::BVLSHR + | Self::BVASHR + | Self::BVNOT + | Self::BVNEG => args.into_iter().next().unwrap(), + Self::BVULT + | Self::BVULE + | Self::BVUGT + | Self::BVUGE + | Self::BVSLT + | Self::BVSLE + | Self::BVSGT + | Self::BVSGE => Sort::bool(), + Self::UBV_TO_INT | Self::SBV_TO_INT => Sort::int(), Self::STORE | Self::SEQ_CONCAT | Self::SEQ_EXTRACT | Self::SEQ_STORE => { args.into_iter().next().unwrap() } @@ -452,6 +483,27 @@ impl Function { pub const OR: Function = Function::infix("or"); pub const NOT: Function = Function::new("not"); pub const NEG: Function = Function::new("-"); + pub const BVADD: Function = Function::new("bvadd"); + pub const BVSUB: Function = Function::new("bvsub"); + pub const BVMUL: Function = Function::new("bvmul"); + pub const BVAND: Function = Function::new("bvand"); + pub const BVOR: Function = Function::new("bvor"); + pub const BVXOR: Function = Function::new("bvxor"); + pub const BVSHL: Function = Function::new("bvshl"); + pub const BVLSHR: Function = Function::new("bvlshr"); + pub const BVASHR: Function = Function::new("bvashr"); + pub const BVNOT: Function = Function::new("bvnot"); + pub const BVNEG: Function = Function::new("bvneg"); + pub const BVULT: Function = Function::new("bvult"); + pub const BVULE: Function = Function::new("bvule"); + pub const BVUGT: Function = Function::new("bvugt"); + pub const BVUGE: Function = Function::new("bvuge"); + pub const BVSLT: Function = Function::new("bvslt"); + pub const BVSLE: Function = Function::new("bvsle"); + pub const BVSGT: Function = Function::new("bvsgt"); + pub const BVSGE: Function = Function::new("bvsge"); + pub const UBV_TO_INT: Function = Function::new("ubv_to_int"); + pub const SBV_TO_INT: Function = Function::new("sbv_to_int"); pub const STORE: Function = Function::new("store"); pub const SELECT: Function = Function::new("select"); pub const SEQ_CONCAT: Function = Function::new("seq.++"); @@ -483,6 +535,10 @@ pub enum Term { TupleProj(Box>, usize), DatatypeCtor(DatatypeSort, DatatypeSymbol, Vec>), DatatypeDiscr(DatatypeSymbol, Box>), + IntToBitVec { + width: u32, + term: Box>, + }, /// Used in [`Formula`] to represent quantified variables appearing in annotations. UserQuantifiedVar(Sort, UserQuantifiedVarId), } @@ -556,6 +612,9 @@ where Term::DatatypeDiscr(_, t) => allocator .text("discriminant") .append(t.pretty(allocator).parens()), + Term::IntToBitVec { width, term } => allocator + .text(format!("int_to_bv{width}")) + .append(term.pretty(allocator).parens()), Term::UserQuantifiedVar(_, var) => allocator.as_string(var), } } @@ -605,6 +664,10 @@ impl Term { args.into_iter().map(|t| t.subst_var(&mut f)).collect(), ), Term::DatatypeDiscr(d_sym, t) => Term::DatatypeDiscr(d_sym, Box::new(t.subst_var(f))), + Term::IntToBitVec { width, term } => Term::IntToBitVec { + width, + term: Box::new(term.subst_var(f)), + }, Term::UserQuantifiedVar(sort, var) => Term::UserQuantifiedVar(sort, var), } } @@ -653,6 +716,7 @@ impl Term { Term::TupleProj(t, i) => t.sort(var_sort).tuple_elem(*i), Term::DatatypeCtor(sort, _, _) => sort.clone().into(), Term::DatatypeDiscr(_, _) => Sort::int(), + Term::IntToBitVec { width, .. } => Sort::bit_vec(*width), Term::UserQuantifiedVar(sort, _) => sort.clone(), } } @@ -676,7 +740,7 @@ impl Term { Term::Tuple(ts) => Box::new(ts.iter().flat_map(|t| t.fv_impl())), Term::TupleProj(t, _) => t.fv_impl(), Term::DatatypeCtor(_, _, args) => Box::new(args.iter().flat_map(|t| t.fv_impl())), - Term::DatatypeDiscr(_, t) => t.fv_impl(), + Term::DatatypeDiscr(_, t) | Term::IntToBitVec { term: t, .. } => t.fv_impl(), } } @@ -711,6 +775,7 @@ impl Term { match sort { Sort::Null => Term::Null, Sort::Int => Term::int(0), + Sort::BitVec { width } => Term::int(0).int_to_bit_vec(*width), Sort::Bool => Term::Bool(false), Sort::String => Term::String(String::new()), Sort::Box(s) => Term::Box(Box::new(Self::default_for(s))), @@ -908,6 +973,13 @@ impl Term { Term::DatatypeCtor(DatatypeSort::new(d_sym, d_args), c_sym, args) } + pub fn int_to_bit_vec(self, width: u32) -> Self { + Term::IntToBitVec { + width, + term: Box::new(self), + } + } + pub fn datatype_discr(d_sym: DatatypeSymbol, t: Term) -> Self { Term::DatatypeDiscr(d_sym, Box::new(t)) } diff --git a/src/chc/format_context.rs b/src/chc/format_context.rs index 3803f2d9..8c6744b9 100644 --- a/src/chc/format_context.rs +++ b/src/chc/format_context.rs @@ -65,7 +65,9 @@ fn term_sorts( term_sorts(var_sorts, arg, sorts); } } - chc::Term::DatatypeDiscr(_, t) => term_sorts(var_sorts, t, sorts), + chc::Term::DatatypeDiscr(_, t) | chc::Term::IntToBitVec { term: t, .. } => { + term_sorts(var_sorts, t, sorts) + } chc::Term::UserQuantifiedVar(_, _) => {} } } @@ -117,6 +119,7 @@ impl<'a> std::fmt::Display for SortSymbol<'a> { match self.inner { chc::Sort::Null => write!(f, "Null"), chc::Sort::Int => write!(f, "Int"), + chc::Sort::BitVec { width } => write!(f, "BitVec{}", width), chc::Sort::Bool => write!(f, "Bool"), chc::Sort::String => write!(f, "String"), chc::Sort::Param(i) => write!(f, "T{}", i), @@ -424,6 +427,7 @@ impl FormatContext { fn fmt_sort_impl(&self, sort: &chc::Sort) -> Box { match sort { + chc::Sort::BitVec { width } => Box::new(format!("(_ BitVec {})", width)), chc::Sort::Seq(elem) => Box::new(format!("(Seq {})", self.fmt_sort(elem))), chc::Sort::Array(s1, s2) => { let s1 = self.fmt_sort(s1); diff --git a/src/chc/smtlib2.rs b/src/chc/smtlib2.rs index f5aa46a2..76dc584e 100644 --- a/src/chc/smtlib2.rs +++ b/src/chc/smtlib2.rs @@ -236,6 +236,14 @@ impl<'ctx, 'a> std::fmt::Display for Term<'ctx, 'a> { Term::new(self.ctx, self.var_sorts, t) ) } + chc::Term::IntToBitVec { width, term } => { + write!( + f, + "((_ int_to_bv {}) {})", + width, + Term::new(self.ctx, self.var_sorts, term) + ) + } chc::Term::UserQuantifiedVar(_, var) => write!(f, "{}", var), } } diff --git a/src/chc/unbox.rs b/src/chc/unbox.rs index 06ca468c..c504181c 100644 --- a/src/chc/unbox.rs +++ b/src/chc/unbox.rs @@ -21,6 +21,10 @@ fn unbox_term(term: Term) -> Term { args.into_iter().map(unbox_term).collect(), ), Term::DatatypeDiscr(sym, arg) => Term::DatatypeDiscr(sym, Box::new(unbox_term(*arg))), + Term::IntToBitVec { width, term } => Term::IntToBitVec { + width, + term: Box::new(unbox_term(*term)), + }, Term::UserQuantifiedVar(sort, var) => Term::UserQuantifiedVar(unbox_sort(sort), var), } } @@ -64,6 +68,7 @@ fn unbox_sort(sort: Sort) -> Sort { match sort { Sort::Null => Sort::Null, Sort::Int => Sort::Int, + Sort::BitVec { width } => Sort::BitVec { width }, Sort::Bool => Sort::Bool, Sort::String => Sort::String, Sort::Param(i) => Sort::Param(i), diff --git a/src/refine/template.rs b/src/refine/template.rs index fd471a26..f0d85fc5 100644 --- a/src/refine/template.rs +++ b/src/refine/template.rs @@ -168,6 +168,23 @@ impl<'tcx> TypeBuilder<'tcx> { ty } + /// Builds the type of `thrust_models::model::BitVec` from its generic `args`. + fn build_bit_vec(&self, args: mir_ty::GenericArgsRef<'tcx>) -> rty::BitVecType { + let width = args + .const_at(0) + .try_to_target_usize(self.tcx) + .expect("BitVec width must be a known constant"); + let signed = args + .const_at(1) + .try_to_value() + .and_then(|value| value.try_to_bool()) + .expect("BitVec signedness must be a known constant"); + rty::BitVecType { + width: width.try_into().unwrap(), + signed, + } + } + // TODO: consolidate two impls fn model_adt( &self, @@ -178,6 +195,10 @@ impl<'tcx> TypeBuilder<'tcx> { return Some(rty::Type::int()); } + if Some(adt.did()) == self.def_ids.bit_vec_model() { + return Some(rty::Type::BitVec(self.build_bit_vec(args))); + } + if Some(adt.did()) == self.def_ids.mut_model() { let elem_ty = self.build(args.type_at(0)); return Some(rty::PointerType::mut_to(elem_ty).into()); @@ -372,6 +393,10 @@ where return Some(rty::Type::int()); } + if Some(adt.did()) == self.inner.def_ids.bit_vec_model() { + return Some(rty::Type::BitVec(self.inner.build_bit_vec(args))); + } + if Some(adt.did()) == self.inner.def_ids.mut_model() { let elem_ty = self.build(args.type_at(0)); return Some(rty::PointerType::mut_to(elem_ty).into()); diff --git a/src/rty.rs b/src/rty.rs index 51ea99eb..5a193273 100644 --- a/src/rty.rs +++ b/src/rty.rs @@ -925,10 +925,26 @@ impl ArrayType { } } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct BitVecType { + pub width: u32, + pub signed: bool, +} + +impl<'a, D> Pretty<'a, D, termcolor::ColorSpec> for &BitVecType +where + D: pretty::DocAllocator<'a, termcolor::ColorSpec>, +{ + fn pretty(self, allocator: &'a D) -> pretty::DocBuilder<'a, D, termcolor::ColorSpec> { + allocator.text(format!("BitVec<{}, {}>", self.width, self.signed)) + } +} + /// An underlying type of a refinement type. #[derive(Debug, Clone)] pub enum Type { Int, + BitVec(BitVecType), Bool, String, Never, @@ -986,6 +1002,7 @@ where fn pretty(self, allocator: &'a D) -> pretty::DocBuilder<'a, D, termcolor::ColorSpec> { match self { Type::Int => allocator.text("int"), + Type::BitVec(ty) => ty.pretty(allocator), Type::Bool => allocator.text("bool"), Type::String => allocator.text("string"), Type::Never => allocator.text("!"), @@ -1120,6 +1137,7 @@ impl Type { pub fn to_sort(&self) -> chc::Sort { match self { Type::Int => chc::Sort::int(), + Type::BitVec(ty) => chc::Sort::bit_vec(ty.width), Type::Bool => chc::Sort::bool(), // TODO: enable string reasoning // currently String sort seems not available in HORN logic of Z3 @@ -1160,6 +1178,7 @@ impl Type { { match self { Type::Int => Type::Int, + Type::BitVec(ty) => Type::BitVec(ty), Type::Bool => Type::Bool, Type::String => Type::String, Type::Never => Type::Never, @@ -1179,6 +1198,7 @@ impl Type { { match self { Type::Int => Type::Int, + Type::BitVec(ty) => Type::BitVec(ty), Type::Bool => Type::Bool, Type::String => Type::String, Type::Never => Type::Never, @@ -1199,6 +1219,7 @@ impl Type { pub fn strip_refinement(self) -> Type { match self { Type::Int => Type::Int, + Type::BitVec(ty) => Type::BitVec(ty), Type::Bool => Type::Bool, Type::String => Type::String, Type::Never => Type::Never, @@ -1214,7 +1235,9 @@ impl Type { pub fn free_ty_params(&self) -> HashSet { match self { - Type::Int | Type::Bool | Type::String | Type::Never => Default::default(), + Type::Int | Type::BitVec(_) | Type::Bool | Type::String | Type::Never => { + Default::default() + } Type::Param(ty) => std::iter::once(ty.index()).collect(), Type::Pointer(ty) => ty.free_ty_params(), Type::Function(ty) => ty.free_ty_params(), @@ -1825,7 +1848,7 @@ impl RefinedType { { self.refinement.subst_ty_params_in_sorts(subst); match &mut self.ty { - Type::Int | Type::Bool | Type::String | Type::Never => {} + Type::Int | Type::BitVec(_) | Type::Bool | Type::String | Type::Never => {} Type::Param(ty) => { if let Some(rty) = subst.get(ty.index()) { let RefinedType { @@ -1861,6 +1884,7 @@ impl RefinedType { { match (self.ty, other.ty) { (Type::Int, Type::Int) + | (Type::BitVec(_), Type::BitVec(_)) | (Type::Bool, Type::Bool) | (Type::String, Type::String) | (Type::Never, Type::Never) => Default::default(), @@ -1899,7 +1923,11 @@ impl RefinedType { /// Substitutes type parameters in a sort. fn subst_ty_params_in_sort(sort: &mut chc::Sort, subst: &TypeParamSubst) { match sort { - chc::Sort::Null | chc::Sort::Int | chc::Sort::Bool | chc::Sort::String => {} + chc::Sort::Null + | chc::Sort::Int + | chc::Sort::BitVec { .. } + | chc::Sort::Bool + | chc::Sort::String => {} chc::Sort::Param(idx) => { let type_param_idx = TypeParamIdx::from_usize(*idx); if let Some(rty) = subst.get(type_param_idx) { @@ -1986,7 +2014,8 @@ fn subst_ty_params_in_term(term: &mut chc::Term, subst: &TypeParamSubst | chc::Term::MutCurrent(t) | chc::Term::MutFinal(t) | chc::Term::TupleProj(t, _) - | chc::Term::DatatypeDiscr(_, t) => { + | chc::Term::DatatypeDiscr(_, t) + | chc::Term::IntToBitVec { term: t, .. } => { subst_ty_params_in_term(t, subst); } chc::Term::Mut(t1, t2) => { diff --git a/src/rty/subtyping.rs b/src/rty/subtyping.rs index ec7e3d91..1340c9b9 100644 --- a/src/rty/subtyping.rs +++ b/src/rty/subtyping.rs @@ -102,6 +102,7 @@ where | (Type::Bool, Type::Bool) | (Type::String, Type::String) | (Type::Never, Type::Never) => {} + (Type::BitVec(got), Type::BitVec(expected)) if got == expected => {} (Type::Enum(got), Type::Enum(expected)) if got.symbol() == expected.symbol() => { for (got_ty, expected_ty) in got.args.iter().zip(expected.args.iter()) { let cs = relate_refined_type(scope, got_ty, expected_ty, relation); diff --git a/std.rs b/std.rs index 36730f95..a625d209 100644 --- a/std.rs +++ b/std.rs @@ -62,6 +62,83 @@ mod thrust_models { } } + /// An SMT-LIB bit-vector of `WIDTH` bits. `SIGNED` selects the signed or unsigned variant + /// of the operations that have both. + #[thrust::def::bit_vec_model] + pub struct BitVec; + + impl BitVec { + #[allow(dead_code)] + #[thrust::def::bit_vec_from_int] + #[thrust::ignored] + pub fn from_int(_n: U) -> Self where U: super::Model { + unimplemented!() + } + + #[allow(dead_code)] + #[thrust::def::bit_vec_to_int] + #[thrust::ignored] + pub fn to_int(self) -> Int { + unimplemented!() + } + } + + impl PartialEq for BitVec { + #[thrust::ignored] + fn eq(&self, _other: &Self) -> bool { + unimplemented!() + } + } + + impl PartialOrd for BitVec { + #[thrust::ignored] + fn partial_cmp(&self, _other: &Self) -> Option { + unimplemented!() + } + } + + macro_rules! bit_vec_binary_op { + ($Trait:ident, $method:ident) => { + impl std::ops::$Trait + for BitVec + { + type Output = Self; + + #[thrust::ignored] + fn $method(self, _rhs: Self) -> Self::Output { + unimplemented!() + } + } + }; + } + + bit_vec_binary_op!(Add, add); + bit_vec_binary_op!(Sub, sub); + bit_vec_binary_op!(Mul, mul); + bit_vec_binary_op!(BitAnd, bitand); + bit_vec_binary_op!(BitOr, bitor); + bit_vec_binary_op!(BitXor, bitxor); + bit_vec_binary_op!(Shl, shl); + bit_vec_binary_op!(Shr, shr); + + impl std::ops::Not for BitVec { + type Output = Self; + + #[thrust::ignored] + fn not(self) -> Self::Output { + unimplemented!() + } + } + + impl std::ops::Neg for BitVec { + type Output = Self; + + #[thrust::ignored] + fn neg(self) -> Self::Output { + unimplemented!() + } + } + #[thrust::def::mut_model] pub struct Mut(PhantomData); @@ -276,6 +353,10 @@ mod thrust_models { type Ty = model::Int; } + impl Model for model::BitVec { + type Ty = model::BitVec; + } + macro_rules! int_model { ($T:ty) => { impl Model for $T { diff --git a/tests/ui/fail/bit_vec_set.rs b/tests/ui/fail/bit_vec_set.rs new file mode 100644 index 00000000..f53bdbd2 --- /dev/null +++ b/tests/ui/fail/bit_vec_set.rs @@ -0,0 +1,45 @@ +//@error-in-other-file: Unsat +//@compile-flags: -C debug-assertions=off + +use thrust_models::model::BitVec; + +struct BitSet64 { + bits: u64, +} + +impl thrust_models::Model for BitSet64 { + type Ty = BitVec<64, false>; +} + +#[thrust_macros::context] +impl BitSet64 { + #[thrust::trusted] + #[thrust_macros::ensures(result == BitVec::from_int(0))] + fn new() -> Self { + BitSet64 { bits: 0 } + } + + #[thrust::trusted] + #[thrust_macros::requires(i < 64)] + #[thrust_macros::ensures(!self == *self | (BitVec::from_int(1) << BitVec::from_int(i)))] + fn insert(&mut self, i: usize) { + self.bits |= 1 << i; + } + + #[thrust::trusted] + #[thrust_macros::requires(i < 64)] + #[thrust_macros::ensures( + result == ((*self >> BitVec::from_int(i)) & BitVec::from_int(1) == BitVec::from_int(1)) + )] + fn contains(&self, i: usize) -> bool { + self.bits & (1 << i) != 0 + } +} + +fn main() { + let mut set = BitSet64::new(); + set.insert(3); + set.insert(5); + assert!(set.contains(3)); + assert!(set.contains(4)); +} diff --git a/tests/ui/fail/bit_vec_wrapping.rs b/tests/ui/fail/bit_vec_wrapping.rs new file mode 100644 index 00000000..9209c276 --- /dev/null +++ b/tests/ui/fail/bit_vec_wrapping.rs @@ -0,0 +1,36 @@ +//@error-in-other-file: Unsat +//@compile-flags: -C debug-assertions=off + +use thrust_models::model::BitVec; + +#[derive(Clone, Copy)] +struct WrappingI32(i32); + +impl thrust_models::Model for WrappingI32 { + type Ty = BitVec<32, true>; +} + +#[thrust::trusted] +#[thrust_macros::ensures(result.to_int() == x)] +fn wrap(x: i32) -> WrappingI32 { + WrappingI32(x) +} + +#[thrust::trusted] +#[thrust_macros::ensures(result == x + y)] +fn add(x: WrappingI32, y: WrappingI32) -> WrappingI32 { + WrappingI32(x.0.wrapping_add(y.0)) +} + +#[thrust::trusted] +#[thrust_macros::ensures(result == (x < y))] +fn lt(x: WrappingI32, y: WrappingI32) -> bool { + x.0 < y.0 +} + +fn main() { + let max = wrap(i32::MAX); + let one = wrap(1); + let sum = add(max, one); + assert!(!lt(sum, max)); +} diff --git a/tests/ui/pass/bit_vec_set.rs b/tests/ui/pass/bit_vec_set.rs new file mode 100644 index 00000000..27aa8aaa --- /dev/null +++ b/tests/ui/pass/bit_vec_set.rs @@ -0,0 +1,45 @@ +//@check-pass +//@compile-flags: -C debug-assertions=off + +use thrust_models::model::BitVec; + +struct BitSet64 { + bits: u64, +} + +impl thrust_models::Model for BitSet64 { + type Ty = BitVec<64, false>; +} + +#[thrust_macros::context] +impl BitSet64 { + #[thrust::trusted] + #[thrust_macros::ensures(result == BitVec::from_int(0))] + fn new() -> Self { + BitSet64 { bits: 0 } + } + + #[thrust::trusted] + #[thrust_macros::requires(i < 64)] + #[thrust_macros::ensures(!self == *self | (BitVec::from_int(1) << BitVec::from_int(i)))] + fn insert(&mut self, i: usize) { + self.bits |= 1 << i; + } + + #[thrust::trusted] + #[thrust_macros::requires(i < 64)] + #[thrust_macros::ensures( + result == ((*self >> BitVec::from_int(i)) & BitVec::from_int(1) == BitVec::from_int(1)) + )] + fn contains(&self, i: usize) -> bool { + self.bits & (1 << i) != 0 + } +} + +fn main() { + let mut set = BitSet64::new(); + set.insert(3); + set.insert(5); + assert!(set.contains(3)); + assert!(!set.contains(4)); +} diff --git a/tests/ui/pass/bit_vec_wrapping.rs b/tests/ui/pass/bit_vec_wrapping.rs new file mode 100644 index 00000000..d4ced040 --- /dev/null +++ b/tests/ui/pass/bit_vec_wrapping.rs @@ -0,0 +1,36 @@ +//@check-pass +//@compile-flags: -C debug-assertions=off + +use thrust_models::model::BitVec; + +#[derive(Clone, Copy)] +struct WrappingI32(i32); + +impl thrust_models::Model for WrappingI32 { + type Ty = BitVec<32, true>; +} + +#[thrust::trusted] +#[thrust_macros::ensures(result.to_int() == x)] +fn wrap(x: i32) -> WrappingI32 { + WrappingI32(x) +} + +#[thrust::trusted] +#[thrust_macros::ensures(result == x + y)] +fn add(x: WrappingI32, y: WrappingI32) -> WrappingI32 { + WrappingI32(x.0.wrapping_add(y.0)) +} + +#[thrust::trusted] +#[thrust_macros::ensures(result == (x < y))] +fn lt(x: WrappingI32, y: WrappingI32) -> bool { + x.0 < y.0 +} + +fn main() { + let max = wrap(i32::MAX); + let one = wrap(1); + let sum = add(max, one); + assert!(lt(sum, max)); +}