Skip to main content

charon_lib/transform/simplify_output/
inline_selected_functions.rs

1use std::{collections::HashMap, mem};
2
3use crate::transform::CowBox;
4use crate::transform::{TransformCtx, ctx::UllbcPass};
5use crate::ullbc_ast::*;
6
7pub struct Transform {
8    to_inline: HashMap<FunDeclId, FunDecl>,
9}
10
11impl Transform {
12    pub fn new(ctx: &mut TransformCtx) -> CowBox<dyn UllbcPass> {
13        let panic_name = Name::from_path(names::EXPLICIT_PANIC_NAME);
14        let panic_explicit = ctx
15            .translated
16            .fun_decls
17            .iter_indexed()
18            .find(|(_, decl)| decl.item_meta.name == panic_name)
19            .map(|(id, _)| id);
20
21        // Collect and remove the functions that we want to inline.
22        let to_inline = ctx
23            .translated
24            .fun_decls
25            .extract(|_, decl| {
26                decl.body.as_unstructured().is_some_and(|body| {
27                    // `panic!` generates a function item named `panic_cold_explicit` that calls
28                    // `panic_explicit`. We inline that function.
29                    let block = &body.body[START_BLOCK_ID];
30                    let is_local_panic_fn = if decl.item_meta.name.short_str()
31                        == Some("panic_cold_explicit")
32                        && matches!(body.body.len(), 2 | 3)
33                        && block.statements.is_empty()
34                        && let TerminatorKind::Call { call, .. } = &block.terminator.kind
35                        && let FnOperand::Regular(fn_ptr) = &call.func
36                        && let FnPtrKind::Fun(id) = fn_ptr.kind.as_ref()
37                        && Some(*id) == panic_explicit
38                    {
39                        body.body.iter_enumerated().all(|(id, block)| {
40                            id == START_BLOCK_ID
41                                || block.statements.is_empty()
42                                    && matches!(
43                                        block.terminator.kind,
44                                        TerminatorKind::UnwindResume
45                                            | TerminatorKind::UnwindTerminate
46                                            | TerminatorKind::UndefinedBehavior
47                                    )
48                        })
49                    } else {
50                        false
51                    };
52                    // The `anon_consts_to_call` pass already transformed references to anon consts
53                    // into calls to their initializers so we only have to inline these.
54                    let is_anon_const_initializer = if let FunSource::GlobalInitializer(global) =
55                        &decl.src
56                        && let Some(gdecl) = ctx.translated.global_decls.get(global.id)
57                    {
58                        matches!(gdecl.global_kind, GlobalKind::AnonConst)
59                    } else {
60                        false
61                    };
62                    let is_vec_construction_fn = decl.item_meta.diagnostic_item.as_deref()
63                        == Some(names::BOX_ASSUME_INIT_INTO_VEC_UNSAFE);
64                    is_local_panic_fn
65                        || (is_anon_const_initializer && ctx.options.inline_anon_consts)
66                        || (is_vec_construction_fn && ctx.options.treat_box_as_builtin)
67                })
68            })
69            .collect();
70
71        CowBox::Owned(Box::new(Transform { to_inline }))
72    }
73}
74impl UllbcPass for Transform {
75    fn should_run(&self, _options: &crate::options::TranslateOptions) -> bool {
76        !self.to_inline.is_empty()
77    }
78    fn apply_preceding_passes(&mut self, ctx: &mut TransformCtx, passes: &[CowBox<dyn UllbcPass>]) {
79        for decl in self.to_inline.values_mut() {
80            for pass in passes {
81                pass.transform_item(ctx, ItemRefMut::Fun(decl));
82            }
83        }
84    }
85    fn transform_body(&self, _ctx: &mut TransformCtx, outer_body: &mut ullbc_ast::ExprBody) {
86        for block_id in outer_body.body.indices() {
87            let Some(block) = outer_body.body.get_mut(block_id) else {
88                continue;
89            };
90            let TerminatorKind::Call {
91                call: Call {
92                    func, args, dest, ..
93                },
94                target,
95                on_unwind,
96            } = &mut block.terminator.kind
97            else {
98                continue;
99            };
100            let target = *target;
101            let on_unwind = *on_unwind;
102            let is_cleanup = block.is_cleanup;
103            let dest_place = dest.clone();
104            let args = args.clone();
105            let FnOperand::Regular(fn_ptr) = &func else {
106                continue;
107            };
108            let FnPtrKind::Fun(fun_id) = fn_ptr.kind.as_ref() else {
109                continue;
110            };
111            let Some(initializer) = self.to_inline.get(fun_id) else {
112                continue;
113            };
114            let span = initializer.item_meta.span;
115            let Some(inner_body) = initializer.body.as_unstructured() else {
116                continue;
117            };
118
119            // We inline the required body by shifting its local ids and block ids
120            // and adding its blocks to the outer body. The inner body's return
121            // local becomes a normal local that we can read from. We redirect some
122            // gotos so that the inner body is executed before the current block.
123            let mut inner_body = {
124                let mut inner_body = inner_body.clone();
125                let inner_bound = inner_body.bound_body_regions;
126
127                // Shift all the body regions in the inner body BEFORE substitution,
128                // so that we only shift the inner body's own regions.
129                inner_body.dyn_visit_mut(|r: &mut Region| {
130                    if let Region::Body(v) = r {
131                        *v += outer_body.bound_body_regions;
132                    }
133                });
134                outer_body.bound_body_regions += inner_bound;
135
136                // Now substitute generics. This may inject outer-body Region::Body
137                // IDs, which is correct since they don't need shifting.
138                inner_body.substitute(&fn_ptr.generics)
139            };
140
141            let return_local = outer_body.locals.locals.next_idx();
142            inner_body.dyn_visit_in_body_mut(|l: &mut LocalId| {
143                *l += return_local;
144            });
145            outer_body
146                .locals
147                .locals
148                .extend(mem::take(&mut inner_body.locals.locals));
149
150            // The inner body assumes the arg places are live; allocate them, and initialize the
151            // args.
152            inner_body.body[0].statements.splice(
153                0..0,
154                args.into_iter()
155                    .enumerate()
156                    .flat_map(|(i, arg)| {
157                        let arg_local = return_local + i + 1;
158                        let arg_place = outer_body.locals.place_for_var(arg_local);
159                        [
160                            StatementKind::StorageLive(arg_local),
161                            StatementKind::Assign(arg_place, Rvalue::Use(arg, WithRetag::Yes)),
162                        ]
163                    })
164                    .map(|kind| Statement::new(span, kind)),
165            );
166
167            let mut final_block = BlockData::new_goto(span, target, is_cleanup);
168
169            // The inner body will write to `return_place`, but the outer body expects the value at
170            // `dest_place`.
171            let return_place = outer_body.locals.place_for_var(return_local);
172            final_block.statements.push(Statement::new(
173                span,
174                StatementKind::Assign(
175                    dest_place,
176                    Rvalue::Use(Operand::Move(return_place), WithRetag::Yes),
177                ),
178            ));
179            let final_block = outer_body.body.push(final_block);
180
181            // Shift all block ids in the inner body and point return/unwind to where they should.
182            let start_block = outer_body.body.next_idx();
183            inner_body.visit_block_ids_mut(|b: &mut BlockId| {
184                *b += start_block;
185            });
186            inner_body
187                .body
188                .dyn_visit_in_body_mut(|t: &mut Terminator| match t.kind {
189                    TerminatorKind::Return => {
190                        t.kind = TerminatorKind::Goto {
191                            target: final_block,
192                        };
193                    }
194                    TerminatorKind::UnwindResume => {
195                        t.kind = TerminatorKind::Goto { target: on_unwind };
196                    }
197                    _ => (),
198                });
199            if is_cleanup {
200                for block in &mut inner_body.body {
201                    block.is_cleanup = true;
202                }
203            }
204            // At the end of the current block, start evaluating the inner body.
205            outer_body.body[block_id].terminator.kind = TerminatorKind::Goto {
206                target: start_block,
207            };
208            // Add the blocks for the inner body.
209            outer_body.body.extend(inner_body.body);
210        }
211    }
212}