Skip to main content

rapx/verify/contract/
pest_conv.rs

1//! Semantic converter: pest `Pairs<Rule>` → `ContractExpr` / `NumericPredicate`
2//! / `CompoundBody`.
3//!
4//! This is the phase-2 counterpart to `pest_grammar.rs`: it turns the parse
5//! tree produced by the pest grammar into the contract AST (`types.rs` /
6//! `compound.rs`).
7//!
8//! Places (fields, projections) are bridged through the `place.rs` / `resolve.rs`
9//! helpers via a `syn` round-trip, since resolving a field name to a `Ty` still
10//! needs the rustc type context.  The arithmetic / call / if / constant layers
11//! are converted directly from the pest tree.
12
13use pest::Parser;
14use pest::iterators::Pair;
15use rustc_hir::def_id::DefId;
16use rustc_middle::ty::TyCtxt;
17
18use super::compound::{CompoundArg, CompoundBody};
19use super::pest_grammar::{ContractParser, Rule};
20use super::place::resolve_place_from_ident;
21use super::types::{
22    ContractExpr, ContractPlace, NumericBinOp, NumericPredicate, NumericUnaryOp, PlaceBase, RelOp,
23};
24use crate::helpers::name::match_ty_with_ident;
25
26fn only_child(pair: Pair<Rule>) -> Pair<Rule> {
27    pair.into_inner()
28        .next()
29        .expect("expected a single child pair")
30}
31
32fn relop_from_str(s: &str) -> Option<RelOp> {
33    match s {
34        "==" => Some(RelOp::Eq),
35        "!=" => Some(RelOp::Ne),
36        "<" => Some(RelOp::Lt),
37        "<=" => Some(RelOp::Le),
38        ">" => Some(RelOp::Gt),
39        ">=" => Some(RelOp::Ge),
40        _ => None,
41    }
42}
43
44/// Parse a numeric expression (no comparison) into a `ContractExpr`.
45pub(crate) fn parse_expr_pest<'tcx>(
46    tcx: TyCtxt<'tcx>,
47    def_id: DefId,
48    text: &str,
49) -> ContractExpr<'tcx> {
50    let Ok(mut pairs) = ContractParser::parse(Rule::expr, text) else {
51        rap_debug!("contract expression not supported by grammar: {text}");
52        return ContractExpr::Unknown;
53    };
54    conv_expr(tcx, def_id, pairs.next().expect("expr pair"))
55}
56
57/// Parse a predicate (comparison / `!x.is_empty()` / bare expr) into a
58/// `NumericPredicate`.
59pub(crate) fn parse_predicate_pest<'tcx>(
60    tcx: TyCtxt<'tcx>,
61    def_id: DefId,
62    text: &str,
63) -> Option<NumericPredicate<'tcx>> {
64    let Ok(mut pairs) = ContractParser::parse(Rule::expr, text) else {
65        rap_debug!("contract predicate not supported by grammar: {text}");
66        return None;
67    };
68    conv_predicate(tcx, def_id, pairs.next().expect("expr pair"))
69}
70
71fn conv_expr<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
72    match pair.as_rule() {
73        Rule::expr => conv_expr(tcx, def_id, only_child(pair)),
74        Rule::if_expr => conv_if(tcx, def_id, pair),
75        Rule::cmp => {
76            // Expression layer carries no comparison operator.
77            let mut inner = pair.into_inner();
78            let lhs = conv_bit_or(tcx, def_id, inner.next().expect("cmp lhs"));
79            if inner.next().is_some() {
80                ContractExpr::Unknown
81            } else {
82                lhs
83            }
84        }
85        Rule::bit_or | Rule::bit_xor | Rule::bit_and | Rule::additive | Rule::multiplicative => {
86            conv_bit_or(tcx, def_id, pair)
87        }
88        Rule::unary => conv_unary(tcx, def_id, pair),
89        Rule::primary => conv_primary(tcx, def_id, pair),
90        Rule::call => conv_call(tcx, def_id, pair),
91        Rule::place => conv_place_bridge(tcx, def_id, pair),
92        Rule::const_path => conv_const_path(tcx, def_id, pair),
93        Rule::int => ContractExpr::Const(pair.as_str().parse::<u128>().unwrap_or(0)),
94        _ => ContractExpr::Unknown,
95    }
96}
97
98fn conv_predicate<'tcx>(
99    tcx: TyCtxt<'tcx>,
100    def_id: DefId,
101    pair: Pair<Rule>,
102) -> Option<NumericPredicate<'tcx>> {
103    match pair.as_rule() {
104        Rule::expr | Rule::cond => conv_predicate(tcx, def_id, only_child(pair)),
105        Rule::cmp => {
106            let mut inner = pair.into_inner();
107            let lhs = conv_bit_or(tcx, def_id, inner.next()?);
108            match inner.next() {
109                Some(relop_pair) => {
110                    let op = relop_from_str(relop_pair.as_str())?;
111                    let rhs = conv_bit_or(tcx, def_id, inner.next()?);
112                    Some(NumericPredicate::new(lhs, op, rhs))
113                }
114                // Bare expression → `expr != 0`.
115                None => Some(NumericPredicate::new(
116                    lhs,
117                    RelOp::Ne,
118                    ContractExpr::Const(0),
119                )),
120            }
121        }
122        Rule::not_is_empty => {
123            let mut inner = pair.into_inner();
124            let base_text = inner.next()?.as_str().to_string();
125            let place = conv_base(tcx, def_id, &base_text);
126            Some(NumericPredicate::new(
127                ContractExpr::Len(Box::new(place)),
128                RelOp::Ne,
129                ContractExpr::Const(0),
130            ))
131        }
132        _ => None,
133    }
134}
135
136fn conv_if<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
137    let mut inner = pair.into_inner();
138    let cond_pair = inner.next().expect("if cond");
139    let then_pair = inner.next().expect("if then");
140    let else_pair = inner.next().expect("if else");
141    let Some(cond) = conv_predicate(tcx, def_id, cond_pair) else {
142        return ContractExpr::Unknown;
143    };
144    let then_expr = conv_expr(tcx, def_id, then_pair);
145    let else_expr = conv_expr(tcx, def_id, else_pair);
146    ContractExpr::If {
147        cond: Box::new(cond),
148        then_expr: Box::new(then_expr),
149        else_expr: Box::new(else_expr),
150    }
151}
152
153fn op_from_str(op: &str) -> Option<NumericBinOp> {
154    match op {
155        "+" => Some(NumericBinOp::Add),
156        "-" => Some(NumericBinOp::Sub),
157        "*" => Some(NumericBinOp::Mul),
158        "/" => Some(NumericBinOp::Div),
159        "%" => Some(NumericBinOp::Rem),
160        "&" => Some(NumericBinOp::BitAnd),
161        "|" => Some(NumericBinOp::BitOr),
162        "^" => Some(NumericBinOp::BitXor),
163        _ => None,
164    }
165}
166
167fn conv_left_assoc<'tcx>(
168    tcx: TyCtxt<'tcx>,
169    def_id: DefId,
170    pair: Pair<Rule>,
171    operand: impl Fn(TyCtxt<'tcx>, DefId, Pair<Rule>) -> ContractExpr<'tcx>,
172) -> ContractExpr<'tcx> {
173    let mut inner = pair.into_inner();
174    let mut acc = operand(tcx, def_id, inner.next().expect("first operand"));
175    while let Some(op_pair) = inner.next() {
176        let Some(op) = op_from_str(op_pair.as_str()) else {
177            return ContractExpr::Unknown;
178        };
179        let rhs = operand(tcx, def_id, inner.next().expect("rhs operand"));
180        acc = ContractExpr::Binary {
181            op,
182            lhs: Box::new(acc),
183            rhs: Box::new(rhs),
184        };
185    }
186    acc
187}
188
189fn conv_bit_or<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
190    match pair.as_rule() {
191        Rule::bit_or => conv_left_assoc(tcx, def_id, pair, conv_bit_xor),
192        _ => conv_bit_xor(tcx, def_id, pair),
193    }
194}
195
196fn conv_bit_xor<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
197    match pair.as_rule() {
198        Rule::bit_xor => conv_left_assoc(tcx, def_id, pair, conv_bit_and),
199        _ => conv_bit_and(tcx, def_id, pair),
200    }
201}
202
203fn conv_bit_and<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
204    match pair.as_rule() {
205        Rule::bit_and => conv_left_assoc(tcx, def_id, pair, conv_additive),
206        _ => conv_additive(tcx, def_id, pair),
207    }
208}
209
210fn conv_additive<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
211    match pair.as_rule() {
212        Rule::additive => conv_left_assoc(tcx, def_id, pair, conv_multiplicative),
213        _ => conv_multiplicative(tcx, def_id, pair),
214    }
215}
216
217fn conv_multiplicative<'tcx>(
218    tcx: TyCtxt<'tcx>,
219    def_id: DefId,
220    pair: Pair<Rule>,
221) -> ContractExpr<'tcx> {
222    match pair.as_rule() {
223        Rule::multiplicative => conv_left_assoc(tcx, def_id, pair, conv_unary),
224        _ => conv_unary(tcx, def_id, pair),
225    }
226}
227
228fn conv_unary<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
229    let mut inner = pair.into_inner();
230    let first = inner.next().expect("unary operand");
231    match first.as_rule() {
232        Rule::unop => {
233            let op = match first.as_str() {
234                "!" => NumericUnaryOp::Not,
235                "-" => NumericUnaryOp::Neg,
236                _ => return ContractExpr::Unknown,
237            };
238            let operand = conv_unary(tcx, def_id, inner.next().expect("unary inner"));
239            ContractExpr::Unary {
240                op,
241                expr: Box::new(operand),
242            }
243        }
244        _ => conv_primary(tcx, def_id, first),
245    }
246}
247
248fn conv_primary<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
249    let inner = only_child(pair);
250    match inner.as_rule() {
251        Rule::int => inner
252            .as_str()
253            .parse::<u128>()
254            .map(ContractExpr::Const)
255            .unwrap_or(ContractExpr::Unknown),
256        Rule::call => conv_call(tcx, def_id, inner),
257        Rule::size_of_call => conv_size_of_call(tcx, def_id, inner),
258        Rule::const_path => conv_const_path(tcx, def_id, inner),
259        Rule::place => conv_place_bridge(tcx, def_id, inner),
260        Rule::expr => conv_expr(tcx, def_id, inner),
261        _ => ContractExpr::Unknown,
262    }
263}
264
265/// Convert `size_of::<T>()` / `align_of::<T>()` (optionally `std::mem::` /
266/// `core::mem::` prefixed) into `SizeOf` / `AlignOf`.
267fn conv_size_of_call<'tcx>(
268    tcx: TyCtxt<'tcx>,
269    def_id: DefId,
270    pair: Pair<Rule>,
271) -> ContractExpr<'tcx> {
272    let text = pair.as_str();
273    let (kind, rest) = if text.contains("align_of") {
274        ("align_of", text.split("align_of").nth(1).unwrap_or(""))
275    } else {
276        ("size_of", text.split("size_of").nth(1).unwrap_or(""))
277    };
278    // rest looks like " :: < usize > ()" — extract the ident between `<` and `>`.
279    let ty_name = rest
280        .find('<')
281        .and_then(|lt| {
282            rest[lt + 1..]
283                .find('>')
284                .map(|gt| rest[lt + 1..lt + 1 + gt].trim().to_string())
285        })
286        .unwrap_or_default();
287    let Some(ty) = match_ty_with_ident(tcx, def_id, ty_name) else {
288        return ContractExpr::Unknown;
289    };
290    match kind {
291        "size_of" => ContractExpr::SizeOf(ty),
292        _ => ContractExpr::AlignOf(ty),
293    }
294}
295
296fn conv_call<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
297    let mut inner = pair.into_inner();
298    let builtin = inner.next().expect("builtin").as_str().to_string();
299    // `call = builtin "(" arg_list? ")"`; arg_list's children are the `arg`s.
300    let args: Vec<Pair<Rule>> = match inner.next() {
301        Some(arg_list) => arg_list.into_inner().collect(),
302        None => Vec::new(),
303    };
304    match builtin.as_str() {
305        "size_of" | "align_of" => {
306            let ty_name = args
307                .first()
308                .map(|a| a.as_str().trim().to_string())
309                .unwrap_or_default();
310            let Some(ty) = match_ty_with_ident(tcx, def_id, ty_name) else {
311                return ContractExpr::Unknown;
312            };
313            match builtin.as_str() {
314                "size_of" => ContractExpr::SizeOf(ty),
315                _ => ContractExpr::AlignOf(ty),
316            }
317        }
318        "len" => {
319            let Some(arg) = args.first() else {
320                return ContractExpr::Unknown;
321            };
322            ContractExpr::Len(Box::new(conv_arg_expr(tcx, def_id, arg.clone())))
323        }
324        "min" | "max" => {
325            if args.len() != 2 {
326                return ContractExpr::Unknown;
327            }
328            let a = conv_arg_expr(tcx, def_id, args[0].clone());
329            let b = conv_arg_expr(tcx, def_id, args[1].clone());
330            let op = if builtin == "min" {
331                NumericBinOp::Min
332            } else {
333                NumericBinOp::Max
334            };
335            ContractExpr::Binary {
336                op,
337                lhs: Box::new(a),
338                rhs: Box::new(b),
339            }
340        }
341        _ => ContractExpr::Unknown,
342    }
343}
344
345fn conv_arg_expr<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, arg: Pair<Rule>) -> ContractExpr<'tcx> {
346    let inner = only_child(arg);
347    match inner.as_rule() {
348        Rule::expr => conv_expr(tcx, def_id, inner),
349        _ => ContractExpr::Unknown,
350    }
351}
352
353fn conv_const_path<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, pair: Pair<Rule>) -> ContractExpr<'tcx> {
354    let text = pair.as_str();
355    let Some((ty_name, which)) = text.rsplit_once("::") else {
356        return ContractExpr::Unknown;
357    };
358    let ty_name = ty_name.trim();
359    let which = which.trim();
360    let Some(ty) = super::resolve::resolve_type_name(tcx, def_id, ty_name) else {
361        return ContractExpr::Unknown;
362    };
363    // `T::BITS` is the bit width, i.e. `size_of::<T>() * 8`.
364    if which == "BITS" {
365        return ContractExpr::Binary {
366            op: NumericBinOp::Mul,
367            lhs: Box::new(ContractExpr::SizeOf(ty)),
368            rhs: Box::new(ContractExpr::Const(8)),
369        };
370    }
371    let Some((min, max)) = super::resolve::int_type_min_max(tcx, ty) else {
372        return ContractExpr::Unknown;
373    };
374    match which {
375        "MAX" => ContractExpr::Const(max),
376        "MIN" => {
377            // Signed integers: `int_type_min_max` returns the negated magnitude
378            // (`-(MIN) == 2^(bits-1)`) as a `u128` because it cannot represent
379            // the negative `MIN`. Emit an explicit negation so `i32::MIN`
380            // resolves to `-2147483648` rather than `+2147483648`.
381            if let rustc_middle::ty::TyKind::Int(_) = ty.kind() {
382                ContractExpr::Unary {
383                    op: NumericUnaryOp::Neg,
384                    expr: Box::new(ContractExpr::Const(min)),
385                }
386            } else {
387                ContractExpr::Const(min)
388            }
389        }
390        _ => ContractExpr::Unknown,
391    }
392}
393
394/// Bridge a place through the existing syn-based parser (handles field
395/// projections, `unwrap_some`, `iter`, and `x.len` sugar uniformly).
396fn conv_place_bridge<'tcx>(
397    tcx: TyCtxt<'tcx>,
398    def_id: DefId,
399    pair: Pair<Rule>,
400) -> ContractExpr<'tcx> {
401    let text = pair.as_str();
402    let Ok(expr) = syn::parse_str::<syn::Expr>(text) else {
403        return ContractExpr::Unknown;
404    };
405    super::resolve::parse_contract_expr(tcx, def_id, &expr, "pest")
406}
407
408/// Convert a `not_is_empty` base (`self` / `return` / `Arg_N` / ident) into a
409/// place expression.
410fn conv_base<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId, base_text: &str) -> ContractExpr<'tcx> {
411    match base_text {
412        "return" => ContractExpr::Place(ContractPlace {
413            base: PlaceBase::Return,
414            projections: Vec::new(),
415        }),
416        s if s.starts_with("Arg_") => {
417            let idx = s[4..].parse::<usize>().unwrap_or(0);
418            ContractExpr::Place(ContractPlace::arg(idx))
419        }
420        _ => {
421            let Some((base, fields, _)) = resolve_place_from_ident(tcx, def_id, base_text, &[])
422            else {
423                // Fall back to a `const` item (e.g. `CAPACITY` in
424                // `ValidNum(len <= CAPACITY)`).
425                if let Some(value) =
426                    crate::helpers::mir_utils::resolve_const_item_value(tcx, base_text)
427                {
428                    return ContractExpr::Const(value);
429                }
430                return ContractExpr::Unknown;
431            };
432            ContractExpr::Place(ContractPlace::local(base, fields))
433        }
434    }
435}
436
437// ── Compound-body conversion (`pred!` / `def_body`) ─────────────────────────
438
439/// Parse a compound-property body into a DNF tree.  `||` binds looser than `&&`.
440pub(crate) fn parse_compound_body(body: &str, params: &[String]) -> Option<CompoundBody> {
441    let mut pairs = ContractParser::parse(Rule::def_body, body).ok()?;
442    let def_body = pairs.next()?;
443    let or_expr = def_body.into_inner().next()?;
444    Some(conv_compound_or(or_expr, params))
445}
446
447fn conv_compound_or(pair: Pair<Rule>, params: &[String]) -> CompoundBody {
448    let parts: Vec<CompoundBody> = pair
449        .into_inner()
450        .map(|p| conv_compound_and(p, params))
451        .collect();
452    singleton_or_wrap(parts, CompoundBody::Or)
453}
454
455fn conv_compound_and(pair: Pair<Rule>, params: &[String]) -> CompoundBody {
456    let parts: Vec<CompoundBody> = pair
457        .into_inner()
458        .map(|p| conv_compound_leaf(p, params))
459        .collect();
460    singleton_or_wrap(parts, CompoundBody::And)
461}
462
463/// Unwrap a single-element list, otherwise wrap it with `wrap`.
464fn singleton_or_wrap(
465    parts: Vec<CompoundBody>,
466    wrap: fn(Vec<CompoundBody>) -> CompoundBody,
467) -> CompoundBody {
468    if parts.len() == 1 {
469        parts.into_iter().next().unwrap()
470    } else {
471        wrap(parts)
472    }
473}
474
475fn conv_compound_leaf(pair: Pair<Rule>, params: &[String]) -> CompoundBody {
476    match pair.into_inner().next() {
477        Some(inner) => match inner.as_rule() {
478            Rule::tag_call => conv_compound_call(inner, params),
479            Rule::or_expr => conv_compound_or(inner, params),
480            _ => CompoundBody::Call {
481                tag: String::new(),
482                args: Vec::new(),
483            },
484        },
485        None => CompoundBody::Call {
486            tag: String::new(),
487            args: Vec::new(),
488        },
489    }
490}
491
492fn conv_compound_call(pair: Pair<Rule>, params: &[String]) -> CompoundBody {
493    let mut inner = pair.into_inner();
494    let Some(tag) = inner.next() else {
495        return CompoundBody::Call {
496            tag: String::new(),
497            args: Vec::new(),
498        };
499    };
500    let tag = tag.as_str().to_string();
501    let args = match inner.next() {
502        Some(arg_list) => arg_list
503            .into_inner()
504            .map(|arg| {
505                let text = arg.as_str().trim().to_string();
506                match params.iter().position(|n| n == &text) {
507                    Some(i) => CompoundArg::Param(i),
508                    None => CompoundArg::Lit(text),
509                }
510            })
511            .collect(),
512        None => Vec::new(),
513    };
514    CompoundBody::Call { tag, args }
515}