Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions src/analyze/annot.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
Expand Down
77 changes: 76 additions & 1 deletion src/analyze/annot_fn.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<rty::BitVecType> {
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<rty::FunctionParamIdx> {
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);
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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);
Expand All @@ -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))
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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();
Expand Down
24 changes: 24 additions & 0 deletions src/analyze/did_cache.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ struct DefIds {

model_ty: OnceCell<Option<DefId>>,
int_model: OnceCell<Option<DefId>>,
bit_vec_model: OnceCell<Option<DefId>>,
mut_model: OnceCell<Option<DefId>>,
box_model: OnceCell<Option<DefId>>,
array_model: OnceCell<Option<DefId>>,
Expand All @@ -23,6 +24,8 @@ struct DefIds {
mut_model_new: OnceCell<Option<DefId>>,
box_model_new: OnceCell<Option<DefId>>,
array_model_store: OnceCell<Option<DefId>>,
bit_vec_from_int: OnceCell<Option<DefId>>,
bit_vec_to_int: OnceCell<Option<DefId>>,

seq_model: OnceCell<Option<DefId>>,
seq_empty: OnceCell<Option<DefId>>,
Expand Down Expand Up @@ -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<DefId> {
*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<DefId> {
*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<DefId> {
*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<DefId> {
*self
.def_ids
Expand Down
76 changes: 74 additions & 2 deletions src/chc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,7 @@ impl DatatypeSort {
pub enum Sort {
Null,
Int,
BitVec { width: u32 },
Bool,
String,
Param(usize),
Expand All @@ -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)),
Expand Down Expand Up @@ -185,7 +187,12 @@ impl Sort {
fn walk_impl<'a, 'b>(&'a self, mut f: Box<dyn FnMut(&'a Sort) + 'b>) {
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 {
Expand Down Expand Up @@ -234,6 +241,10 @@ impl Sort {
Sort::Int
}

pub fn bit_vec(width: u32) -> Self {
Sort::BitVec { width }
}

pub fn string() -> Self {
Sort::String
}
Expand Down Expand Up @@ -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()
}
Expand Down Expand Up @@ -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.++");
Expand Down Expand Up @@ -483,6 +535,10 @@ pub enum Term<V = TermVarIdx> {
TupleProj(Box<Term<V>>, usize),
DatatypeCtor(DatatypeSort, DatatypeSymbol, Vec<Term<V>>),
DatatypeDiscr(DatatypeSymbol, Box<Term<V>>),
IntToBitVec {
width: u32,
term: Box<Term<V>>,
},
/// Used in [`Formula`] to represent quantified variables appearing in annotations.
UserQuantifiedVar(Sort, UserQuantifiedVarId),
}
Expand Down Expand Up @@ -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),
}
}
Expand Down Expand Up @@ -605,6 +664,10 @@ impl<V> Term<V> {
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),
}
}
Expand Down Expand Up @@ -653,6 +716,7 @@ impl<V> Term<V> {
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(),
}
}
Expand All @@ -676,7 +740,7 @@ impl<V> Term<V> {
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(),
}
}

Expand Down Expand Up @@ -711,6 +775,7 @@ impl<V> Term<V> {
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))),
Expand Down Expand Up @@ -908,6 +973,13 @@ impl<V> Term<V> {
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<V>) -> Self {
Term::DatatypeDiscr(d_sym, Box::new(t))
}
Expand Down
6 changes: 5 additions & 1 deletion src/chc/format_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(_, _) => {}
}
}
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -424,6 +427,7 @@ impl FormatContext {

fn fmt_sort_impl(&self, sort: &chc::Sort) -> Box<dyn std::fmt::Display> {
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);
Expand Down
Loading
Loading