use crate::expr::*;
use crate::mc::bmc::start_bmc_or_pdr;
use crate::mc::{
ModelCheckResult, TransitionSystemEncoding, bmc, check_assuming, check_assuming_end,
get_smt_value,
};
use crate::smt::*;
use crate::system::TransitionSystem;
use baa::{BitVecOps, Value};
use std::collections::BinaryHeap;
type Step = u64;
const FROM_STEP: Step = 1;
const TO_STEP: Step = 2;
const MAX_FRAMES: usize = 1000;
#[derive(Debug, Clone)]
struct Cube {
literals: Vec<ExprRef>,
}
impl Cube {
const fn tru() -> Self {
Self { literals: vec![] }
}
fn to_expr(&self, ctx: &mut Context) -> ExprRef {
self.literals
.iter()
.copied()
.fold(ctx.get_true(), |acc, e| ctx.and(acc, e))
}
fn negate(&self, ctx: &mut Context) -> ExprRef {
self.literals
.iter()
.copied()
.fold(ctx.get_false(), |acc, e| {
let neg_lit = ctx.not(e);
ctx.or(acc, neg_lit)
})
}
}
type FrameId = usize;
#[derive(Debug, Clone)]
struct TimedCube {
cube: Cube,
frame: usize,
}
impl Eq for TimedCube {}
impl PartialEq for TimedCube {
fn eq(&self, other: &Self) -> bool {
self.frame.eq(&other.frame)
}
}
impl Ord for TimedCube {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other.frame.cmp(&self.frame)
}
}
impl PartialOrd for TimedCube {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
#[derive(Debug, Copy, Clone, PartialEq, Eq)]
enum RelIndType {
Standard,
Extended,
}
fn expr_at_step(
ctx: &mut Context,
enc: &impl TransitionSystemEncoding,
expr: ExprRef,
step: Step,
) -> ExprRef {
simple_transform_expr(ctx, expr, |ctx, e, _| {
if ctx[e].is_symbol() {
Some(enc.get_at(ctx, e, step))
} else {
None
}
})
}
fn extract_state_values(
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
sys: &TransitionSystem,
enc: &impl TransitionSystemEncoding,
step: Step,
) -> Result<Vec<(ExprRef, Value)>> {
let mut state_vals = Vec::with_capacity(sys.states.len());
for state in &sys.states {
let sym = enc.get_at(ctx, state.symbol, step);
state_vals.push((state.symbol, get_smt_value(ctx, smt_ctx, sym)?));
}
Ok(state_vals)
}
fn get_bit_level_cube(
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
sys: &TransitionSystem,
enc: &impl TransitionSystemEncoding,
step: Step,
) -> Result<Cube> {
let mut literals = Vec::new();
let vals = extract_state_values(ctx, smt_ctx, sys, enc, step)?;
assert_eq!(vals.len(), sys.states.len());
for (sym, val) in vals {
match val {
Value::BitVec(bv) => {
let width = bv.width();
for idx in 0..width {
let bit = ctx.slice(sym, idx, idx);
let bit_val = if bv.is_bit_set(idx) {
ctx.get_true()
} else {
ctx.get_false()
};
let lit = ctx.equal(bit, bit_val);
literals.push(lit);
}
}
Value::Array(_av) => todo!("Add array support"),
}
}
Ok(Cube { literals })
}
fn query(
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
sys: &TransitionSystem,
enc: &impl TransitionSystemEncoding,
assumptions: impl IntoIterator<Item = ExprRef>,
) -> Result<(CheckSatResponse, Option<Cube>)> {
let smt_res = check_assuming(ctx, smt_ctx, assumptions)?;
let model = if smt_res == CheckSatResponse::Sat {
let cex = get_bit_level_cube(ctx, smt_ctx, sys, enc, FROM_STEP)?;
Some(cex)
} else {
None
};
check_assuming_end(smt_ctx)?;
Ok((smt_res, model))
}
struct BasePdr {
init_frame: ExprRef,
frames: Vec<Vec<Cube>>,
}
impl BasePdr {
fn init(
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
enc: &impl TransitionSystemEncoding,
sys: &TransitionSystem,
) -> Result<Self> {
let mut init_cube = Cube::tru();
for state in &sys.states {
if let Some(init) = state.init {
let lit = ctx.equal(state.symbol, init);
init_cube.literals.push(lit);
}
}
let init_act = ctx.bv_symbol("act_0", 1);
let init_expr = init_cube.to_expr(ctx);
let init_expr = expr_at_step(ctx, enc, init_expr, FROM_STEP);
let imp = ctx.implies(init_act, init_expr);
smt_ctx.declare_const(ctx, init_act)?;
smt_ctx.assert(ctx, imp)?;
Ok(Self {
init_frame: init_act,
frames: vec![vec![]], })
}
fn frame_assumptions(
&self,
ctx: &mut Context,
enc: &impl TransitionSystemEncoding,
frame: FrameId,
) -> ExprRef {
if frame == 0 {
self.init_frame
} else {
assert!(frame < self.frames.len());
let expr = self.frames[frame].iter().fold(ctx.get_true(), |acc, cube| {
let clause = cube.negate(ctx);
ctx.and(acc, clause)
});
expr_at_step(ctx, enc, expr, FROM_STEP)
}
}
fn rel_ind(
&self,
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
sys: &TransitionSystem,
enc: &impl TransitionSystemEncoding,
cube: &TimedCube,
query_type: RelIndType,
) -> Result<(CheckSatResponse, Option<Cube>)> {
let mut assumptions = Vec::new();
let frame_assumption = self.frame_assumptions(ctx, enc, cube.frame - 1);
assumptions.push(frame_assumption);
let cube_expr = cube.cube.to_expr(ctx);
let cube_nxt = expr_at_step(ctx, enc, cube_expr, TO_STEP);
assumptions.push(cube_nxt);
if query_type == RelIndType::Extended {
let neg_cube_expr = cube.cube.negate(ctx);
let neg_cube_cur = expr_at_step(ctx, enc, neg_cube_expr, FROM_STEP);
assumptions.push(neg_cube_cur);
}
query(ctx, smt_ctx, sys, enc, assumptions)
}
fn add_blocked_cube(&mut self, cube: &TimedCube) {
let front = cube.frame;
for idx in 1..=front {
self.frames[idx].push(cube.cube.clone());
}
}
const fn frontier(&self) -> FrameId {
self.frames.len() - 1
}
fn get_bad_cube(
&self,
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
sys: &TransitionSystem,
enc: &impl TransitionSystemEncoding,
) -> Result<Option<Cube>> {
let front = self.frontier();
let bad_lits: Vec<ExprRef> = sys
.bad_states
.iter()
.map(|&b| expr_at_step(ctx, enc, b, FROM_STEP))
.collect();
let bad_expr = bad_lits
.iter()
.fold(ctx.get_false(), |acc, &b| ctx.or(acc, b));
let front_assumption = self.frame_assumptions(ctx, enc, front);
match query(ctx, smt_ctx, sys, enc, vec![front_assumption, bad_expr])? {
(CheckSatResponse::Sat, Some(cube)) => {
Ok(Some(cube))
}
(CheckSatResponse::Unsat, _) => Ok(None), (CheckSatResponse::Unknown, _) => Err(
Error::UnexpectedResponse(
"`get_bad_cube` in `BasePdr`".into(),
"unknown query".into(),
),
),
_ => unreachable!(),
}
}
fn add_frame(&mut self) {
self.frames.push(Vec::new());
}
fn block_cube(
&mut self,
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
sys: &TransitionSystem,
enc: &impl TransitionSystemEncoding,
cube: &TimedCube,
) -> Result<bool> {
let mut worklist = BinaryHeap::from([cube.clone()]);
while let Some(obj) = worklist.pop() {
if obj.frame == 0 {
return Ok(false);
}
let res = match self.rel_ind(ctx, smt_ctx, sys, enc, cube, RelIndType::Extended)? {
(CheckSatResponse::Sat, Some(cube)) => Some(cube), (CheckSatResponse::Unsat, _) => None,
(CheckSatResponse::Unknown, _) => {
return Err(Error::UnexpectedResponse(
"`block_cube` in `BasePdr`".into(),
"unknown query".into(),
));
}
_ => unreachable!(),
};
if let Some(wit) = res {
worklist.push(TimedCube {
cube: wit,
frame: obj.frame - 1,
});
worklist.push(obj);
} else {
self.add_blocked_cube(&obj);
}
}
Ok(true)
}
fn propagate_blocked_cubes(
&mut self,
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
sys: &TransitionSystem,
enc: &impl TransitionSystemEncoding,
) -> Result<bool> {
let front = self.frontier();
for idx in 1..front {
let mut num_left = self.frames[idx].len();
for cube_idx in 0..self.frames[idx].len() {
let cube = self.frames[idx][cube_idx].clone();
let query_cube = TimedCube {
cube: cube.clone(),
frame: idx + 1,
};
if self
.rel_ind(ctx, smt_ctx, sys, enc, &query_cube, RelIndType::Standard)?
.0
== CheckSatResponse::Unsat
{
self.frames[idx + 1].push(cube.clone());
num_left -= 1;
}
}
if num_left == 0 {
return Ok(true);
}
}
Ok(false)
}
}
pub fn pdr(
ctx: &mut Context,
smt_ctx: &mut impl SolverContext,
sys: &TransitionSystem,
) -> Result<ModelCheckResult> {
let mut enc = match start_bmc_or_pdr(ctx, smt_ctx, sys)? {
(r, None) => return Ok(r),
(_, Some(enc)) => enc,
};
assert!(sys.constraints.is_empty());
enc.init_at(ctx, smt_ctx, FROM_STEP)?;
enc.unroll(ctx, smt_ctx)?;
let mut state = BasePdr::init(ctx, smt_ctx, &enc, sys)?;
while state.frontier() <= MAX_FRAMES {
let bad_cube = state.get_bad_cube(ctx, smt_ctx, sys, &enc)?;
if let Some(bad) = bad_cube {
if !state.block_cube(
ctx,
smt_ctx,
sys,
&enc,
&TimedCube {
cube: bad,
frame: state.frontier(),
},
)? {
smt_ctx.restart()?;
let ModelCheckResult::Fail(wit) =
bmc(ctx, smt_ctx, sys, false, false, MAX_FRAMES as u64)?
else {
unreachable!()
};
return Ok(ModelCheckResult::Fail(wit));
}
} else {
state.add_frame();
if state.propagate_blocked_cubes(ctx, smt_ctx, sys, &enc)? {
return Ok(ModelCheckResult::Success);
}
}
}
Ok(ModelCheckResult::Unknown)
}