Skip to main content

rapx/analysis/range/domain/
symbolic_expr.rs

1#![allow(unused_imports)]
2#![allow(unused_variables)]
3#![allow(dead_code)]
4#![allow(unused_assignments)]
5#![allow(unused_parens)]
6#![allow(non_snake_case)]
7use rust_intervals::NothingBetween;
8
9use crate::analysis::range::domain::ConstraintGraph;
10use crate::analysis::range::domain::domain::{ConstConvert, IntervalArithmetic, VarNode, VarNodes};
11use crate::analysis::range::{Range, RangeType};
12use crate::compat::FxHashMap;
13use crate::{rap_debug, rap_trace};
14use num_traits::{Bounded, CheckedAdd, CheckedSub, One, ToPrimitive, Zero, ops};
15use rustc_abi::Size;
16use rustc_hir::def_id::DefId;
17use rustc_middle::mir::coverage::Op;
18use rustc_middle::mir::{
19    BasicBlock, BinOp, BorrowKind, CastKind, Const, Local, LocalDecl, Operand, Place, Rvalue,
20    Statement, StatementKind, Terminator, UnOp,
21};
22use rustc_middle::ty::{ScalarInt, Ty};
23use rustc_span::sym::no_default_passes;
24use std::cell::RefCell;
25use std::cmp::PartialEq;
26use std::collections::{HashMap, HashSet};
27use std::fmt::Debug;
28use std::hash::Hash;
29use std::ops::{Add, Mul, Sub};
30use std::rc::Rc;
31use std::{fmt, mem};
32#[derive(Debug, Clone, Copy, PartialEq)]
33pub enum BoundMode {
34    Lower,
35    Upper,
36}
37
38impl BoundMode {
39    fn flip(self) -> Self {
40        match self {
41            BoundMode::Lower => BoundMode::Upper,
42            BoundMode::Upper => BoundMode::Lower,
43        }
44    }
45}
46#[derive(Debug, Clone, PartialEq, Eq)]
47pub enum SymbExpr<'tcx> {
48    Constant(Const<'tcx>),
49
50    Place(&'tcx Place<'tcx>),
51
52    Binary(BinOp, Box<SymbExpr<'tcx>>, Box<SymbExpr<'tcx>>),
53
54    Unary(UnOp, Box<SymbExpr<'tcx>>),
55
56    Cast(CastKind, Box<SymbExpr<'tcx>>, Ty<'tcx>),
57
58    Unknown,
59}
60impl<'tcx> fmt::Display for SymbExpr<'tcx> {
61    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
62        write!(f, "{:?}", self)
63    }
64}
65impl<'tcx> SymbExpr<'tcx> {
66    pub fn from_operand(op: &'tcx Operand<'tcx>, place_ctx: &Vec<&'tcx Place<'tcx>>) -> Self {
67        match op {
68            Operand::Copy(place) | Operand::Move(place) => {
69                let found_base = place_ctx
70                    .iter()
71                    .find(|&&p| p.local == place.local && p.projection.is_empty());
72
73                match found_base {
74                    Some(&base_place) => SymbExpr::Place(base_place),
75
76                    None => SymbExpr::Place(place),
77                }
78            }
79            Operand::Constant(c) => SymbExpr::Constant(c.const_),
80            #[cfg(rapx_ge_95)]
81            Operand::RuntimeChecks(_) => SymbExpr::Unknown,
82        }
83    }
84
85    pub fn from_rvalue(rvalue: &'tcx Rvalue<'tcx>, place_ctx: Vec<&'tcx Place<'tcx>>) -> Self {
86        match rvalue {
87            Rvalue::Use(op, ..) => Self::from_operand(op, &place_ctx),
88            Rvalue::BinaryOp(bin_op, pair) => {
89                let (lhs, rhs) = &**pair;
90                let left = Self::from_operand(lhs, &place_ctx);
91                let right = Self::from_operand(rhs, &place_ctx);
92
93                if matches!(left, SymbExpr::Unknown) || matches!(right, SymbExpr::Unknown) {
94                    return SymbExpr::Unknown;
95                }
96
97                SymbExpr::Binary(*bin_op, Box::new(left), Box::new(right))
98            }
99            Rvalue::UnaryOp(un_op, op) => {
100                let expr = Self::from_operand(op, &place_ctx);
101                if matches!(expr, SymbExpr::Unknown) {
102                    return SymbExpr::Unknown;
103                }
104                SymbExpr::Unary(*un_op, Box::new(expr))
105            }
106            Rvalue::Cast(kind, op, ty) => {
107                let expr = Self::from_operand(op, &place_ctx);
108                if matches!(expr, SymbExpr::Unknown) {
109                    return SymbExpr::Unknown;
110                }
111                SymbExpr::Cast(*kind, Box::new(expr), *ty)
112            }
113            Rvalue::Ref(..)
114            | Rvalue::ThreadLocalRef(..)
115            | Rvalue::Aggregate(..)
116            | Rvalue::Repeat(..)
117            | Rvalue::Discriminant(..)
118            | Rvalue::CopyForDeref(..) => SymbExpr::Unknown,
119            #[cfg(not(rapx_ge_99))]
120            Rvalue::ShallowInitBox(..) => SymbExpr::Unknown,
121            #[cfg(not(rapx_ge_95))]
122            Rvalue::NullaryOp(..) => SymbExpr::Unknown,
123            #[cfg(rapx_ge_99)]
124            Rvalue::Reborrow(..) => SymbExpr::Unknown,
125            Rvalue::RawPtr(raw_ptr_kind, place) => todo!(),
126            Rvalue::WrapUnsafeBinder(operand, ty) => todo!(),
127        }
128    }
129
130    pub fn resolve_upper_bound<T: IntervalArithmetic + ConstConvert + Debug + Clone + PartialEq>(
131        &mut self,
132        vars: &VarNodes<'tcx, T>,
133    ) {
134        self.resolve_recursive(vars, 0, BoundMode::Upper);
135    }
136    pub fn resolve_lower_bound<T: IntervalArithmetic + ConstConvert + Debug + Clone + PartialEq>(
137        &mut self,
138        vars: &VarNodes<'tcx, T>,
139    ) {
140        self.resolve_recursive(vars, 0, BoundMode::Lower);
141    }
142
143    fn resolve_recursive<T: IntervalArithmetic + ConstConvert + Debug + Clone + PartialEq>(
144        &mut self,
145        vars: &VarNodes<'tcx, T>,
146        depth: usize,
147        mode: BoundMode,
148    ) {
149        const MAX_DEPTH: usize = 10;
150        if depth > MAX_DEPTH {
151            *self = SymbExpr::Unknown;
152            return;
153        }
154
155        match self {
156            SymbExpr::Binary(op, lhs, rhs) => {
157                lhs.resolve_recursive(vars, depth + 1, mode);
158
159                match op {
160                    BinOp::Add | BinOp::AddUnchecked | BinOp::AddWithOverflow => {
161                        rhs.resolve_recursive(vars, depth + 1, mode);
162                    }
163                    BinOp::Sub | BinOp::SubUnchecked | BinOp::SubWithOverflow => {
164                        rhs.resolve_recursive(vars, depth + 1, mode.flip());
165                    }
166                    _ => rhs.resolve_recursive(vars, depth + 1, mode),
167                }
168            }
169            SymbExpr::Unary(op, inner) => match op {
170                UnOp::Neg => {
171                    inner.resolve_recursive(vars, depth + 1, mode.flip());
172                }
173                _ => inner.resolve_recursive(vars, depth + 1, mode),
174            },
175            SymbExpr::Cast(_, inner, _) => {
176                inner.resolve_recursive(vars, depth + 1, mode);
177            }
178            _ => {}
179        }
180
181        rap_trace!("symexpr {}", self);
182        if let SymbExpr::Place(place) = self {
183            if let Some(node) = vars.get(place) {
184                if let IntervalType::Basic(basic) = &node.interval {
185                    rap_trace!("node {:?}", *node);
186
187                    let target_expr = if basic.lower == basic.upper {
188                        &basic.upper
189                    } else {
190                        match mode {
191                            BoundMode::Upper => &basic.upper,
192                            BoundMode::Lower => &basic.lower,
193                        }
194                    };
195
196                    match target_expr {
197                        SymbExpr::Unknown => *self = SymbExpr::Unknown,
198                        SymbExpr::Constant(c) => *self = SymbExpr::Constant(c.clone()),
199                        expr => {
200                            if let SymbExpr::Place(target_place) = expr {
201                                if target_place == place {
202                                    return;
203                                }
204                            }
205
206                            *self = expr.clone();
207                            self.resolve_recursive(vars, depth + 1, mode);
208                        }
209                    }
210                }
211            }
212        }
213    }
214    pub fn simplify(&mut self) {
215        match self {
216            SymbExpr::Binary(_, lhs, rhs) => {
217                lhs.simplify();
218                rhs.simplify();
219            }
220            SymbExpr::Unary(_, inner) => {
221                inner.simplify();
222            }
223            SymbExpr::Cast(_, inner, _) => {
224                inner.simplify();
225            }
226            _ => {}
227        }
228
229        if let SymbExpr::Binary(op, lhs, rhs) = self {
230            match op {
231                BinOp::Sub | BinOp::SubUnchecked | BinOp::SubWithOverflow => {
232                    if let SymbExpr::Binary(inner_op, inner_lhs, inner_rhs) = lhs.as_ref() {
233                        match inner_op {
234                            BinOp::Add | BinOp::AddUnchecked | BinOp::AddWithOverflow => {
235                                if inner_lhs == rhs {
236                                    *self = *inner_rhs.clone();
237                                } else if inner_rhs == rhs {
238                                    *self = *inner_lhs.clone();
239                                }
240                            }
241                            _ => {}
242                        }
243                    }
244                }
245                BinOp::Add | BinOp::AddUnchecked | BinOp::AddWithOverflow => {
246                    if let SymbExpr::Binary(inner_op, inner_lhs, inner_rhs) = lhs.as_ref() {
247                        match inner_op {
248                            BinOp::Sub | BinOp::SubUnchecked | BinOp::SubWithOverflow => {
249                                if inner_rhs == rhs {
250                                    *self = *inner_lhs.clone();
251                                }
252                            }
253                            _ => {}
254                        }
255                    }
256                }
257                _ => {}
258            }
259        }
260    }
261}
262#[derive(Debug, Clone)]
263pub enum IntervalType<'tcx, T: IntervalArithmetic + ConstConvert + Debug> {
264    Basic(BasicInterval<'tcx, T>),
265    Symb(SymbInterval<'tcx, T>),
266}
267
268impl<'tcx, T: IntervalArithmetic + ConstConvert + Debug> fmt::Display for IntervalType<'tcx, T> {
269    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
270        match self {
271            IntervalType::Basic(b) => write!(
272                f,
273                "BasicInterval: {:?} {:?} {:?} ",
274                b.get_range(),
275                b.lower,
276                b.upper
277            ),
278            IntervalType::Symb(b) => write!(
279                f,
280                "SymbInterval: {:?} {:?} {:?} ",
281                b.get_range(),
282                b.lower,
283                b.upper
284            ),
285        }
286    }
287}
288pub trait IntervalTypeTrait<'tcx, T: IntervalArithmetic + ConstConvert + Debug> {
289    fn get_range(&self) -> &Range<T>;
290    fn set_range(&mut self, new_range: Range<T>);
291    fn get_lower_expr(&self) -> &SymbExpr<'tcx>;
292    fn get_upper_expr(&self) -> &SymbExpr<'tcx>;
293}
294impl<'tcx, T: IntervalArithmetic + ConstConvert + Debug> IntervalTypeTrait<'tcx, T>
295    for IntervalType<'tcx, T>
296{
297    fn get_range(&self) -> &Range<T> {
298        match self {
299            IntervalType::Basic(b) => b.get_range(),
300            IntervalType::Symb(s) => s.get_range(),
301        }
302    }
303
304    fn set_range(&mut self, new_range: Range<T>) {
305        match self {
306            IntervalType::Basic(b) => b.set_range(new_range),
307            IntervalType::Symb(s) => s.set_range(new_range),
308        }
309    }
310    fn get_lower_expr(&self) -> &SymbExpr<'tcx> {
311        match self {
312            IntervalType::Basic(b) => b.get_lower_expr(),
313            IntervalType::Symb(s) => s.get_lower_expr(),
314        }
315    }
316
317    fn get_upper_expr(&self) -> &SymbExpr<'tcx> {
318        match self {
319            IntervalType::Basic(b) => b.get_upper_expr(),
320            IntervalType::Symb(s) => s.get_upper_expr(),
321        }
322    }
323}
324#[derive(Debug, Clone)]
325
326pub struct BasicInterval<'tcx, T: IntervalArithmetic + ConstConvert + Debug> {
327    pub range: Range<T>,
328    pub lower: SymbExpr<'tcx>,
329    pub upper: SymbExpr<'tcx>,
330}
331
332impl<'tcx, T: IntervalArithmetic + ConstConvert + Debug> BasicInterval<'tcx, T> {
333    pub fn new(range: Range<T>) -> Self {
334        Self {
335            range,
336            lower: SymbExpr::Unknown,
337            upper: SymbExpr::Unknown,
338        }
339    }
340    pub fn new_symb(range: Range<T>, lower: SymbExpr<'tcx>, upper: SymbExpr<'tcx>) -> Self {
341        Self {
342            range,
343            lower,
344            upper,
345        }
346    }
347    pub fn default() -> Self {
348        Self {
349            range: Range::bottom(),
350            lower: SymbExpr::Unknown,
351            upper: SymbExpr::Unknown,
352        }
353    }
354}
355
356impl<'tcx, T: IntervalArithmetic + ConstConvert + Debug> IntervalTypeTrait<'tcx, T>
357    for BasicInterval<'tcx, T>
358{
359    fn get_range(&self) -> &Range<T> {
360        &self.range
361    }
362
363    fn set_range(&mut self, new_range: Range<T>) {
364        self.range = new_range;
365        if self.range.get_lower() > self.range.get_upper() {
366            self.range.set_empty();
367        }
368    }
369    fn get_lower_expr(&self) -> &SymbExpr<'tcx> {
370        &self.lower
371    }
372
373    fn get_upper_expr(&self) -> &SymbExpr<'tcx> {
374        &self.upper
375    }
376}
377
378#[derive(Debug, Clone)]
379
380pub struct SymbInterval<'tcx, T: IntervalArithmetic + ConstConvert + Debug> {
381    range: Range<T>,
382    symbound: &'tcx Place<'tcx>,
383    predicate: BinOp,
384    lower: SymbExpr<'tcx>,
385    upper: SymbExpr<'tcx>,
386}
387
388impl<'tcx, T: IntervalArithmetic + ConstConvert + Debug> SymbInterval<'tcx, T> {
389    pub fn new(range: Range<T>, symbound: &'tcx Place<'tcx>, predicate: BinOp) -> Self {
390        Self {
391            range,
392            symbound,
393            predicate,
394            lower: SymbExpr::Unknown,
395            upper: SymbExpr::Unknown,
396        }
397    }
398
399    pub fn get_operation(&self) -> BinOp {
400        self.predicate
401    }
402
403    pub fn get_bound(&self) -> &'tcx Place<'tcx> {
404        self.symbound
405    }
406
407    pub fn sym_fix_intersects(
408        &self,
409        bound: &VarNode<'tcx, T>,
410        sink: &VarNode<'tcx, T>,
411    ) -> Range<T> {
412        let l = bound.get_range().get_lower().clone();
413        let u = bound.get_range().get_upper().clone();
414
415        let lower = sink.get_range().get_lower().clone();
416        let upper = sink.get_range().get_upper().clone();
417
418        match self.predicate {
419            BinOp::Eq => Range::new(l, u, RangeType::Regular),
420
421            BinOp::Le => Range::new(lower, u, RangeType::Regular),
422
423            BinOp::Lt => {
424                if u != T::max_value() {
425                    let u_minus_1 = u.checked_sub(&T::one()).unwrap_or(u);
426                    Range::new(lower, u_minus_1, RangeType::Regular)
427                } else {
428                    Range::new(lower, u, RangeType::Regular)
429                }
430            }
431
432            BinOp::Ge => Range::new(l, upper, RangeType::Regular),
433
434            BinOp::Gt => {
435                if l != T::min_value() {
436                    let l_plus_1 = l.checked_add(&T::one()).unwrap_or(l);
437                    Range::new(l_plus_1, upper, RangeType::Regular)
438                } else {
439                    Range::new(l, upper, RangeType::Regular)
440                }
441            }
442
443            BinOp::Ne => Range::top(),
444
445            _ => Range::top(),
446        }
447    }
448}
449
450impl<'tcx, T: IntervalArithmetic + ConstConvert + Debug> IntervalTypeTrait<'tcx, T>
451    for SymbInterval<'tcx, T>
452{
453    fn get_range(&self) -> &Range<T> {
454        &self.range
455    }
456
457    fn set_range(&mut self, new_range: Range<T>) {
458        self.range = new_range;
459    }
460    fn get_lower_expr(&self) -> &SymbExpr<'tcx> {
461        &self.lower
462    }
463
464    fn get_upper_expr(&self) -> &SymbExpr<'tcx> {
465        &self.upper
466    }
467}