Skip to content
Draft
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
46 changes: 46 additions & 0 deletions src/analyze.rs
Original file line number Diff line number Diff line change
Expand Up @@ -235,6 +235,8 @@ pub struct Analyzer<'tcx> {

/// Collection of functions with `#[thrust::formula_fn]` attribute.
formula_fns: HashMap<LocalDefId, DeferredFormulaFnDef<'tcx>>,
predicate_instances:
Rc<RefCell<HashMap<(DefId, mir_ty::GenericArgsRef<'tcx>), chc::UserDefinedPred>>>,

/// Resulting CHC system.
system: Rc<RefCell<chc::System>>,
Expand Down Expand Up @@ -274,6 +276,7 @@ impl<'tcx> Analyzer<'tcx> {
tcx,
defs,
formula_fns,
predicate_instances: Default::default(),
system,
basic_blocks,
def_ids: did_cache::DefIdCache::new(tcx),
Expand Down Expand Up @@ -446,6 +449,49 @@ impl<'tcx> Analyzer<'tcx> {
def_ty.ty.as_function().cloned()
}

pub fn predicate_with_args(
&self,
def_id: DefId,
generic_args: mir_ty::GenericArgsRef<'tcx>,
) -> chc::UserDefinedPred {
let Some(local_def_id) = def_id
.as_local()
.filter(|id| self.formula_fns.contains_key(id))
else {
return refine::user_defined_pred(self.tcx, def_id);
};
let key = (def_id, generic_args);
if let Some(pred) = self.predicate_instances.borrow().get(&key) {
return pred.clone();
}
let pred = chc::UserDefinedPred::new(format!(
"{}_{}",
refine::user_defined_pred(self.tcx, def_id),
self.predicate_instances.borrow().len(),
));
self.predicate_instances
.borrow_mut()
.insert(key, pred.clone());

let formula_fn = self
.formula_fn_with_args(local_def_id, generic_args)
.unwrap();
let type_builder = TypeBuilder::new(self.tcx, self.def_ids(), def_id);
let arg_sorts = formula_fn
.params()
.iter()
.map(|ty| type_builder.build(*ty).to_sort())
.collect();
let formula = formula_fn
.formula()
.clone()
.map_var(|idx| chc::TermVarIdx::from(idx.index()));
self.system
.borrow_mut()
.push_pred_define_formula(pred.clone(), arg_sorts, formula);
pred
}

pub fn formula_fn_with_args(
&self,
local_def_id: LocalDefId,
Expand Down
10 changes: 5 additions & 5 deletions src/analyze/annot_fn.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ use rustc_middle::ty::{self as mir_ty, TyCtxt};

use crate::analyze::{self, did_cache::DefIdCache};
use crate::chc;
use crate::refine::{self, TypeBuilder};
use crate::refine::TypeBuilder;
use crate::rty;

#[derive(Debug, Clone)]
Expand Down Expand Up @@ -981,12 +981,12 @@ impl<'a, 'tcx> AnnotFnTranslator<'a, 'tcx> {
generic_args,
)
.unwrap();
let pred_def_id = if let Some(instance) = instance {
instance.def_id()
let (pred_def_id, pred_args) = if let Some(instance) = instance {
(instance.def_id(), instance.args)
} else {
def_id
(def_id, generic_args)
};
let pred = refine::user_defined_pred(self.tcx, pred_def_id);
let pred = self.analyzer.predicate_with_args(pred_def_id, pred_args);
let arg_terms = args.iter().map(|e| self.to_term(e)).collect();
let atom = chc::Atom::new(pred.into(), arg_terms);
return FormulaOrTerm::Formula(chc::Formula::Atom(atom));
Expand Down
18 changes: 11 additions & 7 deletions src/analyze/crate_.rs
Original file line number Diff line number Diff line change
Expand Up @@ -88,13 +88,17 @@ impl<'tcx, 'ctx> Analyzer<'tcx, 'ctx> {
self.skip_analysis.insert(*local_def_id);
keys.swap_remove(local_def_id);
}
if analyzer.is_annotated_as_predicate() {
analyzer.analyze_predicate_definition();
self.skip_analysis.insert(*local_def_id);
keys.swap_remove(local_def_id);
}
if analyzer.is_annotated_as_formula_fn() {
self.ctx.register_formula_fn(*local_def_id);
let is_predicate = analyzer.is_annotated_as_predicate();
let is_formula_fn = analyzer.is_annotated_as_formula_fn();
if is_predicate || is_formula_fn {
if is_formula_fn {
self.ctx.register_formula_fn(*local_def_id);
}
if is_predicate && !is_formula_fn {
self.ctx
.local_def_analyzer(*local_def_id)
.analyze_predicate_definition();
}
self.skip_analysis.insert(*local_def_id);
keys.swap_remove(local_def_id);
}
Expand Down
12 changes: 6 additions & 6 deletions src/analyze/local_def.rs
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,12 @@ impl<'tcx, 'ctx> Analyzer<'tcx, 'ctx> {
}

fn define_as_predicate(&self, pred: chc::UserDefinedPred) {
let sig = self.ctx.fn_sig(self.local_def_id.to_def_id());
let arg_sorts = sig
.inputs()
.iter()
.map(|input_ty| self.type_builder.build(*input_ty).to_sort());

// function's body
use rustc_hir::{Block, Expr, ExprKind};

Expand Down Expand Up @@ -89,12 +95,6 @@ impl<'tcx, 'ctx> Analyzer<'tcx, 'ctx> {
.to_string()
});

let sig = self.ctx.fn_sig(self.local_def_id.to_def_id());
let arg_sorts = sig
.inputs()
.iter()
.map(|input_ty| self.type_builder.build(*input_ty).to_sort());

let arg_name_and_sorts = arg_names.into_iter().zip(arg_sorts).collect::<Vec<_>>();

self.ctx.system.borrow_mut().push_pred_define(
Expand Down
64 changes: 57 additions & 7 deletions src/chc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1911,10 +1911,6 @@ impl Clause {
pub fn is_nop(&self) -> bool {
self.head.is_top() || self.body.is_bottom()
}

fn term_sort(&self, term: &Term<TermVarIdx>) -> Sort {
term.sort(|v| self.vars[*v].clone())
}
}

/// A command specified using `thrust::raw_command` attribute
Expand Down Expand Up @@ -1979,11 +1975,22 @@ pub struct PredVarDef {

pub type UserDefinedPredSig = Vec<(String, Sort)>;

/// The body of a user-defined predicate.
///
/// A predicate can be defined either by a raw SMT-LIB2 string (inserted into the
/// generated `define-fun` verbatim) or by a [`Formula`] translated from a Rust
/// expression via the `formula_fn` infrastructure.
#[derive(Debug, Clone)]
pub enum UserDefinedPredBody {
Raw(String),
Formula(Formula<TermVarIdx>),
}

#[derive(Debug, Clone)]
pub struct UserDefinedPredDef {
symbol: UserDefinedPred,
sig: UserDefinedPredSig,
body: String,
body: UserDefinedPredBody,
}

/// A CHC system.
Expand All @@ -1997,6 +2004,29 @@ pub struct System {
}

impl System {
fn user_defined_preds_in_dependency_order(&self) -> Vec<&UserDefinedPredDef> {
let mut remaining: Vec<_> = self.user_defined_pred_defs.iter().collect();
let mut ordered = Vec::with_capacity(remaining.len());
while !remaining.is_empty() {
let next = remaining
.iter()
.position(|def| match &def.body {
UserDefinedPredBody::Raw(_) => true,
UserDefinedPredBody::Formula(formula) => formula.iter_atoms().all(|atom| {
let Pred::UserDefined(pred) = &atom.pred else {
return true;
};
!remaining
.iter()
.any(|dependency| dependency.symbol == *pred)
}),
})
.expect("recursive predicate definitions are not supported");
ordered.push(remaining.remove(next));
}
ordered
}

pub fn new_pred_var(&mut self, sig: PredSig, debug_info: DebugInfo) -> PredVarId {
self.pred_vars.push(PredVarDef { sig, debug_info })
}
Expand All @@ -2011,8 +2041,28 @@ impl System {
sig: UserDefinedPredSig,
body: String,
) {
self.user_defined_pred_defs
.push(UserDefinedPredDef { symbol, sig, body })
self.user_defined_pred_defs.push(UserDefinedPredDef {
symbol,
sig,
body: UserDefinedPredBody::Raw(body),
})
}

pub fn push_pred_define_formula(
&mut self,
symbol: UserDefinedPred,
arg_sorts: IndexVec<TermVarIdx, Sort>,
formula: Formula<TermVarIdx>,
) {
let sig = arg_sorts
.into_iter_enumerated()
.map(|(var, sort)| (var.to_string(), sort))
.collect();
self.user_defined_pred_defs.push(UserDefinedPredDef {
symbol,
sig,
body: UserDefinedPredBody::Formula(formula),
})
}

pub fn push_clause(&mut self, clause: Clause) -> Option<ClauseId> {
Expand Down
2 changes: 1 addition & 1 deletion src/chc/format_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ pub struct FormatContext {

// FIXME: this is obviously ineffective and should be replaced
fn term_sorts(clause: &chc::Clause, t: &chc::Term, sorts: &mut BTreeSet<chc::Sort>) {
sorts.insert(clause.term_sort(t));
sorts.insert(t.sort(|v| clause.vars[*v].clone()));
match t {
chc::Term::Null => {}
chc::Term::Var(_) => {}
Expand Down
Loading
Loading