Skip to main content

rapx/verify/contract/
resolve.rs

1//! Expression / argument resolution: `syn::Expr` → semantic values.
2//!
3//! The numeric-expression layer is parsed by the pest grammar (`pest_conv.rs`);
4//! everything that still needs rustc's type context or `syn` structure lives
5//! here: places (via `place.rs`), const generics, builtin integer bounds, the
6//! `x.len()` sugar, tag argument types/targets, and `ValidNum` predicates.
7
8use quote::ToTokens;
9use rustc_hir::def::DefKind;
10use rustc_hir::def_id::DefId;
11use rustc_middle::ty::{GenericParamDefKind, Ty, TyCtxt};
12use syn::{Expr, Lit};
13
14use crate::helpers::fn_info::parse_expr_into_number;
15use crate::helpers::name::{access_ident_recursive, match_ty_with_ident};
16
17use super::place;
18use super::types::{ContractExpr, ContractPlace, NumericPredicate, PlaceBase, PropertyArg, RelOp};
19
20pub(crate) fn parse_contract_expr<'tcx>(
21    tcx: TyCtxt<'tcx>,
22    def_id: DefId,
23    expr: &Expr,
24    sp: &str,
25) -> ContractExpr<'tcx> {
26    // `x.len()` sugar -> len(x).  A plain `.len` *field* access is deliberately
27    // NOT turned into `len(x)`: a struct such as `LinkedList` carries `len` as an
28    // ordinary `usize` field, and routing it through `len()` would wrongly
29    // reconstruct the length from the backing allocation's size (which, for an
30    // external allocation, is not the field's value).  `.len` therefore stays a
31    // field projection (`self.len` → field 2) resolved by `parse_contract_place`.
32    if let Expr::MethodCall(expr_method) = expr
33        && expr_method.method == "len"
34        && expr_method.args.is_empty()
35    {
36        return ContractExpr::Len(Box::new(parse_contract_expr(
37            tcx,
38            def_id,
39            &expr_method.receiver,
40            sp,
41        )));
42    }
43
44    // A place (fields, projections), a const generic, or a builtin constant.
45    if let Some(place) = place::parse_contract_place(tcx, def_id, expr) {
46        return ContractExpr::Place(place);
47    }
48    if let Some(e) = parse_const_param(tcx, def_id, expr) {
49        return e;
50    }
51    if let Some(value) = parse_builtin_const(tcx, expr) {
52        return ContractExpr::Const(value);
53    }
54    if let Some(value) = parse_expr_into_number(expr) {
55        return ContractExpr::new_value(value);
56    }
57    // A `const` item (e.g. `CAPACITY` in `ValidNum(len <= CAPACITY)`).
58    if let Expr::Path(expr_path) = expr
59        && let Some(ident) = expr_path.path.get_ident()
60        && let Some(value) =
61            crate::helpers::mir_utils::resolve_const_item_value(tcx, &ident.to_string())
62    {
63        return ContractExpr::Const(value);
64    }
65    rap_debug!(
66        "Numeric expression in {:?} could not be resolved: {:?}",
67        sp,
68        expr
69    );
70    ContractExpr::Unknown
71}
72
73pub(crate) fn resolve_type_name<'tcx>(
74    tcx: TyCtxt<'tcx>,
75    def_id: DefId,
76    name: &str,
77) -> Option<Ty<'tcx>> {
78    if name == "Self" {
79        // `Self` refers to the type owning `def_id`: for an ADT (struct/enum/
80        // union) it is the type itself (`tcx.type_of`), for a function it is
81        // the receiver (the first input of the signature).
82        return match tcx.def_kind(def_id) {
83            DefKind::Struct | DefKind::Enum | DefKind::Union => {
84                Some(tcx.type_of(def_id).skip_binder())
85            }
86            _ => {
87                let sig = tcx.fn_sig(def_id).skip_binder();
88                sig.inputs().skip_binder().first().copied()
89            }
90        };
91    }
92    match_ty_with_ident(tcx, def_id, name.to_string())
93}
94
95pub(crate) fn int_type_min_max<'tcx>(tcx: TyCtxt<'tcx>, ty: Ty<'tcx>) -> Option<(u128, u128)> {
96    use rustc_middle::ty::IntTy;
97    use rustc_middle::ty::UintTy;
98    let bits: u32 = match ty.kind() {
99        rustc_middle::ty::TyKind::Uint(ut) => match ut {
100            UintTy::U8 => 8,
101            UintTy::U16 => 16,
102            UintTy::U32 => 32,
103            UintTy::U64 => 64,
104            UintTy::U128 => 128,
105            UintTy::Usize => tcx.data_layout.pointer_size().bits() as u32,
106        },
107        rustc_middle::ty::TyKind::Int(it) => match it {
108            IntTy::I8 => 8,
109            IntTy::I16 => 16,
110            IntTy::I32 => 32,
111            IntTy::I64 => 64,
112            IntTy::I128 => 128,
113            IntTy::Isize => tcx.data_layout.pointer_size().bits() as u32,
114        },
115        _ => return None,
116    };
117    if bits == 0 {
118        return None;
119    }
120    match ty.kind() {
121        rustc_middle::ty::TyKind::Uint(_) => {
122            let max = if bits == 128 {
123                u128::MAX
124            } else {
125                (1u128 << bits) - 1
126            };
127            Some((0, max))
128        }
129        rustc_middle::ty::TyKind::Int(_) => {
130            let max = (1u128 << (bits - 1)) - 1;
131            let min = max + 1;
132            Some((min, max))
133        }
134        _ => None,
135    }
136}
137
138fn parse_builtin_const<'tcx>(tcx: TyCtxt<'tcx>, expr: &Expr) -> Option<u128> {
139    let Expr::Path(expr_path) = expr else {
140        return None;
141    };
142    let mut segments = expr_path.path.segments.iter();
143    let first = segments.next()?.ident.to_string();
144    let second = segments.next()?.ident.to_string();
145    if segments.next().is_some() || second != "MAX" {
146        return None;
147    }
148
149    let pointer_bits = tcx.data_layout.pointer_size().bits();
150    match first.as_str() {
151        "isize" => Some((1_u128 << (pointer_bits - 1)) - 1),
152        "usize" => Some((1_u128 << pointer_bits) - 1),
153        _ => None,
154    }
155}
156
157fn parse_const_param<'tcx>(
158    tcx: TyCtxt<'tcx>,
159    def_id: DefId,
160    expr: &Expr,
161) -> Option<ContractExpr<'tcx>> {
162    let Expr::Path(expr_path) = expr else {
163        return None;
164    };
165    let ident = expr_path.path.get_ident()?.to_string();
166    let mut generics = Some(tcx.generics_of(def_id));
167    while let Some(current) = generics {
168        if let Some(param) = current.own_params.iter().find(|param| {
169            matches!(param.kind, GenericParamDefKind::Const { .. }) && param.name.as_str() == ident
170        }) {
171            return Some(ContractExpr::ConstParam {
172                index: param.index,
173                name: ident,
174            });
175        }
176        generics = current.parent.map(|parent| tcx.generics_of(parent));
177    }
178    None
179}
180
181pub(crate) fn parse_type<'tcx>(
182    tcx: TyCtxt<'tcx>,
183    def_id: DefId,
184    expr: &Expr,
185    sp: &str,
186) -> Option<Ty<'tcx>> {
187    if let Expr::Verbatim(ts) = expr {
188        let syn_ty = syn::parse2::<syn::Type>(ts.clone()).ok();
189        // A slice type `[T]` (e.g. `ValidTransmute([u8], str)`): resolve the
190        // element type and rebuild it as `TyKind::Slice`.
191        if let Some(syn::Type::Slice(slice)) = &syn_ty {
192            let Some(elem_name) = outermost_type_ident(&slice.elem) else {
193                rap_debug!("Incorrect expression for the type of {:?} Tag!", sp);
194                return None;
195            };
196            let Some(elem) = resolve_ty_ident(tcx, def_id, &elem_name) else {
197                rap_debug!("Cannot get type in {:?} Tag!", sp);
198                return None;
199            };
200            return Some(Ty::new_slice(tcx, elem));
201        }
202        // An array type `[T; N]` (e.g. `ValidTransmute([MaybeUninit<T>; N], [T; N])`):
203        // resolve the element type; the const length is kept symbolic (0), which
204        // `check_valid_transmute` handles by treating size 0 as "trust".
205        if let Some(syn::Type::Array(array)) = &syn_ty {
206            let Some(elem_name) = outermost_type_ident(&array.elem) else {
207                rap_debug!("Incorrect expression for the type of {:?} Tag!", sp);
208                return None;
209            };
210            let Some(elem) = resolve_ty_ident(tcx, def_id, &elem_name) else {
211                rap_debug!("Cannot get type in {:?} Tag!", sp);
212                return None;
213            };
214            return Some(Ty::new_array(tcx, elem, 0));
215        }
216        let name = syn_ty.and_then(|ty| outermost_type_ident(&ty));
217        let Some(name) = name else {
218            rap_debug!("Incorrect expression for the type of {:?} Tag!", sp);
219            return None;
220        };
221        let ty = resolve_ty_ident(tcx, def_id, &name);
222        if ty.is_none() {
223            rap_debug!("Cannot get type in {:?} Tag!", sp);
224        }
225        return ty;
226    }
227
228    // A multi-segment path type (`std::ascii::Char`) — resolve its last segment.
229    if let Expr::Path(expr_path) = expr
230        && expr_path.path.segments.len() > 1
231        && let Some(last) = expr_path.path.segments.last()
232    {
233        return resolve_ty_ident(tcx, def_id, &last.ident.to_string());
234    }
235
236    let ty_ident_full = access_ident_recursive(expr);
237    if ty_ident_full.is_none() {
238        rap_debug!("Incorrect expression for the type of {:?} Tag!", sp);
239        return None;
240    }
241    let ty_ident = ty_ident_full.unwrap().0;
242    let ty = resolve_ty_ident(tcx, def_id, &ty_ident);
243    if ty.is_none() {
244        rap_debug!("Cannot get type in {:?} Tag!", sp);
245    }
246    ty
247}
248
249/// Resolve a type identifier to a `Ty`, handling the `Self` keyword (which
250/// [`match_ty_with_ident`] does not understand) by delegating to
251/// [`resolve_type_name`].
252fn resolve_ty_ident<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, name: &str) -> Option<Ty<'tcx>> {
253    if name == "Self" {
254        resolve_type_name(tcx, def_id, name)
255    } else {
256        match_ty_with_ident(tcx, def_id, name.to_string())
257    }
258}
259
260/// Extract the outermost path segment name from a `syn::Type`, e.g. `Option`
261/// from `Option<NonZero<T>>` or `NonZero` from `NonZero<T>`.
262fn outermost_type_ident(ty: &syn::Type) -> Option<String> {
263    match ty {
264        syn::Type::Path(tp) if tp.qself.is_none() => {
265            tp.path.segments.last().map(|s| s.ident.to_string())
266        }
267        _ => None,
268    }
269}
270
271pub(crate) fn parse_target_arg<'tcx>(
272    tcx: TyCtxt<'tcx>,
273    def_id: DefId,
274    expr: &Expr,
275) -> PropertyArg<'tcx> {
276    // `return` parses as `syn::Expr::Return { expr: None }` (a bare `return`),
277    // which the place parser below does not recognise — handle it directly.
278    if matches!(expr, Expr::Return(_)) {
279        return PropertyArg::Expr(ContractExpr::Place(ContractPlace {
280            base: PlaceBase::Return,
281            projections: Vec::new(),
282        }));
283    }
284    // For simple identifiers that aren't local variables (e.g., lifetime param
285    // 'a parsed as ident `a`), store as Ident rather than Expr (which would
286    // become Unknown).
287    if let Expr::Path(expr_path) = expr {
288        if let Some(ident) = expr_path.path.get_ident() {
289            let s = ident.to_string();
290            if s != "return"
291                && !s.starts_with("Arg_")
292                && place::parse_expr_into_local_and_ty(tcx, def_id, expr).is_none()
293            {
294                return PropertyArg::Ident(s);
295            }
296        }
297    }
298    place::parse_contract_place(tcx, def_id, expr)
299        .map(|p| PropertyArg::Expr(ContractExpr::Place(p)))
300        .unwrap_or_else(|| PropertyArg::Expr(parse_contract_expr(tcx, def_id, expr, "target")))
301}
302
303pub(crate) fn parse_valid_num<'tcx>(
304    tcx: TyCtxt<'tcx>,
305    def_id: DefId,
306    exprs: &[Expr],
307) -> Vec<NumericPredicate<'tcx>> {
308    match exprs {
309        [] => Vec::new(),
310        [expr] => parse_numeric_predicate(tcx, def_id, expr)
311            .into_iter()
312            .collect(),
313        [value, range, ..] => {
314            if let Some(predicates) = parse_interval_predicates(tcx, def_id, value, range) {
315                predicates
316            } else {
317                parse_numeric_predicate(tcx, def_id, value)
318                    .into_iter()
319                    .collect()
320            }
321        }
322    }
323}
324
325fn parse_numeric_predicate<'tcx>(
326    tcx: TyCtxt<'tcx>,
327    def_id: DefId,
328    expr: &Expr,
329) -> Option<NumericPredicate<'tcx>> {
330    let text = expr.to_token_stream().to_string();
331    super::pest_conv::parse_predicate_pest(tcx, def_id, &text)
332}
333
334pub(crate) fn expr_to_pest<'tcx>(
335    tcx: TyCtxt<'tcx>,
336    def_id: DefId,
337    expr: &Expr,
338) -> ContractExpr<'tcx> {
339    let text = expr.to_token_stream().to_string();
340    super::pest_conv::parse_expr_pest(tcx, def_id, &text)
341}
342
343fn parse_interval_predicates<'tcx>(
344    tcx: TyCtxt<'tcx>,
345    def_id: DefId,
346    value: &Expr,
347    range: &Expr,
348) -> Option<Vec<NumericPredicate<'tcx>>> {
349    match range {
350        Expr::Array(array) if array.elems.len() == 2 => {
351            let mut elems = array.elems.iter();
352            let lower = elems.next().unwrap();
353            let upper = elems.next().unwrap();
354            Some(build_interval_predicates(
355                tcx, def_id, value, lower, true, upper, true,
356            ))
357        }
358        Expr::Lit(expr_lit) => match &expr_lit.lit {
359            Lit::Str(range_lit) => parse_string_interval(tcx, def_id, value, &range_lit.value()),
360            Lit::Int(int_lit) => {
361                // A bare integer `ValidNum(v, n)` is shorthand for the singleton
362                // interval `[n, n]`, i.e. `v == n`.
363                let n = int_lit.base10_parse::<u64>().ok()?;
364                let n_expr = syn::parse_str::<Expr>(&n.to_string()).ok()?;
365                Some(build_interval_predicates(
366                    tcx, def_id, value, &n_expr, true, &n_expr, true,
367                ))
368            }
369            _ => None,
370        },
371        _ => None,
372    }
373}
374
375fn parse_string_interval<'tcx>(
376    tcx: TyCtxt<'tcx>,
377    def_id: DefId,
378    value: &Expr,
379    raw_range: &str,
380) -> Option<Vec<NumericPredicate<'tcx>>> {
381    let trimmed = raw_range.trim();
382    if trimmed.len() < 3 {
383        return None;
384    }
385
386    let lower_inclusive = trimmed.starts_with('[');
387    let upper_inclusive = trimmed.ends_with(']');
388    if !(lower_inclusive || trimmed.starts_with('('))
389        || !(upper_inclusive || trimmed.ends_with(')'))
390    {
391        return None;
392    }
393
394    let body = &trimmed[1..trimmed.len() - 1];
395    let (lower_raw, upper_raw) = body.split_once(',')?;
396    let lower_raw = lower_raw.trim();
397    let upper_raw = upper_raw.trim();
398
399    // An unbounded side is written as an empty bound, e.g. `[1,)` (no upper
400    // bound) or `(,5]` (no lower bound). Reject an entirely empty interval.
401    if lower_raw.is_empty() && upper_raw.is_empty() {
402        return None;
403    }
404
405    let value_expr = expr_to_pest(tcx, def_id, value);
406    let mut predicates = Vec::with_capacity(2);
407
408    if !lower_raw.is_empty() {
409        let lower = syn::parse_str::<Expr>(lower_raw).ok()?;
410        predicates.push(NumericPredicate::new(
411            expr_to_pest(tcx, def_id, &lower),
412            if lower_inclusive {
413                RelOp::Le
414            } else {
415                RelOp::Lt
416            },
417            value_expr.clone(),
418        ));
419    }
420
421    if !upper_raw.is_empty() {
422        let upper = syn::parse_str::<Expr>(upper_raw).ok()?;
423        predicates.push(NumericPredicate::new(
424            value_expr,
425            if upper_inclusive {
426                RelOp::Le
427            } else {
428                RelOp::Lt
429            },
430            expr_to_pest(tcx, def_id, &upper),
431        ));
432    }
433
434    Some(predicates)
435}
436
437fn build_interval_predicates<'tcx>(
438    tcx: TyCtxt<'tcx>,
439    def_id: DefId,
440    value: &Expr,
441    lower: &Expr,
442    lower_inclusive: bool,
443    upper: &Expr,
444    upper_inclusive: bool,
445) -> Vec<NumericPredicate<'tcx>> {
446    let value_expr = expr_to_pest(tcx, def_id, value);
447    let lower_expr = expr_to_pest(tcx, def_id, lower);
448    let upper_expr = expr_to_pest(tcx, def_id, upper);
449    vec![
450        NumericPredicate::new(
451            lower_expr,
452            if lower_inclusive {
453                RelOp::Le
454            } else {
455                RelOp::Lt
456            },
457            value_expr.clone(),
458        ),
459        NumericPredicate::new(
460            value_expr,
461            if upper_inclusive {
462                RelOp::Le
463            } else {
464                RelOp::Lt
465            },
466            upper_expr,
467        ),
468    ]
469}
470
471/// Extract the inner type from `[T]` (the `SplitTransmute([T], [U])` notation),
472/// then resolve it via `parse_type`.  The `[T]` argument arrives as either
473/// `Expr::Array` (JSON path, parsed via `syn::parse_str::<Expr>`) or
474/// `Expr::Verbatim` (source annotation path, where `parse_property_arg`'s
475/// type-first parse turns `Type::Slice` into `Verbatim`); both are unwrapped to
476/// the element type `T`.
477pub(crate) fn unwrap_array_expr<'tcx>(
478    tcx: TyCtxt<'tcx>,
479    def_id: DefId,
480    expr: &Expr,
481) -> Option<Ty<'tcx>> {
482    if let Expr::Array(arr) = expr
483        && arr.elems.len() == 1
484    {
485        return parse_type(tcx, def_id, &arr.elems[0], "SplitTransmute");
486    }
487    if let Expr::Verbatim(ts) = expr
488        && let Ok(syn::Type::Slice(slice)) = syn::parse2::<syn::Type>(ts.clone())
489        && let Some(name) = outermost_type_ident(slice.elem.as_ref())
490    {
491        return resolve_ty_ident(tcx, def_id, &name);
492    }
493    parse_type(tcx, def_id, expr, "SplitTransmute")
494}