Skip to main content

rapx/verify/property_checker/
numeric.rs

1//! Checker for `ValidNum`: numeric-interval and predicate reasoning.
2//!
3//! Evaluates each `NumericPredicate` to a Z3 comparison and discharges it with
4//! `assert_all` plus Euclidean-division (NIA) axioms injected for both the
5//! contract expression and the VM's computed terms.
6
7use crate::helpers::mir_scan::Checkpoint;
8use crate::verify::api_classify;
9use crate::verify::contract::{
10    ContractExpr, NumericBinOp, PlaceBase, Property, PropertyArg, RelOp,
11};
12use crate::verify::def_use::PlaceKey;
13use crate::verify::report::{CheckResult, UnknownReason};
14use crate::verify::vm::state::VmState;
15use rustc_hash::FxHashSet;
16use rustc_middle::mir::Operand;
17use rustc_middle::ty::TyKind;
18use z3::{
19    SatResult, Solver,
20    ast::{Ast, Bool, Int},
21};
22
23use super::PropertyChecker;
24
25impl PropertyChecker {
26    pub(super) fn check_valid_num<'z3, 'tcx>(
27        &self,
28        vm_state: &VmState<'z3, 'tcx>,
29        solver: &Solver<'z3>,
30        checkpoint: &Checkpoint<'tcx>,
31        property: &Property<'tcx>,
32    ) -> CheckResult {
33        if let Some(PropertyArg::Predicates(predicates)) = property.args().first() {
34            if self.all_predicates_are_slice_size_invariant(vm_state, checkpoint, predicates) {
35                return CheckResult::ProvedByRule;
36            }
37            for pred in predicates {
38                if let Some(r) =
39                    self.eval_numeric_predicate(vm_state, solver, Some(checkpoint), pred)
40                {
41                    if !r.is_proved() {
42                        return r;
43                    }
44                }
45            }
46            return CheckResult::ProvedByRule;
47        }
48        let Some(value) = self.target_value(vm_state, checkpoint, property) else {
49            return CheckResult::Unknown(UnknownReason::Unimplemented);
50        };
51        let ty = Self::ty_arg(property, 1);
52        if let Some(ty) = ty {
53            let size_bits = vm_state.size_of_ty(ty) * 8;
54            if size_bits > 0 && size_bits < 128 {
55                if let TyKind::Int(_) = ty.kind() {
56                    let half = 1u128 << (size_bits - 1);
57                    let min = -(half as i128);
58                    let max = (half - 1) as i128;
59                    solver.push();
60                    let below = value.z3_term.lt(&Int::from_i64(vm_state.z3_ctx, min as i64));
61                    let above = value.z3_term.gt(&Int::from_i64(vm_state.z3_ctx, max as i64));
62                    solver.assert(&Bool::or(vm_state.z3_ctx, &[&below, &above]));
63                    let r = match solver.check() {
64                        SatResult::Unsat => CheckResult::ProvedBySmt,
65                        SatResult::Sat => CheckResult::Failed,
66                        _ => CheckResult::Unknown(UnknownReason::SmtTimeout),
67                    };
68                    solver.pop(1);
69                    return r;
70                }
71                let max = Int::from_u64(
72                    vm_state.z3_ctx,
73                    ((1u128 << size_bits) - 1).min(u64::MAX as u128) as u64,
74                );
75                solver.push();
76                solver.assert(&value.z3_term.gt(&max));
77                let r = match solver.check() {
78                    SatResult::Unsat => CheckResult::ProvedBySmt,
79                    SatResult::Sat => CheckResult::Failed,
80                    _ => CheckResult::Unknown(UnknownReason::SmtTimeout),
81                };
82                solver.pop(1);
83                return r;
84            }
85        }
86        CheckResult::ProvedByRule
87    }
88
89    /// If `expr` is a `SliceIndex` range parameter (e.g. `..n`), return its
90    /// *exclusive* end term, so `ValidNum(index < CAPACITY)` compares `n` (not
91    /// the opaque range value).
92    fn range_end_of_lhs<'z3, 'tcx>(
93        &self,
94        vm_state: &VmState<'z3, 'tcx>,
95        checkpoint: Option<&Checkpoint<'tcx>>,
96        expr: &ContractExpr<'tcx>,
97    ) -> Option<Int<'z3>> {
98        let cp = match expr {
99            ContractExpr::Place(cp) => cp,
100            _ => return None,
101        };
102        if !cp.projections.is_empty() {
103            return None;
104        }
105        let ck = checkpoint?;
106        let op: Option<&Operand<'tcx>> = match cp.base {
107            PlaceBase::Arg(n) => ck.args.get(n),
108            PlaceBase::Local(n) => super::util::local_param_operand(vm_state, ck, n),
109            _ => None,
110        };
111        let op = op?;
112        let end = self.extract_range_end(vm_state, op)?;
113        Some(end.z3_term.clone())
114    }
115
116    pub(super) fn eval_numeric_predicate<'z3, 'tcx>(
117        &self,
118        vm_state: &VmState<'z3, 'tcx>,
119        solver: &Solver<'z3>,
120        checkpoint: Option<&Checkpoint<'tcx>>,
121        pred: &crate::verify::contract::NumericPredicate<'tcx>,
122    ) -> Option<CheckResult> {
123        // For a `SliceIndex` range (e.g. `..n`), `ValidNum(index < CAPACITY)`
124        // must compare the range's *exclusive end* (`n <= CAPACITY`) rather than
125        // the opaque range value. Detect ranges and adjust `<` to `<=`.
126        let (lhs, is_range) = match self.range_end_of_lhs(vm_state, checkpoint, &pred.lhs) {
127            Some(end) => (end, true),
128            None => (
129                self.eval_contract_expr(vm_state, checkpoint, &pred.lhs)?,
130                false,
131            ),
132        };
133        let rhs = self.eval_contract_expr(vm_state, checkpoint, &pred.rhs)?;
134        let condition = match pred.op {
135            RelOp::Le => lhs.le(&rhs),
136            RelOp::Lt if is_range => lhs.le(&rhs),
137            RelOp::Lt => lhs.lt(&rhs),
138            RelOp::Ge => lhs.ge(&rhs),
139            RelOp::Gt => lhs.gt(&rhs),
140            RelOp::Eq => lhs._eq(&rhs),
141            RelOp::Ne => lhs._eq(&rhs).not(),
142        };
143        solver.push();
144        vm_state.assert_all(solver);
145        // Bridge the iter_ptr_offset (tracked by post_inc_start) to the
146        // predicate's LHS (typically the loop counter `i` in position).
147        // At the assert_unchecked(i < n) point, tracked_offset == i + 1
148        // because post_inc_start(1) was just called before the check.
149        for (_, (off, _)) in vm_state.constraints.term_caches.iter_ptr_offset.iter() {
150            let one = Int::from_u64(vm_state.z3_ctx, 1);
151            solver.assert(&off._eq(&Int::add(vm_state.z3_ctx, &[&lhs, &one])));
152        }
153        // For Iter/IterMut Le predicates with rhs computed from fields,
154        // inject a lower-bound: the field-based len is >= 1 when the
155        // struct's entry contract contains !self.is_empty().
156        // Without this, Z3 cannot deduce (end-ptr)/sz >= 1 from != 0.
157        if matches!(pred.op, RelOp::Le) {
158            if let Some(term) = self.try_get_iter_len_term(vm_state, &pred.rhs) {
159                if let Some(one) = lhs.as_u64().or(rhs.as_u64()) {
160                    if one == 1 {
161                        let one_term = Int::from_u64(vm_state.z3_ctx, 1);
162                        solver.assert(&term.ge(&one_term));
163                    }
164                }
165            }
166        }
167        // NIA helper: inject Euclidean division identity for div operands
168        // to help Z3 prove (X/N)*N <= X via X = (X/N)*N + X%N, X%N >= 0.
169        self.inject_nia_axioms(vm_state, solver, checkpoint, &pred.lhs);
170        self.inject_nia_axioms(vm_state, solver, checkpoint, &pred.rhs);
171        // Walk VM binary op sources for Div/Rem terms used in the
172        // predicate — these are not visible in the ContractExpr tree.
173        self.inject_vm_div_axioms(vm_state, solver, &pred.lhs);
174        self.inject_vm_div_axioms(vm_state, solver, &pred.rhs);
175        // For Ne(pred != 0): assert path conditions so that layout
176        // constants (e.g. align_of >= 1) constrain the solver,
177        // while regular parameters remain unconstrained.
178        if matches!(pred.op, RelOp::Ne) && rhs.as_u64() == Some(0) {
179            // Concrete 0 from compiler-evaluated AlignOf/SizeOf for
180            // generic types — semantically always >= 1 for non-ZST.
181            if lhs.as_u64() == Some(0) {
182                return Some(CheckResult::ProvedByRule);
183            }
184            vm_state.assert_all(solver);
185        }
186        solver.assert(&condition.not());
187        let r0 = solver.check();
188        let mut r = match r0 {
189            SatResult::Unsat => Some(CheckResult::ProvedBySmt),
190            SatResult::Sat => Some(CheckResult::Failed),
191            _ => None,
192        };
193        // If Le/Ge check failed, inject a path-condition-level NIA axiom
194        // based on the VM's binary op sources and retry.
195        if matches!(r, Some(CheckResult::Failed))
196            && matches!(pred.op, RelOp::Le | RelOp::Ge | RelOp::Lt | RelOp::Gt)
197        {
198            solver.pop(1);
199            solver.push();
200            vm_state.assert_all(solver);
201            self.inject_nia_axioms(vm_state, solver, checkpoint, &pred.lhs);
202            self.inject_nia_axioms(vm_state, solver, checkpoint, &pred.rhs);
203            self.inject_vm_div_axioms(vm_state, solver, &pred.lhs);
204            self.inject_vm_div_axioms(vm_state, solver, &pred.rhs);
205            solver.assert(&condition.not());
206            r = match solver.check() {
207                SatResult::Unsat => Some(CheckResult::ProvedBySmt),
208                SatResult::Sat => Some(CheckResult::Failed),
209                _ => r,
210            };
211        }
212        solver.pop(1);
213        r
214    }
215
216    pub(super) fn inject_nia_axioms<'z3, 'tcx>(
217        &self,
218        vm_state: &VmState<'z3, 'tcx>,
219        solver: &Solver<'z3>,
220        checkpoint: Option<&Checkpoint<'tcx>>,
221        expr: &ContractExpr<'tcx>,
222    ) {
223        match expr {
224            ContractExpr::Binary {
225                op: NumericBinOp::Div,
226                lhs,
227                rhs,
228            } => {
229                if let (Some(l), Some(r)) = (
230                    self.eval_contract_expr(vm_state, checkpoint, lhs),
231                    self.eval_contract_expr(vm_state, checkpoint, rhs),
232                ) {
233                    let zero = Int::from_u64(vm_state.z3_ctx, 0);
234                    let mul_term = Int::mul(vm_state.z3_ctx, &[&l.div(&r), &r]);
235                    let rem_term = l.rem(&r);
236                    let sum_term = Int::add(vm_state.z3_ctx, &[&mul_term, &rem_term]);
237                    solver.assert(&l._eq(&sum_term));
238                    solver.assert(&rem_term.ge(&zero));
239                }
240            }
241            ContractExpr::Binary {
242                op: NumericBinOp::Mul,
243                lhs,
244                rhs,
245            } => {
246                // Recurse into mul operands in case one is a div
247                self.inject_nia_axioms(vm_state, solver, checkpoint, lhs);
248                self.inject_nia_axioms(vm_state, solver, checkpoint, rhs);
249            }
250            ContractExpr::Binary { lhs, rhs, .. } => {
251                self.inject_nia_axioms(vm_state, solver, checkpoint, lhs);
252                self.inject_nia_axioms(vm_state, solver, checkpoint, rhs);
253            }
254            ContractExpr::Unary { expr: inner, .. } => {
255                self.inject_nia_axioms(vm_state, solver, checkpoint, inner);
256            }
257            _ => {}
258        }
259    }
260
261    pub(super) fn inject_vm_div_axioms<'z3, 'tcx>(
262        &self,
263        vm_state: &VmState<'z3, 'tcx>,
264        solver: &Solver<'z3>,
265        expr: &ContractExpr<'tcx>,
266    ) {
267        let Some(val) = self.eval_contract_expr(vm_state, None, expr) else {
268            return;
269        };
270        self.inject_div_axioms_for_term(vm_state, solver, &val, 4);
271    }
272
273    pub(super) fn inject_div_axioms_for_term<'z3, 'tcx>(
274        &self,
275        vm_state: &VmState<'z3, 'tcx>,
276        solver: &Solver<'z3>,
277        target: &Int<'z3>,
278        depth: usize,
279    ) {
280        if depth == 0 {
281            return;
282        }
283
284        // Walk binary-op producers (recorded on the value) for destinations
285        // whose term matches target.
286        let op_sources: Vec<(Option<PlaceKey>, Option<PlaceKey>)> = {
287            let mut src: Vec<(Option<PlaceKey>, Option<PlaceKey>)> = Vec::new();
288            for (_, val) in vm_state.all_local_values() {
289                if let Some((lhs, rhs, _)) = val.source.operands() {
290                    if val.z3_term == *target {
291                        src.push((lhs.clone(), rhs.clone()));
292                    }
293                }
294            }
295            src
296        };
297
298        let mut already_seen = FxHashSet::default();
299
300        // ── Also recurse through Use / Cast chains: search ALL locals
301        // whose term equals target, and for each binary-op entry
302        // that *consumes* that local as an operand, walk the destination.
303        for local_idx in 0..vm_state.body().local_decls.len() {
304            let local = rustc_middle::mir::Local::from_usize(local_idx);
305            let Some(val) = vm_state.local_value(local) else {
306                continue;
307            };
308            if val.z3_term != *target {
309                continue;
310            }
311
312            for (_, dest_val) in vm_state.all_local_values() {
313                let Some((lhs, rhs, _)) = dest_val.source.operands() else {
314                    continue;
315                };
316                let lhs_local = lhs.as_ref().and_then(|pk| pk.local());
317                let rhs_local = rhs.as_ref().and_then(|pk| pk.local());
318                if (lhs_local == Some(local) || rhs_local == Some(local))
319                    && !already_seen.contains(&dest_val.z3_term)
320                {
321                    already_seen.insert(dest_val.z3_term.clone());
322                    self.inject_div_axioms_for_term(
323                        vm_state,
324                        solver,
325                        &dest_val.z3_term,
326                        depth - 1,
327                    );
328                }
329            }
330        }
331
332        // ── Process direct matches ──
333        for (lhs_pk, rhs_pk) in &op_sources {
334            let (Some(lhs_pk), Some(rhs_pk)) = (lhs_pk, rhs_pk) else {
335                continue;
336            };
337            let (Some(lhs_local), Some(rhs_local)) = (lhs_pk.local(), rhs_pk.local()) else {
338                continue;
339            };
340            let (Some(lhs_val), Some(rhs_val)) = (
341                vm_state.local_value(lhs_local),
342                vm_state.local_value(rhs_local),
343            ) else {
344                continue;
345            };
346
347            // Check if lhs is itself a Div / Rem result
348            if let Some((div_lhs_pk, div_rhs_pk)) = lhs_pk
349                .local()
350                .and_then(|l| vm_state.local_value(l))
351                .and_then(|v| v.source.operands().map(|(l, r, _)| (l.clone(), r.clone())))
352            {
353                let Some(div_lhs_local) = div_lhs_pk.and_then(|pk| pk.local()) else {
354                    continue;
355                };
356                let Some(div_rhs_local) = div_rhs_pk.and_then(|pk| pk.local()) else {
357                    continue;
358                };
359                let Some(div_lhs_val) = vm_state.local_value(div_lhs_local) else {
360                    continue;
361                };
362                let Some(div_rhs_val) = vm_state.local_value(div_rhs_local) else {
363                    continue;
364                };
365
366                let quot = div_lhs_val.z3_term.div(&div_rhs_val.z3_term);
367                let rem = div_lhs_val.z3_term.rem(&div_rhs_val.z3_term);
368                let mul_term = Int::mul(vm_state.z3_ctx, &[&quot, &div_rhs_val.z3_term]);
369                let sum_term = Int::add(vm_state.z3_ctx, &[&mul_term, &rem]);
370                solver.assert(&div_lhs_val.z3_term._eq(&sum_term));
371                let zero = Int::from_u64(vm_state.z3_ctx, 0);
372                solver.assert(&rem.ge(&zero));
373                solver.assert(&mul_term.le(&div_lhs_val.z3_term));
374            }
375
376            // Recurse into operands
377            self.inject_div_axioms_for_term(vm_state, solver, &lhs_val.z3_term, depth - 1);
378            self.inject_div_axioms_for_term(vm_state, solver, &rhs_val.z3_term, depth - 1);
379        }
380    }
381
382    pub(super) fn try_get_iter_len_term<'z3, 'tcx>(
383        &self,
384        vm_state: &VmState<'z3, 'tcx>,
385        expr: &ContractExpr<'tcx>,
386    ) -> Option<Int<'z3>> {
387        let ContractExpr::Len(_) = expr else {
388            return None;
389        };
390        for (_, val) in vm_state.all_local_values() {
391            let is_iter = match val.ty.kind() {
392                TyKind::Ref(_, pointee, _) => match pointee.kind() {
393                    TyKind::Adt(adt_def, _) => api_classify::is_std_iter_or_itermut(adt_def.did()),
394                    _ => false,
395                },
396                _ => false,
397            };
398            if !is_iter {
399                continue;
400            }
401            let alloc_id = val.provenance_alloc_id()?;
402            for (l, lv) in vm_state.all_local_values() {
403                if lv.provenance_alloc_id() != Some(alloc_id) {
404                    continue;
405                }
406                if let (Some(ptr), Some(end)) =
407                    (vm_state.field_value(l, &[0]), vm_state.field_value(l, &[1]))
408                {
409                    if let Some(len) = vm_state.iter_len_from_ptrs(ptr, end) {
410                        return Some(len);
411                    }
412                }
413            }
414        }
415        None
416    }
417
418    pub(super) fn try_iter_len_from_fields<'z3, 'tcx>(
419        &self,
420        vm_state: &VmState<'z3, 'tcx>,
421        checkpoint: &Checkpoint<'tcx>,
422        expr: &ContractExpr<'tcx>,
423    ) -> Option<Int<'z3>> {
424        use rustc_middle::mir::Place;
425        let ContractExpr::Place(cp) = expr else {
426            return None;
427        };
428        let op: &Operand<'tcx> = match cp.base {
429            PlaceBase::Arg(n) => checkpoint.args.get(n)?,
430            PlaceBase::Local(n) => {
431                let callee = checkpoint.callee?;
432                let idx = crate::helpers::mir_utils::callee_param_index_for_local(
433                    vm_state.tcx,
434                    callee,
435                    n,
436                )?;
437                checkpoint.args.get(idx)?
438            }
439            _ => return None,
440        };
441        let place: &Place<'tcx> = match op {
442            Operand::Copy(p) | Operand::Move(p) => p,
443            _ => return None,
444        };
445        let local = place.local;
446        let local_val = vm_state.local_value(local)?;
447        let is_iter = match local_val.ty.kind() {
448            TyKind::Ref(_, pointee, _) => match pointee.kind() {
449                TyKind::Adt(adt_def, _) => api_classify::is_std_iter_or_itermut(adt_def.did()),
450                _ => false,
451            },
452            _ => false,
453        };
454        if !is_iter {
455            return None;
456        }
457        // Direct lookup first, then scan by alloc_id for temp copies.
458        if let (Some(ptr), Some(end)) = (
459            vm_state.field_value(local, &[0]),
460            vm_state.field_value(local, &[1]),
461        ) {
462            if let Some(len) = vm_state.iter_len_from_ptrs(ptr, end) {
463                return Some(len);
464            }
465        }
466        // Fallback: scan all locals for one with same struct alloc.
467        let target_alloc = local_val.provenance_alloc_id()?;
468        for (scan_local, scan_val) in vm_state.all_local_values() {
469            if scan_val.provenance_alloc_id() != Some(target_alloc) {
470                continue;
471            }
472            if let (Some(ptr), Some(end)) = (
473                vm_state.field_value(scan_local, &[0]),
474                vm_state.field_value(scan_local, &[1]),
475            ) {
476                if let Some(len) = vm_state.iter_len_from_ptrs(ptr, end) {
477                    return Some(len);
478                }
479            }
480        }
481        None
482    }
483}