Skip to main content

rustc_ast_lowering/
contract.rs

1use std::sync::Arc;
2
3use rustc_attr_ir::lang_items::LangItem;
4use rustc_attr_ir::target::Target;
5use rustc_span::sym;
6use thin_vec::thin_vec;
7
8use crate::LoweringContext;
9
10impl<'hir> LoweringContext<'_, 'hir> {
11    /// Lowered contracts are guarded with the `contract_checks` compiler flag,
12    /// i.e. the flag turns into a boolean guard in the lowered HIR. The reason
13    /// for not eliminating the contract code entirely when the `contract_checks`
14    /// flag is disabled is so that contracts can be type checked, even when
15    /// they are disabled, which avoids them becoming stale (i.e. out of sync
16    /// with the codebase) over time.
17    ///
18    /// The optimiser should be able to eliminate all contract code guarded
19    /// by `if false`, leaving the original body intact when runtime contract
20    /// checks are disabled.
21    pub(super) fn lower_contract(
22        &mut self,
23        body: impl FnOnce(&mut Self) -> rustc_hir::Expr<'hir>,
24        contract: &rustc_ast::FnContract,
25    ) -> rustc_hir::Expr<'hir> {
26        // The order in which things are lowered is important! I.e to
27        // refer to variables in contract_decls from postcond/precond,
28        // we must lower it first!
29        let contract_decls = self.lower_decls(contract);
30
31        match (&contract.requires, &contract.ensures) {
32            (Some(req), Some(ens)) => {
33                // Lower the fn contract, which turns:
34                //
35                // { body }
36                //
37                // into:
38                //
39                // let __postcond = if contract_checks {
40                //     CONTRACT_DECLARATIONS;
41                //     contract_check_requires(PRECOND);
42                //     Some(|ret_val| POSTCOND)
43                // } else {
44                //     None
45                // };
46                // {
47                //     let ret = { body };
48                //
49                //     if contract_checks {
50                //         contract_check_ensures(__postcond, ret)
51                //     } else {
52                //         ret
53                //     }
54                // }
55
56                let precond = self.lower_precond(req);
57                let postcond_checker = self.lower_postcond_checker(ens);
58
59                let contract_check = self.lower_contract_check_with_postcond(
60                    contract_decls,
61                    Some(precond),
62                    postcond_checker,
63                );
64
65                let wrapped_body =
66                    self.wrap_body_with_contract_check(body, contract_check, postcond_checker.span);
67                self.expr_block(wrapped_body)
68            }
69            (None, Some(ens)) => {
70                // Lower the fn contract, which turns:
71                //
72                // { body }
73                //
74                // into:
75                //
76                // let __postcond = if contract_checks {
77                //     Some(|ret_val| POSTCOND)
78                // } else {
79                //     None
80                // };
81                // {
82                //     let ret = { body };
83                //
84                //     if contract_checks {
85                //         CONTRACT_DECLARATIONS;
86                //         contract_check_ensures(__postcond, ret)
87                //     } else {
88                //         ret
89                //     }
90                // }
91                let postcond_checker = self.lower_postcond_checker(ens);
92                let contract_check =
93                    self.lower_contract_check_with_postcond(contract_decls, None, postcond_checker);
94
95                let wrapped_body =
96                    self.wrap_body_with_contract_check(body, contract_check, postcond_checker.span);
97                self.expr_block(wrapped_body)
98            }
99            (Some(req), None) => {
100                // Lower the fn contract, which turns:
101                //
102                // { body }
103                //
104                // into:
105                //
106                // {
107                //      if contracts_checks {
108                //          CONTRACT_DECLARATIONS;
109                //          contract_requires(PRECOND);
110                //      }
111                //      body
112                // }
113                let precond = self.lower_precond(req);
114                let precond_check = self.lower_contract_check_just_precond(contract_decls, precond);
115
116                let body = self.arena.alloc(body(self));
117
118                // Flatten the body into precond check, then body.
119                let wrapped_body = self.block_all(
120                    body.span,
121                    self.arena.alloc_from_iter([precond_check].into_iter()),
122                    Some(body),
123                );
124                self.expr_block(wrapped_body)
125            }
126            (None, None) => body(self),
127        }
128    }
129
130    fn lower_decls(&mut self, contract: &rustc_ast::FnContract) -> &'hir [rustc_hir::Stmt<'hir>] {
131        let (decls, decls_tail) = self.lower_stmts(&contract.declarations);
132
133        if let Some(e) = decls_tail {
134            // include the tail expression in the declaration statements
135            let tail = self.stmt_expr(e.span, *e);
136            self.arena.alloc_from_iter(decls.into_iter().map(|d| *d).chain([tail].into_iter()))
137        } else {
138            decls
139        }
140    }
141
142    /// Lower the precondition check intrinsic.
143    fn lower_precond(&mut self, req: &Box<rustc_ast::Expr>) -> rustc_hir::Stmt<'hir> {
144        let lowered_req = self.lower_expr_mut(&req);
145        let req_span = self.mark_span_with_reason(
146            rustc_span::DesugaringKind::Contract,
147            lowered_req.span,
148            Some(Arc::clone(&crate::ALLOW_CONTRACTS)),
149        );
150        let precond = self.expr_call_lang_item_fn_mut(
151            req_span,
152            LangItem::ContractCheckRequires,
153            &*self.arena.alloc_from_iter([lowered_req])arena_vec![self; lowered_req],
154        );
155        self.stmt_expr(req.span, precond)
156    }
157
158    fn lower_postcond_checker(
159        &mut self,
160        ens: &Box<rustc_ast::Expr>,
161    ) -> &'hir rustc_hir::Expr<'hir> {
162        let ens_span = self.lower_span(ens.span);
163        let ens_span = self.mark_span_with_reason(
164            rustc_span::DesugaringKind::Contract,
165            ens_span,
166            Some(Arc::clone(&crate::ALLOW_CONTRACTS)),
167        );
168        let lowered_ens = self.lower_expr_mut(&ens);
169        self.expr_call_lang_item_fn(
170            ens_span,
171            LangItem::ContractBuildCheckEnsures,
172            &*self.arena.alloc_from_iter([lowered_ens])arena_vec![self; lowered_ens],
173        )
174    }
175
176    fn lower_contract_check_just_precond(
177        &mut self,
178        contract_decls: &'hir [rustc_hir::Stmt<'hir>],
179        precond: rustc_hir::Stmt<'hir>,
180    ) -> rustc_hir::Stmt<'hir> {
181        let stmts = self
182            .arena
183            .alloc_from_iter(contract_decls.into_iter().map(|d| *d).chain([precond].into_iter()));
184
185        let then_block_stmts = self.block_all(precond.span, stmts, None);
186        let then_block = self.arena.alloc(self.expr_block(&then_block_stmts));
187
188        let precond_check = rustc_hir::ExprKind::If(
189            self.arena.alloc(self.expr_bool_literal(precond.span, self.tcx.sess.contract_checks())),
190            then_block,
191            None,
192        );
193
194        let precond_check = self.expr(precond.span, precond_check);
195        self.stmt_expr(precond.span, precond_check)
196    }
197
198    fn lower_contract_check_with_postcond(
199        &mut self,
200        contract_decls: &'hir [rustc_hir::Stmt<'hir>],
201        precond: Option<rustc_hir::Stmt<'hir>>,
202        postcond_checker: &'hir rustc_hir::Expr<'hir>,
203    ) -> &'hir rustc_hir::Expr<'hir> {
204        let stmts = self
205            .arena
206            .alloc_from_iter(contract_decls.into_iter().map(|d| *d).chain(precond.into_iter()));
207        let span = match precond {
208            Some(precond) => precond.span,
209            None => postcond_checker.span,
210        };
211
212        let postcond_checker = self.arena.alloc(self.expr_enum_variant_lang_item(
213            postcond_checker.span,
214            LangItem::OptionSome,
215            &*self.arena.alloc_from_iter([*postcond_checker])arena_vec![self; *postcond_checker],
216        ));
217        let then_block_stmts = self.block_all(span, stmts, Some(postcond_checker));
218        let then_block = self.arena.alloc(self.expr_block(&then_block_stmts));
219
220        let none_expr = self.arena.alloc(self.expr_enum_variant_lang_item(
221            postcond_checker.span,
222            LangItem::OptionNone,
223            Default::default(),
224        ));
225        let else_block = self.block_expr(none_expr);
226        let else_block = self.arena.alloc(self.expr_block(else_block));
227
228        let contract_check = rustc_hir::ExprKind::If(
229            self.arena.alloc(self.expr_bool_literal(span, self.tcx.sess.contract_checks())),
230            then_block,
231            Some(else_block),
232        );
233        self.arena.alloc(self.expr(span, contract_check))
234    }
235
236    fn wrap_body_with_contract_check(
237        &mut self,
238        body: impl FnOnce(&mut Self) -> rustc_hir::Expr<'hir>,
239        contract_check: &'hir rustc_hir::Expr<'hir>,
240        postcond_span: rustc_span::Span,
241    ) -> &'hir rustc_hir::Block<'hir> {
242        let check_ident: rustc_span::Ident =
243            rustc_span::Ident::new(sym::__ensures_checker, postcond_span);
244        let (check_hir_id, postcond_decl) = {
245            // Set up the postcondition `let` statement.
246            let (checker_pat, check_hir_id) = self.pat_ident_binding_mode_mut(
247                postcond_span,
248                check_ident,
249                rustc_hir::BindingMode::NONE,
250            );
251            (
252                check_hir_id,
253                self.stmt_let_pat(
254                    None,
255                    postcond_span,
256                    Some(contract_check),
257                    self.arena.alloc(checker_pat),
258                    rustc_hir::LocalSource::Contract,
259                ),
260            )
261        };
262
263        // Install contract_ensures so we will intercept `return` statements,
264        // then lower the body.
265        self.contract_ensures = Some((postcond_span, check_ident, check_hir_id));
266        let body = self.arena.alloc(body(self));
267
268        // Finally, inject an ensures check on the implicit return of the body.
269        let body = self.inject_ensures_check(body, postcond_span, check_ident, check_hir_id);
270
271        // Flatten the body into precond, then postcond, then wrapped body.
272        let wrapped_body = self.block_all(
273            body.span,
274            self.arena.alloc_from_iter([postcond_decl].into_iter()),
275            Some(body),
276        );
277        wrapped_body
278    }
279
280    /// Create an `ExprKind::Ret` that is optionally wrapped by a call to check
281    /// a contract ensures clause, if it exists.
282    pub(super) fn checked_return(
283        &mut self,
284        opt_expr: Option<&'hir rustc_hir::Expr<'hir>>,
285    ) -> rustc_hir::ExprKind<'hir> {
286        let checked_ret =
287            if let Some((check_span, check_ident, check_hir_id)) = self.contract_ensures {
288                let expr = opt_expr.unwrap_or_else(|| self.expr_unit(check_span));
289                Some(self.inject_ensures_check(expr, check_span, check_ident, check_hir_id))
290            } else {
291                opt_expr
292            };
293        rustc_hir::ExprKind::Ret(checked_ret)
294    }
295
296    /// Wraps an expression with a call to the ensures check before it gets returned.
297    pub(super) fn inject_ensures_check(
298        &mut self,
299        expr: &'hir rustc_hir::Expr<'hir>,
300        span: rustc_span::Span,
301        cond_ident: rustc_span::Ident,
302        cond_hir_id: rustc_hir::HirId,
303    ) -> &'hir rustc_hir::Expr<'hir> {
304        // {
305        //     let ret = { body };
306        //
307        //     if contract_checks {
308        //         contract_check_ensures(__postcond, ret)
309        //     } else {
310        //         ret
311        //     }
312        // }
313        let ret_ident: rustc_span::Ident = rustc_span::Ident::new(sym::__ret, span);
314
315        // Set up the return `let` statement.
316        let (ret_pat, ret_hir_id) =
317            self.pat_ident_binding_mode_mut(span, ret_ident, rustc_hir::BindingMode::NONE);
318
319        let ret_stmt = self.stmt_let_pat(
320            None,
321            span,
322            Some(expr),
323            self.arena.alloc(ret_pat),
324            rustc_hir::LocalSource::Contract,
325        );
326
327        let ret = self.expr_ident(span, ret_ident, ret_hir_id);
328
329        let cond_fn = self.expr_ident(span, cond_ident, cond_hir_id);
330        let contract_check = self.expr_call_lang_item_fn_mut(
331            span,
332            LangItem::ContractCheckEnsures,
333            self.arena.alloc_from_iter([*cond_fn, *ret])arena_vec![self; *cond_fn, *ret],
334        );
335        let contract_check = self.arena.alloc(contract_check);
336        let call_expr = self.block_expr_block(contract_check);
337
338        // same ident can't be used in 2 places, so we create a new one for the
339        // else branch
340        let ret = self.expr_ident(span, ret_ident, ret_hir_id);
341        let ret_block = self.block_expr_block(ret);
342
343        let contracts_enabled: rustc_hir::Expr<'_> =
344            self.expr_bool_literal(span, self.tcx.sess.contract_checks());
345        let contract_check = self.arena.alloc(self.expr(
346            span,
347            rustc_hir::ExprKind::If(
348                self.arena.alloc(contracts_enabled),
349                call_expr,
350                Some(ret_block),
351            ),
352        ));
353
354        let attrs: rustc_ast::AttrVec = {
    let len = [()].len();
    let mut vec = ::thin_vec::ThinVec::with_capacity(len);
    vec.push(self.unreachable_code_attr(span));
    vec
}thin_vec![self.unreachable_code_attr(span)];
355        self.lower_attrs(contract_check.hir_id, &attrs, span, Target::Expression);
356
357        let ret_block = self.block_all(span, self.arena.alloc_from_iter([ret_stmt])arena_vec![self; ret_stmt], Some(contract_check));
358        self.arena.alloc(self.expr_block(self.arena.alloc(ret_block)))
359    }
360}